diff --git a/README.md b/README.md index 5c57541..126fd21 100644 --- a/README.md +++ b/README.md @@ -743,6 +743,20 @@ PyMongoSQL can be used as a database driver in Apache Superset for querying and This allows seamless integration between MongoDB data and Superset's BI capabilities without requiring data migration to traditional SQL databases. +**Time grains and decimals:** + +- `DATE_TRUNC('', field)` is translated to MongoDB's `$dateTrunc` (MongoDB 5.0+), in + projections and `GROUP BY`. Units: `second`, `minute`, `hour`, `day`, `week` (starting + Sunday), `week_monday`, `month`, `quarter`, `year`, `week_ending_saturday` and + `week_ending_sunday`. Truncation is in UTC. The same function is available in the + superset-mode SQLite stage, so virtual datasets group by time the same way. +- In superset mode, a subquery's result is loaded into an in-memory SQLite database. Columns + holding `Decimal128` values are evaluated exactly there (install `pymongosql[superset]`, + which adds `sqlglot`): the column itself, `SUM`/`AVG`/`MIN`/`MAX` over it, `GROUP BY`, + `ORDER BY` and comparisons with numeric literals, with Decimal128's 34 significant + digits. Any other use of such a column (arithmetic, other functions, `DISTINCT` + aggregates) raises `NotSupportedError` rather than computing with doubles. + **Important Note on Collection Names:** When using collection names containing special characters (`.`, `-`, `:`), you must wrap them in double quotes to prevent Superset's SQL parser from incorrectly interpreting them. diff --git a/pymongosql/executor.py b/pymongosql/executor.py index 9d27b65..1b3cc4e 100644 --- a/pymongosql/executor.py +++ b/pymongosql/executor.py @@ -41,6 +41,14 @@ def _run_db_command(db: Any, command: Dict[str, Any], connection: Any, operation ) +def _paging(limit: Any, skip: Any) -> Any: + """Validate bound LIMIT/OFFSET values; returns (limit, skip).""" + for name, value in (("LIMIT", limit), ("OFFSET", skip)): + if value is not None and (isinstance(value, bool) or not isinstance(value, int) or value < 0): + raise ProgrammingError(f"{name} must be a non-negative integer, got {value!r}") + return limit, skip + + @dataclass class ExecutionContext: """Manages execution context for a single query""" @@ -149,9 +157,16 @@ def _execute_find_plan( # Replace placeholders with parameters in filter_stage only (not in projection) filter_stage = execution_plan.filter_stage or {} - if parameters: - # Positional parameters with ? (named parameters are converted to positional in execute()) - filter_stage = self._replace_placeholders(filter_stage, parameters) + # Positional parameters (named ones are converted to positional in execute()), + # in statement order: WHERE, then LIMIT, then OFFSET + bound, _ = SQLHelper.bind_filter( + {"filter": filter_stage, "limit": execution_plan.limit_stage, "skip": execution_plan.skip_stage}, + parameters, + ) + filter_stage = bound["filter"] + limit, skip = _paging(bound["limit"], bound["skip"]) + if limit == 0: + return {"cursor": {"id": 0, "firstBatch": []}, "ok": 1} projection_stage = execution_plan.projection_stage or {} @@ -170,13 +185,11 @@ def _execute_find_plan( sort_spec[field_name] = direction find_command["sort"] = sort_spec - # Apply skip if specified - if execution_plan.skip_stage: - find_command["skip"] = execution_plan.skip_stage - - # Apply limit if specified - if execution_plan.limit_stage: - find_command["limit"] = execution_plan.limit_stage + # Apply skip and limit if specified (MongoDB reads limit 0 as "no limit") + if skip: + find_command["skip"] = skip + if limit is not None: + find_command["limit"] = limit _logger.debug(f"Executing MongoDB command: {find_command}") @@ -223,7 +236,12 @@ def _execute_aggregate_plan( # Parse pipeline and options from JSON strings try: - pipeline = json.loads(execution_plan.aggregate_pipeline or "[]") + if execution_plan.aggregate_parameterized: + from bson import json_util + + pipeline = json_util.loads(execution_plan.aggregate_pipeline or "[]") + else: + pipeline = json.loads(execution_plan.aggregate_pipeline or "[]") options = json.loads(execution_plan.aggregate_options or "{}") except json.JSONDecodeError as e: raise ProgrammingError(f"Invalid JSON in aggregate pipeline or options: {e}") @@ -232,6 +250,14 @@ def _execute_aggregate_plan( _logger.debug(f"Pipeline: {pipeline}") _logger.debug(f"Options: {options}") + # A pipeline generated from SQL carries parameter markers (WHERE, HAVING), then + # LIMIT and OFFSET + limit, skip = execution_plan.limit_stage, execution_plan.skip_stage + if execution_plan.aggregate_parameterized: + bound, _ = SQLHelper.bind_filter({"pipeline": pipeline, "limit": limit, "skip": skip}, parameters) + pipeline, limit, skip = bound["pipeline"], bound["limit"], bound["skip"] + limit, skip = _paging(limit, skip) + # Get collection and call aggregate() collection = db[execution_plan.collection] @@ -258,11 +284,11 @@ def _execute_aggregate_plan( results = sorted(results, key=lambda x: x.get(field_name), reverse=reverse) # Apply skip and limit - if execution_plan.skip_stage: - results = results[execution_plan.skip_stage :] + if skip: + results = results[skip:] - if execution_plan.limit_stage: - results = results[: execution_plan.limit_stage] + if limit is not None: + results = results[:limit] # Apply projection if specified if execution_plan.projection_stage: @@ -513,11 +539,8 @@ def _execute_execution_plan( filter_conditions = execution_plan.filter_conditions or {} - # Replace placeholders in filter if parameters provided - if parameters and filter_conditions: - filter_conditions = SQLHelper.replace_placeholders_generic( - filter_conditions, parameters, execution_plan.parameter_style - ) + # Bind the WHERE clause's parameter markers; every parameter must be used + filter_conditions, _ = SQLHelper.bind_filter(filter_conditions, parameters) command = {"delete": execution_plan.collection, "deletes": [{"q": filter_conditions, "limit": 0}]} @@ -600,12 +623,10 @@ def _execute_execution_plan( # Replace placeholders if parameters provided # Note: We need to replace both update_fields and filter_conditions in one pass # to maintain correct parameter ordering (SET clause first, then WHERE clause) - if parameters: - # Combine structures for replacement in correct order - combined = {"update_fields": update_fields, "filter_conditions": filter_conditions} - replaced = SQLHelper.replace_placeholders_generic(combined, parameters, execution_plan.parameter_style) - update_fields = replaced["update_fields"] - filter_conditions = replaced["filter_conditions"] + # SET values and the WHERE clause carry parameter markers; bind them in + # statement order (SET first) and require every parameter to be used + bound, _ = SQLHelper.bind_filter({"u": update_fields, "q": filter_conditions}, parameters) + update_fields, filter_conditions = bound["u"], bound["q"] # MongoDB update command format # https://www.mongodb.com/docs/manual/reference/command/update/ diff --git a/pymongosql/helper.py b/pymongosql/helper.py index 6c1d2cb..94c2bee 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,66 @@ 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 bind_filter(value: Any, parameters: Any, exact: bool = True) -> Tuple[Any, int]: + """Bind positional parameters to the parameter markers of a translated filter. + + Translated WHERE clauses mark each ``?`` with a dict marker, so a string literal + '?' is never taken for a parameter. With ``exact``, every parameter must be used: + a parameter the translation dropped (e.g. a LIMIT it could not read) must not + silently widen the query. Returns the bound value and the number used. + """ + from .sql.where_tree import LIKE_KEY, is_like, is_param, like_operator + + params = [] if parameters is None else parameters + if isinstance(params, dict) or not isinstance(params, Sequence) or isinstance(params, (str, bytes)): + raise ProgrammingError("Positional parameters must be provided as a sequence") + idx = [0] + + def take() -> Any: + if idx[0] >= len(params): + raise ProgrammingError("Not enough parameters provided") + out = params[idx[0]] + idx[0] += 1 + return out + + def replace(val: Any) -> Any: + if is_param(val): + return SQLHelper.to_bson_value(take()) + if is_like(val): + like = val[LIKE_KEY] + parts = [] + for part in like["parts"]: + part = take() if is_param(part) else part + if not isinstance(part, str): + raise ProgrammingError(f"A LIKE pattern parameter must be a string, got {part!r}") + parts.append(part) + return like_operator("".join(parts), like["escape"], like["i"]) + if isinstance(val, dict): + return {k: replace(v) for k, v in val.items()} + if isinstance(val, list): + return [replace(v) for v in val] + return val + + bound = replace(value) + if exact and idx[0] != len(params): + raise ProgrammingError(f"{len(params)} parameters were given but the statement uses {idx[0]}") + return bound, idx[0] + @staticmethod def replace_placeholders_generic(value: Any, parameters: Any, style: Optional[str]) -> Any: """Recursively replace placeholders in nested structures for qmark or named styles.""" @@ -120,7 +183,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 +201,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/result_set.py b/pymongosql/result_set.py index c1c1e70..aaf34bb 100644 --- a/pymongosql/result_set.py +++ b/pymongosql/result_set.py @@ -4,6 +4,7 @@ from typing import Any, Dict, List, Optional, Sequence, Tuple import jmespath +from bson import Decimal128, Int64 from pymongo.errors import PyMongoError from . import STRING @@ -63,7 +64,7 @@ def _process_and_cache_batch(self, batch: List[Dict[str, Any]]) -> None: if not batch: return # Process results through projection mapping - processed_batch = [self._process_document(doc) for doc in batch] + processed_batch = [self._to_python(self._process_document(doc)) for doc in batch] # Convert dictionaries to output format (sequence or dict) formatted_batch = [self._format_result(doc) for doc in processed_batch] self._cached_results.extend(formatted_batch) @@ -159,6 +160,24 @@ def _process_document(self, doc: Dict[str, Any]) -> Dict[str, Any]: return processed + @classmethod + def _to_python(cls, value: Any) -> Any: + """Return standard Python types for BSON-specific numbers. + + DB API 2.0 consumers expect ``decimal.Decimal`` and ``int``; PyMongo returns + ``bson.Decimal128`` (which most libraries cannot sum or serialise) and, in + command responses, ``bson.Int64``. + """ + if isinstance(value, Decimal128): + return value.to_decimal() + if isinstance(value, Int64): + return int(value) + if isinstance(value, dict): + return {k: cls._to_python(v) for k, v in value.items()} + if isinstance(value, list): + return [cls._to_python(v) for v in value] + return value + def _mongo_to_bracket_key(self, field_path: str) -> str: """Convert Mongo dot-index notation to bracket notation. diff --git a/pymongosql/sql/ast.py b/pymongosql/sql/ast.py index 9d2372f..783dd30 100644 --- a/pymongosql/sql/ast.py +++ b/pymongosql/sql/ast.py @@ -4,12 +4,12 @@ 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 from .partiql.PartiQLParserVisitor import PartiQLParserVisitor -from .query_handler import QueryParseResult +from .query_handler import QueryParseResult, SelectHandler from .update_handler import UpdateParseResult _logger = logging.getLogger(__name__) @@ -275,6 +275,11 @@ 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) + if sort_spec.expr() is not None: + kind, detail = SelectHandler._classify_item(sort_spec.expr()) + if kind == "aggregate": + self._query_parse_result.sort_aggregates[field_name] = detail # Check for ASC/DESC (default is ASC = 1) direction = 1 # ASC if hasattr(sort_spec, "DESC") and sort_spec.DESC(): @@ -289,39 +294,68 @@ 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") + text = ContextUtilsMixin.normalize_field_path(key.exprSelect().getText()) + try: + truncated = SelectHandler.date_trunc(key.exprSelect()) + except ValueError as e: + self._query_parse_result.unsupported_clauses.append(f"GROUP BY {text} ({e})") + truncated = None + if truncated is not None: + self._query_parse_result.computed[text] = truncated + keys.append(text) + if ctx.PARTIAL() is not None: + self._query_parse_result.unsupported_clauses.append("GROUP PARTIAL BY") + self._query_parse_result.group_by = keys + return None + + def visitHavingClause(self, ctx: PartiQLParser.HavingClauseContext) -> Any: + """Keep the HAVING expression; it is translated after the $group stage.""" + self._query_parse_result.having = ctx.arg + return None + def visitLimitClause(self, ctx: PartiQLParser.LimitClauseContext) -> Any: - """Handle LIMIT clause for result limiting""" - _logger.debug("Processing LIMIT clause") - try: - if hasattr(ctx, "exprSelect") and ctx.exprSelect(): - limit_text = ctx.exprSelect().getText() - try: - limit_value = int(limit_text) - self._query_parse_result.limit_value = limit_value - _logger.debug(f"Extracted limit value: {limit_value}") - except ValueError as e: - _logger.warning(f"Invalid LIMIT value '{limit_text}': {e}") - return self.visitChildren(ctx) - except Exception as e: - _logger.warning(f"Error processing LIMIT clause: {e}") - return self.visitChildren(ctx) + """Handle LIMIT: a non-negative integer literal or a bound parameter.""" + from .where_tree import is_param, operand + + if hasattr(ctx, "exprSelect") and ctx.exprSelect(): + try: + value = operand(ctx.exprSelect()) + except Exception: + value = None + if is_param(value) or (isinstance(value, int) and not isinstance(value, bool) and value >= 0): + self._query_parse_result.limit_value = value + else: + # Dropping it would return every row + text = ctx.exprSelect().getText() + self._query_parse_result.unsupported_clauses.append( + f"LIMIT {text} (needs a non-negative integer or a parameter)" + ) + return None def visitOffsetByClause(self, ctx: PartiQLParser.OffsetByClauseContext) -> Any: - """Handle OFFSET clause for result skipping""" - _logger.debug("Processing OFFSET clause") - try: - if hasattr(ctx, "exprSelect") and ctx.exprSelect(): - offset_text = ctx.exprSelect().getText() - try: - offset_value = int(offset_text) - self._query_parse_result.offset_value = offset_value - _logger.debug(f"Extracted offset value: {offset_value}") - except ValueError as e: - _logger.warning(f"Invalid OFFSET value '{offset_text}': {e}") - return self.visitChildren(ctx) - except Exception as e: - _logger.warning(f"Error processing OFFSET clause: {e}") - return self.visitChildren(ctx) + """Handle OFFSET: a non-negative integer literal or a bound parameter.""" + from .where_tree import is_param, operand + + if hasattr(ctx, "exprSelect") and ctx.exprSelect(): + try: + value = operand(ctx.exprSelect()) + except Exception: + value = None + if is_param(value) or (isinstance(value, int) and not isinstance(value, bool) and value >= 0): + self._query_parse_result.offset_value = value + else: + # Dropping it would return every row + text = ctx.exprSelect().getText() + self._query_parse_result.unsupported_clauses.append( + f"OFFSET {text} (needs a non-negative integer or a parameter)" + ) + return None def visitUpdateClause(self, ctx: PartiQLParser.UpdateClauseContext) -> Any: """Handle UPDATE clause to extract collection/table name.""" diff --git a/pymongosql/sql/builder.py b/pymongosql/sql/builder.py index d9dff9d..db8cd73 100644 --- a/pymongosql/sql/builder.py +++ b/pymongosql/sql/builder.py @@ -2,7 +2,9 @@ import json import logging from dataclasses import dataclass -from typing import TYPE_CHECKING, Any, Dict, Optional, Union +from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union + +from bson import json_util if TYPE_CHECKING: from .delete_builder import DeleteExecutionPlan @@ -110,19 +112,75 @@ 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 + # 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: + 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: + 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"]) + parse_result.group_by = [strip(name) for name in parse_result.group_by] + for expression in parse_result.computed.values(): + expression["field"] = strip(expression["field"]) + for item in parse_result.select_items: + if "field" in item: + item["field"] = strip(item["field"]) + @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.), GROUP BY and HAVING + if parse_result.aggregate_functions or parse_result.group_by or parse_result.having is not None: return ExecutionPlanBuilder._build_sql_aggregate_plan(parse_result) + if any("computed" in item for item in parse_result.select_items): + return ExecutionPlanBuilder._build_computed_plan(parse_result) + + # ORDER BY may name a column by its SELECT alias; find() sorts on the field + field_for_alias = {alias: name for name, alias in parse_result.column_aliases.items()} + 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: @@ -134,9 +192,139 @@ def _build_query_plan(parse_result: "QueryParseResult") -> "QueryExecutionPlan": plan = builder.build() return plan + @staticmethod + def _build_computed_plan(parse_result: "QueryParseResult") -> "QueryExecutionPlan": + """A SELECT with computed columns (DATE_TRUNC) and no grouping. + + Pipeline: $match (WHERE), $addFields (computed columns), $sort, then $project of + the SELECT list in order; OFFSET/LIMIT are applied after binding. + """ + from ..superset_mongodb.time_grain import mongo_expression + + builder = BuilderFactory.create_query_builder().collection(parse_result.collection) + pipeline: List[Dict[str, Any]] = [] + if parse_result.filter_conditions: + pipeline.append({"$match": parse_result.filter_conditions}) + added: Dict[str, Any] = {} + project: Dict[str, Any] = {} + sources: Dict[str, str] = {} # names ORDER BY may use -> sortable field + outputs = [] + for index, item in enumerate(parse_result.select_items): + if "computed" in item: + expression = parse_result.computed[item["computed"]] + hidden = f"__computed{index}" + added[hidden] = mongo_expression(expression["unit"], expression["field"]) + source, text = hidden, item["computed"] + else: + source = text = item["field"] + output = item["alias"] or text + project[output] = f"${source}" + sources[output] = sources[text] = sources[text.upper()] = source + outputs.append(output) + if "_id" not in project: + project["_id"] = 0 + if added: + pipeline.append({"$addFields": added}) + sort = {} + for spec in parse_result.sort_fields: + for name, direction in spec.items(): + sort[sources.get(name, sources.get(name.upper(), name))] = direction + if sort: + pipeline.append({"$sort": sort}) + pipeline.append({"$project": project}) + builder.skip(parse_result.offset_value).limit(parse_result.limit_value) + builder._execution_plan.is_aggregate_query = True + builder._execution_plan.aggregate_parameterized = True + builder._execution_plan.aggregate_pipeline = json_util.dumps(pipeline) + builder._execution_plan.aggregate_options = json.dumps({}) + builder._execution_plan.projection_stage = {name: 1 for name in outputs} + return builder.build() + + @staticmethod + def _aggregate_source(func_info: Dict[str, Any], key: str, index: int) -> Any: + """$project expression for an accumulator: the value, or the reduced DISTINCT set. + + SUM of no (non-NULL) values is NULL, not $sum's 0. + """ + if not func_info.get("distinct"): + if func_info["function"] == "SUM": + return {"$cond": [{"$gt": [f"$__numbers{index}", 0]}, f"${key}", None]} + return f"${key}" + # SQL ignores NULL in DISTINCT aggregates + values = {"$setDifference": [f"${key}", [None]]} + reducer = {"COUNT": "$size", "SUM": "$sum", "AVG": "$avg", "MIN": "$min", "MAX": "$max"} + if func_info["function"] == "SUM": + numbers = {"$filter": {"input": values, "cond": {"$isNumber": "$$this"}}} + return {"$cond": [{"$gt": [{"$size": numbers}, 0]}, {"$sum": values}, None]} + return {reducer[func_info["function"]]: values} + + @staticmethod + def _translate_having(parse_result: "QueryParseResult", group_keys: Dict[str, str], hidden: Dict[str, Any]) -> Any: + """Translate HAVING into a $match on the grouped outputs (SQL three-valued logic).""" + from .partiql.PartiQLParser import PartiQLParser + from .query_handler import SelectHandler + from .where_tree import _PATH_NODES, WhereTreeBuilder, _field_path + + prefixes = [f"{q}." for q in (parse_result.collection_alias, parse_result.collection) if q] + + def strip(name: str) -> str: + for prefix in prefixes: + if name.startswith(prefix) and len(name) > len(prefix): + return name[len(prefix) :] + return name + + outputs = {} + for item in parse_result.select_items: + if "aggregate" in item: + info = parse_result.aggregate_functions[item["aggregate"]] + outputs[(info["function"], info["argument"], bool(info.get("distinct")))] = info["alias"] + else: + name = item.get("field") or item["computed"] + outputs[name] = item["alias"] or name + aliases = set(outputs.values()) + + def resolve(node: Any) -> Any: + if isinstance(node, (PartiQLParser.CountAllContext, PartiQLParser.AggregateBaseContext)): + kind, detail = SelectHandler._classify_item(node) + if kind != "aggregate": + raise ValueError(f"Unsupported aggregate in HAVING: {node.getText()}") + func, arg, distinct = detail + signature = (func, strip(arg) if arg != "*" else arg, distinct) + if signature in outputs: + return outputs[signature] + name = f"__having{len(hidden)}" + parse_result.aggregate_functions.append( + {"function": func, "argument": signature[1], "distinct": distinct, "alias": name, "expression": ""} + ) + hidden[name] = len(parse_result.aggregate_functions) - 1 + outputs[signature] = name + return name + if isinstance(node, _PATH_NODES): + name = strip(_field_path(node)) + if name in aliases: + return name + if name in outputs: + return outputs[name] + if name in group_keys: + hidden_name = f"__having{len(hidden)}" + hidden[hidden_name] = f"$_id.{group_keys[name]}" + outputs[name] = hidden_name + return hidden_name + raise ValueError(f"HAVING column '{name}' must be grouped, aggregated or a SELECT alias") + return None + + return WhereTreeBuilder(resolver=resolve).build(parse_result.having) + @staticmethod def _build_sql_aggregate_plan(parse_result: "QueryParseResult") -> "QueryExecutionPlan": - """Build an aggregate execution plan from SQL aggregate functions 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", @@ -153,37 +341,144 @@ 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"] + from ..superset_mongodb.time_grain import mongo_expression + + group_keys = {name: f"g{i}" for i, name in enumerate(parse_result.group_by)} + + def key_source(name: str) -> Any: + computed = parse_result.computed.get(name) + return mongo_expression(computed["unit"], computed["field"]) if computed else f"${name}" + + group_stage = {"_id": {key: key_source(name) for name, key in group_keys.items()} if group_keys else None} + + # HAVING may name select-list outputs, grouped columns or aggregates; the ones + # not in the SELECT list are computed as hidden outputs and removed afterwards. + hidden: Dict[str, Any] = {} + having_filter = None + if parse_result.having is not None: + having_filter = ExecutionPlanBuilder._translate_having(parse_result, group_keys, hidden) + + # ORDER BY may name an aggregate that is not selected (e.g. a top-N query ordered + # by another metric); it is computed as a hidden output like HAVING's. + order_outputs: Dict[str, str] = {} + prefixes = [f"{q}." for q in (parse_result.collection_alias, parse_result.collection) if q] + for text, (func, arg, distinct) in parse_result.sort_aggregates.items(): + for prefix in prefixes: + if arg.startswith(prefix) and len(arg) > len(prefix): + arg = arg[len(prefix) :] + selected = [ + info["alias"] + for info in parse_result.aggregate_functions + if (info["function"], info["argument"], bool(info.get("distinct"))) == (func, arg, distinct) + and not info["alias"].startswith("__") + ] + if selected: + order_outputs[text] = selected[0] + continue + name = f"__order{len(order_outputs)}" + parse_result.aggregate_functions.append( + {"function": func, "argument": arg, "distinct": distinct, "alias": name, "expression": ""} + ) + hidden[name] = len(parse_result.aggregate_functions) - 1 + order_outputs[text] = name + + accumulator_keys = [] + for i, func_info in enumerate(parse_result.aggregate_functions): func_name = func_info["function"] 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 func_info.get("distinct") or "." in key or key.startswith("$") or key == "_id" or key in group_stage: + key = f"__agg{i}" + accumulator_keys.append(key) + + if func_info.get("distinct"): + # Collect the distinct values; the $project stage reduces the set + group_stage[key] = {"$addToSet": f"${arg}"} + elif func_name == "COUNT" and arg == "*": + group_stage[key] = {accumulator: 1} + elif func_name == "COUNT": + # COUNT(field) counts documents where the field is present and not null + group_stage[key] = {"$sum": {"$cond": [{"$gt": [f"${arg}", None]}, 1, 0]}} else: - group_stage[alias] = {accumulator: f"${arg}"} - - pipeline.append({"$group": group_stage}) - - # Add $project to exclude _id + group_stage[key] = {accumulator: f"${arg}"} + if func_name == "SUM": + # $sum of no numbers is 0; SQL's SUM of no values is NULL + group_stage[f"__numbers{i}"] = {"$sum": {"$cond": [{"$isNumber": f"${arg}"}, 1, 0]}} + + if group_keys: + pipeline.append({"$group": group_stage}) + else: + # An aggregate without GROUP BY returns one row even for no input rows + # (COUNT 0, other aggregates NULL); $group alone would return none. + empty = {"_id": None} + for name, accumulator_spec in group_stage.items(): + if name != "_id": + operator = next(iter(accumulator_spec)) + empty[name] = {"$sum": 0, "$addToSet": []}.get(operator) + pipeline.append({"$facet": {"row": [{"$group": group_stage}]}}) + pipeline.append( + {"$project": {"row": {"$cond": [{"$eq": [{"$size": "$row"}, 0]}, {"$literal": [empty]}, "$row"]}}} + ) + pipeline.append({"$unwind": "$row"}) + pipeline.append({"$replaceRoot": {"newRoot": "$row"}}) + + # Map every SELECT item, in order, to its output name and source project_stage = {"_id": 0} - 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 = ExecutionPlanBuilder._aggregate_source(func_info, key, item["aggregate"]) + if source == f"${key}" and key == output: + source = 1 + output_for[func_info["expression"].upper()] = output + else: + name = item.get("field") or item["computed"] + 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) + for name, source in hidden.items(): + if isinstance(source, int): # a hidden aggregate: index into aggregate_functions + project_stage[name] = ExecutionPlanBuilder._aggregate_source( + parse_result.aggregate_functions[source], accumulator_keys[source], source + ) + else: + project_stage[name] = source pipeline.append({"$project": project_stage}) + if having_filter is not None: + pipeline.append({"$match": having_filter}) + + 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(), order_outputs.get(name))) + if output is None: + raise NotSupportedError(f"ORDER BY '{name}' must name a selected column or its alias") + sort_stage[output] = direction + if sort_stage: + pipeline.append({"$sort": sort_stage}) + if hidden: + pipeline.append({"$project": {name: 0 for name in hidden}}) + # OFFSET/LIMIT (integers or parameters) are applied by the executor after binding + builder.skip(parse_result.offset_value).limit(parse_result.limit_value) # Configure the execution plan as an aggregate query builder._execution_plan.is_aggregate_query = True - builder._execution_plan.aggregate_pipeline = json.dumps(pipeline) + builder._execution_plan.aggregate_parameterized = True + # Extended JSON keeps Decimal128/datetime literals and parameter markers intact + builder._execution_plan.aggregate_pipeline = json_util.dumps(pipeline) builder._execution_plan.aggregate_options = json.dumps({}) - # Set projection for ResultSet description - 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 @@ -212,6 +507,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}" @@ -226,6 +526,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/explain_builder.py b/pymongosql/sql/explain_builder.py index 204b3dd..f643080 100644 --- a/pymongosql/sql/explain_builder.py +++ b/pymongosql/sql/explain_builder.py @@ -91,8 +91,7 @@ def build_inner_command( raise ProgrammingError("No collection specified in query") filter_stage = inner_plan.filter_stage or {} - if parameters: - filter_stage = SQLHelper.replace_placeholders_generic(filter_stage, parameters, "qmark") + filter_stage, _ = SQLHelper.bind_filter(filter_stage, parameters) command = {"find": inner_plan.collection, "filter": filter_stage} diff --git a/pymongosql/sql/handler.py b/pymongosql/sql/handler.py index e01370d..8f1b3ae 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 @@ -250,19 +257,27 @@ 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": - 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 +302,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 +369,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 +411,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 +558,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: @@ -533,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/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/pymongosql/sql/query_builder.py b/pymongosql/sql/query_builder.py index fb3b7cf..2f59075 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""" @@ -51,10 +53,20 @@ def validate(self) -> bool: else: errors = self.validate_base() - if self.limit_stage is not None and (not isinstance(self.limit_stage, int) or self.limit_stage < 0): + from .where_tree import is_param + + if ( + self.limit_stage is not None + and not is_param(self.limit_stage) + and (not isinstance(self.limit_stage, int) or self.limit_stage < 0) + ): errors.append("Limit must be a non-negative integer") - if self.skip_stage is not None and (not isinstance(self.skip_stage, int) or self.skip_stage < 0): + if ( + self.skip_stage is not None + and not is_param(self.skip_stage) + and (not isinstance(self.skip_stage, int) or self.skip_stage < 0) + ): errors.append("Skip must be a non-negative integer") if errors: @@ -76,6 +88,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, ) @@ -153,18 +166,22 @@ def sort(self, specs: List[Dict[str, int]]) -> "MongoQueryBuilder": return self - def limit(self, count: int) -> "MongoQueryBuilder": - """Set limit for results""" - if not isinstance(count, int) or count < 0: + def limit(self, count: Any) -> "MongoQueryBuilder": + """Set limit for results (an integer or a parameter marker bound at execution)""" + from .where_tree import is_param + + if not is_param(count) and (not isinstance(count, int) or count < 0): return self self._execution_plan.limit_stage = count _logger.debug(f"Set limit to: {count}") return self - def skip(self, count: int) -> "MongoQueryBuilder": - """Set skip count for pagination""" - if not isinstance(count, int) or count < 0: + def skip(self, count: Any) -> "MongoQueryBuilder": + """Set skip count for pagination (an integer or a parameter marker bound at execution)""" + from .where_tree import is_param + + if not is_param(count) and (not isinstance(count, int) or count < 0): return self self._execution_plan.skip_stage = count diff --git a/pymongosql/sql/query_handler.py b/pymongosql/sql/query_handler.py index 3a09db8..d2a9d02 100644 --- a/pymongosql/sql/query_handler.py +++ b/pymongosql/sql/query_handler.py @@ -34,6 +34,20 @@ 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 (or the text of a computed expression, see ``computed``) + group_by: List[str] = field(default_factory=list) + # Computed expressions by their SQL text: {"unit": ..., "field": ...} for DATE_TRUNC + computed: Dict[str, Dict[str, str]] = field(default_factory=dict) + # ORDER BY aggregates by their SQL text: (function, argument, distinct) + sort_aggregates: Dict[str, Tuple[str, str, bool]] = field(default_factory=dict) + # Clauses that are parsed but cannot be translated faithfully + unsupported_clauses: List[str] = field(default_factory=list) + # FROM alias (FROM users AS u / FROM users u) + collection_alias: Optional[str] = None + # HAVING expression (parse-tree node), translated after grouping + having: Any = None # Subquery info (for wrapped subqueries, e.g., Superset outering) subquery_plan: Optional[Any] = None @@ -131,21 +145,32 @@ def handle_visitor(self, ctx: PartiQLParser.SelectItemsContext, parse_result: "Q if hasattr(ctx, "projectionItems") and ctx.projectionItems(): for item in ctx.projectionItems().projectionItem(): field_name, alias = self._extract_field_and_alias(item) + kind, detail = self._classify_item(item) - # Check if this is an aggregate function (COUNT, SUM, etc.) - agg_match = self._AGGREGATE_PATTERN.match(field_name) - if agg_match: - func_name = agg_match.group(1).upper() - func_arg = agg_match.group(2) + if kind == "aggregate": + func_name, func_arg, distinct = detail + parse_result.select_items.append({"aggregate": len(parse_result.aggregate_functions)}) parse_result.aggregate_functions.append( { "function": func_name, "argument": func_arg, + "distinct": distinct, "alias": alias or field_name, + "expression": field_name, } ) continue + if kind == "computed": + parse_result.computed[field_name] = detail + parse_result.select_items.append({"computed": field_name, "alias": alias}) + continue + if kind == "unsupported": + # e.g. a + 1 or lower(a): projecting it as a field would silently read NULL + parse_result.unsupported_clauses.append(f"SELECT {detail}") + continue + field_name = detail + parse_result.select_items.append({"field": field_name, "alias": alias}) # Use MongoDB standard projection format: {field: 1} to include field projection[field_name] = 1 # Store alias if present @@ -156,6 +181,58 @@ def handle_visitor(self, ctx: PartiQLParser.SelectItemsContext, parse_result: "Q parse_result.column_aliases = column_aliases return projection + _AGGREGATES = ("COUNT", "SUM", "AVG", "MIN", "MAX") + + @staticmethod + def date_trunc(node: Any) -> Optional[Dict[str, str]]: + """{"unit", "field"} for DATE_TRUNC('', ), else None.""" + from ..superset_mongodb.time_grain import UNITS + from .where_tree import _PATH_NODES, _field_path, _string_literal, _unwrap + + node = _unwrap(node) + if not isinstance(node, PartiQLParser.FunctionCallContext): + return None + if node.functionName().getText().lower() != "date_trunc" or len(node.expr()) != 2: + return None + unit, target = _unwrap(node.expr()[0]), _unwrap(node.expr()[1]) + if not isinstance(unit, PartiQLParser.LiteralStringContext) or not isinstance(target, _PATH_NODES): + raise ValueError(f"DATE_TRUNC needs a unit literal and a field: {node.getText()}") + name = _string_literal(unit.getText()).lower() + if name not in UNITS: + raise ValueError(f"Unsupported DATE_TRUNC unit {name!r}; use one of {', '.join(UNITS)}") + return {"unit": name, "field": _field_path(target)} + + @staticmethod + def _classify_item(item) -> Tuple[str, Any]: + """("field", path) | ("aggregate", (function, argument, distinct)) | ("unsupported", text).""" + from .where_tree import _PATH_NODES, _field_path, _unwrap + + # A projection item's first child is its expression; other nodes are classified as-is + is_item = isinstance(item, PartiQLParser.ProjectionItemContext) + expr = item.children[0] if is_item and getattr(item, "children", None) else item + if not hasattr(expr, "getRuleIndex"): + return "unsupported", str(expr) + node = _unwrap(expr) + try: + if isinstance(node, PartiQLParser.CountAllContext): + return "aggregate", ("COUNT", "*", False) + if isinstance(node, PartiQLParser.AggregateBaseContext): + func = node.func.text.upper() + quantifier = node.setQuantifierStrategy() + argument = _unwrap(node.expr()) + if func not in SelectHandler._AGGREGATES or not isinstance(argument, _PATH_NODES): + return "unsupported", node.getText() + distinct = quantifier is not None and quantifier.getText().upper() == "DISTINCT" + return "aggregate", (func, _field_path(argument), distinct) + if isinstance(node, _PATH_NODES): + return "field", _field_path(node) + truncated = SelectHandler.date_trunc(node) + if truncated is not None: + return "computed", truncated + except Exception as e: + return "unsupported", f"{node.getText()} ({e})" + return "unsupported", node.getText() + def _extract_field_and_alias(self, item) -> Tuple[str, Optional[str]]: """Extract field name and alias from projection item context with nested field support""" if not hasattr(item, "children") or not item.children: @@ -182,7 +259,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): @@ -260,6 +337,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(): @@ -282,12 +388,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 @@ -305,16 +415,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..f6fbfec 100644 --- a/pymongosql/sql/update_handler.py +++ b/pymongosql/sql/update_handler.py @@ -133,15 +133,27 @@ def _extract_set_assignment(self, ctx: Any) -> tuple[Optional[str], Any]: field_name = None field_value = None - # Extract field name from pathSimple + # Extract field name from pathSimple ("my field" -> my field, a[0] -> a.0) if hasattr(ctx, "pathSimple") and ctx.pathSimple(): - field_name = ctx.pathSimple().getText() + from .handler import ContextUtilsMixin - # Extract value from expr + field_name = ContextUtilsMixin.normalize_field_path(ctx.pathSimple().getText()) + + # Extract value from expr: a literal, value function or parameter marker, read + # from the parse tree so a string literal '?' stays a string if hasattr(ctx, "expr") and ctx.expr(): - expr_text = ctx.expr().getText() - # Parse the expression to get the actual value - field_value = self._parse_value(expr_text) + from ..error import NotSupportedError + from .where_tree import _coerce, _Field, operand + + try: + field_value = operand(ctx.expr()) + except NotSupportedError: + field_value = self._parse_value(ctx.expr().getText()) + if field_value == "?": + raise + if isinstance(field_value, _Field): + raise NotSupportedError(f"SET needs a value, not a column: {ctx.getText()}") + field_value = _coerce(field_value, as_double=True) # a fractional literal is stored as a double return field_name, field_value except Exception as e: @@ -190,6 +202,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..6e617df --- /dev/null +++ b/pymongosql/sql/where_tree.py @@ -0,0 +1,450 @@ +# -*- 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. + +Predicates are read from the parse tree, never from concatenated token text: the +field is the path node on one side of the operator and the value is the literal, +parameter or value function on the other. Anything else raises +``NotSupportedError`` rather than matching different documents. +""" + +import datetime +import re +from typing import Any, Dict, List, Optional, Tuple + +from bson import Decimal128 + +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} + +# Stands for a bound parameter (``?``) inside a translated filter. A dict cannot be +# produced by any SQL literal, so a string literal '?' is never mistaken for it. +PARAM_KEY = "$pymongosqlParam" + + +def param_marker() -> Dict[str, Any]: + return {PARAM_KEY: True} + + +def is_param(value: Any) -> bool: + return isinstance(value, dict) and list(value) == [PARAM_KEY] + + +# A LIKE whose pattern contains a bound parameter ('%' || ? || '%', as SQLAlchemy renders +# contains()); the regex is built when the parameters are bound. +LIKE_KEY = "$pymongosqlLike" + + +def is_like(value: Any) -> bool: + return isinstance(value, dict) and list(value) == [LIKE_KEY] + + +def like_operator(pattern: str, escape: Optional[str], case_insensitive: bool) -> Dict[str, Any]: + """The $regex operator document for a LIKE pattern.""" + regex, dotall = _like_regex(pattern, escape) + options = ("s" if dotall else "") + ("i" if case_insensitive else "") + return {"$regex": regex, "$options": options} if options else {"$regex": regex} + + +class _Pattern: + """A LIKE pattern built from string literals and bound parameters.""" + + def __init__(self, parts: List[Any]): + self.parts = parts + + +_LEAVES = ( + PartiQLParser.PredicateComparisonContext, + PartiQLParser.PredicateIsContext, + PartiQLParser.PredicateInContext, + PartiQLParser.PredicateLikeContext, + PartiQLParser.PredicateBetweenContext, +) +_PATH_NODES = ( + PartiQLParser.VariableIdentifierContext, + PartiQLParser.VariableKeywordContext, + PartiQLParser.ExprPrimaryPathContext, +) +_SWAP = {"<": ">=", ">=": "<", ">": "<=", "<=": ">"} +_MIRROR = {"<": ">", ">": "<", "<=": ">=", ">=": "<=", "=": "=", "!=": "!=", "<>": "<>"} +_MONGO = {"<": "$lt", "<=": "$lte", ">": "$gt", ">=": "$gte"} + + +class _Field(str): + """A document field path (as opposed to a string value).""" + + +class _InexactDecimal: + """A decimal literal that no double represents exactly (e.g. 0.1). + + SQL compares a literal in the column's type. A double field is compared with the + nearest double and every other numeric type with the exact Decimal128, so neither + a double 0.1 nor a Decimal128 0.1 is missed. + """ + + def __init__(self, text: str): + self.as_double = float(text) + self.as_decimal = Decimal128(text) + + def __neg__(self) -> "_InexactDecimal": + negated = _InexactDecimal("0") + negated.as_double, negated.as_decimal = -self.as_double, Decimal128(-self.as_decimal.to_decimal()) + return negated + + +def _decimal_literal(text: str) -> Any: + from decimal import Decimal + + exact = Decimal(text) + as_double = float(text) + return as_double if Decimal(as_double) == exact else _InexactDecimal(text) + + +def _coerce(value: Any, as_double: bool) -> Any: + if isinstance(value, _InexactDecimal): + return value.as_double if as_double else value.as_decimal + if isinstance(value, (list, tuple)): + return type(value)(_coerce(v, as_double) for v in value) + return value + + +def _has_inexact(value: Any) -> bool: + if isinstance(value, (list, tuple)): + return any(_has_inexact(v) for v in value) + return isinstance(value, _InexactDecimal) + + +def _all(parts: List[Tuple[Any, Filter]], key: str, chain: type) -> Filter: + """Combine filters under ``key``, flattening only an unparenthesized chain of the same operator.""" + items: List[Filter] = [] + 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 _unwrap(ctx: Any) -> Any: + """Descend through single-child pass-through rules (MathOp00 > ... > ExprTermBase).""" + while True: + children = getattr(ctx, "children", None) or [] + if len(children) == 1 and hasattr(children[0], "getRuleIndex"): + ctx = children[0] + else: + return ctx + + +def _string_literal(token_text: str) -> str: + return token_text[1:-1].replace("''", "'") + + +def _field_path(ctx: Any) -> str: + """Dot path for a variable reference or path expression, quotes removed.""" + if isinstance(ctx, (PartiQLParser.VariableIdentifierContext, PartiQLParser.VariableKeywordContext)): + if getattr(ctx, "qualifier", None) is not None: + raise NotSupportedError(f"Unsupported variable reference: {ctx.getText()}") + text = ctx.getText() + return text[1:-1].replace('""', '"') if text.startswith('"') else text + parts = [_field_path(_unwrap(ctx.getChild(0)))] + for step in ctx.children[1:]: + if isinstance(step, PartiQLParser.PathStepDotExprContext): + key = step.key.getText() + parts.append(key[1:-1].replace('""', '"') if key.startswith('"') else key) + elif isinstance(step, PartiQLParser.PathStepIndexExprContext): + key = _unwrap(step.key) + if isinstance(key, PartiQLParser.LiteralIntegerContext): + parts.append(key.getText()) + elif isinstance(key, PartiQLParser.LiteralStringContext): + parts.append(_string_literal(key.getText())) + else: + raise NotSupportedError(f"Unsupported path step: {step.getText()}") + else: + raise NotSupportedError(f"Unsupported path step: {step.getText()}") + path = ".".join(parts) + # A quoted identifier with dots ("user.name") keeps this driver's nested-path meaning + if any(not segment or segment.startswith("$") for segment in path.split(".")): + raise NotSupportedError(f"Unsupported field name: {path!r}") + return path + + +def operand(ctx: Any, resolver: Any = None) -> Any: + """Evaluate one side of a predicate: a _Field, a Python value or a parameter marker. + + ``resolver`` may map a node (e.g. an aggregate call in HAVING) to a field name. + """ + node = _unwrap(ctx) + if resolver is not None: + resolved = resolver(node) + if resolved is not None: + return _Field(resolved) + if isinstance(node, _PATH_NODES): + return _Field(_field_path(node)) + if isinstance(node, PartiQLParser.ParameterContext): + return param_marker() + if isinstance(node, PartiQLParser.LiteralStringContext): + return _string_literal(node.getText()) + if isinstance(node, PartiQLParser.LiteralIntegerContext): + return int(node.getText()) + if isinstance(node, PartiQLParser.LiteralDecimalContext): + return _decimal_literal(node.getText()) + if isinstance(node, PartiQLParser.LiteralTrueContext): + return True + if isinstance(node, PartiQLParser.LiteralFalseContext): + return False + if isinstance(node, (PartiQLParser.LiteralNullContext, PartiQLParser.LiteralMissingContext)): + return None + if isinstance(node, PartiQLParser.LiteralDateContext): + return datetime.datetime.fromisoformat(_string_literal(node.LITERAL_STRING().getText())) + if isinstance(node, PartiQLParser.ValueExprContext) and node.sign is not None: + value = operand(node.rhs) + if isinstance(value, bool) or not isinstance(value, (int, float, _InexactDecimal)): + raise NotSupportedError(f"Unsupported signed expression: {node.getText()}") + return value if node.sign.text == "+" else -value + if isinstance(node, PartiQLParser.MathOp00Context) and node.op is not None and node.op.text == "||": + parts = [] + for side in (operand(node.lhs), operand(node.rhs)): + if isinstance(side, _Pattern): + parts.extend(side.parts) + elif (isinstance(side, str) and not isinstance(side, _Field)) or is_param(side): + parts.append(side) + else: + raise NotSupportedError(f"Only strings and parameters can be concatenated: {node.getText()}") + if all(isinstance(p, str) for p in parts): + return "".join(parts) + return _Pattern(parts) + if isinstance(node, PartiQLParser.FunctionCallContext): + from .value_function_registry import get_default_registry + + name = node.functionName().getText() + registry = get_default_registry() + if registry.has_function(name): + args = [_coerce(operand(arg), as_double=True) for arg in node.expr()] + if any(isinstance(a, _Field) or is_param(a) for a in args): + raise NotSupportedError(f"Value functions take literal arguments: {node.getText()}") + return registry.execute(name, args) + raise NotSupportedError(f"Unsupported WHERE operand: {node.getText()}") + + +def _like_regex(pattern: str, escape: Optional[str]) -> Tuple[str, bool]: + """Regex for a LIKE pattern and whether it needs DOTALL. + + ``escape`` makes the next character literal. A leading or trailing ``%`` leaves + that end unanchored; a wildcard anywhere else must also match newlines. + """ + tokens, i = [], 0 + while i < len(pattern): + char = pattern[i] + if escape is not None and char == escape: + if i + 1 >= len(pattern): + raise NotSupportedError("LIKE pattern ends with its escape character") + tokens.append(re.escape(pattern[i + 1])) + i += 2 + continue + tokens.append(".*" if char == "%" else "." if char == "_" else re.escape(char)) + i += 1 + body = "".join(tokens) + inner = tokens[1 if tokens[:1] == [".*"] else 0 : len(tokens) - (1 if tokens[-1:] == [".*"] else 0)] + needs_dotall = any(t in (".*", ".") for t in inner) + regex = ("" if body.startswith(".*") else "^") + body + ("" if body.endswith(".*") else "$") + return regex, needs_dotall + + +def leaf_filters(field: str, operator: str, value: Any, escape: Optional[str] = None) -> Pair: + """TRUE and FALSE filters for one predicate on ``field``.""" + if _has_inexact(value): + double = {field: {"$type": "double"}} + other = {field: {"$not": {"$type": "double"}}} + t1, f1 = leaf_filters(field, operator, _coerce(value, True), escape) + t2, f2 = leaf_filters(field, operator, _coerce(value, False), escape) + return ( + {"$or": [{"$and": [double, t1]}, {"$and": [other, t2]}]}, + {"$or": [{"$and": [double, f1]}, {"$and": [other, f2]}]}, + ) + op = operator.upper() + if op == "IS NULL": + return {field: {"$eq": None}}, {field: {"$ne": None}} + if op == "IS NOT NULL": + return {field: {"$ne": None}}, {field: {"$eq": None}} + if op == "IS MISSING": + return {field: {"$exists": False}}, {field: {"$exists": True}} + if op == "IS NOT MISSING": + return {field: {"$exists": True}}, {field: {"$exists": False}} + if op in ("IN", "NOT IN"): + values = value if isinstance(value, list) else [value] + present = [v for v in values if v is not None] + 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", "ILIKE", "NOT ILIKE"): + case_insensitive = "ILIKE" in op + if isinstance(value, str): + regex: Dict[str, Any] = like_operator(value, escape, case_insensitive) + elif is_param(value) or isinstance(value, _Pattern): + parts = value.parts if isinstance(value, _Pattern) else [value] + regex = {LIKE_KEY: {"parts": parts, "escape": escape, "i": case_insensitive}} + else: + raise NotSupportedError("LIKE needs a string pattern") + true = {field: regex} + false = {"$and": [{field: {"$not": regex}}, {field: {"$ne": None}}]} + return (true, false) if op in ("LIKE", "ILIKE") else (false, true) + if op in ("BETWEEN", "NOT BETWEEN"): + low, high = value + true = {"$and": [{field: {"$gte": low}}, {field: {"$lte": high}}]} + false = {"$or": [{field: {"$lt": low}}, {field: {"$gt": high}}]} + return (true, false) if op == "BETWEEN" else (false, true) + if value is None: + # This dialect has always read "= NULL" / "<> NULL" as IS NULL / IS NOT NULL; + # any other comparison with NULL is UNKNOWN. + 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}") + + +def _value(ctx: Any, resolver: Any = None) -> Any: + value = operand(ctx, resolver) + if isinstance(value, _Field): + raise NotSupportedError(f"Comparing two fields is not supported: {ctx.getText()}") + return value + + +class WhereTreeBuilder: + """Build a MongoDB filter for a WHERE (or HAVING) expression from its parse tree.""" + + def __init__(self, resolver: Any = None): + self._resolver = resolver + + def build(self, ctx: Any) -> Filter: + return self._pair(ctx)[0] + + def _pair(self, ctx: Any, substitute: Optional[Tuple[Any, Pair]] = None) -> Pair: + if substitute is not None and ctx is substitute[0]: + return substitute[1] + if isinstance(ctx, PartiQLParser.NotContext): + true, false = self._pair(ctx.rhs, substitute) + return false, true + if isinstance(ctx, (PartiQLParser.AndContext, PartiQLParser.OrContext)): + is_and = isinstance(ctx, PartiQLParser.AndContext) + (t1, f1), (t2, f2) = self._pair(ctx.lhs, substitute), self._pair(ctx.rhs, substitute) + 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(), substitute) + if isinstance(ctx, PartiQLParser.PredicateLikeContext): + return self._like(ctx) + if isinstance(ctx, _LEAVES): + return self._leaf(ctx) + children = [c for c in getattr(ctx, "children", None) or [] if hasattr(c, "getRuleIndex")] + if len(children) == 1 and len(ctx.children) == 1: + return self._pair(children[0], substitute) + if isinstance(ctx, _PATH_NODES): + # A bare boolean field: WHERE flag / WHERE NOT flag + field = str(operand(ctx, self._resolver)) + return {field: True}, {field: False} + if isinstance(ctx, (PartiQLParser.LiteralTrueContext, PartiQLParser.LiteralFalseContext)): + everything: Filter = {} + return (everything, NOTHING) if isinstance(ctx, PartiQLParser.LiteralTrueContext) else (NOTHING, everything) + raise NotSupportedError(f"Unsupported WHERE expression: {ctx.getText()}") + + def _field_and_value(self, lhs: Any, rhs: Any, text: str) -> Tuple[str, Any, bool]: + """(field, value, mirrored) for ``lhs op rhs`` with the field on either side.""" + left, right = operand(lhs, self._resolver), operand(rhs, self._resolver) + if isinstance(left, _Field) and not isinstance(right, _Field): + return str(left), right, False + if isinstance(right, _Field) and not isinstance(left, _Field): + return str(right), left, True + raise NotSupportedError(f"A predicate needs exactly one field and one value: {text}") + + def _leaf(self, ctx: Any) -> Pair: + text = ctx.getText() + if isinstance(ctx, PartiQLParser.PredicateComparisonContext): + field, value, mirrored = self._field_and_value(ctx.lhs, ctx.rhs, text) + op = ctx.op.text + return leaf_filters(field, _MIRROR[op] if mirrored else op, value) + negated = ctx.NOT() is not None + field = operand(ctx.lhs, self._resolver) + if not isinstance(field, _Field): + raise NotSupportedError(f"The left side must be a field: {text}") + if isinstance(ctx, PartiQLParser.PredicateIsContext): + kind = ctx.type_().getText().upper() + if kind not in ("NULL", "MISSING"): + raise NotSupportedError(f"Unsupported IS type test: {text}") + return leaf_filters(field, f"IS {'NOT ' if negated else ''}{kind}", None) + if isinstance(ctx, PartiQLParser.PredicateInContext): + target = _unwrap(ctx.rhs) if ctx.rhs is not None else None + if isinstance(target, PartiQLParser.ValueListContext): + values = [_value(item, self._resolver) for item in target.expr()] + elif ctx.expr() is not None and ctx.rhs is None: + values = [_value(ctx.expr(), self._resolver)] # IN (single value) + else: + raise NotSupportedError(f"Unsupported IN list: {text}") + return leaf_filters(field, "NOT IN" if negated else "IN", values) + if isinstance(ctx, PartiQLParser.PredicateBetweenContext): + bounds = (_value(ctx.lower, self._resolver), _value(ctx.upper, self._resolver)) + return leaf_filters(field, "NOT BETWEEN" if negated else "BETWEEN", bounds) + raise NotSupportedError(f"Unsupported predicate: {text}") + + @staticmethod + def _case_folded(ctx: Any) -> Any: + """The argument of lower(x)/upper(x), as SQLAlchemy renders ILIKE; else None.""" + node = _unwrap(ctx) + if isinstance(node, PartiQLParser.FunctionCallContext) and len(node.expr()) == 1: + if node.functionName().getText().lower() in ("lower", "upper"): + return node.expr()[0] + return None + + def _like(self, ctx: Any) -> Pair: + lhs, rhs, case_insensitive = ctx.lhs, ctx.rhs, False + folded_lhs, folded_rhs = self._case_folded(ctx.lhs), self._case_folded(ctx.rhs) + if folded_lhs is not None: + # lower(field) LIKE lower(pattern) or lower(field) LIKE 'pattern': case-insensitive + lhs, rhs, case_insensitive = folded_lhs, (folded_rhs if folded_rhs is not None else ctx.rhs), True + field = operand(lhs, self._resolver) + if not isinstance(field, _Field): + raise NotSupportedError(f"The left side of LIKE must be a field: {ctx.getText()}") + pattern = _value(rhs, self._resolver) + op = ("NOT " if ctx.NOT() is not None else "") + ("ILIKE" if case_insensitive else "LIKE") + if ctx.escape is None: + return leaf_filters(field, op, pattern) + # The grammar lets ESCAPE take a whole expression, so "a LIKE p ESCAPE '/' AND b = 1" + # parses the AND chain as the escape. Its leftmost operand is the escape character; + # the LIKE predicate takes that operand's place in the chain. + leftmost = _unwrap(ctx.escape) + while isinstance(leftmost, (PartiQLParser.AndContext, PartiQLParser.OrContext)): + leftmost = _unwrap(leftmost.lhs) + escape = _value(leftmost, self._resolver) + if not isinstance(escape, str) or len(escape) != 1: + raise NotSupportedError(f"LIKE ESCAPE must be a single character: {ctx.getText()}") + like = leaf_filters(field, op, pattern, escape) + if leftmost is _unwrap(ctx.escape): + return like + return self._pair(ctx.escape, substitute=(leftmost, like)) diff --git a/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py b/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py index 540a0f1..a9cef5e 100644 --- a/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py +++ b/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py @@ -1,11 +1,13 @@ # -*- coding: utf-8 -*- import logging +import re +import uuid from typing import Any, Dict, List, Optional, Tuple, Type from urllib.parse import quote_plus 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 @@ -32,6 +34,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 +54,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: @@ -84,14 +102,48 @@ 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) + + def limit_clause(self, select, **kw): + """Render LIMIT/OFFSET as integer literals: PyMongoSQL reads them while parsing.""" + text = "" + kw = {**kw, "literal_binds": True} + if select._limit_clause is not None: + text += "\n LIMIT " + self.process(select._limit_clause, **kw) + if select._offset_clause is not None: + text += "\n OFFSET " + self.process(select._offset_clause, **kw) + return text + + # A quoted SQL string literal or quoted identifier ('' and "" escape the quote) + _QUOTED = re.compile(r"'(?:[^']|'')*'|\"(?:[^\"]|\"\")*\"") + + def _process_positional(self): + """Convert bind markers to qmark without touching quoted literals. + + SQLAlchemy 2.0 renders every bind as ``%(name)s`` and then rewrites the + whole statement with a regular expression. That also rewrites the text of a + string literal rendered inline (``literal_binds``), so ``x = '%(k)s'`` became + ``x = '?'``. Quoted segments are masked while the markers are converted. + """ + masked: List[str] = [] + + def mask(match: "re.Match[str]") -> str: + masked.append(match.group(0)) + return "\x00%d\x00" % (len(masked) - 1) + + self.string = self._QUOTED.sub(mask, self.string) + try: + super()._process_positional() + finally: + self.string = re.sub("\x00(\\d+)\x00", lambda m: masked[int(m.group(1))], self.string) class PyMongoSQLDDLCompiler(compiler.DDLCompiler): @@ -151,6 +203,80 @@ 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 _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.""" + + 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. @@ -174,6 +300,11 @@ 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. + supports_native_uuid = True # BSON binary subtype 4 + colspecs = {sqltypes.Numeric: _MongoNumeric, sqltypes.Float: _MongoFloat, sqltypes.Integer: _MongoInteger} + 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 @@ -384,11 +515,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 +562,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 +570,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 +607,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/pymongosql/superset_mongodb/exact_decimal.py b/pymongosql/superset_mongodb/exact_decimal.py new file mode 100644 index 0000000..3b13a4b --- /dev/null +++ b/pymongosql/superset_mongodb/exact_decimal.py @@ -0,0 +1,277 @@ +# -*- coding: utf-8 -*- +"""Exact decimal evaluation for the superset-mode SQLite stage. + +SQLite has no decimal type: a Decimal128 value stored there becomes a double (15-17 +significant digits) or text (which sorts and compares as a string). A column whose +values include decimals is therefore stored twice: as REAL in the query table (so +SQLite can still filter and sort it approximately) and as exact decimal text in a +side table keyed by rowid. Before a query runs, it is rewritten so that projections +of such a column, SUM/AVG/MIN/MAX over it, GROUP BY, ORDER BY and comparisons with a +numeric literal are evaluated from the exact text with Decimal arithmetic in the +precision of MongoDB's Decimal128 (34 significant digits). A query that uses such a +column in any other way (arithmetic, functions, DISTINCT aggregates) raises instead +of silently computing with doubles. +""" + +from decimal import ROUND_HALF_EVEN, Context, Decimal +from typing import Any, Dict, List, Optional, Sequence, Tuple + +from ..error import NotSupportedError + +SHADOW_PREFIX = "__pymongosql_exact__" +# Exact enough for any sum of Decimal128 values; results are then rounded like MongoDB's +# own Decimal128 arithmetic. +_SUM_CONTEXT = Context(prec=13000, rounding=ROUND_HALF_EVEN, Emin=-99999, Emax=99999) +_DECIMAL128 = Context(prec=34, rounding=ROUND_HALF_EVEN, Emin=-6143, Emax=6144) + + +def _decimal(value: Optional[str]) -> Optional[Decimal]: + return None if value is None else Decimal(value) + + +def _round(value: Decimal) -> str: + return str(_DECIMAL128.plus(value)) + + +class _Collect: + def __init__(self) -> None: + self.values: List[Decimal] = [] + + def step(self, value: Optional[str]) -> None: + if value is not None: + self.values.append(Decimal(value)) + + def total(self) -> Decimal: + total = Decimal(0) + for value in self.values: + total = _SUM_CONTEXT.add(total, value) + return total + + +class _Sum(_Collect): + def finalize(self) -> Optional[str]: + return _round(self.total()) if self.values else None + + +class _Avg(_Collect): + def finalize(self) -> Optional[str]: + if not self.values: + return None + return str(_DECIMAL128.divide(self.total(), Decimal(len(self.values)))) + + +class _Min(_Collect): + def finalize(self) -> Optional[str]: + return str(min(self.values)) if self.values else None + + +class _Max(_Collect): + def finalize(self) -> Optional[str]: + return str(max(self.values)) if self.values else None + + +def sort_key(value: Optional[str]) -> Optional[str]: + """Text whose byte order is the numeric order of the decimal ``value``.""" + number = _decimal(value) + if number is None: + return None + if number == 0: + return "1" + digits = "".join(map(str, number.normalize().as_tuple().digits)) + exponent = number.adjusted() + if number > 0: + return "2%06d%s" % (exponent + 100000, digits) + # Negative: larger magnitude first; "~" makes a prefix (-1) sort after -1.2 + complement = "".join(str(9 - int(d)) for d in digits) + return "0%06d%s~" % (100000 - exponent, complement) + + +def compare(value: Optional[str], other: Optional[str]) -> Optional[int]: + left, right = _decimal(value), _decimal(other) + if left is None or right is None: + return None + return (left > right) - (left < right) + + +AGGREGATES = {"SUM": ("__pymongosql_exact_sum", _Sum), "AVG": ("__pymongosql_exact_avg", _Avg)} +AGGREGATES.update({"MIN": ("__pymongosql_exact_min", _Min), "MAX": ("__pymongosql_exact_max", _Max)}) +FUNCTIONS = {"__pymongosql_exact_key": sort_key, "__pymongosql_exact_cmp": compare} + + +def register(connection: Any) -> None: + for name, aggregate in AGGREGATES.values(): + connection.create_aggregate(name, 1, aggregate) + for name, function in FUNCTIONS.items(): + connection.create_function(name, 2 if name.endswith("cmp") else 1, function, deterministic=True) + + +def rewrite(sql: str, table: str, exact_columns: Sequence[str], columns: Sequence[str] = ()) -> Tuple[str, List[int]]: + """Rewrite ``sql`` to evaluate the exact columns of ``table`` exactly. + + Returns the SQL and the positions of the outer projections that return exact + decimal text. ``columns`` (all columns of ``table``, in order) expands ``SELECT *``. + Raises NotSupportedError when an exact column is used in a way that + cannot be evaluated exactly. + """ + try: + import sqlglot + from sqlglot import exp + except ImportError as e: # pragma: no cover - sqlglot ships with Superset + raise NotSupportedError("Exact decimal evaluation needs the sqlglot package") from e + + exact = set(exact_columns) + try: + tree = sqlglot.parse_one(sql, read="sqlite") + except sqlglot.errors.ParseError as e: + raise NotSupportedError(f"Cannot evaluate decimal columns exactly in: {sql}") from e + if not isinstance(tree, exp.Select): + raise NotSupportedError(f"Cannot evaluate decimal columns exactly in: {sql}") + + comparisons = (exp.EQ, exp.NEQ, exp.GT, exp.GTE, exp.LT, exp.LTE) + mirrored = {exp.GT: exp.LT, exp.GTE: exp.LTE, exp.LT: exp.GT, exp.LTE: exp.GTE} + aggregates = {exp.Sum: "SUM", exp.Avg: "AVG", exp.Min: "MIN", exp.Max: "MAX"} + + def reads_table(select: Any) -> bool: + from_ = select.args.get("from") or select.args.get("from_") + return ( + from_ is not None + and isinstance(from_.this, exp.Table) + and not from_.this.db + and from_.this.name == table + and not select.args.get("joins") + ) + + def numeric_literal(node: Any) -> Optional[str]: + if isinstance(node, exp.Literal) and not node.is_string: + return str(Decimal(node.this)) + if isinstance(node, exp.Neg): + inner = numeric_literal(node.this) + return None if inner is None else str(-Decimal(inner)) + return None + + def shadow(name: str) -> Any: + return exp.column(SHADOW_PREFIX + name, quoted=True) + + def call(name: str, *args: Any) -> Any: + return exp.Anonymous(this=name, expressions=list(args)) + + def expand_star(select: Any) -> None: + expanded = [] + for projection in select.expressions: + star = isinstance(projection, exp.Star) or ( + isinstance(projection, exp.Column) and isinstance(projection.this, exp.Star) + ) + if star and columns: + expanded.extend(exp.column(name, quoted=True) for name in columns) + else: + expanded.append(projection) + select.set("expressions", expanded) + + def process(select: Any, outer: bool) -> List[int]: + expand_star(select) + + def column(node: Any) -> Optional[str]: + return node.name if isinstance(node, exp.Column) and node.name in exact else None + + def value(node: Any) -> Optional[Any]: + name = column(node) + if name is not None: + return shadow(name) + if type(node) in aggregates and not isinstance(node.this, exp.Distinct): + name = column(node.this) + if name is not None: + return call(AGGREGATES[aggregates[type(node)]][0], shadow(name)) + return None + + by_name: Dict[str, Any] = {} + by_position: Dict[int, Any] = {} + for index, projection in enumerate(select.expressions): + inner = projection.this if isinstance(projection, exp.Alias) else projection + by_position[index] = inner + by_name.setdefault(projection.alias_or_name, inner) + + def referenced(node: Any) -> Any: + if isinstance(node, exp.Literal) and node.is_int: + return by_position.get(int(node.this) - 1, node) + if isinstance(node, exp.Column) and not node.table: + # SQLite resolves a bare name to an output alias before a column + return by_name.get(node.name, node) + return node + + for key in ("where", "having"): + clause = select.args.get(key) + if clause is None: + continue + for comparison in list(clause.find_all(*comparisons)): + if comparison.find_ancestor(exp.Select) is not select: + continue + left, right, kind = comparison.this, comparison.expression, type(comparison) + if numeric_literal(left) is not None: + left, right, kind = right, left, mirrored.get(kind, kind) + literal = numeric_literal(right) + exact_left = value(referenced(left)) if literal is not None else None + if exact_left is not None: + comparison.replace( + kind( + this=call("__pymongosql_exact_cmp", exact_left, exp.Literal.string(literal)), + expression=exp.Literal.number(0), + ) + ) + group = select.args.get("group") + if group is not None: + group.set( + "expressions", + [shadow(column(referenced(k))) if column(referenced(k)) else k for k in group.expressions], + ) + order = select.args.get("order") + if order is not None: + for ordered in order.expressions: + exact_order = value(referenced(ordered.this)) + if exact_order is not None: + ordered.set("this", call("__pymongosql_exact_key", exact_order)) + positions: List[int] = [] + projections = [] + for index, projection in enumerate(select.expressions): + inner = projection.this if isinstance(projection, exp.Alias) else projection + exact_projection = value(inner) + if exact_projection is None: + projections.append(projection) + continue + if outer: + positions.append(index) + name = projection.alias_or_name + projections.append(exp.alias_(exact_projection, name, quoted=True)) + select.set("expressions", projections) + return positions + + selects = [s for s in tree.find_all(exp.Select) if reads_table(s)] + positions: List[int] = [] + for select in selects: + found = process(select, outer=select is tree) + if select is tree: + positions = found + stars = [s for s in tree.find_all(exp.Star) if not isinstance(s.parent, exp.Count)] + if stars: + # A * over a derived table would return the exact text untyped + raise NotSupportedError("SELECT * over a derived table with decimal columns: name the columns") + # Any remaining reference to an exact column outside COUNT or IS NULL would be + # computed with doubles + for node in tree.find_all(exp.Column): + if node.name in exact: + allowed = node.find_ancestor(exp.Count, exp.Is) + if allowed is None: + raise NotSupportedError( + f"Decimal column {node.name!r} is used in an expression that cannot be evaluated exactly" + ) + columns = ", ".join(f'x."{c}" AS "{SHADOW_PREFIX}{c}"' for c in exact_columns) + source = f'(SELECT q.*, {columns} FROM "{table}" AS q LEFT JOIN "{table}_exact" AS x ON x.rid = q.rowid)' + for node in list(tree.find_all(exp.Table)): + if node.name == table and not node.db: + alias = node.alias or table + node.replace( + exp.Subquery( + this=sqlglot.parse_one(source[1:-1], read="sqlite"), + alias=exp.TableAlias(this=exp.to_identifier(alias, quoted=True)), + ) + ) + return tree.sql(dialect="sqlite"), positions diff --git a/pymongosql/superset_mongodb/executor.py b/pymongosql/superset_mongodb/executor.py index f7090bc..47064b7 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( @@ -101,7 +104,11 @@ def execute( querydb_query = context.query table_name = "virtual_table" - query_db.insert_records(table_name, mongo_dicts) + if mongo_dicts: + query_db.insert_records(table_name, mongo_dicts) + elif column_names: + # No rows: the outer query still reads the table (COUNT(*) is 0, not an error) + query_db.create_table(table_name, {name: "" for name in column_names}) # Execute outer query against intermediate DB _logger.debug(f"Stage 2: Executing QueryDBSQLite query: {querydb_query}") diff --git a/pymongosql/superset_mongodb/query_db_sqlite.py b/pymongosql/superset_mongodb/query_db_sqlite.py index 5e3d356..d08aaa5 100644 --- a/pymongosql/superset_mongodb/query_db_sqlite.py +++ b/pymongosql/superset_mongodb/query_db_sqlite.py @@ -1,12 +1,28 @@ # -*- coding: utf-8 -*- +import datetime import logging import sqlite3 +from decimal import Decimal from typing import Any, Dict, List, Optional +from ..error import NotSupportedError +from . import exact_decimal from .query_db import QueryDatabase +from .time_grain import date_trunc_text, datetime_positions, datetime_text, parse_text, str_to_datetime_text _logger = logging.getLogger(__name__) +# 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) +# Datetimes are stored as fixed-width UTC text (so text order is time order) and read +# back as datetime when a query selects the column itself. +DATETIME_DECLTYPE = "PYMONGOSQL_DATETIME" +sqlite3.register_converter(DATETIME_DECLTYPE, lambda raw: datetime.datetime.fromisoformat(raw.decode())) +_NUMERIC = ("INTEGER", "REAL") + class SQLiteTypeMapper: """Maps Python/MongoDB data types to SQLite3 types""" @@ -16,7 +32,9 @@ 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 + Decimal: "REAL", # approximate copy; the exact text is kept in a side table + datetime.datetime: DATETIME_DECLTYPE, bytes: "BLOB", type(None): "NULL", dict: "TEXT", # Store as JSON string @@ -51,15 +69,16 @@ 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 in _NUMERIC and current in _NUMERIC: + schema[col_name] = "REAL" if "REAL" in (new_type, current) else "INTEGER" + elif new_type not in ("NULL", current): # Upgrade to TEXT if types differ (safest option) - if new_type != schema[col_name]: - schema[col_name] = "TEXT" + schema[col_name] = "TEXT" return schema @@ -69,7 +88,9 @@ def convert_value(cls, value: Any, target_type: str) -> Any: if value is None: return None - if target_type == "INTEGER": + if target_type == DATETIME_DECLTYPE: + return datetime_text(value) + if target_type in ("INTEGER", BOOLEAN_DECLTYPE): return int(value) if value is not None else None elif target_type == "REAL": return float(value) if value is not None else None @@ -98,6 +119,7 @@ def __init__(self) -> None: """Initialize SQLite3 bridge with in-memory database""" self._connection: Optional[sqlite3.Connection] = None self._tables: Dict[str, Dict[str, str]] = {} # table_name -> schema + self._exact_columns: Dict[str, List[str]] = {} # table_name -> columns with decimals self._is_closed = False def _ensure_connection(self) -> sqlite3.Connection: @@ -107,7 +129,11 @@ 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) + # Time-grain functions the engine spec's expressions use + self._connection.create_function("date_trunc", 2, date_trunc_text, deterministic=True) + self._connection.create_function("str_to_datetime", -1, str_to_datetime_text, deterministic=True) + exact_decimal.register(self._connection) # Enable row factory to get dict-like rows self._connection.row_factory = sqlite3.Row _logger.debug("Created in-memory SQLite3 database") @@ -165,7 +191,8 @@ def insert_records( # Build INSERT statement columns = list(records[0].keys()) placeholders = ", ".join(["?" for _ in columns]) - insert_sql = f"INSERT INTO {table_name} ({', '.join(columns)}) VALUES ({placeholders})" + quoted = ", ".join('"%s"' % col.replace('"', '""') for col in columns) + insert_sql = f'INSERT INTO "{table_name}" ({quoted}) VALUES ({placeholders})' # Convert values to appropriate types schema = self._tables[table_name] @@ -178,7 +205,9 @@ def insert_records( converted_records.append(converted_row) try: + first_rowid = conn.execute(f'SELECT COALESCE(MAX(rowid), 0) + 1 FROM "{table_name}"').fetchone()[0] conn.executemany(insert_sql, converted_records) + self._write_exact_copies(table_name, columns, records, first_rowid) conn.commit() _logger.debug(f"Inserted {len(records)} records into {table_name}") return len(records) @@ -186,6 +215,33 @@ def insert_records( _logger.error(f"Error inserting records into {table_name}: {e}") raise + def _write_exact_copies( + self, table_name: str, columns: List[str], records: List[Dict[str, Any]], first_rowid: int + ) -> None: + """Keep the exact text of every column holding decimals (see exact_decimal).""" + exact = [] + for col in columns: + values = [r.get(col) for r in records if r.get(col) is not None] + if any(isinstance(v, Decimal) for v in values) and all( + isinstance(v, (int, float, Decimal)) and not isinstance(v, bool) for v in values + ): + exact.append(col) + if not exact: + return + if table_name in self._exact_columns: + raise NotSupportedError("Decimal columns can be loaded into a query table only once") + self._exact_columns[table_name] = exact + conn = self._ensure_connection() + side = f'"{table_name}_exact"' + conn.execute(f"CREATE TABLE {side} (rid INTEGER PRIMARY KEY, %s)" % ", ".join('"%s" TEXT' % c for c in exact)) + conn.executemany( + f"INSERT INTO {side} VALUES ({', '.join('?' * (len(exact) + 1))})", + [ + (rid,) + tuple(None if r.get(c) is None else str(r.get(c)) for c in exact) + for rid, r in enumerate(records, start=first_rowid) + ], + ) + def execute_query(self, query: str) -> List[Dict[str, Any]]: """ Execute a query against the SQLite3 database. @@ -198,13 +254,29 @@ def execute_query(self, query: str) -> List[Dict[str, Any]]: """ conn = self._ensure_connection() + positions: List[int] = [] + dates = datetime_positions(query) + for table, exact in self._exact_columns.items(): + if table in query: + query, positions = exact_decimal.rewrite(query, table, exact, list(self._tables[table])) try: cursor = conn.execute(query) # Fetch all rows and convert from sqlite3.Row to dict rows = cursor.fetchall() column_names = [desc[0] for desc in cursor.description] if cursor.description else [] - return [dict(zip(column_names, row)) for row in rows] + + def convert(index: int, value: Any) -> Any: + if value is None: + return None + if index in positions: + return Decimal(value) + return parse_text(value) if index in dates else value + + return [ + {name: convert(index, value) for index, (name, value) in enumerate(zip(column_names, row))} + for row in rows + ] except sqlite3.Error as e: _logger.error(f"Error executing query: {e}") raise diff --git a/pymongosql/superset_mongodb/time_grain.py b/pymongosql/superset_mongodb/time_grain.py new file mode 100644 index 0000000..ec05134 --- /dev/null +++ b/pymongosql/superset_mongodb/time_grain.py @@ -0,0 +1,140 @@ +# -*- coding: utf-8 -*- +"""DATE_TRUNC(unit, value): the time-grain function Superset's engine spec emits. + +The same function is evaluated in two places with the same meaning: + +- on a collection, translated to MongoDB's ``$dateTrunc`` (see ``mongo_expression``); +- in the superset-mode SQLite stage, registered as a SQL function over the stage's + fixed-width UTC datetime text (see ``date_trunc_text``). + +Units: second, minute, hour, day, week (starting Sunday), week_monday, month, +quarter, year, week_ending_saturday and week_ending_sunday (the last day of a week +that starts on Sunday or Monday respectively). +""" + +import datetime +from typing import Any, Dict, Optional + +UNITS = ( + "second", + "minute", + "hour", + "day", + "week", + "week_monday", + "month", + "quarter", + "year", + "week_ending_saturday", + "week_ending_sunday", +) +_TEXT = "%Y-%m-%d %H:%M:%S.%f" + + +def _check(unit: str) -> str: + unit = unit.lower() + if unit not in UNITS: + raise ValueError(f"Unsupported DATE_TRUNC unit: {unit!r}") + return unit + + +def mongo_expression(unit: str, field: str) -> Dict[str, Any]: + """The aggregation expression for DATE_TRUNC(unit, field).""" + unit = _check(unit) + date = f"${field}" + if unit in ("week", "week_ending_saturday"): + truncated = {"$dateTrunc": {"date": date, "unit": "week", "startOfWeek": "sunday"}} + elif unit in ("week_monday", "week_ending_sunday"): + truncated = {"$dateTrunc": {"date": date, "unit": "week", "startOfWeek": "monday"}} + else: + truncated = {"$dateTrunc": {"date": date, "unit": unit}} + if unit.startswith("week_ending"): + return {"$dateAdd": {"startDate": truncated, "unit": "day", "amount": 6}} + return truncated + + +def truncate(unit: str, value: datetime.datetime) -> datetime.datetime: + unit = _check(unit) + if unit == "second": + return value.replace(microsecond=0) + if unit == "minute": + return value.replace(second=0, microsecond=0) + if unit == "hour": + return value.replace(minute=0, second=0, microsecond=0) + day = value.replace(hour=0, minute=0, second=0, microsecond=0) + if unit == "day": + return day + if unit in ("week", "week_ending_saturday"): + start = day - datetime.timedelta(days=(day.weekday() + 1) % 7) # back to Sunday + elif unit in ("week_monday", "week_ending_sunday"): + start = day - datetime.timedelta(days=day.weekday()) # back to Monday + elif unit == "month": + return day.replace(day=1) + elif unit == "quarter": + return day.replace(month=(day.month - 1) // 3 * 3 + 1, day=1) + else: # year + return day.replace(month=1, day=1) + return start + datetime.timedelta(days=6) if unit.startswith("week_ending") else start + + +def datetime_text(value: Any) -> Optional[str]: + """Fixed-width UTC text for a datetime (or ISO string) value.""" + if value is None: + return None + if isinstance(value, str): + value = datetime.datetime.fromisoformat(value.strip().replace("Z", "+00:00")) + if isinstance(value, datetime.datetime): + if value.tzinfo is not None: + value = value.astimezone(datetime.timezone.utc).replace(tzinfo=None) + return value.strftime(_TEXT) + if isinstance(value, datetime.date): + return datetime.datetime(value.year, value.month, value.day).strftime(_TEXT) + raise ValueError(f"Not a datetime: {value!r}") + + +def date_trunc_text(unit: str, value: Any) -> Optional[str]: + """SQLite function DATE_TRUNC(unit, datetime text).""" + text = datetime_text(value) + if text is None: + return None + return truncate(unit, datetime.datetime.strptime(text, _TEXT)).strftime(_TEXT) + + +def str_to_datetime_text(*args: Any) -> Optional[str]: + """SQLite function STR_TO_DATETIME(text[, format]), as the value function of the same name.""" + from ..sql.value_function_registry import ValueFunctionRegistry + + if args and args[0] is None: + return None + return datetime_text(ValueFunctionRegistry.str_to_datetime(*args)) + + +def datetime_positions(sql: str) -> list: + """Positions of the outer projections that are DATE_TRUNC/STR_TO_DATETIME calls. + + Their SQLite result is datetime text (an expression has no declared type); the + caller converts it back to datetime. + """ + lowered = sql.lower() + if "date_trunc" not in lowered and "str_to_datetime" not in lowered: + return [] + try: + import sqlglot + from sqlglot import exp + + tree = sqlglot.parse_one(sql, read="sqlite") + except Exception: + return [] + if not isinstance(tree, exp.Select): + return [] + positions = [] + for index, projection in enumerate(tree.expressions): + inner = projection.this if isinstance(projection, exp.Alias) else projection + name = inner.name.lower() if isinstance(inner, exp.Anonymous) else "" + if isinstance(inner, exp.DateTrunc) or name in ("date_trunc", "str_to_datetime"): + positions.append(index) + return positions + + +def parse_text(value: Any) -> Any: + return datetime.datetime.strptime(value, _TEXT) if isinstance(value, str) else value diff --git a/pyproject.toml b/pyproject.toml index 8e5f8ce..914f6ad 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -41,6 +41,7 @@ dependencies = [ [project.optional-dependencies] retry = ["tenacity>=9.0.0"] sqlalchemy = ["sqlalchemy>=1.4.0"] +superset = ["sqlglot>=25.0.0"] pandas_sqlalchemy14 = [ "sqlalchemy>=1.4.0,<2.0.0", "pandas>=2.1.0,<2.2.0; python_version < '3.13'", diff --git a/requirements-optional.txt b/requirements-optional.txt index 3ebd3bb..787d1a2 100644 --- a/requirements-optional.txt +++ b/requirements-optional.txt @@ -3,3 +3,6 @@ tenacity>=9.0.0 # SQLAlchemy support (optional) - supports 1.4+ and 2.x sqlalchemy>=1.4.0,<3.0.0 + +# Superset mode: exact decimal evaluation in the SQLite stage (optional) +sqlglot>=25.0.0 diff --git a/tests/test_cursor_delete.py b/tests/test_cursor_delete.py index 045ea76..d2d1aef 100644 --- a/tests/test_cursor_delete.py +++ b/tests/test_cursor_delete.py @@ -89,7 +89,7 @@ def test_delete_with_and_condition(self, conn): def test_delete_with_qmark_parameters(self, conn): """Test DELETE with qmark (?) placeholder parameters.""" cursor = conn.cursor() - result = cursor.execute(f"DELETE FROM {self.TEST_COLLECTION} WHERE artist = '?'", ["Charlie"]) + result = cursor.execute(f"DELETE FROM {self.TEST_COLLECTION} WHERE artist = ?", ["Charlie"]) assert result == cursor @@ -104,7 +104,7 @@ def test_delete_with_qmark_parameters(self, conn): def test_delete_with_multiple_parameters(self, conn): """Test DELETE with multiple qmark parameters.""" cursor = conn.cursor() - result = cursor.execute(f"DELETE FROM {self.TEST_COLLECTION} WHERE genre = '?' AND year = '?'", ["Pop", 2019]) + result = cursor.execute(f"DELETE FROM {self.TEST_COLLECTION} WHERE genre = ? AND year = ?", ["Pop", 2019]) assert result == cursor @@ -184,7 +184,7 @@ def test_delete_followed_by_insert(self, conn): def test_delete_executemany_with_parameters(self, conn): """Test executemany for bulk delete operations with parameters.""" cursor = conn.cursor() - sql = f"DELETE FROM {self.TEST_COLLECTION} WHERE artist = '?'" + sql = f"DELETE FROM {self.TEST_COLLECTION} WHERE artist = ?" # Delete multiple artists using executemany params = [["Alice"], ["Charlie"], ["Eve"]] diff --git a/tests/test_cursor_update.py b/tests/test_cursor_update.py index bb94461..5d7a16b 100644 --- a/tests/test_cursor_update.py +++ b/tests/test_cursor_update.py @@ -214,7 +214,7 @@ def test_update_set_null(self, conn): def test_update_executemany_with_parameters(self, conn): """Test executemany for bulk update operations with parameters.""" cursor = conn.cursor() - sql = f"UPDATE {self.TEST_COLLECTION} SET price = '?' WHERE title = '?'" + sql = f"UPDATE {self.TEST_COLLECTION} SET price = ? WHERE title = ?" # Update prices for multiple books using executemany params = [[25.99, "Book A"], [35.99, "Book B"], [45.99, "Book D"]] diff --git a/tests/test_decimals_time_grains_literal_binds.py b/tests/test_decimals_time_grains_literal_binds.py new file mode 100644 index 0000000..f848929 --- /dev/null +++ b/tests/test_decimals_time_grains_literal_binds.py @@ -0,0 +1,356 @@ +# -*- coding: utf-8 -*- +"""Superset-mode decimals, DATE_TRUNC time grains and literal binds. + +- The superset-mode SQLite stage evaluates decimal columns exactly (SUM/AVG/MIN/MAX, + GROUP BY, ORDER BY, comparisons) instead of as doubles or text. +- DATE_TRUNC(unit, field) is translated to MongoDB's $dateTrunc on a collection and + evaluated with the same meaning in the SQLite stage. +- Compiling with literal_binds keeps a string literal containing ``%(name)s`` intact. +""" + +import datetime +from decimal import Decimal, localcontext + +import pytest +from bson import Decimal128 + +from pymongosql.error import Error, NotSupportedError +from pymongosql.sql.parser import SQLParser +from pymongosql.superset_mongodb.query_db_sqlite import QueryDBSQLite +from pymongosql.superset_mongodb.time_grain import UNITS, truncate +from tests.conftest import HAS_SQLALCHEMY, make_superset_conn + +if HAS_SQLALCHEMY: + import sqlalchemy as sa + + from pymongosql.sqlalchemy_mongodb.sqlalchemy_dialect import PyMongoSQLDialect + +needs_sqlalchemy = pytest.mark.skipif(not HAS_SQLALCHEMY, reason="SQLAlchemy not available") +sqlglot = pytest.importorskip("sqlglot") + +COLLECTION = "test_decimals_time_grains" +BIG = Decimal("123456789012345678901.1234567891") +AMOUNTS = [BIG, Decimal("0.0000000001"), Decimal("-5.5"), Decimal("10"), None, Decimal("9.99")] +GROUPS = ["a", "a", "b", "b", "b", "a"] +TIMES = [ + datetime.datetime(2025, 12, 31, 23, 59, 59, 999000), # Wednesday + datetime.datetime(2026, 1, 3, 10, 15, 30, 250000), # Saturday + datetime.datetime(2026, 1, 4, 0, 0, 0), # Sunday + datetime.datetime(2026, 1, 5, 8, 30, 0), # Monday + datetime.datetime(2026, 4, 1, 12, 0, 1), + datetime.datetime(2026, 9, 26, 18, 45, 12, 5000), +] + + +def records(): + return [{"g": g, "amt": a, "ts": t} for g, a, t in zip(GROUPS, AMOUNTS, TIMES)] + + +def stage(sql): + db = QueryDBSQLite() + try: + db.insert_records("virtual_table", records()) + return [tuple(r.values()) for r in db.execute_query(sql)] + finally: + db.close() + + +A_SUM = Decimal("123456789012345678911.1134567892") # BIG + 0.0000000001 + 9.99, exact + + +def exact(values): + return [v for v in values if v is not None] + + +class TestExactDecimalsInTheSQLiteStage: + def test_sum_avg_min_max_are_exact(self): + rows = stage('SELECT SUM(amt) AS "SUM(amt)", AVG(amt) AS a, MIN(amt) AS lo, MAX(amt) AS hi FROM virtual_table') + # Decimal128 arithmetic: 34 significant digits, as MongoDB's own $sum and $avg + with localcontext() as ctx: + ctx.prec = 34 + total = sum(exact(AMOUNTS)) + average = total / 5 + assert total == Decimal("123456789012345678915.6134567892") + assert rows == [(total, average, Decimal("-5.5"), BIG)] + assert all(type(v) is Decimal for v in rows[0]) + + def test_group_by_with_exact_sums(self): + rows = stage('SELECT g AS g, SUM(amt) AS "SUM(amt)" FROM virtual_table GROUP BY g ORDER BY g') + assert rows == [("a", A_SUM), ("b", Decimal("4.5"))] + + def test_group_by_the_decimal_column(self): + rows = stage("SELECT amt AS amt, COUNT(*) AS n FROM virtual_table GROUP BY amt ORDER BY amt") + assert [r[0] for r in rows] == [ + None, + Decimal("-5.5"), + Decimal("0.0000000001"), + Decimal("9.99"), + Decimal("10"), + BIG, + ] + + def test_order_by_is_numeric_not_text(self): + rows = stage("SELECT amt FROM virtual_table WHERE amt IS NOT NULL ORDER BY amt DESC") + assert [r[0] for r in rows] == [BIG, Decimal("10"), Decimal("9.99"), Decimal("0.0000000001"), Decimal("-5.5")] + + def test_top_n_by_aggregate(self): + rows = stage('SELECT g, SUM(amt) AS "SUM(amt)" FROM virtual_table GROUP BY g ORDER BY "SUM(amt)" ASC LIMIT 1') + assert rows == [("b", Decimal("4.5"))] + + def test_comparisons_with_numeric_literals_are_exact(self): + # As doubles, 123456789012345678901.1234567891 equals 123456789012345678901.1234567890 + assert stage("SELECT COUNT(*) FROM virtual_table WHERE amt > 123456789012345678901.123456789") == [(1,)] + assert stage("SELECT COUNT(*) FROM virtual_table WHERE amt = 123456789012345678901.123456789") == [(0,)] + assert stage("SELECT COUNT(*) FROM virtual_table WHERE 0.00000000005 < amt AND amt < 1") == [(1,)] + rows = stage('SELECT g, SUM(amt) AS "SUM(amt)" FROM virtual_table GROUP BY g HAVING SUM(amt) > 4.5') + assert [r[0] for r in rows] == ["a"] + + def test_select_star_returns_decimals(self): + rows = stage("SELECT * FROM virtual_table WHERE g = 'b' ORDER BY amt") + assert [r[1] for r in rows] == [None, Decimal("-5.5"), Decimal("10")] + + def test_count_and_null_checks_are_allowed(self): + assert stage("SELECT COUNT(amt), COUNT(*) FROM virtual_table WHERE amt IS NULL OR amt IS NOT NULL") == [(5, 6)] + + @pytest.mark.parametrize( + "sql", + [ + "SELECT amt * 2 FROM virtual_table", + "SELECT ROUND(amt, 2) FROM virtual_table", + "SELECT SUM(DISTINCT amt) FROM virtual_table", + "SELECT * FROM (SELECT amt FROM virtual_table) AS x", + ], + ) + def test_other_uses_are_refused_rather_than_approximated(self, sql): + with pytest.raises(NotSupportedError): + stage(sql) + + def test_integer_and_float_columns_are_unchanged(self): + db = QueryDBSQLite() + try: + db.insert_records("virtual_table", [{"i": 1, "f": 0.5}, {"i": 2, "f": 1.25}]) + assert db.execute_query("SELECT SUM(i) AS i, SUM(f) AS f FROM virtual_table") == [{"i": 3, "f": 1.75}] + finally: + db.close() + + +class TestDateTruncTranslation: + def plan(self, sql): + return SQLParser(sql).get_execution_plan() + + def test_grouped_time_grain(self): + plan = self.plan( + "SELECT DATE_TRUNC('month', ts) AS __timestamp, COUNT(*) AS \"count\" FROM t " + "GROUP BY DATE_TRUNC('month', ts) ORDER BY \"count\" DESC LIMIT 10" + ) + assert '"$dateTrunc": {"date": "$ts", "unit": "month"}' in plan.aggregate_pipeline + + def test_week_ending_adds_six_days(self): + plan = self.plan( + "SELECT DATE_TRUNC('week_ending_sunday', t.ts) AS x FROM t GROUP BY DATE_TRUNC('week_ending_sunday', t.ts)" + ) + assert '"startOfWeek": "monday"' in plan.aggregate_pipeline + assert '"$dateAdd"' in plan.aggregate_pipeline and '"$t.ts"' not in plan.aggregate_pipeline + + def test_ungrouped_projection(self): + plan = self.plan("SELECT DATE_TRUNC('year', ts) AS y, name FROM t ORDER BY y DESC") + assert '"$addFields"' in plan.aggregate_pipeline and '"$sort": {"__computed0": -1}' in plan.aggregate_pipeline + + @pytest.mark.parametrize( + "sql", ["SELECT DATE_TRUNC('fortnight', ts) FROM t", "SELECT DATE_TRUNC(ts, 'day') FROM t"] + ) + def test_invalid_date_trunc_is_refused(self, sql): + with pytest.raises(Error): + self.plan(sql) + + @pytest.mark.parametrize("unit", UNITS) + def test_sqlite_stage_date_trunc(self, unit): + rows = stage(f"SELECT DATE_TRUNC('{unit}', ts) AS t FROM virtual_table ORDER BY ts") + assert [r[0] for r in rows] == [truncate(unit, t) for t in TIMES] + + +@needs_sqlalchemy +class TestLiteralBinds: + def compile(self, stmt, **kw): + return " ".join(str(stmt.compile(dialect=PyMongoSQLDialect(), **kw)).split()) + + def test_percent_name_literal_is_kept(self): + t = sa.table("t", sa.column("x"), sa.column("n")) + stmt = sa.select(t.c.x).where(t.c.x == "%(k)s", t.c.n == 5) + literal = self.compile(stmt, compile_kwargs={"literal_binds": True}) + assert literal == "SELECT x FROM t WHERE x = '%(k)s' AND n = 5" + bound = stmt.compile(dialect=PyMongoSQLDialect()) + assert " ".join(str(bound).split()) == "SELECT x FROM t WHERE x = ? AND n = ?" + assert list(bound.positiontup) == ["x_1", "n_1"] + + def test_quotes_inside_literals(self): + t = sa.table("t", sa.column("x")) + stmt = sa.select(t.c.x).where(t.c.x == 'it\'s %(a)s "q"') + assert self.compile(stmt, compile_kwargs={"literal_binds": True}) == ( + "SELECT x FROM t WHERE x = 'it''s %(a)s \"q\"'" + ) + + def test_limit_with_literal_binds(self): + t = sa.table("t", sa.column("x")) + stmt = sa.select(t.c.x).where(t.c.x.contains("%(k)s", autoescape=True)).limit(3).offset(1) + assert self.compile(stmt, compile_kwargs={"literal_binds": True}) == ( + "SELECT x FROM t WHERE (x LIKE '%' || '/%(k)s' || '%' ESCAPE '/') LIMIT 3 OFFSET 1" + ) + + +@pytest.fixture +def grain_docs(conn): + conn.database.drop_collection(COLLECTION) + conn.database[COLLECTION].insert_many( + [ + {"_id": i, "g": g, "amt": None if a is None else Decimal128(a), "ts": t} + for i, (g, a, t) in enumerate(zip(GROUPS, AMOUNTS, TIMES)) + ] + ) + yield conn + conn.database.drop_collection(COLLECTION) + + +class TestLive: + @pytest.mark.parametrize("unit", UNITS) + def test_physical_time_grain(self, grain_docs, unit): + cursor = grain_docs.cursor() + cursor.execute( + f"SELECT DATE_TRUNC('{unit}', ts) AS __timestamp, COUNT(*) AS \"count\" FROM {COLLECTION} " + f"WHERE ts >= STR_TO_DATETIME('2025-01-01T00:00:00') GROUP BY DATE_TRUNC('{unit}', ts) " + "ORDER BY __timestamp" + ) + expected = {} + for t in TIMES: + expected[truncate(unit, t)] = expected.get(truncate(unit, t), 0) + 1 + assert [tuple(r) for r in cursor.fetchall()] == sorted(expected.items()) + + def test_physical_ungrouped_time_grain(self, grain_docs): + cursor = grain_docs.cursor() + cursor.execute(f"SELECT DATE_TRUNC('week', ts) AS w, g FROM {COLLECTION} ORDER BY w DESC LIMIT 2") + assert [tuple(r) for r in cursor.fetchall()] == [ + (datetime.datetime(2026, 9, 20), "a"), + (datetime.datetime(2026, 3, 29), "b"), + ] + + @pytest.mark.parametrize("unit", UNITS) + def test_virtual_time_grain(self, grain_docs, unit): + superset = make_superset_conn() + try: + cursor = superset.cursor() + cursor.execute( + f"SELECT DATE_TRUNC('{unit}', ts) AS __timestamp, COUNT(*) AS \"count\" " + f"FROM (SELECT ts FROM {COLLECTION}) AS virtual_table " + "WHERE ts >= STR_TO_DATETIME('2025-01-01T00:00:00') " + f"GROUP BY DATE_TRUNC('{unit}', ts) ORDER BY __timestamp" + ) + expected = {} + for t in TIMES: + expected[truncate(unit, t)] = expected.get(truncate(unit, t), 0) + 1 + assert [tuple(r) for r in cursor.fetchall()] == sorted(expected.items()) + finally: + superset.close() + + def test_virtual_dataset_decimals_are_exact(self, grain_docs): + superset = make_superset_conn() + try: + cursor = superset.cursor() + cursor.execute( + 'SELECT g AS g, SUM(amt) AS "SUM(amt)" ' + f'FROM (SELECT g, amt FROM {COLLECTION}) AS virtual_table GROUP BY g ORDER BY "SUM(amt)" DESC' + ) + assert [tuple(r) for r in cursor.fetchall()] == [ + ("a", A_SUM), + ("b", Decimal("4.5")), + ] + cursor.execute(f"SELECT amt FROM (SELECT amt FROM {COLLECTION}) AS virtual_table ORDER BY amt DESC LIMIT 3") + assert [r[0] for r in cursor.fetchall()] == [BIG, Decimal("10"), Decimal("9.99")] + finally: + superset.close() + + @needs_sqlalchemy + def test_literal_binds_statement_runs(self, sqlalchemy_engine, grain_docs): + grain_docs.database[COLLECTION].insert_one({"_id": 99, "g": "%(k)s"}) + t = sa.table(COLLECTION, sa.column("_id"), sa.column("g")) + stmt = sa.select(t.c._id).where(t.c.g == "%(k)s") + sql = str(stmt.compile(sqlalchemy_engine, compile_kwargs={"literal_binds": True})) + with sqlalchemy_engine.connect() as c: + assert list(c.exec_driver_sql(sql).scalars()) == [99] + assert list(c.execute(stmt).scalars()) == [99] + + +class TestLiveAggregateNullSemantics: + @pytest.fixture + def null_docs(self, conn): + name = COLLECTION + "_nulls" + conn.database.drop_collection(name) + conn.database[name].insert_many( + [ + {"_id": 1, "g": "a", "amt": None}, + {"_id": 2, "g": "b", "amt": Decimal128("1.5")}, + {"_id": 3, "g": "b"}, + {"_id": 4, "g": "c", "amt": 0}, + ] + ) + yield conn, name + conn.database.drop_collection(name) + + def rows(self, conn, sql): + cursor = conn.cursor() + cursor.execute(sql) + return [tuple(r) for r in cursor.fetchall()] + + def test_sum_of_no_values_is_null(self, null_docs): + conn, name = null_docs + rows = self.rows(conn, f"SELECT g, SUM(amt) AS s, SUM(DISTINCT amt) AS d FROM {name} GROUP BY g ORDER BY g") + assert rows == [("a", None, None), ("b", Decimal("1.5"), Decimal("1.5")), ("c", 0, 0)] + assert self.rows(conn, f"SELECT SUM(amt) AS s FROM {name} WHERE g = 'a'") == [(None,)] + + def test_aggregate_without_group_by_returns_one_row_for_no_input(self, null_docs): + conn, name = null_docs + empty = f"FROM {name} WHERE g = 'none'" + assert self.rows( + conn, f"SELECT SUM(amt) AS s, COUNT(*) AS n, COUNT(DISTINCT g) AS d, MAX(amt) AS m {empty}" + ) == [(None, 0, 0, None)] + assert self.rows(conn, f"SELECT COUNT(*) AS n {empty} HAVING COUNT(*) > 0") == [] + assert self.rows(conn, f"SELECT g, COUNT(*) AS n {empty} GROUP BY g") == [] + assert self.rows(conn, f"SELECT COUNT(*) AS n FROM {name} LIMIT 0") == [] + + +class TestLiveEmptyVirtualDataset: + def test_inner_query_without_rows(self, grain_docs): + superset = make_superset_conn() + inner = f"(SELECT g, amt FROM {COLLECTION} WHERE g = 'none') AS virtual_table" + try: + cursor = superset.cursor() + cursor.execute(f'SELECT COUNT(*) AS "count", SUM(amt) AS s FROM {inner}') + assert [tuple(r) for r in cursor.fetchall()] == [(0, None)] + cursor.execute(f'SELECT g, COUNT(*) AS "count" FROM {inner} GROUP BY g') + assert cursor.fetchall() == [] + cursor.execute(f"SELECT g FROM {inner}") + assert cursor.fetchall() == [] and [d[0] for d in cursor.description] == ["g"] + finally: + superset.close() + + +class TestLiveOrderByAggregateNotSelected: + def rows(self, conn, sql): + cursor = conn.cursor() + cursor.execute(sql) + return [tuple(r) for r in cursor.fetchall()] + + def test_top_n_ordered_by_another_metric(self, grain_docs): + # Superset's series-limit pre-query: group by the series, order by the limit metric + sql = f'SELECT g AS g, COUNT(*) AS "count" FROM {COLLECTION} GROUP BY g ORDER BY sum(amt) DESC LIMIT 1' + assert self.rows(grain_docs, sql) == [("a", 3)] + sql = f"SELECT g, COUNT(*) AS n FROM {COLLECTION} GROUP BY g ORDER BY SUM({COLLECTION}.amt) ASC" + assert self.rows(grain_docs, sql) == [("b", 3), ("a", 3)] + sql = f"SELECT g, COUNT(*) AS n FROM {COLLECTION} GROUP BY g HAVING SUM(amt) > 5 ORDER BY MAX(amt) DESC" + assert self.rows(grain_docs, sql) == [("a", 3)] + + def test_time_grain_with_having(self, grain_docs): + sql = ( + f"SELECT DATE_TRUNC('month', ts) AS m, COUNT(*) AS n FROM {COLLECTION} " + "GROUP BY DATE_TRUNC('month', ts) HAVING COUNT(*) > 1 ORDER BY m" + ) + assert self.rows(grain_docs, sql) == [(datetime.datetime(2026, 1, 1), 3)] diff --git a/tests/test_parse_tree_predicates.py b/tests/test_parse_tree_predicates.py new file mode 100644 index 0000000..4726dca --- /dev/null +++ b/tests/test_parse_tree_predicates.py @@ -0,0 +1,330 @@ +# -*- coding: utf-8 -*- +"""Predicates are read from the parse tree: field names, operators and values are never +recovered by searching concatenated token text. + +The DML tests check the exact documents a DELETE or UPDATE touches: a truncated +field name there removes or rewrites the wrong documents. +""" + +from decimal import Decimal + +import pytest +from bson import Decimal128, Int64 + +from pymongosql.error import Error +from pymongosql.sql.parser import SQLParser +from tests.conftest import HAS_SQLALCHEMY + +if HAS_SQLALCHEMY: + import sqlalchemy as sa + + from pymongosql.sqlalchemy_mongodb.sqlalchemy_dialect import PyMongoSQLDialect + +needs_sqlalchemy = pytest.mark.skipif(not HAS_SQLALCHEMY, reason="SQLAlchemy not available") +PARAM = {"$pymongosqlParam": True} + + +def plan(sql): + return SQLParser(sql).get_execution_plan() + + +def where(sql_where): + return plan(f"SELECT _id FROM t WHERE {sql_where}").filter_stage + + +def compiled(stmt): + return " ".join(str(stmt.compile(dialect=PyMongoSQLDialect())).split()) + + +class TestFieldNamesAreNotTruncated: + @pytest.mark.parametrize( + "sql_where,expected", + [ + ("dislikes IS NULL", {"dislikes": {"$eq": None}}), + ("unlike = 5", {"unlike": 5}), + ("likes IN (1, 2)", {"likes": {"$in": [1, 2]}}), + ("isnull = 1", {"isnull": 1}), + ("between_x BETWEEN 1 AND 2", {"$and": [{"between_x": {"$gte": 1}}, {"between_x": {"$lte": 2}}]}), + ("login_count > 3", {"login_count": {"$gt": 3}}), + ("inx <> 'a'", {"inx": {"$nin": ["a", None]}}), + ], + ) + def test_select(self, sql_where, expected): + assert where(sql_where) == expected + + def test_delete_and_update(self): + assert plan("DELETE FROM t WHERE dislikes IS NULL").filter_conditions == {"dislikes": {"$eq": None}} + update = plan("UPDATE t SET x = 1 WHERE unlike = 5") + assert update.filter_conditions == {"unlike": 5} + + def test_string_literal_containing_keywords(self): + assert where("name = 'x IN(1) LIKE y BETWEEN'") == {"name": "x IN(1) LIKE y BETWEEN"} + + def test_reversed_operands(self): + assert where("5 < age") == {"age": {"$gt": 5}} + assert where("'x' = name") == {"name": "x"} + assert where("10 >= age") == {"age": {"$lte": 10}} + + @pytest.mark.parametrize( + "sql_where,field", + [('"my field" = 1', "my field"), ('"a-b" = 1', "a-b"), ('"año" = 1', "año"), ('"a"."b c" = 1', "a.b c")], + ) + def test_quoted_field_names(self, sql_where, field): + assert where(sql_where) == {field: 1} + + @pytest.mark.parametrize("sql_where", ["a = b", "lower(a) = 'x'", "a + 1 = 2", "a IN (SELECT b FROM u)"]) + def test_untranslatable_predicates_raise(self, sql_where): + with pytest.raises(Error): + where(sql_where) + + +class TestLike: + def test_concatenated_pattern(self): + assert where("n LIKE '%' || 'ab' || '%'") == {"n": {"$regex": ".*ab.*"}} + assert where("n LIKE 'ab' || '%'") == {"n": {"$regex": "^ab.*"}} + + def test_escape(self): + assert where("n LIKE '50/%%' ESCAPE '/'") == {"n": {"$regex": "^50%.*"}} + assert where("n LIKE 'a/_b' ESCAPE '/'") == {"n": {"$regex": "^a_b$"}} + + def test_escape_followed_by_more_conditions(self): + assert where("n LIKE 'a/_%' ESCAPE '/' AND b = 1 OR c = 2") == { + "$or": [{"$and": [{"n": {"$regex": "^a_.*"}}, {"b": 1}]}, {"c": 2}] + } + + def test_inner_wildcards_match_newlines(self): + assert where("n LIKE 'a%b'") == {"n": {"$regex": "^a.*b$", "$options": "s"}} + + def test_bound_pattern_is_translated_when_bound(self): + from pymongosql.helper import SQLHelper + + f = where("n LIKE '%' || ? || '%' ESCAPE '/'") + bound, used = SQLHelper.bind_filter(f, ["5/%(k)s"]) + assert (bound, used) == ({"n": {"$regex": ".*5%\\(k\\)s.*"}}, 1) + with pytest.raises(Error): + SQLHelper.bind_filter(where("n LIKE ?"), [5]) + + def test_case_insensitive_like_from_lower(self): + assert where("lower(n) LIKE lower('A%b')") == {"n": {"$regex": "^A.*b$", "$options": "si"}} + assert where("lower(n) LIKE 'a%'") == {"n": {"$regex": "^a.*", "$options": "i"}} + + @needs_sqlalchemy + def test_sqlalchemy_helpers_translate_after_binding(self): + from pymongosql.helper import SQLHelper + + t = sa.table("t", sa.column("n")) + for expr, regex in ( + (t.c.n.contains("ab"), {"$regex": ".*ab.*"}), + (t.c.n.startswith("ab"), {"$regex": "^ab.*"}), + (t.c.n.endswith("ab"), {"$regex": ".*ab$"}), + (t.c.n.contains("5%_ %(k)s", autoescape=True), {"$regex": ".*5%_\\ %\\(k\\)s.*"}), + (t.c.n.like("a/_%", escape="/"), {"$regex": "^a_.*"}), + (t.c.n.ilike("A%"), {"$regex": "^A.*", "$options": "i"}), + ): + compiled_stmt = sa.select(t.c.n).where(expr).compile(dialect=PyMongoSQLDialect()) + params = [compiled_stmt.params[name] for name in compiled_stmt.positiontup] + bound, _ = SQLHelper.bind_filter(plan(str(compiled_stmt)).filter_stage, params) + assert bound == {"n": regex}, str(compiled_stmt) + + +class TestParameters: + def test_literal_question_mark_is_a_value(self): + assert where("a = '?' AND b = ?") == {"$and": [{"a": "?"}, {"b": PARAM}]} + + def test_literal_question_mark_in_aggregate(self): + p = plan("SELECT g, COUNT(*) AS n FROM t WHERE a = '?' AND b = ? GROUP BY g") + assert '"a": "?"' in p.aggregate_pipeline and '"$pymongosqlParam"' in p.aggregate_pipeline + + def test_limit_and_offset_parameters_are_kept(self): + p = plan("SELECT _id FROM t WHERE a = ? LIMIT ? OFFSET ?") + assert (p.filter_stage, p.limit_stage, p.skip_stage) == ({"a": PARAM}, PARAM, PARAM) + + @pytest.mark.parametrize("clause", ["LIMIT -1", "LIMIT 'x'", "OFFSET 1.5", "LIMIT a"]) + def test_invalid_limit_or_offset_raises_instead_of_returning_everything(self, clause): + with pytest.raises(Error): + plan(f"SELECT _id FROM t {clause}") + + @needs_sqlalchemy + def test_sqlalchemy_limit_and_offset_are_literals(self): + t = sa.table("t", sa.column("a")) + assert compiled(sa.select(t.c.a).limit(5).offset(10)) == "SELECT a FROM t LIMIT 5 OFFSET 10" + assert compiled(sa.select(t.c.a).offset(3)) == "SELECT a FROM t OFFSET 3" + + +class TestDecimalLiterals: + def test_exact_decimal_literal_is_a_double(self): + assert where("t > 36.5") == {"t": {"$gt": 36.5}} + + def test_inexact_decimal_literal_compares_in_the_field_type(self): + assert where("a = 0.1") == { + "$or": [ + {"$and": [{"a": {"$type": "double"}}, {"a": 0.1}]}, + {"$and": [{"a": {"$not": {"$type": "double"}}}, {"a": Decimal128("0.1")}]}, + ] + } + + +class TestAggregates: + def test_count_distinct(self): + p = plan('SELECT g, COUNT(DISTINCT a) AS "COUNT_DISTINCT(a)" FROM t GROUP BY g') + assert '"$addToSet": "$a"' in p.aggregate_pipeline + assert '"COUNT_DISTINCT(a)": {"$size": {"$setDifference": ["$__agg0", [null]]}}' in p.aggregate_pipeline + + @pytest.mark.parametrize("sql", ["SELECT a + 1 FROM t", "SELECT lower(a) FROM t", "SELECT EVERY(a) FROM t"]) + def test_untranslatable_select_expressions_raise(self, sql): + with pytest.raises(Error): + plan(sql) + + def test_having(self): + p = plan("SELECT g, COUNT(*) AS n FROM t GROUP BY g HAVING n > 1 AND SUM(v) >= 3") + assert '{"$match": {"$and": [{"n": {"$gt": 1}}, {"__having0": {"$gte": 3}}]}}' in p.aggregate_pipeline + assert '{"$project": {"__having0": 0}}' in p.aggregate_pipeline + + +@needs_sqlalchemy +class TestFloatColumns: + """Float returns float; Float(asdecimal=True) returns Decimal as SQLAlchemy documents.""" + + def processor(self, type_): + dialect = PyMongoSQLDialect() + return type_.dialect_impl(dialect).result_processor(dialect, None) + + def test_float_returns_float(self): + for type_ in (sa.Float(), sa.FLOAT(), sa.REAL()): + value = self.processor(type_)(Decimal128("1.25")) + assert value == 1.25 and type(value) is float + + def test_float_asdecimal_returns_decimal(self): + value = self.processor(sa.Float(asdecimal=True))(Decimal128("1.25")) + assert value == Decimal("1.25") and type(value) is Decimal + + +# ---------------------------------------------------------------- live MongoDB + +COLLECTION = "test_parse_tree_predicates" +DOCS = [ + {"_id": 1, "dislikes": None, "dis": 1, "unlike": 5, "un": "x", "n": "ab-50%", "g": "a", "v": 1, "f": 0.1}, + {"_id": 2, "dislikes": 3, "unlike": 6, "un": "=5", "n": "zab", "g": "a", "v": 2, "f": 0.2}, + {"_id": 3, "dislikes": 4, "dis": 2, "unlike": 5, "n": "a_b", "g": "b", "v": 2, "f": Decimal128("0.1")}, + {"_id": 4, "my field": "x", "año": 1, "n": "?", "g": "b", "v": None, "f": None}, +] + + +@pytest.fixture +def docs(conn): + conn.database.drop_collection(COLLECTION) + conn.database[COLLECTION].insert_many([dict(d) for d in DOCS]) + yield conn + conn.database.drop_collection(COLLECTION) + + +def run(conn, sql, params=None): + cursor = conn.cursor() + cursor.execute(sql, params) if params is not None else cursor.execute(sql) + return cursor + + +def ids(conn, sql_where, params=None): + return [r[0] for r in run(conn, f"SELECT _id FROM {COLLECTION} WHERE {sql_where} ORDER BY _id", params).fetchall()] + + +def remaining(conn): + return [d["_id"] for d in conn.database[COLLECTION].find({}, sort=[("_id", 1)])] + + +class TestLiveDML: + def test_delete_is_null_touches_only_null_dislikes(self, docs): + # dislikes IS NULL: _id 1 (null) and 4 (missing); the old translation matched + # every document lacking a field named "dis" (1 was spared, 2 and 4 deleted) + run(docs, f"DELETE FROM {COLLECTION} WHERE dislikes IS NULL") + assert remaining(docs) == [2, 3] + + def test_update_touches_only_matching_unlike(self, docs): + run(docs, f"UPDATE {COLLECTION} SET v = 99 WHERE unlike = 5") + assert {d["_id"]: d.get("v") for d in docs.database[COLLECTION].find()} == {1: 99, 2: 2, 3: 99, 4: None} + + def test_delete_with_parameters_and_literal_question_mark(self, docs): + run(docs, f"DELETE FROM {COLLECTION} WHERE n = '?' OR unlike = ?", [6]) + assert remaining(docs) == [1, 3] + + def test_delete_reversed_operand(self, docs): + run(docs, f"DELETE FROM {COLLECTION} WHERE 6 <= unlike") + assert remaining(docs) == [1, 3, 4] + + def test_unused_parameter_refuses_the_delete(self, docs): + with pytest.raises(Error): + run(docs, f"DELETE FROM {COLLECTION} WHERE unlike = ?", [5, 6]) + assert remaining(docs) == [1, 2, 3, 4] + + def test_untranslatable_update_changes_nothing(self, docs): + with pytest.raises(Error): + run(docs, f"UPDATE {COLLECTION} SET v = 0 WHERE lower(n) = 'ab'") + assert [d.get("v") for d in docs.database[COLLECTION].find({}, sort=[("_id", 1)])] == [1, 2, 2, None] + + +class TestLiveSelect: + @pytest.mark.parametrize( + "sql_where,expected", + [ + ("dislikes IS NULL", [1, 4]), + ("unlike = 5", [1, 3]), + ("5 < unlike", [2]), + ("\"my field\" = 'x'", [4]), + ('"año" = 1', [4]), + ("n LIKE '%' || 'ab' || '%'", [1, 2]), + ("n LIKE '%/%' ESCAPE '/'", [1]), + ("n LIKE 'a/_b' ESCAPE '/' AND g = 'b'", [3]), + ("n = '?'", [4]), + ("f = 0.1", [1, 3]), + ], + ) + def test_rows(self, docs, sql_where, expected): + assert ids(docs, sql_where) == expected + + def test_bound_limit_offset_and_limit_zero(self, docs): + page = run(docs, f"SELECT _id FROM {COLLECTION} WHERE v >= ? ORDER BY _id LIMIT ? OFFSET ?", [1, 2, 1]) + assert [r[0] for r in page.fetchall()] == [2, 3] + assert run(docs, f"SELECT _id FROM {COLLECTION} LIMIT 0").fetchall() == [] + assert run(docs, f"SELECT g, COUNT(*) AS n FROM {COLLECTION} GROUP BY g LIMIT ?", [0]).fetchall() == [] + grouped = run(docs, f"SELECT g, COUNT(*) AS n FROM {COLLECTION} GROUP BY g ORDER BY g LIMIT ? OFFSET ?", [1, 1]) + assert [tuple(r) for r in grouped.fetchall()] == [("b", 2)] + with pytest.raises(Error): + run(docs, f"SELECT _id FROM {COLLECTION} LIMIT ?", [-1]) + with pytest.raises(Error): + run(docs, f"SELECT _id FROM {COLLECTION} WHERE v = ?", [1, 2]) + + def test_count_distinct_and_having(self, docs): + rows = run( + docs, + f"SELECT g, COUNT(DISTINCT v) AS d, COUNT(*) AS n FROM {COLLECTION} GROUP BY g HAVING SUM(v) >= 2 ORDER BY g", + ).fetchall() + assert [tuple(r) for r in rows] == [("a", 2, 2), ("b", 1, 2)] + rows = run(docs, f"SELECT g FROM {COLLECTION} GROUP BY g HAVING COUNT(*) > 1 AND MAX(v) > ?", [1]).fetchall() + assert sorted(tuple(r) for r in rows) == [("a",), ("b",)] + + def test_decimal128_and_int64_are_python_types(self, docs): + docs.database[COLLECTION].insert_one({"_id": 5, "d": Decimal128("1.10"), "i": Int64(2**40)}) + row = run(docs, f"SELECT d, i FROM {COLLECTION} WHERE _id = 5").fetchone() + assert row == (Decimal("1.10"), 2**40) + assert type(row[0]) is Decimal and type(row[1]) is int + + +@needs_sqlalchemy +class TestLiveSQLAlchemy: + def test_like_helpers_limit_offset_and_count_distinct(self, sqlalchemy_engine, docs): + t = sa.table(COLLECTION, sa.column("_id"), sa.column("n"), sa.column("g"), sa.column("v")) + with sqlalchemy_engine.connect() as c: + assert list(c.execute(sa.select(t.c._id).where(t.c.n.contains("ab")).order_by(t.c._id)).scalars()) == [1, 2] + assert list(c.execute(sa.select(t.c._id).where(t.c.n.startswith("a")).order_by(t.c._id)).scalars()) == [ + 1, + 3, + ] + assert list(c.execute(sa.select(t.c._id).where(t.c.n.endswith("b")).order_by(t.c._id)).scalars()) == [2, 3] + assert list(c.execute(sa.select(t.c._id).where(t.c.n.contains("50%", autoescape=True))).scalars()) == [1] + assert list(c.execute(sa.select(t.c._id).where(t.c.n.ilike("AB%"))).scalars()) == [1] + assert list(c.execute(sa.select(t.c._id).where(t.c.n.like("%(k)s"))).scalars()) == [] + page = c.execute(sa.select(t.c._id).order_by(t.c._id).limit(2).offset(1)).scalars() + assert list(page) == [2, 3] + distinct = sa.func.count(sa.distinct(t.c.v)).label("d") + rows = c.execute(sa.select(t.c.g, distinct).group_by(t.c.g).order_by(t.c.g)).fetchall() + assert [tuple(r) for r in rows] == [("a", 2), ("b", 1)] diff --git a/tests/test_sql_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_sql_grouping_filters_aliases.py b/tests/test_sql_grouping_filters_aliases.py new file mode 100644 index 0000000..d6e4030 --- /dev/null +++ b/tests/test_sql_grouping_filters_aliases.py @@ -0,0 +1,178 @@ +# -*- coding: utf-8 -*- +"""GROUP BY, IN/NOT IN, LIKE, quoted literals and keyword aliases must return correct rows.""" + +import json + +import pytest + +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): + p = plan("SELECT dept, SUM(amount) AS total FROM t GROUP BY dept ORDER BY total DESC LIMIT 1 OFFSET 1") + assert json.loads(p.aggregate_pipeline)[-1] == {"$sort": {"total": -1}} + # OFFSET/LIMIT are applied after binding (they may be parameters) + assert (p.skip_stage, p.limit_stage) == (1, 1) + + def test_order_by_aggregate_expression_uses_its_output(self): + stages = pipeline("SELECT dept, SUM(amount) AS total FROM t GROUP BY dept ORDER BY SUM(amount)") + 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_filters_groups(self): + stages = pipeline("SELECT flag, COUNT(*) AS n FROM t GROUP BY flag HAVING COUNT(*) > 1") + assert stages[2] == {"$match": {"n": {"$gt": 1}}} + + def test_in_keeps_literal_types_and_quoted_commas(self): + p = plan("SELECT _id FROM t WHERE _id IN (1, 2.5, 'a,b', 'it''s', TRUE, ?)") + assert p.filter_stage == {"_id": {"$in": [1, 2.5, "a,b", "it's", True, {"$pymongosqlParam": True}]}} + + def test_not_in(self): + assert plan("SELECT _id FROM t WHERE _id NOT IN (1, 2)").filter_stage == {"_id": {"$nin": [1, 2, None]}} + + 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 == {"$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") + 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"} + + +@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_having_returns_only_matching_groups(self, grouping_collection): + got = rows(grouping_collection, f"SELECT flag, COUNT(*) AS n FROM {COLLECTION} GROUP BY flag HAVING n > 1") + assert got == [(True, 3)] + + +@pytest.mark.skipif(not HAS_SQLALCHEMY, reason="SQLAlchemy not available") +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)] 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..faeae0b --- /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 1", + ], + ) + 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_bound(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 ? AND b NOT LIKE ?" + + 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_group.py b/tests/test_sql_parser_group.py index 1f69545..2413ab4 100644 --- a/tests/test_sql_parser_group.py +++ b/tests/test_sql_parser_group.py @@ -4,6 +4,20 @@ from pymongosql.sql.parser import SQLParser +def group_of(pipeline): + """The $group stage; an aggregate without GROUP BY wraps it in $facet (one row even for no input).""" + for stage in pipeline: + if "$group" in stage: + return stage["$group"] + if "$facet" in stage: + return stage["$facet"]["row"][0]["$group"] + raise AssertionError("no $group stage") + + +def project_of(pipeline): + return next(stage["$project"] for stage in pipeline if "$project" in stage and "row" not in stage["$project"]) + + class TestCountStarParsing: """Test that COUNT(*) in SQL is translated to a MongoDB aggregate pipeline.""" @@ -17,11 +31,16 @@ def test_count_star_basic(self): assert plan.collection == "users" pipeline = json.loads(plan.aggregate_pipeline) - # Should have $group and $project stages - assert len(pipeline) == 2 - assert "$group" in pipeline[0] - assert pipeline[0]["$group"]["_id"] is None - assert pipeline[0]["$group"]["COUNT(*)"] == {"$sum": 1} + # $facet($group), the empty-input substitution, then $project + assert [next(iter(stage)) for stage in pipeline] == [ + "$facet", + "$project", + "$unwind", + "$replaceRoot", + "$project", + ] + assert group_of(pipeline)["_id"] is None + assert group_of(pipeline)["COUNT(*)"] == {"$sum": 1} def test_count_star_with_alias(self): """SELECT COUNT(*) AS total FROM users → alias used in $group""" @@ -31,10 +50,10 @@ def test_count_star_with_alias(self): assert plan.is_aggregate_query is True pipeline = json.loads(plan.aggregate_pipeline) - assert pipeline[0]["$group"]["total"] == {"$sum": 1} + assert group_of(pipeline)["total"] == {"$sum": 1} # $project should expose the alias - assert pipeline[1]["$project"]["total"] == 1 - assert pipeline[1]["$project"]["_id"] == 0 + assert project_of(pipeline)["total"] == 1 + assert project_of(pipeline)["_id"] == 0 def test_count_star_with_alias_no_as(self): """SELECT COUNT(*) total FROM users → alias without AS keyword""" @@ -44,7 +63,7 @@ def test_count_star_with_alias_no_as(self): assert plan.is_aggregate_query is True pipeline = json.loads(plan.aggregate_pipeline) - assert pipeline[0]["$group"]["total"] == {"$sum": 1} + assert group_of(pipeline)["total"] == {"$sum": 1} def test_count_star_with_where(self): """SELECT COUNT(*) AS total FROM users WHERE age > 25 → $match before $group""" @@ -54,11 +73,10 @@ def test_count_star_with_where(self): assert plan.is_aggregate_query is True pipeline = json.loads(plan.aggregate_pipeline) - # Should have $match, $group, $project - assert len(pipeline) == 3 + # $match first, then the grouping assert "$match" in pipeline[0] - assert "$group" in pipeline[1] - assert pipeline[1]["$group"]["total"] == {"$sum": 1} + assert "$facet" in pipeline[1] + assert group_of(pipeline)["total"] == {"$sum": 1} def test_count_star_projection_stage(self): """Projection stage should reflect aggregate output fields.""" @@ -83,7 +101,7 @@ def test_sum(self): assert plan.is_aggregate_query is True pipeline = json.loads(plan.aggregate_pipeline) - assert pipeline[0]["$group"]["total_price"] == {"$sum": "$price"} + assert group_of(pipeline)["total_price"] == {"$sum": "$price"} def test_avg(self): """SELECT AVG(age) AS avg_age FROM users""" @@ -93,7 +111,7 @@ def test_avg(self): assert plan.is_aggregate_query is True pipeline = json.loads(plan.aggregate_pipeline) - assert pipeline[0]["$group"]["avg_age"] == {"$avg": "$age"} + assert group_of(pipeline)["avg_age"] == {"$avg": "$age"} def test_min(self): """SELECT MIN(price) AS cheapest FROM products""" @@ -103,7 +121,7 @@ def test_min(self): assert plan.is_aggregate_query is True pipeline = json.loads(plan.aggregate_pipeline) - assert pipeline[0]["$group"]["cheapest"] == {"$min": "$price"} + assert group_of(pipeline)["cheapest"] == {"$min": "$price"} def test_max(self): """SELECT MAX(price) AS most_expensive FROM products""" @@ -113,7 +131,7 @@ def test_max(self): assert plan.is_aggregate_query is True pipeline = json.loads(plan.aggregate_pipeline) - assert pipeline[0]["$group"]["most_expensive"] == {"$max": "$price"} + assert group_of(pipeline)["most_expensive"] == {"$max": "$price"} def test_multiple_aggregates(self): """SELECT COUNT(*) AS cnt, AVG(price) AS avg_price, MAX(price) AS max_price FROM products""" @@ -123,13 +141,13 @@ def test_multiple_aggregates(self): assert plan.is_aggregate_query is True pipeline = json.loads(plan.aggregate_pipeline) - group = pipeline[0]["$group"] + group = group_of(pipeline) assert group["_id"] is None assert group["cnt"] == {"$sum": 1} assert group["avg_price"] == {"$avg": "$price"} assert group["max_price"] == {"$max": "$price"} # $project exposes all three - project = pipeline[1]["$project"] + project = project_of(pipeline) assert project == {"_id": 0, "cnt": 1, "avg_price": 1, "max_price": 1} def test_aggregate_with_nested_field(self): @@ -140,7 +158,7 @@ def test_aggregate_with_nested_field(self): assert plan.is_aggregate_query is True pipeline = json.loads(plan.aggregate_pipeline) - assert pipeline[0]["$group"]["revenue"] == {"$sum": "$details.total"} + assert group_of(pipeline)["revenue"] == {"$sum": "$details.total"} def test_aggregate_no_alias_uses_raw_text(self): """SELECT SUM(price) FROM products → alias defaults to SUM(price)""" @@ -149,7 +167,7 @@ def test_aggregate_no_alias_uses_raw_text(self): plan = parser.get_execution_plan() pipeline = json.loads(plan.aggregate_pipeline) - assert "SUM(price)" in pipeline[0]["$group"] + assert "SUM(price)" in group_of(pipeline) def test_regular_select_unaffected(self): """Regular SELECT without aggregate functions should not be affected.""" 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: diff --git a/tests/test_sqlalchemy_numeric.py b/tests/test_sqlalchemy_numeric.py new file mode 100644 index 0000000..509d779 --- /dev/null +++ b/tests/test_sqlalchemy_numeric.py @@ -0,0 +1,96 @@ +# -*- 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_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 + + +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) + + 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) 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,)] 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) 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) 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""" 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)