diff --git a/pipekit/drivers/base.py b/pipekit/drivers/base.py index 0441b23..c972269 100644 --- a/pipekit/drivers/base.py +++ b/pipekit/drivers/base.py @@ -167,6 +167,18 @@ class Driver(abc.ABC): """DDL to create new_table with the same columns as source_table.""" return f"CREATE TABLE {new_table} (LIKE {source_table} INCLUDING ALL);" + # ---- reconciliation aggregates ---- + # Used by reconcile.py to build the same aggregate on both sides of a sync. + # They must agree numerically across dialects, not merely parse. + + def sum_expression(self, type_raw: str, expr: str) -> str: + """SUM over ``expr``, widened where the dialect would overflow.""" + return f"SUM({expr})" + + def length_expression(self, expr: str) -> str: + """Character length of ``expr``.""" + return f"LENGTH({expr})" + def build_add_column_sql(self, qualified_table: str, column: dict) -> str: """DDL to append one column to an existing table. ALTER can only add at the end — which keeps the positional load aligned as long as the diff --git a/pipekit/drivers/mssql.py b/pipekit/drivers/mssql.py index 8b8c729..9772a96 100644 --- a/pipekit/drivers/mssql.py +++ b/pipekit/drivers/mssql.py @@ -209,6 +209,20 @@ class MSSQLDriver(Driver): return f"RTRIM({col})" return col + # T-SQL's SUM keeps the operand's type, so summing an int column overflows + # at 2^31 where Postgres would have promoted to bigint. Widen the integer + # types only — decimals already promote to decimal(38,s), and casting them + # would risk losing scale. + _INT_TYPES = {"int", "integer", "smallint", "tinyint", "bit"} + + def sum_expression(self, type_raw: str, expr: str) -> str: + if _base(type_raw) in self._INT_TYPES: + return f"SUM(CAST({expr} AS BIGINT))" + return f"SUM({expr})" + + def length_expression(self, expr: str) -> str: + return f"LEN({expr})" # T-SQL has no LENGTH + def map_type(self, type_raw: str) -> str: base = _base(type_raw) mapped = _TYPE_MAP.get(base, "text") diff --git a/reconcile.py b/reconcile.py index b0dfbfd..6f58f29 100755 --- a/reconcile.py +++ b/reconcile.py @@ -65,6 +65,19 @@ def source_expression(drv, column: dict) -> str: return drv.default_expression(column["source_type"], name) +# One name part: "quoted", [bracketed], or bare. A table reference is 1-4 of +# them dot-joined — SQL Server's linked-server form (server.db.schema.table) is +# the reason for 4, GP's CHG.dbo.SOP30200 the reason 2 was never enough. +_PART = r'(?:"[^"]+"|\[[^\]]+\]|[\w#$]+)' +_FROM_RE = re.compile(rf'FROM\s+({_PART}(?:\.{_PART}){{0,3}})', re.IGNORECASE) +_OPENQUERY_RE = re.compile(r'OPENQUERY\s*\(\s*([\w#$]+)\s*,', re.IGNORECASE) + + +def _base_name(ref: str) -> str: + """Last part of a dotted table reference, unquoted and lowercased.""" + return ref.split(".")[-1].strip('"').strip("[]").lower() + + def detect_source_from(source_query: str, dest_table: str) -> str | None: """Best-effort: find the base source table in the module's SELECT. @@ -72,18 +85,70 @@ def detect_source_from(source_query: str, dest_table: str) -> str | None: base name (the modules here name the dest after the source); falls back to the last FROM match. Returns None if nothing looks like a table. """ - matches = re.findall(r'FROM\s+("?[\w#]+"?\.?"?[\w#]*"?)', source_query, - flags=re.IGNORECASE) - matches = [m.strip() for m in matches if "." in m] + matches = [m.strip() for m in _FROM_RE.findall(source_query) if "." in m] if not matches: return None - dest_base = dest_table.split(".")[-1].strip('"').lower() + dest_base = _base_name(dest_table) for m in matches: - if m.split(".")[-1].strip('"').lower() == dest_base: + if _base_name(m) == dest_base: return m return matches[-1] +def detect_openquery_server(source_query: str) -> str | None: + """Linked-server name if the module reads through ``OPENQUERY(SRV, '...')``. + + The GP modules wrap their whole SELECT in OPENQUERY so the work runs on the + remote server. The table named inside that string (``CHG.dbo.SOP30200``) + does not resolve on the local connection, so reconcile has to push its + aggregate through the same wrapper rather than querying the table directly. + """ + m = _OPENQUERY_RE.search(source_query) + return m.group(1) if m else None + + +def build_openquery_sql(server: str, table: str, exprs: list[str], *, + limit: int = 7000) -> str: + """Aggregate ``exprs`` over ``table`` on linked server ``server``. + + Columns are aliased ``c0..cN`` because OPENQUERY rejects a passthrough + result set with unnamed columns (msg 8155); the aliases are positional and + results are still zipped by position, never by name. + + OPENQUERY's passthrough string is capped at 8000 characters and a wide + module overruns it (sop30200: 291 metrics), so the metric list is split + into chunks that each fit and CROSS JOINed back into the single row the + caller expects. `SELECT *` over the join preserves left-to-right order, so + the columns still line up with the labels. Each chunk costs one extra + remote scan — hence chunks as large as the cap allows. + + (The uncapped alternative, ``EXEC(@sql) AT server``, needs RPC Out enabled + on the linked server; GPSERVER has it off.) + """ + overhead = len("SELECT FROM ") + len(table) + chunks: list[list[str]] = [] + cur: list[str] = [] + cur_len = overhead + for i, e in enumerate(exprs): + piece = f"{e} AS c{i}" + add = len(piece) + (2 if cur else 0) + if cur and cur_len + add > limit: + chunks.append(cur) + cur, cur_len = [], overhead + add = len(piece) + cur.append(piece) + cur_len += add + if cur: + chunks.append(cur) + + parts = [] + for n, pieces in enumerate(chunks): + inner = f"SELECT {', '.join(pieces)} FROM {table}" + escaped = inner.replace("'", "''") + parts.append(f"OPENQUERY({server}, '{escaped}') q{n}") + return "SELECT * FROM " + "\n CROSS JOIN ".join(parts) + + def build_metrics(columns, key_names, drv, *, source: bool, quick: bool, super_quick: bool = False): """Return (list of (label, kind) metric descriptors, list of SQL exprs). @@ -118,11 +183,13 @@ def build_metrics(columns, key_names, drv, *, source: bool, quick: bool, e = drv.quote_identifier(c["dest_name"]) name = c["dest_name"] + col_type = c["source_type"] if source else c["dest_type"] + if kind == "numeric": labels.append((f"{name}: count", "int")) exprs.append(f"COUNT({e})") labels.append((f"{name}: sum", "num")) - exprs.append(f"SUM({e})") + exprs.append(drv.sum_expression(col_type, e)) if not quick: labels.append((f"{name}: min", "str")) exprs.append(f"MIN({e})") @@ -141,7 +208,7 @@ def build_metrics(columns, key_names, drv, *, source: bool, quick: bool, exprs.append(f"COUNT({e})") if not quick: labels.append((f"{name}: len_sum", "int")) - exprs.append(f"SUM(LENGTH({e}))") + exprs.append(drv.sum_expression("int", drv.length_expression(e))) return labels, exprs @@ -216,8 +283,19 @@ def main() -> int: src_sql = "SELECT\n " + "\n , ".join(src_exprs) + f"\nFROM {source_from}" dst_sql = "SELECT\n " + "\n , ".join(dst_exprs) + f"\nFROM {m['dest_table']}" + # An explicit --source-from is taken verbatim (the caller may already have + # written their own OPENQUERY). Otherwise, if the module reads through a + # linked server, push the aggregate down the same way: computing it here + # would drag every row across the link, and the inner table name does not + # resolve locally at all. + oq_server = None if args.source_from else detect_openquery_server(m["source_query"]) + if oq_server: + src_sql = build_openquery_sql(oq_server, source_from, src_exprs) + print(f"module {m['name']} (id {m['id']})") - print(f" source {src_conn['name']} FROM {source_from}") + src_label = (f"OPENQUERY({oq_server}) -> {source_from}" if oq_server + else source_from) + print(f" source {src_conn['name']} FROM {src_label}") print(f" dest {dst_conn['name']} FROM {m['dest_table']}") mode = (" [super-quick]" if args.super_quick else " [quick]" if args.quick else "")