diff --git a/src/dataworkbench/datacatalogue.py b/src/dataworkbench/datacatalogue.py index d739e3f..e6f2176 100644 --- a/src/dataworkbench/datacatalogue.py +++ b/src/dataworkbench/datacatalogue.py @@ -145,6 +145,108 @@ def save( except Exception as e: return {"error": str(e), "error_type": type(e).__name__} + def resolve_base_databricks_full_table_name(self, view_name: str) -> str: + """ + Resolve the base table a write-shared view ultimately reads from. + + A view shared with write access carries a ``source_dataset_id`` tag holding the id + of the root dataset. This method reads that tag and resolves the Unity Catalog + object tagged ``dataset_id`` with the same value, which must be an external table. + + Args: + view_name: Fully qualified view name, ``catalog.schema.view``. Must be a + plain view; a materialized view is rejected. + + Returns: + str: Fully qualified name of the base external table, with each identifier + backtick quoted so the result can be used directly in Spark SQL + + Raises: + TypeError: If view_name is not a non-empty string + ValueError: If the view is not valid, does not carry a ``source_dataset_id`` + tag, or its base object is missing or is not an external table + + Example: + >>> catalogue = DataCatalogue() + >>> catalogue.resolve_base_databricks_full_table_name( + ... "receiver_cat.default.shared_view" + ... ) + '`source_cat`.`default`.`sales_2024`' + """ + if not isinstance(view_name, str) or not view_name: + raise TypeError("view_name must be a non-empty string") + + parts = [part.strip().strip("`") for part in view_name.split(".")] + if len(parts) != 3 or not all(parts): + raise ValueError("View is not valid") + + catalog, schema, view = parts + spark = self.storage.spark + view_args = {"catalog": catalog, "schema": schema, "view": view} + + view_rows = spark.sql( + """ + SELECT table_type FROM system.information_schema.tables + WHERE table_catalog = :catalog + AND table_schema = :schema + AND table_name = :view + """, + args=view_args, + ).collect() + + if not view_rows or view_rows[0]["table_type"] != "VIEW": + raise ValueError("View is not valid") + + tag_rows = spark.sql( + """ + SELECT tag_value FROM system.information_schema.table_tags + WHERE catalog_name = :catalog + AND schema_name = :schema + AND table_name = :view + AND tag_name = 'source_dataset_id' + """, + args=view_args, + ).collect() + + source_dataset_id = tag_rows[0]["tag_value"] if tag_rows else None + if not source_dataset_id: + raise ValueError("this view doesn't have share with Write access on it") + + logger.info( + "Resolving base table for view %s via source_dataset_id %s", + view_name, + source_dataset_id, + ) + + # The base table lives in the source workspace's catalog, so this lookup is + # deliberately not scoped to the catalog the view was found in. + base_rows = spark.sql( + """ + SELECT t.table_catalog, t.table_schema, t.table_name, t.table_type + FROM system.information_schema.table_tags tg + JOIN system.information_schema.tables t + ON t.table_catalog = tg.catalog_name + AND t.table_schema = tg.schema_name + AND t.table_name = tg.table_name + WHERE tg.tag_name = 'dataset_id' + AND tg.tag_value = :dataset_id + """, + args={"dataset_id": source_dataset_id}, + ).collect() + + if not base_rows: + raise ValueError("no base table found for this view") + + for row in base_rows: + if row["table_type"] == "EXTERNAL": + parts = (row["table_catalog"], row["table_schema"], row["table_name"]) + # Unity Catalog escapes a backtick inside an identifier by doubling it. + return ".".join(f"`{part.replace('`', '``')}`" for part in parts) + + raise ValueError( + "the base for this view is not an external table. Invalid viewName given as input" + ) + def _rollback_write(self, folder_id: uuid.UUID) -> None: """ Delete table from storage to rollback changes when an operation fails. diff --git a/tests/test_datacatalogue.py b/tests/test_datacatalogue.py index ad8b7d6..b7ceb7d 100644 --- a/tests/test_datacatalogue.py +++ b/tests/test_datacatalogue.py @@ -170,3 +170,122 @@ def test_rollback_write_delete_fails_logs_error(storage_handler): storage_handler._rollback_write(folder_id) storage_handler.storage.delete.assert_called_once_with(target_path, recursive=True) + + +VIEW_NAME = "receiver_cat.default.shared_view" +SOURCE_DATASET_ID = "11111111-1111-1111-1111-111111111111" +VIEW_ROWS = [{"table_type": "VIEW"}] +TAG_ROWS = [{"tag_value": SOURCE_DATASET_ID}] + + +def base_rows(table_type): + return [{ + "table_catalog": "source_cat", + "table_schema": "default", + "table_name": "sales", + "table_type": table_type, + }] + + +def spark_returns(*result_sets): + """One mocked spark.sql(...).collect() result per successive call.""" + return [MagicMock(collect=MagicMock(return_value=rows)) for rows in result_sets] + + +@pytest.fixture +def resolver(): + """DataCatalogue with a mocked Spark session for base table resolution.""" + with patch("dataworkbench.auth.TokenManager.get_token", return_value="mock_token"): + catalogue = DataCatalogue() + catalogue.storage = MagicMock() + return catalogue + + +@pytest.mark.parametrize("view_name", ["", 123, None]) +def test_resolve_base_table_invalid_view_name_type(resolver, view_name): + with pytest.raises(TypeError): + resolver.resolve_base_databricks_full_table_name(view_name) + + +@pytest.mark.parametrize("view_name", ["shared_view", "default.shared_view", "a.b.c.d", "cat..view"]) +def test_resolve_base_table_not_fully_qualified(resolver, view_name): + with pytest.raises(ValueError, match="View is not valid"): + resolver.resolve_base_databricks_full_table_name(view_name) + + +@pytest.mark.parametrize("view_rows", [[], [{"table_type": "MATERIALIZED_VIEW"}], [{"table_type": "EXTERNAL"}]]) +def test_resolve_base_table_input_must_be_a_plain_view(resolver, view_rows): + resolver.storage.spark.sql.side_effect = spark_returns(view_rows) + + with pytest.raises(ValueError, match="View is not valid"): + resolver.resolve_base_databricks_full_table_name(VIEW_NAME) + + +def test_resolve_base_table_without_source_dataset_id_tag(resolver): + resolver.storage.spark.sql.side_effect = spark_returns(VIEW_ROWS, []) + + with pytest.raises(ValueError, match="doesn't have share with Write access on it"): + resolver.resolve_base_databricks_full_table_name(VIEW_NAME) + + +def test_resolve_base_table_no_base_table_found(resolver): + resolver.storage.spark.sql.side_effect = spark_returns(VIEW_ROWS, TAG_ROWS, []) + + with pytest.raises(ValueError, match="no base table found for this view"): + resolver.resolve_base_databricks_full_table_name(VIEW_NAME) + + +@pytest.mark.parametrize("table_type", ["VIEW", "MATERIALIZED_VIEW", "MANAGED", "STREAMING_TABLE"]) +def test_resolve_base_table_base_is_not_external(resolver, table_type): + resolver.storage.spark.sql.side_effect = spark_returns(VIEW_ROWS, TAG_ROWS, base_rows(table_type)) + + with pytest.raises(ValueError, match="is not an external table"): + resolver.resolve_base_databricks_full_table_name(VIEW_NAME) + + +@pytest.mark.parametrize("view_name", [VIEW_NAME, " `receiver_cat`.`default`.`shared_view` "]) +def test_resolve_base_table_returns_external_table_full_name(resolver, view_name): + resolver.storage.spark.sql.side_effect = spark_returns(VIEW_ROWS, TAG_ROWS, base_rows("EXTERNAL")) + + result = resolver.resolve_base_databricks_full_table_name(view_name) + + assert result == "`source_cat`.`default`.`sales`" + + +def test_resolve_base_table_escapes_backticks_in_identifiers(resolver): + quirky = [{ + "table_catalog": "source_cat", + "table_schema": "default", + "table_name": "we`ird", + "table_type": "EXTERNAL", + }] + resolver.storage.spark.sql.side_effect = spark_returns(VIEW_ROWS, TAG_ROWS, quirky) + + result = resolver.resolve_base_databricks_full_table_name(VIEW_NAME) + + assert result == "`source_cat`.`default`.`we``ird`" + + +def test_resolve_base_table_never_interpolates_the_view_name(resolver): + resolver.storage.spark.sql.side_effect = spark_returns(VIEW_ROWS, TAG_ROWS, base_rows("EXTERNAL")) + + resolver.resolve_base_databricks_full_table_name(VIEW_NAME) + + for call in resolver.storage.spark.sql.call_args_list: + query = call.args[0] + assert "receiver_cat" not in query + assert SOURCE_DATASET_ID not in query + assert call.kwargs["args"] + + +def test_resolve_base_table_passes_identifiers_to_args_unescaped(resolver): + # Spark binds args as literals, so a quote must reach it verbatim -- escaping it + # the way an interpolated WHERE clause would need is what breaks the match. + resolver.storage.spark.sql.side_effect = spark_returns(VIEW_ROWS, TAG_ROWS, base_rows("EXTERNAL")) + + resolver.resolve_base_databricks_full_table_name("receiver_cat.default.o'brien_view") + + args = resolver.storage.spark.sql.call_args_list[0].kwargs["args"] + assert args["view"] == "o'brien_view" + assert "\\'" not in args["view"] +