From f8b8c58c5cd11be432d13c39b79f973f72683034 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Mon, 8 Jun 2026 17:49:24 +0000 Subject: [PATCH 01/12] feat(core): carry declared_parameters slot through SQL lifecycle Add a _declared_parameters slot to SQL and thread it through all seven construction/reset paths (__init__, _init_from_sql_object, copy full + fast path via _create_empty_copy, reset, _create_cached_direct, plus the declared_parameters property). Proven mypyc-safe and pool-leak-free against the compiled module. Spike smgc.7 for sqlspec-smgc (gh-491). --- sqlspec/core/statement.py | 11 ++++ tests/unit/core/test_declared_params_spike.py | 59 +++++++++++++++++++ 2 files changed, 70 insertions(+) create mode 100644 tests/unit/core/test_declared_params_spike.py diff --git a/sqlspec/core/statement.py b/sqlspec/core/statement.py index 1577f8dfc..1de465853 100644 --- a/sqlspec/core/statement.py +++ b/sqlspec/core/statement.py @@ -169,6 +169,7 @@ def _parse_order_item(order_item: str, dialect: "str | None", enable_parsing: bo SQL_SLOTS: Final = ( "_compiled_from_cache", + "_declared_parameters", "_dialect", "_filters", "_hash", @@ -335,6 +336,7 @@ def __init__( self._is_script = False self._raw_expression: exp.Expr | None = None self._rebind_processor: ParameterProcessor | None = None + self._declared_parameters: "tuple[Any, ...]" = () if isinstance(statement, SQL): self._init_from_sql_object(statement) @@ -443,6 +445,7 @@ def reset(self) -> None: self._statement_config = get_default_config() self._dialect = self._normalize_dialect(self._statement_config.dialect) self._rebind_processor = None + self._declared_parameters = () def _normalize_dialect(self, dialect: "DialectType") -> "str | None": """Convert dialect to string representation. @@ -480,6 +483,7 @@ def _init_from_sql_object(self, sql_obj: "SQL") -> None: self._sql_param_counters = sql_obj._sql_param_counters.copy() self._is_many = sql_obj.is_many self._is_script = sql_obj.is_script + self._declared_parameters = sql_obj._declared_parameters if sql_obj.is_processed: self._processed_state = sql_obj.get_processed_state() @@ -611,6 +615,11 @@ def original_parameters(self) -> Any: """Get original parameters (public API).""" return self._original_parameters + @property + def declared_parameters(self) -> "tuple[Any, ...]": + """Get declared parameter metadata carried with this statement (public API).""" + return self._declared_parameters + @property def operation_type(self) -> "OperationType": """SQL operation type.""" @@ -911,6 +920,7 @@ def copy(self, statement: "str | exp.Expr | None" = None, parameters: Any | None new_sql._named_parameters.update(self._named_parameters) new_sql._positional_parameters = self._positional_parameters.copy() new_sql._filters = self._filters.copy() + new_sql._declared_parameters = self._declared_parameters return new_sql def _create_empty_copy(self) -> "SQL": @@ -933,6 +943,7 @@ def _create_empty_copy(self) -> "SQL": new_sql._named_parameters = {} new_sql._positional_parameters = [] new_sql._sql_param_counters = self._sql_param_counters.copy() + new_sql._declared_parameters = self._declared_parameters return new_sql diff --git a/tests/unit/core/test_declared_params_spike.py b/tests/unit/core/test_declared_params_spike.py new file mode 100644 index 000000000..fb35c7008 --- /dev/null +++ b/tests/unit/core/test_declared_params_spike.py @@ -0,0 +1,59 @@ +"""SPIKE (throwaway): prove declared_parameters slot carriage + pool-leak safety. + +De-risks the compiled/pooled SQL slot mechanics in isolation before Ch3 wires the +real ParameterDeclaration type. Validates all 7 propagation/reset sites. Folded into +Ch3 (sqlspec-smgc.3) and reverted afterward. +""" + +from sqlspec.core._pool import get_sql_pool +from sqlspec.core.statement import SQL + +_SENTINEL = ("declared-sentinel",) + + +def test_default_declared_parameters_is_empty_tuple() -> None: + sql = SQL("select 1") + assert sql.declared_parameters == () + + +def test_copy_full_path_preserves_declared_parameters() -> None: + sql = SQL("select :a") + sql._declared_parameters = _SENTINEL + new = sql.copy(statement="select :b") + assert new.declared_parameters == _SENTINEL + + +def test_copy_fast_path_preserves_declared_parameters() -> None: + # parameters-only fast path -> _create_empty_copy + sql = SQL("select :a") + sql._declared_parameters = _SENTINEL + new = sql.copy(parameters={"a": 1}) + assert new.declared_parameters == _SENTINEL + + +def test_init_from_sql_object_preserves_declared_parameters() -> None: + sql = SQL("select :a") + sql._declared_parameters = _SENTINEL + new = SQL(sql) + assert new.declared_parameters == _SENTINEL + + +def test_reset_clears_declared_parameters() -> None: + sql = SQL("select :a") + sql._declared_parameters = _SENTINEL + sql.reset() + assert sql.declared_parameters == () + + +def test_pool_recycle_does_not_leak_declared_parameters() -> None: + """PRIMARY leak vector: a recycled SQL must NOT inherit a prior query's declarations.""" + pool = get_sql_pool() + leaky = SQL("select :a") + leaky._declared_parameters = _SENTINEL + pool.release(leaky) # resetter is SQL.reset -> must clear the slot + + recycled = pool.acquire() + try: + assert recycled._declared_parameters == () + finally: + pool.release(recycled) From d49cb2db6a8dde02f22d7b27ce1776f93aed540b Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Mon, 8 Jun 2026 17:55:19 +0000 Subject: [PATCH 02/12] feat(core): add ParameterDeclaration value type + type registry Introduce core/parameters/_declared.py with the mypyc-safe ParameterDeclaration value object and an extensible type registry (register_param_type / resolve_param_type) that resolves declared type strings to Python types via a fixed allowlist plus user registration, never evaluating the string. Exported through sqlspec.core and the top-level package. Ch1 sqlspec-smgc.1 (gh-491). --- sqlspec/__init__.py | 6 ++ sqlspec/core/__init__.py | 6 ++ sqlspec/core/parameters/__init__.py | 4 + sqlspec/core/parameters/_declared.py | 95 +++++++++++++++++++++ tests/unit/core/parameters/test_declared.py | 87 +++++++++++++++++++ 5 files changed, 198 insertions(+) create mode 100644 sqlspec/core/parameters/_declared.py create mode 100644 tests/unit/core/parameters/test_declared.py diff --git a/sqlspec/__init__.py b/sqlspec/__init__.py index 47bf7fa97..4ac4bca01 100644 --- a/sqlspec/__init__.py +++ b/sqlspec/__init__.py @@ -51,6 +51,7 @@ CacheConfig, CacheStats, ParameterConverter, + ParameterDeclaration, ParameterProcessor, ParameterStyle, ParameterStyleConfig, @@ -61,6 +62,8 @@ Statement, StatementConfig, StatementStack, + register_param_type, + resolve_param_type, ) from sqlspec.core import filters as filters from sqlspec.driver import AsyncDriverAdapterBase, ExecutionResult, SyncDriverAdapterBase @@ -113,6 +116,7 @@ "ObservabilityConfig", "ObservabilityRuntime", "ParameterConverter", + "ParameterDeclaration", "ParameterProcessor", "ParameterStyle", "ParameterStyleConfig", @@ -159,6 +163,8 @@ "format_statement_event", "loader", "migrations", + "register_param_type", + "resolve_param_type", "sql", "typing", "utils", diff --git a/sqlspec/core/__init__.py b/sqlspec/core/__init__.py index 3be7ad76d..53512468e 100644 --- a/sqlspec/core/__init__.py +++ b/sqlspec/core/__init__.py @@ -153,6 +153,7 @@ PARAMETER_REGEX, DriverParameterProfile, ParameterConverter, + ParameterDeclaration, ParameterInfo, ParameterProcessingResult, ParameterProcessor, @@ -170,8 +171,10 @@ looks_like_execute_many, normalize_parameter_key, register_driver_profile, + register_param_type, replace_null_parameters_with_literals, replace_placeholders_with_literals, + resolve_param_type, validate_parameter_alignment, wrap_with_type, ) @@ -279,6 +282,7 @@ "OperationType", "OrderByFilter", "ParameterConverter", + "ParameterDeclaration", "ParameterInfo", "ParameterProcessingResult", "ParameterProcessor", @@ -366,10 +370,12 @@ "parse_column_for_condition", "parse_datetime_rfc3339", "register_driver_profile", + "register_param_type", "replace_null_parameters_with_literals", "replace_placeholders_with_literals", "reset_pipeline_registry", "reset_stats_only", + "resolve_param_type", "safe_modify_with_cte", "split_sql_script", "update_cache_config", diff --git a/sqlspec/core/parameters/__init__.py b/sqlspec/core/parameters/__init__.py index d9e0379b2..a04be7b66 100644 --- a/sqlspec/core/parameters/__init__.py +++ b/sqlspec/core/parameters/__init__.py @@ -8,6 +8,7 @@ validate_parameter_alignment, ) from sqlspec.core.parameters._converter import ParameterConverter +from sqlspec.core.parameters._declared import ParameterDeclaration, register_param_type, resolve_param_type from sqlspec.core.parameters._processor import ParameterProcessor, structural_fingerprint, value_fingerprint from sqlspec.core.parameters._registry import ( DRIVER_PARAMETER_PROFILES, @@ -43,6 +44,7 @@ "PARAMETER_REGEX", "DriverParameterProfile", "ParameterConverter", + "ParameterDeclaration", "ParameterInfo", "ParameterMapping", "ParameterPayload", @@ -63,8 +65,10 @@ "looks_like_execute_many", "normalize_parameter_key", "register_driver_profile", + "register_param_type", "replace_null_parameters_with_literals", "replace_placeholders_with_literals", + "resolve_param_type", "structural_fingerprint", "validate_parameter_alignment", "value_fingerprint", diff --git a/sqlspec/core/parameters/_declared.py b/sqlspec/core/parameters/_declared.py new file mode 100644 index 000000000..e8dd7ebe6 --- /dev/null +++ b/sqlspec/core/parameters/_declared.py @@ -0,0 +1,95 @@ +"""Declared parameter metadata for SQL-file ``-- param:`` annotations. + +Carries the name, declared type string, required flag, and description parsed from +``-- param: [?] [description]`` directives, plus an extensible registry +that resolves declared type strings to Python types for validation. Resolution is a +pure lookup; declared type strings are never evaluated. +""" + +from datetime import date, datetime, time +from decimal import Decimal + +__all__ = ("ParameterDeclaration", "register_param_type", "resolve_param_type") + + +class ParameterDeclaration: + """A single parameter declared in a SQL file header.""" + + __slots__ = ("description", "name", "required", "type_str") + + def __init__( + self, name: str, type_str: str, required: bool = True, description: "str | None" = None + ) -> None: + self.name = name + self.type_str = type_str + self.required = required + self.description = description + + def __eq__(self, other: object) -> bool: + if not isinstance(other, ParameterDeclaration): + return NotImplemented + return ( + self.name == other.name + and self.type_str == other.type_str + and self.required == other.required + and self.description == other.description + ) + + def __hash__(self) -> int: + return hash((self.name, self.type_str, self.required, self.description)) + + def __repr__(self) -> str: + return ( + f"ParameterDeclaration(name={self.name!r}, type_str={self.type_str!r}, " + f"required={self.required!r}, description={self.description!r})" + ) + + +_TYPE_REGISTRY: "dict[str, type]" = { + "str": str, + "int": int, + "float": float, + "bool": bool, + "bytes": bytes, + "date": date, + "datetime": datetime, + "time": time, + "decimal": Decimal, + "list[int]": list, + "list[str]": list, + "list[float]": list, + "list[bool]": list, + "list": list, + "tuple": tuple, +} + + +def _normalize_type_key(type_str: str) -> str: + """Normalize a declared type string to its registry lookup key.""" + return "".join(type_str.split()).lower() + + +def register_param_type(name: str, py_type: type) -> None: + """Register or override a declared-type-string to Python-type mapping. + + Args: + name: The declared type string as written in ``-- param:`` (case-insensitive). + py_type: The Python type used for ``isinstance`` validation. + """ + _TYPE_REGISTRY[_normalize_type_key(name)] = py_type + + +def resolve_param_type(type_str: str) -> "type | None": + """Resolve a declared type string to a Python type, or ``None`` if unknown. + + Unknown type strings are documentation-only and skipped during validation. + The declared string is looked up, never evaluated. Parameterized containers + (``list[int]``) resolve to their origin type (``list``). + + Args: + type_str: The declared type string from a ``-- param:`` directive. + + Returns: + The resolved Python type, or ``None`` when not in the registry. + """ + return _TYPE_REGISTRY.get(_normalize_type_key(type_str)) diff --git a/tests/unit/core/parameters/test_declared.py b/tests/unit/core/parameters/test_declared.py new file mode 100644 index 000000000..a86b48cfb --- /dev/null +++ b/tests/unit/core/parameters/test_declared.py @@ -0,0 +1,87 @@ +"""Tests for declared parameter metadata + type registry (Ch1, sqlspec-smgc.1).""" + +from datetime import date, datetime +from decimal import Decimal + +import pytest + +from sqlspec.core.parameters._declared import ( + ParameterDeclaration, + register_param_type, + resolve_param_type, +) + + +def test_declaration_defaults_to_required() -> None: + decl = ParameterDeclaration(name="status_cd", type_str="str") + assert decl.name == "status_cd" + assert decl.type_str == "str" + assert decl.required is True + assert decl.description is None + + +def test_declaration_optional_with_description() -> None: + decl = ParameterDeclaration("limit", "int", required=False, description="Max rows") + assert decl.required is False + assert decl.description == "Max rows" + + +def test_declaration_equality_and_hash() -> None: + a = ParameterDeclaration("a", "int") + b = ParameterDeclaration("a", "int") + c = ParameterDeclaration("a", "int", required=False) + assert a == b + assert a != c + assert a != "not-a-declaration" + assert hash(a) == hash(b) + + +@pytest.mark.parametrize( + ("type_str", "expected"), + [ + ("str", str), + ("int", int), + ("float", float), + ("bool", bool), + ("bytes", bytes), + ("date", date), + ("datetime", datetime), + ("Decimal", Decimal), + ("list[int]", list), + ("list[str]", list), + ("list", list), + ], +) +def test_resolve_known_types(type_str: str, expected: type) -> None: + assert resolve_param_type(type_str) is expected + + +def test_resolve_is_case_and_whitespace_insensitive() -> None: + assert resolve_param_type(" LIST[ INT ] ") is list + assert resolve_param_type("INT") is int + + +def test_resolve_unknown_returns_none() -> None: + assert resolve_param_type("Money") is None + assert resolve_param_type("frobnicate") is None + + +def test_register_param_type_adds_and_resolves() -> None: + assert resolve_param_type("Money") is None + register_param_type("Money", Decimal) + try: + assert resolve_param_type("Money") is Decimal + assert resolve_param_type("money") is Decimal # case-insensitive + finally: + # keep global registry clean for other tests + from sqlspec.core.parameters._declared import _TYPE_REGISTRY + + _TYPE_REGISTRY.pop("money", None) + + +def test_public_exports() -> None: + from sqlspec import ParameterDeclaration as TopDecl + from sqlspec import register_param_type as top_register + + assert TopDecl is ParameterDeclaration + assert top_register is register_param_type From 7c07a7756ef7d8f6a641cc70ab672afae841a136 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Mon, 8 Jun 2026 18:01:46 +0000 Subject: [PATCH 03/12] feat(loader): parse -- param: directives and expose declarations Scan each named statement's leading comment block for -- param: [?] [description] directives (alongside -- dialect:), storing them on NamedStatement.parameters. Malformed directives warn and skip by default, or raise when strict_parameter_annotations is set. Surface declarations via SQLFileLoader.get_query_parameters / SQLSpec.get_query_parameters and accept them in add_named_sql. Declarations ride the SQLFileCacheEntry, surviving the file-cache roundtrip. Ch2 sqlspec-smgc.2 (gh-491). --- sqlspec/base.py | 26 +++- sqlspec/loader.py | 144 ++++++++++++++++++--- tests/unit/loader/test_param_directives.py | 104 +++++++++++++++ tests/unit/loader/test_sql_file_loader.py | 2 +- 4 files changed, 251 insertions(+), 25 deletions(-) create mode 100644 tests/unit/loader/test_param_directives.py diff --git a/sqlspec/base.py b/sqlspec/base.py index 7e0130961..0e23fae4f 100644 --- a/sqlspec/base.py +++ b/sqlspec/base.py @@ -35,10 +35,11 @@ from sqlspec.utils.type_guards import has_name if TYPE_CHECKING: + from collections.abc import Sequence from pathlib import Path from types import TracebackType - from sqlspec.core import SQL + from sqlspec.core import SQL, ParameterDeclaration from sqlspec.typing import PoolT @@ -512,16 +513,23 @@ def load_sql_files(self, *paths: "str | Path") -> None: loader.load_sql(*paths) logger.debug("Loaded SQL files: %s", paths) - def add_named_sql(self, name: str, sql: str, dialect: "str | None" = None) -> None: + def add_named_sql( + self, + name: str, + sql: str, + dialect: "str | None" = None, + parameters: "Sequence[ParameterDeclaration] | None" = None, + ) -> None: """Add a named SQL query directly. Args: name: Name for the SQL query. sql: Raw SQL content. dialect: Optional dialect for the SQL statement. + parameters: Optional declared parameter metadata for the query. """ loader = self._ensure_loader() - loader.add_named_sql(name, sql, dialect) + loader.add_named_sql(name, sql, dialect, parameters) logger.debug("Added named SQL: %s", name) def get_sql(self, name: str) -> "SQL": @@ -536,6 +544,18 @@ def get_sql(self, name: str) -> "SQL": """ return self._ensure_loader().get_sql(name) + def get_query_parameters(self, name: str) -> "tuple[ParameterDeclaration, ...]": + """Get declared parameter metadata for a query. + + Args: + name: Name of the statement from SQL file comments. + Hyphens in names are converted to underscores. + + Returns: + Tuple of declared parameters; empty if the query declares none. + """ + return self._ensure_loader().get_query_parameters(name) + def list_sql_queries(self) -> "list[str]": """List all available query names. diff --git a/sqlspec/loader.py b/sqlspec/loader.py index d75b3a7ee..abb5f3e47 100644 --- a/sqlspec/loader.py +++ b/sqlspec/loader.py @@ -16,7 +16,7 @@ from typing import TYPE_CHECKING, Any, Final from urllib.parse import unquote, urlparse -from sqlspec.core import SQL, get_cache, get_cache_config +from sqlspec.core import SQL, ParameterDeclaration, get_cache, get_cache_config from sqlspec.exceptions import ( FileNotFoundInStorageError, SQLFileNotFoundError, @@ -31,6 +31,8 @@ from sqlspec.utils.type_guards import is_local_path if TYPE_CHECKING: + from collections.abc import Sequence + from sqlspec.observability import ObservabilityRuntime from sqlspec.storage.registry import StorageRegistry @@ -43,6 +45,13 @@ DIALECT_PATTERN = re.compile(r"^\s*--\s*dialect\s*:\s*(?P[a-zA-Z0-9_]+)\s*$", re.IGNORECASE | re.MULTILINE) +PARAM_PATTERN = re.compile( + r"^\s*--\s*param\s*:\s*(?P\w+)\s+(?P[\w.]+(?:\[[\w., ]+\])?)(?P\?)?(?:\s+(?P.*\S))?\s*$", + re.IGNORECASE, +) + +PARAM_PREFIX_PATTERN = re.compile(r"^\s*--\s*param\s*:", re.IGNORECASE) + DIALECT_ALIASES: Final = { "postgresql": "postgres", @@ -64,13 +73,21 @@ class NamedStatement: and line position for error reporting. """ - __slots__ = ("dialect", "name", "sql", "start_line") + __slots__ = ("dialect", "name", "parameters", "sql", "start_line") - def __init__(self, name: str, sql: str, dialect: "str | None" = None, start_line: int = 0) -> None: + def __init__( + self, + name: str, + sql: str, + dialect: "str | None" = None, + start_line: int = 0, + parameters: "tuple[ParameterDeclaration, ...]" = (), + ) -> None: self.name = name self.sql = sql self.dialect = dialect self.start_line = start_line + self.parameters = parameters class SQLFile: @@ -136,6 +153,7 @@ class SQLFileLoader: "_runtime", "encoding", "storage_registry", + "strict_parameter_annotations", ) def __init__( @@ -144,6 +162,7 @@ def __init__( encoding: str = "utf-8", storage_registry: "StorageRegistry | None" = None, runtime: "ObservabilityRuntime | None" = None, + strict_parameter_annotations: bool = False, ) -> None: """Initialize the SQL file loader. @@ -151,8 +170,11 @@ def __init__( encoding: Text encoding for reading SQL files. storage_registry: Storage registry for handling file URIs. runtime: Observability runtime for instrumentation. + strict_parameter_annotations: When True, a malformed ``-- param:`` directive + raises instead of emitting a warning and skipping the line. """ self.encoding = encoding + self.strict_parameter_annotations = strict_parameter_annotations self.storage_registry = storage_registry or default_storage_registry self._compiled_statements: dict[str, SQL] = {} @@ -309,7 +331,63 @@ def _strip_leading_comments(sql_text: str) -> str: return "\n".join(lines[first_sql_line_index:]).strip() @staticmethod - def _parse_sql_content(content: str, file_path: str) -> "dict[str, NamedStatement]": + def _parse_directive_block( + statement_section: str, file_path: str, strict: bool + ) -> "tuple[str | None, tuple[ParameterDeclaration, ...], str]": + """Scan a statement's leading comment block for ``dialect``/``param`` directives. + + Args: + statement_section: The statement body including any leading directive lines. + file_path: File path for error reporting. + strict: When True, a malformed ``-- param:`` line raises instead of warning. + + Returns: + The resolved dialect, the declared parameters, and the SQL body with the + leading directive/comment lines removed. + + Raises: + SQLFileParseError: If ``strict`` and a ``-- param:`` line is malformed. + """ + dialect: str | None = None + params: list[ParameterDeclaration] = [] + raw_lines = statement_section.split("\n") + body_start = len(raw_lines) + for idx, raw in enumerate(raw_lines): + stripped = raw.strip() + if not stripped: + continue + if not stripped.startswith("--"): + body_start = idx + break + dialect_match = DIALECT_PATTERN.match(stripped) + if dialect_match: + dialect = _normalize_dialect(dialect_match.group("dialect").lower()) + continue + param_match = PARAM_PATTERN.match(stripped) + if param_match: + params.append( + ParameterDeclaration( + name=param_match.group("name"), + type_str=param_match.group("type"), + required=param_match.group("opt") is None, + description=param_match.group("desc"), + ) + ) + continue + if PARAM_PREFIX_PATTERN.match(stripped): + if strict: + raise SQLFileParseError( + file_path, file_path, ValueError(f"Malformed -- param: directive: {stripped}") + ) + log_with_context( + logger, logging.WARNING, "sql.parse.param", file_path=file_path, line=stripped, status="malformed" + ) + return dialect, tuple(params), "\n".join(raw_lines[body_start:]) + + @staticmethod + def _parse_sql_content( + content: str, file_path: str, strict_parameter_annotations: bool = False + ) -> "dict[str, NamedStatement]": """Parse SQL content and extract named statements with dialect specifications. Files without any named statement markers are gracefully skipped by returning @@ -345,19 +423,9 @@ def _parse_sql_content(content: str, file_path: str) -> "dict[str, NamedStatemen if not raw_statement_name or not statement_section: continue - dialect = None - statement_sql = statement_section - - section_lines = [line.strip() for line in statement_section.split("\n") if line.strip()] - if section_lines: - first_line = section_lines[0] - dialect_match = DIALECT_PATTERN.match(first_line) - if dialect_match: - declared_dialect = dialect_match.group("dialect").lower() - - dialect = _normalize_dialect(declared_dialect) - remaining_lines = section_lines[1:] - statement_sql = "\n".join(remaining_lines) + dialect, declared_params, statement_sql = SQLFileLoader._parse_directive_block( + statement_section, file_path, strict_parameter_annotations + ) clean_sql = SQLFileLoader._strip_leading_comments(statement_sql) if clean_sql: @@ -368,7 +436,11 @@ def _parse_sql_content(content: str, file_path: str) -> "dict[str, NamedStatemen ) statements[normalized_name] = NamedStatement( - name=normalized_name, sql=clean_sql, dialect=dialect, start_line=statement_start_line + name=normalized_name, + sql=clean_sql, + dialect=dialect, + start_line=statement_start_line, + parameters=declared_params, ) log_with_context( logger, logging.DEBUG, "sql.parse", file_path=file_path, query_name=normalized_name, dialect=dialect @@ -539,7 +611,7 @@ def _load_file_without_cache( runtime = self._runtime if content is None: content = self._read_file_content(file_path) - statements = self._parse_sql_content(content, path_str) + statements = self._parse_sql_content(content, path_str, self.strict_parameter_annotations) if not statements: log_with_context( @@ -569,13 +641,20 @@ def _load_file_without_cache( runtime.increment_metric("loader.files.loaded") runtime.increment_metric("loader.statements.loaded", len(statements)) - def add_named_sql(self, name: str, sql: str, dialect: "str | None" = None) -> None: + def add_named_sql( + self, + name: str, + sql: str, + dialect: "str | None" = None, + parameters: "Sequence[ParameterDeclaration] | None" = None, + ) -> None: """Add a named SQL query directly without loading from a file. Args: name: Name for the SQL query. sql: Raw SQL content. dialect: Optional dialect for the SQL statement. + parameters: Optional declared parameter metadata for the query. Raises: ValueError: If query name already exists. @@ -591,10 +670,33 @@ def add_named_sql(self, name: str, sql: str, dialect: "str | None" = None) -> No if dialect is not None: dialect = _normalize_dialect(dialect) - statement = NamedStatement(name=normalized_name, sql=sql.strip(), dialect=dialect, start_line=0) + statement = NamedStatement( + name=normalized_name, + sql=sql.strip(), + dialect=dialect, + start_line=0, + parameters=tuple(parameters) if parameters else (), + ) self._queries[normalized_name] = statement self._query_to_file[normalized_name] = "" + def get_query_parameters(self, name: str) -> "tuple[ParameterDeclaration, ...]": + """Get declared parameter metadata for a query. + + Args: + name: Query name (hyphens are converted to underscores). + + Returns: + Tuple of declared parameters; empty if the query declares none. + + Raises: + SQLStatementNotFoundError: If the query does not exist. + """ + safe_name = _normalize_query_name(name) + if safe_name not in self._queries: + self._raise_statement_not_found(name, safe_name) + return self._queries[safe_name].parameters + def get_file(self, path: str | Path) -> "SQLFile | None": """Get a loaded SQLFile object by path. diff --git a/tests/unit/loader/test_param_directives.py b/tests/unit/loader/test_param_directives.py new file mode 100644 index 000000000..de66ff62e --- /dev/null +++ b/tests/unit/loader/test_param_directives.py @@ -0,0 +1,104 @@ +"""Tests for ``-- param:`` directive parsing + introspection (Ch2, sqlspec-smgc.2).""" + +import logging + +import pytest + +from sqlspec.core import ParameterDeclaration +from sqlspec.exceptions import SQLFileParseError, SQLStatementNotFoundError +from sqlspec.loader import PARAM_PATTERN, SQLFileLoader + + +@pytest.mark.parametrize( + ("line", "name", "type_str", "required", "description"), + [ + ("-- param: status_cd str The status code", "status_cd", "str", True, "The status code"), + ("-- param: limit int?", "limit", "int", False, None), + ("-- param: offer_ids list[int] List of ids", "offer_ids", "list[int]", True, "List of ids"), + ("--param:x bool", "x", "bool", True, None), + ("-- PARAM: Y Decimal? money", "Y", "Decimal", False, "money"), + ], +) +def test_param_pattern(line: str, name: str, type_str: str, required: bool, description: "str | None") -> None: + m = PARAM_PATTERN.match(line) + assert m is not None + assert m.group("name") == name + assert m.group("type") == type_str + assert (m.group("opt") is None) is required + assert m.group("desc") == description + + +def test_parse_declared_params_interleaved_with_dialect() -> None: + content = """ +-- name: get_offers +-- dialect: oracle +-- param: status_cd str? The status code +-- param: offer_ids list[int] List of offer IDs +-- param: limit int? Maximum rows +select offer_id from offers where status_cd = :status_cd and offer_id in (:offer_ids) +""" + statements = SQLFileLoader._parse_sql_content(content, "test.sql") + stmt = statements["get_offers"] + assert stmt.dialect == "oracle" + assert stmt.parameters == ( + ParameterDeclaration("status_cd", "str", required=False, description="The status code"), + ParameterDeclaration("offer_ids", "list[int]", required=True, description="List of offer IDs"), + ParameterDeclaration("limit", "int", required=False, description="Maximum rows"), + ) + assert stmt.sql.startswith("select offer_id from offers") + assert "-- param" not in stmt.sql + + +def test_query_without_params_is_unchanged() -> None: + content = "-- name: plain\nselect 1\n" + statements = SQLFileLoader._parse_sql_content(content, "test.sql") + assert statements["plain"].parameters == () + assert statements["plain"].sql == "select 1" + + +def test_malformed_param_warns_and_skips_by_default(caplog: pytest.LogCaptureFixture) -> None: + content = "-- name: q\n-- param: oops\nselect 1\n" + with caplog.at_level(logging.WARNING): + statements = SQLFileLoader._parse_sql_content(content, "test.sql") + assert statements["q"].parameters == () + assert statements["q"].sql == "select 1" + assert any("malformed" in r.message or "param" in r.message.lower() for r in caplog.records) + + +def test_malformed_param_raises_in_strict_mode() -> None: + content = "-- name: q\n-- param: oops\nselect 1\n" + with pytest.raises(SQLFileParseError): + SQLFileLoader._parse_sql_content(content, "test.sql", strict_parameter_annotations=True) + + +def test_add_named_sql_with_parameters() -> None: + loader = SQLFileLoader() + decls = [ParameterDeclaration("a", "int")] + loader.add_named_sql("q", "select :a", parameters=decls) + assert loader.get_query_parameters("q") == (ParameterDeclaration("a", "int"),) + + +def test_get_query_parameters_empty_and_missing() -> None: + loader = SQLFileLoader() + loader.add_named_sql("plain", "select 1") + assert loader.get_query_parameters("plain") == () + with pytest.raises(SQLStatementNotFoundError): + loader.get_query_parameters("nope") + + +def test_declared_params_survive_file_cache_roundtrip(tmp_path: "object") -> None: + """Declarations ride on NamedStatement inside SQLFileCacheEntry, so a second + loader that hits the file cache must see the same declarations.""" + from pathlib import Path + + sql_path = Path(str(tmp_path)) / "q.sql" + sql_path.write_text("-- name: q\n-- param: a int Identifier\nselect :a\n") + + first = SQLFileLoader() + first.load_sql(sql_path) + assert first.get_query_parameters("q") == (ParameterDeclaration("a", "int", description="Identifier"),) + + # A fresh loader reuses the populated file cache (same content hash). + second = SQLFileLoader() + second.load_sql(sql_path) + assert second.get_query_parameters("q") == (ParameterDeclaration("a", "int", description="Identifier"),) diff --git a/tests/unit/loader/test_sql_file_loader.py b/tests/unit/loader/test_sql_file_loader.py index d5c280f96..85db94f6f 100644 --- a/tests/unit/loader/test_sql_file_loader.py +++ b/tests/unit/loader/test_sql_file_loader.py @@ -128,7 +128,7 @@ def test_named_statement_slots() -> None: stmt = NamedStatement("test", "SELECT 1") assert hasattr(stmt.__class__, "__slots__") - assert stmt.__class__.__slots__ == ("dialect", "name", "sql", "start_line") + assert stmt.__class__.__slots__ == ("dialect", "name", "parameters", "sql", "start_line") with pytest.raises(AttributeError): stmt.arbitrary_attr = "value" # pyright: ignore[reportAttributeAccessIssue] From bef8caa67afdda4fe483b6f6e3304be750036e1d Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Mon, 8 Jun 2026 18:17:15 +0000 Subject: [PATCH 04/12] feat(core): expose declared_parameters on SQL and preserve through driver Tighten the SQL._declared_parameters slot to tuple[ParameterDeclaration, ...], add a declared_parameters constructor keyword, and populate it from get_sql(). Thread the slot through every self-derivation (as_script, add_named_parameter, _create_modified_copy_with_expression) and through the driver's _prepare_from_sql rebuild so declarations survive prepare_statement + filter application across all adapters. Fold the spike proof into a permanent carriage test. Ch3 sqlspec-smgc.3 (gh-491). --- sqlspec/core/statement.py | 28 +++++- sqlspec/driver/_common.py | 31 +++++- sqlspec/loader.py | 2 +- .../core/test_declared_params_carriage.py | 94 +++++++++++++++++++ tests/unit/core/test_declared_params_spike.py | 59 ------------ 5 files changed, 145 insertions(+), 69 deletions(-) create mode 100644 tests/unit/core/test_declared_params_carriage.py delete mode 100644 tests/unit/core/test_declared_params_spike.py diff --git a/sqlspec/core/statement.py b/sqlspec/core/statement.py index 1de465853..ea7a3abce 100644 --- a/sqlspec/core/statement.py +++ b/sqlspec/core/statement.py @@ -19,6 +19,7 @@ from sqlspec.core.hashing import hash_filters from sqlspec.core.parameters import ( ParameterConverter, + ParameterDeclaration, ParameterProcessor, ParameterProfile, ParameterStyle, @@ -310,6 +311,7 @@ def __init__( *parameters: "Any | StatementFilter | list[Any | StatementFilter]", statement_config: "StatementConfig | None" = None, is_many: bool | None = None, + declared_parameters: "tuple[ParameterDeclaration, ...]" = (), **kwargs: Any, ) -> None: """Initialize SQL statement. @@ -336,7 +338,7 @@ def __init__( self._is_script = False self._raw_expression: exp.Expr | None = None self._rebind_processor: ParameterProcessor | None = None - self._declared_parameters: "tuple[Any, ...]" = () + self._declared_parameters: "tuple[ParameterDeclaration, ...]" = declared_parameters if isinstance(statement, SQL): self._init_from_sql_object(statement) @@ -616,7 +618,7 @@ def original_parameters(self) -> Any: return self._original_parameters @property - def declared_parameters(self) -> "tuple[Any, ...]": + def declared_parameters(self) -> "tuple[ParameterDeclaration, ...]": """Get declared parameter metadata carried with this statement (public API).""" return self._declared_parameters @@ -884,7 +886,13 @@ def as_script(self) -> "SQL": config = self._statement_config is_many = self._is_many statement_seed = self._raw_expression or self._raw_sql - new_sql = SQL(statement_seed, *original_params, statement_config=config, is_many=is_many) + new_sql = SQL( + statement_seed, + *original_params, + statement_config=config, + is_many=is_many, + declared_parameters=self._declared_parameters, + ) new_sql._named_parameters.update(self._named_parameters) new_sql._positional_parameters = self._positional_parameters.copy() new_sql._filters = self._filters.copy() @@ -1083,7 +1091,11 @@ def _create_modified_copy_with_expression(self, new_expr: "exp.Expr") -> "SQL": New SQL instance with the expression and copied state """ new_sql = SQL( - new_expr, *self._original_parameters, statement_config=self._statement_config, is_many=self._is_many + new_expr, + *self._original_parameters, + statement_config=self._statement_config, + is_many=self._is_many, + declared_parameters=self._declared_parameters, ) new_sql._named_parameters.update(self._named_parameters) new_sql._positional_parameters = self._positional_parameters.copy() @@ -1105,7 +1117,13 @@ def add_named_parameter(self, name: str, value: Any) -> "SQL": config = self._statement_config is_many = self._is_many statement_seed = self._raw_expression or self._raw_sql - new_sql = SQL(statement_seed, *original_params, statement_config=config, is_many=is_many) + new_sql = SQL( + statement_seed, + *original_params, + statement_config=config, + is_many=is_many, + declared_parameters=self._declared_parameters, + ) new_sql._named_parameters.update(self._named_parameters) new_sql._named_parameters[name] = value new_sql._positional_parameters = self._positional_parameters.copy() diff --git a/sqlspec/driver/_common.py b/sqlspec/driver/_common.py index 406940887..6be92d54f 100644 --- a/sqlspec/driver/_common.py +++ b/sqlspec/driver/_common.py @@ -1456,6 +1456,7 @@ def _prepare_from_sql( statement_config: "StatementConfig", kwargs: "dict[str, Any]", ) -> "SQL": + declared = sql_statement.declared_parameters if data_parameters or kwargs: merged_parameters = ( (*sql_statement.positional_parameters, *tuple(data_parameters)) @@ -1463,7 +1464,13 @@ def _prepare_from_sql( else sql_statement.positional_parameters ) statement_seed = sql_statement.raw_expression or sql_statement.raw_sql - return SQL(statement_seed, *merged_parameters, statement_config=statement_config, **kwargs) + return SQL( + statement_seed, + *merged_parameters, + statement_config=statement_config, + declared_parameters=declared, + **kwargs, + ) needs_rebuild = False if statement_config.dialect and ( @@ -1481,10 +1488,26 @@ def _prepare_from_sql( if needs_rebuild: statement_seed = sql_statement.raw_expression or sql_statement.raw_sql if sql_statement.is_many and sql_statement.parameters: - return SQL(statement_seed, sql_statement.parameters, statement_config=statement_config, is_many=True) + return SQL( + statement_seed, + sql_statement.parameters, + statement_config=statement_config, + is_many=True, + declared_parameters=declared, + ) if sql_statement.named_parameters: - return SQL(statement_seed, statement_config=statement_config, **sql_statement.named_parameters) - return SQL(statement_seed, *sql_statement.positional_parameters, statement_config=statement_config) + return SQL( + statement_seed, + statement_config=statement_config, + declared_parameters=declared, + **sql_statement.named_parameters, + ) + return SQL( + statement_seed, + *sql_statement.positional_parameters, + statement_config=statement_config, + declared_parameters=declared, + ) return sql_statement def _prepare_from_string( diff --git a/sqlspec/loader.py b/sqlspec/loader.py index abb5f3e47..df17f7f69 100644 --- a/sqlspec/loader.py +++ b/sqlspec/loader.py @@ -806,7 +806,7 @@ def get_sql(self, name: str) -> "SQL": if parsed_statement.dialect: sqlglot_dialect = _normalize_dialect(parsed_statement.dialect) - sql = SQL(parsed_statement.sql, dialect=sqlglot_dialect) + sql = SQL(parsed_statement.sql, dialect=sqlglot_dialect, declared_parameters=parsed_statement.parameters) try: sql.compile() except Exception as exc: diff --git a/tests/unit/core/test_declared_params_carriage.py b/tests/unit/core/test_declared_params_carriage.py new file mode 100644 index 000000000..0db1f7393 --- /dev/null +++ b/tests/unit/core/test_declared_params_carriage.py @@ -0,0 +1,94 @@ +"""declared_parameters slot carriage + pool-leak safety (Ch3, sqlspec-smgc.3). + +Validates all seven propagation/reset sites for the SQL._declared_parameters slot, +plus loader population and driver-derivation preservation. +""" + +from sqlspec.core import ParameterDeclaration +from sqlspec.core._pool import get_sql_pool +from sqlspec.core.statement import SQL + +_SENTINEL = (ParameterDeclaration("a", "int"),) + + +def test_default_declared_parameters_is_empty_tuple() -> None: + sql = SQL("select 1") + assert sql.declared_parameters == () + + +def test_copy_full_path_preserves_declared_parameters() -> None: + sql = SQL("select :a") + sql._declared_parameters = _SENTINEL + new = sql.copy(statement="select :b") + assert new.declared_parameters == _SENTINEL + + +def test_copy_fast_path_preserves_declared_parameters() -> None: + # parameters-only fast path -> _create_empty_copy + sql = SQL("select :a") + sql._declared_parameters = _SENTINEL + new = sql.copy(parameters={"a": 1}) + assert new.declared_parameters == _SENTINEL + + +def test_init_from_sql_object_preserves_declared_parameters() -> None: + sql = SQL("select :a") + sql._declared_parameters = _SENTINEL + new = SQL(sql) + assert new.declared_parameters == _SENTINEL + + +def test_reset_clears_declared_parameters() -> None: + sql = SQL("select :a") + sql._declared_parameters = _SENTINEL + sql.reset() + assert sql.declared_parameters == () + + +def test_pool_recycle_does_not_leak_declared_parameters() -> None: + """PRIMARY leak vector: a recycled SQL must NOT inherit a prior query's declarations.""" + pool = get_sql_pool() + leaky = SQL("select :a") + leaky._declared_parameters = _SENTINEL + pool.release(leaky) # resetter is SQL.reset -> must clear the slot + + recycled = pool.acquire() + try: + assert recycled._declared_parameters == () + finally: + pool.release(recycled) + + +def test_get_sql_populates_declared_parameters() -> None: + from sqlspec.loader import SQLFileLoader + + loader = SQLFileLoader() + loader.add_named_sql("q", "select :a", parameters=[ParameterDeclaration("a", "int", description="id")]) + sql = loader.get_sql("q") + assert sql.declared_parameters == (ParameterDeclaration("a", "int", description="id"),) + + +def test_undeclared_get_sql_has_empty_declarations() -> None: + from sqlspec.loader import SQLFileLoader + + loader = SQLFileLoader() + loader.add_named_sql("plain", "select 1") + assert loader.get_sql("plain").declared_parameters == () + + +def test_constructor_kwarg_sets_declarations() -> None: + sql = SQL("select :a", {"a": 1}, declared_parameters=_SENTINEL) + assert sql.declared_parameters == _SENTINEL + + +def test_declarations_survive_driver_prepare_with_filter() -> None: + """Declarations must survive prepare_statement rebuild + filter application.""" + from sqlspec.adapters.sqlite import SqliteConfig + from sqlspec.core.filters import LimitOffsetFilter + + config = SqliteConfig(pool_config={"database": ":memory:"}) + with config.provide_session() as session: + base = SQL("select :a") + base._declared_parameters = _SENTINEL + prepared = session.prepare_statement(base, ({"a": 1}, LimitOffsetFilter(limit=5, offset=0))) + assert prepared.declared_parameters == _SENTINEL diff --git a/tests/unit/core/test_declared_params_spike.py b/tests/unit/core/test_declared_params_spike.py deleted file mode 100644 index fb35c7008..000000000 --- a/tests/unit/core/test_declared_params_spike.py +++ /dev/null @@ -1,59 +0,0 @@ -"""SPIKE (throwaway): prove declared_parameters slot carriage + pool-leak safety. - -De-risks the compiled/pooled SQL slot mechanics in isolation before Ch3 wires the -real ParameterDeclaration type. Validates all 7 propagation/reset sites. Folded into -Ch3 (sqlspec-smgc.3) and reverted afterward. -""" - -from sqlspec.core._pool import get_sql_pool -from sqlspec.core.statement import SQL - -_SENTINEL = ("declared-sentinel",) - - -def test_default_declared_parameters_is_empty_tuple() -> None: - sql = SQL("select 1") - assert sql.declared_parameters == () - - -def test_copy_full_path_preserves_declared_parameters() -> None: - sql = SQL("select :a") - sql._declared_parameters = _SENTINEL - new = sql.copy(statement="select :b") - assert new.declared_parameters == _SENTINEL - - -def test_copy_fast_path_preserves_declared_parameters() -> None: - # parameters-only fast path -> _create_empty_copy - sql = SQL("select :a") - sql._declared_parameters = _SENTINEL - new = sql.copy(parameters={"a": 1}) - assert new.declared_parameters == _SENTINEL - - -def test_init_from_sql_object_preserves_declared_parameters() -> None: - sql = SQL("select :a") - sql._declared_parameters = _SENTINEL - new = SQL(sql) - assert new.declared_parameters == _SENTINEL - - -def test_reset_clears_declared_parameters() -> None: - sql = SQL("select :a") - sql._declared_parameters = _SENTINEL - sql.reset() - assert sql.declared_parameters == () - - -def test_pool_recycle_does_not_leak_declared_parameters() -> None: - """PRIMARY leak vector: a recycled SQL must NOT inherit a prior query's declarations.""" - pool = get_sql_pool() - leaky = SQL("select :a") - leaky._declared_parameters = _SENTINEL - pool.release(leaky) # resetter is SQL.reset -> must clear the slot - - recycled = pool.acquire() - try: - assert recycled._declared_parameters == () - finally: - pool.release(recycled) From e61e7be7c327a91b78724c40db2d0dba8d63a0fb Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Mon, 8 Jun 2026 18:47:48 +0000 Subject: [PATCH 05/12] feat(loader): load-time validation of declared parameters When a query declares params, cross-check them against the SQL's actual placeholders at load time: declared names must be a subset of the named placeholders (drift), or for positional binding the declared count must equal the placeholder count. Raises SQLFileParseError early; declaration-driven, so queries without -- param: directives are unaffected. Ch4 sqlspec-smgc.4 (gh-491). --- sqlspec/loader.py | 60 +++++++++++++++++-- tests/unit/loader/test_param_directives.py | 1 + .../unit/loader/test_param_load_validation.py | 51 ++++++++++++++++ 3 files changed, 106 insertions(+), 6 deletions(-) create mode 100644 tests/unit/loader/test_param_load_validation.py diff --git a/sqlspec/loader.py b/sqlspec/loader.py index df17f7f69..282dd50b5 100644 --- a/sqlspec/loader.py +++ b/sqlspec/loader.py @@ -16,7 +16,7 @@ from typing import TYPE_CHECKING, Any, Final from urllib.parse import unquote, urlparse -from sqlspec.core import SQL, ParameterDeclaration, get_cache, get_cache_config +from sqlspec.core import SQL, ParameterDeclaration, ParameterValidator, get_cache, get_cache_config from sqlspec.exceptions import ( FileNotFoundInStorageError, SQLFileNotFoundError, @@ -384,6 +384,50 @@ def _parse_directive_block( ) return dialect, tuple(params), "\n".join(raw_lines[body_start:]) + @staticmethod + def _validate_declared_parameters( + clean_sql: str, declared: "tuple[ParameterDeclaration, ...]", statement_name: str, file_path: str + ) -> None: + """Validate declared parameters against the query's actual placeholders. + + For named binding, every declared name must appear among the SQL placeholders + (declared names may be a subset; filters and undeclared params are allowed). For + positional binding, the declared count must equal the placeholder count. + + Args: + clean_sql: The SQL body with directives/comments stripped. + declared: Declared parameters for the query. + statement_name: Raw query name for error messages. + file_path: File path for error reporting. + + Raises: + SQLFileParseError: On name drift (named) or count mismatch (positional). + """ + if not declared: + return + infos = ParameterValidator().extract_parameters(clean_sql) + named = {info.name for info in infos if info.name and not info.name.isdigit()} + if named: + for decl in declared: + if decl.name not in named: + raise SQLFileParseError( + file_path, + file_path, + ValueError( + f"Declared parameter '{decl.name}' for query '{statement_name}' is not present in the " + f"SQL placeholders {sorted(named)}" + ), + ) + elif len(declared) != len(infos): + raise SQLFileParseError( + file_path, + file_path, + ValueError( + f"Query '{statement_name}' declares {len(declared)} parameter(s) but the SQL has " + f"{len(infos)} positional placeholder(s)" + ), + ) + @staticmethod def _parse_sql_content( content: str, file_path: str, strict_parameter_annotations: bool = False @@ -435,6 +479,10 @@ def _parse_sql_content( file_path, file_path, ValueError(f"Duplicate statement name: {raw_statement_name}") ) + SQLFileLoader._validate_declared_parameters( + clean_sql, declared_params, raw_statement_name, file_path + ) + statements[normalized_name] = NamedStatement( name=normalized_name, sql=clean_sql, @@ -670,12 +718,12 @@ def add_named_sql( if dialect is not None: dialect = _normalize_dialect(dialect) + declared = tuple(parameters) if parameters else () + clean_sql = sql.strip() + self._validate_declared_parameters(clean_sql, declared, name, "") + statement = NamedStatement( - name=normalized_name, - sql=sql.strip(), - dialect=dialect, - start_line=0, - parameters=tuple(parameters) if parameters else (), + name=normalized_name, sql=clean_sql, dialect=dialect, start_line=0, parameters=declared ) self._queries[normalized_name] = statement self._query_to_file[normalized_name] = "" diff --git a/tests/unit/loader/test_param_directives.py b/tests/unit/loader/test_param_directives.py index de66ff62e..e8c5d98de 100644 --- a/tests/unit/loader/test_param_directives.py +++ b/tests/unit/loader/test_param_directives.py @@ -36,6 +36,7 @@ def test_parse_declared_params_interleaved_with_dialect() -> None: -- param: offer_ids list[int] List of offer IDs -- param: limit int? Maximum rows select offer_id from offers where status_cd = :status_cd and offer_id in (:offer_ids) +fetch first :limit rows only """ statements = SQLFileLoader._parse_sql_content(content, "test.sql") stmt = statements["get_offers"] diff --git a/tests/unit/loader/test_param_load_validation.py b/tests/unit/loader/test_param_load_validation.py new file mode 100644 index 000000000..de050e15e --- /dev/null +++ b/tests/unit/loader/test_param_load_validation.py @@ -0,0 +1,51 @@ +"""Load-time validation: declared-name drift + positional count (Ch4, sqlspec-smgc.4).""" + +import pytest + +from sqlspec.core import ParameterDeclaration +from sqlspec.exceptions import SQLFileParseError +from sqlspec.loader import SQLFileLoader + + +def test_named_drift_raises() -> None: + content = "-- name: q\n-- param: status_cd str The code\nselect 1 from t where status = :status\n" + with pytest.raises(SQLFileParseError, match="status_cd"): + SQLFileLoader._parse_sql_content(content, "test.sql") + + +def test_named_all_present_loads() -> None: + content = "-- name: q\n-- param: a int\n-- param: b int\nselect :a, :b\n" + statements = SQLFileLoader._parse_sql_content(content, "test.sql") + assert len(statements["q"].parameters) == 2 + + +def test_undeclared_placeholder_is_allowed() -> None: + # declared subset of placeholders -> OK (filters/undeclared params are legal) + content = "-- name: q\n-- param: a int\nselect :a, :b\n" + statements = SQLFileLoader._parse_sql_content(content, "test.sql") + assert statements["q"].parameters == (ParameterDeclaration("a", "int"),) + + +def test_positional_count_mismatch_raises() -> None: + content = "-- name: q\n-- param: a int\n-- param: b int\n-- param: c int\nselect ?, ?\n" + with pytest.raises(SQLFileParseError, match="positional"): + SQLFileLoader._parse_sql_content(content, "test.sql") + + +def test_positional_count_match_loads() -> None: + content = "-- name: q\n-- param: a int\n-- param: b int\nselect ?, ?\n" + statements = SQLFileLoader._parse_sql_content(content, "test.sql") + assert len(statements["q"].parameters) == 2 + + +def test_no_declarations_skips_validation() -> None: + # mismatched counts but no declarations -> no validation, loads fine + content = "-- name: q\nselect ?, ?, ?\n" + statements = SQLFileLoader._parse_sql_content(content, "test.sql") + assert statements["q"].parameters == () + + +def test_add_named_sql_validates_drift() -> None: + loader = SQLFileLoader() + with pytest.raises(SQLFileParseError, match="nope"): + loader.add_named_sql("q", "select :a", parameters=[ParameterDeclaration("nope", "int")]) From 6eec00d04f0c30278e92fa69622b5d4ae6f23f3a Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Mon, 8 Jun 2026 18:47:56 +0000 Subject: [PATCH 06/12] refactor(core): rename generated parameter prefix to param_ Auto-generated parameter names (builder where/in/between helpers) now use the param_ prefix instead of parameter_, aligning with the param_{ordinal} fallback already used in the converter/processor and with the new -- param: file annotations. gh-491. --- sqlspec/core/statement.py | 6 ++-- tests/unit/core/test_sql_modifiers.py | 48 +++++++++++++-------------- 2 files changed, 27 insertions(+), 27 deletions(-) diff --git a/sqlspec/core/statement.py b/sqlspec/core/statement.py index ea7a3abce..85a92d157 100644 --- a/sqlspec/core/statement.py +++ b/sqlspec/core/statement.py @@ -1015,9 +1015,9 @@ def _handle_compile_failure(self, error: Exception) -> ProcessedState: # ========================================================================== def _generate_sql_param_name(self, base_name: str) -> str: - """Generate unique parameter name with parameter_ prefix. + """Generate unique parameter name with param_ prefix. - Uses parameter_ prefix to avoid collision with user-provided parameters. + Uses param_ prefix to avoid collision with user-provided parameters. Auto-generated parameters are namespaced to prevent conflicts. Args: @@ -1026,7 +1026,7 @@ def _generate_sql_param_name(self, base_name: str) -> str: Returns: A unique parameter name that doesn't exist in current parameters """ - prefixed_base = f"parameter_{base_name}" + prefixed_base = f"param_{base_name}" current_index = self._sql_param_counters.get(prefixed_base, 0) if prefixed_base not in self._named_parameters: diff --git a/tests/unit/core/test_sql_modifiers.py b/tests/unit/core/test_sql_modifiers.py index 5531c01a1..cc8c1cd6a 100644 --- a/tests/unit/core/test_sql_modifiers.py +++ b/tests/unit/core/test_sql_modifiers.py @@ -28,8 +28,8 @@ def test_sql_where_eq_where_eq_creates_equality_condition() -> None: modified = stmt.where_eq("status", "active") assert "WHERE" in modified.raw_sql assert "status" in modified.raw_sql - assert "parameter_status" in modified.named_parameters - assert modified.named_parameters["parameter_status"] == "active" + assert "param_status" in modified.named_parameters + assert modified.named_parameters["param_status"] == "active" def test_sql_where_eq_where_eq_preserves_original() -> None: @@ -46,8 +46,8 @@ def test_sql_where_eq_where_eq_chains_with_and() -> None: stmt = SQL("SELECT * FROM users") modified = stmt.where_eq("status", "active").where_eq("role", "admin") assert "AND" in modified.raw_sql - assert "parameter_status" in modified.named_parameters - assert "parameter_role" in modified.named_parameters + assert "param_status" in modified.named_parameters + assert "param_role" in modified.named_parameters def test_sql_where_neq_where_neq_creates_not_equal_condition() -> None: @@ -56,8 +56,8 @@ def test_sql_where_neq_where_neq_creates_not_equal_condition() -> None: modified = stmt.where_neq("status", "deleted") assert "WHERE" in modified.raw_sql assert "<>" in modified.raw_sql or "!=" in modified.raw_sql - assert "parameter_status" in modified.named_parameters - assert modified.named_parameters["parameter_status"] == "deleted" + assert "param_status" in modified.named_parameters + assert modified.named_parameters["param_status"] == "deleted" def test_sql_where_comparisons_where_lt() -> None: @@ -66,7 +66,7 @@ def test_sql_where_comparisons_where_lt() -> None: modified = stmt.where_lt("price", 100) assert "WHERE" in modified.raw_sql assert "<" in modified.raw_sql - assert modified.named_parameters["parameter_price"] == 100 + assert modified.named_parameters["param_price"] == 100 def test_sql_where_comparisons_where_lte() -> None: @@ -75,7 +75,7 @@ def test_sql_where_comparisons_where_lte() -> None: modified = stmt.where_lte("price", 100) assert "WHERE" in modified.raw_sql assert "<=" in modified.raw_sql - assert modified.named_parameters["parameter_price"] == 100 + assert modified.named_parameters["param_price"] == 100 def test_sql_where_comparisons_where_gt() -> None: @@ -84,7 +84,7 @@ def test_sql_where_comparisons_where_gt() -> None: modified = stmt.where_gt("price", 50) assert "WHERE" in modified.raw_sql assert ">" in modified.raw_sql - assert modified.named_parameters["parameter_price"] == 50 + assert modified.named_parameters["param_price"] == 50 def test_sql_where_comparisons_where_gte() -> None: @@ -93,7 +93,7 @@ def test_sql_where_comparisons_where_gte() -> None: modified = stmt.where_gte("price", 50) assert "WHERE" in modified.raw_sql assert ">=" in modified.raw_sql - assert modified.named_parameters["parameter_price"] == 50 + assert modified.named_parameters["param_price"] == 50 def test_sql_where_like_where_like() -> None: @@ -102,7 +102,7 @@ def test_sql_where_like_where_like() -> None: modified = stmt.where_like("name", "%john%") assert "WHERE" in modified.raw_sql assert "LIKE" in modified.raw_sql - assert modified.named_parameters["parameter_name"] == "%john%" + assert modified.named_parameters["param_name"] == "%john%" def test_sql_where_like_where_ilike() -> None: @@ -111,7 +111,7 @@ def test_sql_where_like_where_ilike() -> None: modified = stmt.where_ilike("name", "%john%") assert "WHERE" in modified.raw_sql assert "ILIKE" in modified.raw_sql - assert modified.named_parameters["parameter_name"] == "%john%" + assert modified.named_parameters["param_name"] == "%john%" def test_sql_where_null_where_is_null() -> None: @@ -179,10 +179,10 @@ def test_sql_where_between_where_between_creates_between_condition() -> None: assert "WHERE" in modified.raw_sql assert "BETWEEN" in modified.raw_sql assert "AND" in modified.raw_sql - assert "parameter_total_low" in modified.named_parameters - assert "parameter_total_high" in modified.named_parameters - assert modified.named_parameters["parameter_total_low"] == 100 - assert modified.named_parameters["parameter_total_high"] == 500 + assert "param_total_low" in modified.named_parameters + assert "param_total_high" in modified.named_parameters + assert modified.named_parameters["param_total_low"] == 100 + assert modified.named_parameters["param_total_high"] == 500 def test_sql_limit_limit_adds_limit_clause() -> None: @@ -284,23 +284,23 @@ def test_parameter_generation_params_dont_collide_with_user_params() -> None: modified = stmt.where_eq("status", "active") assert "status" in modified.named_parameters assert modified.named_parameters["status"] == 1 - assert "parameter_status" in modified.named_parameters + assert "param_status" in modified.named_parameters def test_parameter_generation_avoids_generated_prefix_collision() -> None: """Test generated params append suffixes when the namespace already exists.""" - stmt = SQL("SELECT * FROM users WHERE id = :parameter_status", {"parameter_status": 1}) + stmt = SQL("SELECT * FROM users WHERE id = :param_status", {"param_status": 1}) modified = stmt.where_eq("status", "active") - assert modified.named_parameters["parameter_status"] == 1 - assert modified.named_parameters["parameter_status_1"] == "active" + assert modified.named_parameters["param_status"] == 1 + assert modified.named_parameters["param_status_1"] == "active" def test_sql_where_in_uses_oracle_safe_generated_names() -> None: """Test where_in creates letter-leading generated parameters.""" stmt = SQL("SELECT * FROM users") modified = stmt.where_in("status", ["active", "pending"]) - assert "parameter_status_in_0" in modified.named_parameters - assert "parameter_status_in_1" in modified.named_parameters + assert "param_status_in_0" in modified.named_parameters + assert "param_status_in_1" in modified.named_parameters assert "_sqlspec_status_in_0" not in modified.named_parameters @@ -309,9 +309,9 @@ def test_generated_parameter_names_are_safe_for_named_colon_placeholders() -> No stmt = SQL("SELECT * FROM users", statement_config=_named_colon_config()) modified = stmt.where_eq("status", "active") compiled_sql, parameters = modified.compile() - assert ":parameter_status" in compiled_sql + assert ":param_status" in compiled_sql assert ":_sqlspec" not in compiled_sql - assert parameters == {"parameter_status": "active"} + assert parameters == {"param_status": "active"} def test_cte_preservation_where_eq_preserves_cte() -> None: From 90ec68fbbf03b15e7f036112fd29b18712492e3b Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Mon, 8 Jun 2026 19:18:16 +0000 Subject: [PATCH 07/12] refactor(core): drop optional parameter declarations Declared parameters are now strictly binary: a declared param is always validated (presence + type); undeclared params are untouched. Removes the ?-suffix grammar and ParameterDeclaration.required, which had an empty domain for loaded SQL files (every declared placeholder is static and always bound, so required=False could never legitimately fire). --- sqlspec/core/parameters/_declared.py | 25 ++++++--------------- sqlspec/loader.py | 3 +-- tests/unit/core/parameters/test_declared.py | 10 ++++----- tests/unit/loader/test_param_directives.py | 25 ++++++++++----------- 4 files changed, 24 insertions(+), 39 deletions(-) diff --git a/sqlspec/core/parameters/_declared.py b/sqlspec/core/parameters/_declared.py index e8dd7ebe6..2804a4bc9 100644 --- a/sqlspec/core/parameters/_declared.py +++ b/sqlspec/core/parameters/_declared.py @@ -1,7 +1,7 @@ """Declared parameter metadata for SQL-file ``-- param:`` annotations. -Carries the name, declared type string, required flag, and description parsed from -``-- param: [?] [description]`` directives, plus an extensible registry +Carries the name, declared type string, and description parsed from +``-- param: [description]`` directives, plus an extensible registry that resolves declared type strings to Python types for validation. Resolution is a pure lookup; declared type strings are never evaluated. """ @@ -15,34 +15,23 @@ class ParameterDeclaration: """A single parameter declared in a SQL file header.""" - __slots__ = ("description", "name", "required", "type_str") + __slots__ = ("description", "name", "type_str") - def __init__( - self, name: str, type_str: str, required: bool = True, description: "str | None" = None - ) -> None: + def __init__(self, name: str, type_str: str, description: "str | None" = None) -> None: self.name = name self.type_str = type_str - self.required = required self.description = description def __eq__(self, other: object) -> bool: if not isinstance(other, ParameterDeclaration): return NotImplemented - return ( - self.name == other.name - and self.type_str == other.type_str - and self.required == other.required - and self.description == other.description - ) + return self.name == other.name and self.type_str == other.type_str and self.description == other.description def __hash__(self) -> int: - return hash((self.name, self.type_str, self.required, self.description)) + return hash((self.name, self.type_str, self.description)) def __repr__(self) -> str: - return ( - f"ParameterDeclaration(name={self.name!r}, type_str={self.type_str!r}, " - f"required={self.required!r}, description={self.description!r})" - ) + return f"ParameterDeclaration(name={self.name!r}, type_str={self.type_str!r}, description={self.description!r})" _TYPE_REGISTRY: "dict[str, type]" = { diff --git a/sqlspec/loader.py b/sqlspec/loader.py index 282dd50b5..2fcd9325c 100644 --- a/sqlspec/loader.py +++ b/sqlspec/loader.py @@ -46,7 +46,7 @@ DIALECT_PATTERN = re.compile(r"^\s*--\s*dialect\s*:\s*(?P[a-zA-Z0-9_]+)\s*$", re.IGNORECASE | re.MULTILINE) PARAM_PATTERN = re.compile( - r"^\s*--\s*param\s*:\s*(?P\w+)\s+(?P[\w.]+(?:\[[\w., ]+\])?)(?P\?)?(?:\s+(?P.*\S))?\s*$", + r"^\s*--\s*param\s*:\s*(?P\w+)\s+(?P[\w.]+(?:\[[\w., ]+\])?)(?:\s+(?P.*\S))?\s*$", re.IGNORECASE, ) @@ -369,7 +369,6 @@ def _parse_directive_block( ParameterDeclaration( name=param_match.group("name"), type_str=param_match.group("type"), - required=param_match.group("opt") is None, description=param_match.group("desc"), ) ) diff --git a/tests/unit/core/parameters/test_declared.py b/tests/unit/core/parameters/test_declared.py index a86b48cfb..fc9a5e5c5 100644 --- a/tests/unit/core/parameters/test_declared.py +++ b/tests/unit/core/parameters/test_declared.py @@ -12,24 +12,22 @@ ) -def test_declaration_defaults_to_required() -> None: +def test_declaration_fields() -> None: decl = ParameterDeclaration(name="status_cd", type_str="str") assert decl.name == "status_cd" assert decl.type_str == "str" - assert decl.required is True assert decl.description is None -def test_declaration_optional_with_description() -> None: - decl = ParameterDeclaration("limit", "int", required=False, description="Max rows") - assert decl.required is False +def test_declaration_with_description() -> None: + decl = ParameterDeclaration("limit", "int", description="Max rows") assert decl.description == "Max rows" def test_declaration_equality_and_hash() -> None: a = ParameterDeclaration("a", "int") b = ParameterDeclaration("a", "int") - c = ParameterDeclaration("a", "int", required=False) + c = ParameterDeclaration("a", "int", description="differs") assert a == b assert a != c assert a != "not-a-declaration" diff --git a/tests/unit/loader/test_param_directives.py b/tests/unit/loader/test_param_directives.py index e8c5d98de..1c9da334a 100644 --- a/tests/unit/loader/test_param_directives.py +++ b/tests/unit/loader/test_param_directives.py @@ -10,21 +10,20 @@ @pytest.mark.parametrize( - ("line", "name", "type_str", "required", "description"), + ("line", "name", "type_str", "description"), [ - ("-- param: status_cd str The status code", "status_cd", "str", True, "The status code"), - ("-- param: limit int?", "limit", "int", False, None), - ("-- param: offer_ids list[int] List of ids", "offer_ids", "list[int]", True, "List of ids"), - ("--param:x bool", "x", "bool", True, None), - ("-- PARAM: Y Decimal? money", "Y", "Decimal", False, "money"), + ("-- param: status_cd str The status code", "status_cd", "str", "The status code"), + ("-- param: limit int", "limit", "int", None), + ("-- param: offer_ids list[int] List of ids", "offer_ids", "list[int]", "List of ids"), + ("--param:x bool", "x", "bool", None), + ("-- PARAM: Y Decimal money", "Y", "Decimal", "money"), ], ) -def test_param_pattern(line: str, name: str, type_str: str, required: bool, description: "str | None") -> None: +def test_param_pattern(line: str, name: str, type_str: str, description: "str | None") -> None: m = PARAM_PATTERN.match(line) assert m is not None assert m.group("name") == name assert m.group("type") == type_str - assert (m.group("opt") is None) is required assert m.group("desc") == description @@ -32,9 +31,9 @@ def test_parse_declared_params_interleaved_with_dialect() -> None: content = """ -- name: get_offers -- dialect: oracle --- param: status_cd str? The status code +-- param: status_cd str The status code -- param: offer_ids list[int] List of offer IDs --- param: limit int? Maximum rows +-- param: limit int Maximum rows select offer_id from offers where status_cd = :status_cd and offer_id in (:offer_ids) fetch first :limit rows only """ @@ -42,9 +41,9 @@ def test_parse_declared_params_interleaved_with_dialect() -> None: stmt = statements["get_offers"] assert stmt.dialect == "oracle" assert stmt.parameters == ( - ParameterDeclaration("status_cd", "str", required=False, description="The status code"), - ParameterDeclaration("offer_ids", "list[int]", required=True, description="List of offer IDs"), - ParameterDeclaration("limit", "int", required=False, description="Maximum rows"), + ParameterDeclaration("status_cd", "str", description="The status code"), + ParameterDeclaration("offer_ids", "list[int]", description="List of offer IDs"), + ParameterDeclaration("limit", "int", description="Maximum rows"), ) assert stmt.sql.startswith("select offer_id from offers") assert "-- param" not in stmt.sql From 4d81c09b42994b36ca444c23278719075f692393 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Mon, 8 Jun 2026 19:24:27 +0000 Subject: [PATCH 08/12] feat(driver): execute-time validation of declared parameters Single shared hook in prepare_statement validates declared params on the original user params before driver style conversion, covering all adapters and execution methods. Declared params must be present (named binding); present non-None values whose declared type resolves via the registry must satisfy isinstance. None is allowed (SQL NULL); unresolved types are documentation-only; extra params (filter-injected limit/offset) are never rejected. execute_many checks the first row only; positional binding is skipped (arity validated at load). No-op unless declarations are present. --- sqlspec/driver/_common.py | 55 ++++++- .../driver/test_declared_param_validation.py | 137 ++++++++++++++++++ 2 files changed, 191 insertions(+), 1 deletion(-) create mode 100644 tests/unit/driver/test_declared_param_validation.py diff --git a/sqlspec/driver/_common.py b/sqlspec/driver/_common.py index 6be92d54f..ec5a639f4 100644 --- a/sqlspec/driver/_common.py +++ b/sqlspec/driver/_common.py @@ -17,6 +17,7 @@ from sqlspec.core import ( SQL, CachedStatement, + ParameterDeclaration, ParameterStyle, SQLResult, Statement, @@ -24,6 +25,7 @@ TypedParameter, get_cache, get_cache_config, + resolve_param_type, split_sql_script, ) from sqlspec.core._pool import get_processed_state_pool, get_sql_pool @@ -43,7 +45,13 @@ coerce_arrow_table, create_storage_job, ) -from sqlspec.exceptions import ImproperConfigurationError, NotFoundError, SQLFileNotFoundError, StorageCapabilityError +from sqlspec.exceptions import ( + ImproperConfigurationError, + NotFoundError, + SQLFileNotFoundError, + SQLSpecError, + StorageCapabilityError, +) from sqlspec.observability import ObservabilityRuntime, get_trace_context, resolve_db_system from sqlspec.protocols import HasDataProtocol, HasExecuteProtocol, StatementProtocol from sqlspec.utils.dispatch import TypeDispatcher @@ -328,6 +336,50 @@ def hash_stack_operations(stack: "StatementStack") -> "tuple[str, ...]": return tuple(hashes) +def _check_declared_named_row( + declared: "tuple[ParameterDeclaration, ...]", supplied: "dict[str, Any]" +) -> None: + """Validate a single named-parameter mapping against declared params. + + Each declared param must be present; a present non-``None`` value whose declared + type resolves via the registry must satisfy ``isinstance``. ``None`` is allowed + (SQL ``NULL``); unresolved types are documentation-only. Extra keys are ignored. + """ + for declaration in declared: + name = declaration.name + if name not in supplied: + raise SQLSpecError(f"Missing required parameter '{name}' for declared SQL statement.") + value = supplied[name] + if value is None: + continue + resolved = resolve_param_type(declaration.type_str) + if resolved is not None and not isinstance(value, resolved): + raise SQLSpecError( + f"Parameter '{name}' expected type '{declaration.type_str}' but got {type(value).__name__}." + ) + + +def _validate_declared_parameters(sql_statement: "SQL") -> None: + """Enforce declared-parameter contracts on the original user params. + + No-op unless the statement carries declarations and is not a script. Runs before + driver style conversion, so declared names and raw values are intact. Named binding + is validated for presence and type; ``execute_many`` checks the first row only; + positional binding is skipped (arity is validated at load time). + """ + declared = sql_statement.declared_parameters + if not declared or sql_statement.is_script: + return + if sql_statement.is_many: + rows = sql_statement.positional_parameters + if rows and isinstance(rows[0], dict): + _check_declared_named_row(declared, rows[0]) + return + if sql_statement.positional_parameters: + return + _check_declared_named_row(declared, sql_statement.named_parameters) + + class StackExecutionObserver: """Context manager that aggregates telemetry for stack execution.""" @@ -1017,6 +1069,7 @@ def prepare_statement( if not filters and not kwargs and isinstance(statement, str): self._statement_cache[statement] = sql_statement + _validate_declared_parameters(sql_statement) return self._apply_filters(sql_statement, filters) def split_script_statements( diff --git a/tests/unit/driver/test_declared_param_validation.py b/tests/unit/driver/test_declared_param_validation.py new file mode 100644 index 000000000..ed22242c8 --- /dev/null +++ b/tests/unit/driver/test_declared_param_validation.py @@ -0,0 +1,137 @@ +"""Execute-time validation of declared parameters (Ch5, sqlspec-smgc.5). + +The single shared hook in ``prepare_statement`` validates declared params on the +original user params before style conversion, for every adapter and every +execution method. Declared => validated (present + typed); undeclared => untouched. +""" + +from typing import Any + +import pytest + +from sqlspec.core import ParameterDeclaration, StatementConfig +from sqlspec.core.filters import LimitOffsetFilter +from sqlspec.core.statement import SQL +from sqlspec.driver import SyncDriverAdapterBase +from sqlspec.exceptions import SQLSpecError +from tests.conftest import requires_interpreted + +# pyright: reportPrivateUsage=false + +pytestmark = requires_interpreted + + +class _MockDriver(SyncDriverAdapterBase): + def __init__(self) -> None: + self.statement_config = StatementConfig() + + @property + def connection(self) -> "Any": + return None + + def dispatch_execute(self, *args: "Any", **kwargs: "Any") -> "Any": + raise NotImplementedError + + def dispatch_execute_many(self, *args: "Any", **kwargs: "Any") -> "Any": + raise NotImplementedError + + def with_cursor(self, *args: "Any", **kwargs: "Any") -> "Any": + raise NotImplementedError + + def handle_database_exceptions(self, *args: "Any", **kwargs: "Any") -> "Any": + raise NotImplementedError + + def begin(self) -> None: + raise NotImplementedError + + def rollback(self) -> None: + raise NotImplementedError + + def commit(self) -> None: + raise NotImplementedError + + +@pytest.fixture +def driver() -> _MockDriver: + return _MockDriver() + + +def _declared(*decls: ParameterDeclaration) -> "tuple[ParameterDeclaration, ...]": + return decls + + +def test_required_missing_raises(driver: _MockDriver) -> None: + sql = SQL("select :a", declared_parameters=_declared(ParameterDeclaration("a", "int"))) + with pytest.raises(SQLSpecError, match="a"): + driver.prepare_statement(sql, ()) + + +def test_required_present_passes(driver: _MockDriver) -> None: + sql = SQL("select :a", declared_parameters=_declared(ParameterDeclaration("a", "int"))) + prepared = driver.prepare_statement(sql, ({"a": 1},)) + assert prepared.named_parameters == {"a": 1} + + +def test_type_mismatch_raises(driver: _MockDriver) -> None: + sql = SQL("select :a", declared_parameters=_declared(ParameterDeclaration("a", "int"))) + with pytest.raises(SQLSpecError, match="a"): + driver.prepare_statement(sql, ({"a": "not-an-int"},)) + + +def test_type_match_passes(driver: _MockDriver) -> None: + sql = SQL("select :a, :b", declared_parameters=_declared(ParameterDeclaration("a", "int"), ParameterDeclaration("b", "str"))) + prepared = driver.prepare_statement(sql, ({"a": 1, "b": "x"},)) + assert prepared.named_parameters == {"a": 1, "b": "x"} + + +def test_none_value_allowed_when_present(driver: _MockDriver) -> None: + """None means SQL NULL; the key is present so the param is supplied; type check skipped.""" + sql = SQL("select :a", declared_parameters=_declared(ParameterDeclaration("a", "int"))) + prepared = driver.prepare_statement(sql, ({"a": None},)) + assert prepared.named_parameters == {"a": None} + + +def test_unresolved_type_is_skipped(driver: _MockDriver) -> None: + """A type string not in the registry is documentation-only; no isinstance check.""" + sql = SQL("select :a", declared_parameters=_declared(ParameterDeclaration("a", "Money"))) + prepared = driver.prepare_statement(sql, ({"a": object()},)) + assert "a" in prepared.named_parameters + + +def test_extra_params_tolerated(driver: _MockDriver) -> None: + sql = SQL("select :a", declared_parameters=_declared(ParameterDeclaration("a", "int"))) + prepared = driver.prepare_statement(sql, ({"a": 1, "unexpected": 99},)) + assert prepared.named_parameters["a"] == 1 + + +def test_filter_injected_params_do_not_trip_validation(driver: _MockDriver) -> None: + """LimitOffsetFilter adds limit/offset after validation; they are never declared.""" + sql = SQL("select :a", declared_parameters=_declared(ParameterDeclaration("a", "int"))) + prepared = driver.prepare_statement(sql, ({"a": 1}, LimitOffsetFilter(limit=10, offset=0))) + assert prepared.named_parameters["a"] == 1 + + +def test_undeclared_query_is_untouched(driver: _MockDriver) -> None: + """No declarations => no validation, even with empty params.""" + sql = SQL("select :a") + prepared = driver.prepare_statement(sql, ()) + assert prepared.declared_parameters == () + + +def test_positional_binding_skips_name_checks(driver: _MockDriver) -> None: + """Positional binding can't be name-matched; arity was checked at load (Ch4).""" + sql = SQL("select ?", 1, declared_parameters=_declared(ParameterDeclaration("a", "int")), statement_config=StatementConfig()) + prepared = driver.prepare_statement(sql, ()) + assert prepared.positional_parameters == [1] + + +def test_execute_many_validates_first_row(driver: _MockDriver) -> None: + bad = SQL("select :a", [{"b": 2}, {"a": 1}], is_many=True, declared_parameters=_declared(ParameterDeclaration("a", "int")), statement_config=StatementConfig()) + with pytest.raises(SQLSpecError, match="a"): + driver.prepare_statement(bad, ()) + + +def test_execute_many_first_row_valid_passes(driver: _MockDriver) -> None: + good = SQL("select :a", [{"a": 1}, {"a": 2}], is_many=True, declared_parameters=_declared(ParameterDeclaration("a", "int")), statement_config=StatementConfig()) + prepared = driver.prepare_statement(good, ()) + assert prepared.is_many From 65294aa8c2afa39f68e88e1c6bfd0b916000d0f6 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Mon, 8 Jun 2026 19:42:37 +0000 Subject: [PATCH 09/12] docs(loader): document -- param: declared parameters Adds a Declared Parameters section to the SQL file loader usage guide (grammar, type vocabulary, register_param_type, load/execute-time validation timing, strict_parameter_annotations, introspection), a runnable declared_params.py example, and ParameterDeclaration/register_param_type/resolve_param_type autodoc in the loader reference. Documents the binary declaration model (no optional marker). Docs build clean under -W. --- docs/examples/sql_files/declared_params.py | 48 ++++++++++++ docs/reference/loader.rst | 15 ++++ docs/usage/sql_files.rst | 85 ++++++++++++++++++++++ 3 files changed, 148 insertions(+) create mode 100644 docs/examples/sql_files/declared_params.py diff --git a/docs/examples/sql_files/declared_params.py b/docs/examples/sql_files/declared_params.py new file mode 100644 index 000000000..3f837f758 --- /dev/null +++ b/docs/examples/sql_files/declared_params.py @@ -0,0 +1,48 @@ +from pathlib import Path + +__all__ = ("test_declared_params",) + + +def test_declared_params(tmp_path: "Path") -> None: + # start-example + from sqlspec import SQLSpec + from sqlspec.adapters.sqlite import SqliteConfig + from sqlspec.exceptions import SQLSpecError + + sql_file = tmp_path / "teams.sql" + sql_file.write_text( + "-- name: get_team_by_name\n" + "-- param: name str The team name to look up\n" + "select id, name from teams where name = :name\n" + ) + + spec = SQLSpec() + config = spec.add_config(SqliteConfig(connection_config={"database": ":memory:"})) + spec.load_sql_files(sql_file) + + # Introspect declared parameters without executing. + declarations = spec.get_query_parameters("get_team_by_name") + assert declarations[0].name == "name" + assert declarations[0].type_str == "str" + assert declarations[0].description == "The team name to look up" + + # The declarations also ride on the SQL object returned by get_sql(). + query = spec.get_sql("get_team_by_name") + assert query.declared_parameters == declarations + + with spec.provide_session(config) as session: + session.execute("create table teams (id integer primary key, name text)") + session.execute("insert into teams (name) values ('Litestar'), ('SQLSpec')") + + # A declared query validates supplied parameters automatically. + row = session.execute(query, {"name": "SQLSpec"}).one() + + # Omitting a declared parameter raises before the query reaches the driver. + try: + session.execute(spec.get_sql("get_team_by_name"), {}) + except SQLSpecError as exc: + missing_error = str(exc) + # end-example + + assert row["name"] == "SQLSpec" + assert "name" in missing_error diff --git a/docs/reference/loader.rst b/docs/reference/loader.rst index 079327a33..092ea5b7a 100644 --- a/docs/reference/loader.rst +++ b/docs/reference/loader.rst @@ -27,3 +27,18 @@ NamedStatement .. autoclass:: NamedStatement :members: :show-inheritance: + +Declared Parameters +=================== + +Parameters declared in SQL files via ``-- param:`` directives are exposed as +:class:`ParameterDeclaration` objects. See :ref:`Declared Parameters ` +for the grammar and validation behavior. + +.. autoclass:: sqlspec.ParameterDeclaration + :members: + :show-inheritance: + +.. autofunction:: sqlspec.register_param_type + +.. autofunction:: sqlspec.resolve_param_type diff --git a/docs/usage/sql_files.rst b/docs/usage/sql_files.rst index f392d958b..f14514157 100644 --- a/docs/usage/sql_files.rst +++ b/docs/usage/sql_files.rst @@ -56,10 +56,95 @@ Available where helpers: Each call returns a new ``SQL`` object (immutable chaining). +.. _declared-parameters: + +Declared Parameters +------------------- + +Declare a query's parameters inline with ``-- param:`` directives in the header +block. Declared queries become self-documenting, introspectable, and +self-validating -- without SQLSpec becoming an ORM. + +.. code-block:: sql + + -- name: get_offers_by_status + -- dialect: oracle + -- param: status_cd str The status code to filter by + -- param: offer_ids list[int] List of offer IDs to include + -- param: limit int Maximum number of rows to return + + select offer_id, offer_name from offers + where status_cd = :status_cd and offer_id in (:offer_ids) + fetch first :limit rows only + +The grammar is ``-- param: [description]``, placed alongside +``-- name:`` and ``-- dialect:`` in the leading comment block. + +.. literalinclude:: /examples/sql_files/declared_params.py + :language: python + :caption: ``declared parameters`` + :start-after: # start-example + :end-before: # end-example + :dedent: 4 + :no-upgrade: + +**Declaration is binary.** A query with **no** ``-- param:`` lines behaves +exactly as before -- same code path, zero overhead. Declaring a parameter opts +*that* query into validation: + +- It must be **supplied** when the query executes. +- If its declared type resolves to a Python type, the supplied value must match + (``isinstance``). Pass ``None`` for SQL ``NULL`` -- the key is still present and + the type check is skipped. +- Extra parameters are never rejected -- statement filters legitimately inject + ``limit``/``offset``, so only *declared* names are checked. + +There is no optional marker: a loaded ``.sql`` file's placeholders are static and +always bound, so a declared parameter is always required. Parameters that can +genuinely be absent (filter-injected ``limit``/``offset``) are simply left +undeclared. + +**Type vocabulary.** Declared types resolve through a fixed allowlist -- +``str``, ``int``, ``float``, ``bool``, ``bytes``, ``date``, ``datetime``, +``time``, ``Decimal``, and the container forms ``list``, ``list[int]``, +``list[str]``, ``list[float]``, ``list[bool]``, ``tuple``. The raw string is +always stored and **never** evaluated. Register custom mappings with +:func:`~sqlspec.register_param_type`: + +.. code-block:: python + + from decimal import Decimal + + from sqlspec import register_param_type + + register_param_type("Money", Decimal) # -- param: price Money + +Type strings that do not resolve are documentation-only -- their values are not +type-checked. + +**Validation timing.** + +- *Load time* -- declared names are cross-checked against the actual + ``:placeholders`` (name drift), and declared count against placeholder count for + positionally-bound queries. Mismatches raise :exc:`~sqlspec.exceptions.SQLSpecError`. +- *Execute time* -- presence and type are enforced for every declared parameter, + uniformly across every adapter. ``execute_many`` checks the first row only. + +A **malformed** ``-- param:`` line (a typo or wrong arity) is a soft warning and +the line is skipped, preserving backward compatibility. Pass +``strict_parameter_annotations=True`` to :class:`~sqlspec.loader.SQLFileLoader` +to escalate malformed annotations to an error. (A genuine *validation mismatch* +-- drift, count, missing, or wrong type -- always raises.) + +**Introspection.** Read declarations without executing via +``spec.get_query_parameters(name)`` or the ``declared_parameters`` tuple on the +``SQL`` object returned by ``spec.get_sql(name)``. + How Query Names Work -------------------- - Name queries with ``-- name: query_name`` comments. - SQLSpec normalizes names to snake_case for Python access. - Add ``-- dialect: postgres`` on the first line of a block to bind SQL to a dialect. +- Declare parameters with ``-- param: [description]`` (see `Declared Parameters`_). - Directory structures become namespaces when you load directories (``reports/daily.sql`` -> ``reports.``). From 5d4863d0df20531810d3f7c59a1d20daeb5129b8 Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Mon, 8 Jun 2026 20:03:24 +0000 Subject: [PATCH 10/12] feat(parameters): add declared parameters for execute-time validation refactor(error-handling): improve error messages for missing and type mismatch parameters feat(loader): add strict parameter annotations for SQL file loading chore(dependencies): update package versions in lock file --- sqlspec/core/statement.py | 3 +- sqlspec/driver/_common.py | 12 +++--- sqlspec/loader.py | 8 ++-- tests/unit/core/parameters/test_declared.py | 6 +-- .../driver/test_declared_param_validation.py | 28 +++++++++++-- uv.lock | 42 +++++++++---------- 6 files changed, 56 insertions(+), 43 deletions(-) diff --git a/sqlspec/core/statement.py b/sqlspec/core/statement.py index 85a92d157..de8273bec 100644 --- a/sqlspec/core/statement.py +++ b/sqlspec/core/statement.py @@ -321,6 +321,7 @@ def __init__( *parameters: Parameters and filters statement_config: Configuration is_many: Mark as execute_many operation + declared_parameters: Parameter declarations to validate against at execute time **kwargs: Additional parameters """ config = statement_config or self._create_auto_config(statement, parameters, kwargs) @@ -338,7 +339,7 @@ def __init__( self._is_script = False self._raw_expression: exp.Expr | None = None self._rebind_processor: ParameterProcessor | None = None - self._declared_parameters: "tuple[ParameterDeclaration, ...]" = declared_parameters + self._declared_parameters: tuple[ParameterDeclaration, ...] = declared_parameters if isinstance(statement, SQL): self._init_from_sql_object(statement) diff --git a/sqlspec/driver/_common.py b/sqlspec/driver/_common.py index ec5a639f4..d6c6ca238 100644 --- a/sqlspec/driver/_common.py +++ b/sqlspec/driver/_common.py @@ -336,9 +336,7 @@ def hash_stack_operations(stack: "StatementStack") -> "tuple[str, ...]": return tuple(hashes) -def _check_declared_named_row( - declared: "tuple[ParameterDeclaration, ...]", supplied: "dict[str, Any]" -) -> None: +def _check_declared_named_row(declared: "tuple[ParameterDeclaration, ...]", supplied: "dict[str, Any]") -> None: """Validate a single named-parameter mapping against declared params. Each declared param must be present; a present non-``None`` value whose declared @@ -348,15 +346,15 @@ def _check_declared_named_row( for declaration in declared: name = declaration.name if name not in supplied: - raise SQLSpecError(f"Missing required parameter '{name}' for declared SQL statement.") + msg = f"Missing required parameter '{name}' for declared SQL statement." + raise SQLSpecError(msg) value = supplied[name] if value is None: continue resolved = resolve_param_type(declaration.type_str) if resolved is not None and not isinstance(value, resolved): - raise SQLSpecError( - f"Parameter '{name}' expected type '{declaration.type_str}' but got {type(value).__name__}." - ) + msg = f"Parameter '{name}' expected type '{declaration.type_str}' but got {type(value).__name__}." + raise SQLSpecError(msg) def _validate_declared_parameters(sql_statement: "SQL") -> None: diff --git a/sqlspec/loader.py b/sqlspec/loader.py index 2fcd9325c..4ed31d7d4 100644 --- a/sqlspec/loader.py +++ b/sqlspec/loader.py @@ -46,8 +46,7 @@ DIALECT_PATTERN = re.compile(r"^\s*--\s*dialect\s*:\s*(?P[a-zA-Z0-9_]+)\s*$", re.IGNORECASE | re.MULTILINE) PARAM_PATTERN = re.compile( - r"^\s*--\s*param\s*:\s*(?P\w+)\s+(?P[\w.]+(?:\[[\w., ]+\])?)(?:\s+(?P.*\S))?\s*$", - re.IGNORECASE, + r"^\s*--\s*param\s*:\s*(?P\w+)\s+(?P[\w.]+(?:\[[\w., ]+\])?)(?:\s+(?P.*\S))?\s*$", re.IGNORECASE ) PARAM_PREFIX_PATTERN = re.compile(r"^\s*--\s*param\s*:", re.IGNORECASE) @@ -440,6 +439,7 @@ def _parse_sql_content( Args: content: Raw SQL file content to parse. file_path: File path for error reporting. + strict_parameter_annotations: Raise on malformed parameter declarations instead of skipping them. Returns: Dictionary mapping normalized statement names to NamedStatement objects. @@ -478,9 +478,7 @@ def _parse_sql_content( file_path, file_path, ValueError(f"Duplicate statement name: {raw_statement_name}") ) - SQLFileLoader._validate_declared_parameters( - clean_sql, declared_params, raw_statement_name, file_path - ) + SQLFileLoader._validate_declared_parameters(clean_sql, declared_params, raw_statement_name, file_path) statements[normalized_name] = NamedStatement( name=normalized_name, diff --git a/tests/unit/core/parameters/test_declared.py b/tests/unit/core/parameters/test_declared.py index fc9a5e5c5..48042b822 100644 --- a/tests/unit/core/parameters/test_declared.py +++ b/tests/unit/core/parameters/test_declared.py @@ -5,11 +5,7 @@ import pytest -from sqlspec.core.parameters._declared import ( - ParameterDeclaration, - register_param_type, - resolve_param_type, -) +from sqlspec.core.parameters._declared import ParameterDeclaration, register_param_type, resolve_param_type def test_declaration_fields() -> None: diff --git a/tests/unit/driver/test_declared_param_validation.py b/tests/unit/driver/test_declared_param_validation.py index ed22242c8..040650a92 100644 --- a/tests/unit/driver/test_declared_param_validation.py +++ b/tests/unit/driver/test_declared_param_validation.py @@ -79,7 +79,10 @@ def test_type_mismatch_raises(driver: _MockDriver) -> None: def test_type_match_passes(driver: _MockDriver) -> None: - sql = SQL("select :a, :b", declared_parameters=_declared(ParameterDeclaration("a", "int"), ParameterDeclaration("b", "str"))) + sql = SQL( + "select :a, :b", + declared_parameters=_declared(ParameterDeclaration("a", "int"), ParameterDeclaration("b", "str")), + ) prepared = driver.prepare_statement(sql, ({"a": 1, "b": "x"},)) assert prepared.named_parameters == {"a": 1, "b": "x"} @@ -120,18 +123,35 @@ def test_undeclared_query_is_untouched(driver: _MockDriver) -> None: def test_positional_binding_skips_name_checks(driver: _MockDriver) -> None: """Positional binding can't be name-matched; arity was checked at load (Ch4).""" - sql = SQL("select ?", 1, declared_parameters=_declared(ParameterDeclaration("a", "int")), statement_config=StatementConfig()) + sql = SQL( + "select ?", + 1, + declared_parameters=_declared(ParameterDeclaration("a", "int")), + statement_config=StatementConfig(), + ) prepared = driver.prepare_statement(sql, ()) assert prepared.positional_parameters == [1] def test_execute_many_validates_first_row(driver: _MockDriver) -> None: - bad = SQL("select :a", [{"b": 2}, {"a": 1}], is_many=True, declared_parameters=_declared(ParameterDeclaration("a", "int")), statement_config=StatementConfig()) + bad = SQL( + "select :a", + [{"b": 2}, {"a": 1}], + is_many=True, + declared_parameters=_declared(ParameterDeclaration("a", "int")), + statement_config=StatementConfig(), + ) with pytest.raises(SQLSpecError, match="a"): driver.prepare_statement(bad, ()) def test_execute_many_first_row_valid_passes(driver: _MockDriver) -> None: - good = SQL("select :a", [{"a": 1}, {"a": 2}], is_many=True, declared_parameters=_declared(ParameterDeclaration("a", "int")), statement_config=StatementConfig()) + good = SQL( + "select :a", + [{"a": 1}, {"a": 2}], + is_many=True, + declared_parameters=_declared(ParameterDeclaration("a", "int")), + statement_config=StatementConfig(), + ) prepared = driver.prepare_statement(good, ()) assert prepared.is_many diff --git a/uv.lock b/uv.lock index 15b571652..dbaccb838 100644 --- a/uv.lock +++ b/uv.lock @@ -1344,11 +1344,11 @@ wheels = [ [[package]] name = "distlib" -version = "0.4.1" +version = "0.4.2" source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/86/b2/d6fc3f2347f43dada79e5ff118493e8109c98400a0e29a1d5264a3aa479b/distlib-0.4.1.tar.gz", hash = "sha256:c3804d0d2d4b5fcd44036eb860cb6660485fcdf5c2aba53dc324d805837ea65b", size = 610526, upload-time = "2026-06-02T11:17:40.691Z" } +sdist = { url = "https://files.pythonhosted.org/packages/46/8d/873e9252ea2c0e0c857884e0a2899ec43ade132345df1925ef24cbe64f18/distlib-0.4.2.tar.gz", hash = "sha256:baeb401c90f27acd15c4861ae0847d1e731c27ac3dbf4210643ba61fa1e813db", size = 614914, upload-time = "2026-06-08T16:24:15.439Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/25/18/3497c4fa83a76dcb154923fd2075522e8dd6995ecee4093c00ae18160046/distlib-0.4.1-py2.py3-none-any.whl", hash = "sha256:9c2c552c68cbadc619f2d0ed3a69e27c351a3f4c9baa9ffb7df9e9cdc3d19a97", size = 469216, upload-time = "2026-06-02T11:17:38.779Z" }, + { url = "https://files.pythonhosted.org/packages/c1/60/aa891c893821d4d127292ed66c6940d1d715894bd5a0ce048056bc641773/distlib-0.4.2-py2.py3-none-any.whl", hash = "sha256:ca4cb11e5d746b5ec13c199cbf19ae27a241f89702b54e153a74332955446067", size = 470510, upload-time = "2026-06-08T16:24:13.208Z" }, ] [[package]] @@ -1482,7 +1482,7 @@ name = "exceptiongroup" version = "1.3.1" source = { registry = "https://pypi.org/simple" } dependencies = [ - { name = "typing-extensions", marker = "python_full_version < '3.11'" }, + { name = "typing-extensions", marker = "python_full_version < '3.12'" }, ] sdist = { url = "https://files.pythonhosted.org/packages/50/79/66800aadf48771f6b62f7eb014e352e5d06856655206165d775e675a02c9/exceptiongroup-1.3.1.tar.gz", hash = "sha256:8b412432c6055b0b7d14c310000ae93352ed6754f70fa8f7c34141f91c4e3219", size = 30371, upload-time = "2025-11-21T23:01:54.787Z" } wheels = [ @@ -2637,14 +2637,14 @@ wheels = [ [[package]] name = "joserfc" -version = "1.7.0" +version = "1.7.1" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "cryptography" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/d3/c3/2f590052b55cbdd0ace470ee7ee1f685f6882051be93a9374891005623e2/joserfc-1.7.0.tar.gz", hash = "sha256:4aced6ab0c47846f0a531402aec2419a874b91e918df9c4c9da8a82fb559d6c4", size = 232967, upload-time = "2026-06-02T09:59:34.506Z" } +sdist = { url = "https://files.pythonhosted.org/packages/44/90/25cb27518750218e4f850be63d8bbb2343efaad1c01c3571aaa4b3c33bd7/joserfc-1.7.1.tar.gz", hash = "sha256:77d0b76514879c68c6f433bc5b7357a4ab72008ff1e33d8379fd11d72bd8ca81", size = 233181, upload-time = "2026-06-08T07:21:33.412Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/5b/83/b6b62a66a06ce872d9429a5eb5ee20b2002fd9c331b953c94381c1f7c9f9/joserfc-1.7.0-py3-none-any.whl", hash = "sha256:17e5d7a5a35e65442b05efc435a3d5d46696ffa2c8a2ed0eea6f63fc268e3224", size = 70387, upload-time = "2026-06-02T09:59:33.264Z" }, + { url = "https://files.pythonhosted.org/packages/b3/00/fa62404c3e347f946faa13aa21085205f9cc06ad17671e37f81a51662ae8/joserfc-1.7.1-py3-none-any.whl", hash = "sha256:b3e3d655612e2e1ef67b2600f2f420e12e537b020208fab1761fad647319c164", size = 70423, upload-time = "2026-06-08T07:21:32.001Z" }, ] [[package]] @@ -7351,19 +7351,19 @@ wheels = [ [[package]] name = "tornado" -version = "6.5.6" +version = "6.5.7" source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/50/57/6d7303a77ae439d9189108f76c0c4fd89ee5e2cc8387bffb55232565c4ed/tornado-6.5.6.tar.gz", hash = "sha256:9a365179fe8ff6b8766f602c0f67c185d778193e9bdd828b19f0b6ed7764177d", size = 518139, upload-time = "2026-05-27T15:35:54.646Z" } +sdist = { url = "https://files.pythonhosted.org/packages/64/24/95ec527ad67b76d59299e5465b3935d05e4294b7e0290a3924b7487df30b/tornado-6.5.7.tar.gz", hash = "sha256:66c513a76cda70d53907bc27cf1447557699c2e95aa48ba27a442ff61c3ddfc2", size = 519252, upload-time = "2026-06-08T17:34:51.232Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/1b/0d/b4f481e18c5a51864e6d12b9a05ecf72919696680b747c958c3fc1f4fbae/tornado-6.5.6-cp39-abi3-macosx_10_9_universal2.whl", hash = "sha256:65fcfaafb079435c2c19dc9e07c0f1cf0fa9051759ed0a7d0a3ba7ea7f64919c", size = 447737, upload-time = "2026-05-27T15:35:38.122Z" }, - { url = "https://files.pythonhosted.org/packages/9e/9c/5430c39fcab1144d35860f457b15e9c08b4bc7ac86764354204e983d6183/tornado-6.5.6-cp39-abi3-macosx_10_9_x86_64.whl", hash = "sha256:38bc01b4acacded2de63ae78023548e41ebe6fbed3ec05a796d7ae3ad893887e", size = 445899, upload-time = "2026-05-27T15:35:40.519Z" }, - { url = "https://files.pythonhosted.org/packages/8b/79/fa7e14a2f939c807a8d30619b4eb604eab219601b78792516ebe22d40cf9/tornado-6.5.6-cp39-abi3-manylinux1_x86_64.manylinux_2_28_x86_64.manylinux_2_5_x86_64.whl", hash = "sha256:b942e6a137fda31ff54bf8e6e2c8d1c37f1f50583f3ed53fb840b53b9601d104", size = 448964, upload-time = "2026-05-27T15:35:42.106Z" }, - { url = "https://files.pythonhosted.org/packages/a7/71/bd67d5f5199f937dafe03a49a37989f60f600ff6fef34c79412a829d97bd/tornado-6.5.6-cp39-abi3-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:8666946e70171b8c3f1fc9b7876fac492e84822c4c7f3746f4e8f8bc9ac92a79", size = 449935, upload-time = "2026-05-27T15:35:43.906Z" }, - { url = "https://files.pythonhosted.org/packages/cc/a4/c24388c9cf5b3c3a513b56a158af9f23092c9a2810d789e294310797df21/tornado-6.5.6-cp39-abi3-musllinux_1_2_aarch64.whl", hash = "sha256:1c34cfab7ad6d104f052f55de06d39bbafc5885cfeb4da688803308dbcfa90b7", size = 449767, upload-time = "2026-05-27T15:35:45.793Z" }, - { url = "https://files.pythonhosted.org/packages/a5/eb/6a07ad550c3f7b37244bd0becdf293ec3d3e961783d8b720a97df50de1b2/tornado-6.5.6-cp39-abi3-musllinux_1_2_x86_64.whl", hash = "sha256:385f35e4e22fb52551dfcda4cdc8c30c61c2c001aef5ddad99cdfe116952efd3", size = 449174, upload-time = "2026-05-27T15:35:47.485Z" }, - { url = "https://files.pythonhosted.org/packages/bb/84/3469e098dccdb6763130e06aacd786bb4363fca7b590a55c101ddf34ed30/tornado-6.5.6-cp39-abi3-win32.whl", hash = "sha256:db475f1b67b2809b10bb16264829087724ca8d24fe4ed47f7b8675cae453ef86", size = 450230, upload-time = "2026-05-27T15:35:49.322Z" }, - { url = "https://files.pythonhosted.org/packages/d2/3c/273a04e0b9dd9016f1685cca0c1c8795a71ac88a34a8c889a0b443483226/tornado-6.5.6-cp39-abi3-win_amd64.whl", hash = "sha256:6739bf1e8eb09230f1280ddbd3236f0309db70f2c551a8dbc40f62babdf82f79", size = 450667, upload-time = "2026-05-27T15:35:51.194Z" }, - { url = "https://files.pythonhosted.org/packages/02/98/0cffe22a224f60c5fb1e3aa0b76f9da2e1ca78b0e9545e3d077c68ce60a7/tornado-6.5.6-cp39-abi3-win_arm64.whl", hash = "sha256:2543597b24a695d72338a9a77818362d72387c03ae173f1f169eadc5c91466ac", size = 449690, upload-time = "2026-05-27T15:35:52.902Z" }, + { url = "https://files.pythonhosted.org/packages/02/dc/c7043cab6fed8ae159fc1923ce829ada35c4dbd797d408a43858ffaf9639/tornado-6.5.7-cp39-abi3-macosx_10_9_universal2.whl", hash = "sha256:148b2eb15c2c765a50796172c1e499649b35f30d2e3c3d3e15913cfa56bfb163", size = 448543, upload-time = "2026-06-08T17:34:38.052Z" }, + { url = "https://files.pythonhosted.org/packages/92/4f/090b1431e5a43df696feceffc268c5383cc079ecb5f08ce58f917109aafe/tornado-6.5.7-cp39-abi3-macosx_10_9_x86_64.whl", hash = "sha256:9da38de27f1da3b78a966f0dae12b5a1ea9afe72ca805d84ff06508272ddf100", size = 446707, upload-time = "2026-06-08T17:34:39.594Z" }, + { url = "https://files.pythonhosted.org/packages/37/d8/ef374952fd5da67d4463122c2b8e5a96536ec10b4b339254c6dcde81d01c/tornado-6.5.7-cp39-abi3-manylinux1_x86_64.manylinux_2_28_x86_64.manylinux_2_5_x86_64.whl", hash = "sha256:8d759e71906ee783f8867b93bf26a265743da4c1e2f4a018464c1ba019862972", size = 449774, upload-time = "2026-06-08T17:34:41.204Z" }, + { url = "https://files.pythonhosted.org/packages/35/37/d434c73f4c6e014b745b9b37085f34f40c022f007efff3d7fe65991899f3/tornado-6.5.7-cp39-abi3-manylinux2014_aarch64.manylinux_2_17_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:8a46347a18f23fb92b396beebe0fb78f61dda0cc302445202c16203d8a18848b", size = 450745, upload-time = "2026-06-08T17:34:42.531Z" }, + { url = "https://files.pythonhosted.org/packages/b6/2b/56b9aff361d7f1ab728a805ec7d7ea835f8807afa9f5cc690ea0e630efb9/tornado-6.5.7-cp39-abi3-musllinux_1_2_aarch64.whl", hash = "sha256:7778b30bef919231265e91c69963ce0f49a1e9c07ac900bbe75b19ce2575ba92", size = 450578, upload-time = "2026-06-08T17:34:43.787Z" }, + { url = "https://files.pythonhosted.org/packages/02/30/a7444fb23aa76860a14198fab96ac79f1866b0a6e19e26c4381b0938e50f/tornado-6.5.7-cp39-abi3-musllinux_1_2_x86_64.whl", hash = "sha256:e726f0c75da7726eec023aa62751ff8878bd2737e34fbdd33b1ae5897d2200f5", size = 449985, upload-time = "2026-06-08T17:34:45.326Z" }, + { url = "https://files.pythonhosted.org/packages/5c/42/5f0e56c01e8d9d36f4e23f367b85ae6cae0c1ecddd5e6977d8388ad27488/tornado-6.5.7-cp39-abi3-win32.whl", hash = "sha256:f8de3bf12d3efdd0cbe7c8887868198f8a91415e3f29fcf258d9b8eb7b1d9ae4", size = 451047, upload-time = "2026-06-08T17:34:46.784Z" }, + { url = "https://files.pythonhosted.org/packages/c9/a4/b393076ffb21b469eec5b328a0534cf03a3b90bfc6b1f09507cdd075d938/tornado-6.5.7-cp39-abi3-win_amd64.whl", hash = "sha256:de942f843533a039ef9fa3d9c88c7cd8a7c94553fb5ad0154270989b3d99a2c4", size = 451485, upload-time = "2026-06-08T17:34:48.248Z" }, + { url = "https://files.pythonhosted.org/packages/71/2e/7b1c769803121b809112cf9a00681c472eae1d80e32d7ec0e0bd61d0d0e1/tornado-6.5.7-cp39-abi3-win_arm64.whl", hash = "sha256:ff934fce95643af5f11efdae618eaa73d469dc588641e5c8d19295a0c65c4796", size = 450506, upload-time = "2026-06-08T17:34:49.702Z" }, ] [[package]] @@ -7985,11 +7985,11 @@ wheels = [ [[package]] name = "wcwidth" -version = "0.8.0" +version = "0.8.1" source = { registry = "https://pypi.org/simple" } -sdist = { url = "https://files.pythonhosted.org/packages/af/44/c833e6b746ffb654e9abacf7ad6c2480a9c8c42e9637c1ae849964fb4dde/wcwidth-0.8.0.tar.gz", hash = "sha256:68a882ff6d14e3d14e0cae590b96a0551be64ce4905408112a8254434a1bdf69", size = 1305357, upload-time = "2026-06-05T21:19:35.667Z" } +sdist = { url = "https://files.pythonhosted.org/packages/49/b4/51fe890511f0f242d07cb1ebe6a5b6db417262b9d2568b460347c57d95cc/wcwidth-0.8.1.tar.gz", hash = "sha256:faf5b4a5366a72dc49cad48cdf21f52bdf63bdda995178e483ba247ff79089b9", size = 1466072, upload-time = "2026-06-08T05:57:23.146Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/fb/17/c68b6cbcfeadbf420b3c3edaf8fda51335bc9c38732adb2d3ba8984dc607/wcwidth-0.8.0-py3-none-any.whl", hash = "sha256:8c75e6099cefd197c4bcc67a486f70b5dbc68f997c05f34a811d853910450d64", size = 324935, upload-time = "2026-06-05T21:19:33.999Z" }, + { url = "https://files.pythonhosted.org/packages/bd/6e/95b0e537de1f4d4301f76f944642c6da50d1511cc7b3d64dc418a66c7509/wcwidth-0.8.1-py3-none-any.whl", hash = "sha256:f453740b1e4a4f3291faa37944c555d71056c4da08d59809b307ef4feba695c8", size = 323092, upload-time = "2026-06-08T05:57:21.413Z" }, ] [[package]] From d3901a545f50838ba1dd6afd17250355c7d7250a Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Mon, 8 Jun 2026 20:14:21 +0000 Subject: [PATCH 11/12] test: fix declared parameter CI checks --- tests/unit/adapters/test_sqlite/test_config.py | 2 +- tests/unit/driver/test_declared_param_validation.py | 4 ++++ 2 files changed, 5 insertions(+), 1 deletion(-) diff --git a/tests/unit/adapters/test_sqlite/test_config.py b/tests/unit/adapters/test_sqlite/test_config.py index bf42ae21c..926ae60bb 100644 --- a/tests/unit/adapters/test_sqlite/test_config.py +++ b/tests/unit/adapters/test_sqlite/test_config.py @@ -14,7 +14,7 @@ def _annotation_contains(annotation: object, expected: object) -> bool: """Return whether an annotation tree contains the expected object.""" - if annotation is expected: + if annotation is expected or annotation == expected: return True return any(_annotation_contains(arg, expected) for arg in get_args(annotation)) diff --git a/tests/unit/driver/test_declared_param_validation.py b/tests/unit/driver/test_declared_param_validation.py index 040650a92..bb2624593 100644 --- a/tests/unit/driver/test_declared_param_validation.py +++ b/tests/unit/driver/test_declared_param_validation.py @@ -29,6 +29,10 @@ def __init__(self) -> None: def connection(self) -> "Any": return None + @property + def data_dictionary(self) -> "Any": + raise NotImplementedError + def dispatch_execute(self, *args: "Any", **kwargs: "Any") -> "Any": raise NotImplementedError From 0daf628f2a246a4bbf1486b8a739da2b48c7549e Mon Sep 17 00:00:00 2001 From: Cody Fincher Date: Mon, 8 Jun 2026 23:00:06 +0000 Subject: [PATCH 12/12] feat: expand declared parameter validation --- docs/examples/sql_files/declared_params.py | 8 ++ docs/usage/sql_files.rst | 37 +++--- sqlspec/__init__.py | 4 + sqlspec/core/__init__.py | 4 + sqlspec/core/parameters/__init__.py | 10 +- sqlspec/core/parameters/_declared.py | 107 ++++++++++++++---- sqlspec/core/statement.py | 10 +- sqlspec/driver/_common.py | 23 +++- sqlspec/loader.py | 24 ++-- tests/unit/core/parameters/test_declared.py | 29 ++++- .../core/test_declared_params_carriage.py | 18 ++- .../driver/test_declared_param_validation.py | 64 +++++++++++ tests/unit/loader/test_param_directives.py | 33 ++++-- 13 files changed, 309 insertions(+), 62 deletions(-) diff --git a/docs/examples/sql_files/declared_params.py b/docs/examples/sql_files/declared_params.py index 3f837f758..604d9d82f 100644 --- a/docs/examples/sql_files/declared_params.py +++ b/docs/examples/sql_files/declared_params.py @@ -14,6 +14,10 @@ def test_declared_params(tmp_path: "Path") -> None: "-- name: get_team_by_name\n" "-- param: name str The team name to look up\n" "select id, name from teams where name = :name\n" + "\n" + "-- name: list_teams\n" + "-- param: name str? Optional team name filter\n" + "select id, name from teams where (:name is null or name = :name) order by id\n" ) spec = SQLSpec() @@ -37,6 +41,9 @@ def test_declared_params(tmp_path: "Path") -> None: # A declared query validates supplied parameters automatically. row = session.execute(query, {"name": "SQLSpec"}).one() + # Optional named parameters are bound as NULL when omitted. + optional_rows = session.execute(spec.get_sql("list_teams")).all() + # Omitting a declared parameter raises before the query reaches the driver. try: session.execute(spec.get_sql("get_team_by_name"), {}) @@ -45,4 +52,5 @@ def test_declared_params(tmp_path: "Path") -> None: # end-example assert row["name"] == "SQLSpec" + assert [team["name"] for team in optional_rows] == ["Litestar", "SQLSpec"] assert "name" in missing_error diff --git a/docs/usage/sql_files.rst b/docs/usage/sql_files.rst index f14514157..5dd3c1065 100644 --- a/docs/usage/sql_files.rst +++ b/docs/usage/sql_files.rst @@ -78,7 +78,9 @@ self-validating -- without SQLSpec becoming an ORM. fetch first :limit rows only The grammar is ``-- param: [description]``, placed alongside -``-- name:`` and ``-- dialect:`` in the leading comment block. +``-- name:`` and ``-- dialect:`` in the leading comment block. Append ``?`` to +the declared type, or end the description with ``(optional)``, to mark a named +parameter as optional. .. literalinclude:: /examples/sql_files/declared_params.py :language: python @@ -88,28 +90,36 @@ The grammar is ``-- param: [description]``, placed alongside :dedent: 4 :no-upgrade: -**Declaration is binary.** A query with **no** ``-- param:`` lines behaves +**Declaration is opt-in.** A query with **no** ``-- param:`` lines behaves exactly as before -- same code path, zero overhead. Declaring a parameter opts *that* query into validation: -- It must be **supplied** when the query executes. +- Required declarations must be **supplied** when the query executes. +- Missing optional named declarations are bound as ``None``, so SQL receives + ``NULL``. The query must still express the intended nullable behavior, for + example ``(:status_cd is null or status_cd = :status_cd)``. +- Positional placeholders still rely on arity and cannot be omitted by name. - If its declared type resolves to a Python type, the supplied value must match (``isinstance``). Pass ``None`` for SQL ``NULL`` -- the key is still present and the type check is skipped. - Extra parameters are never rejected -- statement filters legitimately inject ``limit``/``offset``, so only *declared* names are checked. -There is no optional marker: a loaded ``.sql`` file's placeholders are static and -always bound, so a declared parameter is always required. Parameters that can -genuinely be absent (filter-injected ``limit``/``offset``) are simply left -undeclared. +.. code-block:: sql + + -- name: list_offers + -- param: status_cd str? Optional status filter + select offer_id, offer_name from offers + where (:status_cd is null or status_cd = :status_cd) **Type vocabulary.** Declared types resolve through a fixed allowlist -- ``str``, ``int``, ``float``, ``bool``, ``bytes``, ``date``, ``datetime``, -``time``, ``Decimal``, and the container forms ``list``, ``list[int]``, -``list[str]``, ``list[float]``, ``list[bool]``, ``tuple``. The raw string is -always stored and **never** evaluated. Register custom mappings with -:func:`~sqlspec.register_param_type`: +``time``, ``Decimal``, ``uuid`` / ``uuid.UUID``, ``dict``, ``dict[str, Any]``, +``json`` / ``jsonb``, and the container forms ``list``, ``list[int]``, +``list[str]``, ``list[float]``, ``list[bool]``, ``tuple``. ``json`` and +``jsonb`` use SQLSpec's existing JSON serializer to validate that values can be +encoded. The raw string is always stored and **never** evaluated. Register +custom mappings with :func:`~sqlspec.register_param_type`: .. code-block:: python @@ -128,7 +138,8 @@ type-checked. ``:placeholders`` (name drift), and declared count against placeholder count for positionally-bound queries. Mismatches raise :exc:`~sqlspec.exceptions.SQLSpecError`. - *Execute time* -- presence and type are enforced for every declared parameter, - uniformly across every adapter. ``execute_many`` checks the first row only. + uniformly across every adapter. ``execute_many`` binds missing optional named + values on each row, then checks the first row only. A **malformed** ``-- param:`` line (a typo or wrong arity) is a soft warning and the line is skipped, preserving backward compatibility. Pass @@ -146,5 +157,5 @@ How Query Names Work - Name queries with ``-- name: query_name`` comments. - SQLSpec normalizes names to snake_case for Python access. - Add ``-- dialect: postgres`` on the first line of a block to bind SQL to a dialect. -- Declare parameters with ``-- param: [description]`` (see `Declared Parameters`_). +- Declare parameters with ``-- param: [?] [description]`` (see `Declared Parameters`_). - Directory structures become namespaces when you load directories (``reports/daily.sql`` -> ``reports.``). diff --git a/sqlspec/__init__.py b/sqlspec/__init__.py index 4ac4bca01..39988db5d 100644 --- a/sqlspec/__init__.py +++ b/sqlspec/__init__.py @@ -55,6 +55,7 @@ ParameterProcessor, ParameterStyle, ParameterStyleConfig, + ParamTypeMatcher, ProcessedState, SQLResult, StackOperation, @@ -62,6 +63,7 @@ Statement, StatementConfig, StatementStack, + matches_param_type, register_param_type, resolve_param_type, ) @@ -115,6 +117,7 @@ "Merge", "ObservabilityConfig", "ObservabilityRuntime", + "ParamTypeMatcher", "ParameterConverter", "ParameterDeclaration", "ParameterProcessor", @@ -162,6 +165,7 @@ "filters", "format_statement_event", "loader", + "matches_param_type", "migrations", "register_param_type", "resolve_param_type", diff --git a/sqlspec/core/__init__.py b/sqlspec/core/__init__.py index 53512468e..60ff4172e 100644 --- a/sqlspec/core/__init__.py +++ b/sqlspec/core/__init__.py @@ -161,6 +161,7 @@ ParameterStyle, ParameterStyleConfig, ParameterValidator, + ParamTypeMatcher, TypedParameter, build_literal_inlining_transform, build_null_pruning_transform, @@ -169,6 +170,7 @@ get_driver_profile, is_iterable_parameters, looks_like_execute_many, + matches_param_type, normalize_parameter_key, register_driver_profile, register_param_type, @@ -281,6 +283,7 @@ "OperationProfile", "OperationType", "OrderByFilter", + "ParamTypeMatcher", "ParameterConverter", "ParameterDeclaration", "ParameterInfo", @@ -366,6 +369,7 @@ "is_iterable_parameters", "log_cache_stats", "looks_like_execute_many", + "matches_param_type", "normalize_parameter_key", "parse_column_for_condition", "parse_datetime_rfc3339", diff --git a/sqlspec/core/parameters/__init__.py b/sqlspec/core/parameters/__init__.py index a04be7b66..33e461693 100644 --- a/sqlspec/core/parameters/__init__.py +++ b/sqlspec/core/parameters/__init__.py @@ -8,7 +8,13 @@ validate_parameter_alignment, ) from sqlspec.core.parameters._converter import ParameterConverter -from sqlspec.core.parameters._declared import ParameterDeclaration, register_param_type, resolve_param_type +from sqlspec.core.parameters._declared import ( + ParameterDeclaration, + ParamTypeMatcher, + matches_param_type, + register_param_type, + resolve_param_type, +) from sqlspec.core.parameters._processor import ParameterProcessor, structural_fingerprint, value_fingerprint from sqlspec.core.parameters._registry import ( DRIVER_PARAMETER_PROFILES, @@ -43,6 +49,7 @@ "EXECUTE_MANY_MIN_ROWS", "PARAMETER_REGEX", "DriverParameterProfile", + "ParamTypeMatcher", "ParameterConverter", "ParameterDeclaration", "ParameterInfo", @@ -63,6 +70,7 @@ "get_driver_profile", "is_iterable_parameters", "looks_like_execute_many", + "matches_param_type", "normalize_parameter_key", "register_driver_profile", "register_param_type", diff --git a/sqlspec/core/parameters/_declared.py b/sqlspec/core/parameters/_declared.py index 2804a4bc9..91cb543f3 100644 --- a/sqlspec/core/parameters/_declared.py +++ b/sqlspec/core/parameters/_declared.py @@ -6,35 +6,40 @@ pure lookup; declared type strings are never evaluated. """ +from collections.abc import Callable from datetime import date, datetime, time from decimal import Decimal +from typing import TypeAlias +from uuid import UUID -__all__ = ("ParameterDeclaration", "register_param_type", "resolve_param_type") +from sqlspec.utils.serializers import to_json +__all__ = ( + "ParamTypeMatcher", + "ParameterDeclaration", + "matches_param_type", + "register_param_type", + "resolve_param_type", +) -class ParameterDeclaration: - """A single parameter declared in a SQL file header.""" +ParamTypeMatcher: TypeAlias = type | tuple[type, ...] | Callable[[object], bool] - __slots__ = ("description", "name", "type_str") - - def __init__(self, name: str, type_str: str, description: "str | None" = None) -> None: - self.name = name - self.type_str = type_str - self.description = description - def __eq__(self, other: object) -> bool: - if not isinstance(other, ParameterDeclaration): - return NotImplemented - return self.name == other.name and self.type_str == other.type_str and self.description == other.description +_JSON_VALUE_TYPES = (dict, list, str, int, float, bool) - def __hash__(self) -> int: - return hash((self.name, self.type_str, self.description)) - def __repr__(self) -> str: - return f"ParameterDeclaration(name={self.name!r}, type_str={self.type_str!r}, description={self.description!r})" +def _is_json_value(value: object) -> bool: + """Return whether a value can be encoded by SQLSpec's JSON serializer.""" + if not isinstance(value, _JSON_VALUE_TYPES): + return False + try: + to_json(value) + except (TypeError, ValueError): + return False + return True -_TYPE_REGISTRY: "dict[str, type]" = { +_TYPE_REGISTRY: dict[str, ParamTypeMatcher] = { "str": str, "int": int, "float": float, @@ -44,6 +49,13 @@ def __repr__(self) -> str: "datetime": datetime, "time": time, "decimal": Decimal, + "uuid": UUID, + "uuid.uuid": UUID, + "dict": dict, + "dict[str,any]": dict, + "dict[str,object]": dict, + "json": _is_json_value, + "jsonb": _is_json_value, "list[int]": list, "list[str]": list, "list[float]": list, @@ -53,23 +65,54 @@ def __repr__(self) -> str: } +class ParameterDeclaration: + """A single parameter declared in a SQL file header.""" + + __slots__ = ("description", "name", "required", "type_str") + + def __init__(self, name: str, type_str: str, description: "str | None" = None, *, required: bool = True) -> None: + self.name = name + self.type_str = type_str + self.description = description + self.required = required + + def __eq__(self, other: object) -> bool: + if not isinstance(other, ParameterDeclaration): + return NotImplemented + return ( + self.name == other.name + and self.type_str == other.type_str + and self.description == other.description + and self.required == other.required + ) + + def __hash__(self) -> int: + return hash((self.name, self.type_str, self.description, self.required)) + + def __repr__(self) -> str: + return ( + f"ParameterDeclaration(name={self.name!r}, type_str={self.type_str!r}, " + f"description={self.description!r}, required={self.required!r})" + ) + + def _normalize_type_key(type_str: str) -> str: """Normalize a declared type string to its registry lookup key.""" return "".join(type_str.split()).lower() -def register_param_type(name: str, py_type: type) -> None: - """Register or override a declared-type-string to Python-type mapping. +def register_param_type(name: str, py_type: ParamTypeMatcher) -> None: + """Register or override a declared-type-string matcher. Args: name: The declared type string as written in ``-- param:`` (case-insensitive). - py_type: The Python type used for ``isinstance`` validation. + py_type: The Python type, tuple of types, or predicate used for validation. """ _TYPE_REGISTRY[_normalize_type_key(name)] = py_type -def resolve_param_type(type_str: str) -> "type | None": - """Resolve a declared type string to a Python type, or ``None`` if unknown. +def resolve_param_type(type_str: str) -> "ParamTypeMatcher | None": + """Resolve a declared type string to a matcher, or ``None`` if unknown. Unknown type strings are documentation-only and skipped during validation. The declared string is looked up, never evaluated. Parameterized containers @@ -79,6 +122,22 @@ def resolve_param_type(type_str: str) -> "type | None": type_str: The declared type string from a ``-- param:`` directive. Returns: - The resolved Python type, or ``None`` when not in the registry. + The resolved matcher, or ``None`` when not in the registry. """ return _TYPE_REGISTRY.get(_normalize_type_key(type_str)) + + +def matches_param_type(type_str: str, value: object) -> bool: + """Return whether a value satisfies a declared type string. + + Unknown type strings are documentation-only and always match. ``None`` is + handled by the driver as SQL ``NULL`` before this helper is called. + """ + resolved = resolve_param_type(type_str) + if resolved is None: + return True + if isinstance(resolved, tuple): + return isinstance(value, resolved) + if isinstance(resolved, type): + return isinstance(value, resolved) + return resolved(value) diff --git a/sqlspec/core/statement.py b/sqlspec/core/statement.py index de8273bec..a181f8c94 100644 --- a/sqlspec/core/statement.py +++ b/sqlspec/core/statement.py @@ -384,6 +384,7 @@ def __reduce__(self) -> "tuple[Any, ...]": self._is_many, self._is_script, dict(self._named_parameters), + self._declared_parameters, ), ) @@ -1970,9 +1971,16 @@ def _rebuild_sql( is_many: bool, is_script: bool, named_parameters: "dict[str, Any]", + declared_parameters: "tuple[ParameterDeclaration, ...]", ) -> "SQL": """Reconstruct a SQL instance from pickled / deepcopied state.""" - new_sql = SQL(raw_sql, *original_parameters, statement_config=statement_config, is_many=is_many) + new_sql = SQL( + raw_sql, + *original_parameters, + statement_config=statement_config, + is_many=is_many, + declared_parameters=declared_parameters, + ) if filters: new_sql._filters.extend(filters) if named_parameters: diff --git a/sqlspec/driver/_common.py b/sqlspec/driver/_common.py index d6c6ca238..7fa43db19 100644 --- a/sqlspec/driver/_common.py +++ b/sqlspec/driver/_common.py @@ -25,7 +25,7 @@ TypedParameter, get_cache, get_cache_config, - resolve_param_type, + matches_param_type, split_sql_script, ) from sqlspec.core._pool import get_processed_state_pool, get_sql_pool @@ -336,13 +336,22 @@ def hash_stack_operations(stack: "StatementStack") -> "tuple[str, ...]": return tuple(hashes) +def _apply_declared_optional_defaults(declared: "tuple[ParameterDeclaration, ...]", supplied: "dict[str, Any]") -> None: + """Bind missing optional named params as SQL NULL.""" + for declaration in declared: + if not declaration.required and declaration.name not in supplied: + supplied[declaration.name] = None + + def _check_declared_named_row(declared: "tuple[ParameterDeclaration, ...]", supplied: "dict[str, Any]") -> None: """Validate a single named-parameter mapping against declared params. - Each declared param must be present; a present non-``None`` value whose declared - type resolves via the registry must satisfy ``isinstance``. ``None`` is allowed - (SQL ``NULL``); unresolved types are documentation-only. Extra keys are ignored. + Required params must be present. Missing optional params are bound as + ``None`` so SQL receives ``NULL``. A present non-``None`` value whose + declared type resolves via the registry must satisfy that matcher. + Unresolved types are documentation-only. Extra keys are ignored. """ + _apply_declared_optional_defaults(declared, supplied) for declaration in declared: name = declaration.name if name not in supplied: @@ -351,8 +360,7 @@ def _check_declared_named_row(declared: "tuple[ParameterDeclaration, ...]", supp value = supplied[name] if value is None: continue - resolved = resolve_param_type(declaration.type_str) - if resolved is not None and not isinstance(value, resolved): + if not matches_param_type(declaration.type_str, value): msg = f"Parameter '{name}' expected type '{declaration.type_str}' but got {type(value).__name__}." raise SQLSpecError(msg) @@ -371,6 +379,9 @@ def _validate_declared_parameters(sql_statement: "SQL") -> None: if sql_statement.is_many: rows = sql_statement.positional_parameters if rows and isinstance(rows[0], dict): + for row in rows: + if isinstance(row, dict): + _apply_declared_optional_defaults(declared, row) _check_declared_named_row(declared, rows[0]) return if sql_statement.positional_parameters: diff --git a/sqlspec/loader.py b/sqlspec/loader.py index 4ed31d7d4..f6f69e66d 100644 --- a/sqlspec/loader.py +++ b/sqlspec/loader.py @@ -46,10 +46,12 @@ DIALECT_PATTERN = re.compile(r"^\s*--\s*dialect\s*:\s*(?P[a-zA-Z0-9_]+)\s*$", re.IGNORECASE | re.MULTILINE) PARAM_PATTERN = re.compile( - r"^\s*--\s*param\s*:\s*(?P\w+)\s+(?P[\w.]+(?:\[[\w., ]+\])?)(?:\s+(?P.*\S))?\s*$", re.IGNORECASE + r"^\s*--\s*param\s*:\s*(?P\w+)\s+(?P[\w.]+(?:\[[\w., ]+\])?)(?P\?)?(?:\s+(?P.*\S))?\s*$", + re.IGNORECASE, ) PARAM_PREFIX_PATTERN = re.compile(r"^\s*--\s*param\s*:", re.IGNORECASE) +PARAM_OPTIONAL_DESCRIPTION_PATTERN = re.compile(r"(?:^|\s)\(optional\)\s*$", re.IGNORECASE) DIALECT_ALIASES: Final = { @@ -64,6 +66,18 @@ MIN_QUERY_PARTS: Final = 3 +def _parse_parameter_declaration(param_match: "re.Match[str]") -> ParameterDeclaration: + """Build a parameter declaration from a matched ``-- param:`` line.""" + description = param_match.group("desc") + required = param_match.group("optional") != "?" + if description is not None and PARAM_OPTIONAL_DESCRIPTION_PATTERN.search(description): + required = False + description = PARAM_OPTIONAL_DESCRIPTION_PATTERN.sub("", description).strip() or None + return ParameterDeclaration( + name=param_match.group("name"), type_str=param_match.group("type"), description=description, required=required + ) + + class NamedStatement: """Represents a parsed SQL statement with metadata. @@ -364,13 +378,7 @@ def _parse_directive_block( continue param_match = PARAM_PATTERN.match(stripped) if param_match: - params.append( - ParameterDeclaration( - name=param_match.group("name"), - type_str=param_match.group("type"), - description=param_match.group("desc"), - ) - ) + params.append(_parse_parameter_declaration(param_match)) continue if PARAM_PREFIX_PATTERN.match(stripped): if strict: diff --git a/tests/unit/core/parameters/test_declared.py b/tests/unit/core/parameters/test_declared.py index 48042b822..cae777cb1 100644 --- a/tests/unit/core/parameters/test_declared.py +++ b/tests/unit/core/parameters/test_declared.py @@ -2,10 +2,16 @@ from datetime import date, datetime from decimal import Decimal +from uuid import UUID import pytest -from sqlspec.core.parameters._declared import ParameterDeclaration, register_param_type, resolve_param_type +from sqlspec.core.parameters._declared import ( + ParameterDeclaration, + matches_param_type, + register_param_type, + resolve_param_type, +) def test_declaration_fields() -> None: @@ -20,12 +26,19 @@ def test_declaration_with_description() -> None: assert decl.description == "Max rows" +def test_optional_declaration_marks_required_false() -> None: + decl = ParameterDeclaration("status_cd", "str", required=False) + assert decl.required is False + + def test_declaration_equality_and_hash() -> None: a = ParameterDeclaration("a", "int") b = ParameterDeclaration("a", "int") c = ParameterDeclaration("a", "int", description="differs") + d = ParameterDeclaration("a", "int", required=False) assert a == b assert a != c + assert a != d assert a != "not-a-declaration" assert hash(a) == hash(b) @@ -41,6 +54,11 @@ def test_declaration_equality_and_hash() -> None: ("date", date), ("datetime", datetime), ("Decimal", Decimal), + ("uuid", UUID), + ("UUID", UUID), + ("uuid.UUID", UUID), + ("dict", dict), + ("dict[str, Any]", dict), ("list[int]", list), ("list[str]", list), ("list", list), @@ -55,6 +73,13 @@ def test_resolve_is_case_and_whitespace_insensitive() -> None: assert resolve_param_type("INT") is int +def test_json_type_uses_serializer_backed_matcher() -> None: + assert matches_param_type("json", {"ok": ["nested"]}) + assert matches_param_type("jsonb", ["ok"]) + assert not matches_param_type("json", {"bad": object()}) + assert not matches_param_type("json", object()) + + def test_resolve_unknown_returns_none() -> None: assert resolve_param_type("Money") is None assert resolve_param_type("frobnicate") is None @@ -75,7 +100,9 @@ def test_register_param_type_adds_and_resolves() -> None: def test_public_exports() -> None: from sqlspec import ParameterDeclaration as TopDecl + from sqlspec import matches_param_type as top_matches from sqlspec import register_param_type as top_register assert TopDecl is ParameterDeclaration + assert top_matches is matches_param_type assert top_register is register_param_type diff --git a/tests/unit/core/test_declared_params_carriage.py b/tests/unit/core/test_declared_params_carriage.py index 0db1f7393..41e473696 100644 --- a/tests/unit/core/test_declared_params_carriage.py +++ b/tests/unit/core/test_declared_params_carriage.py @@ -8,7 +8,7 @@ from sqlspec.core._pool import get_sql_pool from sqlspec.core.statement import SQL -_SENTINEL = (ParameterDeclaration("a", "int"),) +_SENTINEL = (ParameterDeclaration("a", "int", required=False),) def test_default_declared_parameters_is_empty_tuple() -> None: @@ -81,6 +81,22 @@ def test_constructor_kwarg_sets_declarations() -> None: assert sql.declared_parameters == _SENTINEL +def test_declarations_survive_deepcopy() -> None: + import copy + + sql = SQL("select :a", {"a": 1}, declared_parameters=_SENTINEL) + new = copy.deepcopy(sql) + assert new.declared_parameters == _SENTINEL + + +def test_declarations_survive_pickle_roundtrip() -> None: + import pickle + + sql = SQL("select :a", {"a": 1}, declared_parameters=_SENTINEL) + new = pickle.loads(pickle.dumps(sql)) + assert new.declared_parameters == _SENTINEL + + def test_declarations_survive_driver_prepare_with_filter() -> None: """Declarations must survive prepare_statement rebuild + filter application.""" from sqlspec.adapters.sqlite import SqliteConfig diff --git a/tests/unit/driver/test_declared_param_validation.py b/tests/unit/driver/test_declared_param_validation.py index bb2624593..360112ebf 100644 --- a/tests/unit/driver/test_declared_param_validation.py +++ b/tests/unit/driver/test_declared_param_validation.py @@ -6,6 +6,7 @@ """ from typing import Any +from uuid import uuid4 import pytest @@ -70,6 +71,37 @@ def test_required_missing_raises(driver: _MockDriver) -> None: driver.prepare_statement(sql, ()) +def test_optional_missing_is_bound_as_none(driver: _MockDriver) -> None: + sql = SQL("select :a", declared_parameters=_declared(ParameterDeclaration("a", "int", required=False))) + prepared = driver.prepare_statement(sql, ()) + assert prepared.named_parameters == {"a": None} + + +def test_optional_missing_preserves_supplied_params(driver: _MockDriver) -> None: + sql = SQL( + "select :a, :b", + declared_parameters=_declared( + ParameterDeclaration("a", "int", required=False), ParameterDeclaration("b", "str") + ), + ) + prepared = driver.prepare_statement(sql, ({"b": "kept"},)) + assert prepared.named_parameters == {"b": "kept", "a": None} + + +def test_execute_many_optional_missing_is_bound_as_none_for_each_row(driver: _MockDriver) -> None: + sql = SQL( + "select :a, :b", + [{"b": "first"}, {"a": 2, "b": "second"}], + is_many=True, + declared_parameters=_declared( + ParameterDeclaration("a", "int", required=False), ParameterDeclaration("b", "str") + ), + statement_config=StatementConfig(), + ) + prepared = driver.prepare_statement(sql, ()) + assert prepared.positional_parameters == [{"b": "first", "a": None}, {"a": 2, "b": "second"}] + + def test_required_present_passes(driver: _MockDriver) -> None: sql = SQL("select :a", declared_parameters=_declared(ParameterDeclaration("a", "int"))) prepared = driver.prepare_statement(sql, ({"a": 1},)) @@ -91,6 +123,38 @@ def test_type_match_passes(driver: _MockDriver) -> None: assert prepared.named_parameters == {"a": 1, "b": "x"} +def test_uuid_type_match_passes(driver: _MockDriver) -> None: + value = uuid4() + sql = SQL("select :a", declared_parameters=_declared(ParameterDeclaration("a", "uuid"))) + prepared = driver.prepare_statement(sql, ({"a": value},)) + assert prepared.named_parameters == {"a": value} + + +def test_uuid_type_mismatch_raises(driver: _MockDriver) -> None: + sql = SQL("select :a", declared_parameters=_declared(ParameterDeclaration("a", "uuid"))) + with pytest.raises(SQLSpecError, match="a"): + driver.prepare_statement(sql, ({"a": "not-a-uuid"},)) + + +def test_dict_type_match_passes(driver: _MockDriver) -> None: + sql = SQL("select :a", declared_parameters=_declared(ParameterDeclaration("a", "dict"))) + prepared = driver.prepare_statement(sql, ({"a": {"ok": True}},)) + assert prepared.named_parameters == {"a": {"ok": True}} + + +def test_json_type_accepts_json_container_and_scalar_values(driver: _MockDriver) -> None: + for value in ({"ok": True}, ["ok"], "ok", 1, 1.5, True): + sql = SQL("select :a", declared_parameters=_declared(ParameterDeclaration("a", "json"))) + prepared = driver.prepare_statement(sql, ({"a": value},)) + assert prepared.named_parameters == {"a": value} + + +def test_json_type_mismatch_raises(driver: _MockDriver) -> None: + sql = SQL("select :a", declared_parameters=_declared(ParameterDeclaration("a", "json"))) + with pytest.raises(SQLSpecError, match="a"): + driver.prepare_statement(sql, ({"a": object()},)) + + def test_none_value_allowed_when_present(driver: _MockDriver) -> None: """None means SQL NULL; the key is present so the param is supplied; type check skipped.""" sql = SQL("select :a", declared_parameters=_declared(ParameterDeclaration("a", "int"))) diff --git a/tests/unit/loader/test_param_directives.py b/tests/unit/loader/test_param_directives.py index 1c9da334a..bff135a19 100644 --- a/tests/unit/loader/test_param_directives.py +++ b/tests/unit/loader/test_param_directives.py @@ -10,23 +10,42 @@ @pytest.mark.parametrize( - ("line", "name", "type_str", "description"), + ("line", "name", "type_str", "required", "description"), [ - ("-- param: status_cd str The status code", "status_cd", "str", "The status code"), - ("-- param: limit int", "limit", "int", None), - ("-- param: offer_ids list[int] List of ids", "offer_ids", "list[int]", "List of ids"), - ("--param:x bool", "x", "bool", None), - ("-- PARAM: Y Decimal money", "Y", "Decimal", "money"), + ("-- param: status_cd str The status code", "status_cd", "str", True, "The status code"), + ("-- param: limit int", "limit", "int", True, None), + ("-- param: offer_ids list[int] List of ids", "offer_ids", "list[int]", True, "List of ids"), + ("-- param: status_cd str? Optional status filter", "status_cd", "str", False, "Optional status filter"), + ("-- param: payload dict[str, Any]? Optional payload", "payload", "dict[str, Any]", False, "Optional payload"), + ("--param:x bool", "x", "bool", True, None), + ("-- PARAM: Y Decimal money", "Y", "Decimal", True, "money"), ], ) -def test_param_pattern(line: str, name: str, type_str: str, description: "str | None") -> None: +def test_param_pattern(line: str, name: str, type_str: str, required: bool, description: "str | None") -> None: m = PARAM_PATTERN.match(line) assert m is not None assert m.group("name") == name assert m.group("type") == type_str + assert (m.group("optional") != "?") is required assert m.group("desc") == description +def test_parse_optional_declared_param_suffix() -> None: + content = "-- name: q\n-- param: status_cd str? Optional status filter\nselect :status_cd\n" + statements = SQLFileLoader._parse_sql_content(content, "test.sql") + assert statements["q"].parameters == ( + ParameterDeclaration("status_cd", "str", required=False, description="Optional status filter"), + ) + + +def test_parse_optional_declared_param_description_marker() -> None: + content = "-- name: q\n-- param: status_cd str Status filter (optional)\nselect :status_cd\n" + statements = SQLFileLoader._parse_sql_content(content, "test.sql") + assert statements["q"].parameters == ( + ParameterDeclaration("status_cd", "str", required=False, description="Status filter"), + ) + + def test_parse_declared_params_interleaved_with_dialect() -> None: content = """ -- name: get_offers