From b1f0c5da0265988ad620df63f4b194e3cd47e35d Mon Sep 17 00:00:00 2001 From: Nico Ritschel Date: Fri, 2 Oct 2026 06:09:19 -0700 Subject: [PATCH 1/3] Preserve Ossie logical expressions and aggregate grains in Python --- sidemantic/interchange/ossie/lowering.py | 202 +++++++++++-- sidemantic/interchange/ossie/portable.py | 274 ++++++++++++++++++ sidemantic/sql/generator.py | 257 +++++++++++++--- .../ossie/fixtures/portable_expressions.json | 136 +++++++++ .../ossie/test_expression_conformance.py | 126 ++++++++ .../ossie/test_portable_expression.py | 124 ++++++++ .../interchange/ossie/test_runtime_queries.py | 142 +++++++++ tests/interchange/ossie/test_synthesis.py | 46 +++ tests/metrics/test_symmetric_aggs.py | 6 +- 9 files changed, 1254 insertions(+), 59 deletions(-) create mode 100644 sidemantic/interchange/ossie/portable.py create mode 100644 tests/interchange/ossie/fixtures/portable_expressions.json create mode 100644 tests/interchange/ossie/test_expression_conformance.py create mode 100644 tests/interchange/ossie/test_portable_expression.py create mode 100644 tests/interchange/ossie/test_runtime_queries.py diff --git a/sidemantic/interchange/ossie/lowering.py b/sidemantic/interchange/ossie/lowering.py index 078a68140..e9f46933c 100644 --- a/sidemantic/interchange/ossie/lowering.py +++ b/sidemantic/interchange/ossie/lowering.py @@ -4,6 +4,7 @@ import hashlib import json +import re from collections import Counter from collections.abc import Mapping, Sequence from dataclasses import dataclass @@ -26,7 +27,7 @@ ) from sidemantic.interchange.ossie.documents import OssieLogicalDocument, OssieOntologyDocument, logical_model_entries from sidemantic.interchange.ossie.expression_validation import scalar_sql_expression_error -from sidemantic.interchange.ossie.identifier import identifier_within_limit, normalize_identifier +from sidemantic.interchange.ossie.identifier import identifier_within_limit, is_quoted_identifier, normalize_identifier from sidemantic.interchange.ossie.parser import OssieParseResult from sidemantic.interchange.ossie.profiles import OssieImportPolicy from sidemantic.interchange.ossie.runtime_extension import decode_runtime_extension, resolve_runtime_graph @@ -163,16 +164,160 @@ def _expression_for_target(expression: object, target_dialect: str) -> tuple[str target_label = _DIALECT_LABELS.get(normalized) if target_label and target_label in by_dialect: return by_dialect[target_label], target_label + if "OSSIE_SQL_2026" in by_dialect: + return by_dialect["OSSIE_SQL_2026"], "OSSIE_SQL_2026" if "ANSI_SQL" in by_dialect: return by_dialect["ANSI_SQL"], "ANSI_SQL" return None -def _sql_expression_error(expression: str, target_dialect: str) -> str | None: +def _sql_expression_error(expression: str, target_dialect: str, *, row_level: bool = False) -> str | None: dialect = _SQLGLOT_DIALECTS.get(_normalize_dialect(target_dialect)) if dialect is None: return f"target dialect {target_dialect!r} has no configured SQL parser" - return scalar_sql_expression_error(expression, sqlglot_dialect=dialect) + return scalar_sql_expression_error(expression, sqlglot_dialect=dialect, row_level=row_level) + + +def _lower_selected_sql( + selected: tuple[str, str], target_dialect: str, *, row_level: bool = False +) -> tuple[str, str | None]: + expression, dialect = selected + if dialect == "OSSIE_SQL_2026": + from sidemantic.interchange.ossie.portable import lower_ossie_sql + + try: + expression = lower_ossie_sql(expression, target_dialect) + except ValueError as exc: + return expression, str(exc) + return expression, _sql_expression_error(expression, target_dialect, row_level=row_level) + + +def _reference_key(identifier: exp.Identifier, expression_dialect: str) -> str: + """Compare a parsed identifier using the source contract, not the warehouse's case rules.""" + + quoted = identifier.args.get("quoted") and expression_dialect not in {"BIGQUERY", "DATABRICKS"} + return identifier.name if quoted else normalize_identifier(identifier.name) + + +def _runtime_names(names: Sequence[str]) -> dict[str, str]: + """Decode quoted declarations without merging distinct source identities. + + Runtime planners also target engines with case-insensitive aliases. Keep + ordinary names when possible; mangle unsafe or colliding quoted names with + a stable FNV-1a suffix (also used by the native importer). + """ + + candidates = {name: normalize_identifier(name) if is_quoted_identifier(name) else name for name in names} + counts = Counter(candidate.lower() for candidate in candidates.values()) + used = {candidate.lower() for candidate in candidates.values()} + result = {} + for name in sorted(names): + candidate = candidates[name] + safe = re.fullmatch(r"[A-Za-z_][A-Za-z0-9_]*", candidate) is not None + if safe and (not is_quoted_identifier(name) or counts[candidate.lower()] == 1): + result[name] = candidate + continue + fingerprint = 0xCBF29CE484222325 + for byte in name.encode(): + fingerprint = ((fingerprint ^ byte) * 0x100000001B3) & 0xFFFFFFFFFFFFFFFF + base = f"__ossie_{fingerprint:016x}" + candidate = base + suffix = 1 + while candidate.lower() in used: + candidate = f"{base}_{suffix}" + suffix += 1 + used.add(candidate.lower()) + result[name] = candidate + return result + + +def _dimension_declarations(model: Model) -> dict[str, str]: + return { + normalize_identifier((dimension.metadata or {}).get("ossie_source_name", dimension.name)): dimension.name + for dimension in model.dimensions + } + + +def _bind_metric_expression( + expression: str, + *, + models: Mapping[str, Model], + metric_expressions: Mapping[str, tuple[str, str, str]], + target_dialect: str, + expression_dialect: str, + active_metrics: tuple[str, ...], +) -> str: + """Resolve logical fields and expand metric references before entering the native graph. + + Bind declared fields before native complete-SQL planning, which otherwise + treats differently spelled references as physical columns. Physical column + references remain available when no logical field matches the source name. + """ + + dialect = _SQLGLOT_DIALECTS[_normalize_dialect(target_dialect)] + parsed = sqlglot.parse_one(expression, read=dialect) + changed = False + for column in list(parsed.find_all(exp.Column)): + if len(column.parts) > 2 or not isinstance(column.this, exp.Identifier): + raise ValueError(f"Unsupported logical field reference {column.sql(dialect=dialect)!r}") + field_key = _reference_key(column.this, expression_dialect) + if column.table: + table = column.args.get("table") + model = models.get(_reference_key(table, expression_dialect)) if isinstance(table, exp.Identifier) else None + if model is None: + raise ValueError(f"Unknown logical dataset in {column.sql(dialect=dialect)!r}") + candidates = [ + (model, dimension) + for dimension in model.dimensions + if _dimension_declarations(model).get(field_key) == dimension.name + ] + if not candidates: + if column.table != model.name: + column.set("table", exp.to_identifier(model.name, quoted=model.name.startswith('"'))) + changed = True + continue + elif field_key in metric_expressions: + name, metric_sql, metric_dialect = metric_expressions[field_key] + if field_key in active_metrics: + raise ValueError(f"Cyclic metric reference involving {name!r}") + bound = _bind_metric_expression( + metric_sql, + models=models, + metric_expressions=metric_expressions, + target_dialect=target_dialect, + expression_dialect=metric_dialect, + active_metrics=(*active_metrics, field_key), + ) + replacement = exp.Paren(this=sqlglot.parse_one(bound, read=dialect)) + if column is parsed: + parsed = replacement + else: + column.replace(replacement) + changed = True + continue + else: + candidates = [ + (model, dimension) + for model in models.values() + for dimension in model.dimensions + if _dimension_declarations(model).get(field_key) == dimension.name + ] + if not candidates and len(models) == 1: + model = next(iter(models.values())) + column.set("table", exp.to_identifier(model.name, quoted=model.name.startswith('"'))) + changed = True + continue + if len(candidates) != 1: + reason = "Ambiguous" if candidates else "Unknown" + raise ValueError(f"{reason} logical field {column.sql(dialect=dialect)!r}") + model, dimension = candidates[0] + # Runtime identifiers are decoded or safely aliased; source identity is + # retained separately and used for every lookup above. + if column.table != model.name or column.name != dimension.name: + column.set("table", exp.to_identifier(model.name, quoted=model.name.startswith('"'))) + column.set("this", exp.to_identifier(dimension.name, quoted=dimension.name.startswith('"'))) + changed = True + return parsed.sql(dialect=dialect) if changed else expression def _classify_source(source: str, source_dialect: str | None) -> tuple[str, str] | None: @@ -271,6 +416,7 @@ def _lower_scope( graph = runtime_override if runtime_override is not None else SemanticGraph() source_dialect = result.options.source_dialect if result.options else None dataset_values = _array(semantic_model.get("datasets")) if runtime_override is None else () + dataset_names = _runtime_names([name for _, _, name in _unique_named_items(dataset_values)]) lowered_models: dict[str, Model] = {} # Model/Metric construction has a legacy auto-registration hook. Lowering @@ -306,6 +452,7 @@ def _lower_scope( dimensions: list[Dimension] = [] fields = _array(dataset.get("fields")) + field_names = _runtime_names([name for _, _, name in _unique_named_items(fields)]) for field_index, field, field_name in _unique_named_items(fields): field_pointer = f"{pointer}/fields/{field_index}" selected = _expression_for_target(field.get("expression"), target_dialect) @@ -323,8 +470,7 @@ def _lower_scope( ) ) continue - expression, _ = selected - expression_error = _sql_expression_error(expression, target_dialect) + expression, expression_error = _lower_selected_sql(selected, target_dialect, row_level=True) if expression_error is not None: diagnostics.append( _diagnostic( @@ -348,13 +494,14 @@ def _lower_scope( ) try: runtime_dimension = Dimension( - name=field_name, + name=field_names[field_name], type=_runtime_dimension_type(logical_type, effective_is_time), logical_data_type=logical_type, declared_is_time=declared_is_time, sql=expression, description=field.get("description") if isinstance(field.get("description"), str) else None, label=field.get("label") if isinstance(field.get("label"), str) else None, + metadata={"ossie_source_name": field_name}, ) except (TypeError, ValueError) as exc: diagnostics.append( @@ -369,7 +516,10 @@ def _lower_scope( continue dimensions.append(runtime_dimension) - field_declarations = _canonical_name_lookup([dimension.name for dimension in dimensions]) + field_declarations = { + normalize_identifier(dimension.metadata["ossie_source_name"]): dimension.name + for dimension in dimensions + } primary_columns = _canonical_columns(_array(dataset.get("primary_key")), field_declarations) primary_key = _key_value(primary_columns) unique_keys: list[list[str]] = [] @@ -380,7 +530,7 @@ def _lower_scope( source_kind, source_text = classified_source try: model = Model( - name=dataset_name, + name=dataset_names[dataset_name], table=source_text if source_kind == "table" else None, sql=source_text if source_kind == "query" else None, description=dataset.get("description") if isinstance(dataset.get("description"), str) else None, @@ -388,7 +538,11 @@ def _lower_scope( unique_keys=unique_keys or None, dimensions=dimensions, default_time_dimension=None, - metadata={"ossie_source_kind": source_kind, "ossie_pointer": pointer}, + metadata={ + "ossie_source_kind": source_kind, + "ossie_pointer": pointer, + "ossie_source_name": dataset_name, + }, ) except (TypeError, ValueError) as exc: diagnostics.append( @@ -421,12 +575,8 @@ def _lower_scope( if to_name and identifier_within_limit(to_name) else None ) - from_declarations = ( - _canonical_name_lookup([dimension.name for dimension in from_model.dimensions]) if from_model else {} - ) - to_declarations = ( - _canonical_name_lookup([dimension.name for dimension in to_model.dimensions]) if to_model else {} - ) + from_declarations = _dimension_declarations(from_model) if from_model else {} + to_declarations = _dimension_declarations(to_model) if to_model else {} canonical_from_columns = _canonical_columns(from_columns, from_declarations) canonical_to_columns = _canonical_columns(to_columns, to_declarations) from_key = _key_value(canonical_from_columns) @@ -468,6 +618,13 @@ def _lower_scope( ) metrics = _array(semantic_model.get("metrics")) if runtime_override is None else () + metric_names = _runtime_names([name for _, _, name in _unique_named_items(metrics)]) + metric_expressions = { + normalize_identifier(name): (name, lowered[0], selected[1]) + for _, metric, name in _unique_named_items(metrics) + if (selected := _expression_for_target(metric.get("expression"), target_dialect)) is not None + if (lowered := _lower_selected_sql(selected, target_dialect))[1] is None + } for metric_index, metric, metric_name in _unique_named_items(metrics): pointer = f"{scope_pointer}/metrics/{metric_index}" selected = _expression_for_target(metric.get("expression"), target_dialect) @@ -485,8 +642,8 @@ def _lower_scope( ) ) continue - expression, expression_dialect = selected - expression_error = _sql_expression_error(expression, target_dialect) + _, expression_dialect = selected + expression, expression_error = _lower_selected_sql(selected, target_dialect) if expression_error is not None: diagnostics.append( _diagnostic( @@ -502,14 +659,23 @@ def _lower_scope( ) continue try: + expression = _bind_metric_expression( + expression, + models=lowered_models, + metric_expressions=metric_expressions, + target_dialect=target_dialect, + expression_dialect=expression_dialect, + active_metrics=(normalize_identifier(metric_name),), + ) metric_object = Metric( - name=metric_name, + name=metric_names[metric_name], sql=expression, # The selected target SQL must bypass Metric's implicit DuckDB extraction. sql_is_complete=True, metadata={ "ossie_expression_dialect": expression_dialect, "ossie_target_dialect": _normalize_dialect(target_dialect), + **({"ossie_source_name": metric_name} if metric_names[metric_name] != metric_name else {}), }, logical_data_type=(metric.get("datatype") if isinstance(metric.get("datatype"), str) else None), description=metric.get("description") if isinstance(metric.get("description"), str) else None, diff --git a/sidemantic/interchange/ossie/portable.py b/sidemantic/interchange/ossie/portable.py new file mode 100644 index 000000000..97966a07c --- /dev/null +++ b/sidemantic/interchange/ossie/portable.py @@ -0,0 +1,274 @@ +"""The OSSIE_SQL_2026 expression dialect, independently of warehouse SQL. + +The source grammar uses ANSI identifiers and the argument order in Ossie's +expression_language.md. Snowflake's parser recognizes that grammar, but its +session-dependent and vendor-specific function semantics are not the contract. +Unknown functions are preserved as extensions, as the proposal requires. +""" + +from __future__ import annotations + +import sqlglot +from sqlglot import exp +from sqlglot.dialects.snowflake import Snowflake +from sqlglot.errors import ErrorLevel, SqlglotError +from sqlglot.tokens import TokenType + +_TARGETS = {"postgresql": "postgres", "ansi_sql": "duckdb", "ansi": "duckdb"} +_EXTRACTIONS = { + exp.Year: "YEAR", + exp.Quarter: "QUARTER", + exp.Month: "MONTH", + exp.Day: "DAY", + exp.DayOfYear: "DAYOFYEAR", + exp.Hour: "HOUR", + exp.Minute: "MINUTE", + exp.Second: "SECOND", +} +_PARTS = { + "YEAR", + "QUARTER", + "MONTH", + "WEEK", + "DAY", + "DAYOFWEEK", + "DAYOFYEAR", + "HOUR", + "MINUTE", + "SECOND", + "MILLISECOND", +} + + +def _ansi_value_frames(sql: str) -> str: + """Spell out ANSI defaults before the source parser adds Snowflake frames. + + Token offsets let us insert only the missing frame, preserving every other + source function and literal. In particular, this avoids round-tripping the + entire expression through an unrelated SQL dialect. + """ + tokens = Snowflake().tokenize(sql) + additions = [] + for index, token in enumerate(tokens): + if token.text.upper() not in {"FIRST_VALUE", "LAST_VALUE", "NTH_VALUE"} or token.token_type in { + TokenType.STRING, + TokenType.IDENTIFIER, + }: + continue + cursor = index + 1 + if cursor >= len(tokens) or tokens[cursor].token_type != TokenType.L_PAREN: + continue + depth = 0 + while cursor < len(tokens): + depth += (tokens[cursor].token_type == TokenType.L_PAREN) - (tokens[cursor].token_type == TokenType.R_PAREN) + cursor += 1 + if depth == 0: + break + if ( + cursor + 1 >= len(tokens) + or tokens[cursor].token_type != TokenType.OVER + or tokens[cursor + 1].token_type != TokenType.L_PAREN + ): + continue + cursor += 2 + depth, has_frame, ordered = 1, False, False + while cursor < len(tokens) and depth: + current = tokens[cursor] + depth += (current.token_type == TokenType.L_PAREN) - (current.token_type == TokenType.R_PAREN) + if depth == 1: + has_frame |= current.text.upper() in {"ROWS", "RANGE", "GROUPS"} + ordered |= current.token_type == TokenType.ORDER_BY + if depth == 0 and not has_frame: + frame = ( + " RANGE BETWEEN UNBOUNDED PRECEDING AND CURRENT ROW" + if ordered + else " ROWS BETWEEN UNBOUNDED PRECEDING AND UNBOUNDED FOLLOWING" + ) + additions.append((current.start, frame)) + cursor += 1 + for position, frame in sorted(additions, reverse=True): + sql = sql[:position] + frame + sql[position:] + return sql + + +def parse_portable_expression(expression: str) -> exp.Expression: + """Parse one expression and reject query/statement constructs before lowering.""" + try: + expressions = sqlglot.parse(_ansi_value_frames(expression), read="snowflake") + except SqlglotError as exc: + raise ValueError(f"Invalid OSSIE_SQL_2026 expression: {exc}") from exc + if len(expressions) != 1 or expressions[0] is None: + raise ValueError("OSSIE_SQL_2026 requires exactly one expression") + root = expressions[0] + if isinstance(root, (exp.Alias, exp.Star, exp.Tuple)): + raise ValueError("OSSIE_SQL_2026 requires a scalar expression without an alias") + for node in root.walk(): + if isinstance( + node, + ( + exp.Query, + exp.DDL, + exp.Drop, + exp.Transaction, + exp.Commit, + exp.Rollback, + exp.DML, + exp.Command, + exp.Where, + exp.Group, + exp.Join, + exp.With, + exp.Array, + exp.Bracket, + ), + ): + raise ValueError(f"OSSIE_SQL_2026 expressions cannot contain {type(node).__name__}") + if isinstance(node, exp.Identifier) and len(node.name) > 128: + raise ValueError("OSSIE_SQL_2026 identifiers cannot exceed 128 characters") + if isinstance(node, exp.Column) and len(node.parts) > 2: + raise ValueError("OSSIE_SQL_2026 field references have at most two identifiers") + return root + + +def _call(name: str, *args: exp.Expression) -> exp.Expression: + return exp.Anonymous(this=name, expressions=[arg.copy() for arg in args]) + + +def _part(node: exp.Expression) -> str: + part = node.name.upper() + if part not in _PARTS: + raise ValueError(f"Unsupported OSSIE_SQL_2026 date part {part!r}") + return part + + +def lower_ossie_sql(expression: str, target_dialect: str | None) -> str: + """Translate portable SQL without silently approximating required operations. + + A warehouse limitation raises ValueError instead of publishing SQL that has + a different meaning. BigQuery exact ordered-set aggregates require a query + rewrite (its exact percentiles are analytic-only), outside scalar lowering. + """ + target = _TARGETS.get((target_dialect or "duckdb").lower(), (target_dialect or "duckdb").lower()) + if target not in {"duckdb", "postgres", "snowflake", "bigquery", "databricks"}: + raise ValueError(f"Unsupported OSSIE_SQL_2026 target {target_dialect!r}") + root = parse_portable_expression(expression) + + def rewrite(node: exp.Expression) -> exp.Expression: + if isinstance(node, exp.Anonymous) and node.name.upper() == "TO_TIMESTAMP": + if len(node.expressions) == 1: + return exp.Cast(this=node.expressions[0].copy(), to=exp.DataType.build("TIMESTAMPNTZ")) + if type(node) in _EXTRACTIONS: + part = _EXTRACTIONS[type(node)] + if part == "DAYOFYEAR" and target == "postgres": + part = "DOY" + return exp.Extract(this=exp.Var(this=part), expression=node.this.copy()) + if isinstance(node, exp.Extract): + part = _part(node.this) + if target == "postgres": + part = {"DAYOFYEAR": "DOY", "DAYOFWEEK": "DOW"}.get(part, part) + return exp.Extract(this=exp.Var(this=part), expression=node.expression.copy()) + if isinstance(node, exp.Log) and node.expression is not None: + # DuckDB has only the unary base-ten LOG; division also avoids + # PostgreSQL's numeric-only two-argument LOG overload. + return exp.Paren( + this=exp.Div( + this=exp.Ln(this=node.expression.copy()), expression=exp.Ln(this=node.this.copy()), safe=False + ) + ) + if isinstance(node, exp.Trunc): + decimals = node.args.get("decimals") or exp.Literal.number(0) + # Multiplying by POWER before FLOOR loses decimal digits (1.15 * + # 100 becomes 114.999...). ROUND retains DECIMAL arithmetic; undo + # its one-unit adjustment only when it rounded away from zero. + rounded = exp.Round(this=node.this.copy(), decimals=decimals.copy()) + places = decimals.sql() + if places.lstrip("-").isdigit() and abs(int(places)) <= 38: + count = int(places) + unit = "0." + "0" * (count - 1) + "1" if count > 0 else "1" + "0" * -count + quantum = exp.Literal.number(unit) + else: + quantum = exp.Pow(this=exp.Literal.number(10), expression=exp.Neg(this=decimals.copy())) + return exp.If( + this=exp.GT(this=exp.Abs(this=rounded.copy()), expression=exp.Abs(this=node.this.copy())), + true=exp.Sub( + this=rounded.copy(), expression=exp.Mul(this=exp.Sign(this=node.this.copy()), expression=quantum) + ), + false=rounded, + ) + if isinstance(node, exp.Contains): + return exp.GT( + this=exp.StrPosition(this=node.this.copy(), substr=node.expression.copy()), + expression=exp.Literal.number(0), + ) + if isinstance(node, (exp.StartsWith, exp.EndsWith)): + side = exp.Left if isinstance(node, exp.StartsWith) else exp.Right + return exp.EQ( + this=side(this=node.this.copy(), expression=exp.Length(this=node.expression.copy())), + expression=node.expression.copy(), + ) + if isinstance(node, exp.RegexpLike): + # The portable spelling denotes a match, not Snowflake's implicit + # whole-string anchoring. Make search behavior explicit there. + node.set("full_match", False) + if target == "snowflake": + return exp.GT(this=_call("REGEXP_INSTR", node.this, node.expression), expression=exp.Literal.number(0)) + if isinstance(node, (exp.TimestampTrunc, exp.DateTrunc)): + part = _part(node.args["unit"]) + if part in {"DAYOFWEEK", "DAYOFYEAR", "MILLISECOND"}: + raise ValueError(f"Unsupported OSSIE_SQL_2026 truncation part {part!r}") + if part == "WEEK" and target == "bigquery": + return _call("DATE_TRUNC", node.this, _call("WEEK", exp.Var(this="MONDAY"))) + if part == "WEEK" and target == "snowflake": + # WEEK_START is a session option; ISO weekdays always start Monday. + return sqlglot.parse_one( + f"DATEADD(day, 1 - DAYOFWEEKISO({node.this.sql(dialect=target)}), " + f"DATE_TRUNC('day', {node.this.sql(dialect=target)}))", + read=target, + ) + if isinstance(node, exp.DateDiff) and target == "postgres": + part = _part(node.args["unit"]) + start, end = node.expression.copy(), node.this.copy() + + def extract(value: exp.Expression, unit: str) -> exp.Expression: + return exp.Extract(this=exp.Var(this=unit), expression=value.copy()) + + if part in {"YEAR", "MONTH", "QUARTER"}: + years = exp.Sub(this=extract(end, "YEAR"), expression=extract(start, "YEAR")) + if part == "YEAR": + return years + unit = "MONTH" if part == "MONTH" else "QUARTER" + return exp.Paren( + this=exp.Add( + this=exp.Mul( + this=exp.Paren(this=years), expression=exp.Literal.number(12 if part == "MONTH" else 4) + ), + expression=exp.Sub(this=extract(end, unit), expression=extract(start, unit)), + ) + ) + if isinstance(node, exp.WithinGroup) and isinstance(node.this, (exp.PercentileCont, exp.PercentileDisc)): + if target == "bigquery": + raise ValueError("BigQuery exact ordered-set percentiles require query-level lowering") + if target == "databricks": + # SQLGlot maps this to PERCENTILE_APPROX, which is not equivalent. + node.set( + "this", + _call( + "PERCENTILE_CONT" if isinstance(node.this, exp.PercentileCont) else "PERCENTILE_DISC", + node.this.this, + ), + ) + if isinstance(node, exp.Median) and target == "bigquery": + raise ValueError("BigQuery exact MEDIAN requires query-level lowering") + return node + + # Bottom-up replacement keeps corrections inside compound expressions. + for node in reversed(list(root.walk())): + replacement = rewrite(node) + if node is root: + root = replacement + elif replacement is not node: + node.replace(replacement) + try: + return root.sql(dialect=target, unsupported_level=ErrorLevel.RAISE) + except SqlglotError as exc: + raise ValueError(f"OSSIE_SQL_2026 cannot be lowered to {target}: {exc}") from exc diff --git a/sidemantic/sql/generator.py b/sidemantic/sql/generator.py index 4354481fe..06255a77a 100644 --- a/sidemantic/sql/generator.py +++ b/sidemantic/sql/generator.py @@ -133,6 +133,145 @@ def __init__( self._generate_cache: dict[tuple[object, ...], str] = {} self._generate_cache_limit = 256 + def _lower_ossie_aggregates(self) -> "SQLGenerator | None": + """Give semantic-model aggregates an explicit grain before native planning. + + Ossie fields denote logical expressions, whereas native model measures + read physical columns. Resolve those expressions while creating private + leaves on a query-local graph; retain the public metric as a formula. + """ + replacements = {} + graph_metrics = self.graph.metrics.copy() + changed = False + aggregate_types = {exp.Sum: "sum", exp.Avg: "avg", exp.Count: "count", exp.Min: "min", exp.Max: "max"} + for name, metric in self.graph.metrics.items(): + if ( + not metric.sql_is_complete + or not metric.sql + or not (metric.metadata or {}).get("ossie_target_dialect") + or (metric.metadata or {}).get("ossie_joined_leaf") + ): + continue + parsed = _parse_fragment(metric.sql, self.dialect) + # A window frame cannot be split into independent aggregate grains. + if parsed.find(exp.Window): + continue + aggregates = list(parsed.find_all(exp.AggFunc)) + if any(aggregate.find_ancestor(exp.AggFunc) is not None for aggregate in aggregates): + continue + aggregates = [ + aggregate.parent if isinstance(aggregate.parent, (exp.Filter, exp.WithinGroup)) else aggregate + for aggregate in aggregates + ] + expression_models = {column.table for column in parsed.find_all(exp.Column)} + owners = [] + for aggregate in aggregates: + columns = list(aggregate.find_all(exp.Column)) + models = {column.table for column in columns} if columns else expression_models + if not models or not models <= self.graph.models.keys() or (not columns and len(models) != 1): + break + owners.append(next(iter(models)) if len(models) == 1 else None) + else: + # Every aggregate has a known row source. A formula with no + # aggregate leaves (e.g. revenue * 2) uses normal dependencies. + for aggregate, owner in zip(aggregates, owners, strict=True): + if owner is None: + # A row expression spanning datasets keeps its joined + # population, independently of sibling aggregate leaves. + index = 0 + occupied = {name.lower() for name in graph_metrics} + while f"__ossie_joined_{index}" in occupied: + index += 1 + leaf_name = f"__ossie_joined_{index}" + graph_metrics[leaf_name] = metric.model_copy( + update={ + "name": leaf_name, + "sql": aggregate.sql(dialect=self.dialect), + "type": None, + "agg": None, + "metadata": {**metric.metadata, "ossie_joined_leaf": True}, + } + ) + reference = exp.column(leaf_name) + if aggregate is parsed: + parsed = reference + else: + aggregate.replace(reference) + continue + model = replacements.get(owner, self.graph.models[owner]) + occupied = {field.name.lower() for field in [*model.dimensions, *model.metrics]} + occupied.update(key.lower() for key in model.primary_key_columns) + index = 0 + while f"__ossie_aggregate_{index}" in occupied or f"__ossie_aggregate_{index}_raw" in occupied: + index += 1 + leaf_name = f"__ossie_aggregate_{index}" + expression = aggregate.copy() + for column in list(expression.find_all(exp.Column)): + dimension = model.get_dimension(column.name) + if dimension is None: + # Source fields may be implicit in a portable graph + # projection; preserve their physical reference. + column.set("table", None) + continue + field_sql = dimension.sql_expr.replace("{model}.", "") + column.replace(exp.Paren(this=_parse_fragment(field_sql, self.dialect))) + aggregation = aggregate_types.get(type(expression)) + argument = expression.this + if isinstance(argument, exp.Distinct): + if aggregation == "count" and len(argument.expressions) == 1: + aggregation = "count_distinct" + argument = argument.expressions[0] + else: + aggregation = None + if aggregation and any( + value for key, value in expression.args.items() if key not in ("this", "big_int") + ): + aggregation = None + leaf = metric.model_copy( + update={ + "name": leaf_name, + "type": "simple" if aggregation else None, + "agg": aggregation, + "sql": argument.sql(dialect=self.dialect) + if aggregation + else expression.sql(dialect=self.dialect), + "sql_is_complete": aggregation is None, + } + ) + replacements[owner] = model.model_copy(update={"metrics": [*model.metrics, leaf]}) + reference = exp.column(leaf_name, table=owner) + if aggregate is parsed: + parsed = reference + else: + aggregate.replace(reference) + graph_metrics[name] = metric.model_copy( + update={"sql": parsed.sql(dialect=self.dialect), "type": "derived", "sql_is_complete": False} + ) + changed = True + if not changed: + return None + graph = copy(self.graph) + graph.models = {**self.graph.models, **replacements} + # Ossie permits unique keys without a primary key. Either declaration + # identifies entity rows for the native deduplication planner. + graph.models = { + name: model.model_copy(update={"primary_key": model.unique_keys[0]}) + if (model.metadata or {}).get("ossie_source_kind") and not model.primary_key_columns and model.unique_keys + else model + for name, model in graph.models.items() + } + graph.metrics = graph_metrics + graph._adjacency_dirty = True + graph._adjacency = {} + graph._role_models = {} + graph._role_owners = {} + graph._relationship_instances = {} + graph._relationship_path_cache = {} + generator = copy(self) + generator.graph = graph + generator._generate_cache = {} + return generator + def _lower_filtered_complete_aggregates(self) -> "SQLGenerator | None": """Reuse filtered aggregate planning without changing the caller's live graph.""" # TSQL count widths require separate qualification; retain its existing @@ -1256,7 +1395,7 @@ def generate( Returns: SQL query string """ - lowered = self._lower_filtered_complete_aggregates() + lowered = self._lower_ossie_aggregates() or self._lower_filtered_complete_aggregates() if lowered is not None: return lowered.generate( metrics=metrics, @@ -2364,7 +2503,14 @@ def _build_model_cte( def add_passthrough_column(column: str) -> None: if column not in columns_added: - select_cols.append(f"{self._quote_identifier(column)} AS {self._quote_alias(column)}") + dimension = model.get_dimension(column) if (model.metadata or {}).get("ossie_source_kind") else None + if dimension is not None: + self._ensure_sql_dimension(model_name, dimension) + expression = self._dimension_base_expr(dimension) + expression = expression.replace("{model}", "t") if model.sql else expression.replace("{model}.", "") + else: + expression = self._quote_identifier(column) + select_cols.append(f"{expression} AS {self._quote_alias(column)}") columns_added.add(column) # Cross joins do not need keys for the join predicate, but exact fan-out aggregation @@ -2841,39 +2987,21 @@ def _has_fanout_joins(self, base_model_name: str, other_models: list[str]) -> di Returns: Dict mapping model names to whether they need symmetric aggregates """ - needs_symmetric = {} - - # Check if there are any one-to-many relationships - one_to_many_count = 0 - many_to_one_models = [] - - for other_model in other_models: - try: - join_path = self.graph.find_relationship_path(base_model_name, other_model) - if not join_path: + models = [base_model_name, *other_models] + needs_symmetric = dict.fromkeys(models, False) + # A dimension can anchor the query on the many side. Assess fanout from + # each metric owner's perspective, not just the chosen FROM model. + for model_name in models: + for other_model in models: + if model_name == other_model: continue - # Check all hops: any one_to_many in the path creates fan-out - has_fanout = any(hop.relationship == "one_to_many" for hop in join_path) - if has_fanout: - one_to_many_count += 1 - elif join_path[0].relationship == "many_to_one": - many_to_one_models.append(other_model) - except (ValueError, KeyError): - pass - - # Base model needs symmetric aggregates if there are any one-to-many joins - needs_symmetric[base_model_name] = one_to_many_count > 0 - - # Models on the "many" side of a many-to-one relationship also need symmetric - # aggregation if they're being joined (because from their perspective, - # they're creating fan-out for the "one" side) - for other_model in other_models: - if other_model in many_to_one_models: - # Check if the "one" side (base) has metrics - if so, it needs symmetric agg - # But we're checking from the perspective of this model, so mark False - needs_symmetric[other_model] = False - else: - needs_symmetric[other_model] = False + try: + path = self.graph.find_relationship_path(model_name, other_model) + if any(hop.relationship == "one_to_many" for hop in path): + needs_symmetric[model_name] = True + break + except (ValueError, KeyError): + pass return needs_symmetric @@ -2928,6 +3056,35 @@ def _explicit_join_type_for_path(self, join_path) -> str | None: "full_outer": "full", }.get(how) + def _complete_graph_aggregates(self, metrics: list[str]) -> set[str]: + """Find joined-row Ossie aggregates that must keep their own row query.""" + complete = set() + visited = set() + + def visit(reference): + if reference in visited: + return + visited.add(reference) + try: + owner, metric = self.graph.resolve_metric_reference(reference) + except KeyError: + return + if ( + owner is None + and metric.sql_is_complete + and metric.sql + and (metric.metadata or {}).get("ossie_target_dialect") + and sql_has_aggregate(metric.sql, self.dialect) + ): + complete.add(reference) + return + for dependency in metric.get_dependencies(self.graph, owner): + visit(dependency) + + for reference in metrics: + visit(reference) + return complete + def _needs_preaggregation_for_fanout(self, metrics: list[str], dimensions: list[str]) -> bool: """Determine if pre-aggregation is needed to avoid fan-out. @@ -2954,6 +3111,9 @@ def _needs_preaggregation_for_fanout(self, metrics: list[str], dimensions: list[ # Calculated metrics can span multiple grains even when only one output # is selected (for example order revenue divided by customer count). metric_models = self._find_aggregate_metric_models(metrics) + complete_graph_metrics = self._complete_graph_aggregates(metrics) + if complete_graph_metrics and (metric_models or len(complete_graph_metrics) > 1): + return True if len(metric_models) < 2: return False @@ -3025,6 +3185,8 @@ def _generate_with_preaggregation( calculations: dict[str, str] = {} leaf_refs: list[str] = [] + complete_groups: dict[str, str] = {} + complete_metrics = self._complete_graph_aggregates(metrics) # Children need stable, distinct output names even when public fields # share a basename or a calculation hides a colliding aggregate leaf. child_aliases = { @@ -3039,12 +3201,26 @@ def expand_metric(reference: str, context: str | None = None, stack: tuple[str, raise ValueError(f"Circular metric dependency involving {reference}") model_name, metric = self.graph.resolve_metric_reference(reference) aggregate_models = self._find_aggregate_metric_models([reference]) - if model_name is not None and aggregate_models == {model_name}: + if reference in self._complete_graph_aggregates([reference]): + complete_metrics.add(reference) + if reference in complete_metrics: + if reference not in complete_groups: + index = len(complete_groups) + while ( + f"__ossie_query_{index}" in self.graph.models + or f"__ossie_query_{index}" in complete_groups.values() + ): + index += 1 + complete_groups[reference] = f"__ossie_query_{index}" + group_name = complete_groups[reference] + else: + group_name = model_name + if reference in complete_metrics or (model_name is not None and aggregate_models == {model_name}): if reference not in leaf_refs: leaf_refs.append(reference) child_aliases[reference] = f"__sidemantic_metric_{len(leaf_refs) - 1}" source_name = child_aliases[reference] - expression = f"{model_name}_preagg.{self._quote_identifier(source_name)}" + expression = f"{group_name}_preagg.{self._quote_identifier(source_name)}" if metric.agg in ("count", "count_distinct", "approx_count_distinct"): # A missing group has an empty count population. Restore # zero before evaluating formulas or their outer defaults. @@ -3099,6 +3275,7 @@ def expand_metric(reference: str, context: str | None = None, stack: tuple[str, model_name, _ = self.graph.resolve_metric_reference(metric_ref) except KeyError: model_name = None + model_name = complete_groups.get(metric_ref, model_name) if model_name: if model_name not in metrics_by_model: metrics_by_model[model_name] = [] @@ -3157,7 +3334,9 @@ def expand_metric(reference: str, context: str | None = None, stack: tuple[str, # Query-level row filters define one population, so every child query must # see them. Otherwise sibling metrics in the final row can describe different # populations. Metric filters remain at the outer aggregate grain. - all_model_names = set(metrics_by_model.keys()) + all_model_names = set(metrics_by_model.keys()) - set(complete_groups.values()) + for reference in complete_groups: + all_model_names.update(self._extract_models_from_sql(self.graph.get_metric(reference).sql)) pushdown_by_model, shared_filters, window_dim_filters = self._classify_filters_for_pushdown( row_or_leaf_filters, all_model_names ) @@ -3190,7 +3369,9 @@ def expand_metric(reference: str, context: str | None = None, stack: tuple[str, # Preserve that source's unmatched rows regardless of dimension # order. An explicit Explore scope still controls the population. child_generator = copy(self) - child_generator.base_model = self.base_model or model_name + child_generator.base_model = self.base_model or ( + None if model_name in complete_groups.values() else model_name + ) child_generator._generate_cache = {} sub_query = child_generator.generate( metrics=model_metrics, diff --git a/tests/interchange/ossie/fixtures/portable_expressions.json b/tests/interchange/ossie/fixtures/portable_expressions.json new file mode 100644 index 000000000..275cfb4a9 --- /dev/null +++ b/tests/interchange/ossie/fixtures/portable_expressions.json @@ -0,0 +1,136 @@ +[ + {"expression":"1 + 2 * 3 - 4 / 2 + 7 % 4", "expected":[8]}, + {"expression":"NOT (2 BETWEEN 3 AND 5) AND 2 IN (1,2) OR FALSE", "expected":[true]}, + {"expression":"2 NOT IN (1,3) AND 2 <> 3 AND 2 != 3 AND 2 <= 3 AND 3 >= 2", "expected":[true]}, + {"expression":"NULL IS NULL AND 1 IS NOT NULL", "expected":[true]}, + {"expression":"NULL IS NOT DISTINCT FROM NULL", "expected":[true]}, + {"expression":"1 IS DISTINCT FROM NULL", "expected":[true]}, + {"expression":"CASE WHEN 2 > 1 THEN 7 ELSE 9 END", "expected":[7]}, + {"expression":"CASE 2 WHEN 1 THEN 7 WHEN 2 THEN 9 END", "expected":[9]}, + {"expression":"CURRENT_DATE IS NOT NULL", "expected":[true]}, + {"expression":"CURRENT_DATE() IS NOT NULL", "expected":[true]}, + {"expression":"CURRENT_TIME IS NOT NULL", "expected":[true]}, + {"expression":"CURRENT_TIME() IS NOT NULL", "expected":[true]}, + {"expression":"CURRENT_TIMESTAMP IS NOT NULL", "expected":[true]}, + {"expression":"CURRENT_TIMESTAMP() IS NOT NULL", "expected":[true]}, + {"expression":"YEAR(DATE '2024-03-01')", "expected":[2024]}, + {"expression":"QUARTER(DATE '2024-05-01')", "expected":[2]}, + {"expression":"MONTH(DATE '2024-03-01')", "expected":[3]}, + {"expression":"DAY(DATE '2024-03-01')", "expected":[1]}, + {"expression":"DAYOFYEAR(DATE '2024-03-01')", "expected":[61]}, + {"expression":"HOUR(TIMESTAMP_NTZ '2024-03-01 12:34:56')", "expected":[12]}, + {"expression":"MINUTE(TIMESTAMP_NTZ '2024-03-01 12:34:56')", "expected":[34]}, + {"expression":"SECOND(TIMESTAMP_NTZ '2024-03-01 12:34:56')", "expected":[56]}, + {"expression":"EXTRACT(YEAR FROM DATE '2024-03-01')", "expected":[2024]}, + {"expression":"EXTRACT(DAYOFYEAR FROM DATE '2024-03-01')", "expected":[61]}, + {"expression":"DATE_PART('month', DATE '2024-03-01')", "expected":[3]}, + {"expression":"CAST(DATE_TRUNC('week', DATE '2024-03-03') AS DATE)", "expected":["2024-02-26"]}, + {"expression":"CAST(DATE_TRUNC('quarter', DATE '2024-05-12') AS DATE)", "expected":["2024-04-01"]}, + {"expression":"CAST(DATEADD(day, 7, DATE '2024-02-25') AS DATE)", "expected":["2024-03-03"]}, + {"expression":"CAST(DATEADD(month, -1, DATE '2024-03-15') AS DATE)", "expected":["2024-02-15"]}, + {"expression":"DATEDIFF(day, DATE '2024-02-25', DATE '2024-03-03')", "expected":[7]}, + {"expression":"DATEDIFF(month, DATE '2023-12-31', DATE '2024-02-01')", "expected":[2]}, + {"expression":"DATEDIFF(year, DATE '2023-12-31', DATE '2024-01-01')", "expected":[1]}, + {"expression":"TIME '12:34:56'", "expected":["12:34:56"]}, + {"expression":"TO_DATE('2024-03-01')", "expected":["2024-03-01"]}, + {"expression":"TO_TIMESTAMP('2024-03-01 12:34:56')", "expected":["2024-03-01T12:34:56"]}, + {"expression":"CAST('2024-03-01 12:34:56' AS TIMESTAMP_NTZ)", "expected":["2024-03-01T12:34:56"]}, + {"expression":"CONCAT('ab', 'cd', 'ef')", "expected":["abcdef"]}, + {"expression":"'ab' || 'cd'", "expected":["abcd"]}, + {"expression":"LENGTH('café')", "expected":[4]}, + {"expression":"LOWER('AbC')", "expected":["abc"]}, + {"expression":"UPPER('AbC')", "expected":["ABC"]}, + {"expression":"TRIM(' abc ')", "expected":["abc"]}, + {"expression":"LTRIM(' abc ')", "expected":["abc "]}, + {"expression":"RTRIM(' abc ')", "expected":[" abc"]}, + {"expression":"LEFT('abcdef', 2)", "expected":["ab"]}, + {"expression":"RIGHT('abcdef', 2)", "expected":["ef"]}, + {"expression":"SUBSTRING('abcdef', 2, 3)", "expected":["bcd"]}, + {"expression":"REPLACE('aba', 'a', 'x')", "expected":["xbx"]}, + {"expression":"SPLIT_PART('a-b-c', '-', 2)", "expected":["b"]}, + {"expression":"POSITION('bc' IN 'abcd')", "expected":[2]}, + {"expression":"CHARINDEX('bc', 'abcd')", "expected":[2]}, + {"expression":"CONTAINS('Abcd', 'bc')", "expected":[true]}, + {"expression":"CONTAINS('Abcd', 'BC')", "expected":[false]}, + {"expression":"STARTSWITH('ab_cd', 'ab_')", "expected":[true]}, + {"expression":"ENDSWITH('ab_cd', '_cd')", "expected":[true]}, + {"expression":"'Abcd' LIKE 'A_c%'", "expected":[true]}, + {"expression":"'Abcd' ILIKE 'a_c%'", "expected":[true]}, + {"expression":"REGEXP_LIKE('abcd', '^a.*d$')", "expected":[true]}, + {"expression":"REGEXP_LIKE('abcd', 'bc')", "expected":[true]}, + {"expression":"ABS(-12.3)", "expected":[12.3]}, + {"expression":"ROUND(12.345, 2)", "expected":[12.35]}, + {"expression":"FLOOR(-12.3)", "expected":[-13]}, + {"expression":"CEIL(-12.3)", "expected":[-12]}, + {"expression":"CEILING(12.3)", "expected":[13]}, + {"expression":"TRUNC(-12.345, 2)", "expected":[-12.34]}, + {"expression":"TRUNC(1.15, 2)", "expected":[1.15]}, + {"expression":"TRUNC(-0.29, 2)", "expected":[-0.29]}, + {"expression":"TRUNCATE(123.45, -1)", "expected":[120]}, + {"expression":"MOD(7, 3)", "expected":[1]}, + {"expression":"SIGN(-12)", "expected":[-1]}, + {"expression":"POWER(2, 3)", "expected":[8]}, + {"expression":"SQRT(9)", "expected":[3]}, + {"expression":"EXP(0)", "expected":[1]}, + {"expression":"LN(1)", "expected":[0]}, + {"expression":"LOG(2, 8)", "expected":[3]}, + {"expression":"LOG10(100)", "expected":[2]}, + {"expression":"GREATEST(1, 7, 3)", "expected":[7]}, + {"expression":"LEAST(1, 7, 3)", "expected":[1]}, + {"expression":"IF(TRUE, 7, 9)", "expected":[7]}, + {"expression":"IFF(FALSE, 7, 9)", "expected":[9]}, + {"expression":"NULLIF(1,1)", "expected":[null]}, + {"expression":"COALESCE(NULL,NULL,3)", "expected":[3]}, + {"expression":"IFNULL(NULL,3)", "expected":[3]}, + {"expression":"NVL(NULL,3)", "expected":[3]}, + {"expression":"NVL2(NULL,3,7)", "expected":[7]}, + {"expression":"NVL2(1,3,7)", "expected":[3]}, + {"expression":"ZEROIFNULL(NULL)", "expected":[0]}, + {"expression":"NULLIFZERO(0)", "expected":[null]}, + {"expression":"CAST(12 AS VARCHAR)", "expected":["12"]}, + {"expression":"CAST('12' AS STRING)", "expected":["12"]}, + {"expression":"CAST('12' AS INTEGER)", "expected":[12]}, + {"expression":"CAST('12' AS BIGINT)", "expected":[12]}, + {"expression":"CAST('12.3' AS DECIMAL(5,2))", "expected":[12.3]}, + {"expression":"CAST('12.3' AS NUMERIC(5,2))", "expected":[12.3]}, + {"expression":"CAST('12.3' AS FLOAT)", "expected":[12.3]}, + {"expression":"CAST('12.3' AS DOUBLE)", "expected":[12.3]}, + {"expression":"CAST('true' AS BOOLEAN)", "expected":[true]}, + {"expression":"CAST('2024-03-01' AS DATE)", "expected":["2024-03-01"]}, + {"expression":"CAST('2024-03-01 12:34:56' AS TIMESTAMP)", "expected":["2024-03-01T12:34:56"]}, + {"expression":"CAST('12:34:56' AS TIME)", "expected":["12:34:56"]}, + {"expression":"SUM(x)", "from_sql":"(VALUES (1),(2),(3),(NULL)) AS t(x)", "expected":[6]}, + {"expression":"COUNT(x)", "from_sql":"(VALUES (1),(2),(3),(NULL)) AS t(x)", "expected":[3]}, + {"expression":"COUNT(*)", "from_sql":"(VALUES (1),(2),(3),(NULL)) AS t(x)", "expected":[4]}, + {"expression":"COUNT(DISTINCT x)", "from_sql":"(VALUES (1),(1),(3),(NULL)) AS t(x)", "expected":[2]}, + {"expression":"SUM(DISTINCT x)", "from_sql":"(VALUES (1),(1),(3),(NULL)) AS t(x)", "expected":[4]}, + {"expression":"AVG(x)", "from_sql":"(VALUES (1),(2),(3),(NULL)) AS t(x)", "expected":[2]}, + {"expression":"MIN(x)", "from_sql":"(VALUES (1),(2),(3),(NULL)) AS t(x)", "expected":[1]}, + {"expression":"MAX(x)", "from_sql":"(VALUES (1),(2),(3),(NULL)) AS t(x)", "expected":[3]}, + {"expression":"STDDEV(x)", "from_sql":"(VALUES (1),(2),(3),(NULL)) AS t(x)", "expected":[1]}, + {"expression":"STDDEV_SAMP(x)", "from_sql":"(VALUES (1),(2),(3),(NULL)) AS t(x)", "expected":[1]}, + {"expression":"STDDEV_POP(x)", "from_sql":"(VALUES (1),(2),(3),(NULL)) AS t(x)", "expected":[0.816496580927726]}, + {"expression":"VARIANCE(x)", "from_sql":"(VALUES (1),(2),(3),(NULL)) AS t(x)", "expected":[1]}, + {"expression":"VAR_SAMP(x)", "from_sql":"(VALUES (1),(2),(3),(NULL)) AS t(x)", "expected":[1]}, + {"expression":"VAR_POP(x)", "from_sql":"(VALUES (1),(2),(3),(NULL)) AS t(x)", "expected":[0.6666666666666666]}, + {"expression":"MEDIAN(x)", "from_sql":"(VALUES (1),(2),(8),(NULL)) AS t(x)", "expected":[2]}, + {"expression":"PERCENTILE_CONT(0.25) WITHIN GROUP (ORDER BY x)", "from_sql":"(VALUES (1),(2),(8),(NULL)) AS t(x)", "expected":[1.5]}, + {"expression":"PERCENTILE_DISC(0.25) WITHIN GROUP (ORDER BY x)", "from_sql":"(VALUES (1),(2),(8),(NULL)) AS t(x)", "expected":[1]}, + {"expression":"PERCENTILE_CONT(0.25) WITHIN GROUP (ORDER BY x DESC)", "from_sql":"(VALUES (1),(2),(8),(NULL)) AS t(x)", "expected":[5]}, + {"expression":"PERCENTILE_DISC(0.25) WITHIN GROUP (ORDER BY x DESC)", "from_sql":"(VALUES (1),(2),(8),(NULL)) AS t(x)", "expected":[8]}, + {"expression":"SUM(CASE WHEN x > 1 THEN x ELSE 0 END)", "from_sql":"(VALUES (1),(2),(3),(NULL)) AS t(x)", "expected":[5]}, + {"expression":"ROW_NUMBER() OVER (ORDER BY x)", "from_sql":"(VALUES (1),(2),(3)) AS t(x)", "expected":[1,2,3]}, + {"expression":"RANK() OVER (ORDER BY x)", "from_sql":"(VALUES (1),(1),(3)) AS t(x)", "expected":[1,1,3]}, + {"expression":"DENSE_RANK() OVER (ORDER BY x)", "from_sql":"(VALUES (1),(1),(3)) AS t(x)", "expected":[1,1,2]}, + {"expression":"NTILE(2) OVER (ORDER BY x)", "from_sql":"(VALUES (1),(2),(3)) AS t(x)", "expected":[1,1,2]}, + {"expression":"PERCENT_RANK() OVER (ORDER BY x)", "from_sql":"(VALUES (1),(2),(3)) AS t(x)", "expected":[0,0.5,1]}, + {"expression":"CUME_DIST() OVER (ORDER BY x)", "from_sql":"(VALUES (1),(2),(3)) AS t(x)", "expected":[0.3333333333333333,0.6666666666666666,1]}, + {"expression":"LAG(x, 1, 0) OVER (ORDER BY x)", "from_sql":"(VALUES (1),(2),(3)) AS t(x)", "expected":[0,1,2]}, + {"expression":"LEAD(x, 1, 0) OVER (ORDER BY x)", "from_sql":"(VALUES (1),(2),(3)) AS t(x)", "expected":[2,3,0]}, + {"expression":"FIRST_VALUE(x) OVER (ORDER BY x)", "from_sql":"(VALUES (1),(2),(3)) AS t(x)", "expected":[1,1,1]}, + {"expression":"LAST_VALUE(x) OVER (ORDER BY x ROWS BETWEEN UNBOUNDED PRECEDING AND CURRENT ROW)", "from_sql":"(VALUES (1),(2),(3)) AS t(x)", "expected":[1,2,3]}, + {"expression":"NTH_VALUE(x, 2) OVER (ORDER BY x)", "from_sql":"(VALUES (1),(2),(3)) AS t(x)", "expected":[null,2,2]}, + {"expression":"SUM(x) OVER (ORDER BY x ROWS BETWEEN 1 PRECEDING AND 1 FOLLOWING)", "from_sql":"(VALUES (1),(2),(3)) AS t(x)", "expected":[3,6,5]}, + {"expression":"AVG(x) OVER (ORDER BY x RANGE BETWEEN UNBOUNDED PRECEDING AND CURRENT ROW)", "from_sql":"(VALUES (1),(2),(3)) AS t(x)", "expected":[1,1.5,2]}, + {"expression":"SUM(x) OVER (PARTITION BY x)", "from_sql":"(VALUES (1),(1),(3)) AS t(x)", "expected":[2,2,3], "sort_results":true} +] diff --git a/tests/interchange/ossie/test_expression_conformance.py b/tests/interchange/ossie/test_expression_conformance.py new file mode 100644 index 000000000..f22465c21 --- /dev/null +++ b/tests/interchange/ossie/test_expression_conformance.py @@ -0,0 +1,126 @@ +from __future__ import annotations + +import json + +import pytest + +from sidemantic import SemanticLayer +from sidemantic.interchange.ossie import OssieParseOptions, lower_ossie_document, parse_ossie_document + + +def _expression(sql: str, dialect: str = "ANSI_SQL") -> dict: + return {"dialects": [{"dialect": dialect, "expression": sql}]} + + +def _lower( + *, + sql: str = "SUM(orders.amount)", + field_sql: str = "amount * 2", + flat: bool = False, + dataset: str = "Orders", + field: str = "Amount", + extra_metrics: list | None = None, + source_sql: str = "SELECT 10 AS amount UNION ALL SELECT 20 AS amount", +): + scope = { + "name": "commerce", + "datasets": [ + { + "name": dataset, + "source": source_sql, + "fields": [{"name": field, "expression": _expression(field_sql)}], + } + ], + "metrics": [{"name": "total", "expression": _expression(sql)}, *(extra_metrics or [])], + } + source = {"version": "0.2.0.dev0", **scope} if flat else {"version": "0.2.0.dev0", "semantic_model": [scope]} + parsed = parse_ossie_document(json.dumps(source).encode(), options=OssieParseOptions(validate_schema=True)) + return lower_ossie_document(parsed, target_dialect="duckdb") + + +@pytest.mark.parametrize("flat", [False, True]) +@pytest.mark.parametrize("sql", ["SUM(orders.amount)", 'SUM("ORDERS"."AMOUNT")', "SUM(Orders.Amount)"]) +def test_metric_references_bind_computed_fields_before_physical_fallback(sql, flat): + result = _lower(sql=sql, flat=flat) + assert result.valid, result.diagnostics + layer = SemanticLayer.from_catalog(result.catalog, engine="python", fallback=False, auto_register=False) + assert layer.query(metrics=["total"]).fetchall() == [(60,)] + + +def test_quoted_declarations_bind_without_losing_their_identity(): + result = _lower(dataset='"orders"', field='"amount"', sql='SUM("orders"."amount")') + assert result.valid, result.diagnostics + layer = SemanticLayer.from_catalog(result.catalog, engine="python", fallback=False, auto_register=False) + assert layer.query(metrics=["total"]).fetchall() == [(60,)] + + +@pytest.mark.parametrize("name", ["order items", "orders.total", "2orders", '"broken', '"bad"quote"', '""']) +@pytest.mark.parametrize("flat", [False, True]) +def test_malformed_identifier_declarations_fail_before_compilation(name, flat): + result = _lower(dataset=name, flat=flat) + assert not result.valid + assert any(d.code == "ossie.semantic.identifier.invalid" for d in result.diagnostics) + + +@pytest.mark.parametrize("source", ["SELECT", "SELECT FROM orders"]) +@pytest.mark.parametrize("flat", [False, True]) +def test_empty_query_source_projection_is_rejected(source, flat): + result = _lower(source_sql=source, flat=flat) + assert not result.valid + + +def test_quoted_dataset_does_not_match_regular_declaration(): + result = _lower(sql='SUM("orders".amount)') + assert not result.valid + assert not result.catalog.scope_ids + assert any(d.code == "ossie.lowering.metric_unexecutable" for d in result.diagnostics) + + +@pytest.mark.parametrize("field_sql", ["SUM(amount)", "AVG(amount)", "COUNT(*)", "amount AS renamed", "*"]) +@pytest.mark.parametrize("flat", [False, True]) +def test_invalid_row_expressions_are_rejected_before_querying(field_sql, flat): + result = _lower(field_sql=field_sql, flat=flat) + assert not result.valid + assert not result.catalog.scope_ids + diagnostic = next(d for d in result.diagnostics if d.code == "ossie.lowering.expression_invalid") + assert diagnostic.json_pointer == ("" if flat else "/semantic_model/0") + "/datasets/0/fields/0/expression" + + +def test_metric_dependencies_expand_after_normalized_lookup(): + result = _lower(sql="BASE * 2", extra_metrics=[{"name": "base", "expression": _expression("SUM(orders.amount)")}]) + assert result.valid, result.diagnostics + layer = SemanticLayer.from_catalog(result.catalog, engine="python", fallback=False, auto_register=False) + assert layer.query(metrics=["total"]).fetchall() == [(120,)] + + +def test_metric_cycles_are_rejected(): + result = _lower(sql="base * 2", extra_metrics=[{"name": "base", "expression": _expression("total / 2")}]) + assert not result.valid + assert any("Cyclic metric reference" in d.message for d in result.diagnostics) + + +def test_quoted_and_regular_dataset_names_do_not_merge_runtime_identity(): + source = { + "version": "0.2.0.dev0", + "name": "commerce", + "datasets": [ + { + "name": name, + "source": f"SELECT {amount} AS amount", + "fields": [{"name": "amount", "expression": _expression("amount")}], + } + for name, amount in [("Orders", 10), ('"Orders"', 20)] + ], + "metrics": [ + {"name": "regular", "expression": _expression("SUM(orders.amount)")}, + {"name": "quoted", "expression": _expression('SUM("Orders".amount)')}, + ], + } + result = lower_ossie_document(parse_ossie_document(json.dumps(source).encode()), target_dialect="duckdb") + assert result.valid, result.diagnostics + graph = result.catalog["commerce"].graph + assert len(graph.models) == 2 + assert {model.metadata["ossie_source_name"] for model in graph.models.values()} == {"Orders", '"Orders"'} + layer = SemanticLayer.from_catalog(result.catalog, engine="python", fallback=False, auto_register=False) + assert layer.query(metrics=["regular"]).fetchall() == [(10,)] + assert layer.query(metrics=["quoted"]).fetchall() == [(20,)] diff --git a/tests/interchange/ossie/test_portable_expression.py b/tests/interchange/ossie/test_portable_expression.py new file mode 100644 index 000000000..2e6cf9a41 --- /dev/null +++ b/tests/interchange/ossie/test_portable_expression.py @@ -0,0 +1,124 @@ +from __future__ import annotations + +import datetime +import decimal +import json +from pathlib import Path + +import duckdb +import pytest +import sqlglot + +from sidemantic.interchange.ossie.portable import lower_ossie_sql + +CORPUS = json.loads((Path(__file__).parent / "fixtures" / "portable_expressions.json").read_text()) + + +def normalize_result(value): + if isinstance(value, (datetime.datetime, datetime.date, datetime.time)): + return value.isoformat() + if isinstance(value, decimal.Decimal): + return float(value) + return value + + +@pytest.mark.parametrize("case", CORPUS, ids=lambda case: case["expression"]) +def test_required_portable_functions_execute_with_independent_expected_results(case): + sql = lower_ossie_sql(case["expression"], "duckdb") + source = f" FROM {case['from_sql']}" if "from_sql" in case else "" + with duckdb.connect() as connection: + results = [normalize_result(row[0]) for row in connection.execute(f"SELECT {sql}{source}").fetchall()] + if case.get("sort_results"): + results.sort() + assert len(results) == len(case["expected"]) + for result, expected in zip(results, case["expected"], strict=True): + if isinstance(expected, float): + assert result == pytest.approx(expected) + else: + assert result == expected + + +@pytest.mark.parametrize( + "expression", + ["CURRENT_DATE", "CURRENT_DATE()", "CURRENT_TIME", "CURRENT_TIME()", "CURRENT_TIMESTAMP", "CURRENT_TIMESTAMP()"], +) +def test_current_date_and_time_forms(expression): + sql = lower_ossie_sql(expression, "duckdb") + with duckdb.connect() as connection: + assert connection.execute(f"SELECT {sql} IS NOT NULL").fetchone() == (True,) + + +@pytest.mark.parametrize( + "expression", + [ + "SELECT x FROM t", + "x IN (SELECT x FROM t)", + "WITH t AS (SELECT 1) SELECT * FROM t", + "DROP TABLE t", + "INSERT INTO t VALUES (1)", + "UPDATE t SET x = 1", + "DELETE FROM t", + "x; y", + "x AS y", + "*", + "a.b.c", + '"' + "a" * 129 + '"', + ], +) +def test_rejects_disallowed_expression_constructs(expression): + with pytest.raises(ValueError): + lower_ossie_sql(expression, "duckdb") + + +@pytest.mark.parametrize("target", ["duckdb", "postgres", "snowflake", "bigquery", "databricks"]) +def test_target_translations_do_not_drop_required_semantics(target): + for expression in ["LOG(2, x)", "TRUNC(x, 2)", "DAYOFYEAR(d)", "TO_TIMESTAMP(s)", "CONTAINS(s,p)", "ENDSWITH(s,p)"]: + sql = lower_ossie_sql(expression, target) + assert sqlglot.parse_one(sql, read=target) is not None + if expression == "LOG(2, x)": + assert "LN(x)" in sql and "/ LN(2)" in sql + if expression == "TRUNC(x, 2)": + assert "ROUND(x, 2)" in sql and "0.01" in sql + if expression == "CONTAINS(s,p)": + assert "CONTAINS_SUBSTR" not in sql + if expression == "TO_TIMESTAMP(s)": + assert "CAST(s AS" in sql + + +def test_exact_percentiles_never_become_approximate(): + expression = "PERCENTILE_CONT(.25) WITHIN GROUP (ORDER BY x)" + assert "PERCENTILE_CONT" in lower_ossie_sql(expression, "databricks") + with pytest.raises(ValueError, match="query-level"): + lower_ossie_sql(expression, "bigquery") + with pytest.raises(ValueError, match="query-level"): + lower_ossie_sql("MEDIAN(x)", "bigquery") + + +def test_unknown_extension_functions_are_preserved(): + assert lower_ossie_sql("MY_EXTENSION(x)", "duckdb") == "MY_EXTENSION(x)" + + +def test_week_truncation_uses_monday_independently_of_target_defaults(): + assert "WEEK(MONDAY)" in lower_ossie_sql("DATE_TRUNC('week',d)", "bigquery") + assert "DAYOFWEEKISO" in lower_ossie_sql("DATE_TRUNC('week',d)", "snowflake") + + +def test_postgres_calendar_difference_arithmetic_keeps_parentheses(): + # The generated expression is ANSI arithmetic, executable in DuckDB too. + # This verifies the calculation, without claiming PostgreSQL execution. + sql = lower_ossie_sql("DATEDIFF(month, DATE '2023-12-31', DATE '2024-02-01')", "postgres") + with duckdb.connect() as connection: + assert connection.execute(f"SELECT {sql}").fetchone() == (2,) + + +def test_ansi_value_window_frames_preserve_explicit_frames_and_other_source_functions(): + source = "ZEROIFNULL(NTH_VALUE(DAYOFYEAR(d), 2) OVER (ORDER BY d))" + sql = lower_ossie_sql(source, "duckdb") + with duckdb.connect() as connection: + assert connection.execute( + f"SELECT {sql} FROM (VALUES (DATE '2024-01-01'), (DATE '2024-02-01')) AS t(d)" + ).fetchall() == [(0,), (32,)] + explicit = lower_ossie_sql( + "LAST_VALUE(x) OVER (ORDER BY x ROWS BETWEEN UNBOUNDED PRECEDING AND UNBOUNDED FOLLOWING)", "duckdb" + ) + assert connection.execute(f"SELECT {explicit} FROM (VALUES (1),(2)) AS t(x)").fetchall() == [(2,), (2,)] diff --git a/tests/interchange/ossie/test_runtime_queries.py b/tests/interchange/ossie/test_runtime_queries.py new file mode 100644 index 000000000..a28c3f170 --- /dev/null +++ b/tests/interchange/ossie/test_runtime_queries.py @@ -0,0 +1,142 @@ +"""Execute imported logical fields and aggregates against adversarial row grains.""" + +import json + +import pytest + +from sidemantic import SemanticLayer +from sidemantic.interchange.ossie import OssieParseOptions, lower_ossie_document, parse_ossie_document + + +def field(name, expression=None): + return {"name": name, "expression": {"dialects": [{"dialect": "ANSI_SQL", "expression": expression or name}]}} + + +def commerce(*, source_key="customer_id", target_key="id", key_kind="primary_key", budgets=(100, 200)): + customers = { + "name": "customers", + "source": "customers", + "fields": [field("id", target_key), field("name"), field("budget", "budget * 1")], + key_kind: ["id"] if key_kind == "primary_key" else [["id"]], + } + document = { + "version": "0.2.0.dev0", + "semantic_model": [ + { + "name": "commerce", + "datasets": [ + { + "name": "orders", + "source": "orders", + "primary_key": ["id"], + "fields": [field("id"), field("customer_ref", source_key), field("amount", "amount * 1")], + }, + customers, + ], + "relationships": [ + { + "name": "customer", + "from": "orders", + "to": "customers", + "from_columns": ["customer_ref"], + "to_columns": ["id"], + } + ], + "metrics": [ + field(name, sql) + for name, sql in [ + ("revenue", "SUM(orders.amount)"), + ("budget", "SUM(customers.budget)"), + ("customer_count", "COUNT(customers.id)"), + ("mean_budget", "AVG(customers.budget)"), + ("ratio", "SUM(orders.amount) / SUM(customers.budget)"), + ("double_revenue", "revenue * 2"), + ("weighted", "SUM(orders.amount * customers.budget)"), + ("combined", "weighted + budget"), + ("filtered_budget", "SUM(CASE WHEN customers.id = 1 THEN customers.budget END)"), + ("average_order", "SUM(orders.amount) / COUNT(*)"), + ] + ], + } + ], + } + parsed = parse_ossie_document(json.dumps(document).encode(), options=OssieParseOptions(target_dialect="duckdb")) + lowered = lower_ossie_document(parsed) + assert lowered.valid, lowered.diagnostics + layer = SemanticLayer.from_catalog(lowered.catalog, engine="python", fallback=False, auto_register=False) + layer.adapter.conn.execute("create table orders(id int, customer_id int, amount int)") + layer.adapter.conn.execute("insert into orders values (1,1,10),(2,1,20),(3,2,30)") + layer.adapter.conn.execute("create table customers(id int, name varchar, budget int)") + layer.adapter.conn.executemany("insert into customers values (?,?,?)", [(1, "A", budgets[0]), (2, "B", budgets[1])]) + return layer + + +@pytest.mark.parametrize("key_kind", ["primary_key", "unique_keys"]) +def test_independent_aggregate_grains_and_ratio(key_kind): + layer = commerce(key_kind=key_kind) + assert layer.query(metrics=["budget", "customer_count", "mean_budget"]).fetchall() == [(300, 2, 150)] + assert layer.query(metrics=["revenue", "budget", "customer_count", "mean_budget", "ratio"]).fetchall() == [ + (60, 300, 2, 150, 0.2) + ] + assert layer.query( + metrics=["budget"], dimensions=["orders.customer_ref"], order_by=["orders.customer_ref"] + ).fetchall() == [(1, 100), (2, 200)] + assert layer.query( + metrics=["revenue", "budget"], dimensions=["customers.name"], order_by=["customers.name"] + ).fetchall() == [("A", 30, 100), ("B", 30, 200)] + + +def test_equal_values_are_distinct_entities_and_filters_keep_row_grain(): + layer = commerce(budgets=(100, 100)) + assert layer.query( + metrics=["budget", "customer_count"], dimensions=["orders.customer_ref"], order_by=["orders.customer_ref"] + ).fetchall() == [(1, 100, 1), (2, 100, 1)] + assert layer.query(metrics=["revenue", "budget", "customer_count"]).fetchall() == [(60, 200, 2)] + assert layer.query(metrics=["revenue", "budget"], filters=["customers.name = 'A'"]).fetchall() == [(30, 100)] + assert layer.query(metrics=["budget"], filters=["orders.amount >= 20"]).fetchall() == [(200,)] + + +@pytest.mark.parametrize("side", ["source", "target"]) +def test_computed_join_fields_with_complete_foreign_keys(side): + layer = commerce( + source_key="customer_id + 1" if side == "source" else "customer_id", + target_key="id + 1" if side == "target" else "id", + ) + layer.adapter.conn.execute("delete from orders") + values = [(1, 0, 10), (2, 1, 20)] if side == "source" else [(1, 2, 10), (2, 3, 20)] + layer.adapter.conn.executemany("insert into orders values (?,?,?)", values) + assert layer.query(metrics=["revenue"], dimensions=["customers.name"], order_by=["customers.name"]).fetchall() == [ + ("A", 10), + ("B", 20), + ] + + +def test_metric_dependencies_and_cross_dataset_row_expressions(): + layer = commerce() + assert layer.query(metrics=["double_revenue"]).fetchall() == [(120,)] + assert layer.query(metrics=["double_revenue", "budget"]).fetchall() == [(120, 300)] + assert layer.query(metrics=["weighted"]).fetchall() == [(9000,)] + assert layer.query(metrics=["weighted", "revenue", "budget"]).fetchall() == [(9000, 60, 300)] + assert layer.query(metrics=["weighted", "revenue", "budget", "combined"]).fetchall() == [(9000, 60, 300, 9300)] + assert layer.query(metrics=["combined"]).fetchall() == [(9300,)] + assert layer.query(metrics=["combined"], dimensions=["customers.name"], order_by=["customers.name"]).fetchall() == [ + ("A", 3100), + ("B", 6200), + ] + assert layer.query(metrics=["combined"], filters=["orders.amount >= 20"]).fetchall() == [(8300,)] + assert layer.query(metrics=["weighted", "budget"]).fetchall() == [(9000, 300)] + assert layer.query( + metrics=["weighted", "budget"], dimensions=["customers.name"], order_by=["customers.name"] + ).fetchall() == [("A", 3000, 100), ("B", 6000, 200)] + assert layer.query(metrics=["weighted", "budget"], filters=["orders.amount >= 20"]).fetchall() == [(8000, 300)] + # Query-local lowering must not mutate retained runtime or compilation state. + assert layer.graph.metrics["revenue"].sql_is_complete + assert layer.graph.models["orders"].metrics == [] + + +def test_filtered_aggregate_and_columnless_leaf_keep_their_population(): + layer = commerce() + assert layer.query(metrics=["revenue", "filtered_budget", "average_order"]).fetchall() == [(60, 100, 20)] + assert layer.query( + metrics=["filtered_budget"], dimensions=["orders.customer_ref"], order_by=["orders.customer_ref"] + ).fetchall() == [(1, 100), (2, None)] diff --git a/tests/interchange/ossie/test_synthesis.py b/tests/interchange/ossie/test_synthesis.py index 02f7bf6b6..3c8caa5ba 100644 --- a/tests/interchange/ossie/test_synthesis.py +++ b/tests/interchange/ossie/test_synthesis.py @@ -704,3 +704,49 @@ def test_stable_version_export_keeps_enveloped_shape(): ) assert result.valid, result.diagnostics assert result.document.to_parsed_data()["semantic_model"][0]["name"] == "commerce" + + +@pytest.mark.parametrize("dataset_name", ["orders", '"orders"', '"Order Items"']) +def test_imported_identifier_provenance_remains_portable_with_bound_runtime_names(dataset_name): + data = { + "version": "0.2.0.dev0", + "name": "commerce", + "datasets": [ + { + "name": dataset_name, + "source": "SELECT 10 AS amount", + "fields": [ + { + "name": '"line amount"', + "expression": {"dialects": [{"dialect": "ANSI_SQL", "expression": "amount"}]}, + } + ], + } + ], + "metrics": [ + { + "name": '"total amount"', + "expression": { + "dialects": [{"dialect": "ANSI_SQL", "expression": f'SUM({dataset_name}."line amount")'}] + }, + } + ], + } + lowered = lower_ossie_document(parse_ossie_document(json.dumps(data).encode()), target_dialect="duckdb") + assert lowered.valid, lowered.diagnostics + graph = lowered.catalog["commerce"].graph + result = synthesize_ossie_document(graph, scope_name="commerce", expression_dialect="ANSI_SQL", portable_only=True) + assert result.valid, result.diagnostics + reloaded = lower_ossie_document( + parse_ossie_document(json.dumps(result.document.to_parsed_data()).encode()), target_dialect="duckdb" + ) + assert reloaded.valid, reloaded.diagnostics + metric_name = next(iter(graph.metrics)) + for catalog in (lowered.catalog, reloaded.catalog): + layer = SemanticLayer.from_catalog(catalog, engine="python", fallback=False, auto_register=False) + assert layer.query(metrics=[metric_name]).fetchall() == [(10,)] + model = next(iter(graph.models.values())) + model.dimensions[0].metadata["custom_behavior"] = "preserve" + refused = synthesize_ossie_document(graph, scope_name="commerce", expression_dialect="ANSI_SQL", portable_only=True) + assert not refused.valid + assert any(item.code == "ossie.synthesis.native_state_unrepresented" for item in refused.diagnostics) diff --git a/tests/metrics/test_symmetric_aggs.py b/tests/metrics/test_symmetric_aggs.py index 6306520ac..bd481e18a 100644 --- a/tests/metrics/test_symmetric_aggs.py +++ b/tests/metrics/test_symmetric_aggs.py @@ -161,12 +161,12 @@ def test_fanout_join_detection_multiple_joins(): generator = SQLGenerator(graph) - # Multiple one-to-many joins SHOULD trigger symmetric aggregates for base model + # Each sibling is also repeated by the other sibling's rows. needs_symmetric = generator._has_fanout_joins("orders", ["order_items", "shipments"]) assert needs_symmetric["orders"] is True - assert needs_symmetric["order_items"] is False - assert needs_symmetric["shipments"] is False + assert needs_symmetric["order_items"] is True + assert needs_symmetric["shipments"] is True def test_symmetric_aggregates_in_sql_generation(): From f837e586abffb978140e8b6eff4e106175079f3f Mon Sep 17 00:00:00 2001 From: Nico Ritschel Date: Fri, 2 Oct 2026 06:34:22 -0700 Subject: [PATCH 2/3] Fix reviewed Ossie expression execution cases --- sidemantic/interchange/ossie/lowering.py | 7 +++- sidemantic/interchange/ossie/portable.py | 29 ++++++++++++++++- sidemantic/sql/generator.py | 29 ++++++++--------- tests/interchange/ossie/test_lowering.py | 32 +++++++++++++++++++ .../ossie/test_portable_expression.py | 22 +++++++++++++ .../interchange/ossie/test_runtime_queries.py | 30 +++++++++++++++++ 6 files changed, 132 insertions(+), 17 deletions(-) diff --git a/sidemantic/interchange/ossie/lowering.py b/sidemantic/interchange/ossie/lowering.py index e9f46933c..fa256043e 100644 --- a/sidemantic/interchange/ossie/lowering.py +++ b/sidemantic/interchange/ossie/lowering.py @@ -165,7 +165,12 @@ def _expression_for_target(expression: object, target_dialect: str) -> tuple[str if target_label and target_label in by_dialect: return by_dialect[target_label], target_label if "OSSIE_SQL_2026" in by_dialect: - return by_dialect["OSSIE_SQL_2026"], "OSSIE_SQL_2026" + from sidemantic.interchange.ossie.portable import supports_ossie_sql_target + + # A supplied ANSI alternative remains usable on targets such as Spark + # that have a SQL parser but no portable-expression lowering yet. + if supports_ossie_sql_target(normalized) or "ANSI_SQL" not in by_dialect: + return by_dialect["OSSIE_SQL_2026"], "OSSIE_SQL_2026" if "ANSI_SQL" in by_dialect: return by_dialect["ANSI_SQL"], "ANSI_SQL" return None diff --git a/sidemantic/interchange/ossie/portable.py b/sidemantic/interchange/ossie/portable.py index 97966a07c..155b18cef 100644 --- a/sidemantic/interchange/ossie/portable.py +++ b/sidemantic/interchange/ossie/portable.py @@ -15,6 +15,7 @@ from sqlglot.tokens import TokenType _TARGETS = {"postgresql": "postgres", "ansi_sql": "duckdb", "ansi": "duckdb"} +_SUPPORTED_TARGETS = {"duckdb", "postgres", "snowflake", "bigquery", "databricks"} _EXTRACTIONS = { exp.Year: "YEAR", exp.Quarter: "QUARTER", @@ -141,6 +142,11 @@ def _part(node: exp.Expression) -> str: return part +def supports_ossie_sql_target(target_dialect: str | None) -> bool: + target = (target_dialect or "duckdb").lower() + return _TARGETS.get(target, target) in _SUPPORTED_TARGETS + + def lower_ossie_sql(expression: str, target_dialect: str | None) -> str: """Translate portable SQL without silently approximating required operations. @@ -149,7 +155,7 @@ def lower_ossie_sql(expression: str, target_dialect: str | None) -> str: rewrite (its exact percentiles are analytic-only), outside scalar lowering. """ target = _TARGETS.get((target_dialect or "duckdb").lower(), (target_dialect or "duckdb").lower()) - if target not in {"duckdb", "postgres", "snowflake", "bigquery", "databricks"}: + if not supports_ossie_sql_target(target): raise ValueError(f"Unsupported OSSIE_SQL_2026 target {target_dialect!r}") root = parse_portable_expression(expression) @@ -245,6 +251,27 @@ def extract(value: exp.Expression, unit: str) -> exp.Expression: expression=exp.Sub(this=extract(end, unit), expression=extract(start, unit)), ) ) + if part in {"DAY", "HOUR", "MINUTE", "SECOND"}: + # DATEDIFF counts boundaries, not elapsed whole units. Truncate + # both endpoints before subtracting, including for negative spans. + def truncate(value: exp.Expression) -> exp.Expression: + return _call( + "DATE_TRUNC", + exp.Literal.string(part.lower()), + exp.Cast(this=value, to=exp.DataType.build("TIMESTAMP")), + ) + + seconds = exp.Paren( + this=exp.Sub(this=extract(truncate(end), "EPOCH"), expression=extract(truncate(start), "EPOCH")) + ) + return exp.Paren( + this=exp.Div( + this=seconds, + expression=exp.Literal.number({"DAY": 86400, "HOUR": 3600, "MINUTE": 60, "SECOND": 1}[part]), + safe=False, + typed=True, + ) + ) if isinstance(node, exp.WithinGroup) and isinstance(node.this, (exp.PercentileCont, exp.PercentileDisc)): if target == "bigquery": raise ValueError("BigQuery exact ordered-set percentiles require query-level lowering") diff --git a/sidemantic/sql/generator.py b/sidemantic/sql/generator.py index 06255a77a..a83e61340 100644 --- a/sidemantic/sql/generator.py +++ b/sidemantic/sql/generator.py @@ -213,7 +213,7 @@ def _lower_ossie_aggregates(self) -> "SQLGenerator | None": # projection; preserve their physical reference. column.set("table", None) continue - field_sql = dimension.sql_expr.replace("{model}.", "") + field_sql = self._strip_model_prefixes([dimension.sql_expr], model.name)[0] column.replace(exp.Paren(this=_parse_fragment(field_sql, self.dialect))) aggregation = aggregate_types.get(type(expression)) argument = expression.this @@ -2501,13 +2501,24 @@ def _build_model_cte( # Track all columns added (not just join keys) to avoid duplicates columns_added = set() + # CTEs read the physical table or a source query aliased as t. Ossie + # field SQL may instead qualify its physical inputs by the dataset name. + model_table_alias = "t" if model.sql else "" + + def replace_model_placeholder(sql_expr: str) -> str: + """Bind a field or measure expression to this CTE's source.""" + if (model.metadata or {}).get("ossie_source_kind"): + sql_expr = self._strip_model_prefixes([sql_expr], model.name)[0] + if model_table_alias: + return sql_expr.replace("{model}", model_table_alias) + return sql_expr.replace("{model}.", "") + def add_passthrough_column(column: str) -> None: if column not in columns_added: dimension = model.get_dimension(column) if (model.metadata or {}).get("ossie_source_kind") else None if dimension is not None: self._ensure_sql_dimension(model_name, dimension) - expression = self._dimension_base_expr(dimension) - expression = expression.replace("{model}", "t") if model.sql else expression.replace("{model}.", "") + expression = replace_model_placeholder(self._dimension_base_expr(dimension)) else: expression = self._quote_identifier(column) select_cols.append(f"{expression} AS {self._quote_alias(column)}") @@ -2613,18 +2624,6 @@ def add_passthrough_column(column: str) -> None: select_cols.append(f"{self._quote_identifier(fk)} AS {self._quote_alias(fk)}") columns_added.add(fk) - # Determine table alias for {model} placeholder replacement - # In CTEs, we're selecting from the raw table (or subquery AS t) - model_table_alias = "t" if model.sql else "" - - def replace_model_placeholder(sql_expr: str) -> str: - """Replace {model} placeholder with appropriate table reference.""" - if model_table_alias: - return sql_expr.replace("{model}", model_table_alias) - else: - # No alias needed - just remove {model}. - return sql_expr.replace("{model}.", "") - # Add only needed dimension columns for dimension in model.dimensions: if dimension.name in needed_dimensions and dimension.name not in columns_added: diff --git a/tests/interchange/ossie/test_lowering.py b/tests/interchange/ossie/test_lowering.py index e2b76e5d8..3df260661 100644 --- a/tests/interchange/ossie/test_lowering.py +++ b/tests/interchange/ossie/test_lowering.py @@ -53,6 +53,38 @@ def test_vendor_expression_alternatives_preserved_while_ansi_sql_is_selected() - assert lowered.document.to_parsed_data() == json.loads(source) +@pytest.mark.parametrize("with_ansi", [False, True]) +def test_spark_selects_supplied_ansi_alternative_to_portable_sql(with_ansi): + def expression(portable, ansi): + dialects = [{"dialect": "OSSIE_SQL_2026", "expression": portable}] + if with_ansi: + dialects.append({"dialect": "ANSI_SQL", "expression": ansi}) + return {"dialects": dialects} + + source = { + "version": "0.2.0.dev0", + "name": "commerce", + "datasets": [ + { + "name": "orders", + "source": "orders", + "fields": [{"name": "amount", "expression": expression("ZEROIFNULL(amount)", "COALESCE(amount, 0)")}], + } + ], + "metrics": [{"name": "revenue", "expression": expression("SUM(orders.amount)", "SUM(orders.amount)")}], + } + lowered = lower_ossie_document(_parse(json.dumps(source)), target_dialect="spark") + + assert lowered.valid is with_ansi, lowered.diagnostics + if with_ansi: + graph = lowered.catalog["commerce"].graph + assert graph.get_model("orders").get_dimension("amount").sql == "COALESCE(amount, 0)" + assert graph.get_metric("revenue").metadata["ossie_expression_dialect"] == "ANSI_SQL" + else: + assert any("Unsupported OSSIE_SQL_2026 target 'spark'" in d.message for d in lowered.diagnostics) + assert lowered.document.to_parsed_data() == source + + def test_lowers_multiple_scopes_without_flattening_duplicate_model_names() -> None: parsed = _parse( """version: 0.2.0.dev0 diff --git a/tests/interchange/ossie/test_portable_expression.py b/tests/interchange/ossie/test_portable_expression.py index 2e6cf9a41..7d6b301a5 100644 --- a/tests/interchange/ossie/test_portable_expression.py +++ b/tests/interchange/ossie/test_portable_expression.py @@ -111,6 +111,28 @@ def test_postgres_calendar_difference_arithmetic_keeps_parentheses(): assert connection.execute(f"SELECT {sql}").fetchone() == (2,) +@pytest.mark.parametrize( + "unit, start, end, expected", + [ + ("day", "2024-01-01 23:59:00", "2024-01-02 00:01:00", 1), + ("hour", "2024-01-01 10:59:00", "2024-01-01 11:01:00", 1), + ("minute", "2024-01-01 10:00:59", "2024-01-01 10:01:01", 1), + ("second", "2024-01-01 10:00:00.999999", "2024-01-01 10:00:01.000001", 1), + ("hour", "2024-01-01 10:01:00", "2024-01-01 10:59:00", 0), + ("hour", "1969-12-31 23:59:00", "1970-01-01 00:01:00", 1), + ], +) +@pytest.mark.parametrize("reverse", [False, True]) +def test_postgres_datediff_counts_boundaries(unit, start, end, expected, reverse): + if reverse: + start, end, expected = end, start, -expected + sql = lower_ossie_sql(f"DATEDIFF({unit}, TIMESTAMP '{start}', TIMESTAMP '{end}')", "postgres") + # Execute the PostgreSQL arithmetic without round-tripping it through a + # transpiler. DuckDB supports the same EXTRACT/DATE_TRUNC timestamp operators. + with duckdb.connect() as connection: + assert connection.execute(f"SELECT {sql}").fetchone() == (expected,) + + def test_ansi_value_window_frames_preserve_explicit_frames_and_other_source_functions(): source = "ZEROIFNULL(NTH_VALUE(DAYOFYEAR(d), 2) OVER (ORDER BY d))" sql = lower_ossie_sql(source, "duckdb") diff --git a/tests/interchange/ossie/test_runtime_queries.py b/tests/interchange/ossie/test_runtime_queries.py index a28c3f170..92901fa49 100644 --- a/tests/interchange/ossie/test_runtime_queries.py +++ b/tests/interchange/ossie/test_runtime_queries.py @@ -140,3 +140,33 @@ def test_filtered_aggregate_and_columnless_leaf_keep_their_population(): assert layer.query( metrics=["filtered_budget"], dimensions=["orders.customer_ref"], order_by=["orders.customer_ref"] ).fetchall() == [(1, 100), (2, None)] + + +@pytest.mark.parametrize("source", ["raw_orders", "SELECT id, amount FROM raw_orders"]) +def test_qualified_fields_bind_to_the_dataset_source(source): + document = { + "version": "0.2.0.dev0", + "name": "commerce", + "datasets": [ + { + "name": "orders", + "source": source, + "primary_key": ["id"], + "fields": [field("id", "orders.id"), field("amount", "orders.amount * 2")], + } + ], + "metrics": [field("revenue", "SUM(orders.amount)")], + } + parsed = parse_ossie_document(json.dumps(document).encode(), options=OssieParseOptions(target_dialect="duckdb")) + lowered = lower_ossie_document(parsed) + assert lowered.valid, lowered.diagnostics + layer = SemanticLayer.from_catalog(lowered.catalog, engine="python", fallback=False, auto_register=False) + layer.adapter.conn.execute("create table raw_orders(id int, amount int)") + layer.adapter.conn.execute("insert into raw_orders values (1, 10), (2, 20)") + + assert layer.query(metrics=["revenue"]).fetchall() == [(60,)] + assert layer.query(metrics=["revenue"], filters=["orders.amount > 20"]).fetchall() == [(40,)] + assert layer.query(dimensions=["orders.id", "orders.amount"], order_by=["orders.id"]).fetchall() == [ + (1, 20), + (2, 40), + ] From c48fb07c61cb79655ce4acecf26ac89cbc12f7c8 Mon Sep 17 00:00:00 2001 From: Nico Ritschel Date: Fri, 2 Oct 2026 06:52:02 -0700 Subject: [PATCH 3/3] Qualify Ossie scope parity fixture inputs --- tests/adapters/osi/test_rust_ossie_forward_parity.py | 7 ++++--- 1 file changed, 4 insertions(+), 3 deletions(-) diff --git a/tests/adapters/osi/test_rust_ossie_forward_parity.py b/tests/adapters/osi/test_rust_ossie_forward_parity.py index 12105e7d5..f53788bf1 100644 --- a/tests/adapters/osi/test_rust_ossie_forward_parity.py +++ b/tests/adapters/osi/test_rust_ossie_forward_parity.py @@ -72,8 +72,8 @@ def test_selected_scope_shape_preserves_parity_fields_and_target_selection(tmp_p datatype: Decimal expression: dialects: - - {dialect: ANSI_SQL, expression: SUM(amount)} - - {dialect: BIGQUERY, expression: SUM(SAFE_CAST(amount AS NUMERIC))} + - {dialect: ANSI_SQL, expression: SUM(orders.amount)} + - {dialect: BIGQUERY, expression: SUM(SAFE_CAST(orders.amount AS NUMERIC))} - name: operations datasets: - name: orders @@ -84,6 +84,7 @@ def test_selected_scope_shape_preserves_parity_fields_and_target_selection(tmp_p rust = rust_ossie_select_scope(content, "yaml", "commerce", target="BIGQUERY") python = OssieAdapter(scope_id="commerce", target_dialect="bigquery").parse_document(source) + assert python.valid, python.diagnostics python_graph = python.catalog["commerce"].graph assert rust["scope_id"] == "commerce" @@ -97,7 +98,7 @@ def test_selected_scope_shape_preserves_parity_fields_and_target_selection(tmp_p assert "declared_is_time" not in rust["models"][0]["dimensions"][1] assert rust["metrics"][0]["logical_data_type"] == "Decimal" assert rust["metrics"][0]["agg"] == "sum" - assert rust["metrics"][0]["sql"] == "SAFE_CAST(amount AS NUMERIC)" + assert rust["metrics"][0]["sql"] == "SAFE_CAST(orders.amount AS NUMERIC)" assert rust["models"][0]["relationships"][0]["edge_id"] == "orders_customer" assert python_graph.get_model("orders").relationships[0].edge_id == "orders_customer" assert rust["models"][0]["primary_key"] == ""