This is an automated email from the ASF dual-hosted git repository.
shahar1 pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/airflow.git
The following commit(s) were added to refs/heads/main by this push:
new d4a47caaf09 Notify downstream Dags when COPY INTO writes a Unity
(Databricks) table (#74191)
d4a47caaf09 is described below
commit d4a47caaf09d1502ba356755ffaafb53e49de4ee
Author: Yossi Eliaz <[email protected]>
AuthorDate: Mon Oct 5 21:26:33 2026 +0300
Notify downstream Dags when COPY INTO writes a Unity (Databricks) table
(#74191)
* Publish the Unity Catalog table written by COPY INTO as an asset
Downstream Dags that schedule on a Databricks table had no way to learn
that COPY INTO refreshed it, unless every caller hand-built outlets and
kept them in sync with table_name.
DatabricksCopyIntoAssetOperator takes a static UnityTableIdentity, so the
asset is known at parse time while table_name stays templated. At run
time it refuses to execute SQL if the rendered target is not that table,
so the published asset cannot drift from the table actually written.
* Validate COPY INTO asset workspace
* Match Unity table assets regardless of name capitalization
---
providers/databricks/docs/operators/copy_into.rst | 81 ++++++++
.../providers/databricks/assets/databricks.py | 37 ++++
.../databricks/operators/databricks_sql.py | 82 ++++++--
.../unit/databricks/assets/test_databricks.py | 55 ++++++
.../databricks/operators/test_databricks_copy.py | 215 ++++++++++++++++++++-
5 files changed, 457 insertions(+), 13 deletions(-)
diff --git a/providers/databricks/docs/operators/copy_into.rst
b/providers/databricks/docs/operators/copy_into.rst
index 618730b6cb3..133d75aee63 100644
--- a/providers/databricks/docs/operators/copy_into.rst
+++ b/providers/databricks/docs/operators/copy_into.rst
@@ -50,3 +50,84 @@ An example usage of the DatabricksCopyIntoOperator to import
CSV data into a tab
:language: python
:start-after: [START howto_operator_databricks_copy_into]
:end-before: [END howto_operator_databricks_copy_into]
+
+.. _howto/operator:DatabricksCopyIntoAssetOperator:
+
+DatabricksCopyIntoAssetOperator
+===============================
+
+Use
:class:`~airflow.providers.databricks.operators.databricks_sql.DatabricksCopyIntoAssetOperator`
+when the ``COPY INTO`` target is a Unity Catalog table that downstream Dags
schedule on.
+It accepts every ``DatabricksCopyIntoOperator`` argument plus a required
``unity_table``.
+``DatabricksCopyIntoOperator`` itself declares no assets.
+
+``unity_table`` is a
:class:`~airflow.providers.databricks.assets.databricks.UnityTableIdentity`
+with ``host``, ``catalog``, ``schema``, and ``table``. All four are required
and static.
+Jinja in any field raises ``ValueError``. ``host`` follows the Databricks
connection rule, so
+``https://my-workspace.cloud.databricks.com/`` becomes
``my-workspace.cloud.databricks.com``.
+The hostname, catalog, schema, and table are normalized to lowercase because
their names
+are case-insensitive. Different capitalization therefore produces the same
asset URI.
+``unity_table.to_asset()`` returns the
``databricks://host/catalog/schema/table`` asset.
+
+.. code-block:: python
+
+ from airflow.providers.databricks.assets.databricks import
UnityTableIdentity
+ from airflow.providers.databricks.operators.databricks_sql import
DatabricksCopyIntoAssetOperator
+
+ users = UnityTableIdentity(
+ host="my-workspace.cloud.databricks.com",
+ catalog="main",
+ schema="default",
+ table="users",
+ )
+
+ load_users = DatabricksCopyIntoAssetOperator(
+ task_id="load_users",
+ sql_endpoint_name="my-endpoint",
+ file_location="/Volumes/main/default/landing/users.csv",
+ file_format="CSV",
+ table_name="main.default.users",
+ unity_table=users,
+ )
+
+Outlets
+-------
+
+When you omit ``outlets``, the operator sets
``outlets=[unity_table.to_asset()]`` at parse time.
+When you pass ``outlets``, including ``outlets=[]``, the operator keeps your
value.
+To add extra outlets, pass the Unity table asset among them, as in
+``outlets=[users.to_asset(), other_asset]``.
+
+Templated table names
+---------------------
+
+``table_name`` stays templated. ``unity_table`` is not templated and fixes the
asset at parse time.
+Before running SQL, the operator resolves the rendered ``table_name`` and
compares it with ``unity_table``.
+The table and connection hostname comparisons are case-insensitive; the SQL
keeps the supplied casing.
+
+* A three-part name ``catalog.schema.table`` is used as is.
+* A two-part name ``schema.table`` takes the catalog from the ``catalog``
argument.
+* A one-part name ``table`` takes the catalog and schema from the ``catalog``
and ``schema`` arguments.
+
+The operator never guesses the workspace default catalog or schema. If a part
is missing or the
+resolved table differs from ``unity_table``, the task raises ``ValueError``
and no SQL runs.
+
+Dynamic task mapping
+--------------------
+
+A mapped task does not run the operator constructor when the Dag is parsed, so
it gets no
+automatic outlet. Pass ``outlets=[users.to_asset()]`` in ``partial()``. All
mapped instances share
+one ``unity_table``, so each rendered ``table_name`` must still resolve to
that table.
+
+Sources and other operators
+---------------------------
+
+``file_location`` may be a Unity Catalog volume path such as
``/Volumes/main/default/landing``.
+The volume is the ``COPY INTO`` source. It is not an outlet. The outlet is the
target table.
+
+:class:`~airflow.providers.databricks.operators.databricks.DatabricksSQLStatementsOperator`
does not
+infer assets from SQL. Pass ``outlets=[users.to_asset()]`` for the tables your
statements write.
+
+Databricks job and pipeline operators, such as
+:class:`~airflow.providers.databricks.operators.databricks.DatabricksRunNowOperator`,
do not infer
+table assets either. Pass ``outlets`` explicitly.
diff --git
a/providers/databricks/src/airflow/providers/databricks/assets/databricks.py
b/providers/databricks/src/airflow/providers/databricks/assets/databricks.py
index 6424aff878c..2fef5428a32 100644
--- a/providers/databricks/src/airflow/providers/databricks/assets/databricks.py
+++ b/providers/databricks/src/airflow/providers/databricks/assets/databricks.py
@@ -17,6 +17,7 @@
from __future__ import annotations
+from dataclasses import dataclass, fields
from typing import TYPE_CHECKING
from airflow.providers.common.compat.assets import Asset
@@ -42,6 +43,42 @@ def create_asset(
return Asset(uri=f"databricks://{host}{port}/{catalog}/{schema}/{table}",
extra=extra)
+@dataclass(frozen=True, slots=True)
+class UnityTableIdentity:
+ """
+ Static identity of a Unity Catalog table in one Databricks workspace.
+
+ :param host: Workspace hostname. A URL such as
``https://xx.cloud.databricks.com`` is
+ reduced to its hostname, as the Databricks connection does.
+ :param catalog: Unity Catalog catalog name.
+ :param schema: Schema name inside ``catalog``.
+ :param table: Table name inside ``schema``.
+ """
+
+ host: str
+ catalog: str
+ schema: str
+ table: str
+
+ def __post_init__(self) -> None:
+ from airflow.providers.databricks.hooks.databricks_base import
BaseDatabricksHook
+
+ object.__setattr__(self, "host",
BaseDatabricksHook._parse_host(self.host))
+ for field in fields(self):
+ value = getattr(self, field.name)
+ if not value:
+ raise ValueError(f"UnityTableIdentity.{field.name} must not be
empty.")
+ if "{{" in value or "{%" in value:
+ raise ValueError(
+ f"UnityTableIdentity.{field.name} must be static, got
Jinja {value!r}. "
+ "It is not a template field and is never rendered."
+ )
+ object.__setattr__(self, field.name, value.lower())
+
+ def to_asset(self) -> Asset:
+ return create_asset(host=self.host, catalog=self.catalog,
schema=self.schema, table=self.table)
+
+
def convert_asset_to_openlineage(asset: Asset, lineage_context) ->
OpenLineageDataset:
"""Translate Asset with valid AIP-60 uri to OpenLineage with assistance
from the hook."""
from urllib.parse import urlsplit
diff --git
a/providers/databricks/src/airflow/providers/databricks/operators/databricks_sql.py
b/providers/databricks/src/airflow/providers/databricks/operators/databricks_sql.py
index f44278982aa..5d97bd5e75d 100644
---
a/providers/databricks/src/airflow/providers/databricks/operators/databricks_sql.py
+++
b/providers/databricks/src/airflow/providers/databricks/operators/databricks_sql.py
@@ -26,7 +26,7 @@ import re
from collections.abc import Sequence
from functools import cached_property
from tempfile import NamedTemporaryFile
-from typing import TYPE_CHECKING, Any, ClassVar
+from typing import TYPE_CHECKING, Any, ClassVar, NamedTuple
from urllib.parse import urlparse
from databricks.sql.utils import ParamEscaper
@@ -42,11 +42,28 @@ from airflow.providers.databricks.utils.query_tags import
build_query_tags
if TYPE_CHECKING:
from airflow.providers.common.compat.sdk import Context
+ from airflow.providers.databricks.assets.databricks import
UnityTableIdentity
_IDENTIFIER_RE = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$")
_DISALLOWED_SQL_TOKENS = (";", "--", "/*", "*/")
+class _CopyIntoTarget(NamedTuple):
+ catalog: str | None
+ schema: str | None
+ table: str
+
+
+def _resolve_copy_into_target(table_name: str, *, catalog: str | None, schema:
str | None) -> _CopyIntoTarget:
+ """Split a 1, 2, or 3-part ``table_name``, filling missing leading parts
from ``catalog`` and ``schema``."""
+ parts = table_name.split(".")
+ if len(parts) == 3:
+ return _CopyIntoTarget(catalog=parts[0], schema=parts[1],
table=parts[2])
+ if len(parts) == 2:
+ return _CopyIntoTarget(catalog=catalog, schema=parts[0],
table=parts[1])
+ return _CopyIntoTarget(catalog=catalog, schema=schema, table=table_name)
+
+
class DatabricksSqlOperator(SQLExecuteQueryOperator):
"""
Executes SQL code in a Databricks SQL endpoint or a Databricks cluster.
@@ -484,7 +501,6 @@ class DatabricksCopyIntoOperator(BaseOperator):
raise ValueError("expression_list must not contain
statement separators or comments.")
def _create_sql_query(self) -> str:
-
self._validate_sql_fragments()
escaper = ParamEscaper()
maybe_with = ""
@@ -582,16 +598,7 @@ FILEFORMAT = {self._file_format}
from airflow.providers.common.compat.openlineage.facet import Dataset,
Error
try:
- table_parts = self.table_name.split(".")
- if len(table_parts) == 3: # catalog.schema.table
- catalog, schema, table = table_parts
- elif len(table_parts) == 2: # schema.table
- catalog = None
- schema, table = table_parts
- else:
- catalog = None
- schema = None
- table = self.table_name
+ catalog, schema, table =
_resolve_copy_into_target(self.table_name, catalog=None, schema=None)
hook = self._get_hook()
schema = schema or hook.get_openlineage_default_schema() #
Fallback to default schema
@@ -655,3 +662,54 @@ FILEFORMAT = {self._file_format}
job_facets={"sql":
SQLJobFacet(query=SQLParser.normalize_sql(self._sql))},
run_facets=run_facets,
)
+
+
+class DatabricksCopyIntoAssetOperator(DatabricksCopyIntoOperator):
+ """
+ Run ``COPY INTO`` and declare the target Unity Catalog table as an asset
outlet.
+
+ Accepts every :class:`DatabricksCopyIntoOperator` argument. ``table_name``
stays templated.
+ ``unity_table`` is static, so the asset is known when the Dag is parsed.
+
+ .. seealso::
+ For more information on how to use this operator, take a look at the
guide:
+ :ref:`howto/operator:DatabricksCopyIntoAssetOperator`
+
+ :param unity_table: Static identity of the ``COPY INTO`` target table.
When ``outlets`` is
+ omitted, the outlets are ``[unity_table.to_asset()]``. When
``outlets`` is passed, including
+ ``outlets=[]``, it is kept as is, so callers that want extra outlets
include
+ ``unity_table.to_asset()`` among them. At execution, the rendered
``table_name`` is resolved
+ with ``catalog`` and ``schema`` and must equal ``unity_table``.
Otherwise ``ValueError`` is
+ raised and no SQL runs.
+ """
+
+ def __init__(self, *, unity_table: UnityTableIdentity, **kwargs) -> None:
+ if "outlets" not in kwargs:
+ kwargs["outlets"] = [unity_table.to_asset()]
+ super().__init__(**kwargs)
+ self.unity_table = unity_table
+
+ def execute(self, context: Context) -> Any:
+ expected = _CopyIntoTarget(
+ catalog=self.unity_table.catalog, schema=self.unity_table.schema,
table=self.unity_table.table
+ )
+ target = _resolve_copy_into_target(self.table_name,
catalog=self._catalog, schema=self._schema)
+ normalized_target = _CopyIntoTarget(
+ catalog=target.catalog.lower() if target.catalog is not None else
None,
+ schema=target.schema.lower() if target.schema is not None else
None,
+ table=target.table.lower(),
+ )
+ if normalized_target != expected:
+ raise ValueError(
+ f"COPY INTO target {target._asdict()} resolved from
table_name={self.table_name!r}, "
+ f"catalog={self._catalog!r}, schema={self._schema!r} does not
match "
+ f"unity_table {expected._asdict()}."
+ )
+
+ hook = self._get_hook()
+ if (hook.host or "").lower() != self.unity_table.host:
+ raise ValueError(
+ f"Databricks connection host {hook.host!r} does not match "
+ f"unity_table host {self.unity_table.host!r}."
+ )
+ return super().execute(context)
diff --git
a/providers/databricks/tests/unit/databricks/assets/test_databricks.py
b/providers/databricks/tests/unit/databricks/assets/test_databricks.py
index 7b072ce1ceb..2a2a2ebe065 100644
--- a/providers/databricks/tests/unit/databricks/assets/test_databricks.py
+++ b/providers/databricks/tests/unit/databricks/assets/test_databricks.py
@@ -23,6 +23,7 @@ import pytest
from airflow.providers.common.compat.assets import Asset
from airflow.providers.databricks.assets.databricks import (
+ UnityTableIdentity,
convert_asset_to_openlineage,
create_asset,
sanitize_uri,
@@ -75,3 +76,57 @@ def test_convert_asset_to_openlineage() -> None:
ol_dataset = convert_asset_to_openlineage(asset=asset,
lineage_context=None)
assert ol_dataset.namespace ==
"databricks://my-workspace.cloud.databricks.com"
assert ol_dataset.name == "main.default.users"
+
+
[email protected](
+ "host",
+ [
+ pytest.param("my-workspace.cloud.databricks.com", id="hostname"),
+ pytest.param("https://my-workspace.cloud.databricks.com", id="url"),
+ pytest.param("https://my-workspace.cloud.databricks.com/",
id="url-trailing-slash"),
+ ],
+)
+def test_unity_table_identity_to_asset(host: str) -> None:
+ identity = UnityTableIdentity(host=host, catalog="main", schema="default",
table="users")
+ assert identity.to_asset() == Asset(
+ uri="databricks://my-workspace.cloud.databricks.com/main/default/users"
+ )
+
+
[email protected](
+ ("fields", "match"),
+ [
+ pytest.param({"host": ""}, "host must not be empty", id="empty-host"),
+ pytest.param({"catalog": ""}, "catalog must not be empty",
id="empty-catalog"),
+ pytest.param({"schema": ""}, "schema must not be empty",
id="empty-schema"),
+ pytest.param({"table": ""}, "table must not be empty",
id="empty-table"),
+ pytest.param({"host": "{{ conn.host }}"}, "host must be static",
id="jinja-host"),
+ pytest.param({"table": "{{ params.table }}"}, "table must be static",
id="jinja-table"),
+ pytest.param({"schema": "{% if x %}a{% endif %}"}, "schema must be
static", id="jinja-block"),
+ ],
+)
+def test_unity_table_identity_rejects_invalid_fields(fields: dict[str, str],
match: str) -> None:
+ valid = {
+ "host": "my-workspace.cloud.databricks.com",
+ "catalog": "main",
+ "schema": "default",
+ "table": "users",
+ }
+ with pytest.raises(ValueError, match=match):
+ UnityTableIdentity(**{**valid, **fields})
+
+
[email protected](
+ "host",
+ [
+ "My-Workspace.cloud.Databricks.com",
+ "https://My-Workspace.cloud.Databricks.com/",
+ ],
+)
+def test_unity_table_identity_normalizes_case(host: str) -> None:
+ identity = UnityTableIdentity(host=host, catalog="Main", schema="Default",
table="Users")
+ canonical = UnityTableIdentity(
+ host="my-workspace.cloud.databricks.com", catalog="main",
schema="default", table="users"
+ )
+ assert identity == canonical
+ assert identity.to_asset() == canonical.to_asset()
diff --git
a/providers/databricks/tests/unit/databricks/operators/test_databricks_copy.py
b/providers/databricks/tests/unit/databricks/operators/test_databricks_copy.py
index 1d369a3343b..564a87a9e4b 100644
---
a/providers/databricks/tests/unit/databricks/operators/test_databricks_copy.py
+++
b/providers/databricks/tests/unit/databricks/operators/test_databricks_copy.py
@@ -21,13 +21,18 @@ from unittest import mock
import pytest
+from airflow.providers.common.compat.assets import Asset
from airflow.providers.common.compat.openlineage.facet import (
Dataset,
ExternalQueryRunFacet,
SQLJobFacet,
)
from airflow.providers.common.compat.sdk import AirflowException
-from airflow.providers.databricks.operators.databricks_sql import
DatabricksCopyIntoOperator
+from airflow.providers.databricks.assets.databricks import UnityTableIdentity
+from airflow.providers.databricks.operators.databricks_sql import (
+ DatabricksCopyIntoAssetOperator,
+ DatabricksCopyIntoOperator,
+)
from airflow.providers.openlineage.extractors import OperatorLineage
DATE = "2017-04-20"
@@ -525,3 +530,211 @@ def test_get_openlineage_facets():
"externalQuery": ExternalQueryRunFacet(externalQueryId="query_id",
source="scheme://host")
}
assert result.job_facets == {"sql": SQLJobFacet(query=op._sql)}
+
+
+USERS_TABLE = UnityTableIdentity(
+ host="https://my-workspace.cloud.databricks.com/", catalog="main",
schema="default", table="users"
+)
+USERS_URI = "databricks://my-workspace.cloud.databricks.com/main/default/users"
+
+
+def run_copy_into(op, context=None):
+ with
mock.patch("airflow.providers.databricks.operators.databricks_sql.DatabricksSqlHook")
as hook_cls:
+ hook_cls.return_value.host = USERS_TABLE.host
+ if context is not None:
+ op.render_template_fields(context)
+ op.execute(context)
+ return hook_cls.return_value.run
+
+
+def test_asset_operator_declares_unity_table_outlet():
+ op = DatabricksCopyIntoAssetOperator(
+ task_id=TASK_ID,
+ file_location=COPY_FILE_LOCATION,
+ file_format="JSON",
+ table_name="main.default.users",
+ unity_table=USERS_TABLE,
+ )
+ assert op.outlets == [Asset(uri=USERS_URI)]
+
+
[email protected](
+ "outlets",
+ (
+ pytest.param([], id="empty"),
+ pytest.param(
+ [Asset(uri=USERS_URI),
Asset(uri="s3://my-bucket/manifests/users")],
+ id="unity-asset-plus-extra",
+ ),
+ ),
+)
+def test_asset_operator_keeps_explicit_outlets(outlets):
+ op = DatabricksCopyIntoAssetOperator(
+ task_id=TASK_ID,
+ file_location=COPY_FILE_LOCATION,
+ file_format="JSON",
+ table_name="main.default.users",
+ unity_table=USERS_TABLE,
+ outlets=outlets,
+ )
+ assert op.outlets == outlets
+
+
+def
test_asset_operator_runs_templated_table_name_that_renders_to_unity_table():
+ op = DatabricksCopyIntoAssetOperator(
+ task_id=TASK_ID,
+ file_location=COPY_FILE_LOCATION,
+ file_format="JSON",
+ table_name="{{ params.catalog }}.default.{{ params.table }}",
+ unity_table=USERS_TABLE,
+ )
+ run = run_copy_into(op, {"params": {"catalog": "main", "table": "users"}})
+ run.assert_called_once_with(
+ f"COPY INTO main.default.users\nFROM
'{COPY_FILE_LOCATION}'\nFILEFORMAT = JSON"
+ )
+
+
[email protected](
+ ("table_name", "catalog", "schema"),
+ (
+ pytest.param("main.default.users", "other_catalog", "other_schema",
id="three-part"),
+ pytest.param("default.users", "main", None, id="two-part"),
+ pytest.param("users", "main", "default", id="one-part"),
+ ),
+)
+def
test_asset_operator_resolves_table_name_with_session_catalog_and_schema(table_name,
catalog, schema):
+ op = DatabricksCopyIntoAssetOperator(
+ task_id=TASK_ID,
+ file_location=COPY_FILE_LOCATION,
+ file_format="JSON",
+ table_name=table_name,
+ catalog=catalog,
+ schema=schema,
+ unity_table=USERS_TABLE,
+ )
+ run = run_copy_into(op)
+ run.assert_called_once_with(f"COPY INTO {table_name}\nFROM
'{COPY_FILE_LOCATION}'\nFILEFORMAT = JSON")
+
+
[email protected](
+ ("table_name", "catalog", "schema"),
+ (
+ pytest.param("main.default.orders", None, None, id="other-table"),
+ pytest.param("default.users", None, None,
id="two-part-without-catalog"),
+ pytest.param("users", "main", None, id="one-part-without-schema"),
+ pytest.param("users", "dev", "default", id="session-catalog-differs"),
+ pytest.param("x.main.default.users", None, None, id="four-part"),
+ ),
+)
+def test_asset_operator_mismatch_raises_before_sql(table_name, catalog,
schema):
+ op = DatabricksCopyIntoAssetOperator(
+ task_id=TASK_ID,
+ file_location=COPY_FILE_LOCATION,
+ file_format="JSON",
+ table_name=table_name,
+ catalog=catalog,
+ schema=schema,
+ unity_table=USERS_TABLE,
+ )
+ with
mock.patch("airflow.providers.databricks.operators.databricks_sql.DatabricksSqlHook")
as hook_cls:
+ with pytest.raises(ValueError, match="does not match unity_table"):
+ op.execute(None)
+ assert hook_cls.return_value.run.call_args_list == []
+ assert op._sql is None
+
+
[email protected]("host", ["other-workspace.cloud.databricks.com",
None, ""])
+def
test_asset_operator_rejects_connection_for_other_workspace_before_sql(host):
+ op = DatabricksCopyIntoAssetOperator(
+ task_id=TASK_ID,
+ file_location=COPY_FILE_LOCATION,
+ file_format="JSON",
+ table_name="main.default.users",
+ unity_table=USERS_TABLE,
+ )
+ with
mock.patch("airflow.providers.databricks.operators.databricks_sql.DatabricksSqlHook")
as hook_cls:
+ hook_cls.return_value.host = host
+ with pytest.raises(ValueError, match="connection host .* does not
match unity_table host"):
+ op.execute(None)
+ assert hook_cls.return_value.run.call_args_list == []
+ assert op._sql is None
+
+
+def test_asset_operator_volume_is_copy_source_and_table_is_outlet():
+ op = DatabricksCopyIntoAssetOperator(
+ task_id=TASK_ID,
+ file_location="/Volumes/main/default/landing/users.csv",
+ file_format="CSV",
+ table_name="main.default.users",
+ unity_table=USERS_TABLE,
+ )
+ run = run_copy_into(op)
+ assert op.outlets == [Asset(uri=USERS_URI)]
+ run.assert_called_once_with(
+ "COPY INTO main.default.users\nFROM
'/Volumes/main/default/landing/users.csv'\nFILEFORMAT = CSV"
+ )
+
+
+def test_base_copy_into_operator_has_no_outlets_and_runs_any_table():
+ op = DatabricksCopyIntoOperator(
+ task_id=TASK_ID,
+ file_location=COPY_FILE_LOCATION,
+ file_format="JSON",
+ table_name="main.default.orders",
+ )
+ run = run_copy_into(op)
+ assert op.outlets == []
+ run.assert_called_once_with(
+ f"COPY INTO main.default.orders\nFROM
'{COPY_FILE_LOCATION}'\nFILEFORMAT = JSON"
+ )
+
+
[email protected](
+ ("table_name", "catalog", "schema"),
+ [
+ pytest.param("MAIN.Default.Users", None, None, id="three-part"),
+ pytest.param("Default.Users", "MAIN", None, id="two-part"),
+ pytest.param("Users", "MAIN", "Default", id="one-part"),
+ pytest.param("{{ params.table }}", None, None, id="templated"),
+ ],
+)
[email protected]("airflow.providers.databricks.operators.databricks_sql.DatabricksSqlHook",
autospec=True)
+def test_asset_operator_matches_table_case_insensitively(hook_cls, table_name,
catalog, schema):
+ hook_cls.return_value.host = USERS_TABLE.host
+ op = DatabricksCopyIntoAssetOperator(
+ task_id=TASK_ID,
+ file_location=COPY_FILE_LOCATION,
+ file_format="JSON",
+ table_name=table_name,
+ catalog=catalog,
+ schema=schema,
+ unity_table=USERS_TABLE,
+ )
+ context = {"params": {"table": "MAIN.Default.Users"}}
+ op.render_template_fields(context)
+ op.execute(context)
+ rendered_table = "MAIN.Default.Users" if table_name.startswith("{{") else
table_name
+ hook_cls.return_value.run.assert_called_once_with(
+ f"COPY INTO {rendered_table}\nFROM '{COPY_FILE_LOCATION}'\nFILEFORMAT
= JSON"
+ )
+ assert op.outlets == [Asset(uri=USERS_URI)]
+
+
[email protected]("host", ["my-workspace.cloud.databricks.com",
"My-Workspace.cloud.Databricks.com"])
[email protected]("airflow.providers.databricks.operators.databricks_sql.DatabricksSqlHook",
autospec=True)
+def test_asset_operator_matches_workspace_case_insensitively(hook_cls, host):
+ hook_cls.return_value.host = host
+ op = DatabricksCopyIntoAssetOperator(
+ task_id=TASK_ID,
+ file_location=COPY_FILE_LOCATION,
+ file_format="JSON",
+ table_name="main.default.users",
+ unity_table=UnityTableIdentity(
+ host="My-Workspace.cloud.Databricks.com", catalog="Main",
schema="Default", table="Users"
+ ),
+ )
+ op.execute(None)
+ hook_cls.return_value.run.assert_called_once_with(
+ f"COPY INTO main.default.users\nFROM
'{COPY_FILE_LOCATION}'\nFILEFORMAT = JSON"
+ )
+ assert op.outlets == [Asset(uri=USERS_URI)]