Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
102 changes: 102 additions & 0 deletions src/dataworkbench/datacatalogue.py
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Comment thread
Anupma110 marked this conversation as resolved.

logger.info(
"Resolving base table for view %s via source_dataset_id %s",
view_name,
source_dataset_id,
)
Comment thread
Copilot marked this conversation as resolved.

# 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"
)
Comment on lines +246 to +248

def _rollback_write(self, folder_id: uuid.UUID) -> None:
"""
Delete table from storage to rollback changes when an operation fails.
Expand Down
119 changes: 119 additions & 0 deletions tests/test_datacatalogue.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"]

Loading