From 357d159bf709a8accf12d6b71706373ee0cca672 Mon Sep 17 00:00:00 2001 From: Amin Ghadersohi Date: Sat, 26 Sep 2026 03:15:57 +0000 Subject: [PATCH 01/14] 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/14] 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/14] 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/14] 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/14] 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/14] 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/14] 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 b93eb514b138f5e2c1ac9ddc49a585c48f33eebe Mon Sep 17 00:00:00 2001 From: Amin Ghadersohi Date: Sat, 26 Sep 2026 03:27:32 +0000 Subject: [PATCH 08/14] ci: publish immutable wheels from the fork Jenkins runs the full test suite against a disposable MongoDB in the build pod, builds a reproducible wheel and uploads it without ever overwriting an existing artifact; a retry is accepted only when the stored archive has identical content. Pull-request builds publish +pr.., normalized per PEP 440 so the filename matches the one the build backend writes; stable versions are published only from master. --- Jenkinsfile | 108 +++++++++++++++++++++++ ci/publish_wheel.py | 71 ++++++++++++++++ ci/release_version.py | 39 +++++++++ tests/ci/__init__.py | 0 tests/ci/test_publish_wheel.py | 142 +++++++++++++++++++++++++++++++ tests/ci/test_release_version.py | 40 +++++++++ 6 files changed, 400 insertions(+) create mode 100644 Jenkinsfile create mode 100644 ci/publish_wheel.py create mode 100644 ci/release_version.py create mode 100644 tests/ci/__init__.py create mode 100644 tests/ci/test_publish_wheel.py create mode 100644 tests/ci/test_release_version.py diff --git a/Jenkinsfile b/Jenkinsfile new file mode 100644 index 0000000..ceb67d5 --- /dev/null +++ b/Jenkinsfile @@ -0,0 +1,108 @@ +// Fork publisher. Pull-request wheels are versioned +pr.. +// (PEP 440-normalized by ci/release_version.py); stable wheels are published only +// from reviewed master. A published artifact is never overwritten. +podTemplate( + imagePullSecrets: ['preset-pull'], + containers: [ + containerTemplate(name: 'ci', image: 'preset/ci:latest', + ttyEnabled: true, command: 'cat'), + containerTemplate(name: 'py-ci', image: 'preset/python:3.9.18-2024-02-21-ci', + ttyEnabled: true, command: 'cat'), + // Disposable server for the test suite; reachable on localhost inside the pod. + containerTemplate(name: 'mongo', image: 'mongo:8.0', + envVars: [ + envVar(key: 'MONGO_INITDB_ROOT_USERNAME', value: 'admin'), + envVar(key: 'MONGO_INITDB_ROOT_PASSWORD', value: 'secret'), + ]) + ] +) { + node(POD_LABEL) { + checkout scm + def revision = sh(script: 'git rev-parse HEAD', returnStdout: true).trim() + boolean isMaster = env.BRANCH_NAME == 'master' + boolean isPR = env.CHANGE_ID != null + if (!isMaster && !isPR) { + error('Only master and pull-request builds publish; use a PR.') + } + + container('py-ci') { + stage('Test and build') { + def args = isMaster ? '' : "${env.CHANGE_ID} ${revision.take(12)}" + sh ''' + set -eu + python -m venv .venv + .venv/bin/pip install 'sqlalchemy==2.0.52' 'pymongo==4.17.0' \ + 'antlr4-python3-runtime==4.13.2' 'jmespath==1.1.0' 'pandas>=2.2,<3' \ + 'tenacity==9.1.2' 'pytest==8.3.5' 'boto3>=1.36,<2' 'packaging==25.0' \ + 'build==1.4.4' 'setuptools==80.9.0' 'setuptools_scm==8.3.1' 'wheel==0.45.1' + ''' + def version = sh(script: ".venv/bin/python ci/release_version.py ${args}", + returnStdout: true).trim() + env.PUBLISH_VERSION = version + env.WHEEL = "pymongosql-${version}-py3-none-any.whl" + env.KEY = "pymongosql/${env.WHEEL}" + sh ''' + set -eu + .venv/bin/pip install --no-deps -e . + for attempt in $(seq 1 30); do + .venv/bin/python -c "import pymongo; pymongo.MongoClient('mongodb://admin:secret@localhost:27017', serverSelectionTimeoutMS=2000).admin.command('ping')" && break + sleep 2 + done + .venv/bin/python tests/run_test_server.py setup + .venv/bin/python -m pytest -q tests + .venv/bin/pip uninstall -y pymongosql + python - <<'PY' +import os +import re +from pathlib import Path +path = Path('pymongosql/__init__.py') +source = path.read_text() +pattern = re.compile(r'^__version__: str = "[^"]+"$', re.MULTILINE) +assert len(pattern.findall(source)) == 1 +path.write_text(pattern.sub('__version__: str = "' + os.environ['PUBLISH_VERSION'] + '"', source)) +PY + SOURCE_DATE_EPOCH=$(git -c safe.directory="$PWD" log -1 --format=%ct) + case "$SOURCE_DATE_EPOCH" in + ''|*[!0-9]*) echo "Invalid commit timestamp for reproducible build" >&2; exit 1 ;; + esac + export SOURCE_DATE_EPOCH + # Pin the build backend and remove stale output for reproducible retries. + rm -rf build dist pymongosql.egg-info + .venv/bin/python -m build --wheel --no-isolation + test -f "dist/$WHEEL" || { echo "missing dist/$WHEEL"; ls -1 dist; exit 1; } + .venv/bin/python -c "import os, sys, zipfile; names = zipfile.ZipFile('dist/' + os.environ['WHEEL']).namelist(); sys.exit('wheel ships tests or ci' if any(n.startswith(('tests/', 'ci/')) for n in names) else 0)" + .venv/bin/pip install --force-reinstall --no-deps "dist/$WHEEL" + .venv/bin/python - <<'PY' +import importlib.metadata as im +import os +import sqlalchemy as sa +assert im.version('pymongosql') == os.environ['PUBLISH_VERSION'] +engine = sa.create_engine('mongodb://user@localhost/db') +assert engine.dialect.name == 'mongodb' +engine.dispose() +PY + sha256sum "dist/$WHEEL" + ''' + } + } + container('ci') { + stage('Publish immutable wheel') { + withCredentials([[ + $class: 'AmazonWebServicesCredentialsBinding', + credentialsId: 'ci-user', + accessKeyVariable: 'AWS_ACCESS_KEY_ID', + secretKeyVariable: 'AWS_SECRET_ACCESS_KEY' + ]]) { + withEnv(["ALLOW_IDENTICAL_PR_ARTIFACT=${isPR && !isMaster}"]) { + sh ''' + set -eu + python -m pip install --quiet 'boto3>=1.36,<2' + python ci/publish_wheel.py + ''' + } + } + } + } + archiveArtifacts artifacts: 'dist/*.whl,published.sha256', fingerprint: true + } +} diff --git a/ci/publish_wheel.py b/ci/publish_wheel.py new file mode 100644 index 0000000..894fdca --- /dev/null +++ b/ci/publish_wheel.py @@ -0,0 +1,71 @@ +"""Publish an immutable wheel; retries may reuse identical archive content.""" + +import hashlib +import io +import os +from pathlib import Path +from zipfile import BadZipFile, ZipFile + +import boto3 +from botocore.exceptions import ClientError + + +def same_wheel_content(stored, fresh): + """Compare sorted member names and bytes, ignoring ZIP metadata such as timestamps.""" + if stored == fresh: + return True + try: + with ZipFile(io.BytesIO(stored)) as old, ZipFile(io.BytesIO(fresh)) as new: + old_members = sorted(old.infolist(), key=lambda member: member.filename) + new_members = sorted(new.infolist(), key=lambda member: member.filename) + return [member.filename for member in old_members] == [member.filename for member in new_members] and all( + old.read(a) == new.read(b) for a, b in zip(old_members, new_members) + ) + except BadZipFile: + return False + + +def publish_wheel(s3, bucket, key, body, is_pr=False): + try: + s3.put_object(Bucket=bucket, Key=key, Body=body, IfNoneMatch="*") + except ClientError as error: + if error.response["Error"]["Code"] != "PreconditionFailed": + raise + print("Artifact already exists; verifying archive content without overwriting.") + + response = s3.get_object(Bucket=bucket, Key=key) + try: + stored = response["Body"].read() + finally: + response["Body"].close() + if not same_wheel_content(stored, body): + raise RuntimeError( + "Stored wheel differs from this commit's freshly built artifact; " + "refusing to overwrite or accept it (local sha256={}, stored sha256={}).".format( + hashlib.sha256(body).hexdigest(), + hashlib.sha256(stored).hexdigest(), + ) + + ("" if is_pr else " Bump __version__ before publishing different content.") + ) + print("Published wheel verified by archive content: " + key) + return hashlib.sha256(stored).hexdigest() + + +def main(): + wheel = os.environ["WHEEL"] + receipt = Path("published.sha256") + # Do not leave a previous build's success receipt after a failed retry. + if receipt.exists(): + receipt.unlink() + digest = publish_wheel( + boto3.client("s3"), + "preset-pypi", + os.environ["KEY"], + Path("dist", wheel).read_bytes(), + is_pr=os.environ.get("ALLOW_IDENTICAL_PR_ARTIFACT") == "true", + ) + receipt.write_text(digest + " " + wheel + "\n") + + +if __name__ == "__main__": + main() diff --git a/ci/release_version.py b/ci/release_version.py new file mode 100644 index 0000000..3b0abfb --- /dev/null +++ b/ci/release_version.py @@ -0,0 +1,39 @@ +"""Compute the published version of this fork, normalized as the wheel will be. + +Stable builds publish the declared ``__version__`` (a four-part Preset release +such as 0.7.4.1). Pull-request builds publish ``+pr..``. +PEP 440 normalizes local segments (lower case; a numeric segment loses leading +zeros), so the filename must be derived from the normalized form or it will not +match the file the build backend writes. +""" + +import re +import sys +from pathlib import Path + +from packaging.version import Version + +INIT = Path(__file__).resolve().parents[1] / "pymongosql" / "__init__.py" + + +def declared_version(source=None): + source = INIT.read_text() if source is None else source + matches = re.findall(r'^__version__: str = "([^"]+)"$', source, re.MULTILINE) + if len(matches) != 1: + raise SystemExit("expected exactly one __version__ declaration") + version = Version(matches[0]) + if str(version) != matches[0] or version.local or len(version.release) != 4: + raise SystemExit(f"__version__ must be a normalized four-part release, got {matches[0]!r}") + return matches[0] + + +def release_version(base, change_id=None, revision=None): + if change_id is None: + return str(Version(base)) + if not re.fullmatch(r"[0-9]+", change_id) or not re.fullmatch(r"[0-9a-fA-F]{7,40}", revision or ""): + raise SystemExit("pull-request builds need a numeric change id and a git revision") + return str(Version(f"{base}+pr.{change_id}.{revision}")) + + +if __name__ == "__main__": + print(release_version(declared_version(), *sys.argv[1:])) diff --git a/tests/ci/__init__.py b/tests/ci/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/ci/test_publish_wheel.py b/tests/ci/test_publish_wheel.py new file mode 100644 index 0000000..7412ec6 --- /dev/null +++ b/tests/ci/test_publish_wheel.py @@ -0,0 +1,142 @@ +"""Exercise the same immutable publisher used by Jenkins without AWS access.""" + +import hashlib +import io +import runpy +from pathlib import Path +from unittest.mock import Mock +from zipfile import ZipFile, ZipInfo + +import pytest +from botocore.exceptions import ClientError + +publisher = runpy.run_path(str(Path(__file__).resolve().parents[2] / "ci" / "publish_wheel.py")) +publish_wheel = publisher["publish_wheel"] + + +def wheel(members=None, year=2020): + output = io.BytesIO() + if members is None: + members = [("package.py", b"code"), ("metadata", b"version")] + with ZipFile(output, "w") as archive: + for name, body in members: + archive.writestr(ZipInfo(name, (year, 1, 1, 0, 0, 0)), body) + return output.getvalue() + + +def client(stored=None, error=None): + s3 = Mock() + s3.get_object.return_value = {"Body": io.BytesIO(wheel() if stored is None else stored)} + if error: + s3.put_object.side_effect = ClientError({"Error": {"Code": error}}, "PutObject") + return s3 + + +def test_new_artifact_is_conditionally_written_and_verified(): + body = wheel() + s3 = client(body) + assert publish_wheel(s3, "bucket", "key", body) == hashlib.sha256(body).hexdigest() + s3.put_object.assert_called_once_with( + Bucket="bucket", + Key="key", + Body=body, + IfNoneMatch="*", + ) + s3.get_object.assert_called_once_with(Bucket="bucket", Key="key") + + +@pytest.mark.parametrize("is_pr", [False, True]) +@pytest.mark.parametrize("variation", ["identical", "timestamps", "member_order"]) +def test_identical_content_retry_succeeds_without_overwrite(is_pr, variation): + stored = wheel() + fresh = wheel(year=2021) if variation == "timestamps" else stored + if variation == "member_order": + fresh = wheel([("metadata", b"version"), ("package.py", b"code")]) + if variation != "identical": + assert fresh != stored + s3 = client(stored, "PreconditionFailed") + assert publish_wheel(s3, "bucket", "key", fresh, is_pr) == hashlib.sha256(stored).hexdigest() + s3.put_object.assert_called_once_with( + Bucket="bucket", + Key="key", + Body=fresh, + IfNoneMatch="*", + ) + assert s3.get_object.return_value["Body"].closed + + +@pytest.mark.parametrize("is_pr", [False, True]) +@pytest.mark.parametrize("error", [None, "PreconditionFailed"]) +@pytest.mark.parametrize( + "members", + [ + [("package.py", b"changed"), ("metadata", b"version")], + [("renamed.py", b"code"), ("metadata", b"version")], + [("package.py", b"code")], + [("package.py", b"code"), ("metadata", b"version"), ("extra", b"")], + ], +) +def test_different_content_fails(is_pr, error, members): + s3 = client(wheel(members), error) + with pytest.raises(RuntimeError, match="Stored wheel differs.*refusing to overwrite") as exc: + publish_wheel(s3, "bucket", "key", wheel(), is_pr) + assert ("Bump __version__" in str(exc.value)) == (not is_pr) + assert s3.put_object.call_count == 1 + assert s3.get_object.return_value["Body"].closed + + +@pytest.mark.parametrize("is_pr", [False, True]) +@pytest.mark.parametrize("error", ["AccessDenied", "ConditionalRequestConflict"]) +def test_other_s3_errors_are_reraised(is_pr, error): + s3 = client(error=error) + with pytest.raises(ClientError) as exc: + publish_wheel(s3, "bucket", "key", wheel(), is_pr) + assert exc.value is s3.put_object.side_effect + s3.get_object.assert_not_called() + + +def test_unreadable_existing_artifact_fails_closed(): + s3 = client(error="PreconditionFailed") + s3.get_object.side_effect = ClientError({"Error": {"Code": "AccessDenied"}}, "GetObject") + with pytest.raises(ClientError, match="AccessDenied"): + publish_wheel(s3, "bucket", "key", wheel(), True) + + +@pytest.mark.parametrize("is_pr", [False, True]) +def test_invalid_existing_archive_fails_closed(is_pr): + s3 = client(b"not a zip", "PreconditionFailed") + with pytest.raises(RuntimeError, match="Stored wheel differs"): + publish_wheel(s3, "bucket", "key", wheel(), is_pr) + + +@pytest.mark.parametrize("is_pr", [False, True]) +def test_receipt_records_stored_digest(tmp_path, monkeypatch, is_pr): + stored, fresh = wheel(), wheel(year=2021) + s3 = client(stored, "PreconditionFailed") + monkeypatch.chdir(tmp_path) + monkeypatch.setenv("WHEEL", "test.whl") + monkeypatch.setenv("KEY", "test-key") + monkeypatch.setenv("ALLOW_IDENTICAL_PR_ARTIFACT", str(is_pr).lower()) + monkeypatch.setattr(publisher["boto3"], "client", lambda service: s3) + Path("dist").mkdir() + Path("dist/test.whl").write_bytes(fresh) + Path("published.sha256").write_text("stale receipt") + publisher["main"]() + assert Path("published.sha256").read_text() == hashlib.sha256(stored).hexdigest() + " test.whl\n" + + +def test_failed_retry_removes_stale_receipt(tmp_path, monkeypatch): + monkeypatch.chdir(tmp_path) + monkeypatch.setenv("WHEEL", "test.whl") + monkeypatch.setenv("KEY", "test-key") + monkeypatch.setattr( + publisher["boto3"], + "client", + lambda service: client(b"bad", "PreconditionFailed"), + ) + Path("dist").mkdir() + Path("dist/test.whl").write_bytes(wheel()) + Path("published.sha256").write_text("stale receipt") + with pytest.raises(RuntimeError): + publisher["main"]() + assert not Path("published.sha256").exists() diff --git a/tests/ci/test_release_version.py b/tests/ci/test_release_version.py new file mode 100644 index 0000000..85eaeaf --- /dev/null +++ b/tests/ci/test_release_version.py @@ -0,0 +1,40 @@ +"""The published version must equal the one the build backend writes.""" + +import runpy +from pathlib import Path + +import pytest + +module = runpy.run_path(str(Path(__file__).resolve().parents[2] / "ci" / "release_version.py")) +release_version = module["release_version"] +declared_version = module["declared_version"] + + +@pytest.mark.parametrize( + "change,revision,expected", + [ + (None, None, "0.7.4.1"), + ("3", "e2bc688f1a2b", "0.7.4.1+pr.3.e2bc688f1a2b"), + ("3", "ABCDEF123456", "0.7.4.1+pr.3.abcdef123456"), + ("3", "012345678901", "0.7.4.1+pr.3.12345678901"), + ], +) +def test_release_version_is_normalized(change, revision, expected): + assert release_version("0.7.4.1", change, revision) == expected + + +@pytest.mark.parametrize("change,revision", [("PR-3", "e2bc688"), ("3", "not-a-sha"), ("3", None)]) +def test_bad_pull_request_inputs_fail(change, revision): + with pytest.raises(SystemExit): + release_version("0.7.4.1", change, revision) + + +def test_declared_version_is_a_four_part_release(): + assert declared_version('__version__: str = "0.7.4.1"\n') == "0.7.4.1" + assert len(declared_version().split(".")) == 4 + + +@pytest.mark.parametrize("source", ['__version__: str = "0.7.4"\n', '__version__: str = "0.7.4.1+x"\n', ""]) +def test_declared_version_rejects_other_forms(source): + with pytest.raises(SystemExit): + declared_version(source) From 4bdff73f3f52a656f4a3cb364ea7ecdb82685f5e Mon Sep 17 00:00:00 2001 From: Amin Ghadersohi Date: Sat, 26 Sep 2026 03:27:32 +0000 Subject: [PATCH 09/14] chore: release 0.7.4.1 Fork release on top of 0.7.4 carrying the reflection, Decimal, qualified-column and GROUP BY/IN/LIKE/alias fixes. --- pymongosql/__init__.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/pymongosql/__init__.py b/pymongosql/__init__.py index a80c6d8..79bc12c 100644 --- a/pymongosql/__init__.py +++ b/pymongosql/__init__.py @@ -6,7 +6,7 @@ if TYPE_CHECKING: from .connection import Connection -__version__: str = "0.7.4" +__version__: str = "0.7.4.1" # Globals https://www.python.org/dev/peps/pep-0249/#globals apilevel: str = "2.0" From a71b04b78361830f185624aa509ecd1640aedc19 Mon Sep 17 00:00:00 2001 From: Amin Ghadersohi Date: Sat, 26 Sep 2026 05:37:09 +0000 Subject: [PATCH 10/14] ci: trust the checkout for git introspection during builds setuptools_scm runs git while resolving build requirements, and git refuses a workspace owned by a different uid (dubious ownership). Mark the checkout as a safe directory through the environment for the build step. --- Jenkinsfile | 2 ++ 1 file changed, 2 insertions(+) diff --git a/Jenkinsfile b/Jenkinsfile index ceb67d5..839565b 100644 --- a/Jenkinsfile +++ b/Jenkinsfile @@ -43,6 +43,8 @@ podTemplate( env.KEY = "pymongosql/${env.WHEEL}" sh ''' set -eu + # The checkout is owned by another uid; setuptools_scm runs git during builds. + export GIT_CONFIG_COUNT=1 GIT_CONFIG_KEY_0=safe.directory GIT_CONFIG_VALUE_0="$PWD" .venv/bin/pip install --no-deps -e . for attempt in $(seq 1 30); do .venv/bin/python -c "import pymongo; pymongo.MongoClient('mongodb://admin:secret@localhost:27017', serverSelectionTimeoutMS=2000).admin.command('ping')" && break From d4cb57747224821e2cdf0c074af6bcc94edfa551 Mon Sep 17 00:00:00 2001 From: Amin Ghadersohi Date: Sat, 26 Sep 2026 05:39:16 +0000 Subject: [PATCH 11/14] ci: mark the checkout safe in the pod's git config The CI image's git predates GIT_CONFIG_COUNT, so the environment override was ignored. Add the workspace to safe.directory in the ephemeral pod's global config instead. --- Jenkinsfile | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/Jenkinsfile b/Jenkinsfile index 839565b..ae09605 100644 --- a/Jenkinsfile +++ b/Jenkinsfile @@ -44,7 +44,9 @@ podTemplate( sh ''' set -eu # The checkout is owned by another uid; setuptools_scm runs git during builds. - export GIT_CONFIG_COUNT=1 GIT_CONFIG_KEY_0=safe.directory GIT_CONFIG_VALUE_0="$PWD" + # The pod is ephemeral, so its global git config is disposable. + git --version + git config --global --add safe.directory "$PWD" .venv/bin/pip install --no-deps -e . for attempt in $(seq 1 30); do .venv/bin/python -c "import pymongo; pymongo.MongoClient('mongodb://admin:secret@localhost:27017', serverSelectionTimeoutMS=2000).admin.command('ping')" && break From 2f7dba345d147c1a85095c171bde65c4ffe31fb6 Mon Sep 17 00:00:00 2001 From: Amin Ghadersohi Date: Sat, 26 Sep 2026 06:00:28 +0000 Subject: [PATCH 12/14] 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 9ad3e14e41d81a642861fe86e652efb83a6726f2 Mon Sep 17 00:00:00 2001 From: Amin Ghadersohi Date: Sat, 26 Sep 2026 06:02:30 +0000 Subject: [PATCH 13/14] 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 b83ab0b6ad2162ad100bc1727e06d4e60a971127 Mon Sep 17 00:00:00 2001 From: Amin Ghadersohi Date: Sat, 26 Sep 2026 06:03:41 +0000 Subject: [PATCH 14/14] 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)