File size: 14,601 Bytes
6bff5d9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
f3025d4
6bff5d9
 
 
 
 
 
 
 
81e5fe7
6bff5d9
 
 
 
 
 
 
cbc8d6a
 
 
 
 
6bff5d9
 
 
f3025d4
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6bff5d9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
f3025d4
 
 
 
 
 
6bff5d9
f3025d4
 
 
 
 
 
 
81e5fe7
6bff5d9
 
 
81e5fe7
 
 
 
6bff5d9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
f873f92
6bff5d9
 
 
 
 
 
 
3743cfe
 
 
6bff5d9
 
 
 
 
 
cbc8d6a
 
 
 
 
 
 
6bff5d9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
81e5fe7
 
 
 
 
 
 
 
 
6bff5d9
 
81e5fe7
 
cbc8d6a
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
81e5fe7
 
 
 
 
 
 
 
 
 
0721bb4
81e5fe7
 
 
 
 
 
 
 
 
 
 
 
 
 
3743cfe
 
 
81e5fe7
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
"""DbExecutor β€” runs a compiled IR against a user's external SQL database.

Pipeline:
  IR β†’ SqlCompiler.compile()  β†’  CompiledSql(sql, params)
       ↓
  sqlglot guard  (defense-in-depth: SELECT-only, no DML / DDL)
       ↓
  resolve creds (catalog.location_ref β†’ dbclient://{client_id} β†’ DatabaseClient
                 row β†’ Fernet decrypt)
       ↓
  asyncio.to_thread(_run_sync)
    β”” db_pipeline_service.engine_scope(db_type, creds)
       β”” session-level: default_transaction_read_only + statement_timeout=30s
                        (postgres / supabase only)
       β”” engine.execute(text(sql), params)
       ↓
  QueryResult (always returned β€” errors populate `.error`, never raised)
"""

from __future__ import annotations

import asyncio
import time
from concurrent.futures import ThreadPoolExecutor
from typing import Any

import sqlglot
import sqlglot.expressions as exp
from sqlalchemy import text

from ...catalog.models import Catalog, Source
from ...database_client.database_client_service import database_client_service
from ...database_client.engine import user_engine_cache
from ...db.postgres.connection import AsyncSessionLocal
from ...middlewares.logging import get_logger
from ...utils.db_credential_encryption import decrypt_credentials_dict
from ..compiler.sql import CompiledSql, SqlCompiler
from ..ir.models import QueryIR
from .base import BaseExecutor, QueryResult

# Orphaned by the F-4 tripwire (2026-07-23): the only consumer was the legacy
# non-postgres branch in `_run_sync`, now commented out there. Restore this import
# alongside that branch (kept out of the block above so import sorting stays clean).
# from ...pipeline.db_pipeline import db_pipeline_service

logger = get_logger("db_executor")

_QUERY_TIMEOUT_SECONDS = 30

# Dedicated pool for blocking customer-DB work (F-5, 2026-07-24).
#
# `asyncio.wait_for` cancels the awaiting COROUTINE; the worker thread underneath is
# not cancellable and runs to completion regardless. So every timed-out query leaves a
# thread occupied until the customer's server gives up (bounded by the connection's
# `statement_timeout`, which is best-effort β€” see engine.py). On the DEFAULT executor
# those abandoned workers accumulate in a pool of `min(32, cpu_count + 4)` that is
# shared with every other `to_thread` caller in the process β€” notably the tabular
# Parquet loader. A handful of slow customer queries could therefore stall unrelated
# work across the whole service.
#
# Isolating them means the blast radius of a slow customer database is queries against
# THAT class of work, not the entire process. Sized to the engine cache's own ceiling
# (_MAX_ENGINES=50 x _POOL_SIZE=1) β€” more threads than that cannot make progress
# anyway, since each needs a pooled connection.
_DB_THREAD_POOL = ThreadPoolExecutor(
    max_workers=50, thread_name_prefix="dbexec"
)
_DBCLIENT_PREFIX = "dbclient://"


class DbExecutor(BaseExecutor):
    """Executes compiled SQL on the user's registered DB.

    Constructed once per query with the user's catalog. The catalog is the
    source of truth for identifiers; the executor never touches the user's
    DB metadata at execution time.
    """

    def __init__(self, catalog: Catalog) -> None:
        self._catalog = catalog
        self._compiler = SqlCompiler(catalog)

    async def run(self, ir: QueryIR) -> QueryResult:
        started = time.perf_counter()
        table_name = ""
        source_name = ""
        try:
            source = self._find_source(ir.source_id)
            source_name = source.name
            table_name = next(
                (t.name for t in source.tables if t.table_id == ir.table_id), ""
            )
            if source.source_type != "schema":
                raise ValueError(
                    f"DbExecutor cannot run on source_type={source.source_type!r}; "
                    "expected 'schema'"
                )

            compiled = self._compiler.compile(ir)
            self._sqlglot_guard(compiled.sql)

            client_id = self._parse_client_id(source.location_ref)
            client = await self._fetch_client(client_id)
            if client.user_id != self._catalog.user_id:
                raise PermissionError(
                    f"DatabaseClient {client_id!r} owner mismatch "
                    f"(client.user_id != catalog.user_id)"
                )
            creds = decrypt_credentials_dict(client.credentials)

            # `run_in_executor` on the dedicated pool, not `asyncio.to_thread` (which
            # always uses the shared default executor). The timeout semantics are
            # unchanged β€” `wait_for` still stops US waiting after 30s β€” but a worker
            # abandoned by that timeout now occupies a DB-only thread. See
            # _DB_THREAD_POOL. (F-5)
            loop = asyncio.get_running_loop()
            columns, rows = await asyncio.wait_for(
                loop.run_in_executor(
                    _DB_THREAD_POOL,
                    self._run_sync,
                    client_id,
                    client.db_type,
                    creds,
                    compiled,
                ),
                timeout=_QUERY_TIMEOUT_SECONDS,
            )

            # The compiler bounded the SQL to `row_cap` (+1 when the IR was
            # unbounded). More than row_cap rows means the result was truncated.
            truncated = len(rows) > compiled.row_cap
            capped = rows[:compiled.row_cap]
            elapsed_ms = int((time.perf_counter() - started) * 1000)
            logger.info(
                "db query complete",
                source_id=ir.source_id,
                rows=len(capped),
                truncated=truncated,
                elapsed_ms=elapsed_ms,
            )
            return QueryResult(
                source_id=ir.source_id,
                backend="sql",
                columns=columns,
                rows=capped,
                row_count=len(capped),
                truncated=truncated,
                elapsed_ms=elapsed_ms,
                table_id=ir.table_id,
                table_name=table_name,
                source_name=source_name,
                query=compiled.sql,  # executed SQL, for traceability (KM-691)
            )

        except Exception as e:
            elapsed_ms = int((time.perf_counter() - started) * 1000)
            logger.error(
                "db executor failed",
                source_id=ir.source_id,
                # repr, not str: some exceptions (e.g. Fernet InvalidToken) have an
                # empty str(), which hides the real failure as error="".
                error=repr(e),
                elapsed_ms=elapsed_ms,
            )
            return QueryResult(
                source_id=ir.source_id,
                backend="sql",
                elapsed_ms=elapsed_ms,
                # `str(e) or repr(e)`, not bare repr: the log above already uses repr,
                # but this payload reaches the assembler prompt, the traceability
                # record and the report caveats β€” where a Fernet InvalidToken arrived
                # as an EMPTY string, so the user-facing artifact said nothing while
                # the log was diagnosable. Falling back only when str() is empty means
                # no existing error text changes. (F-26)
                error=str(e) or repr(e),
                table_id=ir.table_id,
                table_name=table_name,
                source_name=source_name,
            )

    # ------------------------------------------------------------------
    # Helpers
    # ------------------------------------------------------------------

    def _find_source(self, source_id: str) -> Source:
        for s in self._catalog.sources:
            if s.source_id == source_id:
                return s
        raise ValueError(f"source_id {source_id!r} not in catalog")

    @staticmethod
    def _parse_client_id(location_ref: str) -> str:
        if not location_ref.startswith(_DBCLIENT_PREFIX):
            raise ValueError(
                f"DbExecutor expects 'dbclient://...' location_ref, got {location_ref!r}"
            )
        client_id = location_ref[len(_DBCLIENT_PREFIX):]
        if not client_id:
            raise ValueError("location_ref is missing client_id after 'dbclient://'")
        return client_id

    @staticmethod
    async def _fetch_client(client_id: str) -> Any:
        async with AsyncSessionLocal() as session:
            client = await database_client_service.get(session, client_id)
        if client is None:
            raise ValueError(f"DatabaseClient {client_id!r} not found")
        if client.status != "active":
            raise ValueError(
                f"DatabaseClient {client_id!r} is not active "
                f"(status={client.status!r})"
            )
        return client

    @staticmethod
    def _sqlglot_guard(sql: str) -> None:
        """Defense-in-depth: ensure the compiled SQL is a SELECT statement.

        The compiler is already deterministic and only constructs SELECTs from
        validated IR, but this guard catches any future bug that could leak
        DML/DDL through.
        """
        try:
            parsed = sqlglot.parse_one(sql, read="postgres")
        except sqlglot.errors.ParseError as e:
            raise ValueError(f"compiled SQL failed to parse: {e}") from e
        if not isinstance(parsed, exp.Select):
            raise ValueError(
                f"compiled SQL is not a SELECT (got {type(parsed).__name__})"
            )
        forbidden = (exp.Insert, exp.Update, exp.Delete, exp.Drop, exp.Alter)
        for node in parsed.find_all(forbidden):
            raise ValueError(
                f"compiled SQL contains forbidden DML/DDL: {type(node).__name__}"
            )

    @staticmethod
    def _run_sync(
        client_id: str, db_type: str, creds: dict, compiled: CompiledSql
    ) -> tuple[list[str], list[dict]]:
        engine = user_engine_cache.get_engine(client_id, db_type, creds)
        if engine is not None:
            # Pooled, reused engine (postgres-like). Read-only + statement_timeout
            # are set once per physical connection (connect event in UserEngineCache),
            # so no per-query SET round-trips and no dispose β€” the connection returns
            # to the pool warm for the next query.
            with engine.connect() as conn:
                result = conn.execute(text(compiled.sql), compiled.params)
                return list(result.keys()), [dict(row) for row in result.mappings()]

        # TRIPWIRE (F-4, 2026-07-23). Below this line was the legacy per-call path for
        # non-postgres db_types, whose own comment conceded "these never set
        # read-only/timeout before, so behavior is unchanged". That means such a source
        # would get only four of the five documented defense layers: IR validation, the
        # compiler whitelist, the sqlglot guard and LIMIT β€” but NO read-only session and
        # NO statement_timeout. `CLAUDE.md` Β§2.5 states all five unconditionally; in
        # truth they were conditional on db_type, and nothing said so where a reader
        # would look.
        #
        # Zero blast radius today: Go's `database_clients.Service.Create` gates on
        # `isSupportedActive`, and only `postgres` is `active` β€” mysql/sqlserver/
        # bigquery/snowflake are all "Coming soon", so no such source can be registered.
        # This refuses loudly the day someone flips that flag, instead of silently
        # executing against a customer's database with two guardrails missing.
        #
        # Note the compiler is built with dialect="postgres" regardless of db_type and
        # the sqlglot guard parses with read="postgres", so these queries would fail on
        # a parse error anyway β€” which the never-throw path would degrade into "data not
        # available", masquerading as a data problem. The danger is someone fixing THAT
        # without noticing the pooling branch. Re-enabling this path requires
        # dialect-correct compilation AND session hardening, not just a dialect string.
        raise ValueError(
            f"source type {db_type!r} is not supported for analysis yet β€” only "
            "PostgreSQL sources can be queried safely (read-only session and query "
            "timeout are not yet implemented for other database types)"
        )
        # with db_pipeline_service.engine_scope(db_type, creds) as eng:
        #     with eng.connect() as conn:
        #         result = conn.execute(text(compiled.sql), compiled.params)
        #         return list(result.keys()), [dict(row) for row in result.mappings()]

    # ------------------------------------------------------------------
    # Speculative pre-connect (DB3)
    # ------------------------------------------------------------------

    @classmethod
    async def prewarm(cls, catalog: Catalog, user_id: str) -> None:
        """Best-effort: warm pooled engines for the catalog's schema sources.

        Called at slow-path entry so the TCP+TLS+auth handshake overlaps the ~4s
        Planner LLM call β€” by the time `retrieve_data` runs, the connection is
        already established. Warming is an optimization, never a requirement, so
        this never raises and per-source failures are swallowed.
        """
        for source in catalog.sources:
            if source.source_type != "schema":
                continue
            try:
                client_id = cls._parse_client_id(source.location_ref)
                client = await cls._fetch_client(client_id)
                if client.user_id != user_id:
                    continue
                creds = decrypt_credentials_dict(client.credentials)
                await asyncio.to_thread(cls._warm_sync, client_id, client.db_type, creds)
            except Exception as exc:  # noqa: BLE001 β€” best-effort warming
                # repr, not str: empty-str exceptions (e.g. Fernet InvalidToken)
                # would otherwise log as error="".
                logger.info("prewarm skipped", source_id=source.source_id, error=repr(exc))

    @staticmethod
    def _warm_sync(client_id: str, db_type: str, creds: dict) -> None:
        engine = user_engine_cache.get_engine(client_id, db_type, creds)
        if engine is not None:
            # Open + return a pooled physical connection: forces the handshake and
            # runs the connect-event session SETs, leaving the pool warm.
            with engine.connect():
                pass