Files
ware_house/async_msssql_query.py

369 lines
12 KiB
Python

"""Async SQL Server access layer used by the warehouse application.
The module centralizes DSN creation and exposes :class:`AsyncMSSQLClient`,
which lazily binds a SQLAlchemy async engine to the running event loop. The
implementation intentionally avoids pooling because the GUI schedules work on a
single shared background loop and pooled connections were a source of
cross-loop errors.
"""
from __future__ import annotations
import asyncio
import inspect
import json
import logging
import os
import time
import urllib.parse
from pathlib import Path
from typing import Any, Dict, Optional
from sqlalchemy import text
from sqlalchemy.ext.asyncio import create_async_engine
from sqlalchemy.pool import NullPool
from version_info import module_version
__version__ = module_version(__name__)
QUERY_PROFILER_LOG = Path(__file__).with_name("warehouse_query_profiler.log")
QUERY_PROFILER_CONFIG_PATH = Path(__file__).with_name("db_connection.json")
def _load_query_profiler_config() -> dict[str, Any]:
"""Read profiler settings from db_connection.json, with env overrides."""
config: dict[str, Any] = {
"enabled": False,
"slow_ms": 0,
"sql_limit": 4000,
}
try:
data = json.loads(QUERY_PROFILER_CONFIG_PATH.read_text(encoding="utf-8"))
section = data.get("query_profiler") if isinstance(data, dict) else None
if isinstance(section, dict):
config.update(section)
except Exception:
pass
if "WAREHOUSE_QUERY_PROFILER" in os.environ:
config["enabled"] = os.environ.get("WAREHOUSE_QUERY_PROFILER", "1").strip().lower() not in {
"0",
"false",
"no",
"off",
}
if "WAREHOUSE_QUERY_PROFILER_SLOW_MS" in os.environ:
config["slow_ms"] = os.environ.get("WAREHOUSE_QUERY_PROFILER_SLOW_MS", "0")
if "WAREHOUSE_QUERY_PROFILER_SQL_LIMIT" in os.environ:
config["sql_limit"] = os.environ.get("WAREHOUSE_QUERY_PROFILER_SQL_LIMIT", "4000")
return config
def _query_profiler_enabled() -> bool:
value = _load_query_profiler_config().get("enabled", False)
if isinstance(value, str):
return value.strip().lower() not in {"0", "false", "no", "off"}
return bool(value)
def _query_profiler_slow_ms() -> float:
try:
return float(_load_query_profiler_config().get("slow_ms", 0) or 0)
except Exception:
return 0.0
def _query_profiler_sql_limit() -> int:
try:
return int(_load_query_profiler_config().get("sql_limit", 4000) or 4000)
except Exception:
return 4000
try:
import pyodbc
# The desktop app opens short-lived SQL connections from a background
# asyncio loop. ODBC pooling can keep native handles alive for a while after
# the GUI closes, which is especially visible with pythonw.
pyodbc.pooling = False
except Exception:
pyodbc = None # type: ignore[assignment]
try:
import orjson as _json
def _dumps(obj: Any) -> str:
"""Serialize an object to JSON using the fastest available backend."""
return _json.dumps(obj, default=str).decode("utf-8")
except Exception:
import json as _json
def _dumps(obj: Any) -> str:
"""Serialize an object to JSON using the standard library fallback."""
return _json.dumps(obj, default=str)
def make_mssql_dsn(
*,
server: str,
database: str,
user: Optional[str] = None,
password: Optional[str] = None,
driver: str = "ODBC Driver 17 for SQL Server",
trust_server_certificate: bool = True,
encrypt: Optional[str] = None,
extra_odbc_kv: Optional[Dict[str, str]] = None,
) -> str:
"""Build a SQLAlchemy ``mssql+aioodbc`` DSN from SQL Server parameters."""
kv = {
"DRIVER": driver,
"SERVER": server,
"DATABASE": database,
"TrustServerCertificate": "Yes" if trust_server_certificate else "No",
}
if user:
kv["UID"] = user
if password:
kv["PWD"] = password
if encrypt:
kv["Encrypt"] = encrypt
if extra_odbc_kv:
kv.update(extra_odbc_kv)
odbc = ";".join(f"{k}={v}" for k, v in kv.items()) + ";"
return f"mssql+aioodbc:///?odbc_connect={urllib.parse.quote_plus(odbc)}"
def _compact_sql(sql: str, *, limit: int | None = None) -> str:
"""Collapse SQL whitespace so profiler entries stay readable."""
if limit is None:
limit = _query_profiler_sql_limit()
text_value = " ".join(str(sql or "").split())
if limit > 0 and len(text_value) > limit:
return text_value[:limit] + "...<truncated>"
return text_value
def _profile_params(params: Optional[Dict[str, Any]]) -> str:
"""Serialize SQL parameters for diagnostics."""
if not params:
return "{}"
try:
return _dumps(params)
except Exception:
return repr(params)
def _query_caller() -> str:
"""Return the first external Python frame that triggered the DB call."""
current_file = Path(__file__).resolve()
for frame in inspect.stack(context=0)[2:]:
try:
frame_file = Path(frame.filename).resolve()
except Exception:
continue
if frame_file == current_file:
continue
return f"{frame_file.name}:{frame.lineno}:{frame.function}"
return "unknown"
def _append_query_profile(
*,
method: str,
elapsed_ms: float,
rows: int | None,
rowcount: int | None,
commit: bool,
ok: bool,
sql: str,
params: Optional[Dict[str, Any]],
caller: str,
error: str = "",
) -> None:
"""Append one query timing line to the local profiler log."""
if not _query_profiler_enabled():
return
slow_ms = _query_profiler_slow_ms()
if ok and slow_ms > 0 and elapsed_ms < slow_ms:
return
try:
timestamp = time.strftime("%Y-%m-%d %H:%M:%S")
row_info = f"rows={rows}" if rows is not None else f"rowcount={rowcount}"
lines = [
(
f"{timestamp} | {elapsed_ms:.3f} ms | {method} | "
f"{row_info} | commit={int(commit)} | ok={int(ok)} | caller={caller}"
),
f"PARAMS: {_profile_params(params)}",
f"SQL: {_compact_sql(sql)}",
]
if error:
lines.append(f"ERROR: {error}")
with QUERY_PROFILER_LOG.open("a", encoding="utf-8") as handle:
handle.write("\n".join(lines) + "\n---\n")
except Exception:
# Profiling must never interfere with warehouse operations.
pass
class AsyncMSSQLClient:
"""Thin async query client for SQL Server.
The engine is created lazily on the currently running event loop and uses
:class:`sqlalchemy.pool.NullPool` to avoid recycling connections across
loops or threads.
"""
def __init__(self, dsn: str, *, echo: bool = False, log: bool = True):
"""Initialize the client without opening any connection immediately."""
self._dsn = dsn
self._echo = echo
self._engine = None
self._engine_loop: Optional[asyncio.AbstractEventLoop] = None
self._logger = logging.getLogger("AsyncMSSQLClient")
if log and not self._logger.handlers:
handler = logging.StreamHandler()
handler.setFormatter(logging.Formatter("[%(levelname)s] %(message)s"))
self._logger.addHandler(handler)
self._enable_log = log
async def _ensure_engine(self):
"""Create the async engine on first use for the current running loop."""
if self._engine is not None:
return
loop = asyncio.get_running_loop()
self._engine = create_async_engine(
self._dsn,
echo=self._echo,
# NullPool avoids reusing connections bound to a different event loop.
poolclass=NullPool,
# aioodbc must explicitly receive the loop to bind to.
connect_args={"loop": loop},
)
self._engine_loop = loop
if self._enable_log:
self._logger.info("Engine created on loop %s", id(loop))
async def dispose(self):
"""Dispose the engine on the loop where it was created."""
if self._engine is None:
return
if asyncio.get_running_loop() is self._engine_loop:
await self._engine.dispose()
else:
fut = asyncio.run_coroutine_threadsafe(self._engine.dispose(), self._engine_loop)
fut.result(timeout=2)
self._engine = None
if self._enable_log:
self._logger.info("Engine disposed")
async def query_json(
self,
sql: str,
params: Optional[Dict[str, Any]] = None,
*,
as_dict_rows: bool = False,
commit: bool = False,
) -> Dict[str, Any]:
"""Execute a query and return a JSON-friendly payload.
Args:
sql: SQL statement to execute.
params: Optional named parameters bound to the statement.
as_dict_rows: When ``True`` returns rows as dictionaries keyed by
column name; otherwise rows are returned as lists.
commit: When ``True`` the statement runs in a transaction that is
committed on success. Useful for SQL batches that both mutate
data and return a final result set.
Returns:
A dictionary containing column names, rows and elapsed execution
time in milliseconds.
"""
await self._ensure_engine()
t0 = time.perf_counter()
caller = _query_caller()
try:
async with (self._engine.begin() if commit else self._engine.connect()) as conn:
res = await conn.execute(text(sql), params or {})
rows = res.fetchall()
cols = list(res.keys())
except Exception as exc:
elapsed_ms = round((time.perf_counter() - t0) * 1000, 3)
_append_query_profile(
method="query_json",
elapsed_ms=elapsed_ms,
rows=None,
rowcount=None,
commit=commit,
ok=False,
sql=sql,
params=params,
caller=caller,
error=repr(exc),
)
raise
if as_dict_rows:
rows_out = [dict(zip(cols, row)) for row in rows]
else:
rows_out = [list(row) for row in rows]
elapsed_ms = round((time.perf_counter() - t0) * 1000, 3)
_append_query_profile(
method="query_json",
elapsed_ms=elapsed_ms,
rows=len(rows_out),
rowcount=None,
commit=commit,
ok=True,
sql=sql,
params=params,
caller=caller,
)
return {
"columns": cols,
"rows": rows_out,
"elapsed_ms": elapsed_ms,
}
async def exec(self, sql: str, params: Optional[Dict[str, Any]] = None, *, commit: bool = False) -> int:
"""Execute a DML statement and return its row count."""
await self._ensure_engine()
t0 = time.perf_counter()
caller = _query_caller()
try:
async with (self._engine.begin() if commit else self._engine.connect()) as conn:
res = await conn.execute(text(sql), params or {})
rowcount = res.rowcount or 0
except Exception as exc:
elapsed_ms = round((time.perf_counter() - t0) * 1000, 3)
_append_query_profile(
method="exec",
elapsed_ms=elapsed_ms,
rows=None,
rowcount=None,
commit=commit,
ok=False,
sql=sql,
params=params,
caller=caller,
error=repr(exc),
)
raise
elapsed_ms = round((time.perf_counter() - t0) * 1000, 3)
_append_query_profile(
method="exec",
elapsed_ms=elapsed_ms,
rows=None,
rowcount=rowcount,
commit=commit,
ok=True,
sql=sql,
params=params,
caller=caller,
)
return rowcount