diff --git a/sqlmesh/core/engine_adapter/fabric.py b/sqlmesh/core/engine_adapter/fabric.py index 7b2f1acd73..2e92abb722 100644 --- a/sqlmesh/core/engine_adapter/fabric.py +++ b/sqlmesh/core/engine_adapter/fabric.py @@ -8,9 +8,11 @@ from sqlglot import exp from tenacity import retry, stop_after_attempt, wait_exponential, retry_if_result from sqlmesh.core.engine_adapter.mssql import MSSQLEngineAdapter +from sqlmesh.core.dialect import to_schema from sqlmesh.core.engine_adapter.shared import ( CommentCreationTable, CommentCreationView, + DataObject, InsertOverwriteStrategy, ) from sqlmesh.utils.errors import SQLMeshError @@ -18,7 +20,6 @@ from sqlmesh.core.schema_diff import TableAlterOperation from sqlmesh.utils import random_id - logger = logging.getLogger(__name__) @@ -85,6 +86,14 @@ def _catalog_state_label(self, catalog_name: t.Optional[str]) -> str: or "" ) + def _resolved_catalog(self) -> t.Optional[str]: + return ( + self.get_current_catalog() + or self._normalize_catalog(self._connected_catalog) + or self._default_catalog + or self._extra_config.get("database") + ) + @property def api_client(self) -> FabricHttpClient: # the requests Session is not guaranteed to be threadsafe @@ -224,6 +233,31 @@ def set_current_catalog(self, catalog_name: t.Optional[str]) -> None: self._target_catalog = target_catalog + def get_data_objects( + self, + schema_name: t.Union[str, exp.Table], + object_names: t.Optional[t.Set[str]] = None, + safe_to_cache: bool = False, + ) -> t.List[DataObject]: + # Fabric uses None as "default catalog" so we skip reconnects. Other engines + # return a real warehouse name here. + # Cache on schema.table due to @set_catalog stripping the catalog. Then put a + # warehouse name on the returned objects so listing matches other engines. + objects = super().get_data_objects(schema_name, object_names, safe_to_cache) + catalog = to_schema(schema_name).catalog or self._resolved_catalog() + if not catalog: + return objects + return [ + DataObject( + catalog=obj.catalog or catalog, + schema=obj.schema_name, + name=obj.name, + type=obj.type, + clustering_key=obj.clustering_key, + ) + for obj in objects + ] + def alter_table( self, alter_expressions: t.Union[t.List[exp.Alter], t.List[TableAlterOperation]] ) -> None: diff --git a/tests/core/engine_adapter/test_fabric.py b/tests/core/engine_adapter/test_fabric.py index d16e973e8a..2d738f2a67 100644 --- a/tests/core/engine_adapter/test_fabric.py +++ b/tests/core/engine_adapter/test_fabric.py @@ -9,7 +9,7 @@ from sqlmesh.core.engine_adapter import FabricEngineAdapter from tests.core.engine_adapter import to_sql_calls -from sqlmesh.core.engine_adapter.shared import DataObject +from sqlmesh.core.engine_adapter.shared import DataObject, DataObjectType pytestmark = [pytest.mark.engine, pytest.mark.fabric] @@ -19,6 +19,18 @@ def adapter(make_mocked_engine_adapter: t.Callable) -> FabricEngineAdapter: return make_mocked_engine_adapter(FabricEngineAdapter) +def _record_catalogs_at_execute(adapter: FabricEngineAdapter) -> t.List[t.Optional[str]]: + catalogs: t.List[t.Optional[str]] = [] + real_execute = adapter.execute + + def execute(*args: t.Any, **kwargs: t.Any) -> None: + catalogs.append(adapter._resolved_catalog()) + return real_execute(*args, **kwargs) + + adapter.execute = execute # type: ignore[method-assign] + return catalogs + + def test_get_current_catalog_uses_only_explicit_target_catalog( make_mocked_engine_adapter: t.Callable, ): @@ -451,3 +463,121 @@ def test_comments(make_mocked_engine_adapter: t.Callable, mocker: MockerFixture) create_table_comment_mock.assert_not_called() create_column_comments_mock.assert_not_called() assert to_sql_calls(adapter) == [] + + +def test_get_data_objects_cache_hits_for_default_catalog( + make_mocked_engine_adapter: t.Callable, + mocker: MockerFixture, +) -> None: + """Listing fills catalog on the return value. Cache must still hit.""" + adapter = make_mocked_engine_adapter( + FabricEngineAdapter, + default_catalog="ci_abc", + database="ci_abc", + patch_get_data_objects=False, + ) + fetchdf = mocker.patch.object( + adapter, + "fetchdf", + return_value=pd.DataFrame([{"name": "t", "schema_name": "dbo", "type": "TABLE"}]), + ) + + first = adapter.get_data_objects("ci_abc.dbo", {"t"}, safe_to_cache=True) + second = adapter.get_data_objects("ci_abc.dbo", {"t"}, safe_to_cache=True) + + assert first[0].catalog == "ci_abc" + assert second[0].catalog == "ci_abc" + assert fetchdf.call_count == 1 + + +def test_get_data_objects_labels_connected_warehouse_after_lazy_restore( + make_mocked_engine_adapter: t.Callable, + mocker: MockerFixture, +) -> None: + """Logical catalog is None; the connection is still on planning. + + Unqualified list must be planning. ci_abc.dbo must still switch. + """ + adapter = make_mocked_engine_adapter( + FabricEngineAdapter, + default_catalog="ci_abc", + database="ci_abc", + patch_get_data_objects=False, + ) + adapter.set_current_catalog("planning") + adapter.set_current_catalog(None) + assert adapter.get_current_catalog() is None + + fetchdf = mocker.patch.object( + adapter, + "fetchdf", + return_value=pd.DataFrame([{"name": "t", "schema_name": "dbo", "type": "TABLE"}]), + ) + + objects = adapter.get_data_objects("dbo") + + assert objects == [ + DataObject( + catalog="planning", + schema="dbo", + name="t", + type=DataObjectType.TABLE, + ) + ] + + assert adapter.get_data_objects("ci_abc.dbo") == [ + DataObject( + catalog="ci_abc", + schema="dbo", + name="t", + type=DataObjectType.TABLE, + ) + ] + assert fetchdf.call_count == 2 + + +def test_drop_data_object_default_catalog_drops_without_requalifying( + make_mocked_engine_adapter: t.Callable, + mocker: MockerFixture, +) -> None: + """DROP SQL has no warehouse. Prove it ran on ci_abc, without reconnecting.""" + adapter = make_mocked_engine_adapter( + FabricEngineAdapter, + default_catalog="ci_abc", + database="ci_abc", + ) + close = mocker.spy(adapter._connection_pool, "close") + catalogs_at_execute = _record_catalogs_at_execute(adapter) + + adapter.drop_data_object( + DataObject(catalog="ci_abc", schema="dbo", name="v", type=DataObjectType.VIEW) + ) + + close.assert_not_called() + assert catalogs_at_execute == ["ci_abc"] + assert to_sql_calls(adapter) == ["DROP VIEW IF EXISTS [dbo].[v];"] + + +def test_drop_data_object_default_catalog_reconnects_when_connected_elsewhere( + make_mocked_engine_adapter: t.Callable, + mocker: MockerFixture, +) -> None: + """DROP SQL has no warehouse. Prove it ran on ci_abc, then restored planning.""" + adapter = make_mocked_engine_adapter( + FabricEngineAdapter, + default_catalog="ci_abc", + database="ci_abc", + ) + adapter.set_current_catalog("planning") + assert adapter.get_current_catalog() == "planning" + close = mocker.spy(adapter._connection_pool, "close") + catalogs_at_execute = _record_catalogs_at_execute(adapter) + + adapter.drop_data_object( + DataObject(catalog="ci_abc", schema="dbo", name="v", type=DataObjectType.VIEW) + ) + + assert catalogs_at_execute == ["ci_abc"] + assert adapter.get_current_catalog() == "planning" + assert close.call_count == 2 + assert "DROP VIEW IF EXISTS [dbo].[v];" in to_sql_calls(adapter)