From 357d159bf709a8accf12d6b71706373ee0cca672 Mon Sep 17 00:00:00 2001 From: Amin Ghadersohi Date: Sat, 26 Sep 2026 03:15:57 +0000 Subject: [PATCH 01/19] fix(sqlalchemy): infer Decimal128, Int64 and binary column types in reflection get_columns typed each field from the first sampled value, and _infer_bson_type had no branch for Decimal128, Int64 or binary values, so they fell through to String. A field whose first sampled document held a null reflected as NullType even when later documents held values. The type map lookup also lowercased its key, so the camel-case objectId and binData entries could never match. Type a field from its first non-null sampled value, recognise Decimal128 (DECIMAL), Int64 (BigInteger) and Binary/bytes (LargeBinary), and look the BSON type up without changing its case. --- .../sqlalchemy_mongodb/sqlalchemy_dialect.py | 15 +++- tests/test_sqlalchemy_reflection_types.py | 73 +++++++++++++++++++ 2 files changed, 84 insertions(+), 4 deletions(-) create mode 100644 tests/test_sqlalchemy_reflection_types.py diff --git a/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py b/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py index 540a0f1..e15bcde 100644 --- a/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py +++ b/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py @@ -384,11 +384,12 @@ def get_columns(self, connection, table_name: str, schema: Optional[str] = None, # Sample a few documents to infer schema sample_docs = list(collection.find().limit(10)) if sample_docs: - # Collect all unique field names and types + # Collect all unique field names and types. A null only + # types a field that no sampled document gives a value. field_types = {} for doc in sample_docs: for field_name, value in doc.items(): - if field_name not in field_types: + if field_types.get(field_name, "null") == "null": field_types[field_name] = self._infer_bson_type(value) # Convert to SQLAlchemy column format @@ -430,7 +431,7 @@ def _infer_bson_type(self, value: Any) -> str: """Infer BSON type from a Python value.""" from datetime import datetime - from bson import ObjectId + from bson import Binary, Decimal128, Int64, ObjectId if isinstance(value, ObjectId): return "objectId" @@ -438,8 +439,14 @@ def _infer_bson_type(self, value: Any) -> str: return "string" elif isinstance(value, bool): return "bool" + elif isinstance(value, Int64): + return "long" elif isinstance(value, int): return "int" + elif isinstance(value, Decimal128): + return "decimal" + elif isinstance(value, (Binary, bytes)): + return "binData" elif isinstance(value, float): return "double" elif isinstance(value, datetime): @@ -469,7 +476,7 @@ def _get_column_type(self, mongo_type: str) -> Type[types.TypeEngine]: "object": types.JSON, "binData": types.LargeBinary, } - return type_map.get(mongo_type.lower(), types.String) + return type_map.get(mongo_type, types.String) def get_pk_constraint(self, connection, table_name: str, schema: Optional[str] = None, **kwargs) -> Dict[str, Any]: """Get primary key constraint info. diff --git a/tests/test_sqlalchemy_reflection_types.py b/tests/test_sqlalchemy_reflection_types.py new file mode 100644 index 0000000..00ffc7e --- /dev/null +++ b/tests/test_sqlalchemy_reflection_types.py @@ -0,0 +1,73 @@ +# -*- coding: utf-8 -*- +"""Offline tests for column type inference in SQLAlchemy reflection.""" + +import uuid +from datetime import datetime +from unittest.mock import MagicMock + +import pytest + +sqlalchemy = pytest.importorskip("sqlalchemy") + +from bson import Binary, Decimal128, Int64, ObjectId # noqa: E402 +from sqlalchemy import types # noqa: E402 +from sqlalchemy.sql.sqltypes import NullType # noqa: E402 + +from pymongosql.sqlalchemy_mongodb.sqlalchemy_dialect import PyMongoSQLDialect # noqa: E402 + + +def reflect(documents): + """Run get_columns against an in-memory sample instead of a server.""" + collection = MagicMock() + collection.find.return_value.limit.return_value = documents + database = MagicMock() + database.__getitem__.return_value = collection + client = MagicMock() + client.__getitem__.return_value = database + connection = MagicMock() + connection.connection._client = client + columns = PyMongoSQLDialect().get_columns(connection, "c", schema="db") + return {c["name"]: c["type"] for c in columns} + + +def is_type(reflected, expected): + reflected = reflected if isinstance(reflected, type) else type(reflected) + return issubclass(reflected, expected) + + +def test_decimal128_reflects_as_numeric(): + reflected = reflect([{"_id": 1, "amount": Decimal128("123456789012345678901.1234567890")}]) + assert is_type(reflected["amount"], types.Numeric) + assert not is_type(reflected["amount"], types.Float) + + +def test_leading_null_uses_the_first_non_null_value(): + reflected = reflect([{"_id": 1, "optional": None}, {"_id": 2, "optional": "present"}]) + assert is_type(reflected["optional"], types.String) + + +def test_all_null_field_stays_null_type(): + reflected = reflect([{"_id": 1, "optional": None}, {"_id": 2, "optional": None}]) + assert is_type(reflected["optional"], NullType) + + +def test_first_non_null_type_is_not_overwritten_by_later_values(): + reflected = reflect([{"_id": 1, "n": 1}, {"_id": 2, "n": "text"}]) + assert is_type(reflected["n"], types.Integer) + + +@pytest.mark.parametrize( + "value,expected", + [ + (ObjectId(), types.String), + (Int64(2**40), types.BigInteger), + (Binary(b"\x00\x01"), types.LargeBinary), + (b"\x00\x01", types.LargeBinary), + (uuid.uuid4(), types.String), + (datetime(2026, 1, 1), types.DateTime), + (True, types.Boolean), + (1.5, types.Float), + ], +) +def test_bson_values_map_to_their_sqlalchemy_types(value, expected): + assert is_type(reflect([{"_id": 1, "v": value}])["v"], expected) From 77877d545a2dc75dcead0ad8eb7e7dd4a3343add Mon Sep 17 00:00:00 2001 From: Amin Ghadersohi Date: Sat, 26 Sep 2026 03:16:54 +0000 Subject: [PATCH 02/19] fix(sqlalchemy): round-trip Decimal values through Decimal128 The dialect declares supports_native_decimal, so SQLAlchemy passes decimal.Decimal parameters straight to the DBAPI and installs no Numeric result processor. PyMongo cannot encode decimal.Decimal, so binding one failed with "cannot encode object", and reads returned bson.Decimal128 instead of the decimal.Decimal a Numeric column promises. Encode bound decimal.Decimal values as Decimal128 when placeholders are replaced, and give Numeric and Float columns result processors that convert Decimal128 before the usual Numeric handling. --- pymongosql/helper.py | 21 ++++- .../sqlalchemy_mongodb/sqlalchemy_dialect.py | 30 +++++++- tests/test_sqlalchemy_numeric.py | 77 +++++++++++++++++++ 3 files changed, 125 insertions(+), 3 deletions(-) create mode 100644 tests/test_sqlalchemy_numeric.py diff --git a/pymongosql/helper.py b/pymongosql/helper.py index 6c1d2cb..6337899 100644 --- a/pymongosql/helper.py +++ b/pymongosql/helper.py @@ -6,9 +6,12 @@ """ import logging +from decimal import Decimal from typing import Any, Optional, Sequence, Tuple from urllib.parse import parse_qs, urlparse +from bson import Decimal128 + from .error import ProgrammingError _logger = logging.getLogger(__name__) @@ -102,6 +105,20 @@ def parse_connection_string(connection_string: Optional[str]) -> Tuple[Optional[ class SQLHelper: """SQL-related helper utilities.""" + @staticmethod + def to_bson_value(value: Any) -> Any: + """Convert a bound parameter to a type BSON can encode. + + ``decimal.Decimal`` has no BSON encoding; ``Decimal128`` stores it exactly. + """ + if isinstance(value, Decimal): + return Decimal128(value) + if isinstance(value, dict): + return {k: SQLHelper.to_bson_value(v) for k, v in value.items()} + if isinstance(value, (list, tuple)): + return [SQLHelper.to_bson_value(v) for v in value] + return value + @staticmethod def replace_placeholders_generic(value: Any, parameters: Any, style: Optional[str]) -> Any: """Recursively replace placeholders in nested structures for qmark or named styles.""" @@ -120,7 +137,7 @@ def replace(val: Any) -> Any: raise ProgrammingError("Not enough parameters provided") out = parameters[idx[0]] idx[0] += 1 - return out + return SQLHelper.to_bson_value(out) if isinstance(val, dict): return {k: replace(v) for k, v in val.items()} if isinstance(val, list): @@ -138,7 +155,7 @@ def replace(val: Any) -> Any: key = val[1:] if key not in parameters: raise ProgrammingError(f"Missing named parameter: {key}") - return parameters[key] + return SQLHelper.to_bson_value(parameters[key]) if isinstance(val, dict): return {k: replace(v) for k, v in val.items()} if isinstance(val, list): diff --git a/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py b/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py index e15bcde..63b44ba 100644 --- a/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py +++ b/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py @@ -5,7 +5,7 @@ from sqlalchemy import pool, types from sqlalchemy.engine import default, url -from sqlalchemy.sql import compiler +from sqlalchemy.sql import compiler, sqltypes from sqlalchemy.sql.sqltypes import NULLTYPE import pymongosql @@ -151,6 +151,32 @@ def visit_BOOLEAN(self, type_, **kwargs): return "BOOL" +def _decode_decimal128(processor): + """Wrap a Numeric result processor so it also accepts BSON Decimal128.""" + from bson import Decimal128 + + def process(value): + if isinstance(value, Decimal128): + value = value.to_decimal() + return processor(value) if processor else value + + return process + + +class _MongoNumeric(sqltypes.Numeric): + """Numeric that returns ``decimal.Decimal`` (or float) for Decimal128 values.""" + + def result_processor(self, dialect, coltype): + return _decode_decimal128(super().result_processor(dialect, coltype)) + + +class _MongoFloat(sqltypes.Float): + """Float that returns ``float`` (or Decimal) for Decimal128 values.""" + + def result_processor(self, dialect, coltype): + return _decode_decimal128(super().result_processor(dialect, coltype)) + + class PyMongoSQLDialect(default.DefaultDialect): """SQLAlchemy dialect for PyMongoSQL. @@ -174,6 +200,8 @@ class PyMongoSQLDialect(default.DefaultDialect): supports_empty_inserts = True supports_multivalues_insert = True supports_native_decimal = True # BSON Decimal128 + # PyMongo returns Decimal128, not decimal.Decimal; convert on the way out. + colspecs = {sqltypes.Numeric: _MongoNumeric, sqltypes.Float: _MongoFloat} supports_native_boolean = True # BSON Boolean supports_sequences = False # No sequences in MongoDB supports_native_enum = False # No native enums diff --git a/tests/test_sqlalchemy_numeric.py b/tests/test_sqlalchemy_numeric.py new file mode 100644 index 0000000..e294a7d --- /dev/null +++ b/tests/test_sqlalchemy_numeric.py @@ -0,0 +1,77 @@ +# -*- coding: utf-8 -*- +"""Decimal values must round-trip exactly through the SQLAlchemy dialect.""" + +from decimal import Decimal + +import pytest + +from pymongosql.helper import SQLHelper +from tests.conftest import HAS_SQLALCHEMY + +pytestmark = pytest.mark.skipif(not HAS_SQLALCHEMY, reason="SQLAlchemy not available") + +if HAS_SQLALCHEMY: + import sqlalchemy as sa + from bson import Decimal128 + + from pymongosql.sqlalchemy_mongodb.sqlalchemy_dialect import PyMongoSQLDialect + +EXACT = Decimal("123456789012345678901.1234567890") +TINY = Decimal("-0.0000000001") + + +def result_processor(type_): + dialect = PyMongoSQLDialect() + return type_.dialect_impl(dialect).result_processor(dialect, None) or (lambda v: v) + + +class TestOffline: + def test_decimal_parameters_are_encoded_as_decimal128(self): + replaced = SQLHelper.replace_placeholders_generic({"a": "?", "b": {"$in": ["?"]}}, [EXACT, TINY], "qmark") + assert replaced == {"a": Decimal128(EXACT), "b": {"$in": [Decimal128(TINY)]}} + + def test_named_decimal_parameters_are_encoded_as_decimal128(self): + replaced = SQLHelper.replace_placeholders_generic({"a": ":v"}, {"v": EXACT}, "named") + assert replaced == {"a": Decimal128(EXACT)} + + def test_other_parameters_are_unchanged(self): + values = [1, 1.5, "s", None, True] + replaced = SQLHelper.replace_placeholders_generic(["?"] * 5, values, "qmark") + assert replaced == values + + def test_numeric_column_returns_decimal(self): + value = result_processor(sa.Numeric(31, 10))(Decimal128(EXACT)) + assert value == EXACT and type(value) is Decimal + + def test_float_column_returns_float(self): + value = result_processor(sa.Float())(Decimal128("1.25")) + assert value == 1.25 and type(value) is float + + def test_numeric_column_passes_other_values_through(self): + processor = result_processor(sa.Numeric(31, 10)) + assert processor(None) is None + + +class TestLive: + def test_decimal_roundtrip(self, sqlalchemy_engine, conn): + table = sa.Table( + "test_decimal_roundtrip", + sa.MetaData(), + sa.Column("id", sa.Integer), + sa.Column("amount", sa.Numeric(31, 10)), + ) + conn.database.drop_collection(table.name) + try: + with sqlalchemy_engine.begin() as connection: + connection.execute(table.insert(), [{"id": 1, "amount": EXACT}, {"id": 2, "amount": TINY}]) + stored = [d["amount"] for d in conn.database[table.name].find({}, sort=[("id", 1)])] + assert stored == [Decimal128(EXACT), Decimal128(TINY)] + with sqlalchemy_engine.connect() as connection: + amounts = connection.execute( + sa.select(sa.column("amount", sa.Numeric(31, 10))).select_from(sa.table(table.name)) + ).scalars() + values = sorted(amounts) + assert values == [TINY, EXACT] + assert all(type(v) is Decimal for v in values) + finally: + conn.database.drop_collection(table.name) From fccb8dc5a11b289474375f4cb788282eee98bd7f Mon Sep 17 00:00:00 2001 From: Amin Ghadersohi Date: Sat, 26 Sep 2026 03:19:40 +0000 Subject: [PATCH 03/19] fix: resolve table-qualified column references to the column SQLAlchemy qualifies every table-bound column (SELECT users.name FROM users), and the compiler only dropped the qualifier for names starting with an underscore. The translator reads a dotted name as an embedded-document path, so users.name looked for a field name inside a field users: projections silently returned NULL, filters matched nothing, and select(table) raised NoSuchColumnError. Render columns without a table qualifier in the SQLAlchemy compiler, and resolve collection-qualified references in hand-written SQL to the field before building the plan. Other dotted names keep their nested-path meaning. --- pymongosql/sql/builder.py | 35 ++++++++++ .../sqlalchemy_mongodb/sqlalchemy_dialect.py | 17 ++--- tests/test_sqlalchemy_qualified_columns.py | 69 +++++++++++++++++++ 3 files changed, 113 insertions(+), 8 deletions(-) create mode 100644 tests/test_sqlalchemy_qualified_columns.py diff --git a/pymongosql/sql/builder.py b/pymongosql/sql/builder.py index d9dff9d..625bd85 100644 --- a/pymongosql/sql/builder.py +++ b/pymongosql/sql/builder.py @@ -110,9 +110,44 @@ def build_from_parse_result( else: # Default to SELECT/query return ExecutionPlanBuilder._build_query_plan(parse_result) + @staticmethod + def _strip_collection_qualifier(parse_result: "QueryParseResult") -> None: + """Resolve ``collection.field`` references to ``field``. + + SQL qualifies a column with the table it belongs to, while MongoDB reads a + dotted name as an embedded-document path. Without this, a qualified + reference such as ``users.name`` reads the missing path ``users.name`` + and silently returns NULL. As in SQL, the collection name takes + precedence over an embedded document of the same name. + """ + collection = parse_result.collection + if not collection: + return + prefix = f"{collection}." + + def strip(name: Any) -> Any: + if isinstance(name, str) and name.startswith(prefix) and len(name) > len(prefix): + return name[len(prefix) :] + return name + + def strip_filter(value: Any) -> Any: + if isinstance(value, dict): + return {strip(k): strip_filter(v) for k, v in value.items()} + if isinstance(value, list): + return [strip_filter(v) for v in value] + return value + + parse_result.projection = {strip(k): v for k, v in parse_result.projection.items()} + parse_result.column_aliases = {strip(k): v for k, v in parse_result.column_aliases.items()} + parse_result.sort_fields = [{strip(k): v for k, v in spec.items()} for spec in parse_result.sort_fields] + parse_result.filter_conditions = strip_filter(parse_result.filter_conditions) + for func_info in parse_result.aggregate_functions: + func_info["argument"] = strip(func_info["argument"]) + @staticmethod def _build_query_plan(parse_result: "QueryParseResult") -> "QueryExecutionPlan": """Build a query execution plan from SELECT parsing.""" + ExecutionPlanBuilder._strip_collection_qualifier(parse_result) # Auto-generate aggregate pipeline for SQL aggregate functions (COUNT, SUM, etc.) if getattr(parse_result, "aggregate_functions", None): diff --git a/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py b/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py index 63b44ba..5c05587 100644 --- a/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py +++ b/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py @@ -84,14 +84,15 @@ class PyMongoSQLCompiler(compiler.SQLCompiler): Handles SQL compilation specific to MongoDB's query patterns. """ - def visit_column(self, column, **kwargs): - """Handle column references for MongoDB field names.""" - name = column.name - # Handle MongoDB-specific field name patterns - if name.startswith("_"): - # MongoDB system fields like _id - return self.preparer.quote(name) - return super().visit_column(column, **kwargs) + def visit_column(self, column, include_table=True, **kwargs): + """Render column references without a table qualifier. + + A statement reads a single collection, and PyMongoSQL resolves a dotted + reference such as ``users.name`` as the embedded-document path ``name`` + inside a field ``users``. SQLAlchemy qualifies every table-bound column, + so a qualified reference would silently read NULL. + """ + return super().visit_column(column, include_table=False, **kwargs) class PyMongoSQLDDLCompiler(compiler.DDLCompiler): diff --git a/tests/test_sqlalchemy_qualified_columns.py b/tests/test_sqlalchemy_qualified_columns.py new file mode 100644 index 0000000..a5046bc --- /dev/null +++ b/tests/test_sqlalchemy_qualified_columns.py @@ -0,0 +1,69 @@ +# -*- coding: utf-8 -*- +"""Table-qualified column references must read the column, not a nested path.""" + +import pytest + +from pymongosql.sql.parser import SQLParser +from tests.conftest import HAS_SQLALCHEMY + +if HAS_SQLALCHEMY: + import sqlalchemy as sa + + from pymongosql.sqlalchemy_mongodb.sqlalchemy_dialect import PyMongoSQLDialect + + TABLE = sa.Table( + "users", + sa.MetaData(), + sa.Column("_id", sa.String, primary_key=True), + sa.Column("name", sa.String), + sa.Column("age", sa.Integer), + ) + +needs_sqlalchemy = pytest.mark.skipif(not HAS_SQLALCHEMY, reason="SQLAlchemy not available") + + +def plan(sql): + return SQLParser(sql).get_execution_plan() + + +class TestParser: + def test_qualified_projection_filter_and_sort(self): + p = plan("SELECT users.name, users.age FROM users WHERE users.age > 30 ORDER BY users.age DESC") + assert p.projection_stage == {"name": 1, "age": 1} + assert p.filter_stage == {"age": {"$gt": 30}} + assert p.sort_stage == [{"age": -1}] + + def test_qualified_nested_path_keeps_the_path(self): + assert plan("SELECT users.profile.bio FROM users").projection_stage == {"profile.bio": 1} + + def test_other_prefix_is_still_a_nested_path(self): + assert plan("SELECT profile.bio FROM users").projection_stage == {"profile.bio": 1} + + def test_qualified_aggregate_argument(self): + p = plan("SELECT SUM(users.age) AS total FROM users") + assert '"$sum": "$age"' in p.aggregate_pipeline + + +@needs_sqlalchemy +class TestCompiler: + def test_core_select_renders_unqualified_columns(self): + stmt = sa.select(TABLE.c.name, TABLE.c.age).where(TABLE.c._id == "x").order_by(TABLE.c.age) + sql = " ".join(str(stmt.compile(dialect=PyMongoSQLDialect())).split()) + assert sql == "SELECT name, age FROM users WHERE _id = ? ORDER BY age" + + +@needs_sqlalchemy +class TestLive: + def test_core_select_returns_values(self, sqlalchemy_engine, conn): + expected = conn.database["users"].find_one({"_id": "1"}, {"name": 1, "age": 1}) + with sqlalchemy_engine.connect() as connection: + row = connection.execute(sa.select(TABLE.c.name, TABLE.c.age).where(TABLE.c._id == "1")).one() + full = connection.execute(sa.select(TABLE).where(TABLE.c._id == "1")).mappings().one() + assert tuple(row) == (expected["name"], expected["age"]) + assert (full["name"], full["age"]) == (expected["name"], expected["age"]) + + def test_raw_qualified_sql_returns_values(self, conn): + expected = conn.database["users"].find_one({"_id": "1"}, {"name": 1})["name"] + cursor = conn.cursor() + cursor.execute("SELECT users.name FROM users WHERE users._id = '1'") + assert cursor.fetchall() == [(expected,)] From 5a541ec438dfa7fd7ac98d0d06f7793194fc17ce Mon Sep 17 00:00:00 2001 From: Amin Ghadersohi Date: Sat, 26 Sep 2026 03:24:08 +0000 Subject: [PATCH 04/19] fix: return correct rows for GROUP BY, IN, LIKE and keyword aliases Several constructs were translated into MongoDB queries that ran without error but returned wrong rows: - GROUP BY was ignored: the generated $group always used _id: null, so SELECT flag, COUNT(*) ... GROUP BY flag returned one global count and dropped the grouped column. ORDER BY, OFFSET and LIMIT were also dropped from aggregate queries, and ? placeholders in their WHERE clause were never replaced. - IN wrapped every literal in quotes, so IN (1, 2) compared numbers with the strings '1' and '2', and a quoted value containing a comma was split in two. NOT IN and NOT LIKE were read as a field named NOT. - LIKE did not escape regex metacharacters, and a SQL-escaped quote ('O''Brien') was kept doubled. - The SQLAlchemy dialect never quoted PartiQL keywords, so a label such as COUNT(*) AS count failed to parse, and a quoted alias kept its quotes in the result description. Group on the GROUP BY keys and project the SELECT list in order, apply ORDER BY/OFFSET/LIMIT and parameters inside the generated pipeline, keep IN literal types, support NOT IN and NOT LIKE, escape LIKE patterns, unescape doubled quotes, quote PartiQL keywords in the dialect and unquote quoted aliases and ORDER BY keys. Columns that are neither grouped nor aggregated, and HAVING, now raise instead of being dropped. Superset-mode subqueries that aggregate are run as aggregates. --- pymongosql/executor.py | 4 + pymongosql/sql/ast.py | 20 +- pymongosql/sql/builder.py | 93 ++++++++-- pymongosql/sql/handler.py | 70 +++++-- pymongosql/sql/query_builder.py | 3 + pymongosql/sql/query_handler.py | 11 +- .../sqlalchemy_mongodb/sqlalchemy_dialect.py | 74 +++++--- pymongosql/superset_mongodb/executor.py | 5 +- tests/test_sql_grouping_filters_aliases.py | 173 ++++++++++++++++++ 9 files changed, 385 insertions(+), 68 deletions(-) create mode 100644 tests/test_sql_grouping_filters_aliases.py diff --git a/pymongosql/executor.py b/pymongosql/executor.py index 9d27b65..da36382 100644 --- a/pymongosql/executor.py +++ b/pymongosql/executor.py @@ -232,6 +232,10 @@ def _execute_aggregate_plan( _logger.debug(f"Pipeline: {pipeline}") _logger.debug(f"Options: {options}") + # A pipeline generated from SQL carries the WHERE clause's ? placeholders + if parameters and execution_plan.aggregate_parameterized: + pipeline = self._replace_placeholders(pipeline, parameters) + # Get collection and call aggregate() collection = db[execution_plan.collection] diff --git a/pymongosql/sql/ast.py b/pymongosql/sql/ast.py index 9d2372f..7a7a931 100644 --- a/pymongosql/sql/ast.py +++ b/pymongosql/sql/ast.py @@ -4,7 +4,7 @@ from ..error import SqlSyntaxError from .delete_handler import DeleteParseResult -from .handler import BaseHandler, HandlerFactory +from .handler import BaseHandler, ContextUtilsMixin, HandlerFactory from .insert_handler import InsertParseResult from .partiql.PartiQLLexer import PartiQLLexer from .partiql.PartiQLParser import PartiQLParser @@ -275,6 +275,7 @@ def visitOrderByClause(self, ctx: PartiQLParser.OrderByClauseContext) -> Any: if hasattr(ctx, "orderSortSpec") and ctx.orderSortSpec(): for sort_spec in ctx.orderSortSpec(): field_name = sort_spec.expr().getText() if sort_spec.expr() else "_id" + field_name = ContextUtilsMixin.normalize_field_path(field_name) # Check for ASC/DESC (default is ASC = 1) direction = 1 # ASC if hasattr(sort_spec, "DESC") and sort_spec.DESC(): @@ -289,6 +290,23 @@ def visitOrderByClause(self, ctx: PartiQLParser.OrderByClauseContext) -> Any: _logger.warning(f"Error processing ORDER BY clause: {e}") return self.visitChildren(ctx) + def visitGroupClause(self, ctx: PartiQLParser.GroupClauseContext) -> Any: + """Handle GROUP BY keys; they become the _id of a $group stage.""" + keys = [] + for key in ctx.groupKey() or []: + if key.symbolPrimitive() is not None: + self._query_parse_result.unsupported_clauses.append("GROUP BY key alias") + keys.append(ContextUtilsMixin.normalize_field_path(key.exprSelect().getText())) + if ctx.PARTIAL() is not None: + self._query_parse_result.unsupported_clauses.append("GROUP PARTIAL BY") + self._query_parse_result.group_by = keys + return None + + def visitHavingClause(self, ctx: PartiQLParser.HavingClauseContext) -> Any: + """HAVING is not translated; record it so the query fails instead of ignoring it.""" + self._query_parse_result.unsupported_clauses.append("HAVING") + return None + def visitLimitClause(self, ctx: PartiQLParser.LimitClauseContext) -> Any: """Handle LIMIT clause for result limiting""" _logger.debug("Processing LIMIT clause") diff --git a/pymongosql/sql/builder.py b/pymongosql/sql/builder.py index 625bd85..68dd7aa 100644 --- a/pymongosql/sql/builder.py +++ b/pymongosql/sql/builder.py @@ -147,17 +147,28 @@ def strip_filter(value: Any) -> Any: @staticmethod def _build_query_plan(parse_result: "QueryParseResult") -> "QueryExecutionPlan": """Build a query execution plan from SELECT parsing.""" + from ..error import NotSupportedError + ExecutionPlanBuilder._strip_collection_qualifier(parse_result) + if parse_result.unsupported_clauses: + raise NotSupportedError(f"Unsupported SQL clause: {', '.join(parse_result.unsupported_clauses)}") - # Auto-generate aggregate pipeline for SQL aggregate functions (COUNT, SUM, etc.) - if getattr(parse_result, "aggregate_functions", None): + # Auto-generate aggregate pipeline for SQL aggregate functions (COUNT, SUM, etc.) and GROUP BY + if parse_result.aggregate_functions or parse_result.group_by: return ExecutionPlanBuilder._build_sql_aggregate_plan(parse_result) + # ORDER BY may name a column by its SELECT alias; find() sorts on the field + field_for_alias = {alias: name for name, alias in parse_result.column_aliases.items()} + sort_fields = [ + {field_for_alias.get(name, name): direction for name, direction in spec.items()} + for spec in parse_result.sort_fields + ] + builder = BuilderFactory.create_query_builder().collection(parse_result.collection) builder.filter(parse_result.filter_conditions).project(parse_result.projection).column_aliases( parse_result.column_aliases - ).sort(parse_result.sort_fields).limit(parse_result.limit_value).skip(parse_result.offset_value) + ).sort(sort_fields).limit(parse_result.limit_value).skip(parse_result.offset_value) # Set aggregate flags BEFORE building (needed for validation) if hasattr(parse_result, "is_aggregate_query") and parse_result.is_aggregate_query: @@ -171,7 +182,14 @@ def _build_query_plan(parse_result: "QueryParseResult") -> "QueryExecutionPlan": @staticmethod def _build_sql_aggregate_plan(parse_result: "QueryParseResult") -> "QueryExecutionPlan": - """Build an aggregate execution plan from SQL aggregate functions like COUNT(*), SUM(), etc.""" + """Build an aggregate execution plan from SQL aggregate functions and GROUP BY. + + Pipeline: $match (WHERE), $group (GROUP BY keys as _id, one accumulator per + aggregate), $project (SELECT list, in order, under its output names), then + $sort, $skip and $limit on those output names. + """ + from ..error import NotSupportedError + _FUNCTION_TO_ACCUMULATOR = { "COUNT": "$sum", "SUM": "$sum", @@ -188,37 +206,72 @@ def _build_sql_aggregate_plan(parse_result: "QueryParseResult") -> "QueryExecuti if parse_result.filter_conditions: pipeline.append({"$match": parse_result.filter_conditions}) - # Build $group stage from aggregate functions - group_stage = {"_id": None} - for func_info in parse_result.aggregate_functions: - alias = func_info["alias"] + group_keys = {name: f"g{i}" for i, name in enumerate(parse_result.group_by)} + group_stage = {"_id": {key: f"${name}" for name, key in group_keys.items()} if group_keys else None} + accumulator_keys = [] + for i, func_info in enumerate(parse_result.aggregate_functions): func_name = func_info["function"] arg = func_info["argument"] accumulator = _FUNCTION_TO_ACCUMULATOR[func_name] - - if func_name == "COUNT": - group_stage[alias] = {accumulator: 1} + # $group output names may not contain "." or start with "$", nor repeat + key = func_info["alias"] + if "." in key or key.startswith("$") or key == "_id" or key in group_stage: + key = f"__agg{i}" + accumulator_keys.append(key) + + if func_name == "COUNT" and arg == "*": + group_stage[key] = {accumulator: 1} + elif func_name == "COUNT": + # COUNT(field) counts documents where the field is present and not null + group_stage[key] = {"$sum": {"$cond": [{"$gt": [f"${arg}", None]}, 1, 0]}} else: - group_stage[alias] = {accumulator: f"${arg}"} + group_stage[key] = {accumulator: f"${arg}"} pipeline.append({"$group": group_stage}) - # Add $project to exclude _id + # Map every SELECT item, in order, to its output name and source project_stage = {"_id": 0} - for func_info in parse_result.aggregate_functions: - project_stage[func_info["alias"]] = 1 + outputs = [] + output_for = {} # names ORDER BY may use -> output name + for item in parse_result.select_items: + if "aggregate" in item: + func_info = parse_result.aggregate_functions[item["aggregate"]] + output, key = func_info["alias"], accumulator_keys[item["aggregate"]] + source = 1 if key == output else f"${key}" + output_for[func_info["expression"].upper()] = output + else: + name = item["field"] + if name not in group_keys: + raise NotSupportedError(f"Column '{name}' must appear in GROUP BY or in an aggregate function") + output, source = item["alias"] or name, f"$_id.{group_keys[name]}" + output_for[name] = output + output_for[output] = output + project_stage[output] = source + outputs.append(output) pipeline.append({"$project": project_stage}) + sort_stage = {} + for spec in parse_result.sort_fields: + for name, direction in spec.items(): + output = output_for.get(name, output_for.get(name.upper())) + if output is None: + raise NotSupportedError(f"ORDER BY '{name}' must name a selected column or its alias") + sort_stage[output] = direction + if sort_stage: + pipeline.append({"$sort": sort_stage}) + if parse_result.offset_value: + pipeline.append({"$skip": parse_result.offset_value}) + if parse_result.limit_value is not None: + pipeline.append({"$limit": parse_result.limit_value}) + # Configure the execution plan as an aggregate query builder._execution_plan.is_aggregate_query = True + builder._execution_plan.aggregate_parameterized = True builder._execution_plan.aggregate_pipeline = json.dumps(pipeline) builder._execution_plan.aggregate_options = json.dumps({}) - # Set projection for ResultSet description - agg_projection = {} - for func_info in parse_result.aggregate_functions: - agg_projection[func_info["alias"]] = 1 - builder._execution_plan.projection_stage = agg_projection + # Set projection for ResultSet description, in SELECT order + builder._execution_plan.projection_stage = {name: 1 for name in outputs} plan = builder.build() return plan diff --git a/pymongosql/sql/handler.py b/pymongosql/sql/handler.py index e01370d..fce8c93 100644 --- a/pymongosql/sql/handler.py +++ b/pymongosql/sql/handler.py @@ -52,6 +52,13 @@ def has_children(ctx: Any) -> bool: """Check if context has children""" return hasattr(ctx, "children") and bool(ctx.children) + @staticmethod + def unquote_identifier(name: Optional[str]) -> Optional[str]: + """Strip the double quotes of a quoted SQL identifier (``"count"`` -> ``count``).""" + if isinstance(name, str) and len(name) >= 2 and name.startswith('"') and name.endswith('"'): + return name[1:-1].replace('""', '"') + return name + @staticmethod def normalize_field_path(path: str) -> str: """Normalize jmspath/bracket notation to MongoDB dot notation. @@ -144,10 +151,10 @@ def _parse_value(self, value_text: str) -> Any: # Remove parentheses from values value_text = value_text.strip("()") - # Remove quotes from string values - if (value_text.startswith("'") and value_text.endswith("'")) or ( - value_text.startswith('"') and value_text.endswith('"') - ): + # Remove quotes from string values; SQL escapes a quote by doubling it + if len(value_text) >= 2 and value_text.startswith("'") and value_text.endswith("'"): + return value_text[1:-1].replace("''", "'") + if len(value_text) >= 2 and value_text.startswith('"') and value_text.endswith('"'): return value_text[1:-1] # Try to parse as number @@ -251,18 +258,20 @@ def _build_mongo_filter(self, field_name: str, operator: str, value: Any) -> Dic return {field_name: value} # Handle special operators - if operator == "IN": - return {field_name: {"$in": value if isinstance(value, list) else [value]}} - elif operator == "LIKE": + if operator in ("IN", "NOT IN"): + values = value if isinstance(value, list) else [value] + return {field_name: {"$in" if operator == "IN" else "$nin": values}} + elif operator in ("LIKE", "NOT LIKE"): # Convert SQL LIKE pattern to regex if isinstance(value, str): - # Replace % with .* and _ with . for regex - regex_pattern = value.replace("%", ".*").replace("_", ".") + regex_pattern = self._like_to_regex(value) # Add anchors based on pattern if not regex_pattern.startswith(".*"): regex_pattern = "^" + regex_pattern if not regex_pattern.endswith(".*"): regex_pattern = regex_pattern + "$" + if operator == "NOT LIKE": + return {field_name: {"$not": {"$regex": regex_pattern}}} return {field_name: {"$regex": regex_pattern}} return {field_name: value} elif operator == "BETWEEN": @@ -287,6 +296,26 @@ def _build_mongo_filter(self, field_name: str, operator: str, value: Any) -> Dic _logger.warning(f"Unknown operator '{operator}', falling back to equality") return {field_name: value} + @staticmethod + def _like_to_regex(pattern: str) -> str: + """Translate a LIKE pattern, escaping every other regex metacharacter.""" + return "".join(".*" if c == "%" else "." if c == "_" else re.escape(c) for c in pattern) + + def _negated_keyword(self, ctx: Any, text: str, keyword: str) -> bool: + """Whether ``keyword`` (IN( or LIKE) is preceded by NOT. + + getText() drops whitespace, so ``a NOT IN (1)`` reads ``aNOTIN(1)``. Use the + parse tree when there is one; otherwise require the upper-case NOT that + generated SQL uses, so a field such as ``cannot`` is not misread. + """ + not_method = getattr(ctx, "NOT", None) + if callable(not_method) and type(ctx).__name__.startswith("Predicate"): + try: + return not_method() is not None + except Exception: + pass + return f"NOT{keyword}" in text + def _is_comparison_context(self, ctx: Any) -> bool: """Check if context is a comparison based on structure""" context_name = self.get_context_type_name(ctx).lower() @@ -334,6 +363,8 @@ def _extract_field_name(self, ctx: Any) -> str: for keyword in sql_keywords: if keyword in text_upper: idx = text_upper.index(keyword) + if keyword in ("IN(", "LIKE") and self._negated_keyword(ctx, text, keyword): + idx = text.index(f"NOT{keyword}") candidate = text[:idx].strip() return self.normalize_field_path(candidate) @@ -374,6 +405,8 @@ def _extract_operator(self, ctx: Any) -> str: for construct, operator in sql_constructs.items(): if construct in text_upper: + if construct in ("IN(", "LIKE") and self._negated_keyword(ctx, text, construct): + return f"NOT {operator}" return operator # Look for comparison operators @@ -519,13 +552,18 @@ def _extract_in_values(self, text: str) -> List[Any]: end = text.rfind(")") if end > start >= 0: - values_text = text[start:end] - values = [] - for val in values_text.split(","): - cleaned_val = val.strip().strip("'\"") - if cleaned_val: # Skip empty values - values.append(self._parse_value(f"'{cleaned_val}'")) - return values + # Split on commas outside quoted strings; keep each literal's type + values, current, in_quote = [], "", False + for char in text[start:end]: + if char == "'": + in_quote = not in_quote + if char == "," and not in_quote: + values.append(current) + current = "" + else: + current += char + values.append(current) + return [self._extract_value_or_function(v) for v in values if v.strip()] return [] def _extract_like_pattern(self, text: str) -> str: diff --git a/pymongosql/sql/query_builder.py b/pymongosql/sql/query_builder.py index fb3b7cf..8e45b82 100644 --- a/pymongosql/sql/query_builder.py +++ b/pymongosql/sql/query_builder.py @@ -22,6 +22,8 @@ class QueryExecutionPlan(ExecutionPlan): aggregate_pipeline: Optional[str] = None # JSON string representation of pipeline aggregate_options: Optional[str] = None # JSON string representation of options is_aggregate_query: bool = False # Flag indicating this is an aggregate() call + # True when the pipeline was generated from SQL and may hold ? placeholders + aggregate_parameterized: bool = False def to_dict(self) -> Dict[str, Any]: """Convert query plan to dictionary representation""" @@ -76,6 +78,7 @@ def copy(self) -> "QueryExecutionPlan": aggregate_pipeline=self.aggregate_pipeline, aggregate_options=self.aggregate_options, is_aggregate_query=self.is_aggregate_query, + aggregate_parameterized=self.aggregate_parameterized, ) diff --git a/pymongosql/sql/query_handler.py b/pymongosql/sql/query_handler.py index 3a09db8..a9bdd50 100644 --- a/pymongosql/sql/query_handler.py +++ b/pymongosql/sql/query_handler.py @@ -34,6 +34,12 @@ class QueryParseResult: # SQL aggregate functions detected in SELECT (COUNT, SUM, AVG, MIN, MAX) aggregate_functions: List[Dict[str, Any]] = field(default_factory=list) + # SELECT items in order: {"field": name, "alias": alias} or {"aggregate": index} + select_items: List[Dict[str, Any]] = field(default_factory=list) + # GROUP BY field paths + group_by: List[str] = field(default_factory=list) + # Clauses that are parsed but cannot be translated faithfully + unsupported_clauses: List[str] = field(default_factory=list) # Subquery info (for wrapped subqueries, e.g., Superset outering) subquery_plan: Optional[Any] = None @@ -137,15 +143,18 @@ def handle_visitor(self, ctx: PartiQLParser.SelectItemsContext, parse_result: "Q if agg_match: func_name = agg_match.group(1).upper() func_arg = agg_match.group(2) + parse_result.select_items.append({"aggregate": len(parse_result.aggregate_functions)}) parse_result.aggregate_functions.append( { "function": func_name, "argument": func_arg, "alias": alias or field_name, + "expression": field_name, } ) continue + parse_result.select_items.append({"field": field_name, "alias": alias}) # Use MongoDB standard projection format: {field: 1} to include field projection[field_name] = 1 # Store alias if present @@ -182,7 +191,7 @@ def _extract_field_and_alias(self, item) -> Tuple[str, Optional[str]]: # Pattern: expr symbolPrimitive (without AS) alias = item.children[1].getText() - return field_name, alias + return field_name, self.unquote_identifier(alias) class FromHandler(BaseHandler): diff --git a/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py b/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py index 5c05587..d7ae983 100644 --- a/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py +++ b/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py @@ -32,6 +32,19 @@ from sqlalchemy.engine.interfaces import Dialect +def _partiql_keywords() -> set: + """Lower-case PartiQL keywords, e.g. ``count`` or ``value``. + + The grammar rejects a keyword used as a bare identifier (``COUNT(*) AS count`` + is a syntax error), so SQLAlchemy must quote these names. + """ + import re + + from pymongosql.sql.partiql.PartiQLLexer import PartiQLLexer + + return {name.strip("'").lower() for name in PartiQLLexer.literalNames if re.fullmatch(r"'[A-Za-z_]+'", name or "")} + + class PyMongoSQLIdentifierPreparer(compiler.IdentifierPreparer): """MongoDB-specific identifier preparer. @@ -39,35 +52,38 @@ class PyMongoSQLIdentifierPreparer(compiler.IdentifierPreparer): from SQL databases. """ - reserved_words = set( - [ - # MongoDB reserved words and operators - "$eq", - "$ne", - "$gt", - "$gte", - "$lt", - "$lte", - "$in", - "$nin", - "$and", - "$or", - "$not", - "$nor", - "$exists", - "$type", - "$mod", - "$regex", - "$text", - "$where", - "$all", - "$elemMatch", - "$size", - "$bitsAllClear", - "$bitsAllSet", - "$bitsAnyClear", - "$bitsAnySet", - ] + reserved_words = ( + set( + [ + # MongoDB reserved words and operators + "$eq", + "$ne", + "$gt", + "$gte", + "$lt", + "$lte", + "$in", + "$nin", + "$and", + "$or", + "$not", + "$nor", + "$exists", + "$type", + "$mod", + "$regex", + "$text", + "$where", + "$all", + "$elemMatch", + "$size", + "$bitsAllClear", + "$bitsAllSet", + "$bitsAnyClear", + "$bitsAnySet", + ] + ) + | _partiql_keywords() ) def __init__(self, dialect: Dialect, **kwargs: Any) -> None: diff --git a/pymongosql/superset_mongodb/executor.py b/pymongosql/superset_mongodb/executor.py index f7090bc..22cf24f 100644 --- a/pymongosql/superset_mongodb/executor.py +++ b/pymongosql/superset_mongodb/executor.py @@ -63,7 +63,10 @@ def execute( _logger.debug(f"Stage 1: Executing MongoDB subquery: {mongo_query}") mongo_execution_plan = self._parse_sql(mongo_query) - mongo_result = self._execute_find_plan(mongo_execution_plan, connection) + if mongo_execution_plan.is_aggregate_query: + mongo_result = self._execute_aggregate_plan(mongo_execution_plan, connection) + else: + mongo_result = self._execute_find_plan(mongo_execution_plan, connection) # Extract result set from MongoDB mongo_result_set = ResultSet( diff --git a/tests/test_sql_grouping_filters_aliases.py b/tests/test_sql_grouping_filters_aliases.py new file mode 100644 index 0000000..dfe13aa --- /dev/null +++ b/tests/test_sql_grouping_filters_aliases.py @@ -0,0 +1,173 @@ +# -*- coding: utf-8 -*- +"""GROUP BY, IN/NOT IN, LIKE, quoted literals and keyword aliases must return correct rows.""" + +import json + +import pytest + +from pymongosql.error import Error +from pymongosql.sql.parser import SQLParser +from tests.conftest import HAS_SQLALCHEMY, make_superset_conn + +COLLECTION = "test_grouping_filters" +DOCS = [ + {"_id": 1, "flag": True, "dept": "a", "amount": 10, "name": "O'Brien"}, + {"_id": 2, "flag": False, "dept": "a", "amount": 20, "name": "x.y"}, + {"_id": 3, "flag": True, "dept": "b", "amount": 40, "name": "xzy"}, + {"_id": 4, "flag": True, "dept": "b", "amount": None, "name": "plain"}, +] + + +def plan(sql): + return SQLParser(sql).get_execution_plan() + + +def pipeline(sql): + return json.loads(plan(sql).aggregate_pipeline) + + +class TestPlans: + def test_group_by_groups_on_the_key(self): + stages = pipeline("SELECT flag, COUNT(*) AS n FROM t GROUP BY flag") + assert stages[0]["$group"]["_id"] == {"g0": "$flag"} + assert stages[1]["$project"] == {"_id": 0, "flag": "$_id.g0", "n": 1} + + def test_aggregate_query_keeps_order_by_skip_and_limit(self): + stages = pipeline("SELECT dept, SUM(amount) AS total FROM t GROUP BY dept ORDER BY total DESC LIMIT 1 OFFSET 1") + assert stages[-3:] == [{"$sort": {"total": -1}}, {"$skip": 1}, {"$limit": 1}] + + def test_order_by_aggregate_expression_uses_its_output(self): + stages = pipeline("SELECT dept, SUM(amount) AS total FROM t GROUP BY dept ORDER BY SUM(amount)") + assert stages[-1] == {"$sort": {"total": 1}} + + def test_quoted_keyword_alias_is_unquoted(self): + p = plan('SELECT flag, COUNT(*) AS "count" FROM t GROUP BY flag ORDER BY "count" DESC') + assert list(p.projection_stage) == ["flag", "count"] + assert json.loads(p.aggregate_pipeline)[-1] == {"$sort": {"count": -1}} + + def test_find_order_by_alias_sorts_on_the_field(self): + p = plan('SELECT amount AS "value" FROM t ORDER BY "value" DESC') + assert p.sort_stage == [{"amount": -1}] + assert p.column_aliases == {"amount": "value"} + + def test_ungrouped_column_is_rejected(self): + with pytest.raises(Exception, match="must appear in GROUP BY"): + plan("SELECT flag, COUNT(*) FROM t") + + def test_having_is_rejected_not_ignored(self): + with pytest.raises(Exception, match="HAVING"): + plan("SELECT flag, COUNT(*) AS n FROM t GROUP BY flag HAVING COUNT(*) > 1") + + def test_in_keeps_literal_types_and_quoted_commas(self): + p = plan("SELECT _id FROM t WHERE _id IN (1, 2.5, 'a,b', 'it''s', TRUE, ?)") + assert p.filter_stage == {"_id": {"$in": [1, 2.5, "a,b", "it's", True, "?"]}} + + def test_not_in(self): + assert plan("SELECT _id FROM t WHERE _id NOT IN (1, 2)").filter_stage == {"_id": {"$nin": [1, 2]}} + + def test_field_ending_in_not_is_not_negated(self): + assert plan("SELECT _id FROM t WHERE cannot IN (1)").filter_stage == {"cannot": {"$in": [1]}} + + def test_not_like_and_regex_metacharacters(self): + p = plan("SELECT _id FROM t WHERE name NOT LIKE 'x.%'") + assert p.filter_stage == {"name": {"$not": {"$regex": "^x\\..*"}}} + + def test_escaped_quote_in_string_literal(self): + assert plan("SELECT _id FROM t WHERE name = 'O''Brien'").filter_stage == {"name": "O'Brien"} + + +@pytest.mark.skipif(not HAS_SQLALCHEMY, reason="SQLAlchemy not available") +class TestDialectQuoting: + def test_keyword_alias_and_column_are_quoted(self): + import sqlalchemy as sa + + from pymongosql.sqlalchemy_mongodb.sqlalchemy_dialect import PyMongoSQLDialect + + t = sa.table("t", sa.column("flag"), sa.column("value")) + count = sa.func.count().label("count") + stmt = sa.select(sa.column("flag"), t.c.value, count).select_from(t).group_by(sa.column("flag")).order_by(count) + sql = " ".join(str(stmt.compile(dialect=PyMongoSQLDialect())).split()) + assert sql == 'SELECT flag, "value", count(*) AS "count" FROM t GROUP BY flag ORDER BY "count"' + + +@pytest.fixture +def grouping_collection(conn): + conn.database.drop_collection(COLLECTION) + conn.database[COLLECTION].insert_many(DOCS) + yield conn + conn.database.drop_collection(COLLECTION) + + +def rows(conn, sql, params=None): + cursor = conn.cursor() + cursor.execute(sql, params) if params is not None else cursor.execute(sql) + return [tuple(r) for r in cursor.fetchall()] + + +class TestLive: + def test_group_by_returns_one_row_per_group(self, grouping_collection): + got = rows(grouping_collection, f"SELECT flag, COUNT(*) AS n FROM {COLLECTION} GROUP BY flag ORDER BY flag") + assert got == [(False, 1), (True, 3)] + + def test_group_by_with_sum_where_parameter_order_and_limit(self, grouping_collection): + sql = ( + f"SELECT dept, SUM(amount) AS total, COUNT(amount) AS counted FROM {COLLECTION} " + "WHERE _id > ? GROUP BY dept ORDER BY total DESC LIMIT 1" + ) + assert rows(grouping_collection, sql, [0]) == [("b", 40, 1)] + + def test_quoted_count_alias(self, grouping_collection): + sql = f'SELECT flag, COUNT(*) AS "count" FROM {COLLECTION} GROUP BY flag ORDER BY "count" DESC' + cursor = grouping_collection.cursor() + cursor.execute(sql) + assert [d[0] for d in cursor.description] == ["flag", "count"] + assert [tuple(r) for r in cursor.fetchall()] == [(True, 3), (False, 1)] + + def test_in_and_not_in(self, grouping_collection): + assert rows(grouping_collection, f"SELECT _id FROM {COLLECTION} WHERE _id IN (1, 3) ORDER BY _id") == [ + (1,), + (3,), + ] + assert rows(grouping_collection, f"SELECT _id FROM {COLLECTION} WHERE _id IN (?, ?) ORDER BY _id", [2, 4]) == [ + (2,), + (4,), + ] + assert rows(grouping_collection, f"SELECT _id FROM {COLLECTION} WHERE _id NOT IN (1, 3) ORDER BY _id") == [ + (2,), + (4,), + ] + + def test_escaped_quote_and_like(self, grouping_collection): + assert rows(grouping_collection, f"SELECT _id FROM {COLLECTION} WHERE name = 'O''Brien'") == [(1,)] + assert rows(grouping_collection, f"SELECT _id FROM {COLLECTION} WHERE name LIKE 'x.%'") == [(2,)] + assert rows(grouping_collection, f"SELECT _id FROM {COLLECTION} WHERE name NOT LIKE 'x%' ORDER BY _id") == [ + (1,), + (4,), + ] + + def test_superset_mode_physical_table_chart_query(self, grouping_collection): + conn = make_superset_conn() + try: + sql = ( + f'SELECT flag AS flag, COUNT(*) AS "count" FROM {COLLECTION} ' + 'GROUP BY flag ORDER BY "count" DESC LIMIT 100' + ) + assert rows(conn, sql) == [(True, 3), (False, 1)] + finally: + conn.close() + + def test_unsupported_having_raises(self, grouping_collection): + with pytest.raises(Error): + rows(grouping_collection, f"SELECT flag, COUNT(*) AS n FROM {COLLECTION} GROUP BY flag HAVING n > 1") + + +@pytest.mark.skipif(not HAS_SQLALCHEMY, reason="SQLAlchemy not available") +class TestLiveSQLAlchemy: + def test_core_group_by_count_label(self, sqlalchemy_engine, grouping_collection): + import sqlalchemy as sa + + t = sa.table(COLLECTION, sa.column("flag"), sa.column("_id")) + count = sa.func.count().label("count") + stmt = sa.select(t.c.flag, count).where(t.c._id.in_([1, 2, 3])).group_by(t.c.flag).order_by(count.desc()) + with sqlalchemy_engine.connect() as connection: + assert [tuple(r) for r in connection.execute(stmt)] == [(True, 2), (False, 1)] From b14ae06bc81bafdce24aa0e7cd96339f9d63409e Mon Sep 17 00:00:00 2001 From: Amin Ghadersohi Date: Sat, 26 Sep 2026 03:30:35 +0000 Subject: [PATCH 05/19] fix(sqlalchemy): read and bind Uuid columns as BSON UUIDs The dialect did not declare native UUID support, so SQLAlchemy 2's Uuid type used its string-based processors. PyMongo returns uuid.UUID (or a subtype-4 Binary), and reading a Uuid column failed with "'UUID' object has no attribute 'replace'". Declare native UUID support and map Uuid to a type that binds standard subtype-4 binaries and reads uuid.UUID, subtype-4 Binary or string values. Legacy subtype 3 is returned unchanged because its byte order depends on the driver that wrote it. --- .../sqlalchemy_mongodb/sqlalchemy_dialect.py | 40 +++++++++++++ tests/test_sqlalchemy_uuid.py | 60 +++++++++++++++++++ 2 files changed, 100 insertions(+) create mode 100644 tests/test_sqlalchemy_uuid.py diff --git a/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py b/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py index d7ae983..0c25d8c 100644 --- a/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py +++ b/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py @@ -1,5 +1,6 @@ # -*- coding: utf-8 -*- import logging +import uuid from typing import Any, Dict, List, Optional, Tuple, Type from urllib.parse import quote_plus @@ -194,6 +195,42 @@ def result_processor(self, dialect, coltype): return _decode_decimal128(super().result_processor(dialect, coltype)) +class _MongoUuid(getattr(sqltypes, "Uuid", sqltypes.TypeEngine)): # Uuid is new in SQLAlchemy 2.0 + """Uuid stored as BSON binary subtype 4 (the standard UUID representation). + + PyMongo returns ``uuid.UUID`` under ``uuidRepresentation=standard`` and a + subtype-4 ``Binary`` otherwise; the generic non-native Uuid processors expect a + hex string and fail on both. Legacy subtype 3 is left as ``Binary``: its byte + order depends on the driver that wrote it. + """ + + def bind_processor(self, dialect): + from bson.binary import Binary + + def process(value): + if value is None: + return None + if not isinstance(value, uuid.UUID): + value = uuid.UUID(str(value)) + return Binary.from_uuid(value) + + return process + + def result_processor(self, dialect, coltype): + from bson.binary import UUID_SUBTYPE, Binary + + def process(value): + if isinstance(value, Binary) and value.subtype == UUID_SUBTYPE: + value = value.as_uuid() + elif isinstance(value, str): + value = uuid.UUID(value) + if isinstance(value, uuid.UUID) and not self.as_uuid: + return str(value) + return value + + return process + + class PyMongoSQLDialect(default.DefaultDialect): """SQLAlchemy dialect for PyMongoSQL. @@ -218,7 +255,10 @@ class PyMongoSQLDialect(default.DefaultDialect): supports_multivalues_insert = True supports_native_decimal = True # BSON Decimal128 # PyMongo returns Decimal128, not decimal.Decimal; convert on the way out. + supports_native_uuid = True # BSON binary subtype 4 colspecs = {sqltypes.Numeric: _MongoNumeric, sqltypes.Float: _MongoFloat} + if hasattr(sqltypes, "Uuid"): + colspecs[sqltypes.Uuid] = _MongoUuid supports_native_boolean = True # BSON Boolean supports_sequences = False # No sequences in MongoDB supports_native_enum = False # No native enums diff --git a/tests/test_sqlalchemy_uuid.py b/tests/test_sqlalchemy_uuid.py new file mode 100644 index 0000000..6f9b7e7 --- /dev/null +++ b/tests/test_sqlalchemy_uuid.py @@ -0,0 +1,60 @@ +# -*- coding: utf-8 -*- +"""SQLAlchemy 2 Uuid columns must round-trip BSON UUIDs.""" + +import uuid + +import pytest + +from tests.conftest import HAS_SQLALCHEMY + +sa = pytest.importorskip("sqlalchemy") +pytestmark = pytest.mark.skipif(not (HAS_SQLALCHEMY and hasattr(sa, "Uuid")), reason="needs SQLAlchemy 2 Uuid") + +from bson.binary import Binary, UuidRepresentation # noqa: E402 + +from pymongosql.sqlalchemy_mongodb.sqlalchemy_dialect import PyMongoSQLDialect # noqa: E402 + +VALUE = uuid.UUID("00000000-0000-4000-8000-000000000001") + + +def processors(type_): + dialect = PyMongoSQLDialect() + impl = type_.dialect_impl(dialect) + return impl.bind_processor(dialect), impl.result_processor(dialect, None) + + +@pytest.mark.parametrize("stored", [VALUE, Binary.from_uuid(VALUE), str(VALUE)]) +def test_uuid_column_reads_uuid(stored): + _, result = processors(sa.Uuid()) + assert result(stored) == VALUE + + +def test_uuid_column_as_string(): + _, result = processors(sa.Uuid(as_uuid=False)) + assert result(VALUE) == str(VALUE) + + +def test_uuid_bind_is_standard_binary(): + bind, _ = processors(sa.Uuid()) + assert bind(VALUE) == Binary.from_uuid(VALUE) + assert bind(str(VALUE)) == Binary.from_uuid(VALUE) + assert bind(None) is None + + +def test_legacy_subtype_3_is_not_guessed(): + legacy = Binary.from_uuid(VALUE, UuidRepresentation.PYTHON_LEGACY) + _, result = processors(sa.Uuid()) + assert result(legacy) == legacy + + +def test_live_roundtrip(sqlalchemy_engine, conn): + table = sa.Table("test_uuid_roundtrip", sa.MetaData(), sa.Column("id", sa.Integer), sa.Column("u", sa.Uuid)) + conn.database.drop_collection(table.name) + try: + with sqlalchemy_engine.begin() as connection: + connection.execute(table.insert(), [{"id": 1, "u": VALUE}]) + assert conn.database[table.name].find_one()["u"] == Binary.from_uuid(VALUE) + with sqlalchemy_engine.connect() as connection: + assert connection.execute(sa.select(table.c.u)).scalar_one() == VALUE + finally: + conn.database.drop_collection(table.name) From 8ec2dbf92e91101a85ee14779e0b0f458acffeb5 Mon Sep 17 00:00:00 2001 From: Amin Ghadersohi Date: Sat, 26 Sep 2026 03:53:07 +0000 Subject: [PATCH 06/19] fix(sqlalchemy): return int for BSON int64 in Integer columns The command responses the DBAPI decodes carry 64-bit integers as bson.Int64, so Integer and BigInteger columns returned that subclass rather than the int SQLAlchemy promises. Convert it in the Integer result processor. --- .../sqlalchemy_mongodb/sqlalchemy_dialect.py | 14 +++++++++++++- tests/test_sqlalchemy_numeric.py | 19 +++++++++++++++++++ 2 files changed, 32 insertions(+), 1 deletion(-) diff --git a/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py b/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py index 0c25d8c..b6a271e 100644 --- a/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py +++ b/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py @@ -188,6 +188,18 @@ def result_processor(self, dialect, coltype): return _decode_decimal128(super().result_processor(dialect, coltype)) +class _MongoInteger(sqltypes.Integer): + """Integer that returns ``int`` for BSON int64 values (``bson.Int64``).""" + + def result_processor(self, dialect, coltype): + from bson import Int64 + + def process(value): + return int(value) if isinstance(value, Int64) else value + + return process + + class _MongoFloat(sqltypes.Float): """Float that returns ``float`` (or Decimal) for Decimal128 values.""" @@ -256,7 +268,7 @@ class PyMongoSQLDialect(default.DefaultDialect): supports_native_decimal = True # BSON Decimal128 # PyMongo returns Decimal128, not decimal.Decimal; convert on the way out. supports_native_uuid = True # BSON binary subtype 4 - colspecs = {sqltypes.Numeric: _MongoNumeric, sqltypes.Float: _MongoFloat} + colspecs = {sqltypes.Numeric: _MongoNumeric, sqltypes.Float: _MongoFloat, sqltypes.Integer: _MongoInteger} if hasattr(sqltypes, "Uuid"): colspecs[sqltypes.Uuid] = _MongoUuid supports_native_boolean = True # BSON Boolean diff --git a/tests/test_sqlalchemy_numeric.py b/tests/test_sqlalchemy_numeric.py index e294a7d..509d779 100644 --- a/tests/test_sqlalchemy_numeric.py +++ b/tests/test_sqlalchemy_numeric.py @@ -47,6 +47,13 @@ def test_float_column_returns_float(self): value = result_processor(sa.Float())(Decimal128("1.25")) assert value == 1.25 and type(value) is float + def test_integer_columns_return_int_for_int64(self): + from bson import Int64 + + for type_ in (sa.Integer(), sa.BigInteger()): + value = result_processor(type_)(Int64(2**62)) + assert value == 2**62 and type(value) is int + def test_numeric_column_passes_other_values_through(self): processor = result_processor(sa.Numeric(31, 10)) assert processor(None) is None @@ -75,3 +82,15 @@ def test_decimal_roundtrip(self, sqlalchemy_engine, conn): assert all(type(v) is Decimal for v in values) finally: conn.database.drop_collection(table.name) + + def test_int64_roundtrip(self, sqlalchemy_engine, conn): + table = sa.Table("test_int64_roundtrip", sa.MetaData(), sa.Column("big", sa.BigInteger)) + conn.database.drop_collection(table.name) + try: + conn.database[table.name].insert_many([{"big": 2**63 - 1}, {"big": -(2**63)}]) + with sqlalchemy_engine.connect() as connection: + values = sorted(connection.execute(sa.select(table.c.big)).scalars()) + assert values == [-(2**63), 2**63 - 1] + assert all(type(v) is int for v in values) + finally: + conn.database.drop_collection(table.name) From 2570e3382a328a37915c1662c8647fedf5cda363 Mon Sep 17 00:00:00 2001 From: Amin Ghadersohi Date: Sat, 26 Sep 2026 03:55:06 +0000 Subject: [PATCH 07/19] fix: keep -- inside quoted literals when stripping comments The preprocessor cut every line at the first "--", including one inside a string literal or quoted identifier, so WHERE name = 'a -- b' failed to parse. Strip a line comment only when it starts outside quotes. --- pymongosql/sql/parser.py | 22 +++++++++++++++------- tests/test_sql_grouping_filters_aliases.py | 4 ++++ 2 files changed, 19 insertions(+), 7 deletions(-) diff --git a/pymongosql/sql/parser.py b/pymongosql/sql/parser.py index 7dc39d6..b3d95e1 100644 --- a/pymongosql/sql/parser.py +++ b/pymongosql/sql/parser.py @@ -88,17 +88,25 @@ def _preprocess(self) -> None: # Remove extra whitespace and normalize sql = self._original_sql.strip() - # Remove comments (basic implementation) - lines = [] - for line in sql.split("\n"): - # Remove single-line comments - if "--" in line: - line = line[: line.index("--")] - lines.append(line) + # Remove single-line comments; "--" inside a quoted literal or identifier is data + lines = [self._strip_line_comment(line) for line in sql.split("\n")] self._preprocessed_sql = " ".join(lines).strip() _logger.debug(f"Preprocessed SQL: {self._preprocessed_sql}") + @staticmethod + def _strip_line_comment(line: str) -> str: + quote = None + for i, char in enumerate(line): + if quote: + if char == quote: + quote = None # a doubled quote closes and reopens: same result + elif char in ("'", '"'): + quote = char + elif line.startswith("--", i): + return line[:i] + return line + def _generate_ast(self) -> None: """Generate Abstract Syntax Tree from SQL""" try: diff --git a/tests/test_sql_grouping_filters_aliases.py b/tests/test_sql_grouping_filters_aliases.py index dfe13aa..65fffec 100644 --- a/tests/test_sql_grouping_filters_aliases.py +++ b/tests/test_sql_grouping_filters_aliases.py @@ -72,6 +72,10 @@ def test_not_like_and_regex_metacharacters(self): p = plan("SELECT _id FROM t WHERE name NOT LIKE 'x.%'") assert p.filter_stage == {"name": {"$not": {"$regex": "^x\\..*"}}} + def test_double_dash_inside_literal_is_not_a_comment(self): + p = plan("SELECT _id FROM t WHERE name = 'a -- b' AND n = 1 -- trailing comment") + assert p.filter_stage == {"$and": [{"name": "a -- b"}, {"n": 1}]} + def test_escaped_quote_in_string_literal(self): assert plan("SELECT _id FROM t WHERE name = 'O''Brien'").filter_stage == {"name": "O'Brien"} From 4abf1301cfd1da7f65cd34b5524040d3d5bd3d59 Mon Sep 17 00:00:00 2001 From: Amin Ghadersohi Date: Sat, 26 Sep 2026 06:00:28 +0000 Subject: [PATCH 08/19] fix: translate NOT and NULL comparisons with SQL three-valued logic WHERE clauses were translated by splitting getText() output, which has no whitespace. NOT a = 1 became a filter on a field named "NOTa" (no rows), NOT (a = 1 OR b = 2) produced no filter at all (every row), and an operand that could not be translated was silently dropped from an AND, or the whole clause fell back to a $text search. <>, NOT IN and NOT LIKE also matched documents where the field was NULL or missing, which SQL never returns. Translate WHERE over the parse tree. Each predicate yields the filter of documents for which it is TRUE and the filter for which it is FALSE (a NULL or missing operand is in neither). NOT swaps them and AND/OR combine them by De Morgan's laws, so NOT a = 1 excludes NULLs as in SQL. A bare boolean field (WHERE flag / WHERE NOT flag) is supported. A predicate on anything but a field path, or a LIKE with a bound pattern, now raises instead of matching the wrong rows; the SQLAlchemy dialect renders LIKE patterns inline so Core like() keeps working. "= NULL" keeps its existing IS NULL meaning. DELETE and UPDATE use the same translation, and a WHERE clause that cannot be translated now fails the statement: it previously became an empty filter and matched every document. --- pymongosql/sql/builder.py | 10 ++ pymongosql/sql/delete_handler.py | 5 + pymongosql/sql/handler.py | 8 +- pymongosql/sql/query_handler.py | 19 ++- pymongosql/sql/update_handler.py | 5 + pymongosql/sql/where_tree.py | 143 ++++++++++++++++++ .../sqlalchemy_mongodb/sqlalchemy_dialect.py | 9 ++ tests/test_sql_grouping_filters_aliases.py | 4 +- tests/test_sql_not_and_null_semantics.py | 142 +++++++++++++++++ tests/test_sql_parser_comprehensive.py | 2 +- tests/test_sql_parser_delete.py | 2 +- tests/test_sql_parser_general.py | 4 +- tests/test_sql_parser_nested_fields.py | 2 +- 13 files changed, 337 insertions(+), 18 deletions(-) create mode 100644 pymongosql/sql/where_tree.py create mode 100644 tests/test_sql_not_and_null_semantics.py diff --git a/pymongosql/sql/builder.py b/pymongosql/sql/builder.py index 68dd7aa..6bead70 100644 --- a/pymongosql/sql/builder.py +++ b/pymongosql/sql/builder.py @@ -300,6 +300,11 @@ def _build_insert_plan(parse_result: "InsertParseResult") -> "InsertExecutionPla @staticmethod def _build_delete_plan(parse_result: "DeleteParseResult") -> "DeleteExecutionPlan": """Build a DELETE execution plan from DELETE parsing.""" + from ..error import SqlSyntaxError + + if parse_result.has_errors: + # An untranslated WHERE must never become an empty filter (every document) + raise SqlSyntaxError(parse_result.error_message or "DELETE parsing failed") _logger.debug( f"Building DELETE plan with collection: {parse_result.collection}, " f"filters: {parse_result.filter_conditions}" @@ -314,6 +319,11 @@ def _build_delete_plan(parse_result: "DeleteParseResult") -> "DeleteExecutionPla @staticmethod def _build_update_plan(parse_result: "UpdateParseResult") -> "UpdateExecutionPlan": """Build an UPDATE execution plan from UPDATE parsing.""" + from ..error import SqlSyntaxError + + if parse_result.has_errors: + # An untranslated WHERE must never become an empty filter (every document) + raise SqlSyntaxError(parse_result.error_message or "UPDATE parsing failed") _logger.debug( f"Building UPDATE plan with collection: {parse_result.collection}, " f"update_fields: {parse_result.update_fields}, " diff --git a/pymongosql/sql/delete_handler.py b/pymongosql/sql/delete_handler.py index c59643b..3724c71 100644 --- a/pymongosql/sql/delete_handler.py +++ b/pymongosql/sql/delete_handler.py @@ -129,6 +129,11 @@ def handle_where_clause( _logger.debug(f"[WHERE_CLAUSE_DEBUG] Expression context type: {type(expression_ctx).__name__}") from .handler import HandlerFactory + from .where_tree import WhereTreeBuilder + + if expression_ctx is not None: + parse_result.filter_conditions = WhereTreeBuilder().build(expression_ctx) + return parse_result.filter_conditions handler = HandlerFactory.get_expression_handler(expression_ctx) diff --git a/pymongosql/sql/handler.py b/pymongosql/sql/handler.py index fce8c93..8f1b3ae 100644 --- a/pymongosql/sql/handler.py +++ b/pymongosql/sql/handler.py @@ -257,6 +257,12 @@ def _build_mongo_filter(self, field_name: str, operator: str, value: Any) -> Dic if operator == "=": return {field_name: value} + # SQL <>, NOT IN and NOT LIKE are never TRUE for a NULL or missing field + if operator in ("!=", "<>", "NOT IN", "NOT LIKE") or (operator == "LIKE" and value == "?"): + from .where_tree import leaf_filters + + return leaf_filters(field_name, operator, value)[0] + # Handle special operators if operator in ("IN", "NOT IN"): values = value if isinstance(value, list) else [value] @@ -571,7 +577,7 @@ def _extract_like_pattern(self, text: str) -> str: idx = text.upper().find("LIKE") if idx == -1: return "" - return text[idx + 4 :].strip().strip("'\"") + return self._parse_value(text[idx + 4 :].strip()) def _extract_between_range(self, text: str) -> Optional[Tuple[Any, Any]]: """Extract range values from BETWEEN clause""" diff --git a/pymongosql/sql/query_handler.py b/pymongosql/sql/query_handler.py index a9bdd50..f440db9 100644 --- a/pymongosql/sql/query_handler.py +++ b/pymongosql/sql/query_handler.py @@ -314,16 +314,15 @@ def can_handle(self, ctx: Any) -> bool: def handle_visitor(self, ctx: PartiQLParser.WhereClauseSelectContext, parse_result: "QueryParseResult") -> Any: if hasattr(ctx, "exprSelect") and ctx.exprSelect(): + from .where_tree import WhereTreeBuilder + + # Translate over the parse tree with SQL three-valued logic. A clause that + # cannot be translated fails the query; it never falls back to a partial + # or text-search filter that would return different rows. try: - # Use enhanced expression handler for better parsing - filter_conditions = self._expression_handler.handle(ctx) - parse_result.filter_conditions = filter_conditions - return filter_conditions + parse_result.filter_conditions = WhereTreeBuilder().build(ctx.exprSelect()) except Exception as e: - _logger.warning(f"Failed to parse WHERE expression, falling back to text search: {e}") - # Fallback to simple text search - filter_text = ctx.exprSelect().getText() - fallback_filter = {"$text": {"$search": filter_text}} - parse_result.filter_conditions = fallback_filter - return fallback_filter + parse_result.unsupported_clauses.append(f"WHERE ({e})") + parse_result.filter_conditions = {} + return parse_result.filter_conditions return {} diff --git a/pymongosql/sql/update_handler.py b/pymongosql/sql/update_handler.py index 6b08f57..e46f45e 100644 --- a/pymongosql/sql/update_handler.py +++ b/pymongosql/sql/update_handler.py @@ -190,6 +190,11 @@ def handle_where_clause(self, ctx: Any, parse_result: UpdateParseResult) -> Dict if expression_ctx: from .handler import HandlerFactory + from .where_tree import WhereTreeBuilder + + if expression_ctx is not None: + parse_result.filter_conditions = WhereTreeBuilder().build(expression_ctx) + return parse_result.filter_conditions handler = HandlerFactory.get_expression_handler(expression_ctx) diff --git a/pymongosql/sql/where_tree.py b/pymongosql/sql/where_tree.py new file mode 100644 index 0000000..012b033 --- /dev/null +++ b/pymongosql/sql/where_tree.py @@ -0,0 +1,143 @@ +# -*- coding: utf-8 -*- +"""WHERE translation over the parse tree with SQL three-valued logic. + +Every predicate yields two MongoDB filters: the documents for which it is TRUE and +those for which it is FALSE. A NULL (or missing) operand makes a comparison +UNKNOWN, which is in neither set. ``NOT p`` swaps the two sets, and De Morgan's +laws combine them for AND and OR, so ``NOT a = 1`` excludes documents where ``a`` +is NULL or missing, exactly like SQL. +""" + +import re +from typing import Any, Dict, List, Tuple + +from ..error import NotSupportedError +from .partiql.PartiQLParser import PartiQLParser + +Filter = Dict[str, Any] +Pair = Tuple[Filter, Filter] + +# Matches no document; used for predicates that can never be TRUE (or FALSE). +NOTHING: Filter = {"$expr": False} + +_LEAVES = ( + PartiQLParser.PredicateComparisonContext, + PartiQLParser.PredicateIsContext, + PartiQLParser.PredicateInContext, + PartiQLParser.PredicateLikeContext, + PartiQLParser.PredicateBetweenContext, +) +_FIELD_PATH = re.compile(r'^(?:"[^"]+"|[A-Za-z_$][\w$]*)(?:\.(?:"[^"]+"|[A-Za-z_$][\w$]*|\d+))*$') +_FIELD = re.compile(r"^[A-Za-z_$][\w$]*(?:\.[\w$]+)*$") +_SWAP = {"<": ">=", ">=": "<", ">": "<=", "<=": ">"} +_MONGO = {"<": "$lt", "<=": "$lte", ">": "$gt", ">=": "$gte"} + + +def _all(parts: List[Tuple[Any, Filter]], key: str, chain: type) -> Filter: + """Combine filters under ``key``, flattening only an unparenthesized chain of the same operator.""" + items: List[Filter] = [] + for ctx, f in parts: + items.extend(f[key] if isinstance(ctx, chain) and list(f) == [key] else [f]) + return {key: items} + + +def contains_not(ctx: Any) -> bool: + """Whether the expression has a boolean NOT (not a NOT IN / NOT LIKE predicate).""" + if isinstance(ctx, PartiQLParser.NotContext): + return True + return any(contains_not(child) for child in getattr(ctx, "children", None) or []) + + +def leaf_filters(field: str, operator: str, value: Any) -> Pair: + """TRUE and FALSE filters for one predicate on ``field``.""" + op = operator.upper() + if op == "IS NULL": + return {field: {"$eq": None}}, {field: {"$ne": None}} + if op == "IS NOT NULL": + return {field: {"$ne": None}}, {field: {"$eq": None}} + if op in ("IN", "NOT IN"): + values = value if isinstance(value, list) else [value] + present = [v for v in values if v is not None] + true = {field: {"$in": present}} + # x NOT IN (..., NULL) is never TRUE; x IN (..., NULL) is never FALSE + false = NOTHING if None in values else {field: {"$nin": present + [None]}} + return (true, false) if op == "IN" else (false, true) + if op in ("LIKE", "NOT LIKE"): + from .handler import ComparisonExpressionHandler + + if value == "?" or not isinstance(value, str): + # The pattern is translated to a regex while parsing, before parameters are bound + raise NotSupportedError("LIKE needs a literal pattern, not a bound parameter") + + pattern = ComparisonExpressionHandler._like_to_regex(value) + pattern = ("" if pattern.startswith(".*") else "^") + pattern + ("" if pattern.endswith(".*") else "$") + true = {field: {"$regex": pattern}} + false = {"$and": [{field: {"$not": {"$regex": pattern}}}, {field: {"$ne": None}}]} + return (true, false) if op == "LIKE" else (false, true) + if op == "BETWEEN": + low, high = value + true = {"$and": [{field: {"$gte": low}}, {field: {"$lte": high}}]} + return true, {"$or": [{field: {"$lt": low}}, {field: {"$gt": high}}]} + if value is None: + # This dialect has always read "= NULL" / "<> NULL" as IS NULL / IS NOT NULL; + # any other comparison with NULL is UNKNOWN. + if op == "=": + return {field: None}, {field: {"$ne": None}} + if op in ("!=", "<>"): + return {field: {"$ne": None}}, {field: None} + return NOTHING, NOTHING + if op == "=": + return {field: value}, {field: {"$nin": [value, None]}} + if op in ("!=", "<>"): + return {field: {"$nin": [value, None]}}, {field: value} + if op in _MONGO: + return {field: {_MONGO[op]: value}}, {field: {_MONGO[_SWAP[op]]: value}} + raise NotSupportedError(f"Unsupported predicate operator: {operator}") + + +class WhereTreeBuilder: + """Build a MongoDB filter for a WHERE expression from its parse tree.""" + + def build(self, ctx: Any) -> Filter: + return self._pair(ctx)[0] + + def _pair(self, ctx: Any) -> Pair: + if isinstance(ctx, PartiQLParser.NotContext): + true, false = self._pair(ctx.rhs) + return false, true + if isinstance(ctx, (PartiQLParser.AndContext, PartiQLParser.OrContext)): + is_and = isinstance(ctx, PartiQLParser.AndContext) + (t1, f1), (t2, f2) = self._pair(ctx.lhs), self._pair(ctx.rhs) + true_key, false_key = ("$and", "$or") if is_and else ("$or", "$and") + chain = type(ctx) + return ( + _all([(ctx.lhs, t1), (ctx.rhs, t2)], true_key, chain), + _all([(ctx.lhs, f1), (ctx.rhs, f2)], false_key, chain), + ) + if isinstance(ctx, PartiQLParser.ExprTermWrappedQueryContext): + return self._pair(ctx.expr()) + if isinstance(ctx, _LEAVES): + return self._leaf(ctx) + children = [c for c in getattr(ctx, "children", None) or [] if hasattr(c, "getRuleIndex")] + if len(children) == 1 and len(ctx.children) == 1: + return self._pair(children[0]) + text = ctx.getText() + if _FIELD_PATH.match(text) and text.upper() not in ("TRUE", "FALSE", "NULL"): + # A bare boolean field: WHERE flag / WHERE NOT flag + from .handler import ContextUtilsMixin + + field = ContextUtilsMixin.normalize_field_path(text) + return {field: True}, {field: False} + raise NotSupportedError(f"Unsupported WHERE expression: {text}") + + @staticmethod + def _leaf(ctx: Any) -> Pair: + from .handler import ComparisonExpressionHandler + + handler = ComparisonExpressionHandler() + field = handler._extract_field_name(ctx) + operator = handler._extract_operator(ctx) + value = handler._extract_value(ctx) + if not _FIELD.match(field): + raise NotSupportedError(f"Unsupported WHERE predicate: {ctx.getText()}") + return leaf_filters(field, operator, value) diff --git a/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py b/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py index b6a271e..4b40fa6 100644 --- a/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py +++ b/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py @@ -111,6 +111,15 @@ def visit_column(self, column, include_table=True, **kwargs): """ return super().visit_column(column, include_table=False, **kwargs) + def visit_like_op_binary(self, binary, operator, **kw): + """Render LIKE patterns inline: PyMongoSQL turns them into a regex while parsing.""" + kw["literal_binds"] = True + return super().visit_like_op_binary(binary, operator, **kw) + + def visit_not_like_op_binary(self, binary, operator, **kw): + kw["literal_binds"] = True + return super().visit_not_like_op_binary(binary, operator, **kw) + class PyMongoSQLDDLCompiler(compiler.DDLCompiler): """MongoDB-specific DDL compiler. diff --git a/tests/test_sql_grouping_filters_aliases.py b/tests/test_sql_grouping_filters_aliases.py index 65fffec..6590a76 100644 --- a/tests/test_sql_grouping_filters_aliases.py +++ b/tests/test_sql_grouping_filters_aliases.py @@ -63,14 +63,14 @@ def test_in_keeps_literal_types_and_quoted_commas(self): assert p.filter_stage == {"_id": {"$in": [1, 2.5, "a,b", "it's", True, "?"]}} def test_not_in(self): - assert plan("SELECT _id FROM t WHERE _id NOT IN (1, 2)").filter_stage == {"_id": {"$nin": [1, 2]}} + assert plan("SELECT _id FROM t WHERE _id NOT IN (1, 2)").filter_stage == {"_id": {"$nin": [1, 2, None]}} def test_field_ending_in_not_is_not_negated(self): assert plan("SELECT _id FROM t WHERE cannot IN (1)").filter_stage == {"cannot": {"$in": [1]}} def test_not_like_and_regex_metacharacters(self): p = plan("SELECT _id FROM t WHERE name NOT LIKE 'x.%'") - assert p.filter_stage == {"name": {"$not": {"$regex": "^x\\..*"}}} + assert p.filter_stage == {"$and": [{"name": {"$not": {"$regex": "^x\\..*"}}}, {"name": {"$ne": None}}]} def test_double_dash_inside_literal_is_not_a_comment(self): p = plan("SELECT _id FROM t WHERE name = 'a -- b' AND n = 1 -- trailing comment") diff --git a/tests/test_sql_not_and_null_semantics.py b/tests/test_sql_not_and_null_semantics.py new file mode 100644 index 0000000..c1f972f --- /dev/null +++ b/tests/test_sql_not_and_null_semantics.py @@ -0,0 +1,142 @@ +# -*- coding: utf-8 -*- +"""NOT and NULL follow SQL three-valued logic; an untranslatable WHERE never widens a query.""" + +import pytest + +from pymongosql.error import Error +from pymongosql.sql.parser import SQLParser +from tests.conftest import HAS_SQLALCHEMY + +COLLECTION = "test_not_semantics" +# _id 4 has a NULL a; _id 5 has no a at all. SQL never returns them for a +# comparison on a, negated or not. +DOCS = [ + {"_id": 1, "a": 1, "b": "x", "flag": True}, + {"_id": 2, "a": 2, "b": "y", "flag": False}, + {"_id": 3, "a": 3, "b": "xz", "flag": True}, + {"_id": 4, "a": None, "b": None, "flag": None}, + {"_id": 5}, +] + + +def plan(sql): + return SQLParser(sql).get_execution_plan() + + +class TestPlans: + def test_not_comparison_excludes_null(self): + assert plan("SELECT _id FROM t WHERE NOT a = 1").filter_stage == {"a": {"$nin": [1, None]}} + + def test_not_or_uses_de_morgan(self): + assert plan("SELECT _id FROM t WHERE NOT (a = 1 OR b = 'x')").filter_stage == { + "$and": [{"a": {"$nin": [1, None]}}, {"b": {"$nin": ["x", None]}}] + } + + def test_double_not(self): + assert plan("SELECT _id FROM t WHERE NOT NOT a = 1").filter_stage == {"a": 1} + + def test_not_between_and_not_is_null(self): + assert plan("SELECT _id FROM t WHERE NOT a BETWEEN 1 AND 2").filter_stage == { + "$or": [{"a": {"$lt": 1}}, {"a": {"$gt": 2}}] + } + assert plan("SELECT _id FROM t WHERE NOT a IS NULL").filter_stage == {"a": {"$ne": None}} + + def test_not_bare_boolean_field(self): + assert plan("SELECT _id FROM t WHERE NOT flag").filter_stage == {"flag": False} + + def test_not_in_list_containing_null_is_never_true(self): + assert plan("SELECT _id FROM t WHERE a NOT IN (1, NULL)").filter_stage == {"$expr": False} + + def test_not_equal_excludes_null(self): + assert plan("SELECT _id FROM t WHERE a <> 1").filter_stage == {"a": {"$nin": [1, None]}} + + @pytest.mark.parametrize( + "sql", + [ + "SELECT _id FROM t WHERE NOT lower(b) = 'x'", + "SELECT _id FROM t WHERE lower(b) = 'x'", + "SELECT _id FROM t WHERE a = 1 AND lower(b) = 'x'", + "SELECT _id FROM t WHERE b LIKE ?", + ], + ) + def test_untranslatable_where_raises_instead_of_widening(self, sql): + with pytest.raises(Error): + plan(sql) + + @pytest.mark.parametrize("sql", ["DELETE FROM t WHERE lower(b) = 'x'", "UPDATE t SET a = 1 WHERE lower(b) = 'x'"]) + def test_untranslatable_dml_where_raises_instead_of_matching_everything(self, sql): + with pytest.raises(Error): + plan(sql) + + def test_not_in_delete(self): + assert plan("DELETE FROM t WHERE NOT a = 1").filter_conditions == {"a": {"$nin": [1, None]}} + + +@pytest.fixture +def docs(conn): + conn.database.drop_collection(COLLECTION) + conn.database[COLLECTION].insert_many(DOCS) + yield conn + conn.database.drop_collection(COLLECTION) + + +def ids(conn, where, params=None): + cursor = conn.cursor() + sql = f"SELECT _id FROM {COLLECTION} WHERE {where} ORDER BY _id" + cursor.execute(sql, params) if params is not None else cursor.execute(sql) + return [r[0] for r in cursor.fetchall()] + + +class TestLive: + @pytest.mark.parametrize( + "where,expected", + [ + ("NOT a = 1", [2, 3]), + ("NOT (a = 1 OR b = 'y')", [3]), + ("NOT (a = 1 AND b = 'x')", [2, 3]), + ("NOT a IN (1, 2)", [3]), + ("NOT a NOT IN (1, 2)", [1, 2]), + ("NOT b LIKE 'x%'", [2]), + ("NOT a BETWEEN 2 AND 3", [1]), + ("NOT a IS NULL", [1, 2, 3]), + ("NOT flag", [2]), + ("a <> 1", [2, 3]), + ("b NOT LIKE 'x%'", [2]), + ("a NOT IN (1)", [2, 3]), + ("a = 1 OR NOT (b = 'x' OR b = 'xz')", [1, 2]), + ], + ) + def test_negation_returns_sql_rows(self, docs, where, expected): + assert ids(docs, where) == expected + + def test_not_with_bound_parameter(self, docs): + assert ids(docs, "NOT a = ?", [2]) == [1, 3] + + def test_untranslatable_delete_deletes_nothing(self, docs): + cursor = docs.cursor() + with pytest.raises(Error): + cursor.execute(f"DELETE FROM {COLLECTION} WHERE lower(b) = 'x'") + assert docs.database[COLLECTION].count_documents({}) == len(DOCS) + + +@pytest.mark.skipif(not HAS_SQLALCHEMY, reason="SQLAlchemy not available") +class TestSQLAlchemy: + def test_like_pattern_is_rendered_inline(self): + import sqlalchemy as sa + + from pymongosql.sqlalchemy_mongodb.sqlalchemy_dialect import PyMongoSQLDialect + + t = sa.table("t", sa.column("b")) + stmt = sa.select(t.c.b).where(t.c.b.like("O'B%"), t.c.b.not_like("x_")) + sql = " ".join(str(stmt.compile(dialect=PyMongoSQLDialect())).split()) + assert sql == "SELECT b FROM t WHERE b LIKE 'O''B%' AND b NOT LIKE 'x_'" + + def test_core_not_and_like_rows(self, sqlalchemy_engine, docs): + import sqlalchemy as sa + + t = sa.table(COLLECTION, sa.column("_id"), sa.column("a"), sa.column("b")) + with sqlalchemy_engine.connect() as connection: + got = connection.execute( + sa.select(t.c._id).where(sa.not_(t.c.a == 1), t.c.b.like("x%")).order_by(t.c._id) + ).scalars() + assert list(got) == [3] diff --git a/tests/test_sql_parser_comprehensive.py b/tests/test_sql_parser_comprehensive.py index ec19beb..1039787 100644 --- a/tests/test_sql_parser_comprehensive.py +++ b/tests/test_sql_parser_comprehensive.py @@ -157,7 +157,7 @@ def test_bool_and_bracketed_or(self): def test_bool_not_equal_and_comparison(self): sql = "SELECT * FROM col WHERE active!=false AND age>25" plan = SQLParser(sql).get_execution_plan() - assert plan.filter_stage == {"$and": [{"active": {"$ne": False}}, {"age": {"$gt": 25}}]} + assert plan.filter_stage == {"$and": [{"active": {"$nin": [False, None]}}, {"age": {"$gt": 25}}]} # --- null mixed with bool --- diff --git a/tests/test_sql_parser_delete.py b/tests/test_sql_parser_delete.py index 395f37b..83bf6f5 100644 --- a/tests/test_sql_parser_delete.py +++ b/tests/test_sql_parser_delete.py @@ -67,7 +67,7 @@ def test_delete_with_not_equal(self): assert isinstance(plan, DeleteExecutionPlan) assert plan.collection == "temp" - assert plan.filter_conditions == {"valid": {"$ne": True}} + assert plan.filter_conditions == {"valid": {"$nin": [True, None]}} def test_delete_with_qmark_parameter(self): """Test DELETE with qmark placeholder.""" diff --git a/tests/test_sql_parser_general.py b/tests/test_sql_parser_general.py index 8354bdf..cd1cdd0 100644 --- a/tests/test_sql_parser_general.py +++ b/tests/test_sql_parser_general.py @@ -101,7 +101,7 @@ def test_select_with_not_equals(self): execution_plan = parser.get_execution_plan() assert execution_plan.collection == "users" - assert execution_plan.filter_stage == {"status": {"$ne": "inactive"}} + assert execution_plan.filter_stage == {"status": {"$nin": ["inactive", None]}} assert execution_plan.projection_stage == {"name": 1} def test_select_with_and_condition(self): @@ -335,7 +335,7 @@ def test_complex_mixed_operators(self): # Verify complex filter structure with mixed AND/OR conditions expected_filter = { "$or": [ - {"$and": [{"age": {"$gt": 25}}, {"status": "active"}, {"name": {"$ne": "John"}}]}, + {"$and": [{"age": {"$gt": 25}}, {"status": "active"}, {"name": {"$nin": ["John", None]}}]}, {"department": {"$in": ["IT", "HR"]}}, ] } diff --git a/tests/test_sql_parser_nested_fields.py b/tests/test_sql_parser_nested_fields.py index eeed246..5cf6999 100644 --- a/tests/test_sql_parser_nested_fields.py +++ b/tests/test_sql_parser_nested_fields.py @@ -146,7 +146,7 @@ def test_nested_with_comparison_operators(self): ("profile.age > 18", {"profile.age": {"$gt": 18}}), ("settings.total < 100", {"settings.total": {"$lt": 100}}), # Changed from 'count' (reserved) ("status.active = true", {"status.active": True}), - ("config.name != 'default'", {"config.name": {"$ne": "default"}}), + ("config.name != 'default'", {"config.name": {"$nin": ["default", None]}}), ] for where_clause, expected_filter in test_cases: From 7cf1fb7b74c00d3dc002bd9844cdb62888d47ae9 Mon Sep 17 00:00:00 2001 From: Amin Ghadersohi Date: Sat, 26 Sep 2026 06:02:30 +0000 Subject: [PATCH 09/19] fix: resolve FROM aliases and reject FROM clauses that cannot be translated The FROM handler used the whole table reference text as the collection name, so FROM users AS u read a collection named "usersASu" and returned no rows, and u.name was read as an embedded path. Joins and, outside superset mode, subqueries were treated the same way and silently returned nothing. Read the collection and its alias from the parse tree, resolve alias-qualified references (u.name, including in GROUP BY and the ordered SELECT list) to the field, and raise NotSupportedError for joins, subqueries in standard mode and AT/BY bindings. Collection- qualified GROUP BY keys are also resolved now; they grouped on a missing nested path before. --- pymongosql/sql/builder.py | 12 +++- pymongosql/sql/query_handler.py | 43 ++++++++++++-- tests/test_sql_from_alias.py | 98 +++++++++++++++++++++++++++++++ tests/test_superset_connection.py | 8 ++- 4 files changed, 151 insertions(+), 10 deletions(-) create mode 100644 tests/test_sql_from_alias.py diff --git a/pymongosql/sql/builder.py b/pymongosql/sql/builder.py index 6bead70..db58f50 100644 --- a/pymongosql/sql/builder.py +++ b/pymongosql/sql/builder.py @@ -123,11 +123,13 @@ def _strip_collection_qualifier(parse_result: "QueryParseResult") -> None: collection = parse_result.collection if not collection: return - prefix = f"{collection}." + # With FROM users AS u, u.name is the column name; so is users.name + prefixes = [f"{q}." for q in (parse_result.collection_alias, collection) if q] def strip(name: Any) -> Any: - if isinstance(name, str) and name.startswith(prefix) and len(name) > len(prefix): - return name[len(prefix) :] + for prefix in prefixes: + if isinstance(name, str) and name.startswith(prefix) and len(name) > len(prefix): + return name[len(prefix) :] return name def strip_filter(value: Any) -> Any: @@ -143,6 +145,10 @@ def strip_filter(value: Any) -> Any: parse_result.filter_conditions = strip_filter(parse_result.filter_conditions) for func_info in parse_result.aggregate_functions: func_info["argument"] = strip(func_info["argument"]) + parse_result.group_by = [strip(name) for name in parse_result.group_by] + for item in parse_result.select_items: + if "field" in item: + item["field"] = strip(item["field"]) @staticmethod def _build_query_plan(parse_result: "QueryParseResult") -> "QueryExecutionPlan": diff --git a/pymongosql/sql/query_handler.py b/pymongosql/sql/query_handler.py index f440db9..5c1b2fa 100644 --- a/pymongosql/sql/query_handler.py +++ b/pymongosql/sql/query_handler.py @@ -40,6 +40,8 @@ class QueryParseResult: group_by: List[str] = field(default_factory=list) # Clauses that are parsed but cannot be translated faithfully unsupported_clauses: List[str] = field(default_factory=list) + # FROM alias (FROM users AS u / FROM users u) + collection_alias: Optional[str] = None # Subquery info (for wrapped subqueries, e.g., Superset outering) subquery_plan: Optional[Any] = None @@ -269,6 +271,35 @@ def _parse_function_call(self, ctx: Any) -> Optional[Dict[str, Any]]: _logger.debug(f"Error parsing function call: {e}") return None + @staticmethod + def _collection_reference(table_ref: Any) -> Tuple[Optional[str], Optional[str], Optional[str]]: + """Return (collection text, alias, problem) for a FROM table reference. + + Only a single collection, optionally aliased, can be translated. Joins, subqueries + and AT/BY bindings are reported as a problem so the query fails instead of reading + a collection named after the whole clause. + """ + if not hasattr(table_ref, "getRuleIndex"): + return table_ref.getText(), None, None # not a parse-tree node: a plain name + while isinstance(table_ref, PartiQLParser.TableWrappedContext): + table_ref = table_ref.tableReference() + if not isinstance(table_ref, PartiQLParser.TableRefBaseContext): + return None, None, "FROM with a join" + base = table_ref.tableNonJoin().tableBaseReference() + if isinstance(base, PartiQLParser.TableBaseRefSymbolContext): + source, alias = base.source, base.symbolPrimitive().getText() + elif isinstance(base, PartiQLParser.TableBaseRefClausesContext): + if base.atIdent() is not None or base.byIdent() is not None: + return None, None, "FROM ... AT/BY" + source = base.source + alias = base.asIdent().symbolPrimitive().getText() if base.asIdent() is not None else None + else: + return None, None, "FROM with UNPIVOT or a graph match" + text = source.getText() + if text.startswith("("): + return None, None, "FROM a subquery (use mode=superset)" + return text, alias, None + def handle_visitor(self, ctx: PartiQLParser.FromClauseContext, parse_result: "QueryParseResult") -> Any: """Handle FROM clause - detect aggregate calls or regular collections""" if hasattr(ctx, "tableReference") and ctx.tableReference(): @@ -291,12 +322,16 @@ def handle_visitor(self, ctx: PartiQLParser.FromClauseContext, parse_result: "Qu _logger.info(f"Parsed aggregate call: collection={func_info['collection']}") return func_info - # Regular collection reference - table_text = ctx.tableReference().getText() + # Regular collection reference, optionally aliased + source, alias, problem = self._collection_reference(ctx.tableReference()) + if problem: + parse_result.unsupported_clauses.append(problem) + return None # Strip surrounding quotes from collection name (e.g., "user.accounts" -> user.accounts) - collection_name = self._strip_collection_quotes(table_text) + collection_name = self._strip_collection_quotes(source) parse_result.collection = collection_name - _logger.debug(f"Parsed regular collection: {collection_name}") + parse_result.collection_alias = ContextUtilsMixin.unquote_identifier(alias) if alias else None + _logger.debug(f"Parsed regular collection: {collection_name} (alias {alias})") return collection_name return None diff --git a/tests/test_sql_from_alias.py b/tests/test_sql_from_alias.py new file mode 100644 index 0000000..82eaefb --- /dev/null +++ b/tests/test_sql_from_alias.py @@ -0,0 +1,98 @@ +# -*- coding: utf-8 -*- +"""FROM aliases resolve to the collection; untranslatable FROM clauses raise.""" + +import pytest + +from pymongosql.error import Error +from pymongosql.sql.parser import SQLParser +from tests.conftest import HAS_SQLALCHEMY + +COLLECTION = "test_from_alias" +DOCS = [ + {"_id": 1, "g": "a", "v": 10, "profile": {"city": "x"}}, + {"_id": 2, "g": "a", "v": 20, "profile": {"city": "y"}}, + {"_id": 3, "g": "b", "v": 35, "profile": {"city": "x"}}, +] + + +def plan(sql): + return SQLParser(sql).get_execution_plan() + + +class TestPlans: + @pytest.mark.parametrize("from_clause", ["t AS x", "t x", 't AS "x"']) + def test_alias_qualified_columns_resolve_to_fields(self, from_clause): + p = plan(f"SELECT x.a, x.b AS bee FROM {from_clause} WHERE x.b = 1 AND NOT x.c = 2 ORDER BY x.a") + assert p.collection == "t" + assert p.projection_stage == {"a": 1, "b": 1} + assert p.column_aliases == {"b": "bee"} + assert p.filter_stage == {"$and": [{"b": 1}, {"c": {"$nin": [2, None]}}]} + assert p.sort_stage == [{"a": 1}] + + def test_alias_with_nested_path(self): + assert plan("SELECT x.profile.city FROM t AS x").projection_stage == {"profile.city": 1} + + def test_quoted_collection_with_alias(self): + p = plan('SELECT ua.a FROM "user.accounts" AS ua') + assert (p.collection, p.projection_stage) == ("user.accounts", {"a": 1}) + + def test_qualified_group_by_keys(self): + for sql in ( + "SELECT t.g, SUM(t.v) AS s FROM t GROUP BY t.g", + "SELECT x.g, SUM(x.v) AS s FROM t x GROUP BY x.g", + ): + assert '"_id": {"g0": "$g"}' in plan(sql).aggregate_pipeline + assert '"$sum": "$v"' in plan(sql).aggregate_pipeline + + @pytest.mark.parametrize( + "sql", + [ + "SELECT a FROM t, u", + "SELECT a FROM t JOIN u ON t.a = u.a", + "SELECT a FROM (SELECT a FROM t) AS v", + "SELECT a FROM t AS x AT i", + ], + ) + def test_untranslatable_from_raises(self, sql): + with pytest.raises(Error, match="Unsupported SQL clause: FROM"): + plan(sql) + + +@pytest.fixture +def docs(conn): + conn.database.drop_collection(COLLECTION) + conn.database[COLLECTION].insert_many(DOCS) + yield conn + conn.database.drop_collection(COLLECTION) + + +def rows(conn, sql): + cursor = conn.cursor() + cursor.execute(sql) + return [tuple(r) for r in cursor.fetchall()], [d[0] for d in cursor.description] + + +class TestLive: + def test_aliased_select_returns_rows(self, docs): + got, names = rows( + docs, f"SELECT x._id, x.profile.city AS city FROM {COLLECTION} AS x WHERE x.v > 10 ORDER BY x._id" + ) + assert got == [(2, "y"), (3, "x")] + assert names == ["_id", "city"] + + def test_aliased_group_by_returns_rows(self, docs): + got, _ = rows(docs, f"SELECT x.g, SUM(x.v) AS s FROM {COLLECTION} x GROUP BY x.g ORDER BY s DESC") + assert got == [("b", 35), ("a", 30)] + got, _ = rows(docs, f"SELECT x.g, COUNT(*) AS n FROM {COLLECTION} x GROUP BY x.g ORDER BY x.g") + assert got == [("a", 2), ("b", 1)] + + +@pytest.mark.skipif(not HAS_SQLALCHEMY, reason="SQLAlchemy not available") +class TestSQLAlchemy: + def test_core_alias(self, sqlalchemy_engine, docs): + import sqlalchemy as sa + + t = sa.table(COLLECTION, sa.column("_id"), sa.column("v")).alias("x") + with sqlalchemy_engine.connect() as connection: + got = connection.execute(sa.select(t.c._id).where(t.c.v >= 20).order_by(t.c._id)).scalars() + assert list(got) == [2, 3] diff --git a/tests/test_superset_connection.py b/tests/test_superset_connection.py index 159f265..71320d2 100644 --- a/tests/test_superset_connection.py +++ b/tests/test_superset_connection.py @@ -1,4 +1,6 @@ # -*- coding: utf-8 -*- +import pytest + from pymongosql.executor import ExecutionContext, ExecutionPlanFactory from pymongosql.helper import ConnectionHelper from pymongosql.superset_mongodb.executor import SupersetExecution @@ -142,9 +144,9 @@ def test_core_connection_with_subqueries(self, conn): cursor = conn.cursor() subquery_sql = "SELECT * FROM (SELECT _id, name FROM users) AS u WHERE u.age > 25" - cursor.execute(subquery_sql) - rows = cursor.fetchall() - assert len(rows) == 0 + # Standard mode cannot evaluate a subquery; it must fail, not return no rows + with pytest.raises(Exception, match="subquery"): + cursor.execute(subquery_sql) def test_core_connection_with_standard_queries(self, conn): """Test simple query on users collection""" From c74dc40fc2eef614b8cc998c6485e7399e83ca21 Mon Sep 17 00:00:00 2001 From: Amin Ghadersohi Date: Sat, 26 Sep 2026 06:03:41 +0000 Subject: [PATCH 10/19] fix(superset): keep booleans and nullable numbers typed in the SQLite stage Superset-mode subqueries load the MongoDB rows into an in-memory SQLite table. SQLite has no boolean type, so boolean columns came back as 1/0, and a single NULL in a column made the whole column TEXT, returning numbers and booleans as strings. Declare boolean columns with a private type that is converted back to bool when a query selects the column (expressions such as SUM stay numeric), and let NULL values fit any column type when inferring the schema. --- .../superset_mongodb/query_db_sqlite.py | 27 ++++---- tests/test_superset_sqlite_types.py | 63 +++++++++++++++++++ 2 files changed, 79 insertions(+), 11 deletions(-) create mode 100644 tests/test_superset_sqlite_types.py diff --git a/pymongosql/superset_mongodb/query_db_sqlite.py b/pymongosql/superset_mongodb/query_db_sqlite.py index 5e3d356..348a485 100644 --- a/pymongosql/superset_mongodb/query_db_sqlite.py +++ b/pymongosql/superset_mongodb/query_db_sqlite.py @@ -7,6 +7,12 @@ _logger = logging.getLogger(__name__) +# SQLite has no boolean type. Boolean columns are declared with this private type +# (NUMERIC affinity, stored as 0/1) and converted back to bool when a query selects +# the column itself; expressions over it (SUM, CASE, ...) keep their numeric result. +BOOLEAN_DECLTYPE = "PYMONGOSQL_BOOL" +sqlite3.register_converter(BOOLEAN_DECLTYPE, lambda raw: int(raw) != 0) + class SQLiteTypeMapper: """Maps Python/MongoDB data types to SQLite3 types""" @@ -16,7 +22,7 @@ class SQLiteTypeMapper: str: "TEXT", int: "INTEGER", float: "REAL", - bool: "INTEGER", # SQLite3 uses 0/1 for boolean + bool: BOOLEAN_DECLTYPE, # stored as 0/1, read back as bool bytes: "BLOB", type(None): "NULL", dict: "TEXT", # Store as JSON string @@ -51,15 +57,14 @@ def infer_schema(cls, records: List[Dict[str, Any]]) -> Dict[str, str]: for record in records: for col_name, value in record.items(): - if col_name not in schema: - # First occurrence, determine type - schema[col_name] = cls.get_sqlite_type(value) - elif schema[col_name] != "TEXT": - # If we've already determined type, check compatibility - new_type = cls.get_sqlite_type(value) + new_type = cls.get_sqlite_type(value) + current = schema.get(col_name, "NULL") + if current == "NULL": + # First non-null value determines the type; NULL fits every type + schema[col_name] = new_type + elif new_type not in ("NULL", current): # Upgrade to TEXT if types differ (safest option) - if new_type != schema[col_name]: - schema[col_name] = "TEXT" + schema[col_name] = "TEXT" return schema @@ -69,7 +74,7 @@ def convert_value(cls, value: Any, target_type: str) -> Any: if value is None: return None - if target_type == "INTEGER": + if target_type in ("INTEGER", BOOLEAN_DECLTYPE): return int(value) if value is not None else None elif target_type == "REAL": return float(value) if value is not None else None @@ -107,7 +112,7 @@ def _ensure_connection(self) -> sqlite3.Connection: if self._connection is None: # Create in-memory database - self._connection = sqlite3.connect(":memory:") + self._connection = sqlite3.connect(":memory:", detect_types=sqlite3.PARSE_DECLTYPES) # Enable row factory to get dict-like rows self._connection.row_factory = sqlite3.Row _logger.debug("Created in-memory SQLite3 database") diff --git a/tests/test_superset_sqlite_types.py b/tests/test_superset_sqlite_types.py new file mode 100644 index 0000000..c524908 --- /dev/null +++ b/tests/test_superset_sqlite_types.py @@ -0,0 +1,63 @@ +# -*- coding: utf-8 -*- +"""Superset-mode subqueries keep booleans and nullable numbers typed through SQLite.""" + +from pymongosql.superset_mongodb.query_db_sqlite import QueryDBSQLite, SQLiteTypeMapper +from tests.conftest import make_superset_conn + +COLLECTION = "test_superset_types" +DOCS = [ + {"_id": 1, "flag": True, "i": 5, "n": None}, + {"_id": 2, "flag": False, "i": -7, "n": 2}, + {"_id": 3, "flag": None, "i": None, "n": 3}, +] + + +def query(sql, records=None): + db = QueryDBSQLite() + try: + db.insert_records("v", records or [{k: v for k, v in d.items() if k != "_id"} for d in DOCS]) + return db.execute_query(sql) + finally: + db.close() + + +class TestSQLiteStage: + def test_selected_boolean_column_is_bool(self): + got = query("SELECT flag AS flag, SUM(i) AS s FROM v GROUP BY flag ORDER BY s DESC") + assert got == [{"flag": True, "s": 5}, {"flag": False, "s": -7}, {"flag": None, "s": None}] + assert [type(r["flag"]) for r in got] == [bool, bool, type(None)] + + def test_expression_over_boolean_stays_numeric(self): + assert query("SELECT SUM(flag) AS n FROM v") == [{"n": 1}] + + def test_null_does_not_turn_a_column_into_text(self): + assert SQLiteTypeMapper.infer_schema([{"n": None}, {"n": 2}, {"n": None}]) == {"n": "INTEGER"} + got = query("SELECT n FROM v ORDER BY n") + assert got == [{"n": None}, {"n": 2}, {"n": 3}] + assert type(got[1]["n"]) is int + + def test_mixed_types_still_fall_back_to_text(self): + assert SQLiteTypeMapper.infer_schema([{"x": 1}, {"x": "a"}]) == {"x": "TEXT"} + + +class TestLive: + def test_virtual_dataset_returns_booleans(self, conn): + conn.database.drop_collection(COLLECTION) + conn.database[COLLECTION].insert_many(DOCS) + superset = make_superset_conn() + try: + cursor = superset.cursor() + cursor.execute( + 'SELECT flag AS flag, SUM(i) AS "sum_i" FROM ' + f"(SELECT flag, i FROM {COLLECTION}) AS virtual_table " + "GROUP BY flag ORDER BY flag DESC" + ) + rows = [tuple(r) for r in cursor.fetchall()] + assert rows == [(True, 5), (False, -7), (None, None)] + assert type(rows[0][0]) is bool and type(rows[1][0]) is bool + cursor.execute(f"SELECT n FROM (SELECT n FROM {COLLECTION}) AS virtual_table ORDER BY n") + values = [r[0] for r in cursor.fetchall()] + assert values == [None, 2, 3] and type(values[1]) is int + finally: + superset.close() + conn.database.drop_collection(COLLECTION) From 6f843f50f97e658b3ba44645104d93ca6bb224b0 Mon Sep 17 00:00:00 2001 From: Amin Ghadersohi Date: Sat, 26 Sep 2026 10:10:26 +0000 Subject: [PATCH 11/19] fix: return Decimal and int instead of BSON Decimal128 and Int64 Result rows carried bson.Decimal128 and, in command responses, bson.Int64 values. DB API consumers expect decimal.Decimal and int; Decimal128 in particular cannot be summed or serialised by most libraries. Convert both, including inside embedded documents and arrays. --- pymongosql/result_set.py | 21 ++++++++++++++++++++- 1 file changed, 20 insertions(+), 1 deletion(-) diff --git a/pymongosql/result_set.py b/pymongosql/result_set.py index c1c1e70..aaf34bb 100644 --- a/pymongosql/result_set.py +++ b/pymongosql/result_set.py @@ -4,6 +4,7 @@ from typing import Any, Dict, List, Optional, Sequence, Tuple import jmespath +from bson import Decimal128, Int64 from pymongo.errors import PyMongoError from . import STRING @@ -63,7 +64,7 @@ def _process_and_cache_batch(self, batch: List[Dict[str, Any]]) -> None: if not batch: return # Process results through projection mapping - processed_batch = [self._process_document(doc) for doc in batch] + processed_batch = [self._to_python(self._process_document(doc)) for doc in batch] # Convert dictionaries to output format (sequence or dict) formatted_batch = [self._format_result(doc) for doc in processed_batch] self._cached_results.extend(formatted_batch) @@ -159,6 +160,24 @@ def _process_document(self, doc: Dict[str, Any]) -> Dict[str, Any]: return processed + @classmethod + def _to_python(cls, value: Any) -> Any: + """Return standard Python types for BSON-specific numbers. + + DB API 2.0 consumers expect ``decimal.Decimal`` and ``int``; PyMongo returns + ``bson.Decimal128`` (which most libraries cannot sum or serialise) and, in + command responses, ``bson.Int64``. + """ + if isinstance(value, Decimal128): + return value.to_decimal() + if isinstance(value, Int64): + return int(value) + if isinstance(value, dict): + return {k: cls._to_python(v) for k, v in value.items()} + if isinstance(value, list): + return [cls._to_python(v) for v in value] + return value + def _mongo_to_bracket_key(self, field_path: str) -> str: """Convert Mongo dot-index notation to bracket notation. From f3d45fb43b8dbe3cc6a83f208cfd38fd2b568238 Mon Sep 17 00:00:00 2001 From: Amin Ghadersohi Date: Sat, 26 Sep 2026 10:10:26 +0000 Subject: [PATCH 12/19] fix: translate predicates, aggregates and paging from the parse tree WHERE predicates recovered the field name by searching the predicate's concatenated token text for IN(, LIKE, ISNULL and similar, so a field whose name contains one of them was truncated: "dislikes IS NULL" filtered on "dis" and "unlike = 5" became a regex on "un". SELECT returned other rows, and DELETE and UPDATE removed or rewrote documents that did not match. The same text search rejected quoted field names with spaces, hyphens or non-ASCII characters, reversed comparisons (5 < age), and string literals containing IN( or LIKE. Read each predicate from its parse-tree node instead: the field is the path on one side of the operator (either side), the value the literal, parameter or value function on the other. Also: - LIKE: fold concatenated string literals ('%' || 'ab' || '%', as SQLAlchemy renders contains/startswith/endswith) and honour ESCAPE (like(..., escape=...), autoescape). ESCAPE is re-attached when the grammar lets it absorb the rest of the WHERE clause. - Parameters: a bound parameter is a marker in the translated filter, so a string literal '?' is compared as a value, and a parameter the statement does not use raises instead of being ignored. - LIMIT/OFFSET must be integer literals (the dialect renders them inline); LIMIT ? was dropped and returned every row. - COUNT/SUM/AVG/MIN/MAX(DISTINCT x) are computed from the set of distinct non-NULL values; COUNT(DISTINCT x) returned 0. - HAVING is translated into a $match after grouping, including aggregates that are not in the SELECT list. - SELECT expressions other than fields and aggregates raise instead of returning a NULL column. - A fractional literal no double represents exactly (0.1) is compared as a double against double fields and exactly (Decimal128) against other numeric types. - Generated pipelines use extended JSON so Decimal128 and date literals survive. --- pymongosql/executor.py | 45 ++- pymongosql/helper.py | 45 +++ pymongosql/sql/ast.py | 17 +- pymongosql/sql/builder.py | 103 +++++- pymongosql/sql/explain_builder.py | 3 +- pymongosql/sql/query_handler.py | 46 ++- pymongosql/sql/where_tree.py | 342 ++++++++++++++++-- .../sqlalchemy_mongodb/sqlalchemy_dialect.py | 9 + tests/test_parse_tree_predicates.py | 297 +++++++++++++++ tests/test_sql_grouping_filters_aliases.py | 15 +- 10 files changed, 838 insertions(+), 84 deletions(-) create mode 100644 tests/test_parse_tree_predicates.py diff --git a/pymongosql/executor.py b/pymongosql/executor.py index da36382..eff11ff 100644 --- a/pymongosql/executor.py +++ b/pymongosql/executor.py @@ -149,9 +149,8 @@ def _execute_find_plan( # Replace placeholders with parameters in filter_stage only (not in projection) filter_stage = execution_plan.filter_stage or {} - if parameters: - # Positional parameters with ? (named parameters are converted to positional in execute()) - filter_stage = self._replace_placeholders(filter_stage, parameters) + # Positional parameters (named ones are converted to positional in execute()) + filter_stage, _ = SQLHelper.bind_filter(filter_stage, parameters) projection_stage = execution_plan.projection_stage or {} @@ -223,7 +222,12 @@ def _execute_aggregate_plan( # Parse pipeline and options from JSON strings try: - pipeline = json.loads(execution_plan.aggregate_pipeline or "[]") + if execution_plan.aggregate_parameterized: + from bson import json_util + + pipeline = json_util.loads(execution_plan.aggregate_pipeline or "[]") + else: + pipeline = json.loads(execution_plan.aggregate_pipeline or "[]") options = json.loads(execution_plan.aggregate_options or "{}") except json.JSONDecodeError as e: raise ProgrammingError(f"Invalid JSON in aggregate pipeline or options: {e}") @@ -232,9 +236,9 @@ def _execute_aggregate_plan( _logger.debug(f"Pipeline: {pipeline}") _logger.debug(f"Options: {options}") - # A pipeline generated from SQL carries the WHERE clause's ? placeholders - if parameters and execution_plan.aggregate_parameterized: - pipeline = self._replace_placeholders(pipeline, parameters) + # A pipeline generated from SQL carries the WHERE clause's parameter markers + if execution_plan.aggregate_parameterized: + pipeline, _ = SQLHelper.bind_filter(pipeline, parameters) # Get collection and call aggregate() collection = db[execution_plan.collection] @@ -517,11 +521,8 @@ def _execute_execution_plan( filter_conditions = execution_plan.filter_conditions or {} - # Replace placeholders in filter if parameters provided - if parameters and filter_conditions: - filter_conditions = SQLHelper.replace_placeholders_generic( - filter_conditions, parameters, execution_plan.parameter_style - ) + # Bind the WHERE clause's parameter markers; every parameter must be used + filter_conditions, _ = SQLHelper.bind_filter(filter_conditions, parameters) command = {"delete": execution_plan.collection, "deletes": [{"q": filter_conditions, "limit": 0}]} @@ -604,12 +605,20 @@ def _execute_execution_plan( # Replace placeholders if parameters provided # Note: We need to replace both update_fields and filter_conditions in one pass # to maintain correct parameter ordering (SET clause first, then WHERE clause) - if parameters: - # Combine structures for replacement in correct order - combined = {"update_fields": update_fields, "filter_conditions": filter_conditions} - replaced = SQLHelper.replace_placeholders_generic(combined, parameters, execution_plan.parameter_style) - update_fields = replaced["update_fields"] - filter_conditions = replaced["filter_conditions"] + if isinstance(parameters, dict): + update_fields = SQLHelper.replace_placeholders_generic( + update_fields, parameters, execution_plan.parameter_style + ) + filter_conditions, _ = SQLHelper.bind_filter(filter_conditions, []) + else: + # SET values carry "?" placeholders; the WHERE clause carries parameter markers + params = list(parameters or []) + set_count = SQLHelper.count_placeholders(update_fields) + if set_count: + update_fields = SQLHelper.replace_placeholders_generic( + update_fields, params[:set_count], execution_plan.parameter_style or "qmark" + ) + filter_conditions, _ = SQLHelper.bind_filter(filter_conditions, params[set_count:]) # MongoDB update command format # https://www.mongodb.com/docs/manual/reference/command/update/ diff --git a/pymongosql/helper.py b/pymongosql/helper.py index 6337899..9ec0c38 100644 --- a/pymongosql/helper.py +++ b/pymongosql/helper.py @@ -119,6 +119,51 @@ def to_bson_value(value: Any) -> Any: return [SQLHelper.to_bson_value(v) for v in value] return value + @staticmethod + def bind_filter(value: Any, parameters: Any, exact: bool = True) -> Tuple[Any, int]: + """Bind positional parameters to the parameter markers of a translated filter. + + Translated WHERE clauses mark each ``?`` with a dict marker, so a string literal + '?' is never taken for a parameter. With ``exact``, every parameter must be used: + a parameter the translation dropped (e.g. a LIMIT it could not read) must not + silently widen the query. Returns the bound value and the number used. + """ + from .sql.where_tree import is_param + + params = [] if parameters is None else parameters + if isinstance(params, dict) or not isinstance(params, Sequence) or isinstance(params, (str, bytes)): + raise ProgrammingError("Positional parameters must be provided as a sequence") + idx = [0] + + def replace(val: Any) -> Any: + if is_param(val): + if idx[0] >= len(params): + raise ProgrammingError("Not enough parameters provided") + out = params[idx[0]] + idx[0] += 1 + return SQLHelper.to_bson_value(out) + if isinstance(val, dict): + return {k: replace(v) for k, v in val.items()} + if isinstance(val, list): + return [replace(v) for v in val] + return val + + bound = replace(value) + if exact and idx[0] != len(params): + raise ProgrammingError(f"{len(params)} parameters were given but the statement uses {idx[0]}") + return bound, idx[0] + + @staticmethod + def count_placeholders(value: Any) -> int: + """Number of legacy "?" placeholders (INSERT values, UPDATE SET) in a structure.""" + if isinstance(value, str): + return int(value == "?") + if isinstance(value, dict): + return sum(SQLHelper.count_placeholders(v) for v in value.values()) + if isinstance(value, list): + return sum(SQLHelper.count_placeholders(v) for v in value) + return 0 + @staticmethod def replace_placeholders_generic(value: Any, parameters: Any, style: Optional[str]) -> Any: """Recursively replace placeholders in nested structures for qmark or named styles.""" diff --git a/pymongosql/sql/ast.py b/pymongosql/sql/ast.py index 7a7a931..89027ea 100644 --- a/pymongosql/sql/ast.py +++ b/pymongosql/sql/ast.py @@ -303,8 +303,8 @@ def visitGroupClause(self, ctx: PartiQLParser.GroupClauseContext) -> Any: return None def visitHavingClause(self, ctx: PartiQLParser.HavingClauseContext) -> Any: - """HAVING is not translated; record it so the query fails instead of ignoring it.""" - self._query_parse_result.unsupported_clauses.append("HAVING") + """Keep the HAVING expression; it is translated after the $group stage.""" + self._query_parse_result.having = ctx.arg return None def visitLimitClause(self, ctx: PartiQLParser.LimitClauseContext) -> Any: @@ -317,8 +317,11 @@ def visitLimitClause(self, ctx: PartiQLParser.LimitClauseContext) -> Any: limit_value = int(limit_text) self._query_parse_result.limit_value = limit_value _logger.debug(f"Extracted limit value: {limit_value}") - except ValueError as e: - _logger.warning(f"Invalid LIMIT value '{limit_text}': {e}") + except ValueError: + # e.g. LIMIT ?: dropping it would return every row + self._query_parse_result.unsupported_clauses.append( + f"LIMIT {limit_text} (needs an integer literal)" + ) return self.visitChildren(ctx) except Exception as e: _logger.warning(f"Error processing LIMIT clause: {e}") @@ -334,8 +337,10 @@ def visitOffsetByClause(self, ctx: PartiQLParser.OffsetByClauseContext) -> Any: offset_value = int(offset_text) self._query_parse_result.offset_value = offset_value _logger.debug(f"Extracted offset value: {offset_value}") - except ValueError as e: - _logger.warning(f"Invalid OFFSET value '{offset_text}': {e}") + except ValueError: + self._query_parse_result.unsupported_clauses.append( + f"OFFSET {offset_text} (needs an integer literal)" + ) return self.visitChildren(ctx) except Exception as e: _logger.warning(f"Error processing OFFSET clause: {e}") diff --git a/pymongosql/sql/builder.py b/pymongosql/sql/builder.py index db58f50..1a4a07a 100644 --- a/pymongosql/sql/builder.py +++ b/pymongosql/sql/builder.py @@ -4,6 +4,8 @@ from dataclasses import dataclass from typing import TYPE_CHECKING, Any, Dict, Optional, Union +from bson import json_util + if TYPE_CHECKING: from .delete_builder import DeleteExecutionPlan from .delete_handler import DeleteParseResult @@ -159,8 +161,8 @@ def _build_query_plan(parse_result: "QueryParseResult") -> "QueryExecutionPlan": if parse_result.unsupported_clauses: raise NotSupportedError(f"Unsupported SQL clause: {', '.join(parse_result.unsupported_clauses)}") - # Auto-generate aggregate pipeline for SQL aggregate functions (COUNT, SUM, etc.) and GROUP BY - if parse_result.aggregate_functions or parse_result.group_by: + # Auto-generate aggregate pipeline for SQL aggregate functions (COUNT, SUM, etc.), GROUP BY and HAVING + if parse_result.aggregate_functions or parse_result.group_by or parse_result.having is not None: return ExecutionPlanBuilder._build_sql_aggregate_plan(parse_result) # ORDER BY may name a column by its SELECT alias; find() sorts on the field @@ -186,6 +188,72 @@ def _build_query_plan(parse_result: "QueryParseResult") -> "QueryExecutionPlan": plan = builder.build() return plan + @staticmethod + def _aggregate_source(func_info: Dict[str, Any], key: str) -> Any: + """$project expression for an accumulator: the value, or the reduced DISTINCT set.""" + if not func_info.get("distinct"): + return f"${key}" + # SQL ignores NULL in DISTINCT aggregates + values = {"$setDifference": [f"${key}", [None]]} + reducer = {"COUNT": "$size", "SUM": "$sum", "AVG": "$avg", "MIN": "$min", "MAX": "$max"} + return {reducer[func_info["function"]]: values} + + @staticmethod + def _translate_having(parse_result: "QueryParseResult", group_keys: Dict[str, str], hidden: Dict[str, Any]) -> Any: + """Translate HAVING into a $match on the grouped outputs (SQL three-valued logic).""" + from .partiql.PartiQLParser import PartiQLParser + from .query_handler import SelectHandler + from .where_tree import _PATH_NODES, WhereTreeBuilder, _field_path + + prefixes = [f"{q}." for q in (parse_result.collection_alias, parse_result.collection) if q] + + def strip(name: str) -> str: + for prefix in prefixes: + if name.startswith(prefix) and len(name) > len(prefix): + return name[len(prefix) :] + return name + + outputs = {} + for item in parse_result.select_items: + if "aggregate" in item: + info = parse_result.aggregate_functions[item["aggregate"]] + outputs[(info["function"], info["argument"], bool(info.get("distinct")))] = info["alias"] + else: + outputs[item["field"]] = item["alias"] or item["field"] + aliases = set(outputs.values()) + + def resolve(node: Any) -> Any: + if isinstance(node, (PartiQLParser.CountAllContext, PartiQLParser.AggregateBaseContext)): + kind, detail = SelectHandler._classify_item(node) + if kind != "aggregate": + raise ValueError(f"Unsupported aggregate in HAVING: {node.getText()}") + func, arg, distinct = detail + signature = (func, strip(arg) if arg != "*" else arg, distinct) + if signature in outputs: + return outputs[signature] + name = f"__having{len(hidden)}" + parse_result.aggregate_functions.append( + {"function": func, "argument": signature[1], "distinct": distinct, "alias": name, "expression": ""} + ) + hidden[name] = len(parse_result.aggregate_functions) - 1 + outputs[signature] = name + return name + if isinstance(node, _PATH_NODES): + name = strip(_field_path(node)) + if name in aliases: + return name + if name in outputs: + return outputs[name] + if name in group_keys: + hidden_name = f"__having{len(hidden)}" + hidden[hidden_name] = f"$_id.{group_keys[name]}" + outputs[name] = hidden_name + return hidden_name + raise ValueError(f"HAVING column '{name}' must be grouped, aggregated or a SELECT alias") + return None + + return WhereTreeBuilder(resolver=resolve).build(parse_result.having) + @staticmethod def _build_sql_aggregate_plan(parse_result: "QueryParseResult") -> "QueryExecutionPlan": """Build an aggregate execution plan from SQL aggregate functions and GROUP BY. @@ -214,6 +282,14 @@ def _build_sql_aggregate_plan(parse_result: "QueryParseResult") -> "QueryExecuti group_keys = {name: f"g{i}" for i, name in enumerate(parse_result.group_by)} group_stage = {"_id": {key: f"${name}" for name, key in group_keys.items()} if group_keys else None} + + # HAVING may name select-list outputs, grouped columns or aggregates; the ones + # not in the SELECT list are computed as hidden outputs and removed afterwards. + hidden: Dict[str, Any] = {} + having_filter = None + if parse_result.having is not None: + having_filter = ExecutionPlanBuilder._translate_having(parse_result, group_keys, hidden) + accumulator_keys = [] for i, func_info in enumerate(parse_result.aggregate_functions): func_name = func_info["function"] @@ -221,11 +297,14 @@ def _build_sql_aggregate_plan(parse_result: "QueryParseResult") -> "QueryExecuti accumulator = _FUNCTION_TO_ACCUMULATOR[func_name] # $group output names may not contain "." or start with "$", nor repeat key = func_info["alias"] - if "." in key or key.startswith("$") or key == "_id" or key in group_stage: + if func_info.get("distinct") or "." in key or key.startswith("$") or key == "_id" or key in group_stage: key = f"__agg{i}" accumulator_keys.append(key) - if func_name == "COUNT" and arg == "*": + if func_info.get("distinct"): + # Collect the distinct values; the $project stage reduces the set + group_stage[key] = {"$addToSet": f"${arg}"} + elif func_name == "COUNT" and arg == "*": group_stage[key] = {accumulator: 1} elif func_name == "COUNT": # COUNT(field) counts documents where the field is present and not null @@ -243,7 +322,7 @@ def _build_sql_aggregate_plan(parse_result: "QueryParseResult") -> "QueryExecuti if "aggregate" in item: func_info = parse_result.aggregate_functions[item["aggregate"]] output, key = func_info["alias"], accumulator_keys[item["aggregate"]] - source = 1 if key == output else f"${key}" + source = 1 if key == output else ExecutionPlanBuilder._aggregate_source(func_info, key) output_for[func_info["expression"].upper()] = output else: name = item["field"] @@ -254,7 +333,18 @@ def _build_sql_aggregate_plan(parse_result: "QueryParseResult") -> "QueryExecuti output_for[output] = output project_stage[output] = source outputs.append(output) + for name, source in hidden.items(): + if isinstance(source, int): # a hidden aggregate: index into aggregate_functions + project_stage[name] = ExecutionPlanBuilder._aggregate_source( + parse_result.aggregate_functions[source], accumulator_keys[source] + ) + else: + project_stage[name] = source pipeline.append({"$project": project_stage}) + if having_filter is not None: + pipeline.append({"$match": having_filter}) + if hidden: + pipeline.append({"$project": {name: 0 for name in hidden}}) sort_stage = {} for spec in parse_result.sort_fields: @@ -273,7 +363,8 @@ def _build_sql_aggregate_plan(parse_result: "QueryParseResult") -> "QueryExecuti # Configure the execution plan as an aggregate query builder._execution_plan.is_aggregate_query = True builder._execution_plan.aggregate_parameterized = True - builder._execution_plan.aggregate_pipeline = json.dumps(pipeline) + # Extended JSON keeps Decimal128/datetime literals and parameter markers intact + builder._execution_plan.aggregate_pipeline = json_util.dumps(pipeline) builder._execution_plan.aggregate_options = json.dumps({}) # Set projection for ResultSet description, in SELECT order diff --git a/pymongosql/sql/explain_builder.py b/pymongosql/sql/explain_builder.py index 204b3dd..f643080 100644 --- a/pymongosql/sql/explain_builder.py +++ b/pymongosql/sql/explain_builder.py @@ -91,8 +91,7 @@ def build_inner_command( raise ProgrammingError("No collection specified in query") filter_stage = inner_plan.filter_stage or {} - if parameters: - filter_stage = SQLHelper.replace_placeholders_generic(filter_stage, parameters, "qmark") + filter_stage, _ = SQLHelper.bind_filter(filter_stage, parameters) command = {"find": inner_plan.collection, "filter": filter_stage} diff --git a/pymongosql/sql/query_handler.py b/pymongosql/sql/query_handler.py index 5c1b2fa..9a5c25c 100644 --- a/pymongosql/sql/query_handler.py +++ b/pymongosql/sql/query_handler.py @@ -42,6 +42,8 @@ class QueryParseResult: unsupported_clauses: List[str] = field(default_factory=list) # FROM alias (FROM users AS u / FROM users u) collection_alias: Optional[str] = None + # HAVING expression (parse-tree node), translated after grouping + having: Any = None # Subquery info (for wrapped subqueries, e.g., Superset outering) subquery_plan: Optional[Any] = None @@ -139,23 +141,27 @@ def handle_visitor(self, ctx: PartiQLParser.SelectItemsContext, parse_result: "Q if hasattr(ctx, "projectionItems") and ctx.projectionItems(): for item in ctx.projectionItems().projectionItem(): field_name, alias = self._extract_field_and_alias(item) + kind, detail = self._classify_item(item) - # Check if this is an aggregate function (COUNT, SUM, etc.) - agg_match = self._AGGREGATE_PATTERN.match(field_name) - if agg_match: - func_name = agg_match.group(1).upper() - func_arg = agg_match.group(2) + if kind == "aggregate": + func_name, func_arg, distinct = detail parse_result.select_items.append({"aggregate": len(parse_result.aggregate_functions)}) parse_result.aggregate_functions.append( { "function": func_name, "argument": func_arg, + "distinct": distinct, "alias": alias or field_name, "expression": field_name, } ) continue + if kind == "unsupported": + # e.g. a + 1 or lower(a): projecting it as a field would silently read NULL + parse_result.unsupported_clauses.append(f"SELECT {detail}") + continue + field_name = detail parse_result.select_items.append({"field": field_name, "alias": alias}) # Use MongoDB standard projection format: {field: 1} to include field projection[field_name] = 1 @@ -167,6 +173,36 @@ def handle_visitor(self, ctx: PartiQLParser.SelectItemsContext, parse_result: "Q parse_result.column_aliases = column_aliases return projection + _AGGREGATES = ("COUNT", "SUM", "AVG", "MIN", "MAX") + + @staticmethod + def _classify_item(item) -> Tuple[str, Any]: + """("field", path) | ("aggregate", (function, argument, distinct)) | ("unsupported", text).""" + from .where_tree import _PATH_NODES, _field_path, _unwrap + + # A projection item's first child is its expression; other nodes are classified as-is + is_item = isinstance(item, PartiQLParser.ProjectionItemContext) + expr = item.children[0] if is_item and getattr(item, "children", None) else item + if not hasattr(expr, "getRuleIndex"): + return "unsupported", str(expr) + node = _unwrap(expr) + try: + if isinstance(node, PartiQLParser.CountAllContext): + return "aggregate", ("COUNT", "*", False) + if isinstance(node, PartiQLParser.AggregateBaseContext): + func = node.func.text.upper() + quantifier = node.setQuantifierStrategy() + argument = _unwrap(node.expr()) + if func not in SelectHandler._AGGREGATES or not isinstance(argument, _PATH_NODES): + return "unsupported", node.getText() + distinct = quantifier is not None and quantifier.getText().upper() == "DISTINCT" + return "aggregate", (func, _field_path(argument), distinct) + if isinstance(node, _PATH_NODES): + return "field", _field_path(node) + except Exception: + pass + return "unsupported", node.getText() + def _extract_field_and_alias(self, item) -> Tuple[str, Optional[str]]: """Extract field name and alias from projection item context with nested field support""" if not hasattr(item, "children") or not item.children: diff --git a/pymongosql/sql/where_tree.py b/pymongosql/sql/where_tree.py index 012b033..525462a 100644 --- a/pymongosql/sql/where_tree.py +++ b/pymongosql/sql/where_tree.py @@ -6,10 +6,18 @@ UNKNOWN, which is in neither set. ``NOT p`` swaps the two sets, and De Morgan's laws combine them for AND and OR, so ``NOT a = 1`` excludes documents where ``a`` is NULL or missing, exactly like SQL. + +Predicates are read from the parse tree, never from concatenated token text: the +field is the path node on one side of the operator and the value is the literal, +parameter or value function on the other. Anything else raises +``NotSupportedError`` rather than matching different documents. """ +import datetime import re -from typing import Any, Dict, List, Tuple +from typing import Any, Dict, List, Optional, Tuple + +from bson import Decimal128 from ..error import NotSupportedError from .partiql.PartiQLParser import PartiQLParser @@ -20,6 +28,19 @@ # Matches no document; used for predicates that can never be TRUE (or FALSE). NOTHING: Filter = {"$expr": False} +# Stands for a bound parameter (``?``) inside a translated filter. A dict cannot be +# produced by any SQL literal, so a string literal '?' is never mistaken for it. +PARAM_KEY = "$pymongosqlParam" + + +def param_marker() -> Dict[str, Any]: + return {PARAM_KEY: True} + + +def is_param(value: Any) -> bool: + return isinstance(value, dict) and list(value) == [PARAM_KEY] + + _LEAVES = ( PartiQLParser.PredicateComparisonContext, PartiQLParser.PredicateIsContext, @@ -27,12 +48,60 @@ PartiQLParser.PredicateLikeContext, PartiQLParser.PredicateBetweenContext, ) -_FIELD_PATH = re.compile(r'^(?:"[^"]+"|[A-Za-z_$][\w$]*)(?:\.(?:"[^"]+"|[A-Za-z_$][\w$]*|\d+))*$') -_FIELD = re.compile(r"^[A-Za-z_$][\w$]*(?:\.[\w$]+)*$") +_PATH_NODES = ( + PartiQLParser.VariableIdentifierContext, + PartiQLParser.VariableKeywordContext, + PartiQLParser.ExprPrimaryPathContext, +) _SWAP = {"<": ">=", ">=": "<", ">": "<=", "<=": ">"} +_MIRROR = {"<": ">", ">": "<", "<=": ">=", ">=": "<=", "=": "=", "!=": "!=", "<>": "<>"} _MONGO = {"<": "$lt", "<=": "$lte", ">": "$gt", ">=": "$gte"} +class _Field(str): + """A document field path (as opposed to a string value).""" + + +class _InexactDecimal: + """A decimal literal that no double represents exactly (e.g. 0.1). + + SQL compares a literal in the column's type. A double field is compared with the + nearest double and every other numeric type with the exact Decimal128, so neither + a double 0.1 nor a Decimal128 0.1 is missed. + """ + + def __init__(self, text: str): + self.as_double = float(text) + self.as_decimal = Decimal128(text) + + def __neg__(self) -> "_InexactDecimal": + negated = _InexactDecimal("0") + negated.as_double, negated.as_decimal = -self.as_double, Decimal128(-self.as_decimal.to_decimal()) + return negated + + +def _decimal_literal(text: str) -> Any: + from decimal import Decimal + + exact = Decimal(text) + as_double = float(text) + return as_double if Decimal(as_double) == exact else _InexactDecimal(text) + + +def _coerce(value: Any, as_double: bool) -> Any: + if isinstance(value, _InexactDecimal): + return value.as_double if as_double else value.as_decimal + if isinstance(value, (list, tuple)): + return type(value)(_coerce(v, as_double) for v in value) + return value + + +def _has_inexact(value: Any) -> bool: + if isinstance(value, (list, tuple)): + return any(_has_inexact(v) for v in value) + return isinstance(value, _InexactDecimal) + + def _all(parts: List[Tuple[Any, Filter]], key: str, chain: type) -> Filter: """Combine filters under ``key``, flattening only an unparenthesized chain of the same operator.""" items: List[Filter] = [] @@ -48,13 +117,148 @@ def contains_not(ctx: Any) -> bool: return any(contains_not(child) for child in getattr(ctx, "children", None) or []) -def leaf_filters(field: str, operator: str, value: Any) -> Pair: +def _unwrap(ctx: Any) -> Any: + """Descend through single-child pass-through rules (MathOp00 > ... > ExprTermBase).""" + while True: + children = getattr(ctx, "children", None) or [] + if len(children) == 1 and hasattr(children[0], "getRuleIndex"): + ctx = children[0] + else: + return ctx + + +def _string_literal(token_text: str) -> str: + return token_text[1:-1].replace("''", "'") + + +def _field_path(ctx: Any) -> str: + """Dot path for a variable reference or path expression, quotes removed.""" + if isinstance(ctx, (PartiQLParser.VariableIdentifierContext, PartiQLParser.VariableKeywordContext)): + if getattr(ctx, "qualifier", None) is not None: + raise NotSupportedError(f"Unsupported variable reference: {ctx.getText()}") + text = ctx.getText() + return text[1:-1].replace('""', '"') if text.startswith('"') else text + parts = [_field_path(_unwrap(ctx.getChild(0)))] + for step in ctx.children[1:]: + if isinstance(step, PartiQLParser.PathStepDotExprContext): + key = step.key.getText() + parts.append(key[1:-1].replace('""', '"') if key.startswith('"') else key) + elif isinstance(step, PartiQLParser.PathStepIndexExprContext): + key = _unwrap(step.key) + if isinstance(key, PartiQLParser.LiteralIntegerContext): + parts.append(key.getText()) + elif isinstance(key, PartiQLParser.LiteralStringContext): + parts.append(_string_literal(key.getText())) + else: + raise NotSupportedError(f"Unsupported path step: {step.getText()}") + else: + raise NotSupportedError(f"Unsupported path step: {step.getText()}") + path = ".".join(parts) + # A quoted identifier with dots ("user.name") keeps this driver's nested-path meaning + if any(not segment or segment.startswith("$") for segment in path.split(".")): + raise NotSupportedError(f"Unsupported field name: {path!r}") + return path + + +def operand(ctx: Any, resolver: Any = None) -> Any: + """Evaluate one side of a predicate: a _Field, a Python value or a parameter marker. + + ``resolver`` may map a node (e.g. an aggregate call in HAVING) to a field name. + """ + node = _unwrap(ctx) + if resolver is not None: + resolved = resolver(node) + if resolved is not None: + return _Field(resolved) + if isinstance(node, _PATH_NODES): + return _Field(_field_path(node)) + if isinstance(node, PartiQLParser.ParameterContext): + return param_marker() + if isinstance(node, PartiQLParser.LiteralStringContext): + return _string_literal(node.getText()) + if isinstance(node, PartiQLParser.LiteralIntegerContext): + return int(node.getText()) + if isinstance(node, PartiQLParser.LiteralDecimalContext): + return _decimal_literal(node.getText()) + if isinstance(node, PartiQLParser.LiteralTrueContext): + return True + if isinstance(node, PartiQLParser.LiteralFalseContext): + return False + if isinstance(node, (PartiQLParser.LiteralNullContext, PartiQLParser.LiteralMissingContext)): + return None + if isinstance(node, PartiQLParser.LiteralDateContext): + return datetime.datetime.fromisoformat(_string_literal(node.LITERAL_STRING().getText())) + if isinstance(node, PartiQLParser.ValueExprContext) and node.sign is not None: + value = operand(node.rhs) + if isinstance(value, bool) or not isinstance(value, (int, float, _InexactDecimal)): + raise NotSupportedError(f"Unsupported signed expression: {node.getText()}") + return value if node.sign.text == "+" else -value + if isinstance(node, PartiQLParser.MathOp00Context) and node.op is not None and node.op.text == "||": + left, right = operand(node.lhs), operand(node.rhs) + if ( + not (isinstance(left, str) and isinstance(right, str)) + or isinstance(left, _Field) + or isinstance(right, _Field) + ): + raise NotSupportedError(f"Only string literals can be concatenated: {node.getText()}") + return left + right + if isinstance(node, PartiQLParser.FunctionCallContext): + from .value_function_registry import get_default_registry + + name = node.functionName().getText() + registry = get_default_registry() + if registry.has_function(name): + args = [operand(arg) for arg in node.expr()] + if any(isinstance(a, _Field) or is_param(a) for a in args): + raise NotSupportedError(f"Value functions take literal arguments: {node.getText()}") + return registry.execute(name, args) + raise NotSupportedError(f"Unsupported WHERE operand: {node.getText()}") + + +def _like_regex(pattern: str, escape: Optional[str]) -> Tuple[str, bool]: + """Regex for a LIKE pattern and whether it needs DOTALL. + + ``escape`` makes the next character literal. A leading or trailing ``%`` leaves + that end unanchored; a wildcard anywhere else must also match newlines. + """ + tokens, i = [], 0 + while i < len(pattern): + char = pattern[i] + if escape is not None and char == escape: + if i + 1 >= len(pattern): + raise NotSupportedError("LIKE pattern ends with its escape character") + tokens.append(re.escape(pattern[i + 1])) + i += 2 + continue + tokens.append(".*" if char == "%" else "." if char == "_" else re.escape(char)) + i += 1 + body = "".join(tokens) + inner = tokens[1 if tokens[:1] == [".*"] else 0 : len(tokens) - (1 if tokens[-1:] == [".*"] else 0)] + needs_dotall = any(t in (".*", ".") for t in inner) + regex = ("" if body.startswith(".*") else "^") + body + ("" if body.endswith(".*") else "$") + return regex, needs_dotall + + +def leaf_filters(field: str, operator: str, value: Any, escape: Optional[str] = None) -> Pair: """TRUE and FALSE filters for one predicate on ``field``.""" + if _has_inexact(value): + double = {field: {"$type": "double"}} + other = {field: {"$not": {"$type": "double"}}} + t1, f1 = leaf_filters(field, operator, _coerce(value, True), escape) + t2, f2 = leaf_filters(field, operator, _coerce(value, False), escape) + return ( + {"$or": [{"$and": [double, t1]}, {"$and": [other, t2]}]}, + {"$or": [{"$and": [double, f1]}, {"$and": [other, f2]}]}, + ) op = operator.upper() if op == "IS NULL": return {field: {"$eq": None}}, {field: {"$ne": None}} if op == "IS NOT NULL": return {field: {"$ne": None}}, {field: {"$eq": None}} + if op == "IS MISSING": + return {field: {"$exists": False}}, {field: {"$exists": True}} + if op == "IS NOT MISSING": + return {field: {"$exists": True}}, {field: {"$exists": False}} if op in ("IN", "NOT IN"): values = value if isinstance(value, list) else [value] present = [v for v in values if v is not None] @@ -63,21 +267,19 @@ def leaf_filters(field: str, operator: str, value: Any) -> Pair: false = NOTHING if None in values else {field: {"$nin": present + [None]}} return (true, false) if op == "IN" else (false, true) if op in ("LIKE", "NOT LIKE"): - from .handler import ComparisonExpressionHandler - - if value == "?" or not isinstance(value, str): - # The pattern is translated to a regex while parsing, before parameters are bound + if is_param(value) or not isinstance(value, str): + # The pattern becomes a regex while parsing, before parameters are bound raise NotSupportedError("LIKE needs a literal pattern, not a bound parameter") - - pattern = ComparisonExpressionHandler._like_to_regex(value) - pattern = ("" if pattern.startswith(".*") else "^") + pattern + ("" if pattern.endswith(".*") else "$") - true = {field: {"$regex": pattern}} - false = {"$and": [{field: {"$not": {"$regex": pattern}}}, {field: {"$ne": None}}]} + pattern, dotall = _like_regex(value, escape) + regex = {"$regex": pattern, "$options": "s"} if dotall else {"$regex": pattern} + true = {field: regex} + false = {"$and": [{field: {"$not": regex}}, {field: {"$ne": None}}]} return (true, false) if op == "LIKE" else (false, true) - if op == "BETWEEN": + if op in ("BETWEEN", "NOT BETWEEN"): low, high = value true = {"$and": [{field: {"$gte": low}}, {field: {"$lte": high}}]} - return true, {"$or": [{field: {"$lt": low}}, {field: {"$gt": high}}]} + false = {"$or": [{field: {"$lt": low}}, {field: {"$gt": high}}]} + return (true, false) if op == "BETWEEN" else (false, true) if value is None: # This dialect has always read "= NULL" / "<> NULL" as IS NULL / IS NOT NULL; # any other comparison with NULL is UNKNOWN. @@ -95,19 +297,31 @@ def leaf_filters(field: str, operator: str, value: Any) -> Pair: raise NotSupportedError(f"Unsupported predicate operator: {operator}") +def _value(ctx: Any, resolver: Any = None) -> Any: + value = operand(ctx, resolver) + if isinstance(value, _Field): + raise NotSupportedError(f"Comparing two fields is not supported: {ctx.getText()}") + return value + + class WhereTreeBuilder: - """Build a MongoDB filter for a WHERE expression from its parse tree.""" + """Build a MongoDB filter for a WHERE (or HAVING) expression from its parse tree.""" + + def __init__(self, resolver: Any = None): + self._resolver = resolver def build(self, ctx: Any) -> Filter: return self._pair(ctx)[0] - def _pair(self, ctx: Any) -> Pair: + def _pair(self, ctx: Any, substitute: Optional[Tuple[Any, Pair]] = None) -> Pair: + if substitute is not None and ctx is substitute[0]: + return substitute[1] if isinstance(ctx, PartiQLParser.NotContext): - true, false = self._pair(ctx.rhs) + true, false = self._pair(ctx.rhs, substitute) return false, true if isinstance(ctx, (PartiQLParser.AndContext, PartiQLParser.OrContext)): is_and = isinstance(ctx, PartiQLParser.AndContext) - (t1, f1), (t2, f2) = self._pair(ctx.lhs), self._pair(ctx.rhs) + (t1, f1), (t2, f2) = self._pair(ctx.lhs, substitute), self._pair(ctx.rhs, substitute) true_key, false_key = ("$and", "$or") if is_and else ("$or", "$and") chain = type(ctx) return ( @@ -115,29 +329,79 @@ def _pair(self, ctx: Any) -> Pair: _all([(ctx.lhs, f1), (ctx.rhs, f2)], false_key, chain), ) if isinstance(ctx, PartiQLParser.ExprTermWrappedQueryContext): - return self._pair(ctx.expr()) + return self._pair(ctx.expr(), substitute) + if isinstance(ctx, PartiQLParser.PredicateLikeContext): + return self._like(ctx) if isinstance(ctx, _LEAVES): return self._leaf(ctx) children = [c for c in getattr(ctx, "children", None) or [] if hasattr(c, "getRuleIndex")] if len(children) == 1 and len(ctx.children) == 1: - return self._pair(children[0]) - text = ctx.getText() - if _FIELD_PATH.match(text) and text.upper() not in ("TRUE", "FALSE", "NULL"): + return self._pair(children[0], substitute) + if isinstance(ctx, _PATH_NODES): # A bare boolean field: WHERE flag / WHERE NOT flag - from .handler import ContextUtilsMixin - - field = ContextUtilsMixin.normalize_field_path(text) + field = str(operand(ctx, self._resolver)) return {field: True}, {field: False} - raise NotSupportedError(f"Unsupported WHERE expression: {text}") - - @staticmethod - def _leaf(ctx: Any) -> Pair: - from .handler import ComparisonExpressionHandler - - handler = ComparisonExpressionHandler() - field = handler._extract_field_name(ctx) - operator = handler._extract_operator(ctx) - value = handler._extract_value(ctx) - if not _FIELD.match(field): - raise NotSupportedError(f"Unsupported WHERE predicate: {ctx.getText()}") - return leaf_filters(field, operator, value) + if isinstance(ctx, (PartiQLParser.LiteralTrueContext, PartiQLParser.LiteralFalseContext)): + everything: Filter = {} + return (everything, NOTHING) if isinstance(ctx, PartiQLParser.LiteralTrueContext) else (NOTHING, everything) + raise NotSupportedError(f"Unsupported WHERE expression: {ctx.getText()}") + + def _field_and_value(self, lhs: Any, rhs: Any, text: str) -> Tuple[str, Any, bool]: + """(field, value, mirrored) for ``lhs op rhs`` with the field on either side.""" + left, right = operand(lhs, self._resolver), operand(rhs, self._resolver) + if isinstance(left, _Field) and not isinstance(right, _Field): + return str(left), right, False + if isinstance(right, _Field) and not isinstance(left, _Field): + return str(right), left, True + raise NotSupportedError(f"A predicate needs exactly one field and one value: {text}") + + def _leaf(self, ctx: Any) -> Pair: + text = ctx.getText() + if isinstance(ctx, PartiQLParser.PredicateComparisonContext): + field, value, mirrored = self._field_and_value(ctx.lhs, ctx.rhs, text) + op = ctx.op.text + return leaf_filters(field, _MIRROR[op] if mirrored else op, value) + negated = ctx.NOT() is not None + field = operand(ctx.lhs, self._resolver) + if not isinstance(field, _Field): + raise NotSupportedError(f"The left side must be a field: {text}") + if isinstance(ctx, PartiQLParser.PredicateIsContext): + kind = ctx.type_().getText().upper() + if kind not in ("NULL", "MISSING"): + raise NotSupportedError(f"Unsupported IS type test: {text}") + return leaf_filters(field, f"IS {'NOT ' if negated else ''}{kind}", None) + if isinstance(ctx, PartiQLParser.PredicateInContext): + target = _unwrap(ctx.rhs) if ctx.rhs is not None else None + if isinstance(target, PartiQLParser.ValueListContext): + values = [_value(item, self._resolver) for item in target.expr()] + elif ctx.expr() is not None and ctx.rhs is None: + values = [_value(ctx.expr(), self._resolver)] # IN (single value) + else: + raise NotSupportedError(f"Unsupported IN list: {text}") + return leaf_filters(field, "NOT IN" if negated else "IN", values) + if isinstance(ctx, PartiQLParser.PredicateBetweenContext): + bounds = (_value(ctx.lower, self._resolver), _value(ctx.upper, self._resolver)) + return leaf_filters(field, "NOT BETWEEN" if negated else "BETWEEN", bounds) + raise NotSupportedError(f"Unsupported predicate: {text}") + + def _like(self, ctx: Any) -> Pair: + field = operand(ctx.lhs, self._resolver) + if not isinstance(field, _Field): + raise NotSupportedError(f"The left side of LIKE must be a field: {ctx.getText()}") + pattern = _value(ctx.rhs, self._resolver) + op = "NOT LIKE" if ctx.NOT() is not None else "LIKE" + if ctx.escape is None: + return leaf_filters(field, op, pattern) + # The grammar lets ESCAPE take a whole expression, so "a LIKE p ESCAPE '/' AND b = 1" + # parses the AND chain as the escape. Its leftmost operand is the escape character; + # the LIKE predicate takes that operand's place in the chain. + leftmost = _unwrap(ctx.escape) + while isinstance(leftmost, (PartiQLParser.AndContext, PartiQLParser.OrContext)): + leftmost = _unwrap(leftmost.lhs) + escape = _value(leftmost, self._resolver) + if not isinstance(escape, str) or len(escape) != 1: + raise NotSupportedError(f"LIKE ESCAPE must be a single character: {ctx.getText()}") + like = leaf_filters(field, op, pattern, escape) + if leftmost is _unwrap(ctx.escape): + return like + return self._pair(ctx.escape, substitute=(leftmost, like)) diff --git a/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py b/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py index 4b40fa6..64363db 100644 --- a/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py +++ b/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py @@ -111,6 +111,15 @@ def visit_column(self, column, include_table=True, **kwargs): """ return super().visit_column(column, include_table=False, **kwargs) + def limit_clause(self, select, **kw): + """Render LIMIT/OFFSET as integer literals: PyMongoSQL reads them while parsing.""" + text = "" + if select._limit_clause is not None: + text += "\n LIMIT " + self.process(select._limit_clause, literal_binds=True, **kw) + if select._offset_clause is not None: + text += "\n OFFSET " + self.process(select._offset_clause, literal_binds=True, **kw) + return text + def visit_like_op_binary(self, binary, operator, **kw): """Render LIKE patterns inline: PyMongoSQL turns them into a regex while parsing.""" kw["literal_binds"] = True diff --git a/tests/test_parse_tree_predicates.py b/tests/test_parse_tree_predicates.py new file mode 100644 index 0000000..3425a88 --- /dev/null +++ b/tests/test_parse_tree_predicates.py @@ -0,0 +1,297 @@ +# -*- coding: utf-8 -*- +"""Predicates are read from the parse tree: field names, operators and values are never +recovered by searching concatenated token text. + +The DML tests check the exact documents a DELETE or UPDATE touches: a truncated +field name there removes or rewrites the wrong documents. +""" + +from decimal import Decimal + +import pytest +from bson import Decimal128, Int64 + +from pymongosql.error import Error +from pymongosql.sql.parser import SQLParser +from tests.conftest import HAS_SQLALCHEMY + +if HAS_SQLALCHEMY: + import sqlalchemy as sa + + from pymongosql.sqlalchemy_mongodb.sqlalchemy_dialect import PyMongoSQLDialect + +needs_sqlalchemy = pytest.mark.skipif(not HAS_SQLALCHEMY, reason="SQLAlchemy not available") +PARAM = {"$pymongosqlParam": True} + + +def plan(sql): + return SQLParser(sql).get_execution_plan() + + +def where(sql_where): + return plan(f"SELECT _id FROM t WHERE {sql_where}").filter_stage + + +def compiled(stmt): + return " ".join(str(stmt.compile(dialect=PyMongoSQLDialect())).split()) + + +class TestFieldNamesAreNotTruncated: + @pytest.mark.parametrize( + "sql_where,expected", + [ + ("dislikes IS NULL", {"dislikes": {"$eq": None}}), + ("unlike = 5", {"unlike": 5}), + ("likes IN (1, 2)", {"likes": {"$in": [1, 2]}}), + ("isnull = 1", {"isnull": 1}), + ("between_x BETWEEN 1 AND 2", {"$and": [{"between_x": {"$gte": 1}}, {"between_x": {"$lte": 2}}]}), + ("login_count > 3", {"login_count": {"$gt": 3}}), + ("inx <> 'a'", {"inx": {"$nin": ["a", None]}}), + ], + ) + def test_select(self, sql_where, expected): + assert where(sql_where) == expected + + def test_delete_and_update(self): + assert plan("DELETE FROM t WHERE dislikes IS NULL").filter_conditions == {"dislikes": {"$eq": None}} + update = plan("UPDATE t SET x = 1 WHERE unlike = 5") + assert update.filter_conditions == {"unlike": 5} + + def test_string_literal_containing_keywords(self): + assert where("name = 'x IN(1) LIKE y BETWEEN'") == {"name": "x IN(1) LIKE y BETWEEN"} + + def test_reversed_operands(self): + assert where("5 < age") == {"age": {"$gt": 5}} + assert where("'x' = name") == {"name": "x"} + assert where("10 >= age") == {"age": {"$lte": 10}} + + @pytest.mark.parametrize( + "sql_where,field", + [('"my field" = 1', "my field"), ('"a-b" = 1', "a-b"), ('"año" = 1', "año"), ('"a"."b c" = 1', "a.b c")], + ) + def test_quoted_field_names(self, sql_where, field): + assert where(sql_where) == {field: 1} + + @pytest.mark.parametrize("sql_where", ["a = b", "lower(a) = 'x'", "a + 1 = 2", "a IN (SELECT b FROM u)"]) + def test_untranslatable_predicates_raise(self, sql_where): + with pytest.raises(Error): + where(sql_where) + + +class TestLike: + def test_concatenated_pattern(self): + assert where("n LIKE '%' || 'ab' || '%'") == {"n": {"$regex": ".*ab.*"}} + assert where("n LIKE 'ab' || '%'") == {"n": {"$regex": "^ab.*"}} + + def test_escape(self): + assert where("n LIKE '50/%%' ESCAPE '/'") == {"n": {"$regex": "^50%.*"}} + assert where("n LIKE 'a/_b' ESCAPE '/'") == {"n": {"$regex": "^a_b$"}} + + def test_escape_followed_by_more_conditions(self): + assert where("n LIKE 'a/_%' ESCAPE '/' AND b = 1 OR c = 2") == { + "$or": [{"$and": [{"n": {"$regex": "^a_.*"}}, {"b": 1}]}, {"c": 2}] + } + + def test_inner_wildcards_match_newlines(self): + assert where("n LIKE 'a%b'") == {"n": {"$regex": "^a.*b$", "$options": "s"}} + + def test_bound_pattern_raises(self): + with pytest.raises(Error): + where("n LIKE ?") + + @needs_sqlalchemy + def test_sqlalchemy_helpers_render_translatable_patterns(self): + t = sa.table("t", sa.column("n")) + for expr, regex in ( + (t.c.n.contains("ab"), ".*ab.*"), + (t.c.n.startswith("ab"), "^ab.*"), + (t.c.n.endswith("ab"), ".*ab$"), + (t.c.n.contains("5%_", autoescape=True), ".*5%_.*"), + (t.c.n.like("a/_%", escape="/"), "^a_.*"), + ): + sql = compiled(sa.select(t.c.n).where(expr)) + assert plan(sql).filter_stage == {"n": {"$regex": regex}}, sql + + +class TestParameters: + def test_literal_question_mark_is_a_value(self): + assert where("a = '?' AND b = ?") == {"$and": [{"a": "?"}, {"b": PARAM}]} + + def test_literal_question_mark_in_aggregate(self): + p = plan("SELECT g, COUNT(*) AS n FROM t WHERE a = '?' AND b = ? GROUP BY g") + assert '"a": "?"' in p.aggregate_pipeline and '"$pymongosqlParam"' in p.aggregate_pipeline + + def test_limit_parameter_raises_instead_of_returning_everything(self): + with pytest.raises(Error): + plan("SELECT _id FROM t LIMIT ?") + + @needs_sqlalchemy + def test_sqlalchemy_limit_and_offset_are_literals(self): + t = sa.table("t", sa.column("a")) + assert compiled(sa.select(t.c.a).limit(5).offset(10)) == "SELECT a FROM t LIMIT 5 OFFSET 10" + assert compiled(sa.select(t.c.a).offset(3)) == "SELECT a FROM t OFFSET 3" + + +class TestDecimalLiterals: + def test_exact_decimal_literal_is_a_double(self): + assert where("t > 36.5") == {"t": {"$gt": 36.5}} + + def test_inexact_decimal_literal_compares_in_the_field_type(self): + assert where("a = 0.1") == { + "$or": [ + {"$and": [{"a": {"$type": "double"}}, {"a": 0.1}]}, + {"$and": [{"a": {"$not": {"$type": "double"}}}, {"a": Decimal128("0.1")}]}, + ] + } + + +class TestAggregates: + def test_count_distinct(self): + p = plan('SELECT g, COUNT(DISTINCT a) AS "COUNT_DISTINCT(a)" FROM t GROUP BY g') + assert '"$addToSet": "$a"' in p.aggregate_pipeline + assert '"COUNT_DISTINCT(a)": {"$size": {"$setDifference": ["$__agg0", [null]]}}' in p.aggregate_pipeline + + @pytest.mark.parametrize("sql", ["SELECT a + 1 FROM t", "SELECT lower(a) FROM t", "SELECT EVERY(a) FROM t"]) + def test_untranslatable_select_expressions_raise(self, sql): + with pytest.raises(Error): + plan(sql) + + def test_having(self): + p = plan("SELECT g, COUNT(*) AS n FROM t GROUP BY g HAVING n > 1 AND SUM(v) >= 3") + assert '{"$match": {"$and": [{"n": {"$gt": 1}}, {"__having0": {"$gte": 3}}]}}' in p.aggregate_pipeline + assert '{"$project": {"__having0": 0}}' in p.aggregate_pipeline + + +@needs_sqlalchemy +class TestFloatColumns: + """Float returns float; Float(asdecimal=True) returns Decimal as SQLAlchemy documents.""" + + def processor(self, type_): + dialect = PyMongoSQLDialect() + return type_.dialect_impl(dialect).result_processor(dialect, None) + + def test_float_returns_float(self): + for type_ in (sa.Float(), sa.FLOAT(), sa.REAL()): + value = self.processor(type_)(Decimal128("1.25")) + assert value == 1.25 and type(value) is float + + def test_float_asdecimal_returns_decimal(self): + value = self.processor(sa.Float(asdecimal=True))(Decimal128("1.25")) + assert value == Decimal("1.25") and type(value) is Decimal + + +# ---------------------------------------------------------------- live MongoDB + +COLLECTION = "test_parse_tree_predicates" +DOCS = [ + {"_id": 1, "dislikes": None, "dis": 1, "unlike": 5, "un": "x", "n": "ab-50%", "g": "a", "v": 1, "f": 0.1}, + {"_id": 2, "dislikes": 3, "unlike": 6, "un": "=5", "n": "zab", "g": "a", "v": 2, "f": 0.2}, + {"_id": 3, "dislikes": 4, "dis": 2, "unlike": 5, "n": "a_b", "g": "b", "v": 2, "f": Decimal128("0.1")}, + {"_id": 4, "my field": "x", "año": 1, "n": "?", "g": "b", "v": None, "f": None}, +] + + +@pytest.fixture +def docs(conn): + conn.database.drop_collection(COLLECTION) + conn.database[COLLECTION].insert_many([dict(d) for d in DOCS]) + yield conn + conn.database.drop_collection(COLLECTION) + + +def run(conn, sql, params=None): + cursor = conn.cursor() + cursor.execute(sql, params) if params is not None else cursor.execute(sql) + return cursor + + +def ids(conn, sql_where, params=None): + return [r[0] for r in run(conn, f"SELECT _id FROM {COLLECTION} WHERE {sql_where} ORDER BY _id", params).fetchall()] + + +def remaining(conn): + return [d["_id"] for d in conn.database[COLLECTION].find({}, sort=[("_id", 1)])] + + +class TestLiveDML: + def test_delete_is_null_touches_only_null_dislikes(self, docs): + # dislikes IS NULL: _id 1 (null) and 4 (missing); the old translation matched + # every document lacking a field named "dis" (1 was spared, 2 and 4 deleted) + run(docs, f"DELETE FROM {COLLECTION} WHERE dislikes IS NULL") + assert remaining(docs) == [2, 3] + + def test_update_touches_only_matching_unlike(self, docs): + run(docs, f"UPDATE {COLLECTION} SET v = 99 WHERE unlike = 5") + assert {d["_id"]: d.get("v") for d in docs.database[COLLECTION].find()} == {1: 99, 2: 2, 3: 99, 4: None} + + def test_delete_with_parameters_and_literal_question_mark(self, docs): + run(docs, f"DELETE FROM {COLLECTION} WHERE n = '?' OR unlike = ?", [6]) + assert remaining(docs) == [1, 3] + + def test_delete_reversed_operand(self, docs): + run(docs, f"DELETE FROM {COLLECTION} WHERE 6 <= unlike") + assert remaining(docs) == [1, 3, 4] + + def test_unused_parameter_refuses_the_delete(self, docs): + with pytest.raises(Error): + run(docs, f"DELETE FROM {COLLECTION} WHERE unlike = ?", [5, 6]) + assert remaining(docs) == [1, 2, 3, 4] + + def test_untranslatable_update_changes_nothing(self, docs): + with pytest.raises(Error): + run(docs, f"UPDATE {COLLECTION} SET v = 0 WHERE lower(n) = 'ab'") + assert [d.get("v") for d in docs.database[COLLECTION].find({}, sort=[("_id", 1)])] == [1, 2, 2, None] + + +class TestLiveSelect: + @pytest.mark.parametrize( + "sql_where,expected", + [ + ("dislikes IS NULL", [1, 4]), + ("unlike = 5", [1, 3]), + ("5 < unlike", [2]), + ("\"my field\" = 'x'", [4]), + ('"año" = 1', [4]), + ("n LIKE '%' || 'ab' || '%'", [1, 2]), + ("n LIKE '%/%' ESCAPE '/'", [1]), + ("n LIKE 'a/_b' ESCAPE '/' AND g = 'b'", [3]), + ("n = '?'", [4]), + ("f = 0.1", [1, 3]), + ], + ) + def test_rows(self, docs, sql_where, expected): + assert ids(docs, sql_where) == expected + + def test_count_distinct_and_having(self, docs): + rows = run( + docs, + f"SELECT g, COUNT(DISTINCT v) AS d, COUNT(*) AS n FROM {COLLECTION} GROUP BY g HAVING SUM(v) >= 2 ORDER BY g", + ).fetchall() + assert [tuple(r) for r in rows] == [("a", 2, 2), ("b", 1, 2)] + rows = run(docs, f"SELECT g FROM {COLLECTION} GROUP BY g HAVING COUNT(*) > 1 AND MAX(v) > ?", [1]).fetchall() + assert sorted(tuple(r) for r in rows) == [("a",), ("b",)] + + def test_decimal128_and_int64_are_python_types(self, docs): + docs.database[COLLECTION].insert_one({"_id": 5, "d": Decimal128("1.10"), "i": Int64(2**40)}) + row = run(docs, f"SELECT d, i FROM {COLLECTION} WHERE _id = 5").fetchone() + assert row == (Decimal("1.10"), 2**40) + assert type(row[0]) is Decimal and type(row[1]) is int + + +@needs_sqlalchemy +class TestLiveSQLAlchemy: + def test_like_helpers_limit_offset_and_count_distinct(self, sqlalchemy_engine, docs): + t = sa.table(COLLECTION, sa.column("_id"), sa.column("n"), sa.column("g"), sa.column("v")) + with sqlalchemy_engine.connect() as c: + assert list(c.execute(sa.select(t.c._id).where(t.c.n.contains("ab")).order_by(t.c._id)).scalars()) == [1, 2] + assert list(c.execute(sa.select(t.c._id).where(t.c.n.startswith("a")).order_by(t.c._id)).scalars()) == [ + 1, + 3, + ] + assert list(c.execute(sa.select(t.c._id).where(t.c.n.endswith("b")).order_by(t.c._id)).scalars()) == [2, 3] + assert list(c.execute(sa.select(t.c._id).where(t.c.n.contains("50%", autoescape=True))).scalars()) == [1] + page = c.execute(sa.select(t.c._id).order_by(t.c._id).limit(2).offset(1)).scalars() + assert list(page) == [2, 3] + distinct = sa.func.count(sa.distinct(t.c.v)).label("d") + rows = c.execute(sa.select(t.c.g, distinct).group_by(t.c.g).order_by(t.c.g)).fetchall() + assert [tuple(r) for r in rows] == [("a", 2), ("b", 1)] diff --git a/tests/test_sql_grouping_filters_aliases.py b/tests/test_sql_grouping_filters_aliases.py index 6590a76..1f1e147 100644 --- a/tests/test_sql_grouping_filters_aliases.py +++ b/tests/test_sql_grouping_filters_aliases.py @@ -5,7 +5,6 @@ import pytest -from pymongosql.error import Error from pymongosql.sql.parser import SQLParser from tests.conftest import HAS_SQLALCHEMY, make_superset_conn @@ -54,13 +53,13 @@ def test_ungrouped_column_is_rejected(self): with pytest.raises(Exception, match="must appear in GROUP BY"): plan("SELECT flag, COUNT(*) FROM t") - def test_having_is_rejected_not_ignored(self): - with pytest.raises(Exception, match="HAVING"): - plan("SELECT flag, COUNT(*) AS n FROM t GROUP BY flag HAVING COUNT(*) > 1") + def test_having_filters_groups(self): + stages = pipeline("SELECT flag, COUNT(*) AS n FROM t GROUP BY flag HAVING COUNT(*) > 1") + assert stages[2] == {"$match": {"n": {"$gt": 1}}} def test_in_keeps_literal_types_and_quoted_commas(self): p = plan("SELECT _id FROM t WHERE _id IN (1, 2.5, 'a,b', 'it''s', TRUE, ?)") - assert p.filter_stage == {"_id": {"$in": [1, 2.5, "a,b", "it's", True, "?"]}} + assert p.filter_stage == {"_id": {"$in": [1, 2.5, "a,b", "it's", True, {"$pymongosqlParam": True}]}} def test_not_in(self): assert plan("SELECT _id FROM t WHERE _id NOT IN (1, 2)").filter_stage == {"_id": {"$nin": [1, 2, None]}} @@ -160,9 +159,9 @@ def test_superset_mode_physical_table_chart_query(self, grouping_collection): finally: conn.close() - def test_unsupported_having_raises(self, grouping_collection): - with pytest.raises(Error): - rows(grouping_collection, f"SELECT flag, COUNT(*) AS n FROM {COLLECTION} GROUP BY flag HAVING n > 1") + def test_having_returns_only_matching_groups(self, grouping_collection): + got = rows(grouping_collection, f"SELECT flag, COUNT(*) AS n FROM {COLLECTION} GROUP BY flag HAVING n > 1") + assert got == [(True, 3)] @pytest.mark.skipif(not HAS_SQLALCHEMY, reason="SQLAlchemy not available") From 5c832294e9760c6a603e7f84deb9c334f8054a51 Mon Sep 17 00:00:00 2001 From: Amin Ghadersohi Date: Sat, 26 Sep 2026 10:17:13 +0000 Subject: [PATCH 13/19] fix: bind SET values, LIMIT and OFFSET as parameters; honour LIMIT 0 - UPDATE SET values are read from the parse tree like WHERE values, so SET x = '?' stores the string and SET x = ? binds a parameter; SET and WHERE parameters are bound together in statement order. - LIMIT and OFFSET accept a bound parameter (validated as a non-negative integer at execution) instead of being dropped, which returned every row. Other non-integer values raise. - LIMIT 0 returns no rows: MongoDB reads limit 0 as "no limit", and $limit: 0 is invalid in a pipeline. Tests that used a quoted '?' as a WHERE or SET parameter now use the documented unquoted ?; a quoted '?' is a string literal. --- pymongosql/executor.py | 66 ++++++++++++--------- pymongosql/helper.py | 11 ---- pymongosql/sql/ast.py | 69 +++++++++++----------- pymongosql/sql/builder.py | 6 +- pymongosql/sql/query_builder.py | 30 +++++++--- pymongosql/sql/update_handler.py | 24 ++++++-- pymongosql/sql/where_tree.py | 2 +- tests/test_cursor_delete.py | 6 +- tests/test_cursor_update.py | 2 +- tests/test_parse_tree_predicates.py | 21 ++++++- tests/test_sql_grouping_filters_aliases.py | 6 +- 11 files changed, 141 insertions(+), 102 deletions(-) diff --git a/pymongosql/executor.py b/pymongosql/executor.py index eff11ff..1b3cc4e 100644 --- a/pymongosql/executor.py +++ b/pymongosql/executor.py @@ -41,6 +41,14 @@ def _run_db_command(db: Any, command: Dict[str, Any], connection: Any, operation ) +def _paging(limit: Any, skip: Any) -> Any: + """Validate bound LIMIT/OFFSET values; returns (limit, skip).""" + for name, value in (("LIMIT", limit), ("OFFSET", skip)): + if value is not None and (isinstance(value, bool) or not isinstance(value, int) or value < 0): + raise ProgrammingError(f"{name} must be a non-negative integer, got {value!r}") + return limit, skip + + @dataclass class ExecutionContext: """Manages execution context for a single query""" @@ -149,8 +157,16 @@ def _execute_find_plan( # Replace placeholders with parameters in filter_stage only (not in projection) filter_stage = execution_plan.filter_stage or {} - # Positional parameters (named ones are converted to positional in execute()) - filter_stage, _ = SQLHelper.bind_filter(filter_stage, parameters) + # Positional parameters (named ones are converted to positional in execute()), + # in statement order: WHERE, then LIMIT, then OFFSET + bound, _ = SQLHelper.bind_filter( + {"filter": filter_stage, "limit": execution_plan.limit_stage, "skip": execution_plan.skip_stage}, + parameters, + ) + filter_stage = bound["filter"] + limit, skip = _paging(bound["limit"], bound["skip"]) + if limit == 0: + return {"cursor": {"id": 0, "firstBatch": []}, "ok": 1} projection_stage = execution_plan.projection_stage or {} @@ -169,13 +185,11 @@ def _execute_find_plan( sort_spec[field_name] = direction find_command["sort"] = sort_spec - # Apply skip if specified - if execution_plan.skip_stage: - find_command["skip"] = execution_plan.skip_stage - - # Apply limit if specified - if execution_plan.limit_stage: - find_command["limit"] = execution_plan.limit_stage + # Apply skip and limit if specified (MongoDB reads limit 0 as "no limit") + if skip: + find_command["skip"] = skip + if limit is not None: + find_command["limit"] = limit _logger.debug(f"Executing MongoDB command: {find_command}") @@ -236,9 +250,13 @@ def _execute_aggregate_plan( _logger.debug(f"Pipeline: {pipeline}") _logger.debug(f"Options: {options}") - # A pipeline generated from SQL carries the WHERE clause's parameter markers + # A pipeline generated from SQL carries parameter markers (WHERE, HAVING), then + # LIMIT and OFFSET + limit, skip = execution_plan.limit_stage, execution_plan.skip_stage if execution_plan.aggregate_parameterized: - pipeline, _ = SQLHelper.bind_filter(pipeline, parameters) + bound, _ = SQLHelper.bind_filter({"pipeline": pipeline, "limit": limit, "skip": skip}, parameters) + pipeline, limit, skip = bound["pipeline"], bound["limit"], bound["skip"] + limit, skip = _paging(limit, skip) # Get collection and call aggregate() collection = db[execution_plan.collection] @@ -266,11 +284,11 @@ def _execute_aggregate_plan( results = sorted(results, key=lambda x: x.get(field_name), reverse=reverse) # Apply skip and limit - if execution_plan.skip_stage: - results = results[execution_plan.skip_stage :] + if skip: + results = results[skip:] - if execution_plan.limit_stage: - results = results[: execution_plan.limit_stage] + if limit is not None: + results = results[:limit] # Apply projection if specified if execution_plan.projection_stage: @@ -605,20 +623,10 @@ def _execute_execution_plan( # Replace placeholders if parameters provided # Note: We need to replace both update_fields and filter_conditions in one pass # to maintain correct parameter ordering (SET clause first, then WHERE clause) - if isinstance(parameters, dict): - update_fields = SQLHelper.replace_placeholders_generic( - update_fields, parameters, execution_plan.parameter_style - ) - filter_conditions, _ = SQLHelper.bind_filter(filter_conditions, []) - else: - # SET values carry "?" placeholders; the WHERE clause carries parameter markers - params = list(parameters or []) - set_count = SQLHelper.count_placeholders(update_fields) - if set_count: - update_fields = SQLHelper.replace_placeholders_generic( - update_fields, params[:set_count], execution_plan.parameter_style or "qmark" - ) - filter_conditions, _ = SQLHelper.bind_filter(filter_conditions, params[set_count:]) + # SET values and the WHERE clause carry parameter markers; bind them in + # statement order (SET first) and require every parameter to be used + bound, _ = SQLHelper.bind_filter({"u": update_fields, "q": filter_conditions}, parameters) + update_fields, filter_conditions = bound["u"], bound["q"] # MongoDB update command format # https://www.mongodb.com/docs/manual/reference/command/update/ diff --git a/pymongosql/helper.py b/pymongosql/helper.py index 9ec0c38..0d41abe 100644 --- a/pymongosql/helper.py +++ b/pymongosql/helper.py @@ -153,17 +153,6 @@ def replace(val: Any) -> Any: raise ProgrammingError(f"{len(params)} parameters were given but the statement uses {idx[0]}") return bound, idx[0] - @staticmethod - def count_placeholders(value: Any) -> int: - """Number of legacy "?" placeholders (INSERT values, UPDATE SET) in a structure.""" - if isinstance(value, str): - return int(value == "?") - if isinstance(value, dict): - return sum(SQLHelper.count_placeholders(v) for v in value.values()) - if isinstance(value, list): - return sum(SQLHelper.count_placeholders(v) for v in value) - return 0 - @staticmethod def replace_placeholders_generic(value: Any, parameters: Any, style: Optional[str]) -> Any: """Recursively replace placeholders in nested structures for qmark or named styles.""" diff --git a/pymongosql/sql/ast.py b/pymongosql/sql/ast.py index 89027ea..4dba2c5 100644 --- a/pymongosql/sql/ast.py +++ b/pymongosql/sql/ast.py @@ -308,43 +308,42 @@ def visitHavingClause(self, ctx: PartiQLParser.HavingClauseContext) -> Any: return None def visitLimitClause(self, ctx: PartiQLParser.LimitClauseContext) -> Any: - """Handle LIMIT clause for result limiting""" - _logger.debug("Processing LIMIT clause") - try: - if hasattr(ctx, "exprSelect") and ctx.exprSelect(): - limit_text = ctx.exprSelect().getText() - try: - limit_value = int(limit_text) - self._query_parse_result.limit_value = limit_value - _logger.debug(f"Extracted limit value: {limit_value}") - except ValueError: - # e.g. LIMIT ?: dropping it would return every row - self._query_parse_result.unsupported_clauses.append( - f"LIMIT {limit_text} (needs an integer literal)" - ) - return self.visitChildren(ctx) - except Exception as e: - _logger.warning(f"Error processing LIMIT clause: {e}") - return self.visitChildren(ctx) + """Handle LIMIT: a non-negative integer literal or a bound parameter.""" + from .where_tree import is_param, operand + + if hasattr(ctx, "exprSelect") and ctx.exprSelect(): + try: + value = operand(ctx.exprSelect()) + except Exception: + value = None + if is_param(value) or (isinstance(value, int) and not isinstance(value, bool) and value >= 0): + self._query_parse_result.limit_value = value + else: + # Dropping it would return every row + text = ctx.exprSelect().getText() + self._query_parse_result.unsupported_clauses.append( + f"LIMIT {text} (needs a non-negative integer or a parameter)" + ) + return None def visitOffsetByClause(self, ctx: PartiQLParser.OffsetByClauseContext) -> Any: - """Handle OFFSET clause for result skipping""" - _logger.debug("Processing OFFSET clause") - try: - if hasattr(ctx, "exprSelect") and ctx.exprSelect(): - offset_text = ctx.exprSelect().getText() - try: - offset_value = int(offset_text) - self._query_parse_result.offset_value = offset_value - _logger.debug(f"Extracted offset value: {offset_value}") - except ValueError: - self._query_parse_result.unsupported_clauses.append( - f"OFFSET {offset_text} (needs an integer literal)" - ) - return self.visitChildren(ctx) - except Exception as e: - _logger.warning(f"Error processing OFFSET clause: {e}") - return self.visitChildren(ctx) + """Handle OFFSET: a non-negative integer literal or a bound parameter.""" + from .where_tree import is_param, operand + + if hasattr(ctx, "exprSelect") and ctx.exprSelect(): + try: + value = operand(ctx.exprSelect()) + except Exception: + value = None + if is_param(value) or (isinstance(value, int) and not isinstance(value, bool) and value >= 0): + self._query_parse_result.offset_value = value + else: + # Dropping it would return every row + text = ctx.exprSelect().getText() + self._query_parse_result.unsupported_clauses.append( + f"OFFSET {text} (needs a non-negative integer or a parameter)" + ) + return None def visitUpdateClause(self, ctx: PartiQLParser.UpdateClauseContext) -> Any: """Handle UPDATE clause to extract collection/table name.""" diff --git a/pymongosql/sql/builder.py b/pymongosql/sql/builder.py index 1a4a07a..029ec4c 100644 --- a/pymongosql/sql/builder.py +++ b/pymongosql/sql/builder.py @@ -355,10 +355,8 @@ def _build_sql_aggregate_plan(parse_result: "QueryParseResult") -> "QueryExecuti sort_stage[output] = direction if sort_stage: pipeline.append({"$sort": sort_stage}) - if parse_result.offset_value: - pipeline.append({"$skip": parse_result.offset_value}) - if parse_result.limit_value is not None: - pipeline.append({"$limit": parse_result.limit_value}) + # OFFSET/LIMIT (integers or parameters) are applied by the executor after binding + builder.skip(parse_result.offset_value).limit(parse_result.limit_value) # Configure the execution plan as an aggregate query builder._execution_plan.is_aggregate_query = True diff --git a/pymongosql/sql/query_builder.py b/pymongosql/sql/query_builder.py index 8e45b82..2f59075 100644 --- a/pymongosql/sql/query_builder.py +++ b/pymongosql/sql/query_builder.py @@ -53,10 +53,20 @@ def validate(self) -> bool: else: errors = self.validate_base() - if self.limit_stage is not None and (not isinstance(self.limit_stage, int) or self.limit_stage < 0): + from .where_tree import is_param + + if ( + self.limit_stage is not None + and not is_param(self.limit_stage) + and (not isinstance(self.limit_stage, int) or self.limit_stage < 0) + ): errors.append("Limit must be a non-negative integer") - if self.skip_stage is not None and (not isinstance(self.skip_stage, int) or self.skip_stage < 0): + if ( + self.skip_stage is not None + and not is_param(self.skip_stage) + and (not isinstance(self.skip_stage, int) or self.skip_stage < 0) + ): errors.append("Skip must be a non-negative integer") if errors: @@ -156,18 +166,22 @@ def sort(self, specs: List[Dict[str, int]]) -> "MongoQueryBuilder": return self - def limit(self, count: int) -> "MongoQueryBuilder": - """Set limit for results""" - if not isinstance(count, int) or count < 0: + def limit(self, count: Any) -> "MongoQueryBuilder": + """Set limit for results (an integer or a parameter marker bound at execution)""" + from .where_tree import is_param + + if not is_param(count) and (not isinstance(count, int) or count < 0): return self self._execution_plan.limit_stage = count _logger.debug(f"Set limit to: {count}") return self - def skip(self, count: int) -> "MongoQueryBuilder": - """Set skip count for pagination""" - if not isinstance(count, int) or count < 0: + def skip(self, count: Any) -> "MongoQueryBuilder": + """Set skip count for pagination (an integer or a parameter marker bound at execution)""" + from .where_tree import is_param + + if not is_param(count) and (not isinstance(count, int) or count < 0): return self self._execution_plan.skip_stage = count diff --git a/pymongosql/sql/update_handler.py b/pymongosql/sql/update_handler.py index e46f45e..f6fbfec 100644 --- a/pymongosql/sql/update_handler.py +++ b/pymongosql/sql/update_handler.py @@ -133,15 +133,27 @@ def _extract_set_assignment(self, ctx: Any) -> tuple[Optional[str], Any]: field_name = None field_value = None - # Extract field name from pathSimple + # Extract field name from pathSimple ("my field" -> my field, a[0] -> a.0) if hasattr(ctx, "pathSimple") and ctx.pathSimple(): - field_name = ctx.pathSimple().getText() + from .handler import ContextUtilsMixin - # Extract value from expr + field_name = ContextUtilsMixin.normalize_field_path(ctx.pathSimple().getText()) + + # Extract value from expr: a literal, value function or parameter marker, read + # from the parse tree so a string literal '?' stays a string if hasattr(ctx, "expr") and ctx.expr(): - expr_text = ctx.expr().getText() - # Parse the expression to get the actual value - field_value = self._parse_value(expr_text) + from ..error import NotSupportedError + from .where_tree import _coerce, _Field, operand + + try: + field_value = operand(ctx.expr()) + except NotSupportedError: + field_value = self._parse_value(ctx.expr().getText()) + if field_value == "?": + raise + if isinstance(field_value, _Field): + raise NotSupportedError(f"SET needs a value, not a column: {ctx.getText()}") + field_value = _coerce(field_value, as_double=True) # a fractional literal is stored as a double return field_name, field_value except Exception as e: diff --git a/pymongosql/sql/where_tree.py b/pymongosql/sql/where_tree.py index 525462a..2787213 100644 --- a/pymongosql/sql/where_tree.py +++ b/pymongosql/sql/where_tree.py @@ -208,7 +208,7 @@ def operand(ctx: Any, resolver: Any = None) -> Any: name = node.functionName().getText() registry = get_default_registry() if registry.has_function(name): - args = [operand(arg) for arg in node.expr()] + args = [_coerce(operand(arg), as_double=True) for arg in node.expr()] if any(isinstance(a, _Field) or is_param(a) for a in args): raise NotSupportedError(f"Value functions take literal arguments: {node.getText()}") return registry.execute(name, args) diff --git a/tests/test_cursor_delete.py b/tests/test_cursor_delete.py index 045ea76..d2d1aef 100644 --- a/tests/test_cursor_delete.py +++ b/tests/test_cursor_delete.py @@ -89,7 +89,7 @@ def test_delete_with_and_condition(self, conn): def test_delete_with_qmark_parameters(self, conn): """Test DELETE with qmark (?) placeholder parameters.""" cursor = conn.cursor() - result = cursor.execute(f"DELETE FROM {self.TEST_COLLECTION} WHERE artist = '?'", ["Charlie"]) + result = cursor.execute(f"DELETE FROM {self.TEST_COLLECTION} WHERE artist = ?", ["Charlie"]) assert result == cursor @@ -104,7 +104,7 @@ def test_delete_with_qmark_parameters(self, conn): def test_delete_with_multiple_parameters(self, conn): """Test DELETE with multiple qmark parameters.""" cursor = conn.cursor() - result = cursor.execute(f"DELETE FROM {self.TEST_COLLECTION} WHERE genre = '?' AND year = '?'", ["Pop", 2019]) + result = cursor.execute(f"DELETE FROM {self.TEST_COLLECTION} WHERE genre = ? AND year = ?", ["Pop", 2019]) assert result == cursor @@ -184,7 +184,7 @@ def test_delete_followed_by_insert(self, conn): def test_delete_executemany_with_parameters(self, conn): """Test executemany for bulk delete operations with parameters.""" cursor = conn.cursor() - sql = f"DELETE FROM {self.TEST_COLLECTION} WHERE artist = '?'" + sql = f"DELETE FROM {self.TEST_COLLECTION} WHERE artist = ?" # Delete multiple artists using executemany params = [["Alice"], ["Charlie"], ["Eve"]] diff --git a/tests/test_cursor_update.py b/tests/test_cursor_update.py index bb94461..5d7a16b 100644 --- a/tests/test_cursor_update.py +++ b/tests/test_cursor_update.py @@ -214,7 +214,7 @@ def test_update_set_null(self, conn): def test_update_executemany_with_parameters(self, conn): """Test executemany for bulk update operations with parameters.""" cursor = conn.cursor() - sql = f"UPDATE {self.TEST_COLLECTION} SET price = '?' WHERE title = '?'" + sql = f"UPDATE {self.TEST_COLLECTION} SET price = ? WHERE title = ?" # Update prices for multiple books using executemany params = [[25.99, "Book A"], [35.99, "Book B"], [45.99, "Book D"]] diff --git a/tests/test_parse_tree_predicates.py b/tests/test_parse_tree_predicates.py index 3425a88..d45f69e 100644 --- a/tests/test_parse_tree_predicates.py +++ b/tests/test_parse_tree_predicates.py @@ -121,9 +121,14 @@ def test_literal_question_mark_in_aggregate(self): p = plan("SELECT g, COUNT(*) AS n FROM t WHERE a = '?' AND b = ? GROUP BY g") assert '"a": "?"' in p.aggregate_pipeline and '"$pymongosqlParam"' in p.aggregate_pipeline - def test_limit_parameter_raises_instead_of_returning_everything(self): + def test_limit_and_offset_parameters_are_kept(self): + p = plan("SELECT _id FROM t WHERE a = ? LIMIT ? OFFSET ?") + assert (p.filter_stage, p.limit_stage, p.skip_stage) == ({"a": PARAM}, PARAM, PARAM) + + @pytest.mark.parametrize("clause", ["LIMIT -1", "LIMIT 'x'", "OFFSET 1.5", "LIMIT a"]) + def test_invalid_limit_or_offset_raises_instead_of_returning_everything(self, clause): with pytest.raises(Error): - plan("SELECT _id FROM t LIMIT ?") + plan(f"SELECT _id FROM t {clause}") @needs_sqlalchemy def test_sqlalchemy_limit_and_offset_are_literals(self): @@ -262,6 +267,18 @@ class TestLiveSelect: def test_rows(self, docs, sql_where, expected): assert ids(docs, sql_where) == expected + def test_bound_limit_offset_and_limit_zero(self, docs): + page = run(docs, f"SELECT _id FROM {COLLECTION} WHERE v >= ? ORDER BY _id LIMIT ? OFFSET ?", [1, 2, 1]) + assert [r[0] for r in page.fetchall()] == [2, 3] + assert run(docs, f"SELECT _id FROM {COLLECTION} LIMIT 0").fetchall() == [] + assert run(docs, f"SELECT g, COUNT(*) AS n FROM {COLLECTION} GROUP BY g LIMIT ?", [0]).fetchall() == [] + grouped = run(docs, f"SELECT g, COUNT(*) AS n FROM {COLLECTION} GROUP BY g ORDER BY g LIMIT ? OFFSET ?", [1, 1]) + assert [tuple(r) for r in grouped.fetchall()] == [("b", 2)] + with pytest.raises(Error): + run(docs, f"SELECT _id FROM {COLLECTION} LIMIT ?", [-1]) + with pytest.raises(Error): + run(docs, f"SELECT _id FROM {COLLECTION} WHERE v = ?", [1, 2]) + def test_count_distinct_and_having(self, docs): rows = run( docs, diff --git a/tests/test_sql_grouping_filters_aliases.py b/tests/test_sql_grouping_filters_aliases.py index 1f1e147..d6e4030 100644 --- a/tests/test_sql_grouping_filters_aliases.py +++ b/tests/test_sql_grouping_filters_aliases.py @@ -32,8 +32,10 @@ def test_group_by_groups_on_the_key(self): assert stages[1]["$project"] == {"_id": 0, "flag": "$_id.g0", "n": 1} def test_aggregate_query_keeps_order_by_skip_and_limit(self): - stages = pipeline("SELECT dept, SUM(amount) AS total FROM t GROUP BY dept ORDER BY total DESC LIMIT 1 OFFSET 1") - assert stages[-3:] == [{"$sort": {"total": -1}}, {"$skip": 1}, {"$limit": 1}] + p = plan("SELECT dept, SUM(amount) AS total FROM t GROUP BY dept ORDER BY total DESC LIMIT 1 OFFSET 1") + assert json.loads(p.aggregate_pipeline)[-1] == {"$sort": {"total": -1}} + # OFFSET/LIMIT are applied after binding (they may be parameters) + assert (p.skip_stage, p.limit_stage) == (1, 1) def test_order_by_aggregate_expression_uses_its_output(self): stages = pipeline("SELECT dept, SUM(amount) AS total FROM t GROUP BY dept ORDER BY SUM(amount)") From 6fff987507a311b71a57a63cff91c28505682827 Mon Sep 17 00:00:00 2001 From: Amin Ghadersohi Date: Sat, 26 Sep 2026 10:55:28 +0000 Subject: [PATCH 14/19] fix: bind LIKE patterns as parameters instead of inlining them; translate ILIKE The dialect rendered LIKE patterns as inline literals so the translator could turn them into a regex while parsing. SQLAlchemy's positional compilation then rewrites any %(name)s inside such a literal as a parameter (contains('%(k)s') failed with KeyError, other patterns could be corrupted). Render LIKE patterns as bound parameters again and translate them when parameters are bound: a pattern built from literals and parameters ('%' || ? || '%' ESCAPE '/', as SQLAlchemy renders contains, startswith and endswith, including autoescape) becomes the regex at execution. A non-string pattern parameter raises. lower(col) LIKE lower(pattern), as SQLAlchemy renders ilike(), becomes a case-insensitive regex. --- pymongosql/helper.py | 24 ++++-- pymongosql/sql/where_tree.py | 79 ++++++++++++++----- .../sqlalchemy_mongodb/sqlalchemy_dialect.py | 9 --- tests/test_parse_tree_predicates.py | 36 ++++++--- tests/test_sql_not_and_null_semantics.py | 6 +- 5 files changed, 108 insertions(+), 46 deletions(-) diff --git a/pymongosql/helper.py b/pymongosql/helper.py index 0d41abe..94c2bee 100644 --- a/pymongosql/helper.py +++ b/pymongosql/helper.py @@ -128,20 +128,32 @@ def bind_filter(value: Any, parameters: Any, exact: bool = True) -> Tuple[Any, i a parameter the translation dropped (e.g. a LIMIT it could not read) must not silently widen the query. Returns the bound value and the number used. """ - from .sql.where_tree import is_param + from .sql.where_tree import LIKE_KEY, is_like, is_param, like_operator params = [] if parameters is None else parameters if isinstance(params, dict) or not isinstance(params, Sequence) or isinstance(params, (str, bytes)): raise ProgrammingError("Positional parameters must be provided as a sequence") idx = [0] + def take() -> Any: + if idx[0] >= len(params): + raise ProgrammingError("Not enough parameters provided") + out = params[idx[0]] + idx[0] += 1 + return out + def replace(val: Any) -> Any: if is_param(val): - if idx[0] >= len(params): - raise ProgrammingError("Not enough parameters provided") - out = params[idx[0]] - idx[0] += 1 - return SQLHelper.to_bson_value(out) + return SQLHelper.to_bson_value(take()) + if is_like(val): + like = val[LIKE_KEY] + parts = [] + for part in like["parts"]: + part = take() if is_param(part) else part + if not isinstance(part, str): + raise ProgrammingError(f"A LIKE pattern parameter must be a string, got {part!r}") + parts.append(part) + return like_operator("".join(parts), like["escape"], like["i"]) if isinstance(val, dict): return {k: replace(v) for k, v in val.items()} if isinstance(val, list): diff --git a/pymongosql/sql/where_tree.py b/pymongosql/sql/where_tree.py index 2787213..6e617df 100644 --- a/pymongosql/sql/where_tree.py +++ b/pymongosql/sql/where_tree.py @@ -41,6 +41,29 @@ def is_param(value: Any) -> bool: return isinstance(value, dict) and list(value) == [PARAM_KEY] +# A LIKE whose pattern contains a bound parameter ('%' || ? || '%', as SQLAlchemy renders +# contains()); the regex is built when the parameters are bound. +LIKE_KEY = "$pymongosqlLike" + + +def is_like(value: Any) -> bool: + return isinstance(value, dict) and list(value) == [LIKE_KEY] + + +def like_operator(pattern: str, escape: Optional[str], case_insensitive: bool) -> Dict[str, Any]: + """The $regex operator document for a LIKE pattern.""" + regex, dotall = _like_regex(pattern, escape) + options = ("s" if dotall else "") + ("i" if case_insensitive else "") + return {"$regex": regex, "$options": options} if options else {"$regex": regex} + + +class _Pattern: + """A LIKE pattern built from string literals and bound parameters.""" + + def __init__(self, parts: List[Any]): + self.parts = parts + + _LEAVES = ( PartiQLParser.PredicateComparisonContext, PartiQLParser.PredicateIsContext, @@ -194,14 +217,17 @@ def operand(ctx: Any, resolver: Any = None) -> Any: raise NotSupportedError(f"Unsupported signed expression: {node.getText()}") return value if node.sign.text == "+" else -value if isinstance(node, PartiQLParser.MathOp00Context) and node.op is not None and node.op.text == "||": - left, right = operand(node.lhs), operand(node.rhs) - if ( - not (isinstance(left, str) and isinstance(right, str)) - or isinstance(left, _Field) - or isinstance(right, _Field) - ): - raise NotSupportedError(f"Only string literals can be concatenated: {node.getText()}") - return left + right + parts = [] + for side in (operand(node.lhs), operand(node.rhs)): + if isinstance(side, _Pattern): + parts.extend(side.parts) + elif (isinstance(side, str) and not isinstance(side, _Field)) or is_param(side): + parts.append(side) + else: + raise NotSupportedError(f"Only strings and parameters can be concatenated: {node.getText()}") + if all(isinstance(p, str) for p in parts): + return "".join(parts) + return _Pattern(parts) if isinstance(node, PartiQLParser.FunctionCallContext): from .value_function_registry import get_default_registry @@ -266,15 +292,18 @@ def leaf_filters(field: str, operator: str, value: Any, escape: Optional[str] = # x NOT IN (..., NULL) is never TRUE; x IN (..., NULL) is never FALSE false = NOTHING if None in values else {field: {"$nin": present + [None]}} return (true, false) if op == "IN" else (false, true) - if op in ("LIKE", "NOT LIKE"): - if is_param(value) or not isinstance(value, str): - # The pattern becomes a regex while parsing, before parameters are bound - raise NotSupportedError("LIKE needs a literal pattern, not a bound parameter") - pattern, dotall = _like_regex(value, escape) - regex = {"$regex": pattern, "$options": "s"} if dotall else {"$regex": pattern} + if op in ("LIKE", "NOT LIKE", "ILIKE", "NOT ILIKE"): + case_insensitive = "ILIKE" in op + if isinstance(value, str): + regex: Dict[str, Any] = like_operator(value, escape, case_insensitive) + elif is_param(value) or isinstance(value, _Pattern): + parts = value.parts if isinstance(value, _Pattern) else [value] + regex = {LIKE_KEY: {"parts": parts, "escape": escape, "i": case_insensitive}} + else: + raise NotSupportedError("LIKE needs a string pattern") true = {field: regex} false = {"$and": [{field: {"$not": regex}}, {field: {"$ne": None}}]} - return (true, false) if op == "LIKE" else (false, true) + return (true, false) if op in ("LIKE", "ILIKE") else (false, true) if op in ("BETWEEN", "NOT BETWEEN"): low, high = value true = {"$and": [{field: {"$gte": low}}, {field: {"$lte": high}}]} @@ -384,12 +413,26 @@ def _leaf(self, ctx: Any) -> Pair: return leaf_filters(field, "NOT BETWEEN" if negated else "BETWEEN", bounds) raise NotSupportedError(f"Unsupported predicate: {text}") + @staticmethod + def _case_folded(ctx: Any) -> Any: + """The argument of lower(x)/upper(x), as SQLAlchemy renders ILIKE; else None.""" + node = _unwrap(ctx) + if isinstance(node, PartiQLParser.FunctionCallContext) and len(node.expr()) == 1: + if node.functionName().getText().lower() in ("lower", "upper"): + return node.expr()[0] + return None + def _like(self, ctx: Any) -> Pair: - field = operand(ctx.lhs, self._resolver) + lhs, rhs, case_insensitive = ctx.lhs, ctx.rhs, False + folded_lhs, folded_rhs = self._case_folded(ctx.lhs), self._case_folded(ctx.rhs) + if folded_lhs is not None: + # lower(field) LIKE lower(pattern) or lower(field) LIKE 'pattern': case-insensitive + lhs, rhs, case_insensitive = folded_lhs, (folded_rhs if folded_rhs is not None else ctx.rhs), True + field = operand(lhs, self._resolver) if not isinstance(field, _Field): raise NotSupportedError(f"The left side of LIKE must be a field: {ctx.getText()}") - pattern = _value(ctx.rhs, self._resolver) - op = "NOT LIKE" if ctx.NOT() is not None else "LIKE" + pattern = _value(rhs, self._resolver) + op = ("NOT " if ctx.NOT() is not None else "") + ("ILIKE" if case_insensitive else "LIKE") if ctx.escape is None: return leaf_filters(field, op, pattern) # The grammar lets ESCAPE take a whole expression, so "a LIKE p ESCAPE '/' AND b = 1" diff --git a/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py b/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py index 64363db..75e028a 100644 --- a/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py +++ b/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py @@ -120,15 +120,6 @@ def limit_clause(self, select, **kw): text += "\n OFFSET " + self.process(select._offset_clause, literal_binds=True, **kw) return text - def visit_like_op_binary(self, binary, operator, **kw): - """Render LIKE patterns inline: PyMongoSQL turns them into a regex while parsing.""" - kw["literal_binds"] = True - return super().visit_like_op_binary(binary, operator, **kw) - - def visit_not_like_op_binary(self, binary, operator, **kw): - kw["literal_binds"] = True - return super().visit_not_like_op_binary(binary, operator, **kw) - class PyMongoSQLDDLCompiler(compiler.DDLCompiler): """MongoDB-specific DDL compiler. diff --git a/tests/test_parse_tree_predicates.py b/tests/test_parse_tree_predicates.py index d45f69e..4726dca 100644 --- a/tests/test_parse_tree_predicates.py +++ b/tests/test_parse_tree_predicates.py @@ -95,22 +95,36 @@ def test_escape_followed_by_more_conditions(self): def test_inner_wildcards_match_newlines(self): assert where("n LIKE 'a%b'") == {"n": {"$regex": "^a.*b$", "$options": "s"}} - def test_bound_pattern_raises(self): + def test_bound_pattern_is_translated_when_bound(self): + from pymongosql.helper import SQLHelper + + f = where("n LIKE '%' || ? || '%' ESCAPE '/'") + bound, used = SQLHelper.bind_filter(f, ["5/%(k)s"]) + assert (bound, used) == ({"n": {"$regex": ".*5%\\(k\\)s.*"}}, 1) with pytest.raises(Error): - where("n LIKE ?") + SQLHelper.bind_filter(where("n LIKE ?"), [5]) + + def test_case_insensitive_like_from_lower(self): + assert where("lower(n) LIKE lower('A%b')") == {"n": {"$regex": "^A.*b$", "$options": "si"}} + assert where("lower(n) LIKE 'a%'") == {"n": {"$regex": "^a.*", "$options": "i"}} @needs_sqlalchemy - def test_sqlalchemy_helpers_render_translatable_patterns(self): + def test_sqlalchemy_helpers_translate_after_binding(self): + from pymongosql.helper import SQLHelper + t = sa.table("t", sa.column("n")) for expr, regex in ( - (t.c.n.contains("ab"), ".*ab.*"), - (t.c.n.startswith("ab"), "^ab.*"), - (t.c.n.endswith("ab"), ".*ab$"), - (t.c.n.contains("5%_", autoescape=True), ".*5%_.*"), - (t.c.n.like("a/_%", escape="/"), "^a_.*"), + (t.c.n.contains("ab"), {"$regex": ".*ab.*"}), + (t.c.n.startswith("ab"), {"$regex": "^ab.*"}), + (t.c.n.endswith("ab"), {"$regex": ".*ab$"}), + (t.c.n.contains("5%_ %(k)s", autoescape=True), {"$regex": ".*5%_\\ %\\(k\\)s.*"}), + (t.c.n.like("a/_%", escape="/"), {"$regex": "^a_.*"}), + (t.c.n.ilike("A%"), {"$regex": "^A.*", "$options": "i"}), ): - sql = compiled(sa.select(t.c.n).where(expr)) - assert plan(sql).filter_stage == {"n": {"$regex": regex}}, sql + compiled_stmt = sa.select(t.c.n).where(expr).compile(dialect=PyMongoSQLDialect()) + params = [compiled_stmt.params[name] for name in compiled_stmt.positiontup] + bound, _ = SQLHelper.bind_filter(plan(str(compiled_stmt)).filter_stage, params) + assert bound == {"n": regex}, str(compiled_stmt) class TestParameters: @@ -307,6 +321,8 @@ def test_like_helpers_limit_offset_and_count_distinct(self, sqlalchemy_engine, d ] assert list(c.execute(sa.select(t.c._id).where(t.c.n.endswith("b")).order_by(t.c._id)).scalars()) == [2, 3] assert list(c.execute(sa.select(t.c._id).where(t.c.n.contains("50%", autoescape=True))).scalars()) == [1] + assert list(c.execute(sa.select(t.c._id).where(t.c.n.ilike("AB%"))).scalars()) == [1] + assert list(c.execute(sa.select(t.c._id).where(t.c.n.like("%(k)s"))).scalars()) == [] page = c.execute(sa.select(t.c._id).order_by(t.c._id).limit(2).offset(1)).scalars() assert list(page) == [2, 3] distinct = sa.func.count(sa.distinct(t.c.v)).label("d") diff --git a/tests/test_sql_not_and_null_semantics.py b/tests/test_sql_not_and_null_semantics.py index c1f972f..faeae0b 100644 --- a/tests/test_sql_not_and_null_semantics.py +++ b/tests/test_sql_not_and_null_semantics.py @@ -56,7 +56,7 @@ def test_not_equal_excludes_null(self): "SELECT _id FROM t WHERE NOT lower(b) = 'x'", "SELECT _id FROM t WHERE lower(b) = 'x'", "SELECT _id FROM t WHERE a = 1 AND lower(b) = 'x'", - "SELECT _id FROM t WHERE b LIKE ?", + "SELECT _id FROM t WHERE b LIKE 1", ], ) def test_untranslatable_where_raises_instead_of_widening(self, sql): @@ -121,7 +121,7 @@ def test_untranslatable_delete_deletes_nothing(self, docs): @pytest.mark.skipif(not HAS_SQLALCHEMY, reason="SQLAlchemy not available") class TestSQLAlchemy: - def test_like_pattern_is_rendered_inline(self): + def test_like_pattern_is_bound(self): import sqlalchemy as sa from pymongosql.sqlalchemy_mongodb.sqlalchemy_dialect import PyMongoSQLDialect @@ -129,7 +129,7 @@ def test_like_pattern_is_rendered_inline(self): t = sa.table("t", sa.column("b")) stmt = sa.select(t.c.b).where(t.c.b.like("O'B%"), t.c.b.not_like("x_")) sql = " ".join(str(stmt.compile(dialect=PyMongoSQLDialect())).split()) - assert sql == "SELECT b FROM t WHERE b LIKE 'O''B%' AND b NOT LIKE 'x_'" + assert sql == "SELECT b FROM t WHERE b LIKE ? AND b NOT LIKE ?" def test_core_not_and_like_rows(self, sqlalchemy_engine, docs): import sqlalchemy as sa From fb75dcb25e252d258991af4e336a6423fade5a81 Mon Sep 17 00:00:00 2001 From: Amin Ghadersohi Date: Sat, 26 Sep 2026 11:19:12 +0000 Subject: [PATCH 15/19] Keep %(name)s in inline string literals; LIMIT with literal_binds SQLAlchemy 2.0 renders every bind as %(name)s and then converts the whole compiled statement to qmark with a regular expression. Under literal_binds that also rewrites the text of an inline string literal, so x = '%(k)s' was sent as x = '?'. The compiler now masks quoted literals and identifiers while the markers are converted. limit_clause passed literal_binds twice when the statement was compiled with literal_binds (as Apache Superset compiles chart queries), raising TypeError for any query with a LIMIT. --- .../sqlalchemy_mongodb/sqlalchemy_dialect.py | 29 +++++++++++++++++-- 1 file changed, 27 insertions(+), 2 deletions(-) diff --git a/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py b/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py index 75e028a..a9cef5e 100644 --- a/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py +++ b/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py @@ -1,5 +1,6 @@ # -*- coding: utf-8 -*- import logging +import re import uuid from typing import Any, Dict, List, Optional, Tuple, Type from urllib.parse import quote_plus @@ -114,12 +115,36 @@ def visit_column(self, column, include_table=True, **kwargs): def limit_clause(self, select, **kw): """Render LIMIT/OFFSET as integer literals: PyMongoSQL reads them while parsing.""" text = "" + kw = {**kw, "literal_binds": True} if select._limit_clause is not None: - text += "\n LIMIT " + self.process(select._limit_clause, literal_binds=True, **kw) + text += "\n LIMIT " + self.process(select._limit_clause, **kw) if select._offset_clause is not None: - text += "\n OFFSET " + self.process(select._offset_clause, literal_binds=True, **kw) + text += "\n OFFSET " + self.process(select._offset_clause, **kw) return text + # A quoted SQL string literal or quoted identifier ('' and "" escape the quote) + _QUOTED = re.compile(r"'(?:[^']|'')*'|\"(?:[^\"]|\"\")*\"") + + def _process_positional(self): + """Convert bind markers to qmark without touching quoted literals. + + SQLAlchemy 2.0 renders every bind as ``%(name)s`` and then rewrites the + whole statement with a regular expression. That also rewrites the text of a + string literal rendered inline (``literal_binds``), so ``x = '%(k)s'`` became + ``x = '?'``. Quoted segments are masked while the markers are converted. + """ + masked: List[str] = [] + + def mask(match: "re.Match[str]") -> str: + masked.append(match.group(0)) + return "\x00%d\x00" % (len(masked) - 1) + + self.string = self._QUOTED.sub(mask, self.string) + try: + super()._process_positional() + finally: + self.string = re.sub("\x00(\\d+)\x00", lambda m: masked[int(m.group(1))], self.string) + class PyMongoSQLDDLCompiler(compiler.DDLCompiler): """MongoDB-specific DDL compiler. From 4a94d00e7b09d7162143e430a30a630b67be4e11 Mon Sep 17 00:00:00 2001 From: Amin Ghadersohi Date: Sat, 26 Sep 2026 11:19:12 +0000 Subject: [PATCH 16/19] DATE_TRUNC time grains and exact decimals in superset mode DATE_TRUNC('', field) in a projection or GROUP BY is translated to $dateTrunc (week-ending units add six days with $dateAdd). Units: second, minute, hour, day, week, week_monday, month, quarter, year, week_ending_saturday, week_ending_sunday; truncation is in UTC. An unknown unit or a non-field argument raises instead of dropping the column. The superset-mode SQLite stage registers the same DATE_TRUNC and STR_TO_DATETIME, stores datetimes as fixed-width UTC text and returns datetime values for them. Decimal128 columns were stored in the SQLite stage as doubles or text, so SUM lost digits and ORDER BY sorted text. They are now stored as REAL plus an exact text copy, and queries are rewritten (sqlglot, the new "superset" extra) so the column, SUM/AVG/MIN/MAX, GROUP BY, ORDER BY and comparisons with numeric literals are evaluated with Decimal arithmetic in Decimal128 precision (34 digits). Other uses of such a column raise NotSupportedError instead of computing with doubles. --- README.md | 14 + pymongosql/sql/ast.py | 12 +- pymongosql/sql/builder.py | 65 +++- pymongosql/sql/query_handler.py | 34 ++- pymongosql/superset_mongodb/exact_decimal.py | 277 +++++++++++++++++ .../superset_mongodb/query_db_sqlite.py | 71 ++++- pymongosql/superset_mongodb/time_grain.py | 140 +++++++++ pyproject.toml | 1 + requirements-optional.txt | 3 + ...test_decimals_time_grains_literal_binds.py | 279 ++++++++++++++++++ 10 files changed, 886 insertions(+), 10 deletions(-) create mode 100644 pymongosql/superset_mongodb/exact_decimal.py create mode 100644 pymongosql/superset_mongodb/time_grain.py create mode 100644 tests/test_decimals_time_grains_literal_binds.py diff --git a/README.md b/README.md index 5c57541..126fd21 100644 --- a/README.md +++ b/README.md @@ -743,6 +743,20 @@ PyMongoSQL can be used as a database driver in Apache Superset for querying and This allows seamless integration between MongoDB data and Superset's BI capabilities without requiring data migration to traditional SQL databases. +**Time grains and decimals:** + +- `DATE_TRUNC('', field)` is translated to MongoDB's `$dateTrunc` (MongoDB 5.0+), in + projections and `GROUP BY`. Units: `second`, `minute`, `hour`, `day`, `week` (starting + Sunday), `week_monday`, `month`, `quarter`, `year`, `week_ending_saturday` and + `week_ending_sunday`. Truncation is in UTC. The same function is available in the + superset-mode SQLite stage, so virtual datasets group by time the same way. +- In superset mode, a subquery's result is loaded into an in-memory SQLite database. Columns + holding `Decimal128` values are evaluated exactly there (install `pymongosql[superset]`, + which adds `sqlglot`): the column itself, `SUM`/`AVG`/`MIN`/`MAX` over it, `GROUP BY`, + `ORDER BY` and comparisons with numeric literals, with Decimal128's 34 significant + digits. Any other use of such a column (arithmetic, other functions, `DISTINCT` + aggregates) raises `NotSupportedError` rather than computing with doubles. + **Important Note on Collection Names:** When using collection names containing special characters (`.`, `-`, `:`), you must wrap them in double quotes to prevent Superset's SQL parser from incorrectly interpreting them. diff --git a/pymongosql/sql/ast.py b/pymongosql/sql/ast.py index 4dba2c5..aecd7ab 100644 --- a/pymongosql/sql/ast.py +++ b/pymongosql/sql/ast.py @@ -9,7 +9,7 @@ from .partiql.PartiQLLexer import PartiQLLexer from .partiql.PartiQLParser import PartiQLParser from .partiql.PartiQLParserVisitor import PartiQLParserVisitor -from .query_handler import QueryParseResult +from .query_handler import QueryParseResult, SelectHandler from .update_handler import UpdateParseResult _logger = logging.getLogger(__name__) @@ -296,7 +296,15 @@ def visitGroupClause(self, ctx: PartiQLParser.GroupClauseContext) -> Any: for key in ctx.groupKey() or []: if key.symbolPrimitive() is not None: self._query_parse_result.unsupported_clauses.append("GROUP BY key alias") - keys.append(ContextUtilsMixin.normalize_field_path(key.exprSelect().getText())) + text = ContextUtilsMixin.normalize_field_path(key.exprSelect().getText()) + try: + truncated = SelectHandler.date_trunc(key.exprSelect()) + except ValueError as e: + self._query_parse_result.unsupported_clauses.append(f"GROUP BY {text} ({e})") + truncated = None + if truncated is not None: + self._query_parse_result.computed[text] = truncated + keys.append(text) if ctx.PARTIAL() is not None: self._query_parse_result.unsupported_clauses.append("GROUP PARTIAL BY") self._query_parse_result.group_by = keys diff --git a/pymongosql/sql/builder.py b/pymongosql/sql/builder.py index 029ec4c..dc6fd30 100644 --- a/pymongosql/sql/builder.py +++ b/pymongosql/sql/builder.py @@ -2,7 +2,7 @@ import json import logging from dataclasses import dataclass -from typing import TYPE_CHECKING, Any, Dict, Optional, Union +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union from bson import json_util @@ -148,6 +148,8 @@ def strip_filter(value: Any) -> Any: for func_info in parse_result.aggregate_functions: func_info["argument"] = strip(func_info["argument"]) parse_result.group_by = [strip(name) for name in parse_result.group_by] + for expression in parse_result.computed.values(): + expression["field"] = strip(expression["field"]) for item in parse_result.select_items: if "field" in item: item["field"] = strip(item["field"]) @@ -164,6 +166,8 @@ def _build_query_plan(parse_result: "QueryParseResult") -> "QueryExecutionPlan": # Auto-generate aggregate pipeline for SQL aggregate functions (COUNT, SUM, etc.), GROUP BY and HAVING if parse_result.aggregate_functions or parse_result.group_by or parse_result.having is not None: return ExecutionPlanBuilder._build_sql_aggregate_plan(parse_result) + if any("computed" in item for item in parse_result.select_items): + return ExecutionPlanBuilder._build_computed_plan(parse_result) # ORDER BY may name a column by its SELECT alias; find() sorts on the field field_for_alias = {alias: name for name, alias in parse_result.column_aliases.items()} @@ -188,6 +192,54 @@ def _build_query_plan(parse_result: "QueryParseResult") -> "QueryExecutionPlan": plan = builder.build() return plan + @staticmethod + def _build_computed_plan(parse_result: "QueryParseResult") -> "QueryExecutionPlan": + """A SELECT with computed columns (DATE_TRUNC) and no grouping. + + Pipeline: $match (WHERE), $addFields (computed columns), $sort, then $project of + the SELECT list in order; OFFSET/LIMIT are applied after binding. + """ + from ..superset_mongodb.time_grain import mongo_expression + + builder = BuilderFactory.create_query_builder().collection(parse_result.collection) + pipeline: List[Dict[str, Any]] = [] + if parse_result.filter_conditions: + pipeline.append({"$match": parse_result.filter_conditions}) + added: Dict[str, Any] = {} + project: Dict[str, Any] = {} + sources: Dict[str, str] = {} # names ORDER BY may use -> sortable field + outputs = [] + for index, item in enumerate(parse_result.select_items): + if "computed" in item: + expression = parse_result.computed[item["computed"]] + hidden = f"__computed{index}" + added[hidden] = mongo_expression(expression["unit"], expression["field"]) + source, text = hidden, item["computed"] + else: + source = text = item["field"] + output = item["alias"] or text + project[output] = f"${source}" + sources[output] = sources[text] = sources[text.upper()] = source + outputs.append(output) + if "_id" not in project: + project["_id"] = 0 + if added: + pipeline.append({"$addFields": added}) + sort = {} + for spec in parse_result.sort_fields: + for name, direction in spec.items(): + sort[sources.get(name, sources.get(name.upper(), name))] = direction + if sort: + pipeline.append({"$sort": sort}) + pipeline.append({"$project": project}) + builder.skip(parse_result.offset_value).limit(parse_result.limit_value) + builder._execution_plan.is_aggregate_query = True + builder._execution_plan.aggregate_parameterized = True + builder._execution_plan.aggregate_pipeline = json_util.dumps(pipeline) + builder._execution_plan.aggregate_options = json.dumps({}) + builder._execution_plan.projection_stage = {name: 1 for name in outputs} + return builder.build() + @staticmethod def _aggregate_source(func_info: Dict[str, Any], key: str) -> Any: """$project expression for an accumulator: the value, or the reduced DISTINCT set.""" @@ -280,8 +332,15 @@ def _build_sql_aggregate_plan(parse_result: "QueryParseResult") -> "QueryExecuti if parse_result.filter_conditions: pipeline.append({"$match": parse_result.filter_conditions}) + from ..superset_mongodb.time_grain import mongo_expression + group_keys = {name: f"g{i}" for i, name in enumerate(parse_result.group_by)} - group_stage = {"_id": {key: f"${name}" for name, key in group_keys.items()} if group_keys else None} + + def key_source(name: str) -> Any: + computed = parse_result.computed.get(name) + return mongo_expression(computed["unit"], computed["field"]) if computed else f"${name}" + + group_stage = {"_id": {key: key_source(name) for name, key in group_keys.items()} if group_keys else None} # HAVING may name select-list outputs, grouped columns or aggregates; the ones # not in the SELECT list are computed as hidden outputs and removed afterwards. @@ -325,7 +384,7 @@ def _build_sql_aggregate_plan(parse_result: "QueryParseResult") -> "QueryExecuti source = 1 if key == output else ExecutionPlanBuilder._aggregate_source(func_info, key) output_for[func_info["expression"].upper()] = output else: - name = item["field"] + name = item.get("field") or item["computed"] if name not in group_keys: raise NotSupportedError(f"Column '{name}' must appear in GROUP BY or in an aggregate function") output, source = item["alias"] or name, f"$_id.{group_keys[name]}" diff --git a/pymongosql/sql/query_handler.py b/pymongosql/sql/query_handler.py index 9a5c25c..2421d4d 100644 --- a/pymongosql/sql/query_handler.py +++ b/pymongosql/sql/query_handler.py @@ -36,8 +36,10 @@ class QueryParseResult: aggregate_functions: List[Dict[str, Any]] = field(default_factory=list) # SELECT items in order: {"field": name, "alias": alias} or {"aggregate": index} select_items: List[Dict[str, Any]] = field(default_factory=list) - # GROUP BY field paths + # GROUP BY field paths (or the text of a computed expression, see ``computed``) group_by: List[str] = field(default_factory=list) + # Computed expressions by their SQL text: {"unit": ..., "field": ...} for DATE_TRUNC + computed: Dict[str, Dict[str, str]] = field(default_factory=dict) # Clauses that are parsed but cannot be translated faithfully unsupported_clauses: List[str] = field(default_factory=list) # FROM alias (FROM users AS u / FROM users u) @@ -156,6 +158,10 @@ def handle_visitor(self, ctx: PartiQLParser.SelectItemsContext, parse_result: "Q } ) continue + if kind == "computed": + parse_result.computed[field_name] = detail + parse_result.select_items.append({"computed": field_name, "alias": alias}) + continue if kind == "unsupported": # e.g. a + 1 or lower(a): projecting it as a field would silently read NULL parse_result.unsupported_clauses.append(f"SELECT {detail}") @@ -175,6 +181,25 @@ def handle_visitor(self, ctx: PartiQLParser.SelectItemsContext, parse_result: "Q _AGGREGATES = ("COUNT", "SUM", "AVG", "MIN", "MAX") + @staticmethod + def date_trunc(node: Any) -> Optional[Dict[str, str]]: + """{"unit", "field"} for DATE_TRUNC('', ), else None.""" + from ..superset_mongodb.time_grain import UNITS + from .where_tree import _PATH_NODES, _field_path, _string_literal, _unwrap + + node = _unwrap(node) + if not isinstance(node, PartiQLParser.FunctionCallContext): + return None + if node.functionName().getText().lower() != "date_trunc" or len(node.expr()) != 2: + return None + unit, target = _unwrap(node.expr()[0]), _unwrap(node.expr()[1]) + if not isinstance(unit, PartiQLParser.LiteralStringContext) or not isinstance(target, _PATH_NODES): + raise ValueError(f"DATE_TRUNC needs a unit literal and a field: {node.getText()}") + name = _string_literal(unit.getText()).lower() + if name not in UNITS: + raise ValueError(f"Unsupported DATE_TRUNC unit {name!r}; use one of {', '.join(UNITS)}") + return {"unit": name, "field": _field_path(target)} + @staticmethod def _classify_item(item) -> Tuple[str, Any]: """("field", path) | ("aggregate", (function, argument, distinct)) | ("unsupported", text).""" @@ -199,8 +224,11 @@ def _classify_item(item) -> Tuple[str, Any]: return "aggregate", (func, _field_path(argument), distinct) if isinstance(node, _PATH_NODES): return "field", _field_path(node) - except Exception: - pass + truncated = SelectHandler.date_trunc(node) + if truncated is not None: + return "computed", truncated + except Exception as e: + return "unsupported", f"{node.getText()} ({e})" return "unsupported", node.getText() def _extract_field_and_alias(self, item) -> Tuple[str, Optional[str]]: diff --git a/pymongosql/superset_mongodb/exact_decimal.py b/pymongosql/superset_mongodb/exact_decimal.py new file mode 100644 index 0000000..3b13a4b --- /dev/null +++ b/pymongosql/superset_mongodb/exact_decimal.py @@ -0,0 +1,277 @@ +# -*- coding: utf-8 -*- +"""Exact decimal evaluation for the superset-mode SQLite stage. + +SQLite has no decimal type: a Decimal128 value stored there becomes a double (15-17 +significant digits) or text (which sorts and compares as a string). A column whose +values include decimals is therefore stored twice: as REAL in the query table (so +SQLite can still filter and sort it approximately) and as exact decimal text in a +side table keyed by rowid. Before a query runs, it is rewritten so that projections +of such a column, SUM/AVG/MIN/MAX over it, GROUP BY, ORDER BY and comparisons with a +numeric literal are evaluated from the exact text with Decimal arithmetic in the +precision of MongoDB's Decimal128 (34 significant digits). A query that uses such a +column in any other way (arithmetic, functions, DISTINCT aggregates) raises instead +of silently computing with doubles. +""" + +from decimal import ROUND_HALF_EVEN, Context, Decimal +from typing import Any, Dict, List, Optional, Sequence, Tuple + +from ..error import NotSupportedError + +SHADOW_PREFIX = "__pymongosql_exact__" +# Exact enough for any sum of Decimal128 values; results are then rounded like MongoDB's +# own Decimal128 arithmetic. +_SUM_CONTEXT = Context(prec=13000, rounding=ROUND_HALF_EVEN, Emin=-99999, Emax=99999) +_DECIMAL128 = Context(prec=34, rounding=ROUND_HALF_EVEN, Emin=-6143, Emax=6144) + + +def _decimal(value: Optional[str]) -> Optional[Decimal]: + return None if value is None else Decimal(value) + + +def _round(value: Decimal) -> str: + return str(_DECIMAL128.plus(value)) + + +class _Collect: + def __init__(self) -> None: + self.values: List[Decimal] = [] + + def step(self, value: Optional[str]) -> None: + if value is not None: + self.values.append(Decimal(value)) + + def total(self) -> Decimal: + total = Decimal(0) + for value in self.values: + total = _SUM_CONTEXT.add(total, value) + return total + + +class _Sum(_Collect): + def finalize(self) -> Optional[str]: + return _round(self.total()) if self.values else None + + +class _Avg(_Collect): + def finalize(self) -> Optional[str]: + if not self.values: + return None + return str(_DECIMAL128.divide(self.total(), Decimal(len(self.values)))) + + +class _Min(_Collect): + def finalize(self) -> Optional[str]: + return str(min(self.values)) if self.values else None + + +class _Max(_Collect): + def finalize(self) -> Optional[str]: + return str(max(self.values)) if self.values else None + + +def sort_key(value: Optional[str]) -> Optional[str]: + """Text whose byte order is the numeric order of the decimal ``value``.""" + number = _decimal(value) + if number is None: + return None + if number == 0: + return "1" + digits = "".join(map(str, number.normalize().as_tuple().digits)) + exponent = number.adjusted() + if number > 0: + return "2%06d%s" % (exponent + 100000, digits) + # Negative: larger magnitude first; "~" makes a prefix (-1) sort after -1.2 + complement = "".join(str(9 - int(d)) for d in digits) + return "0%06d%s~" % (100000 - exponent, complement) + + +def compare(value: Optional[str], other: Optional[str]) -> Optional[int]: + left, right = _decimal(value), _decimal(other) + if left is None or right is None: + return None + return (left > right) - (left < right) + + +AGGREGATES = {"SUM": ("__pymongosql_exact_sum", _Sum), "AVG": ("__pymongosql_exact_avg", _Avg)} +AGGREGATES.update({"MIN": ("__pymongosql_exact_min", _Min), "MAX": ("__pymongosql_exact_max", _Max)}) +FUNCTIONS = {"__pymongosql_exact_key": sort_key, "__pymongosql_exact_cmp": compare} + + +def register(connection: Any) -> None: + for name, aggregate in AGGREGATES.values(): + connection.create_aggregate(name, 1, aggregate) + for name, function in FUNCTIONS.items(): + connection.create_function(name, 2 if name.endswith("cmp") else 1, function, deterministic=True) + + +def rewrite(sql: str, table: str, exact_columns: Sequence[str], columns: Sequence[str] = ()) -> Tuple[str, List[int]]: + """Rewrite ``sql`` to evaluate the exact columns of ``table`` exactly. + + Returns the SQL and the positions of the outer projections that return exact + decimal text. ``columns`` (all columns of ``table``, in order) expands ``SELECT *``. + Raises NotSupportedError when an exact column is used in a way that + cannot be evaluated exactly. + """ + try: + import sqlglot + from sqlglot import exp + except ImportError as e: # pragma: no cover - sqlglot ships with Superset + raise NotSupportedError("Exact decimal evaluation needs the sqlglot package") from e + + exact = set(exact_columns) + try: + tree = sqlglot.parse_one(sql, read="sqlite") + except sqlglot.errors.ParseError as e: + raise NotSupportedError(f"Cannot evaluate decimal columns exactly in: {sql}") from e + if not isinstance(tree, exp.Select): + raise NotSupportedError(f"Cannot evaluate decimal columns exactly in: {sql}") + + comparisons = (exp.EQ, exp.NEQ, exp.GT, exp.GTE, exp.LT, exp.LTE) + mirrored = {exp.GT: exp.LT, exp.GTE: exp.LTE, exp.LT: exp.GT, exp.LTE: exp.GTE} + aggregates = {exp.Sum: "SUM", exp.Avg: "AVG", exp.Min: "MIN", exp.Max: "MAX"} + + def reads_table(select: Any) -> bool: + from_ = select.args.get("from") or select.args.get("from_") + return ( + from_ is not None + and isinstance(from_.this, exp.Table) + and not from_.this.db + and from_.this.name == table + and not select.args.get("joins") + ) + + def numeric_literal(node: Any) -> Optional[str]: + if isinstance(node, exp.Literal) and not node.is_string: + return str(Decimal(node.this)) + if isinstance(node, exp.Neg): + inner = numeric_literal(node.this) + return None if inner is None else str(-Decimal(inner)) + return None + + def shadow(name: str) -> Any: + return exp.column(SHADOW_PREFIX + name, quoted=True) + + def call(name: str, *args: Any) -> Any: + return exp.Anonymous(this=name, expressions=list(args)) + + def expand_star(select: Any) -> None: + expanded = [] + for projection in select.expressions: + star = isinstance(projection, exp.Star) or ( + isinstance(projection, exp.Column) and isinstance(projection.this, exp.Star) + ) + if star and columns: + expanded.extend(exp.column(name, quoted=True) for name in columns) + else: + expanded.append(projection) + select.set("expressions", expanded) + + def process(select: Any, outer: bool) -> List[int]: + expand_star(select) + + def column(node: Any) -> Optional[str]: + return node.name if isinstance(node, exp.Column) and node.name in exact else None + + def value(node: Any) -> Optional[Any]: + name = column(node) + if name is not None: + return shadow(name) + if type(node) in aggregates and not isinstance(node.this, exp.Distinct): + name = column(node.this) + if name is not None: + return call(AGGREGATES[aggregates[type(node)]][0], shadow(name)) + return None + + by_name: Dict[str, Any] = {} + by_position: Dict[int, Any] = {} + for index, projection in enumerate(select.expressions): + inner = projection.this if isinstance(projection, exp.Alias) else projection + by_position[index] = inner + by_name.setdefault(projection.alias_or_name, inner) + + def referenced(node: Any) -> Any: + if isinstance(node, exp.Literal) and node.is_int: + return by_position.get(int(node.this) - 1, node) + if isinstance(node, exp.Column) and not node.table: + # SQLite resolves a bare name to an output alias before a column + return by_name.get(node.name, node) + return node + + for key in ("where", "having"): + clause = select.args.get(key) + if clause is None: + continue + for comparison in list(clause.find_all(*comparisons)): + if comparison.find_ancestor(exp.Select) is not select: + continue + left, right, kind = comparison.this, comparison.expression, type(comparison) + if numeric_literal(left) is not None: + left, right, kind = right, left, mirrored.get(kind, kind) + literal = numeric_literal(right) + exact_left = value(referenced(left)) if literal is not None else None + if exact_left is not None: + comparison.replace( + kind( + this=call("__pymongosql_exact_cmp", exact_left, exp.Literal.string(literal)), + expression=exp.Literal.number(0), + ) + ) + group = select.args.get("group") + if group is not None: + group.set( + "expressions", + [shadow(column(referenced(k))) if column(referenced(k)) else k for k in group.expressions], + ) + order = select.args.get("order") + if order is not None: + for ordered in order.expressions: + exact_order = value(referenced(ordered.this)) + if exact_order is not None: + ordered.set("this", call("__pymongosql_exact_key", exact_order)) + positions: List[int] = [] + projections = [] + for index, projection in enumerate(select.expressions): + inner = projection.this if isinstance(projection, exp.Alias) else projection + exact_projection = value(inner) + if exact_projection is None: + projections.append(projection) + continue + if outer: + positions.append(index) + name = projection.alias_or_name + projections.append(exp.alias_(exact_projection, name, quoted=True)) + select.set("expressions", projections) + return positions + + selects = [s for s in tree.find_all(exp.Select) if reads_table(s)] + positions: List[int] = [] + for select in selects: + found = process(select, outer=select is tree) + if select is tree: + positions = found + stars = [s for s in tree.find_all(exp.Star) if not isinstance(s.parent, exp.Count)] + if stars: + # A * over a derived table would return the exact text untyped + raise NotSupportedError("SELECT * over a derived table with decimal columns: name the columns") + # Any remaining reference to an exact column outside COUNT or IS NULL would be + # computed with doubles + for node in tree.find_all(exp.Column): + if node.name in exact: + allowed = node.find_ancestor(exp.Count, exp.Is) + if allowed is None: + raise NotSupportedError( + f"Decimal column {node.name!r} is used in an expression that cannot be evaluated exactly" + ) + columns = ", ".join(f'x."{c}" AS "{SHADOW_PREFIX}{c}"' for c in exact_columns) + source = f'(SELECT q.*, {columns} FROM "{table}" AS q LEFT JOIN "{table}_exact" AS x ON x.rid = q.rowid)' + for node in list(tree.find_all(exp.Table)): + if node.name == table and not node.db: + alias = node.alias or table + node.replace( + exp.Subquery( + this=sqlglot.parse_one(source[1:-1], read="sqlite"), + alias=exp.TableAlias(this=exp.to_identifier(alias, quoted=True)), + ) + ) + return tree.sql(dialect="sqlite"), positions diff --git a/pymongosql/superset_mongodb/query_db_sqlite.py b/pymongosql/superset_mongodb/query_db_sqlite.py index 348a485..d08aaa5 100644 --- a/pymongosql/superset_mongodb/query_db_sqlite.py +++ b/pymongosql/superset_mongodb/query_db_sqlite.py @@ -1,9 +1,14 @@ # -*- coding: utf-8 -*- +import datetime import logging import sqlite3 +from decimal import Decimal from typing import Any, Dict, List, Optional +from ..error import NotSupportedError +from . import exact_decimal from .query_db import QueryDatabase +from .time_grain import date_trunc_text, datetime_positions, datetime_text, parse_text, str_to_datetime_text _logger = logging.getLogger(__name__) @@ -12,6 +17,11 @@ # the column itself; expressions over it (SUM, CASE, ...) keep their numeric result. BOOLEAN_DECLTYPE = "PYMONGOSQL_BOOL" sqlite3.register_converter(BOOLEAN_DECLTYPE, lambda raw: int(raw) != 0) +# Datetimes are stored as fixed-width UTC text (so text order is time order) and read +# back as datetime when a query selects the column itself. +DATETIME_DECLTYPE = "PYMONGOSQL_DATETIME" +sqlite3.register_converter(DATETIME_DECLTYPE, lambda raw: datetime.datetime.fromisoformat(raw.decode())) +_NUMERIC = ("INTEGER", "REAL") class SQLiteTypeMapper: @@ -23,6 +33,8 @@ class SQLiteTypeMapper: int: "INTEGER", float: "REAL", bool: BOOLEAN_DECLTYPE, # stored as 0/1, read back as bool + Decimal: "REAL", # approximate copy; the exact text is kept in a side table + datetime.datetime: DATETIME_DECLTYPE, bytes: "BLOB", type(None): "NULL", dict: "TEXT", # Store as JSON string @@ -62,6 +74,8 @@ def infer_schema(cls, records: List[Dict[str, Any]]) -> Dict[str, str]: if current == "NULL": # First non-null value determines the type; NULL fits every type schema[col_name] = new_type + elif new_type in _NUMERIC and current in _NUMERIC: + schema[col_name] = "REAL" if "REAL" in (new_type, current) else "INTEGER" elif new_type not in ("NULL", current): # Upgrade to TEXT if types differ (safest option) schema[col_name] = "TEXT" @@ -74,6 +88,8 @@ def convert_value(cls, value: Any, target_type: str) -> Any: if value is None: return None + if target_type == DATETIME_DECLTYPE: + return datetime_text(value) if target_type in ("INTEGER", BOOLEAN_DECLTYPE): return int(value) if value is not None else None elif target_type == "REAL": @@ -103,6 +119,7 @@ def __init__(self) -> None: """Initialize SQLite3 bridge with in-memory database""" self._connection: Optional[sqlite3.Connection] = None self._tables: Dict[str, Dict[str, str]] = {} # table_name -> schema + self._exact_columns: Dict[str, List[str]] = {} # table_name -> columns with decimals self._is_closed = False def _ensure_connection(self) -> sqlite3.Connection: @@ -113,6 +130,10 @@ def _ensure_connection(self) -> sqlite3.Connection: if self._connection is None: # Create in-memory database self._connection = sqlite3.connect(":memory:", detect_types=sqlite3.PARSE_DECLTYPES) + # Time-grain functions the engine spec's expressions use + self._connection.create_function("date_trunc", 2, date_trunc_text, deterministic=True) + self._connection.create_function("str_to_datetime", -1, str_to_datetime_text, deterministic=True) + exact_decimal.register(self._connection) # Enable row factory to get dict-like rows self._connection.row_factory = sqlite3.Row _logger.debug("Created in-memory SQLite3 database") @@ -170,7 +191,8 @@ def insert_records( # Build INSERT statement columns = list(records[0].keys()) placeholders = ", ".join(["?" for _ in columns]) - insert_sql = f"INSERT INTO {table_name} ({', '.join(columns)}) VALUES ({placeholders})" + quoted = ", ".join('"%s"' % col.replace('"', '""') for col in columns) + insert_sql = f'INSERT INTO "{table_name}" ({quoted}) VALUES ({placeholders})' # Convert values to appropriate types schema = self._tables[table_name] @@ -183,7 +205,9 @@ def insert_records( converted_records.append(converted_row) try: + first_rowid = conn.execute(f'SELECT COALESCE(MAX(rowid), 0) + 1 FROM "{table_name}"').fetchone()[0] conn.executemany(insert_sql, converted_records) + self._write_exact_copies(table_name, columns, records, first_rowid) conn.commit() _logger.debug(f"Inserted {len(records)} records into {table_name}") return len(records) @@ -191,6 +215,33 @@ def insert_records( _logger.error(f"Error inserting records into {table_name}: {e}") raise + def _write_exact_copies( + self, table_name: str, columns: List[str], records: List[Dict[str, Any]], first_rowid: int + ) -> None: + """Keep the exact text of every column holding decimals (see exact_decimal).""" + exact = [] + for col in columns: + values = [r.get(col) for r in records if r.get(col) is not None] + if any(isinstance(v, Decimal) for v in values) and all( + isinstance(v, (int, float, Decimal)) and not isinstance(v, bool) for v in values + ): + exact.append(col) + if not exact: + return + if table_name in self._exact_columns: + raise NotSupportedError("Decimal columns can be loaded into a query table only once") + self._exact_columns[table_name] = exact + conn = self._ensure_connection() + side = f'"{table_name}_exact"' + conn.execute(f"CREATE TABLE {side} (rid INTEGER PRIMARY KEY, %s)" % ", ".join('"%s" TEXT' % c for c in exact)) + conn.executemany( + f"INSERT INTO {side} VALUES ({', '.join('?' * (len(exact) + 1))})", + [ + (rid,) + tuple(None if r.get(c) is None else str(r.get(c)) for c in exact) + for rid, r in enumerate(records, start=first_rowid) + ], + ) + def execute_query(self, query: str) -> List[Dict[str, Any]]: """ Execute a query against the SQLite3 database. @@ -203,13 +254,29 @@ def execute_query(self, query: str) -> List[Dict[str, Any]]: """ conn = self._ensure_connection() + positions: List[int] = [] + dates = datetime_positions(query) + for table, exact in self._exact_columns.items(): + if table in query: + query, positions = exact_decimal.rewrite(query, table, exact, list(self._tables[table])) try: cursor = conn.execute(query) # Fetch all rows and convert from sqlite3.Row to dict rows = cursor.fetchall() column_names = [desc[0] for desc in cursor.description] if cursor.description else [] - return [dict(zip(column_names, row)) for row in rows] + + def convert(index: int, value: Any) -> Any: + if value is None: + return None + if index in positions: + return Decimal(value) + return parse_text(value) if index in dates else value + + return [ + {name: convert(index, value) for index, (name, value) in enumerate(zip(column_names, row))} + for row in rows + ] except sqlite3.Error as e: _logger.error(f"Error executing query: {e}") raise diff --git a/pymongosql/superset_mongodb/time_grain.py b/pymongosql/superset_mongodb/time_grain.py new file mode 100644 index 0000000..ec05134 --- /dev/null +++ b/pymongosql/superset_mongodb/time_grain.py @@ -0,0 +1,140 @@ +# -*- coding: utf-8 -*- +"""DATE_TRUNC(unit, value): the time-grain function Superset's engine spec emits. + +The same function is evaluated in two places with the same meaning: + +- on a collection, translated to MongoDB's ``$dateTrunc`` (see ``mongo_expression``); +- in the superset-mode SQLite stage, registered as a SQL function over the stage's + fixed-width UTC datetime text (see ``date_trunc_text``). + +Units: second, minute, hour, day, week (starting Sunday), week_monday, month, +quarter, year, week_ending_saturday and week_ending_sunday (the last day of a week +that starts on Sunday or Monday respectively). +""" + +import datetime +from typing import Any, Dict, Optional + +UNITS = ( + "second", + "minute", + "hour", + "day", + "week", + "week_monday", + "month", + "quarter", + "year", + "week_ending_saturday", + "week_ending_sunday", +) +_TEXT = "%Y-%m-%d %H:%M:%S.%f" + + +def _check(unit: str) -> str: + unit = unit.lower() + if unit not in UNITS: + raise ValueError(f"Unsupported DATE_TRUNC unit: {unit!r}") + return unit + + +def mongo_expression(unit: str, field: str) -> Dict[str, Any]: + """The aggregation expression for DATE_TRUNC(unit, field).""" + unit = _check(unit) + date = f"${field}" + if unit in ("week", "week_ending_saturday"): + truncated = {"$dateTrunc": {"date": date, "unit": "week", "startOfWeek": "sunday"}} + elif unit in ("week_monday", "week_ending_sunday"): + truncated = {"$dateTrunc": {"date": date, "unit": "week", "startOfWeek": "monday"}} + else: + truncated = {"$dateTrunc": {"date": date, "unit": unit}} + if unit.startswith("week_ending"): + return {"$dateAdd": {"startDate": truncated, "unit": "day", "amount": 6}} + return truncated + + +def truncate(unit: str, value: datetime.datetime) -> datetime.datetime: + unit = _check(unit) + if unit == "second": + return value.replace(microsecond=0) + if unit == "minute": + return value.replace(second=0, microsecond=0) + if unit == "hour": + return value.replace(minute=0, second=0, microsecond=0) + day = value.replace(hour=0, minute=0, second=0, microsecond=0) + if unit == "day": + return day + if unit in ("week", "week_ending_saturday"): + start = day - datetime.timedelta(days=(day.weekday() + 1) % 7) # back to Sunday + elif unit in ("week_monday", "week_ending_sunday"): + start = day - datetime.timedelta(days=day.weekday()) # back to Monday + elif unit == "month": + return day.replace(day=1) + elif unit == "quarter": + return day.replace(month=(day.month - 1) // 3 * 3 + 1, day=1) + else: # year + return day.replace(month=1, day=1) + return start + datetime.timedelta(days=6) if unit.startswith("week_ending") else start + + +def datetime_text(value: Any) -> Optional[str]: + """Fixed-width UTC text for a datetime (or ISO string) value.""" + if value is None: + return None + if isinstance(value, str): + value = datetime.datetime.fromisoformat(value.strip().replace("Z", "+00:00")) + if isinstance(value, datetime.datetime): + if value.tzinfo is not None: + value = value.astimezone(datetime.timezone.utc).replace(tzinfo=None) + return value.strftime(_TEXT) + if isinstance(value, datetime.date): + return datetime.datetime(value.year, value.month, value.day).strftime(_TEXT) + raise ValueError(f"Not a datetime: {value!r}") + + +def date_trunc_text(unit: str, value: Any) -> Optional[str]: + """SQLite function DATE_TRUNC(unit, datetime text).""" + text = datetime_text(value) + if text is None: + return None + return truncate(unit, datetime.datetime.strptime(text, _TEXT)).strftime(_TEXT) + + +def str_to_datetime_text(*args: Any) -> Optional[str]: + """SQLite function STR_TO_DATETIME(text[, format]), as the value function of the same name.""" + from ..sql.value_function_registry import ValueFunctionRegistry + + if args and args[0] is None: + return None + return datetime_text(ValueFunctionRegistry.str_to_datetime(*args)) + + +def datetime_positions(sql: str) -> list: + """Positions of the outer projections that are DATE_TRUNC/STR_TO_DATETIME calls. + + Their SQLite result is datetime text (an expression has no declared type); the + caller converts it back to datetime. + """ + lowered = sql.lower() + if "date_trunc" not in lowered and "str_to_datetime" not in lowered: + return [] + try: + import sqlglot + from sqlglot import exp + + tree = sqlglot.parse_one(sql, read="sqlite") + except Exception: + return [] + if not isinstance(tree, exp.Select): + return [] + positions = [] + for index, projection in enumerate(tree.expressions): + inner = projection.this if isinstance(projection, exp.Alias) else projection + name = inner.name.lower() if isinstance(inner, exp.Anonymous) else "" + if isinstance(inner, exp.DateTrunc) or name in ("date_trunc", "str_to_datetime"): + positions.append(index) + return positions + + +def parse_text(value: Any) -> Any: + return datetime.datetime.strptime(value, _TEXT) if isinstance(value, str) else value diff --git a/pyproject.toml b/pyproject.toml index 8e5f8ce..914f6ad 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -41,6 +41,7 @@ dependencies = [ [project.optional-dependencies] retry = ["tenacity>=9.0.0"] sqlalchemy = ["sqlalchemy>=1.4.0"] +superset = ["sqlglot>=25.0.0"] pandas_sqlalchemy14 = [ "sqlalchemy>=1.4.0,<2.0.0", "pandas>=2.1.0,<2.2.0; python_version < '3.13'", diff --git a/requirements-optional.txt b/requirements-optional.txt index 3ebd3bb..787d1a2 100644 --- a/requirements-optional.txt +++ b/requirements-optional.txt @@ -3,3 +3,6 @@ tenacity>=9.0.0 # SQLAlchemy support (optional) - supports 1.4+ and 2.x sqlalchemy>=1.4.0,<3.0.0 + +# Superset mode: exact decimal evaluation in the SQLite stage (optional) +sqlglot>=25.0.0 diff --git a/tests/test_decimals_time_grains_literal_binds.py b/tests/test_decimals_time_grains_literal_binds.py new file mode 100644 index 0000000..9e6fc21 --- /dev/null +++ b/tests/test_decimals_time_grains_literal_binds.py @@ -0,0 +1,279 @@ +# -*- coding: utf-8 -*- +"""Superset-mode decimals, DATE_TRUNC time grains and literal binds. + +- The superset-mode SQLite stage evaluates decimal columns exactly (SUM/AVG/MIN/MAX, + GROUP BY, ORDER BY, comparisons) instead of as doubles or text. +- DATE_TRUNC(unit, field) is translated to MongoDB's $dateTrunc on a collection and + evaluated with the same meaning in the SQLite stage. +- Compiling with literal_binds keeps a string literal containing ``%(name)s`` intact. +""" + +import datetime +from decimal import Decimal, localcontext + +import pytest +from bson import Decimal128 + +from pymongosql.error import Error, NotSupportedError +from pymongosql.sql.parser import SQLParser +from pymongosql.superset_mongodb.query_db_sqlite import QueryDBSQLite +from pymongosql.superset_mongodb.time_grain import UNITS, truncate +from tests.conftest import HAS_SQLALCHEMY, make_superset_conn + +if HAS_SQLALCHEMY: + import sqlalchemy as sa + + from pymongosql.sqlalchemy_mongodb.sqlalchemy_dialect import PyMongoSQLDialect + +needs_sqlalchemy = pytest.mark.skipif(not HAS_SQLALCHEMY, reason="SQLAlchemy not available") +sqlglot = pytest.importorskip("sqlglot") + +COLLECTION = "test_decimals_time_grains" +BIG = Decimal("123456789012345678901.1234567891") +AMOUNTS = [BIG, Decimal("0.0000000001"), Decimal("-5.5"), Decimal("10"), None, Decimal("9.99")] +GROUPS = ["a", "a", "b", "b", "b", "a"] +TIMES = [ + datetime.datetime(2025, 12, 31, 23, 59, 59, 999000), # Wednesday + datetime.datetime(2026, 1, 3, 10, 15, 30, 250000), # Saturday + datetime.datetime(2026, 1, 4, 0, 0, 0), # Sunday + datetime.datetime(2026, 1, 5, 8, 30, 0), # Monday + datetime.datetime(2026, 4, 1, 12, 0, 1), + datetime.datetime(2026, 9, 26, 18, 45, 12, 5000), +] + + +def records(): + return [{"g": g, "amt": a, "ts": t} for g, a, t in zip(GROUPS, AMOUNTS, TIMES)] + + +def stage(sql): + db = QueryDBSQLite() + try: + db.insert_records("virtual_table", records()) + return [tuple(r.values()) for r in db.execute_query(sql)] + finally: + db.close() + + +A_SUM = Decimal("123456789012345678911.1134567892") # BIG + 0.0000000001 + 9.99, exact + + +def exact(values): + return [v for v in values if v is not None] + + +class TestExactDecimalsInTheSQLiteStage: + def test_sum_avg_min_max_are_exact(self): + rows = stage('SELECT SUM(amt) AS "SUM(amt)", AVG(amt) AS a, MIN(amt) AS lo, MAX(amt) AS hi FROM virtual_table') + # Decimal128 arithmetic: 34 significant digits, as MongoDB's own $sum and $avg + with localcontext() as ctx: + ctx.prec = 34 + total = sum(exact(AMOUNTS)) + average = total / 5 + assert total == Decimal("123456789012345678915.6134567892") + assert rows == [(total, average, Decimal("-5.5"), BIG)] + assert all(type(v) is Decimal for v in rows[0]) + + def test_group_by_with_exact_sums(self): + rows = stage('SELECT g AS g, SUM(amt) AS "SUM(amt)" FROM virtual_table GROUP BY g ORDER BY g') + assert rows == [("a", A_SUM), ("b", Decimal("4.5"))] + + def test_group_by_the_decimal_column(self): + rows = stage("SELECT amt AS amt, COUNT(*) AS n FROM virtual_table GROUP BY amt ORDER BY amt") + assert [r[0] for r in rows] == [ + None, + Decimal("-5.5"), + Decimal("0.0000000001"), + Decimal("9.99"), + Decimal("10"), + BIG, + ] + + def test_order_by_is_numeric_not_text(self): + rows = stage("SELECT amt FROM virtual_table WHERE amt IS NOT NULL ORDER BY amt DESC") + assert [r[0] for r in rows] == [BIG, Decimal("10"), Decimal("9.99"), Decimal("0.0000000001"), Decimal("-5.5")] + + def test_top_n_by_aggregate(self): + rows = stage('SELECT g, SUM(amt) AS "SUM(amt)" FROM virtual_table GROUP BY g ORDER BY "SUM(amt)" ASC LIMIT 1') + assert rows == [("b", Decimal("4.5"))] + + def test_comparisons_with_numeric_literals_are_exact(self): + # As doubles, 123456789012345678901.1234567891 equals 123456789012345678901.1234567890 + assert stage("SELECT COUNT(*) FROM virtual_table WHERE amt > 123456789012345678901.123456789") == [(1,)] + assert stage("SELECT COUNT(*) FROM virtual_table WHERE amt = 123456789012345678901.123456789") == [(0,)] + assert stage("SELECT COUNT(*) FROM virtual_table WHERE 0.00000000005 < amt AND amt < 1") == [(1,)] + rows = stage('SELECT g, SUM(amt) AS "SUM(amt)" FROM virtual_table GROUP BY g HAVING SUM(amt) > 4.5') + assert [r[0] for r in rows] == ["a"] + + def test_select_star_returns_decimals(self): + rows = stage("SELECT * FROM virtual_table WHERE g = 'b' ORDER BY amt") + assert [r[1] for r in rows] == [None, Decimal("-5.5"), Decimal("10")] + + def test_count_and_null_checks_are_allowed(self): + assert stage("SELECT COUNT(amt), COUNT(*) FROM virtual_table WHERE amt IS NULL OR amt IS NOT NULL") == [(5, 6)] + + @pytest.mark.parametrize( + "sql", + [ + "SELECT amt * 2 FROM virtual_table", + "SELECT ROUND(amt, 2) FROM virtual_table", + "SELECT SUM(DISTINCT amt) FROM virtual_table", + "SELECT * FROM (SELECT amt FROM virtual_table) AS x", + ], + ) + def test_other_uses_are_refused_rather_than_approximated(self, sql): + with pytest.raises(NotSupportedError): + stage(sql) + + def test_integer_and_float_columns_are_unchanged(self): + db = QueryDBSQLite() + try: + db.insert_records("virtual_table", [{"i": 1, "f": 0.5}, {"i": 2, "f": 1.25}]) + assert db.execute_query("SELECT SUM(i) AS i, SUM(f) AS f FROM virtual_table") == [{"i": 3, "f": 1.75}] + finally: + db.close() + + +class TestDateTruncTranslation: + def plan(self, sql): + return SQLParser(sql).get_execution_plan() + + def test_grouped_time_grain(self): + plan = self.plan( + "SELECT DATE_TRUNC('month', ts) AS __timestamp, COUNT(*) AS \"count\" FROM t " + "GROUP BY DATE_TRUNC('month', ts) ORDER BY \"count\" DESC LIMIT 10" + ) + assert '"$dateTrunc": {"date": "$ts", "unit": "month"}' in plan.aggregate_pipeline + + def test_week_ending_adds_six_days(self): + plan = self.plan( + "SELECT DATE_TRUNC('week_ending_sunday', t.ts) AS x FROM t GROUP BY DATE_TRUNC('week_ending_sunday', t.ts)" + ) + assert '"startOfWeek": "monday"' in plan.aggregate_pipeline + assert '"$dateAdd"' in plan.aggregate_pipeline and '"$t.ts"' not in plan.aggregate_pipeline + + def test_ungrouped_projection(self): + plan = self.plan("SELECT DATE_TRUNC('year', ts) AS y, name FROM t ORDER BY y DESC") + assert '"$addFields"' in plan.aggregate_pipeline and '"$sort": {"__computed0": -1}' in plan.aggregate_pipeline + + @pytest.mark.parametrize( + "sql", ["SELECT DATE_TRUNC('fortnight', ts) FROM t", "SELECT DATE_TRUNC(ts, 'day') FROM t"] + ) + def test_invalid_date_trunc_is_refused(self, sql): + with pytest.raises(Error): + self.plan(sql) + + @pytest.mark.parametrize("unit", UNITS) + def test_sqlite_stage_date_trunc(self, unit): + rows = stage(f"SELECT DATE_TRUNC('{unit}', ts) AS t FROM virtual_table ORDER BY ts") + assert [r[0] for r in rows] == [truncate(unit, t) for t in TIMES] + + +@needs_sqlalchemy +class TestLiteralBinds: + def compile(self, stmt, **kw): + return " ".join(str(stmt.compile(dialect=PyMongoSQLDialect(), **kw)).split()) + + def test_percent_name_literal_is_kept(self): + t = sa.table("t", sa.column("x"), sa.column("n")) + stmt = sa.select(t.c.x).where(t.c.x == "%(k)s", t.c.n == 5) + literal = self.compile(stmt, compile_kwargs={"literal_binds": True}) + assert literal == "SELECT x FROM t WHERE x = '%(k)s' AND n = 5" + bound = stmt.compile(dialect=PyMongoSQLDialect()) + assert " ".join(str(bound).split()) == "SELECT x FROM t WHERE x = ? AND n = ?" + assert list(bound.positiontup) == ["x_1", "n_1"] + + def test_quotes_inside_literals(self): + t = sa.table("t", sa.column("x")) + stmt = sa.select(t.c.x).where(t.c.x == 'it\'s %(a)s "q"') + assert self.compile(stmt, compile_kwargs={"literal_binds": True}) == ( + "SELECT x FROM t WHERE x = 'it''s %(a)s \"q\"'" + ) + + def test_limit_with_literal_binds(self): + t = sa.table("t", sa.column("x")) + stmt = sa.select(t.c.x).where(t.c.x.contains("%(k)s", autoescape=True)).limit(3).offset(1) + assert self.compile(stmt, compile_kwargs={"literal_binds": True}) == ( + "SELECT x FROM t WHERE (x LIKE '%' || '/%(k)s' || '%' ESCAPE '/') LIMIT 3 OFFSET 1" + ) + + +@pytest.fixture +def grain_docs(conn): + conn.database.drop_collection(COLLECTION) + conn.database[COLLECTION].insert_many( + [ + {"_id": i, "g": g, "amt": None if a is None else Decimal128(a), "ts": t} + for i, (g, a, t) in enumerate(zip(GROUPS, AMOUNTS, TIMES)) + ] + ) + yield conn + conn.database.drop_collection(COLLECTION) + + +class TestLive: + @pytest.mark.parametrize("unit", UNITS) + def test_physical_time_grain(self, grain_docs, unit): + cursor = grain_docs.cursor() + cursor.execute( + f"SELECT DATE_TRUNC('{unit}', ts) AS __timestamp, COUNT(*) AS \"count\" FROM {COLLECTION} " + f"WHERE ts >= STR_TO_DATETIME('2025-01-01T00:00:00') GROUP BY DATE_TRUNC('{unit}', ts) " + "ORDER BY __timestamp" + ) + expected = {} + for t in TIMES: + expected[truncate(unit, t)] = expected.get(truncate(unit, t), 0) + 1 + assert [tuple(r) for r in cursor.fetchall()] == sorted(expected.items()) + + def test_physical_ungrouped_time_grain(self, grain_docs): + cursor = grain_docs.cursor() + cursor.execute(f"SELECT DATE_TRUNC('week', ts) AS w, g FROM {COLLECTION} ORDER BY w DESC LIMIT 2") + assert [tuple(r) for r in cursor.fetchall()] == [ + (datetime.datetime(2026, 9, 20), "a"), + (datetime.datetime(2026, 3, 29), "b"), + ] + + @pytest.mark.parametrize("unit", UNITS) + def test_virtual_time_grain(self, grain_docs, unit): + superset = make_superset_conn() + try: + cursor = superset.cursor() + cursor.execute( + f"SELECT DATE_TRUNC('{unit}', ts) AS __timestamp, COUNT(*) AS \"count\" " + f"FROM (SELECT ts FROM {COLLECTION}) AS virtual_table " + "WHERE ts >= STR_TO_DATETIME('2025-01-01T00:00:00') " + f"GROUP BY DATE_TRUNC('{unit}', ts) ORDER BY __timestamp" + ) + expected = {} + for t in TIMES: + expected[truncate(unit, t)] = expected.get(truncate(unit, t), 0) + 1 + assert [tuple(r) for r in cursor.fetchall()] == sorted(expected.items()) + finally: + superset.close() + + def test_virtual_dataset_decimals_are_exact(self, grain_docs): + superset = make_superset_conn() + try: + cursor = superset.cursor() + cursor.execute( + 'SELECT g AS g, SUM(amt) AS "SUM(amt)" ' + f'FROM (SELECT g, amt FROM {COLLECTION}) AS virtual_table GROUP BY g ORDER BY "SUM(amt)" DESC' + ) + assert [tuple(r) for r in cursor.fetchall()] == [ + ("a", A_SUM), + ("b", Decimal("4.5")), + ] + cursor.execute(f"SELECT amt FROM (SELECT amt FROM {COLLECTION}) AS virtual_table ORDER BY amt DESC LIMIT 3") + assert [r[0] for r in cursor.fetchall()] == [BIG, Decimal("10"), Decimal("9.99")] + finally: + superset.close() + + @needs_sqlalchemy + def test_literal_binds_statement_runs(self, sqlalchemy_engine, grain_docs): + grain_docs.database[COLLECTION].insert_one({"_id": 99, "g": "%(k)s"}) + t = sa.table(COLLECTION, sa.column("_id"), sa.column("g")) + stmt = sa.select(t.c._id).where(t.c.g == "%(k)s") + sql = str(stmt.compile(sqlalchemy_engine, compile_kwargs={"literal_binds": True})) + with sqlalchemy_engine.connect() as c: + assert list(c.exec_driver_sql(sql).scalars()) == [99] + assert list(c.execute(stmt).scalars()) == [99] From 4ad11badbbef92fb3e4803d1f770612d99e5e122 Mon Sep 17 00:00:00 2001 From: Amin Ghadersohi Date: Sat, 26 Sep 2026 11:29:19 +0000 Subject: [PATCH 17/19] SQL NULL semantics for SUM and for aggregates over no rows SUM over a group whose values are all NULL or missing returned 0 ($sum of no numbers); SQL returns NULL. SUM(DISTINCT) likewise. An aggregate without GROUP BY over no input rows returned no row; SQL returns one row (COUNT 0, other aggregates NULL), which is what a chart showing a total expects. The $group is wrapped in $facet so the empty input yields that row; HAVING, LIMIT and grouped queries are unchanged. --- pymongosql/sql/builder.py | 40 +++++++++++-- ...test_decimals_time_grains_literal_binds.py | 38 ++++++++++++ tests/test_sql_parser_group.py | 60 ++++++++++++------- 3 files changed, 111 insertions(+), 27 deletions(-) diff --git a/pymongosql/sql/builder.py b/pymongosql/sql/builder.py index dc6fd30..f29f36b 100644 --- a/pymongosql/sql/builder.py +++ b/pymongosql/sql/builder.py @@ -241,13 +241,21 @@ def _build_computed_plan(parse_result: "QueryParseResult") -> "QueryExecutionPla return builder.build() @staticmethod - def _aggregate_source(func_info: Dict[str, Any], key: str) -> Any: - """$project expression for an accumulator: the value, or the reduced DISTINCT set.""" + def _aggregate_source(func_info: Dict[str, Any], key: str, index: int) -> Any: + """$project expression for an accumulator: the value, or the reduced DISTINCT set. + + SUM of no (non-NULL) values is NULL, not $sum's 0. + """ if not func_info.get("distinct"): + if func_info["function"] == "SUM": + return {"$cond": [{"$gt": [f"$__numbers{index}", 0]}, f"${key}", None]} return f"${key}" # SQL ignores NULL in DISTINCT aggregates values = {"$setDifference": [f"${key}", [None]]} reducer = {"COUNT": "$size", "SUM": "$sum", "AVG": "$avg", "MIN": "$min", "MAX": "$max"} + if func_info["function"] == "SUM": + numbers = {"$filter": {"input": values, "cond": {"$isNumber": "$$this"}}} + return {"$cond": [{"$gt": [{"$size": numbers}, 0]}, {"$sum": values}, None]} return {reducer[func_info["function"]]: values} @staticmethod @@ -370,8 +378,26 @@ def key_source(name: str) -> Any: group_stage[key] = {"$sum": {"$cond": [{"$gt": [f"${arg}", None]}, 1, 0]}} else: group_stage[key] = {accumulator: f"${arg}"} - - pipeline.append({"$group": group_stage}) + if func_name == "SUM": + # $sum of no numbers is 0; SQL's SUM of no values is NULL + group_stage[f"__numbers{i}"] = {"$sum": {"$cond": [{"$isNumber": f"${arg}"}, 1, 0]}} + + if group_keys: + pipeline.append({"$group": group_stage}) + else: + # An aggregate without GROUP BY returns one row even for no input rows + # (COUNT 0, other aggregates NULL); $group alone would return none. + empty = {"_id": None} + for name, accumulator_spec in group_stage.items(): + if name != "_id": + operator = next(iter(accumulator_spec)) + empty[name] = {"$sum": 0, "$addToSet": []}.get(operator) + pipeline.append({"$facet": {"row": [{"$group": group_stage}]}}) + pipeline.append( + {"$project": {"row": {"$cond": [{"$eq": [{"$size": "$row"}, 0]}, {"$literal": [empty]}, "$row"]}}} + ) + pipeline.append({"$unwind": "$row"}) + pipeline.append({"$replaceRoot": {"newRoot": "$row"}}) # Map every SELECT item, in order, to its output name and source project_stage = {"_id": 0} @@ -381,7 +407,9 @@ def key_source(name: str) -> Any: if "aggregate" in item: func_info = parse_result.aggregate_functions[item["aggregate"]] output, key = func_info["alias"], accumulator_keys[item["aggregate"]] - source = 1 if key == output else ExecutionPlanBuilder._aggregate_source(func_info, key) + source = ExecutionPlanBuilder._aggregate_source(func_info, key, item["aggregate"]) + if source == f"${key}" and key == output: + source = 1 output_for[func_info["expression"].upper()] = output else: name = item.get("field") or item["computed"] @@ -395,7 +423,7 @@ def key_source(name: str) -> Any: for name, source in hidden.items(): if isinstance(source, int): # a hidden aggregate: index into aggregate_functions project_stage[name] = ExecutionPlanBuilder._aggregate_source( - parse_result.aggregate_functions[source], accumulator_keys[source] + parse_result.aggregate_functions[source], accumulator_keys[source], source ) else: project_stage[name] = source diff --git a/tests/test_decimals_time_grains_literal_binds.py b/tests/test_decimals_time_grains_literal_binds.py index 9e6fc21..657f8df 100644 --- a/tests/test_decimals_time_grains_literal_binds.py +++ b/tests/test_decimals_time_grains_literal_binds.py @@ -277,3 +277,41 @@ def test_literal_binds_statement_runs(self, sqlalchemy_engine, grain_docs): with sqlalchemy_engine.connect() as c: assert list(c.exec_driver_sql(sql).scalars()) == [99] assert list(c.execute(stmt).scalars()) == [99] + + +class TestLiveAggregateNullSemantics: + @pytest.fixture + def null_docs(self, conn): + name = COLLECTION + "_nulls" + conn.database.drop_collection(name) + conn.database[name].insert_many( + [ + {"_id": 1, "g": "a", "amt": None}, + {"_id": 2, "g": "b", "amt": Decimal128("1.5")}, + {"_id": 3, "g": "b"}, + {"_id": 4, "g": "c", "amt": 0}, + ] + ) + yield conn, name + conn.database.drop_collection(name) + + def rows(self, conn, sql): + cursor = conn.cursor() + cursor.execute(sql) + return [tuple(r) for r in cursor.fetchall()] + + def test_sum_of_no_values_is_null(self, null_docs): + conn, name = null_docs + rows = self.rows(conn, f"SELECT g, SUM(amt) AS s, SUM(DISTINCT amt) AS d FROM {name} GROUP BY g ORDER BY g") + assert rows == [("a", None, None), ("b", Decimal("1.5"), Decimal("1.5")), ("c", 0, 0)] + assert self.rows(conn, f"SELECT SUM(amt) AS s FROM {name} WHERE g = 'a'") == [(None,)] + + def test_aggregate_without_group_by_returns_one_row_for_no_input(self, null_docs): + conn, name = null_docs + empty = f"FROM {name} WHERE g = 'none'" + assert self.rows( + conn, f"SELECT SUM(amt) AS s, COUNT(*) AS n, COUNT(DISTINCT g) AS d, MAX(amt) AS m {empty}" + ) == [(None, 0, 0, None)] + assert self.rows(conn, f"SELECT COUNT(*) AS n {empty} HAVING COUNT(*) > 0") == [] + assert self.rows(conn, f"SELECT g, COUNT(*) AS n {empty} GROUP BY g") == [] + assert self.rows(conn, f"SELECT COUNT(*) AS n FROM {name} LIMIT 0") == [] diff --git a/tests/test_sql_parser_group.py b/tests/test_sql_parser_group.py index 1f69545..2413ab4 100644 --- a/tests/test_sql_parser_group.py +++ b/tests/test_sql_parser_group.py @@ -4,6 +4,20 @@ from pymongosql.sql.parser import SQLParser +def group_of(pipeline): + """The $group stage; an aggregate without GROUP BY wraps it in $facet (one row even for no input).""" + for stage in pipeline: + if "$group" in stage: + return stage["$group"] + if "$facet" in stage: + return stage["$facet"]["row"][0]["$group"] + raise AssertionError("no $group stage") + + +def project_of(pipeline): + return next(stage["$project"] for stage in pipeline if "$project" in stage and "row" not in stage["$project"]) + + class TestCountStarParsing: """Test that COUNT(*) in SQL is translated to a MongoDB aggregate pipeline.""" @@ -17,11 +31,16 @@ def test_count_star_basic(self): assert plan.collection == "users" pipeline = json.loads(plan.aggregate_pipeline) - # Should have $group and $project stages - assert len(pipeline) == 2 - assert "$group" in pipeline[0] - assert pipeline[0]["$group"]["_id"] is None - assert pipeline[0]["$group"]["COUNT(*)"] == {"$sum": 1} + # $facet($group), the empty-input substitution, then $project + assert [next(iter(stage)) for stage in pipeline] == [ + "$facet", + "$project", + "$unwind", + "$replaceRoot", + "$project", + ] + assert group_of(pipeline)["_id"] is None + assert group_of(pipeline)["COUNT(*)"] == {"$sum": 1} def test_count_star_with_alias(self): """SELECT COUNT(*) AS total FROM users → alias used in $group""" @@ -31,10 +50,10 @@ def test_count_star_with_alias(self): assert plan.is_aggregate_query is True pipeline = json.loads(plan.aggregate_pipeline) - assert pipeline[0]["$group"]["total"] == {"$sum": 1} + assert group_of(pipeline)["total"] == {"$sum": 1} # $project should expose the alias - assert pipeline[1]["$project"]["total"] == 1 - assert pipeline[1]["$project"]["_id"] == 0 + assert project_of(pipeline)["total"] == 1 + assert project_of(pipeline)["_id"] == 0 def test_count_star_with_alias_no_as(self): """SELECT COUNT(*) total FROM users → alias without AS keyword""" @@ -44,7 +63,7 @@ def test_count_star_with_alias_no_as(self): assert plan.is_aggregate_query is True pipeline = json.loads(plan.aggregate_pipeline) - assert pipeline[0]["$group"]["total"] == {"$sum": 1} + assert group_of(pipeline)["total"] == {"$sum": 1} def test_count_star_with_where(self): """SELECT COUNT(*) AS total FROM users WHERE age > 25 → $match before $group""" @@ -54,11 +73,10 @@ def test_count_star_with_where(self): assert plan.is_aggregate_query is True pipeline = json.loads(plan.aggregate_pipeline) - # Should have $match, $group, $project - assert len(pipeline) == 3 + # $match first, then the grouping assert "$match" in pipeline[0] - assert "$group" in pipeline[1] - assert pipeline[1]["$group"]["total"] == {"$sum": 1} + assert "$facet" in pipeline[1] + assert group_of(pipeline)["total"] == {"$sum": 1} def test_count_star_projection_stage(self): """Projection stage should reflect aggregate output fields.""" @@ -83,7 +101,7 @@ def test_sum(self): assert plan.is_aggregate_query is True pipeline = json.loads(plan.aggregate_pipeline) - assert pipeline[0]["$group"]["total_price"] == {"$sum": "$price"} + assert group_of(pipeline)["total_price"] == {"$sum": "$price"} def test_avg(self): """SELECT AVG(age) AS avg_age FROM users""" @@ -93,7 +111,7 @@ def test_avg(self): assert plan.is_aggregate_query is True pipeline = json.loads(plan.aggregate_pipeline) - assert pipeline[0]["$group"]["avg_age"] == {"$avg": "$age"} + assert group_of(pipeline)["avg_age"] == {"$avg": "$age"} def test_min(self): """SELECT MIN(price) AS cheapest FROM products""" @@ -103,7 +121,7 @@ def test_min(self): assert plan.is_aggregate_query is True pipeline = json.loads(plan.aggregate_pipeline) - assert pipeline[0]["$group"]["cheapest"] == {"$min": "$price"} + assert group_of(pipeline)["cheapest"] == {"$min": "$price"} def test_max(self): """SELECT MAX(price) AS most_expensive FROM products""" @@ -113,7 +131,7 @@ def test_max(self): assert plan.is_aggregate_query is True pipeline = json.loads(plan.aggregate_pipeline) - assert pipeline[0]["$group"]["most_expensive"] == {"$max": "$price"} + assert group_of(pipeline)["most_expensive"] == {"$max": "$price"} def test_multiple_aggregates(self): """SELECT COUNT(*) AS cnt, AVG(price) AS avg_price, MAX(price) AS max_price FROM products""" @@ -123,13 +141,13 @@ def test_multiple_aggregates(self): assert plan.is_aggregate_query is True pipeline = json.loads(plan.aggregate_pipeline) - group = pipeline[0]["$group"] + group = group_of(pipeline) assert group["_id"] is None assert group["cnt"] == {"$sum": 1} assert group["avg_price"] == {"$avg": "$price"} assert group["max_price"] == {"$max": "$price"} # $project exposes all three - project = pipeline[1]["$project"] + project = project_of(pipeline) assert project == {"_id": 0, "cnt": 1, "avg_price": 1, "max_price": 1} def test_aggregate_with_nested_field(self): @@ -140,7 +158,7 @@ def test_aggregate_with_nested_field(self): assert plan.is_aggregate_query is True pipeline = json.loads(plan.aggregate_pipeline) - assert pipeline[0]["$group"]["revenue"] == {"$sum": "$details.total"} + assert group_of(pipeline)["revenue"] == {"$sum": "$details.total"} def test_aggregate_no_alias_uses_raw_text(self): """SELECT SUM(price) FROM products → alias defaults to SUM(price)""" @@ -149,7 +167,7 @@ def test_aggregate_no_alias_uses_raw_text(self): plan = parser.get_execution_plan() pipeline = json.loads(plan.aggregate_pipeline) - assert "SUM(price)" in pipeline[0]["$group"] + assert "SUM(price)" in group_of(pipeline) def test_regular_select_unaffected(self): """Regular SELECT without aggregate functions should not be affected.""" From bb8c9b6e468831c291988ba7b840afc6ec4f5e7c Mon Sep 17 00:00:00 2001 From: Amin Ghadersohi Date: Sat, 26 Sep 2026 11:34:05 +0000 Subject: [PATCH 18/19] Superset mode: a subquery without rows yields an empty table When a virtual dataset's own query matched no documents the SQLite stage created no table, so the outer query failed with "no such table" instead of returning COUNT 0 or no rows. The table is now created from the subquery's columns. --- pymongosql/superset_mongodb/executor.py | 6 +++++- tests/test_decimals_time_grains_literal_binds.py | 16 ++++++++++++++++ 2 files changed, 21 insertions(+), 1 deletion(-) diff --git a/pymongosql/superset_mongodb/executor.py b/pymongosql/superset_mongodb/executor.py index 22cf24f..47064b7 100644 --- a/pymongosql/superset_mongodb/executor.py +++ b/pymongosql/superset_mongodb/executor.py @@ -104,7 +104,11 @@ def execute( querydb_query = context.query table_name = "virtual_table" - query_db.insert_records(table_name, mongo_dicts) + if mongo_dicts: + query_db.insert_records(table_name, mongo_dicts) + elif column_names: + # No rows: the outer query still reads the table (COUNT(*) is 0, not an error) + query_db.create_table(table_name, {name: "" for name in column_names}) # Execute outer query against intermediate DB _logger.debug(f"Stage 2: Executing QueryDBSQLite query: {querydb_query}") diff --git a/tests/test_decimals_time_grains_literal_binds.py b/tests/test_decimals_time_grains_literal_binds.py index 657f8df..7131965 100644 --- a/tests/test_decimals_time_grains_literal_binds.py +++ b/tests/test_decimals_time_grains_literal_binds.py @@ -315,3 +315,19 @@ def test_aggregate_without_group_by_returns_one_row_for_no_input(self, null_docs assert self.rows(conn, f"SELECT COUNT(*) AS n {empty} HAVING COUNT(*) > 0") == [] assert self.rows(conn, f"SELECT g, COUNT(*) AS n {empty} GROUP BY g") == [] assert self.rows(conn, f"SELECT COUNT(*) AS n FROM {name} LIMIT 0") == [] + + +class TestLiveEmptyVirtualDataset: + def test_inner_query_without_rows(self, grain_docs): + superset = make_superset_conn() + inner = f"(SELECT g, amt FROM {COLLECTION} WHERE g = 'none') AS virtual_table" + try: + cursor = superset.cursor() + cursor.execute(f'SELECT COUNT(*) AS "count", SUM(amt) AS s FROM {inner}') + assert [tuple(r) for r in cursor.fetchall()] == [(0, None)] + cursor.execute(f'SELECT g, COUNT(*) AS "count" FROM {inner} GROUP BY g') + assert cursor.fetchall() == [] + cursor.execute(f"SELECT g FROM {inner}") + assert cursor.fetchall() == [] and [d[0] for d in cursor.description] == ["g"] + finally: + superset.close() From a7cafa1dbb22297a5982bc36dab8f4222cd1c62c Mon Sep 17 00:00:00 2001 From: Amin Ghadersohi Date: Sat, 26 Sep 2026 11:39:03 +0000 Subject: [PATCH 19/19] ORDER BY an aggregate that is not selected; HAVING next to DATE_TRUNC A grouped query ordered by an aggregate that is not in the SELECT list (Apache Superset's series-limit pre-query: GROUP BY the series, ORDER BY the limit metric) raised. The aggregate is now computed as a hidden output, like HAVING's, and removed after $sort. HAVING raised KeyError when the SELECT list had a DATE_TRUNC column. --- pymongosql/sql/ast.py | 4 +++ pymongosql/sql/builder.py | 33 ++++++++++++++++--- pymongosql/sql/query_handler.py | 2 ++ ...test_decimals_time_grains_literal_binds.py | 23 +++++++++++++ 4 files changed, 58 insertions(+), 4 deletions(-) diff --git a/pymongosql/sql/ast.py b/pymongosql/sql/ast.py index aecd7ab..783dd30 100644 --- a/pymongosql/sql/ast.py +++ b/pymongosql/sql/ast.py @@ -276,6 +276,10 @@ def visitOrderByClause(self, ctx: PartiQLParser.OrderByClauseContext) -> Any: for sort_spec in ctx.orderSortSpec(): field_name = sort_spec.expr().getText() if sort_spec.expr() else "_id" field_name = ContextUtilsMixin.normalize_field_path(field_name) + if sort_spec.expr() is not None: + kind, detail = SelectHandler._classify_item(sort_spec.expr()) + if kind == "aggregate": + self._query_parse_result.sort_aggregates[field_name] = detail # Check for ASC/DESC (default is ASC = 1) direction = 1 # ASC if hasattr(sort_spec, "DESC") and sort_spec.DESC(): diff --git a/pymongosql/sql/builder.py b/pymongosql/sql/builder.py index f29f36b..db8cd73 100644 --- a/pymongosql/sql/builder.py +++ b/pymongosql/sql/builder.py @@ -279,7 +279,8 @@ def strip(name: str) -> str: info = parse_result.aggregate_functions[item["aggregate"]] outputs[(info["function"], info["argument"], bool(info.get("distinct")))] = info["alias"] else: - outputs[item["field"]] = item["alias"] or item["field"] + name = item.get("field") or item["computed"] + outputs[name] = item["alias"] or name aliases = set(outputs.values()) def resolve(node: Any) -> Any: @@ -357,6 +358,30 @@ def key_source(name: str) -> Any: if parse_result.having is not None: having_filter = ExecutionPlanBuilder._translate_having(parse_result, group_keys, hidden) + # ORDER BY may name an aggregate that is not selected (e.g. a top-N query ordered + # by another metric); it is computed as a hidden output like HAVING's. + order_outputs: Dict[str, str] = {} + prefixes = [f"{q}." for q in (parse_result.collection_alias, parse_result.collection) if q] + for text, (func, arg, distinct) in parse_result.sort_aggregates.items(): + for prefix in prefixes: + if arg.startswith(prefix) and len(arg) > len(prefix): + arg = arg[len(prefix) :] + selected = [ + info["alias"] + for info in parse_result.aggregate_functions + if (info["function"], info["argument"], bool(info.get("distinct"))) == (func, arg, distinct) + and not info["alias"].startswith("__") + ] + if selected: + order_outputs[text] = selected[0] + continue + name = f"__order{len(order_outputs)}" + parse_result.aggregate_functions.append( + {"function": func, "argument": arg, "distinct": distinct, "alias": name, "expression": ""} + ) + hidden[name] = len(parse_result.aggregate_functions) - 1 + order_outputs[text] = name + accumulator_keys = [] for i, func_info in enumerate(parse_result.aggregate_functions): func_name = func_info["function"] @@ -430,18 +455,18 @@ def key_source(name: str) -> Any: pipeline.append({"$project": project_stage}) if having_filter is not None: pipeline.append({"$match": having_filter}) - if hidden: - pipeline.append({"$project": {name: 0 for name in hidden}}) sort_stage = {} for spec in parse_result.sort_fields: for name, direction in spec.items(): - output = output_for.get(name, output_for.get(name.upper())) + output = output_for.get(name, output_for.get(name.upper(), order_outputs.get(name))) if output is None: raise NotSupportedError(f"ORDER BY '{name}' must name a selected column or its alias") sort_stage[output] = direction if sort_stage: pipeline.append({"$sort": sort_stage}) + if hidden: + pipeline.append({"$project": {name: 0 for name in hidden}}) # OFFSET/LIMIT (integers or parameters) are applied by the executor after binding builder.skip(parse_result.offset_value).limit(parse_result.limit_value) diff --git a/pymongosql/sql/query_handler.py b/pymongosql/sql/query_handler.py index 2421d4d..d2a9d02 100644 --- a/pymongosql/sql/query_handler.py +++ b/pymongosql/sql/query_handler.py @@ -40,6 +40,8 @@ class QueryParseResult: group_by: List[str] = field(default_factory=list) # Computed expressions by their SQL text: {"unit": ..., "field": ...} for DATE_TRUNC computed: Dict[str, Dict[str, str]] = field(default_factory=dict) + # ORDER BY aggregates by their SQL text: (function, argument, distinct) + sort_aggregates: Dict[str, Tuple[str, str, bool]] = field(default_factory=dict) # Clauses that are parsed but cannot be translated faithfully unsupported_clauses: List[str] = field(default_factory=list) # FROM alias (FROM users AS u / FROM users u) diff --git a/tests/test_decimals_time_grains_literal_binds.py b/tests/test_decimals_time_grains_literal_binds.py index 7131965..f848929 100644 --- a/tests/test_decimals_time_grains_literal_binds.py +++ b/tests/test_decimals_time_grains_literal_binds.py @@ -331,3 +331,26 @@ def test_inner_query_without_rows(self, grain_docs): assert cursor.fetchall() == [] and [d[0] for d in cursor.description] == ["g"] finally: superset.close() + + +class TestLiveOrderByAggregateNotSelected: + def rows(self, conn, sql): + cursor = conn.cursor() + cursor.execute(sql) + return [tuple(r) for r in cursor.fetchall()] + + def test_top_n_ordered_by_another_metric(self, grain_docs): + # Superset's series-limit pre-query: group by the series, order by the limit metric + sql = f'SELECT g AS g, COUNT(*) AS "count" FROM {COLLECTION} GROUP BY g ORDER BY sum(amt) DESC LIMIT 1' + assert self.rows(grain_docs, sql) == [("a", 3)] + sql = f"SELECT g, COUNT(*) AS n FROM {COLLECTION} GROUP BY g ORDER BY SUM({COLLECTION}.amt) ASC" + assert self.rows(grain_docs, sql) == [("b", 3), ("a", 3)] + sql = f"SELECT g, COUNT(*) AS n FROM {COLLECTION} GROUP BY g HAVING SUM(amt) > 5 ORDER BY MAX(amt) DESC" + assert self.rows(grain_docs, sql) == [("a", 3)] + + def test_time_grain_with_having(self, grain_docs): + sql = ( + f"SELECT DATE_TRUNC('month', ts) AS m, COUNT(*) AS n FROM {COLLECTION} " + "GROUP BY DATE_TRUNC('month', ts) HAVING COUNT(*) > 1 ORDER BY m" + ) + assert self.rows(grain_docs, sql) == [(datetime.datetime(2026, 1, 1), 3)]