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
Original file line number Diff line number Diff line change
Expand Up @@ -64,6 +64,9 @@ class SnowparkContainerJobOperator(BaseOperator):
This is separate from the ``warehouse`` parameter used by the operator's
own SQL commands
:param replicas: (Optional) number of job replicas to run. (default value: 1)
:param external_access_integrations: (Optional) Names of the external access
integrations that allow your job to access external sites. Names are
case-sensitive (default value: None)
:param wait_for_completion: poll until the job reaches a terminal state.
When disabled, the job is submitted and the operator returns
immediately. (default value: True)
Expand Down Expand Up @@ -100,6 +103,7 @@ class SnowparkContainerJobOperator(BaseOperator):
"name",
"query_warehouse",
"snowflake_conn_id",
"external_access_integrations",
)

def __init__(
Expand All @@ -113,6 +117,7 @@ def __init__(
name: str | None = None,
query_warehouse: str | None = None,
replicas: int = 1,
external_access_integrations: list[str] | None = None,
wait_for_completion: bool = True,
drop_on_completion: bool = True,
poll_interval: int = 10,
Expand All @@ -136,6 +141,7 @@ def __init__(
self.name = name
self.query_warehouse = query_warehouse
self.replicas = replicas
self.external_access_integrations = external_access_integrations
self.wait_for_completion = wait_for_completion
self.drop_on_completion = drop_on_completion
self.poll_interval = poll_interval
Expand Down Expand Up @@ -172,6 +178,9 @@ def _build_sql(self) -> str:
sql += f" REPLICAS = {self.replicas}"
if self.query_warehouse:
sql += f" QUERY_WAREHOUSE = {self.query_warehouse}"
if self.external_access_integrations:
eais = ", ".join(self.external_access_integrations)
sql += f" EXTERNAL_ACCESS_INTEGRATIONS = ({eais})"
if self.spec_text:
sql += f" FROM SPECIFICATION $${self.spec_text}$$"
else:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -120,12 +120,27 @@ def test_build_sql_with_spec_text(self):
pytest.param(
{"query_warehouse": "COMPUTE_WH"}, "QUERY_WAREHOUSE = COMPUTE_WH", id="query_warehouse"
),
pytest.param(
{"external_access_integrations": ["test_eai"]},
"EXTERNAL_ACCESS_INTEGRATIONS = (test_eai)",
id="external_access_integrations_single",
),
pytest.param(
{"external_access_integrations": ["test_eai", "test_eai_2"]},
"EXTERNAL_ACCESS_INTEGRATIONS = (test_eai, test_eai_2)",
id="external_access_integrations_multiple",
),
),
)
def test_build_sql_optional_params(self, kwargs, expected):
op = _make_operator(**kwargs)
assert expected in op._build_sql()

def test_external_access_integrations_in_template_fields(self):
op = _make_operator(external_access_integrations=["test_eai"])
assert "external_access_integrations" in op.template_fields
assert hasattr(op, "external_access_integrations")

@mock.patch(MOCK_HOOK_PATH)
def test_submit_job_parses_job_name(self, mock_hook_cls):
mock_hook = mock_hook_cls.return_value
Expand Down