diff --git a/.github/workflows/schema-hierarchy.yml b/.github/workflows/schema-hierarchy.yml new file mode 100644 index 00000000..077bc6fe --- /dev/null +++ b/.github/workflows/schema-hierarchy.yml @@ -0,0 +1,32 @@ +name: Schema hierarchy integration + +on: + pull_request: + paths: + - 'sqlit/**' + - 'tests/integration/schema_hierarchy/**' + - 'tools/run_schema_provider_integration.py' + - '.github/workflows/schema-hierarchy.yml' + workflow_dispatch: + +permissions: + contents: read + +jobs: + providers: + runs-on: ubuntu-latest + timeout-minutes: 20 + steps: + - uses: actions/checkout@v4 + - uses: astral-sh/setup-uv@v5 + - name: Install integration dependencies + run: uv sync --group test --extra postgres --extra mssql --extra snowflake + - name: Run all supported providers and registry fallback contracts + run: uv run --no-sync python tools/run_schema_provider_integration.py --output /tmp/schema-hierarchy-evidence + - name: Keep test reports and actual app captures + if: always() + uses: actions/upload-artifact@v4 + with: + name: schema-hierarchy-evidence + path: /tmp/schema-hierarchy-evidence + if-no-files-found: error diff --git a/sqlit/domains/connections/app/mock_provider.py b/sqlit/domains/connections/app/mock_provider.py index 11970026..86035814 100644 --- a/sqlit/domains/connections/app/mock_provider.py +++ b/sqlit/domains/connections/app/mock_provider.py @@ -95,6 +95,7 @@ def _display_info(config: ConnectionConfig) -> str: default_schema=str(getattr(adapter, "default_schema", "")), system_databases=frozenset(getattr(adapter, "system_databases", frozenset())), supports_foreign_keys=bool(getattr(adapter, "supports_foreign_keys", False)), + supports_schema_grouping=bool(getattr(adapter, "supports_schema_grouping", False)), ) def apply_database_override(config: ConnectionConfig, database: str | None) -> ConnectionConfig: diff --git a/sqlit/domains/connections/providers/adapter_provider.py b/sqlit/domains/connections/providers/adapter_provider.py index 1bfe19eb..42e5822c 100644 --- a/sqlit/domains/connections/providers/adapter_provider.py +++ b/sqlit/domains/connections/providers/adapter_provider.py @@ -93,6 +93,7 @@ def build_adapter_provider(spec: ProviderSpec, schema: ConnectionSchema, adapter default_schema=str(getattr(adapter, "default_schema", "")), system_databases=frozenset(getattr(adapter, "system_databases", frozenset())), supports_foreign_keys=bool(getattr(adapter, "supports_foreign_keys", False)), + supports_schema_grouping=bool(getattr(adapter, "supports_schema_grouping", False)), ) def display_info(config: ConnectionConfig) -> str: diff --git a/sqlit/domains/connections/providers/adapters/base.py b/sqlit/domains/connections/providers/adapters/base.py index d7258cdc..353a6dcf 100644 --- a/sqlit/domains/connections/providers/adapters/base.py +++ b/sqlit/domains/connections/providers/adapters/base.py @@ -117,6 +117,7 @@ class IndexInfo: name: str table_name: str is_unique: bool = False + schema: str = "" @dataclass @@ -125,6 +126,7 @@ class TriggerInfo: name: str table_name: str + schema: str = "" @dataclass @@ -132,6 +134,7 @@ class SequenceInfo: """Information about a database sequence.""" name: str + schema: str = "" @dataclass(frozen=True) diff --git a/sqlit/domains/connections/providers/model.py b/sqlit/domains/connections/providers/model.py index 00899e77..ab490840 100644 --- a/sqlit/domains/connections/providers/model.py +++ b/sqlit/domains/connections/providers/model.py @@ -39,6 +39,7 @@ class SchemaCapabilities: default_schema: str system_databases: frozenset[str] supports_foreign_keys: bool = False + supports_schema_grouping: bool = False @runtime_checkable diff --git a/sqlit/domains/connections/providers/mssql/adapter.py b/sqlit/domains/connections/providers/mssql/adapter.py index 57fa08e5..05bc2197 100644 --- a/sqlit/domains/connections/providers/mssql/adapter.py +++ b/sqlit/domains/connections/providers/mssql/adapter.py @@ -510,14 +510,26 @@ def get_columns( ) return [ColumnInfo(name=row[0], data_type=row[1], is_primary_key=row[0] in pk_columns) for row in cursor.fetchall()] + supports_schema_grouping = True + + def get_schemas(self, conn: Any, database: str | None = None) -> list[str]: + cursor = self._get_cursor_for_database(conn, database) + cursor.execute( + "SELECT name FROM sys.schemas WHERE schema_id < 16384 " + "AND name NOT IN ('guest', 'sys', 'INFORMATION_SCHEMA') ORDER BY name" + ) + return [row[0] for row in cursor.fetchall()] + def get_procedures(self, conn: Any, database: str | None = None) -> list[str]: """Get stored procedures from SQL Server.""" cursor = self._get_cursor_for_database(conn, database) cursor.execute( - "SELECT ROUTINE_NAME FROM INFORMATION_SCHEMA.ROUTINES " + "SELECT ROUTINE_NAME, ROUTINE_SCHEMA FROM INFORMATION_SCHEMA.ROUTINES " "WHERE ROUTINE_TYPE = 'PROCEDURE' ORDER BY ROUTINE_NAME" ) - return [row[0] for row in cursor.fetchall()] + from sqlit.domains.connections.providers.adapters.base import RoutineInfo + + return [RoutineInfo(row[0], schema=row[1]) for row in cursor.fetchall()] def get_completion_routines( self, conn: Any, database: str | None = None @@ -560,30 +572,30 @@ def get_indexes(self, conn: Any, database: str | None = None) -> list[IndexInfo] """Get indexes from SQL Server.""" cursor = self._get_cursor_for_database(conn, database) cursor.execute( - "SELECT i.name, t.name, i.is_unique " + "SELECT i.name, t.name, i.is_unique, SCHEMA_NAME(t.schema_id) " "FROM sys.indexes i " "JOIN sys.tables t ON i.object_id = t.object_id " "WHERE i.name IS NOT NULL AND i.type > 0 AND i.is_primary_key = 0 " "ORDER BY t.name, i.name" ) - return [IndexInfo(name=row[0], table_name=row[1], is_unique=row[2]) for row in cursor.fetchall()] + return [IndexInfo(name=row[0], table_name=row[1], is_unique=row[2], schema=row[3]) for row in cursor.fetchall()] def get_triggers(self, conn: Any, database: str | None = None) -> list[TriggerInfo]: """Get triggers from SQL Server.""" cursor = self._get_cursor_for_database(conn, database) cursor.execute( - "SELECT tr.name, OBJECT_NAME(tr.parent_id) " + "SELECT tr.name, OBJECT_NAME(tr.parent_id), OBJECT_SCHEMA_NAME(tr.parent_id) " "FROM sys.triggers tr " "WHERE tr.is_ms_shipped = 0 AND tr.parent_id > 0 " "ORDER BY OBJECT_NAME(tr.parent_id), tr.name" ) - return [TriggerInfo(name=row[0], table_name=row[1] or "") for row in cursor.fetchall()] + return [TriggerInfo(name=row[0], table_name=row[1] or "", schema=row[2]) for row in cursor.fetchall()] def get_sequences(self, conn: Any, database: str | None = None) -> list[SequenceInfo]: """Get sequences from SQL Server (2012+).""" cursor = self._get_cursor_for_database(conn, database) - cursor.execute("SELECT name FROM sys.sequences ORDER BY name") - return [SequenceInfo(name=row[0]) for row in cursor.fetchall()] + cursor.execute("SELECT name, SCHEMA_NAME(schema_id) FROM sys.sequences ORDER BY name") + return [SequenceInfo(name=row[0], schema=row[1]) for row in cursor.fetchall()] def get_foreign_keys( self, @@ -665,7 +677,7 @@ def get_referencing_foreign_keys( ] def get_index_definition( - self, conn: Any, index_name: str, table_name: str, database: str | None = None + self, conn: Any, index_name: str, table_name: str, database: str | None = None, schema: str | None = None ) -> dict[str, Any]: """Get detailed information about a SQL Server index.""" cursor = self._get_cursor_for_database(conn, database) @@ -676,8 +688,9 @@ def get_index_definition( "JOIN sys.columns c ON ic.object_id = c.object_id AND ic.column_id = c.column_id " "JOIN sys.tables t ON i.object_id = t.object_id " "WHERE i.name = ? AND t.name = ? " - "ORDER BY ic.key_ordinal", - (index_name, table_name), + + ("AND SCHEMA_NAME(t.schema_id) = ? " if schema is not None else "") + + "ORDER BY ic.key_ordinal", + (index_name, table_name) + ((schema,) if schema is not None else ()), ) rows = cursor.fetchall() is_unique = rows[0][0] if rows else False @@ -692,12 +705,14 @@ def get_index_definition( "type": index_type, "definition": ( f"CREATE {'UNIQUE ' if is_unique else ''}{index_type} INDEX " - f"[{index_name}] ON [{table_name}] ({', '.join(f'[{c}]' for c in columns)})" + f"{self.quote_identifier(index_name)} ON " + f"{self.quote_identifier(schema) + '.' if schema is not None else ''}" + f"{self.quote_identifier(table_name)} ({', '.join(self.quote_identifier(c) for c in columns)})" ), } def get_trigger_definition( - self, conn: Any, trigger_name: str, table_name: str, database: str | None = None + self, conn: Any, trigger_name: str, table_name: str, database: str | None = None, schema: str | None = None ) -> dict[str, Any]: """Get detailed information about a SQL Server trigger.""" cursor = self._get_cursor_for_database(conn, database) @@ -707,8 +722,9 @@ def get_trigger_definition( " ELSE 'AFTER' END as timing " "FROM sys.triggers tr " "JOIN sys.tables t ON tr.parent_id = t.object_id " - "WHERE tr.name = ? AND t.name = ?", - (trigger_name, table_name), + "WHERE tr.name = ? AND t.name = ?" + + (" AND SCHEMA_NAME(t.schema_id) = ?" if schema is not None else ""), + (trigger_name, table_name) + ((schema,) if schema is not None else ()), ) row = cursor.fetchone() if row: @@ -742,15 +758,16 @@ def get_trigger_definition( } def get_sequence_definition( - self, conn: Any, sequence_name: str, database: str | None = None + self, conn: Any, sequence_name: str, database: str | None = None, schema: str | None = None ) -> dict[str, Any]: """Get detailed information about a SQL Server sequence.""" cursor = self._get_cursor_for_database(conn, database) cursor.execute( "SELECT CAST(start_value AS BIGINT), CAST(increment AS BIGINT), " "CAST(minimum_value AS BIGINT), CAST(maximum_value AS BIGINT), is_cycling " - "FROM sys.sequences WHERE name = ?", - (sequence_name,), + "FROM sys.sequences WHERE name = ?" + + (" AND SCHEMA_NAME(schema_id) = ?" if schema is not None else ""), + (sequence_name,) + ((schema,) if schema is not None else ()), ) row = cursor.fetchone() if row: diff --git a/sqlit/domains/connections/providers/postgresql/adapter.py b/sqlit/domains/connections/providers/postgresql/adapter.py index b2a542c6..9da657c2 100644 --- a/sqlit/domains/connections/providers/postgresql/adapter.py +++ b/sqlit/domains/connections/providers/postgresql/adapter.py @@ -145,12 +145,26 @@ def get_databases(self, conn: Any) -> list[str]: cursor.execute("SELECT datname FROM pg_database " "WHERE datistemplate = false ORDER BY datname") return [row[0] for row in cursor.fetchall()] + supports_schema_grouping = True + + def get_schemas(self, conn: Any, database: str | None = None) -> list[str]: + cursor = conn.cursor() + cursor.execute( + "SELECT schema_name FROM information_schema.schemata " + "WHERE schema_name NOT IN ('pg_catalog', 'information_schema') " + "AND schema_name NOT LIKE 'pg_toast%' AND schema_name NOT LIKE 'pg_temp_%' " + "ORDER BY schema_name" + ) + return [row[0] for row in cursor.fetchall()] + def get_procedures(self, conn: Any, database: str | None = None) -> list[str]: """Get stored procedures/functions from PostgreSQL.""" cursor = conn.cursor() cursor.execute( - "SELECT routine_name FROM information_schema.routines " - "WHERE routine_schema = 'public' AND routine_type = 'FUNCTION' " + "SELECT routine_name, routine_schema FROM information_schema.routines " + "WHERE routine_schema NOT IN ('pg_catalog', 'information_schema') " "ORDER BY routine_name" ) - return [row[0] for row in cursor.fetchall()] + from sqlit.domains.connections.providers.adapters.base import RoutineInfo + + return [RoutineInfo(row[0], schema=row[1]) for row in cursor.fetchall()] diff --git a/sqlit/domains/connections/providers/postgresql/base.py b/sqlit/domains/connections/providers/postgresql/base.py index 88303225..e903fb25 100644 --- a/sqlit/domains/connections/providers/postgresql/base.py +++ b/sqlit/domains/connections/providers/postgresql/base.py @@ -155,13 +155,13 @@ def get_indexes(self, conn: Any, database: str | None = None) -> list[IndexInfo] cursor = conn.cursor() cursor.execute( "SELECT indexname, tablename, " - " CASE WHEN indexdef LIKE '%UNIQUE%' THEN true ELSE false END as is_unique " + " CASE WHEN indexdef LIKE '%UNIQUE%' THEN true ELSE false END as is_unique, schemaname " "FROM pg_indexes " "WHERE schemaname NOT IN ('pg_catalog', 'information_schema') " "ORDER BY tablename, indexname" ) return [ - IndexInfo(name=row[0], table_name=row[1], is_unique=row[2]) + IndexInfo(name=row[0], table_name=row[1], is_unique=row[2], schema=row[3]) for row in cursor.fetchall() ] @@ -169,7 +169,7 @@ def get_triggers(self, conn: Any, database: str | None = None) -> list[TriggerIn """Get triggers from PostgreSQL.""" cursor = conn.cursor() cursor.execute( - "SELECT trigger_name, event_object_table " + "SELECT trigger_name, event_object_table, trigger_schema " "FROM information_schema.triggers " "WHERE trigger_schema NOT IN ('pg_catalog', 'information_schema') " "ORDER BY event_object_table, trigger_name" @@ -178,25 +178,25 @@ def get_triggers(self, conn: Any, database: str | None = None) -> list[TriggerIn seen = set() results = [] for row in cursor.fetchall(): - key = (row[0], row[1]) + key = (row[0], row[1], row[2]) if key not in seen: seen.add(key) - results.append(TriggerInfo(name=row[0], table_name=row[1])) + results.append(TriggerInfo(name=row[0], table_name=row[1], schema=row[2])) return results def get_sequences(self, conn: Any, database: str | None = None) -> list[SequenceInfo]: """Get sequences from PostgreSQL.""" cursor = conn.cursor() cursor.execute( - "SELECT sequence_name " + "SELECT sequence_name, sequence_schema " "FROM information_schema.sequences " "WHERE sequence_schema NOT IN ('pg_catalog', 'information_schema') " "ORDER BY sequence_name" ) - return [SequenceInfo(name=row[0]) for row in cursor.fetchall()] + return [SequenceInfo(name=row[0], schema=row[1]) for row in cursor.fetchall()] def get_index_definition( - self, conn: Any, index_name: str, table_name: str, database: str | None = None + self, conn: Any, index_name: str, table_name: str, database: str | None = None, schema: str | None = None ) -> dict[str, Any]: """Get detailed information about a PostgreSQL index.""" cursor = conn.cursor() @@ -204,8 +204,9 @@ def get_index_definition( "SELECT indexdef, " " CASE WHEN indexdef LIKE '%%UNIQUE%%' THEN true ELSE false END as is_unique " "FROM pg_indexes " - "WHERE indexname = %s AND tablename = %s", - (index_name, table_name), + "WHERE indexname = %s AND tablename = %s" + + (" AND schemaname = %s" if schema is not None else ""), + (index_name, table_name) + ((schema,) if schema is not None else ()), ) row = cursor.fetchone() if row: @@ -225,7 +226,7 @@ def get_index_definition( } def get_trigger_definition( - self, conn: Any, trigger_name: str, table_name: str, database: str | None = None + self, conn: Any, trigger_name: str, table_name: str, database: str | None = None, schema: str | None = None ) -> dict[str, Any]: """Get detailed information about a PostgreSQL trigger.""" cursor = conn.cursor() @@ -233,8 +234,9 @@ def get_trigger_definition( "SELECT action_timing, event_manipulation, action_statement " "FROM information_schema.triggers " "WHERE trigger_name = %s AND event_object_table = %s " - "LIMIT 1", - (trigger_name, table_name), + + ("AND trigger_schema = %s " if schema is not None else "") + + "LIMIT 1", + (trigger_name, table_name) + ((schema,) if schema is not None else ()), ) row = cursor.fetchone() if row: @@ -244,8 +246,9 @@ def get_trigger_definition( "SELECT pg_get_triggerdef(t.oid) " "FROM pg_trigger t " "JOIN pg_class c ON t.tgrelid = c.oid " - "WHERE t.tgname = %s AND c.relname = %s", - (trigger_name, table_name), + "WHERE t.tgname = %s AND c.relname = %s" + + (" AND c.relnamespace = (SELECT oid FROM pg_namespace WHERE nspname = %s)" if schema is not None else ""), + (trigger_name, table_name) + ((schema,) if schema is not None else ()), ) def_row = cursor.fetchone() definition = def_row[0] if def_row else row[2] @@ -367,7 +370,7 @@ def get_referencing_foreign_keys( ] def get_sequence_definition( - self, conn: Any, sequence_name: str, database: str | None = None + self, conn: Any, sequence_name: str, database: str | None = None, schema: str | None = None ) -> dict[str, Any]: """Get detailed information about a PostgreSQL sequence.""" cursor = conn.cursor() @@ -375,8 +378,9 @@ def get_sequence_definition( "SELECT start_value, increment, minimum_value, maximum_value, cycle_option " "FROM information_schema.sequences " "WHERE sequence_name = %s " - "AND sequence_schema NOT IN ('pg_catalog', 'information_schema')", - (sequence_name,), + "AND sequence_schema NOT IN ('pg_catalog', 'information_schema')" + + (" AND sequence_schema = %s" if schema is not None else ""), + (sequence_name,) + ((schema,) if schema is not None else ()), ) row = cursor.fetchone() if row: diff --git a/sqlit/domains/connections/providers/schema_explorer.py b/sqlit/domains/connections/providers/schema_explorer.py new file mode 100644 index 00000000..b0e7b17d --- /dev/null +++ b/sqlit/domains/connections/providers/schema_explorer.py @@ -0,0 +1,37 @@ +"""Schema-scoped catalog loading shared by the UI and process worker.""" + +from __future__ import annotations + +from typing import Any + + +def load_schema_folder_items( + inspector: Any, + conn: Any, + database: str | None, + folder_type: str, + schema: str | None, +) -> list[Any]: + """Return existing explorer tuples, retaining exact schema ownership. + + Only providers advertising supports_schema_grouping use this path. Their + ancillary metadata must carry schema names; unqualified names are never + guessed to belong to a schema, even when table names happen to match. + """ + if folder_type == "schemas": + return list(inspector.get_schemas(conn, database)) + if schema is None: + raise ValueError("A schema is required for schema-scoped folders") + if folder_type in {"tables", "views"}: + getter = inspector.get_tables if folder_type == "tables" else inspector.get_views + kind = "table" if folder_type == "tables" else "view" + return [(kind, owner, name) for owner, name in getter(conn, database) if owner == schema] + if folder_type in {"indexes", "triggers"}: + getter = inspector.get_indexes if folder_type == "indexes" else inspector.get_triggers + kind = "index" if folder_type == "indexes" else "trigger" + return [(kind, item.name, item.table_name) for item in getter(conn, database) if item.schema == schema] + if folder_type == "sequences": + return [("sequence", item.name, "") for item in inspector.get_sequences(conn, database) if item.schema == schema] + if folder_type == "procedures": + return [("procedure", schema, str(item)) for item in inspector.get_procedures(conn, database) if getattr(item, "schema", None) == schema] + raise ValueError(f"Unsupported schema folder: {folder_type}") diff --git a/sqlit/domains/connections/providers/snowflake/adapter.py b/sqlit/domains/connections/providers/snowflake/adapter.py index 61b409da..74abf105 100644 --- a/sqlit/domains/connections/providers/snowflake/adapter.py +++ b/sqlit/domains/connections/providers/snowflake/adapter.py @@ -204,19 +204,65 @@ def build_select_query(self, table: str, limit: int, database: str | None = None schema = schema or "PUBLIC" return f'SELECT * FROM "{schema}"."{table}" LIMIT {limit}' + supports_schema_grouping = True + + def get_schemas(self, conn: Any, database: str | None = None) -> list[str]: + """List visible schemas without requiring a running warehouse. + + SHOW has a row limit, so consume every page before returning. + https://docs.snowflake.com/en/sql-reference/sql/show-schemas + """ + cursor = conn.cursor() + scope = f" {self.quote_identifier(database)}" if database else "" + page_size = 1000 + names: list[str] = [] + after: str | None = None + try: + while True: + query = f"SHOW SCHEMAS IN DATABASE{scope} LIMIT {page_size}" + if after is not None: + query += f" FROM {self.quote_literal(after)}" + cursor.execute(query) + name_index = next(i for i, column in enumerate(cursor.description) if column[0].lower() == "name") + rows = cursor.fetchall() + names.extend(row[name_index] for row in rows if row[name_index].upper() != "INFORMATION_SCHEMA") + if len(rows) < page_size: + return names + next_after = rows[-1][name_index] + if next_after == after: + raise RuntimeError("Schema catalog pagination did not advance") + after = next_after + finally: + cursor.close() + def get_procedures(self, conn: Any, database: str | None = None) -> list[str]: - """Get stored procedures.""" + """List procedures with their owning schema, excluding built-ins.""" + from sqlit.domains.connections.providers.adapters.base import RoutineInfo + cursor = conn.cursor() - db_prefix = f"{self.quote_identifier(database)}." if database else "" - sql = ( - "SELECT routine_name FROM " - f"{db_prefix}information_schema.routines " - "WHERE routine_type = 'PROCEDURE' AND routine_schema != 'INFORMATION_SCHEMA' " - "ORDER BY routine_name" - ) - cursor.execute(sql) - # deduplicate - return sorted(list({row[0] for row in cursor.fetchall()})) + scope = f" {self.quote_identifier(database)}" if database else "" + try: + cursor.execute(f"SHOW PROCEDURES IN DATABASE{scope}") + columns = {column[0].lower(): i for i, column in enumerate(cursor.description)} + rows = cursor.fetchall() + if len(rows) >= 10000: + # SHOW PROCEDURES has a hard result cap and no pagination. + # The documented view is PROCEDURES, not ANSI ROUTINES. + prefix = f"{self.quote_identifier(database)}." if database else "" + cursor.execute( + f"SELECT procedure_name, procedure_schema FROM {prefix}information_schema.procedures " + "WHERE procedure_schema != 'INFORMATION_SCHEMA' ORDER BY procedure_schema, procedure_name" + ) + pairs = set(cursor.fetchall()) + else: + pairs = { + (row[columns["name"]], row[columns["schema_name"]]) + for row in rows + if row[columns["is_builtin"]] != "Y" and row[columns["schema_name"]] + } + return [RoutineInfo(name, schema=schema) for name, schema in sorted(pairs)] + finally: + cursor.close() def get_indexes(self, conn: Any, database: str | None = None) -> list[IndexInfo]: """Get indexes.""" @@ -235,9 +281,9 @@ def get_sequences(self, conn: Any, database: str | None = None) -> list[Sequence """Get sequences.""" cursor = conn.cursor() db_prefix = f"{self.quote_identifier(database)}." if database else "" - sql = f"SELECT sequence_name FROM {db_prefix}information_schema.sequences WHERE sequence_schema != 'INFORMATION_SCHEMA'" + sql = f"SELECT sequence_name, sequence_schema FROM {db_prefix}information_schema.sequences WHERE sequence_schema != 'INFORMATION_SCHEMA'" cursor.execute(sql) - return [SequenceInfo(name=row[0]) for row in cursor.fetchall()] + return [SequenceInfo(name=row[0], schema=row[1]) for row in cursor.fetchall()] def get_foreign_keys( self, diff --git a/sqlit/domains/explorer/app/schema_service.py b/sqlit/domains/explorer/app/schema_service.py index 823ea087..31e22964 100644 --- a/sqlit/domains/explorer/app/schema_service.py +++ b/sqlit/domains/explorer/app/schema_service.py @@ -120,7 +120,7 @@ def list_referencing_foreign_keys( database, ) - def list_folder_items(self, folder_type: str, database: str | None) -> list[Any]: + def list_folder_items(self, folder_type: str, database: str | None, schema: str | None = None) -> list[Any]: inspector = self.session.provider.schema_inspector caps = self.session.provider.capabilities db_arg = self._resolve_db_arg(database) @@ -138,6 +138,19 @@ def cached(key: str, loader: Callable[[], Any], *, allow_empty: bool = True) -> obj_cache[cache_key][key] = data return data + if folder_type == "schemas" or schema is not None: + from sqlit.domains.connections.providers.schema_explorer import load_schema_folder_items + + if not caps.supports_schema_grouping: + raise ValueError("This provider does not support schema grouping") + return cached( + f"schema-folder:{schema!r}:{folder_type}", + lambda: self._run_with_retry( + lambda: load_schema_folder_items(inspector, self.session.connection, db_arg, folder_type, schema), + database, + ), + ) + if folder_type == "tables": raw_data = cached( "tables", @@ -210,32 +223,32 @@ def cached(key: str, loader: Callable[[], Any], *, allow_empty: bool = True) -> return [] return [] - def get_index_definition(self, database: str | None, name: str, table_name: str) -> dict[str, Any] | None: + def get_index_definition(self, database: str | None, name: str, table_name: str, schema: str | None = None) -> dict[str, Any] | None: inspector = self.session.provider.schema_inspector if not isinstance(inspector, IndexInspector): return None db_arg = self._resolve_db_arg(database) return self._run_with_retry( - lambda: inspector.get_index_definition(self.session.connection, name, table_name, db_arg), + lambda: inspector.get_index_definition(self.session.connection, name, table_name, db_arg, **({"schema": schema} if schema is not None else {})), database, ) - def get_trigger_definition(self, database: str | None, name: str, table_name: str) -> dict[str, Any] | None: + def get_trigger_definition(self, database: str | None, name: str, table_name: str, schema: str | None = None) -> dict[str, Any] | None: inspector = self.session.provider.schema_inspector if not isinstance(inspector, TriggerInspector): return None db_arg = self._resolve_db_arg(database) return self._run_with_retry( - lambda: inspector.get_trigger_definition(self.session.connection, name, table_name, db_arg), + lambda: inspector.get_trigger_definition(self.session.connection, name, table_name, db_arg, **({"schema": schema} if schema is not None else {})), database, ) - def get_sequence_definition(self, database: str | None, name: str) -> dict[str, Any] | None: + def get_sequence_definition(self, database: str | None, name: str, schema: str | None = None) -> dict[str, Any] | None: inspector = self.session.provider.schema_inspector if not isinstance(inspector, SequenceInspector): return None db_arg = self._resolve_db_arg(database) return self._run_with_retry( - lambda: inspector.get_sequence_definition(self.session.connection, name, db_arg), + lambda: inspector.get_sequence_definition(self.session.connection, name, db_arg, **({"schema": schema} if schema is not None else {})), database, ) diff --git a/sqlit/domains/explorer/domain/tree_nodes.py b/sqlit/domains/explorer/domain/tree_nodes.py index d50c10dd..6e2e706d 100644 --- a/sqlit/domains/explorer/domain/tree_nodes.py +++ b/sqlit/domains/explorer/domain/tree_nodes.py @@ -67,6 +67,7 @@ class FolderNode: folder_type: str # "databases", "tables", "views", "indexes", "triggers", "sequences", "procedures" database: str | None = None + schema: str | None = None def get_label_text(self) -> str: return self.folder_type @@ -84,7 +85,7 @@ class SchemaNode: database: str | None schema: str - folder_type: str + folder_type: str = "" def get_label_text(self) -> str: return self.schema @@ -139,6 +140,8 @@ class ProcedureNode: database: str | None name: str + schema: str | None = None + def get_label_text(self) -> str: return self.name @@ -157,6 +160,8 @@ class IndexNode: name: str table_name: str + schema: str | None = None + def get_label_text(self) -> str: return self.name @@ -175,6 +180,8 @@ class TriggerNode: name: str table_name: str + schema: str | None = None + def get_label_text(self) -> str: return self.name @@ -192,6 +199,8 @@ class SequenceNode: database: str | None name: str + schema: str | None = None + def get_label_text(self) -> str: return self.name diff --git a/sqlit/domains/explorer/ui/mixins/tree.py b/sqlit/domains/explorer/ui/mixins/tree.py index 5e682134..65bb5046 100644 --- a/sqlit/domains/explorer/ui/mixins/tree.py +++ b/sqlit/domains/explorer/ui/mixins/tree.py @@ -95,7 +95,12 @@ def on_tree_node_expanded(self: TreeMixinHost, event: Tree.NodeExpanded) -> None if self._get_node_kind(node) == "database": self._collapse_other_database_nodes(node) - self._ensure_database_connection_async(data.name) + def populate_database() -> None: + if not list(node.children): + tree_builder.add_database_object_nodes(self, node, data.name) + + self._ensure_database_connection_async(data.name, populate_database) + return if self._get_node_kind(node) == "connection": config = getattr(data, "config", None) diff --git a/sqlit/domains/explorer/ui/tree/builder.py b/sqlit/domains/explorer/ui/tree/builder.py index b6c59a8c..0ca0195b 100644 --- a/sqlit/domains/explorer/ui/tree/builder.py +++ b/sqlit/domains/explorer/ui/tree/builder.py @@ -11,6 +11,7 @@ from sqlit.domains.explorer.domain.tree_nodes import ( ConnectionFolderNode, ConnectionNode, + DatabaseNode, FolderNode, SavedQueryFileNode, SavedQueryFolderNode, @@ -721,18 +722,29 @@ def add_saved_query_nodes(host: TreeMixinHost, parent_node: Any) -> None: ) -def add_database_object_nodes(host: TreeMixinHost, parent_node: Any, database: str | None) -> None: +def add_database_object_nodes(host: TreeMixinHost, parent_node: Any, database: str | None, schema: str | None = None) -> None: """Add Tables, Views, Indexes, Triggers, Sequences, and Stored Procedures nodes.""" if not host.current_provider: return caps = host.current_provider.capabilities node_provider = host.current_provider.explorer_nodes + settings = getattr(getattr(host, "services", None), "settings_store", None) + if schema is None and caps.supports_schema_grouping and settings and settings.get("explorer_hierarchy") == "schema": + from . import loaders + + # Database branches load schemas only when opened. The connected + # single-database root is already active and can load immediately. + if isinstance(parent_node.data, DatabaseNode) and not parent_node.is_expanded: + return + loaders.add_loading_placeholder(host, parent_node) + loaders.load_folder_async(host, parent_node, FolderNode(folder_type="schemas", database=database)) + return for folder in node_provider.get_root_folders(caps): if folder.requires(caps): folder_node = parent_node.add(escape_markup(folder.label)) - folder_node.data = FolderNode(folder_type=folder.kind, database=database) + folder_node.data = FolderNode(folder_type=folder.kind, database=database, schema=schema) folder_node.allow_expand = True else: parent_node.add_leaf(f"[dim]{folder.label} (Not available)[/]") diff --git a/sqlit/domains/explorer/ui/tree/loaders.py b/sqlit/domains/explorer/ui/tree/loaders.py index 1340e712..d546c5d3 100644 --- a/sqlit/domains/explorer/ui/tree/loaders.py +++ b/sqlit/domains/explorer/ui/tree/loaders.py @@ -14,6 +14,7 @@ LoadingNode, ProcedureNode, SequenceNode, + SchemaNode, TableNode, TriggerNode, ViewNode, @@ -199,6 +200,32 @@ def load_folder_async(host: TreeMixinHost, node: Any, data: FolderNode) -> None: """Spawn worker to load folder contents (tables/views/indexes/triggers/sequences/procedures).""" folder_type = data.folder_type db_name = data.database + schema = data.schema + session = host._session + refresh_token = getattr(host, "_tree_refresh_token", None) + tokens = getattr(host, "_folder_load_tokens", None) + if tokens is None: + tokens = {} + setattr(host, "_folder_load_tokens", tokens) + token = object() + tokens[id(node)] = token + + def deliver(callback: Any) -> None: + # A refresh, connection switch, or second load owns the new tree. Late + # metadata must never append into it or clear its loading indicator. + if tokens.get(id(node)) is not token: + return + tokens.pop(id(node), None) + if host._session is not session or getattr(host, "_tree_refresh_token", None) is not refresh_token: + return + current = node + while current.parent is not None: + if current not in current.parent.children: + return + current = current.parent + if current is not host.object_tree.root: + return + callback() async def work_async() -> None: import asyncio @@ -227,6 +254,7 @@ async def work_async() -> None: config=host.current_config, database=db_name, folder_type=folder_type, + **({"schema": schema} if schema is not None else {}), ) if getattr(outcome, "cancelled", False): return @@ -241,17 +269,18 @@ async def work_async() -> None: schema_service.list_folder_items, folder_type, db_name, + *([schema] if schema is not None else []), ) host.set_timer( MIN_TIMER_DELAY_S, - lambda: on_folder_loaded(host, node, db_name, folder_type, items), + lambda: deliver(lambda: on_folder_loaded(host, node, db_name, folder_type, items)), ) except Exception as error: error_message = f"Error loading: {error}" host.set_timer( MIN_TIMER_DELAY_S, - lambda: on_tree_load_error(host, node, error_message), + lambda: deliver(lambda: on_tree_load_error(host, node, error_message)), ) host.run_worker(work_async(), name=f"load-folder-{folder_type}", exclusive=False) @@ -272,6 +301,18 @@ def on_folder_loaded( empty_child.data = LoadingNode() return + if folder_type == "schemas": + default = provider.capabilities.default_schema + for schema_name in sorted(set(items), key=lambda name: (name != default, name.casefold(), name)): + schema_node = node.add(f"\\[{escape_markup(schema_name)}]") + schema_node.data = SchemaNode(database=db_name, schema=schema_name) + schema_node.allow_expand = True + tree_builder.add_database_object_nodes(host, schema_node, db_name, schema_name) + expansion_state.restore_subtree_expansion_with_paths(host, node, getattr(host, "_expanded_paths", set())) + ensure_expanded_nodes_loaded(host, node) + tree_builder.restore_pending_cursor(host) + return + if folder_type == "databases": active_db = None if hasattr(host, "_get_effective_database"): @@ -300,21 +341,22 @@ def on_folder_loaded( tree_builder.restore_pending_cursor(host) return + schema = getattr(node.data, "schema", None) for item in items: if item[0] == "procedure": child = node.add_leaf(escape_markup(item[2])) - child.data = ProcedureNode(database=db_name, name=item[2]) + child.data = ProcedureNode(database=db_name, name=item[2], schema=schema) elif item[0] == "index": display = f"{escape_markup(item[1])} [dim]({escape_markup(item[2])})[/]" child = node.add_leaf(display) - child.data = IndexNode(database=db_name, name=item[1], table_name=item[2]) + child.data = IndexNode(database=db_name, name=item[1], table_name=item[2], schema=schema) elif item[0] == "trigger": display = f"{escape_markup(item[1])} [dim]({escape_markup(item[2])})[/]" child = node.add_leaf(display) - child.data = TriggerNode(database=db_name, name=item[1], table_name=item[2]) + child.data = TriggerNode(database=db_name, name=item[1], table_name=item[2], schema=schema) elif item[0] == "sequence": child = node.add_leaf(escape_markup(item[1])) - child.data = SequenceNode(database=db_name, name=item[1]) + child.data = SequenceNode(database=db_name, name=item[1], schema=schema) tree_builder.restore_pending_cursor(host) diff --git a/sqlit/domains/explorer/ui/tree/object_info.py b/sqlit/domains/explorer/ui/tree/object_info.py index 808f2304..88d57002 100644 --- a/sqlit/domains/explorer/ui/tree/object_info.py +++ b/sqlit/domains/explorer/ui/tree/object_info.py @@ -15,7 +15,7 @@ def show_index_info(host: TreeMixinHost, data: IndexNode) -> None: return try: - info = schema_service.get_index_definition(data.database, data.name, data.table_name) + info = schema_service.get_index_definition(data.database, data.name, data.table_name, **({"schema": data.schema} if data.schema is not None else {})) if info is None: host.notify("Indexes not supported for this database.", severity="warning") return @@ -31,7 +31,7 @@ def show_trigger_info(host: TreeMixinHost, data: TriggerNode) -> None: return try: - info = schema_service.get_trigger_definition(data.database, data.name, data.table_name) + info = schema_service.get_trigger_definition(data.database, data.name, data.table_name, **({"schema": data.schema} if data.schema is not None else {})) if info is None: host.notify("Triggers not supported for this database.", severity="warning") return @@ -47,7 +47,7 @@ def show_sequence_info(host: TreeMixinHost, data: SequenceNode) -> None: return try: - info = schema_service.get_sequence_definition(data.database, data.name) + info = schema_service.get_sequence_definition(data.database, data.name, **({"schema": data.schema} if data.schema is not None else {})) if info is None: host.notify("Sequences not supported for this database.", severity="warning") return diff --git a/sqlit/domains/explorer/ui/tree/schema_render.py b/sqlit/domains/explorer/ui/tree/schema_render.py index 84c75ee0..344a6350 100644 --- a/sqlit/domains/explorer/ui/tree/schema_render.py +++ b/sqlit/domains/explorer/ui/tree/schema_render.py @@ -38,12 +38,13 @@ def schema_sort_key(schema: str) -> tuple[int, str]: schema_nodes: dict[str, Any] = {} items_to_add: list[tuple[Any, str, str, str]] = [] expanded_paths = getattr(host, "_expanded_paths", set()) + schema_scoped = getattr(getattr(node, "data", None), "schema", None) is not None for schema in sorted_schemas: schema_items = by_schema[schema] is_default = not schema or schema == default_schema - if is_default and not has_multiple_schemas: + if schema_scoped or (is_default and not has_multiple_schemas): parent = node else: if schema not in schema_nodes: diff --git a/sqlit/domains/process_worker/app/process_worker.py b/sqlit/domains/process_worker/app/process_worker.py index 2339a6c3..6dfe926d 100644 --- a/sqlit/domains/process_worker/app/process_worker.py +++ b/sqlit/domains/process_worker/app/process_worker.py @@ -377,6 +377,7 @@ def _start_schema_folder_items(self, message: dict[str, Any]) -> None: ) return database = message.get("database") + schema = message.get("schema") config_payload = message.get("config", {}) config = ConnectionConfig.from_dict(config_payload) config = normalize_connection_config(config) @@ -431,7 +432,13 @@ def run() -> None: pass inspector = provider.schema_inspector items: list[Any] = [] - if folder_type == "tables": + if folder_type == "schemas" or schema is not None: + from sqlit.domains.connections.providers.schema_explorer import load_schema_folder_items + + if not caps.supports_schema_grouping: + raise ValueError("This provider does not support schema grouping") + items = load_schema_folder_items(inspector, conn, db_arg, folder_type, schema) + elif folder_type == "tables": raw_data = inspector.get_tables(conn, db_arg) items = [("table", schema, name) for schema, name in raw_data] elif folder_type == "views": diff --git a/sqlit/domains/process_worker/app/process_worker_client.py b/sqlit/domains/process_worker/app/process_worker_client.py index 3b45e710..8944ea22 100644 --- a/sqlit/domains/process_worker/app/process_worker_client.py +++ b/sqlit/domains/process_worker/app/process_worker_client.py @@ -206,6 +206,7 @@ def list_folder_items( config: ConnectionConfig, database: str | None, folder_type: str, + schema: str | None = None, ) -> ProcessFolderOutcome: with self._execute_lock: if self._closed: @@ -223,6 +224,7 @@ def list_folder_items( "db_type": config.db_type, "database": database, "folder_type": folder_type, + "schema": schema, } self._send(payload) diff --git a/sqlit/domains/shell/app/commands/__init__.py b/sqlit/domains/shell/app/commands/__init__.py index ad201b1e..c217bb9a 100644 --- a/sqlit/domains/shell/app/commands/__init__.py +++ b/sqlit/domains/shell/app/commands/__init__.py @@ -6,6 +6,7 @@ from . import alert as _alert from . import credentials as _credentials from . import debug as _debug +from . import explorer as _explorer from . import watchdog as _watchdog from . import worker as _worker diff --git a/sqlit/domains/shell/app/commands/explorer.py b/sqlit/domains/shell/app/commands/explorer.py new file mode 100644 index 00000000..5853c9a4 --- /dev/null +++ b/sqlit/domains/shell/app/commands/explorer.py @@ -0,0 +1,75 @@ +"""Persistent explorer layout selection.""" + +from __future__ import annotations + +from typing import Any + +from .router import register_command_handler + + +def apply_explorer_layout(app: Any, layout: str) -> None: + if layout not in {"schema", "type"}: + raise ValueError("Explorer layout must be 'schema' or 'type'") + if getattr(app, "_tree_filter_visible", False): + app.action_tree_filter_close() + store = app.services.settings_store + settings = store.load_all() + previous = settings.get("explorer_hierarchy", "type") + if previous not in {"schema", "type"}: + previous = "type" + if previous != layout: + # Keep both layouts' expansion paths, since their parent order differs. + raw_paths = settings.get("explorer_expanded_by_hierarchy", {}) + paths = dict(raw_paths) if isinstance(raw_paths, dict) else {} + paths[previous] = sorted(app._expanded_paths) + restored = paths.get(layout, []) + app._expanded_paths = {path for path in restored if isinstance(path, str)} if isinstance(restored, list) else set() + settings["explorer_expanded_by_hierarchy"] = paths + settings["expanded_nodes"] = sorted(app._expanded_paths) + settings["explorer_hierarchy"] = layout + store.save_all(settings) + if previous != layout: + # Old paths cannot identify a cursor after the hierarchy is inverted. + # Start at its connection; same-layout refresh still restores exact paths. + node = app.object_tree.cursor_node + while node is not None and app._get_node_kind(node) != "connection": + node = node.parent + if node is not None: + app.object_tree.move_cursor(node) + app._loading_nodes.clear() + app.refresh_tree() + provider = app.current_provider + label = "schema" if layout == "schema" else "object type" + message = f"Explorer grouped by {label}" + if layout == "schema" and provider and not provider.capabilities.supports_schema_grouping: + message = "Schema layout saved; this connection keeps its existing layout" + app.notify(message) + + +def show_explorer_layout(app: Any) -> None: + from sqlit.domains.shell.ui.screens.explorer_layout import ExplorerLayoutScreen + + current = app.services.settings_store.get("explorer_hierarchy", "type") + provider = app.current_provider + supported = provider.capabilities.supports_schema_grouping if provider else None + + def selected(layout: str | None) -> None: + if layout: + apply_explorer_layout(app, layout) + + app.push_screen(ExplorerLayoutScreen(current, supported=supported), selected) + + +def _handle_explorer_command(app: Any, cmd: str, args: list[str]) -> bool: + if cmd != "explorer": + return False + if not args: + show_explorer_layout(app) + elif len(args) == 1 and args[0].lower() in {"schema", "type"}: + apply_explorer_layout(app, args[0].lower()) + else: + app.notify("Usage: :explorer [schema|type]", severity="warning") + return True + + +register_command_handler(_handle_explorer_command) diff --git a/sqlit/domains/shell/app/main.py b/sqlit/domains/shell/app/main.py index 2f85d91b..250910f4 100644 --- a/sqlit/domains/shell/app/main.py +++ b/sqlit/domains/shell/app/main.py @@ -841,6 +841,7 @@ def _show_command_list(self) -> None: "Show process worker status", "Displays worker mode, active state, and last activity.", ), + ("Explorer", ":explorer [schema|type]", "Choose explorer hierarchy", "No argument opens the layout picker; choice is saved."), ("Settings", ":set process_worker_warm on|off", "Warm worker on idle", ""), ("Settings", ":set process_worker_lazy on|off", "Lazy worker start", ""), ("Settings", ":set process_worker_auto_shutdown ", "Auto-shutdown worker", ""), diff --git a/sqlit/domains/shell/state/machine.py b/sqlit/domains/shell/state/machine.py index cea1cb8c..b7b8342d 100644 --- a/sqlit/domains/shell/state/machine.py +++ b/sqlit/domains/shell/state/machine.py @@ -392,6 +392,7 @@ def lk(action: str, menu: str, fallback: str) -> str: # SETTINGS s = HelpSection(id="settings", title="SETTINGS") + s.binding(":explorer [schema|type]", "Explorer hierarchy") s.binding(":alert off|delete|write", "Confirm risky queries") s.binding(":set ln on|off|relative", "Line numbers") sections.append(s) diff --git a/sqlit/domains/shell/ui/screens/explorer_layout.py b/sqlit/domains/shell/ui/screens/explorer_layout.py new file mode 100644 index 00000000..3161e224 --- /dev/null +++ b/sqlit/domains/shell/ui/screens/explorer_layout.py @@ -0,0 +1,49 @@ +"""Choose the explorer hierarchy without editing a configuration file.""" + +from textual.app import ComposeResult +from textual.binding import Binding +from textual.screen import ModalScreen +from textual.widgets import OptionList, Static +from textual.widgets.option_list import Option + +from sqlit.shared.ui.widgets import Dialog + + +class ExplorerLayoutScreen(ModalScreen[str | None]): + BINDINGS = [Binding("escape", "cancel", "Cancel", priority=True)] + CSS = """ + ExplorerLayoutScreen { align: center middle; background: transparent; } + #explorer-layout-dialog { width: 72; height: auto; } + #explorer-layout-description { margin: 0 1 1 1; color: $text-muted; } + #explorer-layout-options { height: 8; border: none; padding: 0 1; } + #explorer-layout-note { margin: 1; color: $text-muted; } + """ + + def __init__(self, current: str, *, supported: bool | None): + super().__init__() + self.current = current + self.supported = supported + + def compose(self) -> ComposeResult: + with Dialog(id="explorer-layout-dialog", title="Explorer Layout", shortcuts=[("Apply", ""), ("Cancel", "")]): + yield Static("Choose how you browse database objects. Your choice is saved.", id="explorer-layout-description") + yield OptionList( + Option("By object type (default)\nDatabase → Tables / Views → Schema\nCompare the same object type across schemas.", id="type"), + Option("By schema\nDatabase → Schema → Tables / Views\nKeep one schema's related objects together.", id="schema"), + id="explorer-layout-options", + ) + note = "SQLite and providers without schema grouping keep their existing layout." + if self.supported is False: + note = "This connection keeps its existing layout. The saved choice applies to supported databases." + yield Static(note, id="explorer-layout-note") + + def on_mount(self) -> None: + options = self.query_one(OptionList) + options.highlighted = 1 if self.current == "schema" else 0 + options.focus() + + def on_option_list_option_selected(self, event: OptionList.OptionSelected) -> None: + self.dismiss(event.option.id) + + def action_cancel(self) -> None: + self.dismiss(None) diff --git a/tests/integration/schema_hierarchy/README.md b/tests/integration/schema_hierarchy/README.md new file mode 100644 index 00000000..034b8965 --- /dev/null +++ b/tests/integration/schema_hierarchy/README.md @@ -0,0 +1,47 @@ +# Schema hierarchy provider integration + +Run the entire lane from the feature checkout: + +```sh +uv sync --group test --extra postgres --extra mssql --extra snowflake +uv run --no-sync python tools/run_schema_provider_integration.py --output /tmp/schema-hierarchy-evidence +``` + +The runner requires all five provider cases and zero skips. Missing drivers, +failed container readiness, missing app screenshots, and test failures make the +lane fail. Normal unit/UI runs do not start containers; integration is opt-in. +The dedicated GitHub Actions workflow runs this same command and uploads evidence. + +| Provider | Execution | Scope | +| --- | --- | --- | +| PostgreSQL | Real PostgreSQL 16, disposable Docker container | Schemas including empty schemas, tables, views, routines, indexes, triggers, sequences, queries, UI, refresh, worker | +| SQL Server | Real SQL Server 2022, disposable Docker container | Same contract; includes copied index DDL retaining schema | +| Snowflake | fakesnow with DuckDB storage | Schema listing, tables, views, queries and app UI; not live Snowflake evidence, empty procedure-folder loading; nonempty routines/sequences not covered by this emulator | +| Supabase | Real Supabase adapter, psycopg2 destination redirected to disposable PostgreSQL | Adapter endpoint construction, schemas, objects, queries and UI; not SaaS authentication/pooler evidence | +| SQLite | Temporary SQLite file | Existing hierarchy, queries and UI remain intact when schema-first preference is enabled | + +The separate registry test exercises every registered provider against the real +explorer builder. Providers without the capability use the existing layout. +That is contract coverage, not a live connection test for every cloud database. + +Each container is uniquely named, loopback-bound and deleted after its module. +SQL Server has 3 GiB memory; 2 GiB exited during this host's startup probe. +Passwords are generated in memory and passed through a mode-0600 temporary Docker +env file that is unlinked immediately. No existing database server is used. +Readiness requires an actual adapter connection and SELECT 1, not an open port. + +Screenshots are generated by assertions against the running Textual app, after +selecting a real fixture table and receiving its rows. The title names the evidence +level. Use the PR-review skill's `render_textual_svgs.cjs` to turn SVGs into PNGs. +Only inspected images should be published; keep the source commit in manifest.json. + +TDD evidence: the SQL Server integration test first failed because copied index +DDL omitted its schema, then passed after the adapter fix. Snowflake schema-list +integration initially hit an unimplemented SCHEMATA view in fakesnow. The adapter +now uses documented SHOW SCHEMAS (which does not require a running warehouse), +with an additional red/green unit test for named output columns and pagination: +https://docs.snowflake.com/en/sql-reference/sql/show-schemas + +An additional red/green integration test opens every advertised object folder. +It exposed Snowflake's old ANSI ROUTINES query; procedure metadata now uses +SHOW PROCEDURES, with the documented PROCEDURES view beyond SHOW's result cap. diff --git a/tests/integration/schema_hierarchy/__init__.py b/tests/integration/schema_hierarchy/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/tests/integration/schema_hierarchy/conftest.py b/tests/integration/schema_hierarchy/conftest.py new file mode 100644 index 00000000..57ec96ae --- /dev/null +++ b/tests/integration/schema_hierarchy/conftest.py @@ -0,0 +1,40 @@ +"""Explicit opt-in; missing requested drivers/services are failures, not skips.""" + +from __future__ import annotations + +import os +from pathlib import Path + +import pytest + +from .support import provider_database + +PROVIDERS = ["postgresql", "mssql", "snowflake", "supabase", "sqlite"] + + +def pytest_generate_tests(metafunc): + if "provider_case" not in metafunc.fixturenames: + return + names = os.environ.get("SQLIT_SCHEMA_PROVIDERS", ",".join(PROVIDERS)).split(",") + unknown = set(names) - set(PROVIDERS) + if unknown: + raise ValueError(f"Unknown schema integration providers: {sorted(unknown)}") + if metafunc.function.__name__ == "test_ancillary_metadata_and_definitions_are_scoped": + names = [name for name in names if name in {"postgresql", "mssql", "supabase"}] + elif metafunc.function.__name__ == "test_process_worker_metadata_matches_direct_provider": + names = [name for name in names if name in {"postgresql", "mssql"}] + metafunc.parametrize("provider_case", names, indirect=True, scope="module") + + +@pytest.fixture(scope="module") +def provider_case(request, tmp_path_factory): + if os.environ.get("SQLIT_SCHEMA_INTEGRATION") != "1": + pytest.skip("Set SQLIT_SCHEMA_INTEGRATION=1 to run disposable provider integration tests") + with provider_database(request.param, tmp_path_factory.mktemp("schema-" + request.param)) as case: + yield case + + +@pytest.fixture +def capture_dir(): + value = os.environ.get("SQLIT_SCHEMA_CAPTURE_DIR") + return Path(value) if value else None diff --git a/tests/integration/schema_hierarchy/support.py b/tests/integration/schema_hierarchy/support.py new file mode 100644 index 00000000..da3ab34b --- /dev/null +++ b/tests/integration/schema_hierarchy/support.py @@ -0,0 +1,207 @@ +"""Disposable integration databases. Never targets an existing database server.""" + +from __future__ import annotations + +import json +import secrets +import sqlite3 +import subprocess +import time +from contextlib import closing, contextmanager +from dataclasses import dataclass +from pathlib import Path +from typing import Any + +from sqlit.domains.connections.domain.config import ConnectionConfig +from sqlit.domains.connections.providers.catalog import get_provider + + +@dataclass(repr=False) +class ProviderCase: + name: str + config: ConnectionConfig + provider: Any + connection: Any + database: str | None + billing: str + analytics: str + empty: str + evidence: str + expected_rows: dict[str, list[tuple]] + + @property + def adapter(self): + return self.provider.connection_factory + + +def sql(conn, statement): + cursor = conn.cursor() + try: + cursor.execute(statement) + finally: + cursor.close() + + +def seed(case): + adapter = case.adapter + q = adapter.quote_identifier + schemas = [""] if case.name == "sqlite" else [case.billing, case.analytics] + if case.name != "sqlite": + for schema in [*schemas, case.empty]: + sql(case.connection, f"CREATE SCHEMA {q(schema)}") + for schema in schemas: + qualified = (q(schema) + ".") if schema else "" + table = qualified + q("orders") + sql(case.connection, f"CREATE TABLE {table} ({q('id')} INTEGER PRIMARY KEY, {q('customer')} VARCHAR(80), {q('total')} DECIMAL(10,2))") + rows = case.expected_rows[schema] + for ident, customer, total in rows: + sql(case.connection, f"INSERT INTO {table} VALUES ({ident}, '{customer}', {total})") + sql(case.connection, f"CREATE VIEW {qualified}{q('open_orders')} AS SELECT * FROM {table}") + if case.name in {"postgresql", "supabase", "mssql"}: + column = "customer" if schema == case.billing else "total" + sql(case.connection, f"CREATE INDEX {q('orders_lookup')} ON {table} ({q(column)})") + start = 10000 if schema == case.billing else 90000 + sql(case.connection, f"CREATE SEQUENCE {qualified}{q('invoice_number')} START WITH {start}") + if case.name in {"postgresql", "supabase"}: + sql(case.connection, f"CREATE FUNCTION {qualified}audit_fn() RETURNS trigger LANGUAGE plpgsql AS $$ BEGIN RETURN NEW; END $$") + sql(case.connection, f"CREATE TRIGGER orders_audit BEFORE INSERT ON {table} FOR EACH ROW EXECUTE FUNCTION {qualified}audit_fn()") + sql(case.connection, f"CREATE PROCEDURE {qualified}close_month() LANGUAGE plpgsql AS $$ BEGIN NULL; END $$") + else: + sql(case.connection, f"CREATE TRIGGER {qualified}orders_audit ON {table} AFTER INSERT AS BEGIN SET NOCOUNT ON; END") + sql(case.connection, f"CREATE PROCEDURE {qualified}close_month AS BEGIN SET NOCOUNT ON; END") + + +@contextmanager +def docker_database(name, root): + container = f"sqlit-schema-{name}-{secrets.token_hex(4)}" + password = "Sqlit-" + secrets.token_hex(12) + "!9" + env_file = root / f"{name}.env" + env_file.write_text("POSTGRES_PASSWORD=" + password + "\nPOSTGRES_DB=sqlit_schema_it\n" if name == "postgresql" else "MSSQL_SA_PASSWORD=" + password + "\nACCEPT_EULA=Y\n") + env_file.chmod(0o600) + image = "postgres:16-alpine" if name == "postgresql" else "mcr.microsoft.com/mssql/server:2022-latest" + internal = 5432 if name == "postgresql" else 1433 + started = False + connection = None + try: + subprocess.run(["docker", "run", "-d", "--name", container, "--cpus=2", "--memory=" + ("512m" if name == "postgresql" else "3g"), "--env-file", str(env_file), "-p", f"127.0.0.1::{internal}", image], check=True, capture_output=True) + started = True + env_file.unlink() + state = json.loads(subprocess.check_output(["docker", "inspect", container]))[0] + port = state["NetworkSettings"]["Ports"][f"{internal}/tcp"][0]["HostPort"] + provider = get_provider(name) + config = ConnectionConfig.from_dict( + dict( + name=f"{provider.metadata.display_name} integration", + db_type=name, + server="127.0.0.1", + port=port, + database="sqlit_schema_it" if name == "postgresql" else "master", + username="postgres" if name == "postgresql" else "sa", + password=password, + options={"tls_trust_server_certificate": True} if name == "mssql" else {}, + ) + ) + deadline = time.monotonic() + 120 + while True: + running = json.loads(subprocess.check_output(["docker", "inspect", container]))[0]["State"] + if not running["Running"]: + logs = subprocess.run(["docker", "logs", "--tail", "35", container], capture_output=True, text=True) + raise RuntimeError(f"{name} exited during startup: " + (logs.stdout + logs.stderr).replace(password, "[redacted]")) + try: + connection = provider.connection_factory.connect(config) + sql(connection, "SELECT 1") + break + except Exception as error: + if connection: + connection.close() + connection = None + if time.monotonic() > deadline: + raise RuntimeError(f"{name} readiness failed: " + str(error).replace(password, "[redacted]")) from None + time.sleep(1) + if name == "mssql": + connection.autocommit = True + sql(connection, "CREATE DATABASE [sqlit_schema_it]") + connection.close() + config = config.with_endpoint(database="sqlit_schema_it") + connection = provider.connection_factory.connect(config) + connection.autocommit = True + yield config, provider, connection, "sqlit_schema_it" + finally: + env_file.unlink(missing_ok=True) + if connection: + connection.close() + if started: + subprocess.run(["docker", "stop", "-t", "20", container], check=True, capture_output=True) + subprocess.run(["docker", "rm", container], check=True, capture_output=True) + + +@contextmanager +def provider_database(name, root): + root = Path(root) + root.mkdir(parents=True, exist_ok=True) + billing = "" if name == "sqlite" else "billing" + rows = {billing: [(1001, "Northwind Bikes", 1240), (1002, "Harbor Coffee", 385)]} + if name != "sqlite": + rows["analytics"] = [(2001, "Monthly revenue", 28500)] + + def case(config, provider, conn, database, evidence): + result = ProviderCase(name, config, provider, conn, database, billing, "analytics", "empty_lab", evidence, rows) + seed(result) + return result + + if name in {"postgresql", "mssql"}: + with docker_database(name, root) as (config, provider, conn, database): + yield case(config, provider, conn, database, "live PostgreSQL 16" if name == "postgresql" else "live SQL Server 2022") + elif name == "supabase": + # Run Supabase's real adapter against PostgreSQL. Replace only the + # cloud network destination; never fabricate catalog/query responses. + from unittest.mock import patch + + import psycopg2 + + with docker_database("postgresql", root) as (local_config, _, conn, database): + provider = get_provider("supabase") + config = ConnectionConfig.from_dict( + dict( + name="Supabase adapter integration", + db_type="supabase", + server="placeholder", + username="fixture", + password=local_config.tcp_endpoint.password, + options={"supabase_project_id": "fixture", "supabase_region": "us-east-1", "supabase_aws_shard": "aws-0"}, + ) + ) + real_connect = psycopg2.connect + + def local_transport(*args, **kwargs): + assert kwargs["host"] == "aws-0-us-east-1.pooler.supabase.com" + assert kwargs["user"] == "postgres.fixture" + assert kwargs["database"] == "postgres" + endpoint = local_config.tcp_endpoint + kwargs.update(host=endpoint.host, port=endpoint.port, user=endpoint.username, password=endpoint.password, database=database) + return real_connect(*args, **kwargs) + + with patch("psycopg2.connect", side_effect=local_transport): + yield case(config, provider, conn, database, "Supabase adapter / local PostgreSQL transport") + elif name == "sqlite": + path = root / "fixture.sqlite" + provider = get_provider(name) + config = ConnectionConfig.from_dict(dict(name="SQLite fallback integration", db_type="sqlite", file_path=str(path))) + with closing(sqlite3.connect(path)) as conn: + result = case(config, provider, conn, None, "local SQLite") + conn.commit() + yield result + elif name == "snowflake": + import fakesnow + + (root / "fakesnow").mkdir() + with fakesnow.patch(db_path=root / "fakesnow"): + config = ConnectionConfig.from_dict(dict(name="Snowflake emulator integration", db_type="snowflake", server="local-emulator", database="SQLIT_SCHEMA_IT", username="fixture", password="fixture", options={"schema": "PUBLIC"})) + provider = get_provider(name) + conn = provider.connection_factory.connect(config) + try: + yield case(config, provider, conn, "SQLIT_SCHEMA_IT", "Snowflake emulation (fakesnow)") + finally: + conn.close() + else: + raise ValueError(f"Unknown integration provider: {name}") diff --git a/tests/integration/schema_hierarchy/test_provider_contract.py b/tests/integration/schema_hierarchy/test_provider_contract.py new file mode 100644 index 00000000..959b7827 --- /dev/null +++ b/tests/integration/schema_hierarchy/test_provider_contract.py @@ -0,0 +1,171 @@ +"""Run real adapter/catalog/UI contracts, not fabricated tree snapshots. + +Opt in with SQLIT_SCHEMA_INTEGRATION=1. Requested providers fail on unavailable +infrastructure rather than reporting a passing, skipped integration lane. +""" + +from __future__ import annotations + +import asyncio + +import pytest + +from sqlit.domains.connections.providers.schema_explorer import load_schema_folder_items +from sqlit.domains.explorer.domain.tree_nodes import FolderNode, SchemaNode, TableNode +from sqlit.domains.shell.app.main import SSMSTUI +from sqlit.shared.app.runtime import RuntimeConfig +from tests.ui.mocks import MockConnectionStore, MockSettingsStore, build_test_services + + +def test_catalog_lists_empty_schema(provider_case): + case = provider_case + if case.name == "sqlite": + assert not case.provider.capabilities.supports_schema_grouping + return + schemas = load_schema_folder_items(case.adapter, case.connection, case.database, "schemas", None) + assert {case.billing, case.analytics, case.empty}.issubset(schemas) + + +@pytest.mark.parametrize("kind,name", [("tables", "orders"), ("views", "open_orders")]) +def test_duplicate_objects_retain_schema_identity(provider_case, kind, name): + case = provider_case + if case.name == "sqlite": + getter = case.adapter.get_tables if kind == "tables" else case.adapter.get_views + assert any(row[1] == name for row in getter(case.connection)) + return + for schema in [case.billing, case.analytics]: + items = load_schema_folder_items(case.adapter, case.connection, case.database, kind, schema) + assert (kind[:-1], schema, name) in items + assert all(item[1] == schema for item in items) + + +def test_select_query_reads_correct_same_named_table(provider_case): + case = provider_case + for schema, expected in case.expected_rows.items(): + query = case.adapter.build_select_query("orders", 100, case.database, schema) + columns, rows, truncated = case.adapter.execute_query(case.connection, query) + assert columns == ["id", "customer", "total"] + assert sorted(rows) == expected + assert not truncated + + +def test_ancillary_metadata_and_definitions_are_scoped(provider_case): + case = provider_case + for schema, column, start in [(case.billing, "customer", 10000), (case.analytics, "total", 90000)]: + for kind, expected in [("indexes", ("index", "orders_lookup", "orders")), ("triggers", ("trigger", "orders_audit", "orders")), ("sequences", ("sequence", "invoice_number", "")), ("procedures", ("procedure", schema, "close_month"))]: + items = load_schema_folder_items(case.adapter, case.connection, case.database, kind, schema) + assert expected in items + index = case.adapter.get_index_definition(case.connection, "orders_lookup", "orders", case.database, schema=schema) + assert column in index["definition"] + # TDD regression: correct metadata is insufficient if copied DDL drops + # its schema and could operate on another same-named table. + qualified = f"[{schema}].[orders]" if case.name == "mssql" else f"{schema}.orders" + assert qualified in index["definition"] + sequence = case.adapter.get_sequence_definition(case.connection, "invoice_number", case.database, schema=schema) + assert int(sequence["start_value"]) == start + trigger = case.adapter.get_trigger_definition(case.connection, "orders_audit", "orders", case.database, schema=schema) + assert case.adapter.quote_identifier(schema) in trigger["definition"] or schema in trigger["definition"] + + +async def settle(app, pilot): + for _ in range(4): + await pilot.pause(0.08) + await app.workers.wait_for_complete() + + +def walk(root): + yield root + for child in root.children: + yield from walk(child) + + +def find(app, cls, **fields): + return next(node for node in walk(app.object_tree.root) if isinstance(node.data, cls) and all(getattr(node.data, key) == value for key, value in fields.items())) + + +@pytest.mark.asyncio +async def test_ui_layout_query_refresh_and_screenshot(provider_case, tmp_path, capture_dir): + case = provider_case + settings = MockSettingsStore({"theme": "vesper", "explorer_hierarchy": "type"}) + services = build_test_services(runtime=RuntimeConfig(settings_path=tmp_path / "settings.json", process_worker=False, process_worker_warm_on_idle=False), settings_store=settings, connection_store=MockConnectionStore([case.config])) + app = SSMSTUI(services=services) + async with app.run_test(size=(140, 44)) as pilot: + await settle(app, pilot) + app.connect_to_server(case.config) + await settle(app, pilot) + assert app.current_connection is not None, getattr(app.screen, "message", "Connection failed") + app._run_command("explorer schema") + await settle(app, pilot) + assert settings.get("explorer_hierarchy") == "schema" + if case.name == "sqlite": + assert not any(isinstance(n.data, SchemaNode) for n in walk(app.object_tree.root)) + folder = find(app, FolderNode, folder_type="tables") + else: + schema = find(app, SchemaNode, schema=case.billing, folder_type="") + schema.expand() + folder = find(app, FolderNode, folder_type="tables", schema=case.billing) + folder.expand() + await settle(app, pilot) + orders = find(app, TableNode, schema=case.billing, name="orders") + orders.expand() + await settle(app, pilot) + assert orders.children + app.object_tree.move_cursor(orders) + app.action_select_table() + await settle(app, pilot) + assert app.results_table.row_count == len(case.expected_rows[case.billing]) + query = app.query_input.text + assert case.adapter.quote_identifier("orders") in query + if case.name != "sqlite": + assert case.adapter.quote_identifier(case.billing) in query + # Exercise every enabled object folder before capturing proof. + # The Snowflake emulator returns an empty procedure catalog; it + # cannot establish live routine creation/execution support. + owner = find(app, SchemaNode, schema=case.billing, folder_type="") + for child in owner.children: + if isinstance(child.data, FolderNode): + child.expand() + await settle(app, pilot) + assert not getattr(app.screen, "message", "") + app.object_tree.focus() + await pilot.pause() + if capture_dir: + capture_dir.mkdir(parents=True, exist_ok=True) + (capture_dir / f"{case.name}.svg").write_text(app.export_screenshot(title=f"sqlit · schema hierarchy · {case.evidence}")) + app._refresh_tree_common(notify=False) + await settle(app, pilot) + assert find(app, TableNode, schema=case.billing, name="orders") + assert app.query_input.text == query + app._run_command("explorer type") + await settle(app, pilot) + assert app.query_input.text == query + assert not any(isinstance(n.data, SchemaNode) and not n.data.folder_type for n in walk(app.object_tree.root)) + + +@pytest.mark.asyncio +async def test_process_worker_metadata_matches_direct_provider(provider_case): + case = provider_case + from sqlit.domains.process_worker.app.process_worker_client import ProcessWorkerClient + + client = ProcessWorkerClient() + try: + for schema in [case.billing, case.analytics]: + for kind in ["tables", "views", "indexes", "triggers", "sequences", "procedures"]: + result = await asyncio.to_thread(client.list_folder_items, config=case.config, database=case.database, folder_type=kind, schema=schema) + assert not result.error and not result.cancelled + assert result.items == load_schema_folder_items(case.adapter, case.connection, case.database, kind, schema) + finally: + client.close() + + +def test_every_advertised_object_folder_loads(provider_case): + """Opening any enabled folder must not fail on its catalog query.""" + case = provider_case + for folder in case.provider.explorer_nodes.get_root_folders(case.provider.capabilities): + if not folder.requires(case.provider.capabilities): + continue + if case.name == "sqlite": + items = case.provider.explorer_nodes.load_folder_items(case.adapter, case.provider.capabilities, case.connection, folder.kind, None) + else: + items = load_schema_folder_items(case.adapter, case.connection, case.database, folder.kind, case.billing) + assert isinstance(items, list) diff --git a/tests/integration/schema_hierarchy/test_registry_fallback.py b/tests/integration/schema_hierarchy/test_registry_fallback.py new file mode 100644 index 00000000..0341c160 --- /dev/null +++ b/tests/integration/schema_hierarchy/test_registry_fallback.py @@ -0,0 +1,34 @@ +"""Every registered provider participates in the explorer capability contract. + +This integrates the real provider registry and tree builder without connecting +to cloud accounts. It is fallback/dispatch coverage, not live database evidence. +""" + +from types import SimpleNamespace + +import pytest +from textual.widgets import Tree + +from sqlit.domains.connections.domain.config import ConnectionConfig +from sqlit.domains.connections.providers.catalog import get_provider, get_supported_db_types +from sqlit.domains.explorer.domain.tree_nodes import ConnectionNode, SchemaNode +from sqlit.domains.explorer.ui.tree import builder, loaders + + +@pytest.mark.parametrize("name", get_supported_db_types()) +def test_every_provider_routes_hierarchy_by_declared_capability(name, monkeypatch): + provider = get_provider(name) + config = ConnectionConfig.from_dict({"name": "Capability check", "db_type": name}) + tree = Tree("root") + parent = tree.root.add("Connection", data=ConnectionNode(config)) + host = SimpleNamespace(current_provider=provider, services=SimpleNamespace(settings_store={"explorer_hierarchy": "schema"})) + loaded = [] + monkeypatch.setattr(loaders, "load_folder_async", lambda *args: loaded.append(args)) + builder.add_database_object_nodes(host, parent, None) + if provider.capabilities.supports_schema_grouping: + assert name in {"postgresql", "mssql", "snowflake", "supabase"} + assert len(loaded) == 1 and loaded[0][2].folder_type == "schemas" + else: + assert not loaded + assert not any(isinstance(n.data, SchemaNode) for n in parent.children) + assert len(parent.children) == len(provider.explorer_nodes.get_root_folders(provider.capabilities)) diff --git a/tests/ui/test_explorer_layout.py b/tests/ui/test_explorer_layout.py new file mode 100644 index 00000000..fda799a7 --- /dev/null +++ b/tests/ui/test_explorer_layout.py @@ -0,0 +1,61 @@ +"""The new layout is optional, cancellable, and harmless on SQLite.""" + +import sqlite3 + +import pytest + +from sqlit.domains.connections.domain.config import ConnectionConfig +from sqlit.domains.explorer.domain.tree_nodes import FolderNode, SchemaNode +from sqlit.domains.shell.app.main import SSMSTUI +from sqlit.domains.shell.ui.screens.explorer_layout import ExplorerLayoutScreen +from sqlit.shared.app.runtime import RuntimeConfig +from tests.ui.mocks import MockConnectionStore, MockSettingsStore, build_test_services + + +@pytest.mark.asyncio +async def test_layout_picker_cancel_and_invalid_value_do_not_change_settings(): + settings = MockSettingsStore({"theme": "vesper", "explorer_hierarchy": "type"}) + app = SSMSTUI(services=build_test_services(settings_store=settings, connection_store=MockConnectionStore())) + async with app.run_test(size=(100, 35)) as pilot: + await pilot.pause() + app._run_command("explorer") + await pilot.pause() + assert isinstance(app.screen, ExplorerLayoutScreen) + await pilot.press("down", "escape") + assert settings.get("explorer_hierarchy") == "type" + app._run_command("explorer invalid") + assert settings.get("explorer_hierarchy") == "type" + app._run_command("explorer schema unexpected") + assert settings.get("explorer_hierarchy") == "type" + + +@pytest.mark.asyncio +async def test_sqlite_stays_flat_and_layout_switch_preserves_query(tmp_path): + path = tmp_path / "db.sqlite" + with sqlite3.connect(path) as db: + db.execute("CREATE TABLE orders (id INTEGER)") + config = ConnectionConfig.from_dict({"name": "SQLite demo", "db_type": "sqlite", "file_path": str(path)}) + settings = MockSettingsStore({"theme": "vesper", "explorer_hierarchy": "schema"}) + services = build_test_services(runtime=RuntimeConfig(process_worker=False), settings_store=settings, connection_store=MockConnectionStore([config])) + app = SSMSTUI(services=services) + async with app.run_test(size=(100, 35)) as pilot: + await pilot.pause() + app.connect_to_server(config) + for _ in range(3): + await pilot.pause(0.05) + await app.workers.wait_for_complete() + assert app.current_connection is not None + connection = app.object_tree.root.children[0] + assert any(isinstance(n.data, FolderNode) and n.data.folder_type == "tables" for n in connection.children) + assert not any(isinstance(n.data, SchemaNode) for n in connection.children) + app.query_input.text = "SELECT * FROM orders;" + app._run_command("explorer") + await pilot.pause() + assert app.screen.supported is False + await pilot.press("escape") + app.action_tree_filter() + app._run_command("explorer type") + await pilot.pause() + assert not app._tree_filter_visible + assert app.query_input.text == "SELECT * FROM orders;" + assert settings.get("explorer_hierarchy") == "type" diff --git a/tests/unit/test_schema_hierarchy.py b/tests/unit/test_schema_hierarchy.py new file mode 100644 index 00000000..1133ad3e --- /dev/null +++ b/tests/unit/test_schema_hierarchy.py @@ -0,0 +1,151 @@ +"""Schema ownership must survive same-named objects and both catalog paths.""" + +from types import SimpleNamespace +from unittest.mock import MagicMock + +import pytest + +from sqlit.domains.connections.providers.adapters.base import IndexInfo, RoutineInfo, SequenceInfo, TriggerInfo +from sqlit.domains.connections.providers.mssql.adapter import SQLServerAdapter +from sqlit.domains.connections.providers.postgresql.adapter import PostgreSQLAdapter +from sqlit.domains.connections.providers.schema_explorer import load_schema_folder_items +from sqlit.domains.connections.providers.snowflake.adapter import SnowflakeAdapter +from sqlit.domains.explorer.app.schema_service import ExplorerSchemaService + + +@pytest.mark.parametrize( + "kind,method,data,expected", + [ + ("tables", "get_tables", [("billing", "orders"), ("analytics", "orders")], [("table", "billing", "orders")]), + ("views", "get_views", [("billing", "summary"), ("analytics", "summary")], [("view", "billing", "summary")]), + ("indexes", "get_indexes", [IndexInfo("lookup", "orders", schema="billing"), IndexInfo("lookup", "orders", schema="analytics"), IndexInfo("unowned", "orders")], [("index", "lookup", "orders")]), + ("triggers", "get_triggers", [TriggerInfo("audit", "orders", schema="billing"), TriggerInfo("audit", "orders", schema="analytics")], [("trigger", "audit", "orders")]), + ("sequences", "get_sequences", [SequenceInfo("id", schema="billing"), SequenceInfo("id", schema="analytics")], [("sequence", "id", "")]), + ("procedures", "get_procedures", [RoutineInfo("close_month", schema="billing"), RoutineInfo("close_month", schema="analytics"), "unknown"], [("procedure", "billing", "close_month")]), + ], +) +def test_schema_scoping_never_guesses_ownership(kind, method, data, expected): + inspector = SimpleNamespace(**{method: lambda conn, database: data}) + assert load_schema_folder_items(inspector, object(), "workbench", kind, "billing") == expected + assert load_schema_folder_items(inspector, object(), "workbench", kind, "BILLING") == [] + + +def test_empty_schema_is_not_inferred_from_tables(): + inspector = SimpleNamespace(get_schemas=lambda conn, database: ["billing", "empty_lab"]) + assert load_schema_folder_items(inspector, object(), None, "schemas", None) == ["billing", "empty_lab"] + + +def test_service_cache_separates_schema_and_database(): + inspector = MagicMock() + inspector.get_tables.side_effect = lambda conn, db: [("billing", db + "_orders"), ("analytics", db + "_report")] + session = SimpleNamespace(provider=SimpleNamespace(schema_inspector=inspector, capabilities=SimpleNamespace(supports_schema_grouping=True)), connection=object()) + service = ExplorerSchemaService(session=session, object_cache={}) + service._run_with_retry = lambda fn, database: fn() + assert service.list_folder_items("tables", "one", "billing") == [("table", "billing", "one_orders")] + assert service.list_folder_items("tables", "one", "analytics") == [("table", "analytics", "one_report")] + assert service.list_folder_items("tables", "two", "billing") == [("table", "billing", "two_orders")] + service.list_folder_items("tables", "one", "billing") + assert inspector.get_tables.call_count == 3 + + +@pytest.mark.parametrize("adapter_cls", [PostgreSQLAdapter, SQLServerAdapter, SnowflakeAdapter]) +def test_advertised_providers_retain_routine_schema(adapter_cls): + adapter = adapter_cls() + conn = MagicMock() + cursor = conn.cursor.return_value + adapter._get_cursor_for_database = lambda conn, db: cursor + cursor.fetchall.return_value = [("close_month", "billing"), ("close_month", "analytics")] + if adapter_cls is SnowflakeAdapter: + cursor.description = [("name",), ("schema_name",), ("is_builtin",)] + cursor.fetchall.return_value = [("close_month", "billing", "N"), ("close_month", "analytics", "N")] + result = adapter.get_procedures(conn, "workbench") + assert sorted((item.schema, str(item)) for item in result) == [("analytics", "close_month"), ("billing", "close_month")] + assert adapter.supports_schema_grouping + + +@pytest.mark.parametrize("adapter_cls", [PostgreSQLAdapter, SQLServerAdapter]) +def test_index_schema_is_parameterized_not_interpolated(adapter_cls): + adapter = adapter_cls() + conn = MagicMock() + cursor = conn.cursor.return_value + adapter._get_cursor_for_database = lambda conn, db: cursor + cursor.fetchall.return_value = [] + cursor.fetchone.return_value = None + schema = "billing' OR 1=1 --" + adapter.get_index_definition(conn, "lookup", "orders", "db", schema=schema) + query, params = cursor.execute.call_args.args + assert schema not in query + assert params[-1] == schema + + +@pytest.mark.parametrize("adapter_cls", [PostgreSQLAdapter, SQLServerAdapter]) +def test_ancillary_metadata_carries_schema(adapter_cls): + adapter = adapter_cls() + conn = MagicMock() + cursor = conn.cursor.return_value + adapter._get_cursor_for_database = lambda conn, db: cursor + cursor.fetchall.return_value = [("same", "orders", True, "billing"), ("same", "orders", False, "analytics")] + assert [i.schema for i in adapter.get_indexes(conn)] == ["billing", "analytics"] + cursor.fetchall.return_value = [("same", "orders", "billing"), ("same", "orders", "analytics")] + assert [i.schema for i in adapter.get_triggers(conn)] == ["billing", "analytics"] + cursor.fetchall.return_value = [("same", "billing"), ("same", "analytics")] + assert [i.schema for i in adapter.get_sequences(conn)] == ["billing", "analytics"] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("stale_by", ["refresh", "session", "removed", "none"]) +async def test_late_folder_response_cannot_mutate_a_replaced_tree(monkeypatch, stale_by): + from textual.widgets import Tree + + from sqlit.domains.explorer.domain.tree_nodes import FolderNode + from sqlit.domains.explorer.ui.tree import loaders + + tree = Tree("root") + node = tree.root.add("Tables", data=FolderNode("tables", "db", "billing")) + jobs = [] + timers = [] + delivered = [] + host = SimpleNamespace( + object_tree=tree, + _session=object(), + _tree_refresh_token=object(), + services=SimpleNamespace(runtime=SimpleNamespace(process_worker=False)), + _get_schema_service=lambda: SimpleNamespace(list_folder_items=lambda *args: [("table", "billing", "orders")]), + run_worker=lambda coro, **kwargs: jobs.append(coro), + set_timer=lambda delay, callback: timers.append(callback), + ) + monkeypatch.setattr(loaders, "on_folder_loaded", lambda *args: delivered.append(args)) + loaders.load_folder_async(host, node, node.data) + await jobs[0] + if stale_by == "refresh": + host._tree_refresh_token = object() + elif stale_by == "session": + host._session = object() + elif stale_by == "removed": + node.remove() + for callback in timers: + callback() + assert len(delivered) == (1 if stale_by == "none" else 0) + assert host._folder_load_tokens == {} + + +def test_closed_database_does_not_load_schemas(monkeypatch): + from textual.widgets import Tree + + from sqlit.domains.explorer.domain.tree_nodes import DatabaseNode + from sqlit.domains.explorer.ui.tree import builder, loaders + + tree = Tree("root") + database = tree.root.add("workbench", data=DatabaseNode("workbench")) + loaded = [] + host = SimpleNamespace( + current_provider=SimpleNamespace(capabilities=SimpleNamespace(supports_schema_grouping=True), explorer_nodes=object()), + services=SimpleNamespace(settings_store={"explorer_hierarchy": "schema"}), + ) + monkeypatch.setattr(loaders, "load_folder_async", lambda *args: loaded.append(args)) + builder.add_database_object_nodes(host, database, "workbench") + assert not loaded and not database.children + database.expand() + builder.add_database_object_nodes(host, database, "workbench") + assert len(loaded) == 1 + assert loaded[0][2].folder_type == "schemas" diff --git a/tests/unit/test_snowflake_adapter.py b/tests/unit/test_snowflake_adapter.py index ff1d68d6..2d1ff4bb 100644 --- a/tests/unit/test_snowflake_adapter.py +++ b/tests/unit/test_snowflake_adapter.py @@ -182,3 +182,19 @@ def test_build_select_query(self): query = adapter.build_select_query("MY_TABLE", 10, schema="MYSCHEMA") assert query == 'SELECT * FROM "MYSCHEMA"."MY_TABLE" LIMIT 10' + + +def test_schema_listing_uses_named_show_column_and_paginates(): + """SHOW works without warehouse compute; pagination must not drop schemas.""" + from sqlit.domains.connections.providers.snowflake.adapter import SnowflakeAdapter + + conn = MagicMock() + cursor = conn.cursor.return_value + cursor.description = [("created_on",), ("name",), ("database_name",)] + first = [(None, "INFORMATION_SCHEMA", "data")] + [(None, f"schema{i:04d}", "data") for i in range(999)] + cursor.fetchall.side_effect = [first, [(None, "z'last", "data")]] + schemas = SnowflakeAdapter().get_schemas(conn, 'data"quoted') + assert schemas == [f"schema{i:04d}" for i in range(999)] + ["z'last"] + calls = [call.args[0] for call in cursor.execute.call_args_list] + assert calls[0] == 'SHOW SCHEMAS IN DATABASE "data""quoted" LIMIT 1000' + assert calls[1].endswith(" FROM 'schema0998'") diff --git a/tools/fixtures/schema_hierarchy.sql b/tools/fixtures/schema_hierarchy.sql new file mode 100644 index 00000000..f98f907c --- /dev/null +++ b/tools/fixtures/schema_hierarchy.sql @@ -0,0 +1,23 @@ +-- Synthetic review data only. Run against a disposable database. +CREATE SCHEMA billing; +CREATE SCHEMA analytics; +CREATE SCHEMA archive; +CREATE SCHEMA empty_lab; +CREATE TABLE billing.orders (id integer PRIMARY KEY, customer text NOT NULL, total numeric(10,2), status text); +INSERT INTO billing.orders VALUES (1001, 'Northwind Bikes', 1240.00, 'paid'), (1002, 'Harbor Coffee', 385.50, 'pending'), (1003, 'Pine Studio', 910.00, 'paid'); +CREATE TABLE billing.invoices (id integer PRIMARY KEY, order_id integer REFERENCES billing.orders(id), issued_on date); +INSERT INTO billing.invoices VALUES (501,1001,'2026-09-01'), (502,1003,'2026-09-02'); +CREATE TABLE analytics.orders (id integer PRIMARY KEY, report_month text, revenue numeric(10,2)); +INSERT INTO analytics.orders VALUES (1,'2026-08',28500.00), (2,'2026-09',19300.00); +CREATE TABLE archive.orders (id integer PRIMARY KEY, archived_on date); +INSERT INTO archive.orders VALUES (400,'2024-12-31'); +CREATE VIEW billing.open_orders AS SELECT * FROM billing.orders WHERE status='pending'; +CREATE VIEW analytics.monthly_revenue AS SELECT report_month, revenue FROM analytics.orders; +CREATE INDEX orders_lookup ON billing.orders(customer); +CREATE INDEX orders_lookup ON analytics.orders(report_month); +CREATE SEQUENCE billing.invoice_number START 10000; +CREATE SEQUENCE analytics.invoice_number START 90000; +CREATE FUNCTION billing.audit_order() RETURNS trigger LANGUAGE plpgsql AS $$ BEGIN RETURN NEW; END $$; +CREATE TRIGGER audit_order BEFORE INSERT ON billing.orders FOR EACH ROW EXECUTE FUNCTION billing.audit_order(); +CREATE PROCEDURE billing.close_month() LANGUAGE plpgsql AS $$ BEGIN NULL; END $$; +CREATE PROCEDURE analytics.close_month() LANGUAGE plpgsql AS $$ BEGIN NULL; END $$; diff --git a/tools/review_schema_hierarchy.py b/tools/review_schema_hierarchy.py new file mode 100644 index 00000000..e04e3d8b --- /dev/null +++ b/tools/review_schema_hierarchy.py @@ -0,0 +1,246 @@ +#!/usr/bin/env python3 +"""Capture and exercise schema hierarchy against the disposable PostgreSQL fixture. + +Create a fresh database and apply tools/fixtures/schema_hierarchy.sql first. +This script never creates or modifies database objects. Output uses only the +fixture's synthetic data. Requires the repository development dependencies. +""" + +from __future__ import annotations + +import argparse +import asyncio +import json +import os +import subprocess +import sys +import tempfile +from pathlib import Path + +ROOT = Path(__file__).resolve().parents[1] +sys.path.insert(0, str(ROOT)) + + +def arguments(): + p = argparse.ArgumentParser(description=__doc__) + p.add_argument("--port", type=int, required=True) + p.add_argument("--output", type=Path, required=True) + p.add_argument("--worker", action="store_true", help="Exercise the process worker as well") + return p.parse_args() + + +async def capture(args): + with tempfile.TemporaryDirectory(prefix="sqlit-schema-review-") as temp: + os.environ["SQLIT_CONFIG_DIR"] = temp + os.environ.pop("NO_COLOR", None) + os.environ["FORCE_COLOR"] = "1" + from sqlit.domains.connections.domain.config import ConnectionConfig + from sqlit.domains.explorer.domain.tree_nodes import ColumnNode, DatabaseNode, FolderNode, SchemaNode, SequenceNode, TableNode + from sqlit.domains.shell.app.main import SSMSTUI + from sqlit.domains.shell.store.settings import SettingsStore + from sqlit.shared.app.runtime import RuntimeConfig + from tests.ui.mocks import MockConnectionStore, build_test_services + + config = ConnectionConfig.from_dict(dict(name="Workbench", db_type="postgresql", server="127.0.0.1", port=str(args.port), database="workbench", username="postgres", password="demo")) + settings = SettingsStore(file_path=Path(temp) / "settings.json") + settings.save_all({"theme": "vesper", "explorer_hierarchy": "type"}) + services = build_test_services( + runtime=RuntimeConfig(settings_path=Path(temp) / "settings.json", process_worker=args.worker, process_worker_warm_on_idle=False), connection_store=MockConnectionStore([config]), settings_store=settings + ) + app = SSMSTUI(services=services) + args.output.mkdir(parents=True, exist_ok=True) + checks = [] + async with app.run_test(size=(140, 46)) as pilot: + + async def settle(): + for _ in range(4): + await pilot.pause(0.08) + await app.workers.wait_for_complete() + + def nodes(): + stack = [app.object_tree.root] + found = [] + while stack: + node = stack.pop() + found.append(node) + stack.extend(reversed(node.children)) + return found + + def find(cls, **fields): + return next(n for n in nodes() if isinstance(n.data, cls) and all(getattr(n.data, k) == v for k, v in fields.items())) + + async def expand(cls, **fields): + node = find(cls, **fields) + node.expand() + await settle() + return node + + async def shot(name): + await settle() + (args.output / (name + ".svg")).write_text(app.export_screenshot(title="sqlit · " + name[3:].replace("-", " "))) + print(name, flush=True) + + await settle() + app.connect_to_server(config) + await settle() + assert app.current_connection is not None, (app._last_notification, type(app.screen).__name__, getattr(app.screen, "message", None), getattr(app.screen, "error_message", None)) + await expand(FolderNode, folder_type="tables", schema=None) + await expand(FolderNode, folder_type="views", schema=None) + for n in nodes(): + if isinstance(n.data, SchemaNode) and n.data.schema == "billing": + n.expand() + await settle() + app.object_tree.focus() + await shot("01-before-object-type") + + app._run_command("explorer") + await settle() + await pilot.press("down") + await shot("02-layout-picker") + await pilot.press("enter") + await settle() + assert settings.get("explorer_hierarchy") == "schema" + billing = await expand(SchemaNode, schema="billing", folder_type="") + assert {n.data.folder_type for n in billing.children if isinstance(n.data, FolderNode)} == {"tables", "views", "indexes", "triggers", "sequences", "procedures"} + for kind in ["tables", "views", "procedures"]: + await expand(FolderNode, folder_type=kind, schema="billing") + app.object_tree.move_cursor(billing) + await shot("03-after-schema-first") + checks.append("Layout picker saves schema preference; all six object types are below billing") + + orders = await expand(TableNode, schema="billing", name="orders") + assert any(isinstance(n.data, ColumnNode) and n.data.schema == "billing" for n in orders.children) + app.object_tree.move_cursor(orders) + app.action_select_table() + await settle() + assert '"billing"."orders"' in app.query_input.text, app.query_input.text + assert app.results_table.row_count == 3 + await shot("04-select-billing-orders") + checks.append("Expanding a table loads columns; selecting billing.orders executes schema-qualified SQL and returns three fixture rows") + + if args.worker: + assert app._process_worker_client is not None + checks.append("Actual process worker started and served schema-scoped catalog and query requests") + query = app.query_input.text + await expand(SchemaNode, schema="analytics", folder_type="") + await expand(FolderNode, folder_type="tables", schema="analytics") + table_names = {(n.data.schema, n.data.name) for n in nodes() if isinstance(n.data, TableNode)} + assert ("billing", "orders") in table_names and ("analytics", "orders") in table_names + app.object_tree.focus() + app.object_tree.move_cursor(find(TableNode, schema="analytics", name="orders")) + await shot("05-duplicate-names") + checks.append("billing.orders and analytics.orders coexist with distinct schema identities") + + app.action_tree_filter() + app._tree_filter_text = "orders" + app._update_tree_filter() + await settle() + assert any(isinstance(n.data, TableNode) and n.data.schema == "analytics" for n in app._tree_filter_matches) + app.action_tree_filter_close() + await settle() + app.object_tree.move_cursor(find(TableNode, schema="analytics", name="orders")) + await settle() + checks.append("Explorer filtering finds objects inside schema-first folders and restores the tree on close") + + app._refresh_tree_common(notify=False) + await settle() + assert app.query_input.text == query + assert find(TableNode, schema="billing", name="orders") + assert find(TableNode, schema="analytics", name="orders") + assert isinstance(app.object_tree.cursor_node.data, TableNode) and app.object_tree.cursor_node.data.schema == "analytics" + checks.append("Refresh restores expanded schemas, folders, columns and the selected analytics.orders; query remains intact") + + # Inspect duplicate ancillary names using schema-qualified metadata. + for owner, expected in [("billing", "10000"), ("analytics", "90000")]: + await expand(FolderNode, folder_type="sequences", schema=owner) + seq = find(SequenceNode, schema=owner, name="invoice_number") + info = await asyncio.to_thread(app._get_schema_service().get_sequence_definition, seq.data.database, seq.data.name, owner) + assert info["start_value"] == expected + checks.append("Same-named sequences resolve to the correct schema (10000 vs 90000)") + + # Empty schemas remain discoverable, and empty folders resolve once. + for n in nodes(): + if isinstance(n.data, SchemaNode): + n.collapse() + await expand(SchemaNode, schema="empty_lab", folder_type="") + empty = await expand(FolderNode, folder_type="tables", schema="empty_lab") + assert "(Empty)" in str(empty.children[0].label) + app.object_tree.move_cursor(empty) + await shot("06-empty-schema") + checks.append("An empty schema is listed and its empty Tables folder displays an explicit empty state") + + app.action_tree_filter() + app._tree_filter_text = "empty_lab" + app._update_tree_filter() + app._run_command("explorer type") + assert not app._tree_filter_visible + await settle() + assert not any(isinstance(n.data, SchemaNode) and not n.data.folder_type for n in nodes()) + assert app.query_input.text == query + app._run_command("explorer schema") + await settle() + assert find(SchemaNode, schema="empty_lab", folder_type="").is_expanded + checks.append("Switching layouts preserves query text and each layout restores its own expanded branches") + + # A new app instance reads the setting from the persisted file. + restarted = SSMSTUI(services=services) + async with restarted.run_test(size=(140, 46)) as pilot: + await pilot.pause() + assert services.settings_store.get("explorer_hierarchy") == "schema" + restarted.connect_to_server(config) + for _ in range(4): + await pilot.pause(0.1) + await restarted.workers.wait_for_complete() + connected = restarted.object_tree.root.children[0] + assert any(isinstance(n.data, SchemaNode) for n in connected.children) + checks.append("New app instance restores the persisted schema layout") + + # Browsing a server must not eagerly enumerate every database's schemas. + multi_config = config.with_endpoint(database="") + multi_services = build_test_services(runtime=RuntimeConfig(process_worker=args.worker, process_worker_warm_on_idle=False), connection_store=MockConnectionStore([multi_config]), settings_store=settings) + multi = SSMSTUI(services=multi_services) + async with multi.run_test(size=(140, 46)) as pilot: + await pilot.pause() + multi.connect_to_server(multi_config) + for _ in range(5): + await pilot.pause(0.1) + await multi.workers.wait_for_complete() + dbs = multi.object_tree.root.children[0].children[0] + assert isinstance(dbs.data, FolderNode) and dbs.data.folder_type == "databases" + assert all(not n.children for n in dbs.children if isinstance(n.data, DatabaseNode)) + workbench = next(n for n in dbs.children if isinstance(n.data, DatabaseNode) and n.data.name == "workbench") + workbench.expand() + for _ in range(5): + await pilot.pause(0.1) + await multi.workers.wait_for_complete() + assert any(isinstance(n.data, SchemaNode) and n.data.schema == "billing" for n in workbench.children) + for db in dbs.children: + if isinstance(db.data, DatabaseNode) and db.data.name != "workbench": + assert not db.children + billing = next(n for n in workbench.children if isinstance(n.data, SchemaNode) and n.data.schema == "billing") + billing.expand() + await pilot.pause() + multi.object_tree.move_cursor(billing) + await pilot.pause() + (args.output / "07-multi-database.svg").write_text(multi.export_screenshot(title="sqlit · database → schema → object type")) + checks.append("Multi-database browsing loads schemas only for the expanded database") + + (args.output / "manifest.json").write_text( + json.dumps( + { + "commit": subprocess.check_output(["git", "-C", str(ROOT), "rev-parse", "HEAD"], text=True).strip(), + "terminal": [140, 46], + "provider": "PostgreSQL 16 in disposable Docker container", + "process_worker": args.worker, + "data": "tools/fixtures/schema_hierarchy.sql", + "checks": checks, + }, + indent=2, + ) + + "\n" + ) + print(json.dumps(checks, indent=2)) + + +if __name__ == "__main__": + asyncio.run(capture(arguments())) diff --git a/tools/run_schema_provider_integration.py b/tools/run_schema_provider_integration.py new file mode 100644 index 00000000..6ff664fd --- /dev/null +++ b/tools/run_schema_provider_integration.py @@ -0,0 +1,61 @@ +#!/usr/bin/env python3 +"""Run the opt-in provider lane; reject skipped/missing provider evidence.""" + +from __future__ import annotations + +import argparse +import json +import os +import subprocess +import sys +import tempfile +import xml.etree.ElementTree as ET +from pathlib import Path + +ROOT = Path(__file__).resolve().parents[1] +PROVIDERS = ("postgresql", "mssql", "snowflake", "supabase", "sqlite") + + +def main(): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--output", required=True, type=Path) + args = parser.parse_args() + output = args.output.resolve() + output.mkdir(parents=True, exist_ok=True) + junit = output / "junit.xml" + with tempfile.TemporaryDirectory(prefix="sqlit-schema-config-") as config: + env = os.environ.copy() + env.pop("NO_COLOR", None) + env.update(SQLIT_SCHEMA_INTEGRATION="1", SQLIT_SCHEMA_PROVIDERS=",".join(PROVIDERS), SQLIT_CONFIG_DIR=config, SQLIT_SCHEMA_CAPTURE_DIR=str(output / "screenshots")) + command = [sys.executable, "-m", "pytest", "tests/integration/schema_hierarchy", "-q", "--tb=short", "--timeout=180", "--junitxml=" + str(junit)] + result = subprocess.run(command, cwd=ROOT, env=env, stdout=subprocess.PIPE, stderr=subprocess.STDOUT, text=True) + (output / "pytest.log").write_text(result.stdout) + print(result.stdout) + if not junit.exists(): + raise SystemExit(result.returncode or "No JUnit results produced") + cases = ET.parse(junit).findall(".//testcase") # noqa: S314 - generated locally by pytest above + summary = {"tests": len(cases), "failures": sum(c.find("failure") is not None for c in cases), "errors": sum(c.find("error") is not None for c in cases), "skips": sum(c.find("skipped") is not None for c in cases)} + names = [case.attrib["name"] for case in cases if "test_provider_contract" in case.attrib.get("classname", "")] + counts = {provider: sum("[" + provider + "]" in name or "[" + provider + "-" in name for name in names) for provider in PROVIDERS} + summary["provider_tests"] = counts + summary["registry_contract_tests"] = sum("test_registry_fallback" in c.attrib.get("classname", "") for c in cases) + summary["head"] = subprocess.check_output(["git", "rev-parse", "HEAD"], cwd=ROOT, text=True).strip() + summary["evidence"] = { + "postgresql": "live PostgreSQL 16 container", + "mssql": "live SQL Server 2022 container", + "snowflake": "fakesnow emulation; not live Snowflake", + "supabase": "real adapter with local PostgreSQL transport; not live Supabase", + "sqlite": "local SQLite file", + } + (output / "manifest.json").write_text(json.dumps(summary, indent=2) + "\n") + screenshots = output / "screenshots" + screenshots.mkdir(exist_ok=True) + (screenshots / "manifest.json").write_text(json.dumps(summary, indent=2) + "\n") + print(json.dumps(summary, indent=2)) + missing = [p for p in PROVIDERS if counts[p] < 5 or not (screenshots / f"{p}.svg").exists()] + if result.returncode or summary["failures"] or summary["errors"] or summary["skips"] or missing: + raise SystemExit(f"Provider lane incomplete: returncode={result.returncode}, missing={missing}, results={summary}") + + +if __name__ == "__main__": + main()