diff --git a/Jenkinsfile b/Jenkinsfile new file mode 100644 index 0000000..ae09605 --- /dev/null +++ b/Jenkinsfile @@ -0,0 +1,112 @@ +// Fork publisher. Pull-request wheels are versioned +pr.. +// (PEP 440-normalized by ci/release_version.py); stable wheels are published only +// from reviewed master. A published artifact is never overwritten. +podTemplate( + imagePullSecrets: ['preset-pull'], + containers: [ + containerTemplate(name: 'ci', image: 'preset/ci:latest', + ttyEnabled: true, command: 'cat'), + containerTemplate(name: 'py-ci', image: 'preset/python:3.9.18-2024-02-21-ci', + ttyEnabled: true, command: 'cat'), + // Disposable server for the test suite; reachable on localhost inside the pod. + containerTemplate(name: 'mongo', image: 'mongo:8.0', + envVars: [ + envVar(key: 'MONGO_INITDB_ROOT_USERNAME', value: 'admin'), + envVar(key: 'MONGO_INITDB_ROOT_PASSWORD', value: 'secret'), + ]) + ] +) { + node(POD_LABEL) { + checkout scm + def revision = sh(script: 'git rev-parse HEAD', returnStdout: true).trim() + boolean isMaster = env.BRANCH_NAME == 'master' + boolean isPR = env.CHANGE_ID != null + if (!isMaster && !isPR) { + error('Only master and pull-request builds publish; use a PR.') + } + + container('py-ci') { + stage('Test and build') { + def args = isMaster ? '' : "${env.CHANGE_ID} ${revision.take(12)}" + sh ''' + set -eu + python -m venv .venv + .venv/bin/pip install 'sqlalchemy==2.0.52' 'pymongo==4.17.0' \ + 'antlr4-python3-runtime==4.13.2' 'jmespath==1.1.0' 'pandas>=2.2,<3' \ + 'tenacity==9.1.2' 'pytest==8.3.5' 'boto3>=1.36,<2' 'packaging==25.0' \ + 'build==1.4.4' 'setuptools==80.9.0' 'setuptools_scm==8.3.1' 'wheel==0.45.1' + ''' + def version = sh(script: ".venv/bin/python ci/release_version.py ${args}", + returnStdout: true).trim() + env.PUBLISH_VERSION = version + env.WHEEL = "pymongosql-${version}-py3-none-any.whl" + env.KEY = "pymongosql/${env.WHEEL}" + sh ''' + set -eu + # The checkout is owned by another uid; setuptools_scm runs git during builds. + # The pod is ephemeral, so its global git config is disposable. + git --version + git config --global --add safe.directory "$PWD" + .venv/bin/pip install --no-deps -e . + for attempt in $(seq 1 30); do + .venv/bin/python -c "import pymongo; pymongo.MongoClient('mongodb://admin:secret@localhost:27017', serverSelectionTimeoutMS=2000).admin.command('ping')" && break + sleep 2 + done + .venv/bin/python tests/run_test_server.py setup + .venv/bin/python -m pytest -q tests + .venv/bin/pip uninstall -y pymongosql + python - <<'PY' +import os +import re +from pathlib import Path +path = Path('pymongosql/__init__.py') +source = path.read_text() +pattern = re.compile(r'^__version__: str = "[^"]+"$', re.MULTILINE) +assert len(pattern.findall(source)) == 1 +path.write_text(pattern.sub('__version__: str = "' + os.environ['PUBLISH_VERSION'] + '"', source)) +PY + SOURCE_DATE_EPOCH=$(git -c safe.directory="$PWD" log -1 --format=%ct) + case "$SOURCE_DATE_EPOCH" in + ''|*[!0-9]*) echo "Invalid commit timestamp for reproducible build" >&2; exit 1 ;; + esac + export SOURCE_DATE_EPOCH + # Pin the build backend and remove stale output for reproducible retries. + rm -rf build dist pymongosql.egg-info + .venv/bin/python -m build --wheel --no-isolation + test -f "dist/$WHEEL" || { echo "missing dist/$WHEEL"; ls -1 dist; exit 1; } + .venv/bin/python -c "import os, sys, zipfile; names = zipfile.ZipFile('dist/' + os.environ['WHEEL']).namelist(); sys.exit('wheel ships tests or ci' if any(n.startswith(('tests/', 'ci/')) for n in names) else 0)" + .venv/bin/pip install --force-reinstall --no-deps "dist/$WHEEL" + .venv/bin/python - <<'PY' +import importlib.metadata as im +import os +import sqlalchemy as sa +assert im.version('pymongosql') == os.environ['PUBLISH_VERSION'] +engine = sa.create_engine('mongodb://user@localhost/db') +assert engine.dialect.name == 'mongodb' +engine.dispose() +PY + sha256sum "dist/$WHEEL" + ''' + } + } + container('ci') { + stage('Publish immutable wheel') { + withCredentials([[ + $class: 'AmazonWebServicesCredentialsBinding', + credentialsId: 'ci-user', + accessKeyVariable: 'AWS_ACCESS_KEY_ID', + secretKeyVariable: 'AWS_SECRET_ACCESS_KEY' + ]]) { + withEnv(["ALLOW_IDENTICAL_PR_ARTIFACT=${isPR && !isMaster}"]) { + sh ''' + set -eu + python -m pip install --quiet 'boto3>=1.36,<2' + python ci/publish_wheel.py + ''' + } + } + } + } + archiveArtifacts artifacts: 'dist/*.whl,published.sha256', fingerprint: true + } +} diff --git a/ci/publish_wheel.py b/ci/publish_wheel.py new file mode 100644 index 0000000..894fdca --- /dev/null +++ b/ci/publish_wheel.py @@ -0,0 +1,71 @@ +"""Publish an immutable wheel; retries may reuse identical archive content.""" + +import hashlib +import io +import os +from pathlib import Path +from zipfile import BadZipFile, ZipFile + +import boto3 +from botocore.exceptions import ClientError + + +def same_wheel_content(stored, fresh): + """Compare sorted member names and bytes, ignoring ZIP metadata such as timestamps.""" + if stored == fresh: + return True + try: + with ZipFile(io.BytesIO(stored)) as old, ZipFile(io.BytesIO(fresh)) as new: + old_members = sorted(old.infolist(), key=lambda member: member.filename) + new_members = sorted(new.infolist(), key=lambda member: member.filename) + return [member.filename for member in old_members] == [member.filename for member in new_members] and all( + old.read(a) == new.read(b) for a, b in zip(old_members, new_members) + ) + except BadZipFile: + return False + + +def publish_wheel(s3, bucket, key, body, is_pr=False): + try: + s3.put_object(Bucket=bucket, Key=key, Body=body, IfNoneMatch="*") + except ClientError as error: + if error.response["Error"]["Code"] != "PreconditionFailed": + raise + print("Artifact already exists; verifying archive content without overwriting.") + + response = s3.get_object(Bucket=bucket, Key=key) + try: + stored = response["Body"].read() + finally: + response["Body"].close() + if not same_wheel_content(stored, body): + raise RuntimeError( + "Stored wheel differs from this commit's freshly built artifact; " + "refusing to overwrite or accept it (local sha256={}, stored sha256={}).".format( + hashlib.sha256(body).hexdigest(), + hashlib.sha256(stored).hexdigest(), + ) + + ("" if is_pr else " Bump __version__ before publishing different content.") + ) + print("Published wheel verified by archive content: " + key) + return hashlib.sha256(stored).hexdigest() + + +def main(): + wheel = os.environ["WHEEL"] + receipt = Path("published.sha256") + # Do not leave a previous build's success receipt after a failed retry. + if receipt.exists(): + receipt.unlink() + digest = publish_wheel( + boto3.client("s3"), + "preset-pypi", + os.environ["KEY"], + Path("dist", wheel).read_bytes(), + is_pr=os.environ.get("ALLOW_IDENTICAL_PR_ARTIFACT") == "true", + ) + receipt.write_text(digest + " " + wheel + "\n") + + +if __name__ == "__main__": + main() diff --git a/ci/release_version.py b/ci/release_version.py new file mode 100644 index 0000000..3b0abfb --- /dev/null +++ b/ci/release_version.py @@ -0,0 +1,39 @@ +"""Compute the published version of this fork, normalized as the wheel will be. + +Stable builds publish the declared ``__version__`` (a four-part Preset release +such as 0.7.4.1). Pull-request builds publish ``+pr..``. +PEP 440 normalizes local segments (lower case; a numeric segment loses leading +zeros), so the filename must be derived from the normalized form or it will not +match the file the build backend writes. +""" + +import re +import sys +from pathlib import Path + +from packaging.version import Version + +INIT = Path(__file__).resolve().parents[1] / "pymongosql" / "__init__.py" + + +def declared_version(source=None): + source = INIT.read_text() if source is None else source + matches = re.findall(r'^__version__: str = "([^"]+)"$', source, re.MULTILINE) + if len(matches) != 1: + raise SystemExit("expected exactly one __version__ declaration") + version = Version(matches[0]) + if str(version) != matches[0] or version.local or len(version.release) != 4: + raise SystemExit(f"__version__ must be a normalized four-part release, got {matches[0]!r}") + return matches[0] + + +def release_version(base, change_id=None, revision=None): + if change_id is None: + return str(Version(base)) + if not re.fullmatch(r"[0-9]+", change_id) or not re.fullmatch(r"[0-9a-fA-F]{7,40}", revision or ""): + raise SystemExit("pull-request builds need a numeric change id and a git revision") + return str(Version(f"{base}+pr.{change_id}.{revision}")) + + +if __name__ == "__main__": + print(release_version(declared_version(), *sys.argv[1:])) diff --git a/pymongosql/__init__.py b/pymongosql/__init__.py index a80c6d8..79bc12c 100644 --- a/pymongosql/__init__.py +++ b/pymongosql/__init__.py @@ -6,7 +6,7 @@ if TYPE_CHECKING: from .connection import Connection -__version__: str = "0.7.4" +__version__: str = "0.7.4.1" # Globals https://www.python.org/dev/peps/pep-0249/#globals apilevel: str = "2.0" diff --git a/pymongosql/executor.py b/pymongosql/executor.py index 9d27b65..da36382 100644 --- a/pymongosql/executor.py +++ b/pymongosql/executor.py @@ -232,6 +232,10 @@ def _execute_aggregate_plan( _logger.debug(f"Pipeline: {pipeline}") _logger.debug(f"Options: {options}") + # A pipeline generated from SQL carries the WHERE clause's ? placeholders + if parameters and execution_plan.aggregate_parameterized: + pipeline = self._replace_placeholders(pipeline, parameters) + # Get collection and call aggregate() collection = db[execution_plan.collection] diff --git a/pymongosql/helper.py b/pymongosql/helper.py index 6c1d2cb..6337899 100644 --- a/pymongosql/helper.py +++ b/pymongosql/helper.py @@ -6,9 +6,12 @@ """ import logging +from decimal import Decimal from typing import Any, Optional, Sequence, Tuple from urllib.parse import parse_qs, urlparse +from bson import Decimal128 + from .error import ProgrammingError _logger = logging.getLogger(__name__) @@ -102,6 +105,20 @@ def parse_connection_string(connection_string: Optional[str]) -> Tuple[Optional[ class SQLHelper: """SQL-related helper utilities.""" + @staticmethod + def to_bson_value(value: Any) -> Any: + """Convert a bound parameter to a type BSON can encode. + + ``decimal.Decimal`` has no BSON encoding; ``Decimal128`` stores it exactly. + """ + if isinstance(value, Decimal): + return Decimal128(value) + if isinstance(value, dict): + return {k: SQLHelper.to_bson_value(v) for k, v in value.items()} + if isinstance(value, (list, tuple)): + return [SQLHelper.to_bson_value(v) for v in value] + return value + @staticmethod def replace_placeholders_generic(value: Any, parameters: Any, style: Optional[str]) -> Any: """Recursively replace placeholders in nested structures for qmark or named styles.""" @@ -120,7 +137,7 @@ def replace(val: Any) -> Any: raise ProgrammingError("Not enough parameters provided") out = parameters[idx[0]] idx[0] += 1 - return out + return SQLHelper.to_bson_value(out) if isinstance(val, dict): return {k: replace(v) for k, v in val.items()} if isinstance(val, list): @@ -138,7 +155,7 @@ def replace(val: Any) -> Any: key = val[1:] if key not in parameters: raise ProgrammingError(f"Missing named parameter: {key}") - return parameters[key] + return SQLHelper.to_bson_value(parameters[key]) if isinstance(val, dict): return {k: replace(v) for k, v in val.items()} if isinstance(val, list): diff --git a/pymongosql/sql/ast.py b/pymongosql/sql/ast.py index 9d2372f..7a7a931 100644 --- a/pymongosql/sql/ast.py +++ b/pymongosql/sql/ast.py @@ -4,7 +4,7 @@ from ..error import SqlSyntaxError from .delete_handler import DeleteParseResult -from .handler import BaseHandler, HandlerFactory +from .handler import BaseHandler, ContextUtilsMixin, HandlerFactory from .insert_handler import InsertParseResult from .partiql.PartiQLLexer import PartiQLLexer from .partiql.PartiQLParser import PartiQLParser @@ -275,6 +275,7 @@ def visitOrderByClause(self, ctx: PartiQLParser.OrderByClauseContext) -> Any: if hasattr(ctx, "orderSortSpec") and ctx.orderSortSpec(): for sort_spec in ctx.orderSortSpec(): field_name = sort_spec.expr().getText() if sort_spec.expr() else "_id" + field_name = ContextUtilsMixin.normalize_field_path(field_name) # Check for ASC/DESC (default is ASC = 1) direction = 1 # ASC if hasattr(sort_spec, "DESC") and sort_spec.DESC(): @@ -289,6 +290,23 @@ def visitOrderByClause(self, ctx: PartiQLParser.OrderByClauseContext) -> Any: _logger.warning(f"Error processing ORDER BY clause: {e}") return self.visitChildren(ctx) + def visitGroupClause(self, ctx: PartiQLParser.GroupClauseContext) -> Any: + """Handle GROUP BY keys; they become the _id of a $group stage.""" + keys = [] + for key in ctx.groupKey() or []: + if key.symbolPrimitive() is not None: + self._query_parse_result.unsupported_clauses.append("GROUP BY key alias") + keys.append(ContextUtilsMixin.normalize_field_path(key.exprSelect().getText())) + if ctx.PARTIAL() is not None: + self._query_parse_result.unsupported_clauses.append("GROUP PARTIAL BY") + self._query_parse_result.group_by = keys + return None + + def visitHavingClause(self, ctx: PartiQLParser.HavingClauseContext) -> Any: + """HAVING is not translated; record it so the query fails instead of ignoring it.""" + self._query_parse_result.unsupported_clauses.append("HAVING") + return None + def visitLimitClause(self, ctx: PartiQLParser.LimitClauseContext) -> Any: """Handle LIMIT clause for result limiting""" _logger.debug("Processing LIMIT clause") diff --git a/pymongosql/sql/builder.py b/pymongosql/sql/builder.py index d9dff9d..db58f50 100644 --- a/pymongosql/sql/builder.py +++ b/pymongosql/sql/builder.py @@ -110,19 +110,71 @@ 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 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.) and GROUP BY + if parse_result.aggregate_functions or parse_result.group_by: return ExecutionPlanBuilder._build_sql_aggregate_plan(parse_result) + # ORDER BY may name a column by its SELECT alias; find() sorts on the field + field_for_alias = {alias: name for name, alias in parse_result.column_aliases.items()} + sort_fields = [ + {field_for_alias.get(name, name): direction for name, direction in spec.items()} + for spec in parse_result.sort_fields + ] + builder = BuilderFactory.create_query_builder().collection(parse_result.collection) builder.filter(parse_result.filter_conditions).project(parse_result.projection).column_aliases( parse_result.column_aliases - ).sort(parse_result.sort_fields).limit(parse_result.limit_value).skip(parse_result.offset_value) + ).sort(sort_fields).limit(parse_result.limit_value).skip(parse_result.offset_value) # Set aggregate flags BEFORE building (needed for validation) if hasattr(parse_result, "is_aggregate_query") and parse_result.is_aggregate_query: @@ -136,7 +188,14 @@ def _build_query_plan(parse_result: "QueryParseResult") -> "QueryExecutionPlan": @staticmethod def _build_sql_aggregate_plan(parse_result: "QueryParseResult") -> "QueryExecutionPlan": - """Build an aggregate execution plan from SQL aggregate functions like COUNT(*), SUM(), etc.""" + """Build an aggregate execution plan from SQL aggregate functions and GROUP BY. + + Pipeline: $match (WHERE), $group (GROUP BY keys as _id, one accumulator per + aggregate), $project (SELECT list, in order, under its output names), then + $sort, $skip and $limit on those output names. + """ + from ..error import NotSupportedError + _FUNCTION_TO_ACCUMULATOR = { "COUNT": "$sum", "SUM": "$sum", @@ -153,37 +212,72 @@ def _build_sql_aggregate_plan(parse_result: "QueryParseResult") -> "QueryExecuti if parse_result.filter_conditions: pipeline.append({"$match": parse_result.filter_conditions}) - # Build $group stage from aggregate functions - group_stage = {"_id": None} - for func_info in parse_result.aggregate_functions: - alias = func_info["alias"] + group_keys = {name: f"g{i}" for i, name in enumerate(parse_result.group_by)} + group_stage = {"_id": {key: f"${name}" for name, key in group_keys.items()} if group_keys else None} + accumulator_keys = [] + for i, func_info in enumerate(parse_result.aggregate_functions): func_name = func_info["function"] arg = func_info["argument"] accumulator = _FUNCTION_TO_ACCUMULATOR[func_name] - - if func_name == "COUNT": - group_stage[alias] = {accumulator: 1} + # $group output names may not contain "." or start with "$", nor repeat + key = func_info["alias"] + if "." in key or key.startswith("$") or key == "_id" or key in group_stage: + key = f"__agg{i}" + accumulator_keys.append(key) + + if func_name == "COUNT" and arg == "*": + group_stage[key] = {accumulator: 1} + elif func_name == "COUNT": + # COUNT(field) counts documents where the field is present and not null + group_stage[key] = {"$sum": {"$cond": [{"$gt": [f"${arg}", None]}, 1, 0]}} else: - group_stage[alias] = {accumulator: f"${arg}"} + group_stage[key] = {accumulator: f"${arg}"} pipeline.append({"$group": group_stage}) - # Add $project to exclude _id + # Map every SELECT item, in order, to its output name and source project_stage = {"_id": 0} - for func_info in parse_result.aggregate_functions: - project_stage[func_info["alias"]] = 1 + outputs = [] + output_for = {} # names ORDER BY may use -> output name + for item in parse_result.select_items: + if "aggregate" in item: + func_info = parse_result.aggregate_functions[item["aggregate"]] + output, key = func_info["alias"], accumulator_keys[item["aggregate"]] + source = 1 if key == output else f"${key}" + output_for[func_info["expression"].upper()] = output + else: + name = item["field"] + if name not in group_keys: + raise NotSupportedError(f"Column '{name}' must appear in GROUP BY or in an aggregate function") + output, source = item["alias"] or name, f"$_id.{group_keys[name]}" + output_for[name] = output + output_for[output] = output + project_stage[output] = source + outputs.append(output) pipeline.append({"$project": project_stage}) + sort_stage = {} + for spec in parse_result.sort_fields: + for name, direction in spec.items(): + output = output_for.get(name, output_for.get(name.upper())) + if output is None: + raise NotSupportedError(f"ORDER BY '{name}' must name a selected column or its alias") + sort_stage[output] = direction + if sort_stage: + pipeline.append({"$sort": sort_stage}) + if parse_result.offset_value: + pipeline.append({"$skip": parse_result.offset_value}) + if parse_result.limit_value is not None: + pipeline.append({"$limit": parse_result.limit_value}) + # Configure the execution plan as an aggregate query builder._execution_plan.is_aggregate_query = True + builder._execution_plan.aggregate_parameterized = True builder._execution_plan.aggregate_pipeline = json.dumps(pipeline) builder._execution_plan.aggregate_options = json.dumps({}) - # Set projection for ResultSet description - agg_projection = {} - for func_info in parse_result.aggregate_functions: - agg_projection[func_info["alias"]] = 1 - builder._execution_plan.projection_stage = agg_projection + # Set projection for ResultSet description, in SELECT order + builder._execution_plan.projection_stage = {name: 1 for name in outputs} plan = builder.build() return plan @@ -212,6 +306,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 +325,11 @@ def _build_delete_plan(parse_result: "DeleteParseResult") -> "DeleteExecutionPla @staticmethod def _build_update_plan(parse_result: "UpdateParseResult") -> "UpdateExecutionPlan": """Build an UPDATE execution plan from UPDATE parsing.""" + from ..error import SqlSyntaxError + + if parse_result.has_errors: + # An untranslated WHERE must never become an empty filter (every document) + raise SqlSyntaxError(parse_result.error_message or "UPDATE parsing failed") _logger.debug( f"Building UPDATE plan with collection: {parse_result.collection}, " f"update_fields: {parse_result.update_fields}, " diff --git a/pymongosql/sql/delete_handler.py b/pymongosql/sql/delete_handler.py index c59643b..3724c71 100644 --- a/pymongosql/sql/delete_handler.py +++ b/pymongosql/sql/delete_handler.py @@ -129,6 +129,11 @@ def handle_where_clause( _logger.debug(f"[WHERE_CLAUSE_DEBUG] Expression context type: {type(expression_ctx).__name__}") from .handler import HandlerFactory + from .where_tree import WhereTreeBuilder + + if expression_ctx is not None: + parse_result.filter_conditions = WhereTreeBuilder().build(expression_ctx) + return parse_result.filter_conditions handler = HandlerFactory.get_expression_handler(expression_ctx) diff --git a/pymongosql/sql/handler.py b/pymongosql/sql/handler.py index 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..8e45b82 100644 --- a/pymongosql/sql/query_builder.py +++ b/pymongosql/sql/query_builder.py @@ -22,6 +22,8 @@ class QueryExecutionPlan(ExecutionPlan): aggregate_pipeline: Optional[str] = None # JSON string representation of pipeline aggregate_options: Optional[str] = None # JSON string representation of options is_aggregate_query: bool = False # Flag indicating this is an aggregate() call + # True when the pipeline was generated from SQL and may hold ? placeholders + aggregate_parameterized: bool = False def to_dict(self) -> Dict[str, Any]: """Convert query plan to dictionary representation""" @@ -76,6 +78,7 @@ def copy(self) -> "QueryExecutionPlan": aggregate_pipeline=self.aggregate_pipeline, aggregate_options=self.aggregate_options, is_aggregate_query=self.is_aggregate_query, + aggregate_parameterized=self.aggregate_parameterized, ) diff --git a/pymongosql/sql/query_handler.py b/pymongosql/sql/query_handler.py index 3a09db8..5c1b2fa 100644 --- a/pymongosql/sql/query_handler.py +++ b/pymongosql/sql/query_handler.py @@ -34,6 +34,14 @@ class QueryParseResult: # SQL aggregate functions detected in SELECT (COUNT, SUM, AVG, MIN, MAX) aggregate_functions: List[Dict[str, Any]] = field(default_factory=list) + # SELECT items in order: {"field": name, "alias": alias} or {"aggregate": index} + select_items: List[Dict[str, Any]] = field(default_factory=list) + # GROUP BY field paths + group_by: List[str] = field(default_factory=list) + # Clauses that are parsed but cannot be translated faithfully + unsupported_clauses: List[str] = field(default_factory=list) + # FROM alias (FROM users AS u / FROM users u) + collection_alias: Optional[str] = None # Subquery info (for wrapped subqueries, e.g., Superset outering) subquery_plan: Optional[Any] = None @@ -137,15 +145,18 @@ def handle_visitor(self, ctx: PartiQLParser.SelectItemsContext, parse_result: "Q if agg_match: func_name = agg_match.group(1).upper() func_arg = agg_match.group(2) + parse_result.select_items.append({"aggregate": len(parse_result.aggregate_functions)}) parse_result.aggregate_functions.append( { "function": func_name, "argument": func_arg, "alias": alias or field_name, + "expression": field_name, } ) continue + parse_result.select_items.append({"field": field_name, "alias": alias}) # Use MongoDB standard projection format: {field: 1} to include field projection[field_name] = 1 # Store alias if present @@ -182,7 +193,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 +271,35 @@ def _parse_function_call(self, ctx: Any) -> Optional[Dict[str, Any]]: _logger.debug(f"Error parsing function call: {e}") return None + @staticmethod + def _collection_reference(table_ref: Any) -> Tuple[Optional[str], Optional[str], Optional[str]]: + """Return (collection text, alias, problem) for a FROM table reference. + + Only a single collection, optionally aliased, can be translated. Joins, subqueries + and AT/BY bindings are reported as a problem so the query fails instead of reading + a collection named after the whole clause. + """ + if not hasattr(table_ref, "getRuleIndex"): + return table_ref.getText(), None, None # not a parse-tree node: a plain name + while isinstance(table_ref, PartiQLParser.TableWrappedContext): + table_ref = table_ref.tableReference() + if not isinstance(table_ref, PartiQLParser.TableRefBaseContext): + return None, None, "FROM with a join" + base = table_ref.tableNonJoin().tableBaseReference() + if isinstance(base, PartiQLParser.TableBaseRefSymbolContext): + source, alias = base.source, base.symbolPrimitive().getText() + elif isinstance(base, PartiQLParser.TableBaseRefClausesContext): + if base.atIdent() is not None or base.byIdent() is not None: + return None, None, "FROM ... AT/BY" + source = base.source + alias = base.asIdent().symbolPrimitive().getText() if base.asIdent() is not None else None + else: + return None, None, "FROM with UNPIVOT or a graph match" + text = source.getText() + if text.startswith("("): + return None, None, "FROM a subquery (use mode=superset)" + return text, alias, None + def handle_visitor(self, ctx: PartiQLParser.FromClauseContext, parse_result: "QueryParseResult") -> Any: """Handle FROM clause - detect aggregate calls or regular collections""" if hasattr(ctx, "tableReference") and ctx.tableReference(): @@ -282,12 +322,16 @@ def handle_visitor(self, ctx: PartiQLParser.FromClauseContext, parse_result: "Qu _logger.info(f"Parsed aggregate call: collection={func_info['collection']}") return func_info - # Regular collection reference - table_text = ctx.tableReference().getText() + # Regular collection reference, optionally aliased + source, alias, problem = self._collection_reference(ctx.tableReference()) + if problem: + parse_result.unsupported_clauses.append(problem) + return None # Strip surrounding quotes from collection name (e.g., "user.accounts" -> user.accounts) - collection_name = self._strip_collection_quotes(table_text) + collection_name = self._strip_collection_quotes(source) parse_result.collection = collection_name - _logger.debug(f"Parsed regular collection: {collection_name}") + parse_result.collection_alias = ContextUtilsMixin.unquote_identifier(alias) if alias else None + _logger.debug(f"Parsed regular collection: {collection_name} (alias {alias})") return collection_name return None @@ -305,16 +349,15 @@ def can_handle(self, ctx: Any) -> bool: def handle_visitor(self, ctx: PartiQLParser.WhereClauseSelectContext, parse_result: "QueryParseResult") -> Any: if hasattr(ctx, "exprSelect") and ctx.exprSelect(): + from .where_tree import WhereTreeBuilder + + # Translate over the parse tree with SQL three-valued logic. A clause that + # cannot be translated fails the query; it never falls back to a partial + # or text-search filter that would return different rows. try: - # Use enhanced expression handler for better parsing - filter_conditions = self._expression_handler.handle(ctx) - parse_result.filter_conditions = filter_conditions - return filter_conditions + parse_result.filter_conditions = WhereTreeBuilder().build(ctx.exprSelect()) except Exception as e: - _logger.warning(f"Failed to parse WHERE expression, falling back to text search: {e}") - # Fallback to simple text search - filter_text = ctx.exprSelect().getText() - fallback_filter = {"$text": {"$search": filter_text}} - parse_result.filter_conditions = fallback_filter - return fallback_filter + parse_result.unsupported_clauses.append(f"WHERE ({e})") + parse_result.filter_conditions = {} + return parse_result.filter_conditions return {} diff --git a/pymongosql/sql/update_handler.py b/pymongosql/sql/update_handler.py index 6b08f57..e46f45e 100644 --- a/pymongosql/sql/update_handler.py +++ b/pymongosql/sql/update_handler.py @@ -190,6 +190,11 @@ def handle_where_clause(self, ctx: Any, parse_result: UpdateParseResult) -> Dict if expression_ctx: from .handler import HandlerFactory + from .where_tree import WhereTreeBuilder + + if expression_ctx is not None: + parse_result.filter_conditions = WhereTreeBuilder().build(expression_ctx) + return parse_result.filter_conditions handler = HandlerFactory.get_expression_handler(expression_ctx) diff --git a/pymongosql/sql/where_tree.py b/pymongosql/sql/where_tree.py new file mode 100644 index 0000000..012b033 --- /dev/null +++ b/pymongosql/sql/where_tree.py @@ -0,0 +1,143 @@ +# -*- coding: utf-8 -*- +"""WHERE translation over the parse tree with SQL three-valued logic. + +Every predicate yields two MongoDB filters: the documents for which it is TRUE and +those for which it is FALSE. A NULL (or missing) operand makes a comparison +UNKNOWN, which is in neither set. ``NOT p`` swaps the two sets, and De Morgan's +laws combine them for AND and OR, so ``NOT a = 1`` excludes documents where ``a`` +is NULL or missing, exactly like SQL. +""" + +import re +from typing import Any, Dict, List, Tuple + +from ..error import NotSupportedError +from .partiql.PartiQLParser import PartiQLParser + +Filter = Dict[str, Any] +Pair = Tuple[Filter, Filter] + +# Matches no document; used for predicates that can never be TRUE (or FALSE). +NOTHING: Filter = {"$expr": False} + +_LEAVES = ( + PartiQLParser.PredicateComparisonContext, + PartiQLParser.PredicateIsContext, + PartiQLParser.PredicateInContext, + PartiQLParser.PredicateLikeContext, + PartiQLParser.PredicateBetweenContext, +) +_FIELD_PATH = re.compile(r'^(?:"[^"]+"|[A-Za-z_$][\w$]*)(?:\.(?:"[^"]+"|[A-Za-z_$][\w$]*|\d+))*$') +_FIELD = re.compile(r"^[A-Za-z_$][\w$]*(?:\.[\w$]+)*$") +_SWAP = {"<": ">=", ">=": "<", ">": "<=", "<=": ">"} +_MONGO = {"<": "$lt", "<=": "$lte", ">": "$gt", ">=": "$gte"} + + +def _all(parts: List[Tuple[Any, Filter]], key: str, chain: type) -> Filter: + """Combine filters under ``key``, flattening only an unparenthesized chain of the same operator.""" + items: List[Filter] = [] + for ctx, f in parts: + items.extend(f[key] if isinstance(ctx, chain) and list(f) == [key] else [f]) + return {key: items} + + +def contains_not(ctx: Any) -> bool: + """Whether the expression has a boolean NOT (not a NOT IN / NOT LIKE predicate).""" + if isinstance(ctx, PartiQLParser.NotContext): + return True + return any(contains_not(child) for child in getattr(ctx, "children", None) or []) + + +def leaf_filters(field: str, operator: str, value: Any) -> Pair: + """TRUE and FALSE filters for one predicate on ``field``.""" + op = operator.upper() + if op == "IS NULL": + return {field: {"$eq": None}}, {field: {"$ne": None}} + if op == "IS NOT NULL": + return {field: {"$ne": None}}, {field: {"$eq": None}} + if op in ("IN", "NOT IN"): + values = value if isinstance(value, list) else [value] + present = [v for v in values if v is not None] + true = {field: {"$in": present}} + # x NOT IN (..., NULL) is never TRUE; x IN (..., NULL) is never FALSE + false = NOTHING if None in values else {field: {"$nin": present + [None]}} + return (true, false) if op == "IN" else (false, true) + if op in ("LIKE", "NOT LIKE"): + from .handler import ComparisonExpressionHandler + + if value == "?" or not isinstance(value, str): + # The pattern is translated to a regex while parsing, before parameters are bound + raise NotSupportedError("LIKE needs a literal pattern, not a bound parameter") + + pattern = ComparisonExpressionHandler._like_to_regex(value) + pattern = ("" if pattern.startswith(".*") else "^") + pattern + ("" if pattern.endswith(".*") else "$") + true = {field: {"$regex": pattern}} + false = {"$and": [{field: {"$not": {"$regex": pattern}}}, {field: {"$ne": None}}]} + return (true, false) if op == "LIKE" else (false, true) + if op == "BETWEEN": + low, high = value + true = {"$and": [{field: {"$gte": low}}, {field: {"$lte": high}}]} + return true, {"$or": [{field: {"$lt": low}}, {field: {"$gt": high}}]} + if value is None: + # This dialect has always read "= NULL" / "<> NULL" as IS NULL / IS NOT NULL; + # any other comparison with NULL is UNKNOWN. + if op == "=": + return {field: None}, {field: {"$ne": None}} + if op in ("!=", "<>"): + return {field: {"$ne": None}}, {field: None} + return NOTHING, NOTHING + if op == "=": + return {field: value}, {field: {"$nin": [value, None]}} + if op in ("!=", "<>"): + return {field: {"$nin": [value, None]}}, {field: value} + if op in _MONGO: + return {field: {_MONGO[op]: value}}, {field: {_MONGO[_SWAP[op]]: value}} + raise NotSupportedError(f"Unsupported predicate operator: {operator}") + + +class WhereTreeBuilder: + """Build a MongoDB filter for a WHERE expression from its parse tree.""" + + def build(self, ctx: Any) -> Filter: + return self._pair(ctx)[0] + + def _pair(self, ctx: Any) -> Pair: + if isinstance(ctx, PartiQLParser.NotContext): + true, false = self._pair(ctx.rhs) + return false, true + if isinstance(ctx, (PartiQLParser.AndContext, PartiQLParser.OrContext)): + is_and = isinstance(ctx, PartiQLParser.AndContext) + (t1, f1), (t2, f2) = self._pair(ctx.lhs), self._pair(ctx.rhs) + true_key, false_key = ("$and", "$or") if is_and else ("$or", "$and") + chain = type(ctx) + return ( + _all([(ctx.lhs, t1), (ctx.rhs, t2)], true_key, chain), + _all([(ctx.lhs, f1), (ctx.rhs, f2)], false_key, chain), + ) + if isinstance(ctx, PartiQLParser.ExprTermWrappedQueryContext): + return self._pair(ctx.expr()) + if isinstance(ctx, _LEAVES): + return self._leaf(ctx) + children = [c for c in getattr(ctx, "children", None) or [] if hasattr(c, "getRuleIndex")] + if len(children) == 1 and len(ctx.children) == 1: + return self._pair(children[0]) + text = ctx.getText() + if _FIELD_PATH.match(text) and text.upper() not in ("TRUE", "FALSE", "NULL"): + # A bare boolean field: WHERE flag / WHERE NOT flag + from .handler import ContextUtilsMixin + + field = ContextUtilsMixin.normalize_field_path(text) + return {field: True}, {field: False} + raise NotSupportedError(f"Unsupported WHERE expression: {text}") + + @staticmethod + def _leaf(ctx: Any) -> Pair: + from .handler import ComparisonExpressionHandler + + handler = ComparisonExpressionHandler() + field = handler._extract_field_name(ctx) + operator = handler._extract_operator(ctx) + value = handler._extract_value(ctx) + if not _FIELD.match(field): + raise NotSupportedError(f"Unsupported WHERE predicate: {ctx.getText()}") + return leaf_filters(field, operator, value) diff --git a/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py b/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py index 540a0f1..4b40fa6 100644 --- a/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py +++ b/pymongosql/sqlalchemy_mongodb/sqlalchemy_dialect.py @@ -1,11 +1,12 @@ # -*- coding: utf-8 -*- import logging +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 +33,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 +53,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 +101,24 @@ 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 visit_like_op_binary(self, binary, operator, **kw): + """Render LIKE patterns inline: PyMongoSQL turns them into a regex while parsing.""" + kw["literal_binds"] = True + return super().visit_like_op_binary(binary, operator, **kw) + + def visit_not_like_op_binary(self, binary, operator, **kw): + kw["literal_binds"] = True + return super().visit_not_like_op_binary(binary, operator, **kw) class PyMongoSQLDDLCompiler(compiler.DDLCompiler): @@ -151,6 +178,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 +275,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 +490,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 +537,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 +545,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 +582,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/executor.py b/pymongosql/superset_mongodb/executor.py index f7090bc..22cf24f 100644 --- a/pymongosql/superset_mongodb/executor.py +++ b/pymongosql/superset_mongodb/executor.py @@ -63,7 +63,10 @@ def execute( _logger.debug(f"Stage 1: Executing MongoDB subquery: {mongo_query}") mongo_execution_plan = self._parse_sql(mongo_query) - mongo_result = self._execute_find_plan(mongo_execution_plan, connection) + if mongo_execution_plan.is_aggregate_query: + mongo_result = self._execute_aggregate_plan(mongo_execution_plan, connection) + else: + mongo_result = self._execute_find_plan(mongo_execution_plan, connection) # Extract result set from MongoDB mongo_result_set = ResultSet( diff --git a/pymongosql/superset_mongodb/query_db_sqlite.py b/pymongosql/superset_mongodb/query_db_sqlite.py index 5e3d356..348a485 100644 --- a/pymongosql/superset_mongodb/query_db_sqlite.py +++ b/pymongosql/superset_mongodb/query_db_sqlite.py @@ -7,6 +7,12 @@ _logger = logging.getLogger(__name__) +# SQLite has no boolean type. Boolean columns are declared with this private type +# (NUMERIC affinity, stored as 0/1) and converted back to bool when a query selects +# the column itself; expressions over it (SUM, CASE, ...) keep their numeric result. +BOOLEAN_DECLTYPE = "PYMONGOSQL_BOOL" +sqlite3.register_converter(BOOLEAN_DECLTYPE, lambda raw: int(raw) != 0) + class SQLiteTypeMapper: """Maps Python/MongoDB data types to SQLite3 types""" @@ -16,7 +22,7 @@ class SQLiteTypeMapper: str: "TEXT", int: "INTEGER", float: "REAL", - bool: "INTEGER", # SQLite3 uses 0/1 for boolean + bool: BOOLEAN_DECLTYPE, # stored as 0/1, read back as bool bytes: "BLOB", type(None): "NULL", dict: "TEXT", # Store as JSON string @@ -51,15 +57,14 @@ def infer_schema(cls, records: List[Dict[str, Any]]) -> Dict[str, str]: for record in records: for col_name, value in record.items(): - if col_name not in schema: - # First occurrence, determine type - schema[col_name] = cls.get_sqlite_type(value) - elif schema[col_name] != "TEXT": - # If we've already determined type, check compatibility - new_type = cls.get_sqlite_type(value) + new_type = cls.get_sqlite_type(value) + current = schema.get(col_name, "NULL") + if current == "NULL": + # First non-null value determines the type; NULL fits every type + schema[col_name] = new_type + elif new_type not in ("NULL", current): # Upgrade to TEXT if types differ (safest option) - if new_type != schema[col_name]: - schema[col_name] = "TEXT" + schema[col_name] = "TEXT" return schema @@ -69,7 +74,7 @@ def convert_value(cls, value: Any, target_type: str) -> Any: if value is None: return None - if target_type == "INTEGER": + if target_type in ("INTEGER", BOOLEAN_DECLTYPE): return int(value) if value is not None else None elif target_type == "REAL": return float(value) if value is not None else None @@ -107,7 +112,7 @@ def _ensure_connection(self) -> sqlite3.Connection: if self._connection is None: # Create in-memory database - self._connection = sqlite3.connect(":memory:") + self._connection = sqlite3.connect(":memory:", detect_types=sqlite3.PARSE_DECLTYPES) # Enable row factory to get dict-like rows self._connection.row_factory = sqlite3.Row _logger.debug("Created in-memory SQLite3 database") diff --git a/tests/ci/__init__.py b/tests/ci/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/tests/ci/test_publish_wheel.py b/tests/ci/test_publish_wheel.py new file mode 100644 index 0000000..7412ec6 --- /dev/null +++ b/tests/ci/test_publish_wheel.py @@ -0,0 +1,142 @@ +"""Exercise the same immutable publisher used by Jenkins without AWS access.""" + +import hashlib +import io +import runpy +from pathlib import Path +from unittest.mock import Mock +from zipfile import ZipFile, ZipInfo + +import pytest +from botocore.exceptions import ClientError + +publisher = runpy.run_path(str(Path(__file__).resolve().parents[2] / "ci" / "publish_wheel.py")) +publish_wheel = publisher["publish_wheel"] + + +def wheel(members=None, year=2020): + output = io.BytesIO() + if members is None: + members = [("package.py", b"code"), ("metadata", b"version")] + with ZipFile(output, "w") as archive: + for name, body in members: + archive.writestr(ZipInfo(name, (year, 1, 1, 0, 0, 0)), body) + return output.getvalue() + + +def client(stored=None, error=None): + s3 = Mock() + s3.get_object.return_value = {"Body": io.BytesIO(wheel() if stored is None else stored)} + if error: + s3.put_object.side_effect = ClientError({"Error": {"Code": error}}, "PutObject") + return s3 + + +def test_new_artifact_is_conditionally_written_and_verified(): + body = wheel() + s3 = client(body) + assert publish_wheel(s3, "bucket", "key", body) == hashlib.sha256(body).hexdigest() + s3.put_object.assert_called_once_with( + Bucket="bucket", + Key="key", + Body=body, + IfNoneMatch="*", + ) + s3.get_object.assert_called_once_with(Bucket="bucket", Key="key") + + +@pytest.mark.parametrize("is_pr", [False, True]) +@pytest.mark.parametrize("variation", ["identical", "timestamps", "member_order"]) +def test_identical_content_retry_succeeds_without_overwrite(is_pr, variation): + stored = wheel() + fresh = wheel(year=2021) if variation == "timestamps" else stored + if variation == "member_order": + fresh = wheel([("metadata", b"version"), ("package.py", b"code")]) + if variation != "identical": + assert fresh != stored + s3 = client(stored, "PreconditionFailed") + assert publish_wheel(s3, "bucket", "key", fresh, is_pr) == hashlib.sha256(stored).hexdigest() + s3.put_object.assert_called_once_with( + Bucket="bucket", + Key="key", + Body=fresh, + IfNoneMatch="*", + ) + assert s3.get_object.return_value["Body"].closed + + +@pytest.mark.parametrize("is_pr", [False, True]) +@pytest.mark.parametrize("error", [None, "PreconditionFailed"]) +@pytest.mark.parametrize( + "members", + [ + [("package.py", b"changed"), ("metadata", b"version")], + [("renamed.py", b"code"), ("metadata", b"version")], + [("package.py", b"code")], + [("package.py", b"code"), ("metadata", b"version"), ("extra", b"")], + ], +) +def test_different_content_fails(is_pr, error, members): + s3 = client(wheel(members), error) + with pytest.raises(RuntimeError, match="Stored wheel differs.*refusing to overwrite") as exc: + publish_wheel(s3, "bucket", "key", wheel(), is_pr) + assert ("Bump __version__" in str(exc.value)) == (not is_pr) + assert s3.put_object.call_count == 1 + assert s3.get_object.return_value["Body"].closed + + +@pytest.mark.parametrize("is_pr", [False, True]) +@pytest.mark.parametrize("error", ["AccessDenied", "ConditionalRequestConflict"]) +def test_other_s3_errors_are_reraised(is_pr, error): + s3 = client(error=error) + with pytest.raises(ClientError) as exc: + publish_wheel(s3, "bucket", "key", wheel(), is_pr) + assert exc.value is s3.put_object.side_effect + s3.get_object.assert_not_called() + + +def test_unreadable_existing_artifact_fails_closed(): + s3 = client(error="PreconditionFailed") + s3.get_object.side_effect = ClientError({"Error": {"Code": "AccessDenied"}}, "GetObject") + with pytest.raises(ClientError, match="AccessDenied"): + publish_wheel(s3, "bucket", "key", wheel(), True) + + +@pytest.mark.parametrize("is_pr", [False, True]) +def test_invalid_existing_archive_fails_closed(is_pr): + s3 = client(b"not a zip", "PreconditionFailed") + with pytest.raises(RuntimeError, match="Stored wheel differs"): + publish_wheel(s3, "bucket", "key", wheel(), is_pr) + + +@pytest.mark.parametrize("is_pr", [False, True]) +def test_receipt_records_stored_digest(tmp_path, monkeypatch, is_pr): + stored, fresh = wheel(), wheel(year=2021) + s3 = client(stored, "PreconditionFailed") + monkeypatch.chdir(tmp_path) + monkeypatch.setenv("WHEEL", "test.whl") + monkeypatch.setenv("KEY", "test-key") + monkeypatch.setenv("ALLOW_IDENTICAL_PR_ARTIFACT", str(is_pr).lower()) + monkeypatch.setattr(publisher["boto3"], "client", lambda service: s3) + Path("dist").mkdir() + Path("dist/test.whl").write_bytes(fresh) + Path("published.sha256").write_text("stale receipt") + publisher["main"]() + assert Path("published.sha256").read_text() == hashlib.sha256(stored).hexdigest() + " test.whl\n" + + +def test_failed_retry_removes_stale_receipt(tmp_path, monkeypatch): + monkeypatch.chdir(tmp_path) + monkeypatch.setenv("WHEEL", "test.whl") + monkeypatch.setenv("KEY", "test-key") + monkeypatch.setattr( + publisher["boto3"], + "client", + lambda service: client(b"bad", "PreconditionFailed"), + ) + Path("dist").mkdir() + Path("dist/test.whl").write_bytes(wheel()) + Path("published.sha256").write_text("stale receipt") + with pytest.raises(RuntimeError): + publisher["main"]() + assert not Path("published.sha256").exists() diff --git a/tests/ci/test_release_version.py b/tests/ci/test_release_version.py new file mode 100644 index 0000000..85eaeaf --- /dev/null +++ b/tests/ci/test_release_version.py @@ -0,0 +1,40 @@ +"""The published version must equal the one the build backend writes.""" + +import runpy +from pathlib import Path + +import pytest + +module = runpy.run_path(str(Path(__file__).resolve().parents[2] / "ci" / "release_version.py")) +release_version = module["release_version"] +declared_version = module["declared_version"] + + +@pytest.mark.parametrize( + "change,revision,expected", + [ + (None, None, "0.7.4.1"), + ("3", "e2bc688f1a2b", "0.7.4.1+pr.3.e2bc688f1a2b"), + ("3", "ABCDEF123456", "0.7.4.1+pr.3.abcdef123456"), + ("3", "012345678901", "0.7.4.1+pr.3.12345678901"), + ], +) +def test_release_version_is_normalized(change, revision, expected): + assert release_version("0.7.4.1", change, revision) == expected + + +@pytest.mark.parametrize("change,revision", [("PR-3", "e2bc688"), ("3", "not-a-sha"), ("3", None)]) +def test_bad_pull_request_inputs_fail(change, revision): + with pytest.raises(SystemExit): + release_version("0.7.4.1", change, revision) + + +def test_declared_version_is_a_four_part_release(): + assert declared_version('__version__: str = "0.7.4.1"\n') == "0.7.4.1" + assert len(declared_version().split(".")) == 4 + + +@pytest.mark.parametrize("source", ['__version__: str = "0.7.4"\n', '__version__: str = "0.7.4.1+x"\n', ""]) +def test_declared_version_rejects_other_forms(source): + with pytest.raises(SystemExit): + declared_version(source) 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..6590a76 --- /dev/null +++ b/tests/test_sql_grouping_filters_aliases.py @@ -0,0 +1,177 @@ +# -*- coding: utf-8 -*- +"""GROUP BY, IN/NOT IN, LIKE, quoted literals and keyword aliases must return correct rows.""" + +import json + +import pytest + +from pymongosql.error import Error +from pymongosql.sql.parser import SQLParser +from tests.conftest import HAS_SQLALCHEMY, make_superset_conn + +COLLECTION = "test_grouping_filters" +DOCS = [ + {"_id": 1, "flag": True, "dept": "a", "amount": 10, "name": "O'Brien"}, + {"_id": 2, "flag": False, "dept": "a", "amount": 20, "name": "x.y"}, + {"_id": 3, "flag": True, "dept": "b", "amount": 40, "name": "xzy"}, + {"_id": 4, "flag": True, "dept": "b", "amount": None, "name": "plain"}, +] + + +def plan(sql): + return SQLParser(sql).get_execution_plan() + + +def pipeline(sql): + return json.loads(plan(sql).aggregate_pipeline) + + +class TestPlans: + def test_group_by_groups_on_the_key(self): + stages = pipeline("SELECT flag, COUNT(*) AS n FROM t GROUP BY flag") + assert stages[0]["$group"]["_id"] == {"g0": "$flag"} + assert stages[1]["$project"] == {"_id": 0, "flag": "$_id.g0", "n": 1} + + def test_aggregate_query_keeps_order_by_skip_and_limit(self): + stages = pipeline("SELECT dept, SUM(amount) AS total FROM t GROUP BY dept ORDER BY total DESC LIMIT 1 OFFSET 1") + assert stages[-3:] == [{"$sort": {"total": -1}}, {"$skip": 1}, {"$limit": 1}] + + def test_order_by_aggregate_expression_uses_its_output(self): + stages = pipeline("SELECT dept, SUM(amount) AS total FROM t GROUP BY dept ORDER BY SUM(amount)") + assert stages[-1] == {"$sort": {"total": 1}} + + def test_quoted_keyword_alias_is_unquoted(self): + p = plan('SELECT flag, COUNT(*) AS "count" FROM t GROUP BY flag ORDER BY "count" DESC') + assert list(p.projection_stage) == ["flag", "count"] + assert json.loads(p.aggregate_pipeline)[-1] == {"$sort": {"count": -1}} + + def test_find_order_by_alias_sorts_on_the_field(self): + p = plan('SELECT amount AS "value" FROM t ORDER BY "value" DESC') + assert p.sort_stage == [{"amount": -1}] + assert p.column_aliases == {"amount": "value"} + + def test_ungrouped_column_is_rejected(self): + with pytest.raises(Exception, match="must appear in GROUP BY"): + plan("SELECT flag, COUNT(*) FROM t") + + def test_having_is_rejected_not_ignored(self): + with pytest.raises(Exception, match="HAVING"): + plan("SELECT flag, COUNT(*) AS n FROM t GROUP BY flag HAVING COUNT(*) > 1") + + def test_in_keeps_literal_types_and_quoted_commas(self): + p = plan("SELECT _id FROM t WHERE _id IN (1, 2.5, 'a,b', 'it''s', TRUE, ?)") + assert p.filter_stage == {"_id": {"$in": [1, 2.5, "a,b", "it's", True, "?"]}} + + def test_not_in(self): + assert plan("SELECT _id FROM t WHERE _id NOT IN (1, 2)").filter_stage == {"_id": {"$nin": [1, 2, 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_unsupported_having_raises(self, grouping_collection): + with pytest.raises(Error): + rows(grouping_collection, f"SELECT flag, COUNT(*) AS n FROM {COLLECTION} GROUP BY flag HAVING n > 1") + + +@pytest.mark.skipif(not HAS_SQLALCHEMY, reason="SQLAlchemy not available") +class TestLiveSQLAlchemy: + def test_core_group_by_count_label(self, sqlalchemy_engine, grouping_collection): + import sqlalchemy as sa + + t = sa.table(COLLECTION, sa.column("flag"), sa.column("_id")) + count = sa.func.count().label("count") + stmt = sa.select(t.c.flag, count).where(t.c._id.in_([1, 2, 3])).group_by(t.c.flag).order_by(count.desc()) + with sqlalchemy_engine.connect() as connection: + assert [tuple(r) for r in connection.execute(stmt)] == [(True, 2), (False, 1)] diff --git a/tests/test_sql_not_and_null_semantics.py b/tests/test_sql_not_and_null_semantics.py new file mode 100644 index 0000000..c1f972f --- /dev/null +++ b/tests/test_sql_not_and_null_semantics.py @@ -0,0 +1,142 @@ +# -*- coding: utf-8 -*- +"""NOT and NULL follow SQL three-valued logic; an untranslatable WHERE never widens a query.""" + +import pytest + +from pymongosql.error import Error +from pymongosql.sql.parser import SQLParser +from tests.conftest import HAS_SQLALCHEMY + +COLLECTION = "test_not_semantics" +# _id 4 has a NULL a; _id 5 has no a at all. SQL never returns them for a +# comparison on a, negated or not. +DOCS = [ + {"_id": 1, "a": 1, "b": "x", "flag": True}, + {"_id": 2, "a": 2, "b": "y", "flag": False}, + {"_id": 3, "a": 3, "b": "xz", "flag": True}, + {"_id": 4, "a": None, "b": None, "flag": None}, + {"_id": 5}, +] + + +def plan(sql): + return SQLParser(sql).get_execution_plan() + + +class TestPlans: + def test_not_comparison_excludes_null(self): + assert plan("SELECT _id FROM t WHERE NOT a = 1").filter_stage == {"a": {"$nin": [1, None]}} + + def test_not_or_uses_de_morgan(self): + assert plan("SELECT _id FROM t WHERE NOT (a = 1 OR b = 'x')").filter_stage == { + "$and": [{"a": {"$nin": [1, None]}}, {"b": {"$nin": ["x", None]}}] + } + + def test_double_not(self): + assert plan("SELECT _id FROM t WHERE NOT NOT a = 1").filter_stage == {"a": 1} + + def test_not_between_and_not_is_null(self): + assert plan("SELECT _id FROM t WHERE NOT a BETWEEN 1 AND 2").filter_stage == { + "$or": [{"a": {"$lt": 1}}, {"a": {"$gt": 2}}] + } + assert plan("SELECT _id FROM t WHERE NOT a IS NULL").filter_stage == {"a": {"$ne": None}} + + def test_not_bare_boolean_field(self): + assert plan("SELECT _id FROM t WHERE NOT flag").filter_stage == {"flag": False} + + def test_not_in_list_containing_null_is_never_true(self): + assert plan("SELECT _id FROM t WHERE a NOT IN (1, NULL)").filter_stage == {"$expr": False} + + def test_not_equal_excludes_null(self): + assert plan("SELECT _id FROM t WHERE a <> 1").filter_stage == {"a": {"$nin": [1, None]}} + + @pytest.mark.parametrize( + "sql", + [ + "SELECT _id FROM t WHERE NOT lower(b) = 'x'", + "SELECT _id FROM t WHERE lower(b) = 'x'", + "SELECT _id FROM t WHERE a = 1 AND lower(b) = 'x'", + "SELECT _id FROM t WHERE b LIKE ?", + ], + ) + def test_untranslatable_where_raises_instead_of_widening(self, sql): + with pytest.raises(Error): + plan(sql) + + @pytest.mark.parametrize("sql", ["DELETE FROM t WHERE lower(b) = 'x'", "UPDATE t SET a = 1 WHERE lower(b) = 'x'"]) + def test_untranslatable_dml_where_raises_instead_of_matching_everything(self, sql): + with pytest.raises(Error): + plan(sql) + + def test_not_in_delete(self): + assert plan("DELETE FROM t WHERE NOT a = 1").filter_conditions == {"a": {"$nin": [1, None]}} + + +@pytest.fixture +def docs(conn): + conn.database.drop_collection(COLLECTION) + conn.database[COLLECTION].insert_many(DOCS) + yield conn + conn.database.drop_collection(COLLECTION) + + +def ids(conn, where, params=None): + cursor = conn.cursor() + sql = f"SELECT _id FROM {COLLECTION} WHERE {where} ORDER BY _id" + cursor.execute(sql, params) if params is not None else cursor.execute(sql) + return [r[0] for r in cursor.fetchall()] + + +class TestLive: + @pytest.mark.parametrize( + "where,expected", + [ + ("NOT a = 1", [2, 3]), + ("NOT (a = 1 OR b = 'y')", [3]), + ("NOT (a = 1 AND b = 'x')", [2, 3]), + ("NOT a IN (1, 2)", [3]), + ("NOT a NOT IN (1, 2)", [1, 2]), + ("NOT b LIKE 'x%'", [2]), + ("NOT a BETWEEN 2 AND 3", [1]), + ("NOT a IS NULL", [1, 2, 3]), + ("NOT flag", [2]), + ("a <> 1", [2, 3]), + ("b NOT LIKE 'x%'", [2]), + ("a NOT IN (1)", [2, 3]), + ("a = 1 OR NOT (b = 'x' OR b = 'xz')", [1, 2]), + ], + ) + def test_negation_returns_sql_rows(self, docs, where, expected): + assert ids(docs, where) == expected + + def test_not_with_bound_parameter(self, docs): + assert ids(docs, "NOT a = ?", [2]) == [1, 3] + + def test_untranslatable_delete_deletes_nothing(self, docs): + cursor = docs.cursor() + with pytest.raises(Error): + cursor.execute(f"DELETE FROM {COLLECTION} WHERE lower(b) = 'x'") + assert docs.database[COLLECTION].count_documents({}) == len(DOCS) + + +@pytest.mark.skipif(not HAS_SQLALCHEMY, reason="SQLAlchemy not available") +class TestSQLAlchemy: + def test_like_pattern_is_rendered_inline(self): + import sqlalchemy as sa + + from pymongosql.sqlalchemy_mongodb.sqlalchemy_dialect import PyMongoSQLDialect + + t = sa.table("t", sa.column("b")) + stmt = sa.select(t.c.b).where(t.c.b.like("O'B%"), t.c.b.not_like("x_")) + sql = " ".join(str(stmt.compile(dialect=PyMongoSQLDialect())).split()) + assert sql == "SELECT b FROM t WHERE b LIKE 'O''B%' AND b NOT LIKE 'x_'" + + def test_core_not_and_like_rows(self, sqlalchemy_engine, docs): + import sqlalchemy as sa + + t = sa.table(COLLECTION, sa.column("_id"), sa.column("a"), sa.column("b")) + with sqlalchemy_engine.connect() as connection: + got = connection.execute( + sa.select(t.c._id).where(sa.not_(t.c.a == 1), t.c.b.like("x%")).order_by(t.c._id) + ).scalars() + assert list(got) == [3] diff --git a/tests/test_sql_parser_comprehensive.py b/tests/test_sql_parser_comprehensive.py index ec19beb..1039787 100644 --- a/tests/test_sql_parser_comprehensive.py +++ b/tests/test_sql_parser_comprehensive.py @@ -157,7 +157,7 @@ def test_bool_and_bracketed_or(self): def test_bool_not_equal_and_comparison(self): sql = "SELECT * FROM col WHERE active!=false AND age>25" plan = SQLParser(sql).get_execution_plan() - assert plan.filter_stage == {"$and": [{"active": {"$ne": False}}, {"age": {"$gt": 25}}]} + assert plan.filter_stage == {"$and": [{"active": {"$nin": [False, None]}}, {"age": {"$gt": 25}}]} # --- null mixed with bool --- diff --git a/tests/test_sql_parser_delete.py b/tests/test_sql_parser_delete.py index 395f37b..83bf6f5 100644 --- a/tests/test_sql_parser_delete.py +++ b/tests/test_sql_parser_delete.py @@ -67,7 +67,7 @@ def test_delete_with_not_equal(self): assert isinstance(plan, DeleteExecutionPlan) assert plan.collection == "temp" - assert plan.filter_conditions == {"valid": {"$ne": True}} + assert plan.filter_conditions == {"valid": {"$nin": [True, None]}} def test_delete_with_qmark_parameter(self): """Test DELETE with qmark placeholder.""" diff --git a/tests/test_sql_parser_general.py b/tests/test_sql_parser_general.py index 8354bdf..cd1cdd0 100644 --- a/tests/test_sql_parser_general.py +++ b/tests/test_sql_parser_general.py @@ -101,7 +101,7 @@ def test_select_with_not_equals(self): execution_plan = parser.get_execution_plan() assert execution_plan.collection == "users" - assert execution_plan.filter_stage == {"status": {"$ne": "inactive"}} + assert execution_plan.filter_stage == {"status": {"$nin": ["inactive", None]}} assert execution_plan.projection_stage == {"name": 1} def test_select_with_and_condition(self): @@ -335,7 +335,7 @@ def test_complex_mixed_operators(self): # Verify complex filter structure with mixed AND/OR conditions expected_filter = { "$or": [ - {"$and": [{"age": {"$gt": 25}}, {"status": "active"}, {"name": {"$ne": "John"}}]}, + {"$and": [{"age": {"$gt": 25}}, {"status": "active"}, {"name": {"$nin": ["John", None]}}]}, {"department": {"$in": ["IT", "HR"]}}, ] } diff --git a/tests/test_sql_parser_nested_fields.py b/tests/test_sql_parser_nested_fields.py index eeed246..5cf6999 100644 --- a/tests/test_sql_parser_nested_fields.py +++ b/tests/test_sql_parser_nested_fields.py @@ -146,7 +146,7 @@ def test_nested_with_comparison_operators(self): ("profile.age > 18", {"profile.age": {"$gt": 18}}), ("settings.total < 100", {"settings.total": {"$lt": 100}}), # Changed from 'count' (reserved) ("status.active = true", {"status.active": True}), - ("config.name != 'default'", {"config.name": {"$ne": "default"}}), + ("config.name != 'default'", {"config.name": {"$nin": ["default", None]}}), ] for where_clause, expected_filter in test_cases: 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)