diff --git a/pyproject.toml b/pyproject.toml index 2d6938b..c0e868b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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] @@ -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 = "*" @@ -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" @@ -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', @@ -211,6 +210,7 @@ lint.select = [ 'RUF', 'S102' ] +target-version = 'py310' [tool.setuptools.package_data] datacustomcode = ["templates/**/*", "config.yaml"] diff --git a/src/datacustomcode/client.py b/src/datacustomcode/client.py index b43e423..5c37d6f 100644 --- a/src/datacustomcode/client.py +++ b/src/datacustomcode/client.py @@ -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. @@ -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): diff --git a/src/datacustomcode/named_credential/spark_base.py b/src/datacustomcode/named_credential/spark_base.py index 2c47743..25715db 100644 --- a/src/datacustomcode/named_credential/spark_base.py +++ b/src/datacustomcode/named_credential/spark_base.py @@ -66,11 +66,12 @@ 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. @@ -78,6 +79,10 @@ def request_col( 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:/?``). A null row value + falls back to ``request.url``; when omitted, ``request.url`` is + used for every row. Returns: A ``Column`` yielding a struct diff --git a/src/datacustomcode/named_credential/spark_default.py b/src/datacustomcode/named_credential/spark_default.py index 8f3ea23..bdccfe7 100644 --- a/src/datacustomcode/named_credential/spark_default.py +++ b/src/datacustomcode/named_credential/spark_default.py @@ -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. @@ -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( diff --git a/tests/spark/test_column_hints_deferred.py b/tests/spark/test_column_hints_deferred.py index edf0622..9c094eb 100644 --- a/tests/spark/test_column_hints_deferred.py +++ b/tests/spark/test_column_hints_deferred.py @@ -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( diff --git a/tests/test_client.py b/tests/test_client.py index 4f3f2f4..e2a68d8 100644 --- a/tests/test_client.py +++ b/tests/test_client.py @@ -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: diff --git a/tests/test_named_credential.py b/tests/test_named_credential.py index d486a74..5e63348 100644 --- a/tests/test_named_credential.py +++ b/tests/test_named_credential.py @@ -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."""