Files
govoplan-reporting/src/govoplan_reporting/backend/postgres_planner.py
T

510 lines
17 KiB
Python

from __future__ import annotations
from collections.abc import Mapping, Sequence
from datetime import date, datetime
from decimal import Decimal
import json
import re
from typing import Any
from sqlalchemy import text
from sqlalchemy.orm import Session
from govoplan_reporting.backend.query_engine import (
QueryResult,
ReportingQueryError,
infer_query_schema,
)
from govoplan_reporting.backend.schemas import (
DatasetDefinition,
DimensionDefinition,
FilterClause,
MeasureDefinition,
ReportQuery,
SemanticModelDefinition,
TypedExpression,
)
POSTGRES_PLANNER_VERSION = "reporting-postgresql-v1"
_IDENTIFIER = re.compile(r"^[a-z0-9._-]{1,120}$")
class PostgresPlanningError(ReportingQueryError):
pass
def execute_postgres_query(
session: Session,
*,
rows: Sequence[Mapping[str, object]],
dataset: DatasetDefinition,
semantic_model: SemanticModelDefinition,
query: ReportQuery,
) -> QueryResult | None:
"""Execute a bounded semantic plan in PostgreSQL, or return None for fallback."""
if session.bind is None or session.bind.dialect.name != "postgresql":
return None
if query.mode == "pivot":
return None
plan = compile_postgres_query(dataset, semantic_model, query)
parameters = {
**plan.parameters,
"rows_json": json.dumps(
[_json_value(dict(item)) for item in rows],
ensure_ascii=True,
separators=(",", ":"),
sort_keys=True,
),
"result_limit": query.limit,
"result_offset": query.offset,
}
result = session.execute(text(plan.sql), parameters).mappings().all()
total_rows = int(result[0]["__reporting_total"]) if result else 0
output = tuple(
{
str(key): _json_value(value)
for key, value in item.items()
if key != "__reporting_total"
}
for item in result
)
return QueryResult(
rows=output,
total_rows=total_rows,
schema=infer_query_schema(output),
truncated=query.offset + len(output) < total_rows,
diagnostics=(
{
"severity": "info",
"code": "postgresql_semantic_plan",
"message": "Filters, grouping, measures, ordering, and bounds were executed by the PostgreSQL Reporting planner.",
"planner_version": POSTGRES_PLANNER_VERSION,
},
),
)
class CompiledPostgresPlan:
__slots__ = ("sql", "parameters")
def __init__(self, sql: str, parameters: Mapping[str, object]) -> None:
self.sql = sql
self.parameters = dict(parameters)
def compile_postgres_query(
dataset: DatasetDefinition,
semantic_model: SemanticModelDefinition,
query: ReportQuery,
) -> CompiledPostgresPlan:
dimensions = {item.key: item for item in semantic_model.dimensions}
measures = {item.key: item for item in semantic_model.measures}
selected_dimensions = tuple(query.dimensions or semantic_model.default_dimensions)
selected_measures = tuple(query.measures or semantic_model.default_measures)
_known(selected_dimensions, dimensions, "dimensions")
_known(selected_measures, measures, "measures")
_known(
tuple(item.dimension for item in query.filters), dimensions, "filter dimensions"
)
selected_keys = set(selected_dimensions)
if query.mode != "detail":
selected_keys.update(selected_measures)
_known(
tuple(item.key for item in query.sort),
{key: True for key in selected_keys},
"sort fields",
)
parameters: dict[str, object] = {}
source = (
"WITH source AS ("
"SELECT value AS source_row "
"FROM jsonb_array_elements(CAST(:rows_json AS jsonb)) AS source_items(value)"
")"
)
where = _filter_sql(query.filters, dimensions, parameters)
if query.mode == "detail":
fields = selected_dimensions
if not fields:
if not dataset.fields:
raise PostgresPlanningError(
"PostgreSQL detail planning requires selected dimensions or a pinned dataset schema."
)
field_types = {item.name: item.type for item in dataset.fields}
projections = [
f"{_source_value(item.name, item.type, parameters, f'detail_{index}')} AS {_quote(item.name)}"
for index, item in enumerate(dataset.fields)
]
selected_keys = set(field_types)
else:
projections = [
f"{_dimension_value(dimensions[key], parameters, f'detail_{index}')} AS {_quote(key)}"
for index, key in enumerate(fields)
]
body = "SELECT " + ", ".join(projections) + " FROM source" + where
else:
dimension_projections = [
(
key,
_dimension_value(dimensions[key], parameters, f"dimension_{index}"),
)
for index, key in enumerate(selected_dimensions)
]
selected_base_keys = [
key
for key in selected_measures
if measures[key].aggregation != "calculated"
]
calculated = [
measures[key]
for key in selected_measures
if measures[key].aggregation == "calculated"
]
dependency_keys = list(
dict.fromkeys(
dependency
for item in calculated
for dependency in _calculated_dependencies(
item.expression, measures, stack=(item.key,)
)
)
)
base_measure_keys = list(dict.fromkeys((*selected_base_keys, *dependency_keys)))
base_measures = [measures[key] for key in base_measure_keys]
grouped_select = [
f"{expression} AS {_quote(key)}"
for key, expression in dimension_projections
] + [
f"{_aggregate_sql(item, parameters, index)} AS {_quote(item.key)}"
for index, item in enumerate(base_measures)
]
if not grouped_select:
raise PostgresPlanningError(
"Summary queries require at least one dimension or measure."
)
grouped = "SELECT " + ", ".join(grouped_select) + " FROM source" + where
if dimension_projections:
grouped += " GROUP BY " + ", ".join(
expression for _key, expression in dimension_projections
)
if calculated:
outer = [_quote(key) for key in selected_dimensions] + [
_quote(key) for key in selected_base_keys
]
outer.extend(
f"{_calculated_sql(item.expression, parameters, f'calculated_{index}', measures=measures, stack=(item.key,))} AS {_quote(item.key)}"
for index, item in enumerate(calculated)
)
body = "SELECT " + ", ".join(outer) + f" FROM ({grouped}) AS grouped"
else:
body = grouped
order = ""
if query.sort:
order = " ORDER BY " + ", ".join(
f"{_quote(item.key)} {item.direction.upper()} NULLS LAST"
for item in query.sort
)
sql = (
source
+ " SELECT planned.*, COUNT(*) OVER() AS __reporting_total FROM ("
+ body
+ ") AS planned"
+ order
+ " LIMIT :result_limit OFFSET :result_offset"
)
return CompiledPostgresPlan(sql, parameters)
def _filter_sql(
filters: Sequence[FilterClause],
dimensions: Mapping[str, DimensionDefinition],
parameters: dict[str, object],
) -> str:
clauses: list[str] = []
for index, clause in enumerate(filters):
value = _dimension_value(
dimensions[clause.dimension], parameters, f"filter_field_{index}"
)
prefix = f"filter_{index}"
if clause.operator == "is_null":
clauses.append(f"{value} IS NULL")
continue
if clause.operator == "not_null":
clauses.append(f"{value} IS NOT NULL")
continue
if clause.operator in {"in", "not_in"}:
if not isinstance(clause.value, (list, tuple)):
raise PostgresPlanningError("Set filters require a list value.")
if not clause.value or len(clause.value) > 500:
raise PostgresPlanningError(
"Set filters require between 1 and 500 values."
)
names: list[str] = []
for item_index, item in enumerate(clause.value):
name = f"{prefix}_{item_index}"
parameters[name] = item
names.append(f":{name}")
operator = "NOT IN" if clause.operator == "not_in" else "IN"
clauses.append(f"{value} {operator} ({', '.join(names)})")
continue
if clause.operator == "between":
if not isinstance(clause.value, (list, tuple)) or len(clause.value) != 2:
raise PostgresPlanningError("Between filters require two values.")
parameters[f"{prefix}_low"] = clause.value[0]
parameters[f"{prefix}_high"] = clause.value[1]
clauses.append(f"{value} BETWEEN :{prefix}_low AND :{prefix}_high")
continue
parameters[prefix] = clause.value
if clause.operator == "contains":
parameters[prefix] = f"%{_like(str(clause.value or ''))}%"
clauses.append(
f"LOWER(CAST({value} AS text)) LIKE LOWER(:{prefix}) ESCAPE '\\'"
)
elif clause.operator == "starts_with":
parameters[prefix] = f"{_like(str(clause.value or ''))}%"
clauses.append(
f"LOWER(CAST({value} AS text)) LIKE LOWER(:{prefix}) ESCAPE '\\'"
)
else:
operator = {
"eq": "=",
"ne": "<>",
"gt": ">",
"gte": ">=",
"lt": "<",
"lte": "<=",
}.get(clause.operator)
if operator is None:
raise PostgresPlanningError(
f"Unsupported PostgreSQL filter operator: {clause.operator}."
)
clauses.append(f"{value} {operator} :{prefix}")
return " WHERE " + " AND ".join(clauses) if clauses else ""
def _aggregate_sql(
measure: MeasureDefinition,
parameters: dict[str, object],
index: int,
) -> str:
if measure.aggregation == "count" and measure.field is None:
return "COUNT(*)"
field = _source_value(
measure.field or "",
"number" if measure.aggregation in {"sum", "average"} else "string",
parameters,
f"measure_{index}",
)
if measure.aggregation == "count":
return f"COUNT({field})"
if measure.aggregation == "count_distinct":
return f"COUNT(DISTINCT {field})"
function = {
"sum": "SUM",
"average": "AVG",
"minimum": "MIN",
"maximum": "MAX",
}.get(measure.aggregation)
if function is None:
raise PostgresPlanningError(
f"Unsupported PostgreSQL aggregation: {measure.aggregation}."
)
return f"{function}({field})"
def _calculated_sql(
expression: TypedExpression | None,
parameters: dict[str, object],
prefix: str,
*,
measures: Mapping[str, MeasureDefinition],
stack: tuple[str, ...],
) -> str:
if expression is None:
return "NULL"
if expression.op == "literal":
parameters[prefix] = expression.value
return f":{prefix}"
if expression.op == "measure":
reference = expression.ref or ""
target = measures.get(reference)
if target is None:
raise PostgresPlanningError(
f"Calculated measure references unknown measure: {reference}."
)
if target.aggregation != "calculated":
return _quote(reference)
if reference in stack:
raise PostgresPlanningError(
"Calculated measure dependency cycle: "
+ " -> ".join((*stack, reference))
)
return _calculated_sql(
target.expression,
parameters,
prefix + "_" + reference,
measures=measures,
stack=(*stack, reference),
)
if expression.op == "field":
raise PostgresPlanningError(
"Calculated aggregate measures may reference measures, not source fields."
)
values = [
_calculated_sql(
item,
parameters,
f"{prefix}_{index}",
measures=measures,
stack=stack,
)
for index, item in enumerate(expression.args)
]
if expression.op in {"add", "multiply", "and", "or"}:
operator = {"add": "+", "multiply": "*", "and": "AND", "or": "OR"}[
expression.op
]
return "(" + f" {operator} ".join(values) + ")"
if expression.op in {"subtract", "divide", "eq", "ne", "gt", "gte", "lt", "lte"}:
if len(values) != 2:
raise PostgresPlanningError(
f"Expression {expression.op} requires exactly two arguments."
)
operator = {
"subtract": "-",
"divide": "/",
"eq": "=",
"ne": "<>",
"gt": ">",
"gte": ">=",
"lt": "<",
"lte": "<=",
}[expression.op]
right = f"NULLIF({values[1]}, 0)" if expression.op == "divide" else values[1]
return f"({values[0]} {operator} {right})"
if expression.op == "not":
if len(values) != 1:
raise PostgresPlanningError("Expression not requires one argument.")
return f"(NOT {values[0]})"
if expression.op == "coalesce":
return "COALESCE(" + ", ".join(values) + ")"
if expression.op == "case":
if len(values) < 3 or len(values) % 2 == 0:
raise PostgresPlanningError(
"Case expressions require condition/value pairs and a default."
)
branches = " ".join(
f"WHEN {values[index]} THEN {values[index + 1]}"
for index in range(0, len(values) - 1, 2)
)
return f"(CASE {branches} ELSE {values[-1]} END)"
raise PostgresPlanningError(
f"Unsupported PostgreSQL expression operator: {expression.op}."
)
def _calculated_dependencies(
expression: TypedExpression | None,
measures: Mapping[str, MeasureDefinition],
*,
stack: tuple[str, ...],
) -> tuple[str, ...]:
if expression is None:
return ()
if expression.op == "measure":
reference = expression.ref or ""
target = measures.get(reference)
if target is None:
raise PostgresPlanningError(
f"Calculated measure references unknown measure: {reference}."
)
if target.aggregation != "calculated":
return (reference,)
if reference in stack:
raise PostgresPlanningError(
"Calculated measure dependency cycle: "
+ " -> ".join((*stack, reference))
)
return _calculated_dependencies(
target.expression,
measures,
stack=(*stack, reference),
)
dependencies: list[str] = []
for item in expression.args:
dependencies.extend(_calculated_dependencies(item, measures, stack=stack))
return tuple(dict.fromkeys(dependencies))
def _dimension_value(
dimension: DimensionDefinition,
parameters: dict[str, object],
prefix: str,
) -> str:
return _source_value(dimension.field, dimension.type, parameters, prefix)
def _source_value(
field: str,
field_type: str,
parameters: dict[str, object],
prefix: str,
) -> str:
parameters[prefix] = field
raw = f"source_row ->> :{prefix}"
if field_type == "integer":
return f"NULLIF({raw}, '')::bigint"
if field_type == "number":
return f"NULLIF({raw}, '')::numeric"
if field_type == "boolean":
return f"NULLIF({raw}, '')::boolean"
if field_type == "date":
return f"NULLIF({raw}, '')::date"
if field_type == "datetime":
return f"NULLIF({raw}, '')::timestamptz"
if field_type == "json":
return f"source_row -> :{prefix}"
return raw
def _known(keys: Sequence[str], available: Mapping[str, object], label: str) -> None:
unknown = set(keys) - set(available)
if unknown:
raise PostgresPlanningError(
f"Report query references unknown {label}: " + ", ".join(sorted(unknown))
)
def _quote(value: str) -> str:
if not _IDENTIFIER.fullmatch(value):
raise PostgresPlanningError(f"Unsafe Reporting identifier: {value!r}.")
return '"' + value.replace('"', '""') + '"'
def _like(value: str) -> str:
return value.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_")
def _json_value(value: object) -> Any:
if isinstance(value, Decimal):
integral = value.to_integral_value()
return int(integral) if value == integral else float(value)
if isinstance(value, (datetime, date)):
return value.isoformat()
if isinstance(value, Mapping):
return {str(key): _json_value(item) for key, item in value.items()}
if isinstance(value, (list, tuple)):
return [_json_value(item) for item in value]
return value
__all__ = [
"POSTGRES_PLANNER_VERSION",
"CompiledPostgresPlan",
"PostgresPlanningError",
"compile_postgres_query",
"execute_postgres_query",
]