Skip to content
Closed
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
12 changes: 6 additions & 6 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -87,10 +87,10 @@ warn_unused_configs = true

[tool.poetry]
include = [
{path = "src/datacustomcode/templates/**/*"},
{path = "src/datacustomcode/config.yaml"}
{ path = "src/datacustomcode/templates/**/*" },
{ path = "src/datacustomcode/config.yaml" }
]
packages = [{include = "datacustomcode", from = "src"}]
packages = [{ include = "datacustomcode", from = "src" }]
version = "0.0.0"

[tool.poetry.build]
Expand All @@ -115,7 +115,7 @@ coverage = ">=7.0.0,<8.0.0"
ipykernel = "^6.29.5"
mypy = "*"
pipreqs = "*"
poetry-dynamic-versioning = {extras = ["plugin"], version = "^1.8.2"}
poetry-dynamic-versioning = { extras = ["plugin"], version = "^1.8.2" }
pre-commit = "*"
pytest = "*"
pytest-cov = "*"
Expand All @@ -127,7 +127,7 @@ types-requests = "*"
[tool.poetry.plugins."poetry.plugin"]

[tool.poetry.requires-plugins]
poetry-dynamic-versioning = {version = ">=1.0.0,<2.0.0", extras = ["plugin"]}
poetry-dynamic-versioning = { version = ">=1.0.0,<2.0.0", extras = ["plugin"] }

[tool.poetry.scripts]
datacustomcode = "datacustomcode.cli:cli"
Expand All @@ -152,7 +152,6 @@ testpaths = ["tests"]
[tool.ruff]
fix = true
line-length = 88
target-version = 'py310'
lint.ignore = [
# do not assign a lambda expression, use a def
'E731',
Expand Down Expand Up @@ -211,6 +210,7 @@ lint.select = [
'RUF',
'S102'
]
target-version = 'py310'

[tool.setuptools.package_data]
datacustomcode = ["templates/**/*", "config.yaml"]
Expand Down
13 changes: 10 additions & 3 deletions src/datacustomcode/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -231,11 +231,12 @@ def einstein_predict_col(
def named_credential_request_col(
request: "HTTPRequest",
body: Optional["Column"] = None,
url: Optional["Column"] = None,
) -> "Column":
"""Build a Spark Column that makes one Named Credential callout per row.

The endpoint, method, and headers are fixed for the call (taken from
``request``); only ``body`` varies per row. Use this instead of
The method and headers are fixed for the call (taken from ``request``);
``body`` and, optionally, ``url`` vary per row. Use this instead of
:meth:`Client.named_credential_request` when the callout runs across a
DataFrame so each row is dispatched independently rather than one-shot on
the driver.
Expand All @@ -253,13 +254,19 @@ def named_credential_request_col(
headers are applied to every row.
body: Optional per-row ``Column`` holding the request body as a
string (or null for no body).
url: Optional per-row ``Column`` holding the full callout url, e.g.
``concat(lit("callout:MyNC/geocode?address="), col("address"))``.
A null row value falls back to ``request.url``; when omitted,
``request.url`` is used for every row.

Returns:
A Spark ``Column`` of ``StructType`` with fields ``status``,
``response``, ``error_code``, and ``error_message``.
"""
named_credential = Client()._get_spark_named_credential()
return named_credential.request_col(request, body=body)
if url is None:
return named_credential.request_col(request, body=body)
return named_credential.request_col(request, body=body, url=url)


class DataCloudObjectType(Enum):
Expand Down
9 changes: 7 additions & 2 deletions src/datacustomcode/named_credential/spark_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -66,18 +66,23 @@ def request_col(
self,
request: HTTPRequest,
body: Optional["Column"] = None,
url: Optional["Column"] = None,
) -> "Column":
"""Build a Spark ``Column`` that makes one external callout per row.

The endpoint, method, and headers are fixed for the call (taken from
``request``); only ``body`` varies per row. Use this instead of
The method and headers are fixed for the call (taken from ``request``);
``body`` and, optionally, ``url`` vary per row. Use this instead of
:meth:`request` when the callout runs across a DataFrame so each row is
dispatched independently rather than one-shot on the driver.

Args:
request: The callout template
body: Optional per-row ``Column`` holding the request body as a
string, sent verbatim (or null for no body).
url: Optional per-row ``Column`` holding the full callout url
(``callout:<NamedCredential>/<path>?<query>``). A null row value
falls back to ``request.url``; when omitted, ``request.url`` is
used for every row.

Returns:
A ``Column`` yielding a struct
Expand Down
29 changes: 25 additions & 4 deletions src/datacustomcode/named_credential/spark_default.py
Original file line number Diff line number Diff line change
Expand Up @@ -79,6 +79,7 @@ def request_col(
self,
request: "HTTPRequest",
body: Optional["Column"] = None,
url: Optional["Column"] = None,
) -> "Column":
"""Per-row callout via a client-side Spark UDF.

Expand Down Expand Up @@ -114,11 +115,31 @@ def request_col(
]
)

def _callout(body_str: Optional[str]) -> Dict[str, Any]:
return _invoke_callout_as_struct(self._named_credential, request, body_str)

body_col = body if body is not None else lit(None).cast(StringType())
return udf(_callout, result_schema)(body_col)

if url is None:

def _callout(body_str: Optional[str]) -> Dict[str, Any]:
return _invoke_callout_as_struct(
self._named_credential, request, body_str
)

return udf(_callout, result_schema)(body_col)

def _callout_with_url(
body_str: Optional[str], url_str: Optional[str]
) -> Dict[str, Any]:
# A null row url falls back to the template's url.
row_request = (
request
if url_str is None
else request.model_copy(update={"url": url_str})
)
return _invoke_callout_as_struct(
self._named_credential, row_request, body_str
)

return udf(_callout_with_url, result_schema)(body_col, url)


def _invoke_callout_as_struct(
Expand Down
3 changes: 2 additions & 1 deletion tests/spark/test_column_hints_deferred.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,7 +38,8 @@ def _run(script: str) -> None:
[sys.executable, "-c", textwrap.dedent(script)],
capture_output=True,
text=True,
timeout=60, check=False,
timeout=60,
check=False,
)
if result.returncode != 0:
pytest.fail(
Expand Down
19 changes: 19 additions & 0 deletions tests/test_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -733,6 +733,25 @@ def test_delegates_to_spark_named_credential(self, mock_build, reset_client):
assert result is sentinel_col
mock_nc.request_col.assert_called_once_with(request, body=body_col)

@patch("datacustomcode.client._build_spark_named_credential")
def test_forwards_per_row_url_column(self, mock_build, reset_client):
mock_nc = MagicMock()
sentinel_col = MagicMock(name="col")
mock_nc.request_col.return_value = sentinel_col
mock_build.return_value = mock_nc

reader = MagicMock(spec=BaseDataCloudReader)
writer = MagicMock(spec=BaseDataCloudWriter)
Client(reader=reader, writer=writer)

request = MagicMock(name="request")
body_col = MagicMock(name="body_col")
url_col = MagicMock(name="url_col")
result = named_credential_request_col(request, body_col, url=url_col)

assert result is sentinel_col
mock_nc.request_col.assert_called_once_with(request, body=body_col, url=url_col)


class TestClientNamedCredentialRequest:

Expand Down
90 changes: 90 additions & 0 deletions tests/test_named_credential.py
Original file line number Diff line number Diff line change
Expand Up @@ -403,6 +403,96 @@ def test_defaults_body_to_typed_null_column(self, mock_lit, mock_udf):
# With no body column, a typed null string column is applied instead.
sentinel_udf.assert_called_once_with(null_col)

@patch("pyspark.sql.functions.udf")
@patch("pyspark.sql.functions.lit")
def test_per_row_url_column_overrides_request_url(self, mock_lit, mock_udf):
from datacustomcode.named_credential.spark_default import (
DefaultSparkNamedCredential,
)

sentinel_udf = MagicMock(name="udf")
mock_udf.return_value = sentinel_udf

underlying = MagicMock()
underlying.request.return_value = HTTPResponse(
status_code=200, headers={}, body="{}"
)
spark_nc = DefaultSparkNamedCredential(named_credential=underlying)

request = (
HTTPRequestBuilder()
.set_url("callout:NC/geocode")
.set_method("GET")
.set_headers({"Accept": "application/json"})
.build()
)
body_col = MagicMock(name="body_col")
url_col = MagicMock(name="url_col")
spark_nc.request_col(request, body_col, url=url_col)

# The UDF is applied over both the body and the per-row url columns.
sentinel_udf.assert_called_once_with(body_col, url_col)

udf_fn = mock_udf.call_args.args[0]
out = udf_fn(None, "callout:NC/geocode?address=1%20Market%20St")

assert out["status"] == "SUCCESS"
sent_request, sent_body = underlying.request.call_args.args
# The row's url replaces the template url; method/headers are kept.
assert sent_request.url == "callout:NC/geocode?address=1%20Market%20St"
assert sent_request.method == "GET"
assert sent_request.headers == {"Accept": "application/json"}
assert sent_body is None
# The caller's template request is left untouched.
assert request.url == "callout:NC/geocode"

@patch("pyspark.sql.functions.udf")
@patch("pyspark.sql.functions.lit")
def test_null_row_url_falls_back_to_request_url(self, mock_lit, mock_udf):
from datacustomcode.named_credential.spark_default import (
DefaultSparkNamedCredential,
)

mock_udf.return_value = MagicMock(name="udf")

underlying = MagicMock()
underlying.request.return_value = HTTPResponse(
status_code=200, headers={}, body=""
)
spark_nc = DefaultSparkNamedCredential(named_credential=underlying)

request = HTTPRequestBuilder().set_url("callout:NC/default").build()
spark_nc.request_col(request, url=MagicMock(name="url_col"))

udf_fn = mock_udf.call_args.args[0]
udf_fn('{"a": 1}', None)

sent_request, sent_body = underlying.request.call_args.args
assert sent_request is request
assert sent_body == '{"a": 1}'

@patch("pyspark.sql.functions.udf")
@patch("pyspark.sql.functions.lit")
def test_url_without_body_applies_typed_null_body(self, mock_lit, mock_udf):
from datacustomcode.named_credential.spark_default import (
DefaultSparkNamedCredential,
)

null_col = MagicMock(name="null_col")
lit_none = MagicMock(name="lit_none")
lit_none.cast.return_value = null_col
mock_lit.return_value = lit_none

sentinel_udf = MagicMock(name="udf")
mock_udf.return_value = sentinel_udf

spark_nc = DefaultSparkNamedCredential(named_credential=MagicMock())
request = HTTPRequestBuilder().set_url("callout:NC/status").build()
url_col = MagicMock(name="url_col")
spark_nc.request_col(request, url=url_col)

sentinel_udf.assert_called_once_with(null_col, url_col)


class TestInvokeCalloutAsStruct:
"""The callout-to-struct shaping shared by every row."""
Expand Down
Loading