|
| 1 | +""" |
| 2 | +Copyright (c) Microsoft Corporation. |
| 3 | +Licensed under the MIT license. |
| 4 | +Selects which ODBC provider (native driver package) mssql-python loads. |
| 5 | +
|
| 6 | +Two providers are supported: ``msodbcsql18`` (the Microsoft ODBC Driver 18, |
| 7 | +shipped by ``mssql_python_odbc``) and ``mssql-odbc`` (the Rust driver, shipped |
| 8 | +inside ``mssql_py_core`` / the ``mssql-python-rs`` wheel). Selection is process-wide and resolved exactly |
| 9 | +once, before the native driver loads, from — in precedence order — the |
| 10 | +``MSSQL_PYTHON_NATIVE_PROVIDER`` environment variable, the ``mssql_python.native_provider`` |
| 11 | +module property, then the release default. An unknown value fails closed rather |
| 12 | +than falling back. |
| 13 | +""" |
| 14 | + |
| 15 | +import os |
| 16 | +import threading |
| 17 | +import warnings |
| 18 | +import importlib |
| 19 | +from typing import Dict, Optional, Tuple |
| 20 | + |
| 21 | +from mssql_python.logging import logger |
| 22 | + |
| 23 | +NATIVE_PROVIDER_ENV_VAR = "MSSQL_PYTHON_NATIVE_PROVIDER" |
| 24 | + |
| 25 | +# Customer-facing provider identifiers. |
| 26 | +PROVIDER_MSODBCSQL18 = "msodbcsql18" |
| 27 | +PROVIDER_MSSQL_ODBC = "mssql-odbc" |
| 28 | + |
| 29 | +# Phase 1 default. Phase 2 flips this to PROVIDER_MSSQL_ODBC via a documented release. |
| 30 | +_DEFAULT_PROVIDER = PROVIDER_MSODBCSQL18 |
| 31 | + |
| 32 | +# Provider -> import package that ships its native binaries. |
| 33 | +# Keep in sync with ddbc_bindings.cpp's ProviderPackageForId / ProviderDistForId. |
| 34 | +_PACKAGE_BY_PROVIDER: Dict[str, str] = { |
| 35 | + PROVIDER_MSODBCSQL18: "mssql_python_odbc", |
| 36 | + PROVIDER_MSSQL_ODBC: "mssql_py_core", |
| 37 | +} |
| 38 | + |
| 39 | +# Provider -> the pip distribution that installs its package (for error hints). |
| 40 | +_DIST_BY_PROVIDER: Dict[str, str] = { |
| 41 | + PROVIDER_MSODBCSQL18: "mssql-python-odbc", |
| 42 | + PROVIDER_MSSQL_ODBC: "mssql-python-rs", |
| 43 | +} |
| 44 | + |
| 45 | + |
| 46 | +def _normalize(value: str) -> str: |
| 47 | + """Return the canonical provider id for ``value`` or raise ``ValueError``. |
| 48 | +
|
| 49 | + An unrecognized selection is rejected so a typo fails closed instead of |
| 50 | + silently loading the default provider. |
| 51 | + """ |
| 52 | + canonical = value.strip().lower() |
| 53 | + if canonical not in _PACKAGE_BY_PROVIDER: |
| 54 | + valid = ", ".join(sorted(_PACKAGE_BY_PROVIDER)) |
| 55 | + raise ValueError(f"Unknown ODBC provider {value!r}. Valid providers are: {valid}.") |
| 56 | + return canonical |
| 57 | + |
| 58 | + |
| 59 | +class ProviderManager: |
| 60 | + """Process-wide, resolve-once selector for the ODBC provider. |
| 61 | +
|
| 62 | + The selection freezes when :meth:`resolve` first runs (at native driver |
| 63 | + load). A later change to the module property is ignored with a warning, |
| 64 | + mirroring the connection-pool configuration model. |
| 65 | + """ |
| 66 | + |
| 67 | + _lock: threading.Lock = threading.Lock() |
| 68 | + _property_value: Optional[str] = None |
| 69 | + _resolved: Optional[str] = None |
| 70 | + _source: Optional[str] = None |
| 71 | + |
| 72 | + @classmethod |
| 73 | + def _compute(cls) -> Tuple[str, str]: |
| 74 | + """Apply precedence env var -> module property -> default (lock-free).""" |
| 75 | + env_value = os.environ.get(NATIVE_PROVIDER_ENV_VAR) |
| 76 | + if env_value and env_value.strip(): |
| 77 | + return _normalize(env_value), "environment" |
| 78 | + if cls._property_value is not None: |
| 79 | + return cls._property_value, "property" |
| 80 | + return _DEFAULT_PROVIDER, "default" |
| 81 | + |
| 82 | + @classmethod |
| 83 | + def set_property(cls, value: Optional[str]) -> None: |
| 84 | + """Set the module-property selection. |
| 85 | +
|
| 86 | + Accepts a provider id or ``None`` to clear. A change after the provider |
| 87 | + has been resolved is ignored with a warning; the env var still takes |
| 88 | + precedence over this value when both are set. |
| 89 | + """ |
| 90 | + with cls._lock: |
| 91 | + canonical = _normalize(value) if value is not None else None |
| 92 | + if cls._resolved is not None: |
| 93 | + if canonical != cls._resolved: |
| 94 | + cls._warn_frozen() |
| 95 | + return |
| 96 | + cls._property_value = canonical |
| 97 | + env_value = os.environ.get(NATIVE_PROVIDER_ENV_VAR) |
| 98 | + if canonical is not None and env_value and env_value.strip(): |
| 99 | + try: |
| 100 | + env_provider = _normalize(env_value) |
| 101 | + except ValueError: |
| 102 | + # Preserve the existing fail-closed error at connection time. |
| 103 | + return |
| 104 | + if canonical != env_provider: |
| 105 | + cls._warn_env_override(canonical, env_provider) |
| 106 | + |
| 107 | + @classmethod |
| 108 | + def resolve(cls) -> str: |
| 109 | + """Resolve and freeze the provider, returning its canonical id.""" |
| 110 | + with cls._lock: |
| 111 | + if cls._resolved is None: |
| 112 | + cls._resolved, cls._source = cls._compute() |
| 113 | + logger.info( |
| 114 | + "ODBC provider resolved to '%s' (source=%s)", |
| 115 | + cls._resolved, |
| 116 | + cls._source, |
| 117 | + ) |
| 118 | + return cls._resolved |
| 119 | + |
| 120 | + @classmethod |
| 121 | + def effective(cls) -> str: |
| 122 | + """Return the provider that would be used, without freezing it. |
| 123 | +
|
| 124 | + Reports the release default for an invalid selection (e.g. a mistyped |
| 125 | + env var) rather than raising - this backs the public getter and |
| 126 | + diagnostics, which must stay safe to read at any time. The hard |
| 127 | + failure for a bad selection surfaces at :meth:`resolve`/ |
| 128 | + :meth:`ensure_available` instead. |
| 129 | + """ |
| 130 | + with cls._lock: |
| 131 | + if cls._resolved is not None: |
| 132 | + return cls._resolved |
| 133 | + try: |
| 134 | + provider, _ = cls._compute() |
| 135 | + except ValueError: |
| 136 | + return _DEFAULT_PROVIDER |
| 137 | + return provider |
| 138 | + |
| 139 | + @classmethod |
| 140 | + def package_name(cls, provider: Optional[str] = None) -> str: |
| 141 | + """Return the import package that ships ``provider``'s native binaries.""" |
| 142 | + provider = provider or cls.effective() |
| 143 | + return _PACKAGE_BY_PROVIDER[provider] |
| 144 | + |
| 145 | + @classmethod |
| 146 | + def ensure_available(cls) -> str: |
| 147 | + """Verify the selected provider's package is installed, then freeze it. |
| 148 | +
|
| 149 | + Called before the native driver loads. Fails closed with an actionable |
| 150 | + error if the package is missing. The selection is only frozen (via |
| 151 | + :meth:`resolve`) once the package has been confirmed importable, so a |
| 152 | + failed check here does not permanently lock in a provider that never |
| 153 | + actually loaded - a later call can still select a different, installed |
| 154 | + provider instead of requiring a process restart. |
| 155 | + """ |
| 156 | + provider = cls.effective() |
| 157 | + package = _PACKAGE_BY_PROVIDER[provider] |
| 158 | + try: |
| 159 | + importlib.import_module(package) |
| 160 | + except ModuleNotFoundError as exc: |
| 161 | + if exc.name != package: |
| 162 | + # A transitive dependency of an installed package is missing, |
| 163 | + # or the package is broken - don't mask it as "not installed". |
| 164 | + raise |
| 165 | + dist = _DIST_BY_PROVIDER[provider] |
| 166 | + raise ImportError( |
| 167 | + f"The '{provider}' ODBC provider is selected but its package " |
| 168 | + f"'{package}' is not installed. Install it with: pip install {dist}" |
| 169 | + ) from exc |
| 170 | + return cls.resolve() |
| 171 | + |
| 172 | + @classmethod |
| 173 | + def is_frozen(cls) -> bool: |
| 174 | + """Whether the provider has been resolved and can no longer change.""" |
| 175 | + return cls._resolved is not None |
| 176 | + |
| 177 | + @classmethod |
| 178 | + def get_info(cls) -> Dict[str, object]: |
| 179 | + """Report the selected provider for diagnostics. |
| 180 | +
|
| 181 | + Never raises: an invalid selection is reported via the ``error`` key |
| 182 | + (with ``id`` falling back to the default) instead of propagating, so |
| 183 | + this stays safe to call at any time, including before a provider is |
| 184 | + chosen or resolvable. |
| 185 | + """ |
| 186 | + with cls._lock: |
| 187 | + if cls._resolved is not None: |
| 188 | + provider, source, error = cls._resolved, cls._source, None |
| 189 | + else: |
| 190 | + try: |
| 191 | + provider, source = cls._compute() |
| 192 | + error = None |
| 193 | + except ValueError as exc: |
| 194 | + provider, source, error = _DEFAULT_PROVIDER, None, str(exc) |
| 195 | + frozen = cls._resolved is not None |
| 196 | + |
| 197 | + version = None |
| 198 | + driver_path = None |
| 199 | + package = _PACKAGE_BY_PROVIDER[provider] |
| 200 | + try: |
| 201 | + provider_module = importlib.import_module(package) |
| 202 | + version = getattr(provider_module, "__version__", None) |
| 203 | + module_file = getattr(provider_module, "__file__", None) |
| 204 | + if module_file: |
| 205 | + from mssql_python import ddbc_bindings |
| 206 | + |
| 207 | + driver_path = ddbc_bindings._get_odbc_driver_path( |
| 208 | + os.path.dirname(os.path.abspath(module_file)), provider |
| 209 | + ) |
| 210 | + except Exception: # pylint: disable=broad-exception-caught |
| 211 | + # Diagnostics must remain safe even for a broken provider package. |
| 212 | + pass |
| 213 | + |
| 214 | + info: Dict[str, object] = { |
| 215 | + "id": provider, |
| 216 | + "package": package, |
| 217 | + "version": version, |
| 218 | + "driver_path": driver_path, |
| 219 | + "source": source, |
| 220 | + "frozen": frozen, |
| 221 | + } |
| 222 | + if error is not None: |
| 223 | + info["error"] = error |
| 224 | + return info |
| 225 | + |
| 226 | + @classmethod |
| 227 | + def _warn_env_override(cls, requested: str, effective: str) -> None: |
| 228 | + message = ( |
| 229 | + f"ODBC provider property was set to '{requested}', but " |
| 230 | + f"{NATIVE_PROVIDER_ENV_VAR} selects '{effective}' and takes precedence." |
| 231 | + ) |
| 232 | + logger.warning(message) |
| 233 | + warnings.warn(message, RuntimeWarning, stacklevel=3) |
| 234 | + |
| 235 | + @classmethod |
| 236 | + def _warn_frozen(cls) -> None: |
| 237 | + message = ( |
| 238 | + f"ODBC provider is already loaded as '{cls._resolved}'; ignoring the " |
| 239 | + f"change. Select a provider before the first connection, or set the " |
| 240 | + f"{NATIVE_PROVIDER_ENV_VAR} environment variable." |
| 241 | + ) |
| 242 | + logger.warning(message) |
| 243 | + warnings.warn(message, RuntimeWarning, stacklevel=3) |
| 244 | + |
| 245 | + @classmethod |
| 246 | + def _reset_for_testing(cls) -> None: |
| 247 | + """Reset selection state - for testing purposes only.""" |
| 248 | + with cls._lock: |
| 249 | + cls._property_value = None |
| 250 | + cls._resolved = None |
| 251 | + cls._source = None |
0 commit comments