From 86bdfbd8282de7695961db481d330f8e5847c24e Mon Sep 17 00:00:00 2001 From: Adithya Samavedhi Date: Thu, 3 Sep 2026 16:24:52 -0700 Subject: [PATCH] use-ttd-auth-from-dataclient-instead-of-passing-it-per-call --- README.md | 11 +++-- pyproject.toml | 4 +- tests/test_placeholder.py | 1 - tests/unit/test_batch_process_early_exit.py | 1 - tests/unit/test_call_api.py | 6 +-- tests/unit/test_client_helpers.py | 1 - tests/unit/test_contexts.py | 5 +-- tests/unit/test_handlers_build_items.py | 30 +++++++------ tests/unit/test_process_partitions.py | 15 +++++-- tests/unit/test_push_data.py | 1 - tests/unit/test_uid2_resolutions.py | 13 +++--- .../ttd_databricks/batching.py | 19 +++----- .../ttd_databricks/handlers/advertiser.py | 2 - .../handlers/deletion_optout_advertiser.py | 2 - .../handlers/deletion_optout_merchant.py | 2 - .../handlers/deletion_optout_thirdparty.py | 2 - .../handlers/offline_conversion.py | 2 - .../ttd_databricks/handlers/third_party.py | 2 - .../ttd_databricks/schemas/advertiser.py | 4 +- .../ttd_databricks/schemas/third_party.py | 4 +- .../ttd_databricks/ttd_client.py | 43 +++++++++++-------- 21 files changed, 79 insertions(+), 91 deletions(-) diff --git a/README.md b/README.md index bf9e929..65def3b 100644 --- a/README.md +++ b/README.md @@ -198,20 +198,21 @@ client = TtdDatabricksClient.from_params( Provide your own [`DataClient`](https://github.com/thetradedesk/ttd-data-python/blob/main/src/ttd_data/sdk.py) instance to control the underlying HTTP transport directly. Use this when you need to configure options not exposed by `from_params()`, or to inject a mock in tests. +The `DataClient` you pass in must carry your API token as `ttd_auth`; every request the SDK makes authenticates with it. ```python from ttd_data import DataClient from ttd_databricks_python.ttd_databricks import TtdDatabricksClient -# Configure DataClient with custom HTTP settings. +# Configure DataClient with your API token and custom HTTP settings. data_client = DataClient( + ttd_auth="", # your TTD platform API token server_url="https://custom-server.example.com", # override default server URL timeout_ms=10000, # request timeout in milliseconds ) client = TtdDatabricksClient( data_api_client=data_client, - api_token="", spark=spark, # optional; spark variable available from the Databricks notebook runtime ) ``` @@ -456,15 +457,13 @@ from ttd_data.utils.retries import BackoffStrategy, RetryConfig from ttd_databricks_python.ttd_databricks import TtdDatabricksClient data_client = DataClient( + ttd_auth="", # your TTD platform API token server_url="https://custom-server.example.com", # override default server URL timeout_ms=10000, # request timeout in milliseconds retry_config=RetryConfig("backoff", BackoffStrategy(1000, 60000, 1.5, 3600000), True), # custom retry config ) -client = TtdDatabricksClient( - data_api_client=data_client, - api_token="", -) +client = TtdDatabricksClient(data_api_client=data_client) ``` In batch processing mode, a `DataClient` singleton is maintained per Spark worker process to enable HTTP connection reuse across batches, reducing overhead during distributed execution. diff --git a/pyproject.toml b/pyproject.toml index ba08aed..75586f6 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta" [project] name = "ttd-databricks" -version = "0.5.0" +version = "0.6.0" description = "Client implementation and helper functions for integrating with the TTD Databricks services." readme = "README.md" requires-python = ">=3.10" @@ -15,7 +15,7 @@ authors = [ ] dependencies = [ - "ttd-data>=0.2.6,<0.3.0", + "ttd-data>=0.3.1,<0.4.0", "pandas>=1.0.5", "pyarrow>=4.0.0", "setuptools>=63.4.1", diff --git a/tests/test_placeholder.py b/tests/test_placeholder.py index 75851a1..773b5ea 100644 --- a/tests/test_placeholder.py +++ b/tests/test_placeholder.py @@ -13,4 +13,3 @@ def test_placeholder(self) -> None: if __name__ == "__main__": unittest.main() - diff --git a/tests/unit/test_batch_process_early_exit.py b/tests/unit/test_batch_process_early_exit.py index 04c1085..aaafcc8 100644 --- a/tests/unit/test_batch_process_early_exit.py +++ b/tests/unit/test_batch_process_early_exit.py @@ -34,7 +34,6 @@ def _make_client(spark: SparkSession) -> TtdDatabricksClient: return TtdDatabricksClient( data_api_client=MagicMock(spec=DataClient), - api_token="test-token", spark=spark, ) diff --git a/tests/unit/test_call_api.py b/tests/unit/test_call_api.py index 0269661..02f929f 100644 --- a/tests/unit/test_call_api.py +++ b/tests/unit/test_call_api.py @@ -17,10 +17,8 @@ from typing import Any from unittest.mock import MagicMock, patch -import pytest - import httpx - +import pytest from ttd_data import DataClient from ttd_data.errors import DataError, NoResponseError, ResponseValidationError @@ -30,7 +28,7 @@ def _make_client() -> TtdDatabricksClient: - return TtdDatabricksClient(data_api_client=MagicMock(spec=DataClient), api_token="test-token") + return TtdDatabricksClient(data_api_client=MagicMock(spec=DataClient)) def _make_rows(*dicts: dict[str, Any]) -> list[MagicMock]: diff --git a/tests/unit/test_client_helpers.py b/tests/unit/test_client_helpers.py index 30200d4..a88ea07 100644 --- a/tests/unit/test_client_helpers.py +++ b/tests/unit/test_client_helpers.py @@ -19,7 +19,6 @@ def _make_client(**kwargs) -> TtdDatabricksClient: # type: ignore[no-untyped-def] return TtdDatabricksClient( data_api_client=MagicMock(spec=DataClient), - api_token="test-token", **kwargs, ) diff --git a/tests/unit/test_contexts.py b/tests/unit/test_contexts.py index 07795c4..03ac8e8 100644 --- a/tests/unit/test_contexts.py +++ b/tests/unit/test_contexts.py @@ -14,7 +14,6 @@ ) from ttd_databricks_python.ttd_databricks.endpoints import TTDEndpoint - _REQUEST_TYPE = PartnerDsrRequestType.OPT_OUT _DATA_ORIGINS = [DataOrigin(id="test-origin", type=DataOriginType.DATA_PROVIDER)] @@ -72,9 +71,7 @@ def test_deletion_optout_advertiser_context_verify_context_pickling(): def test_deletion_optout_thirdparty_context_verify_context_pickling(): - ctx = DeletionOptOutThirdPartyContext( - data_provider_id="prov123", request_type=_REQUEST_TYPE, brand_id="brand99" - ) + ctx = DeletionOptOutThirdPartyContext(data_provider_id="prov123", request_type=_REQUEST_TYPE, brand_id="brand99") restored = pickle.loads(pickle.dumps(ctx)) assert restored.data_provider_id == "prov123" assert restored.request_type == _REQUEST_TYPE diff --git a/tests/unit/test_handlers_build_items.py b/tests/unit/test_handlers_build_items.py index f7b0e2f..0319228 100644 --- a/tests/unit/test_handlers_build_items.py +++ b/tests/unit/test_handlers_build_items.py @@ -4,20 +4,19 @@ """ from datetime import datetime, timezone -from typing import Union import numpy as np import pytest +from ttd_data.models import AdvertiserDataItem, OfflineConversionDataItem, PartnerDsrDataItem, ThirdPartyDataItem +from ttd_data.types import UNSET import ttd_databricks_python.ttd_databricks.handlers.advertiser as adv_handler -from ttd_databricks_python.ttd_databricks.id_types import normalize_id_type import ttd_databricks_python.ttd_databricks.handlers.deletion_optout_advertiser as del_adv_handler import ttd_databricks_python.ttd_databricks.handlers.deletion_optout_merchant as del_merch_handler import ttd_databricks_python.ttd_databricks.handlers.deletion_optout_thirdparty as del_tp_handler import ttd_databricks_python.ttd_databricks.handlers.offline_conversion as oc_handler import ttd_databricks_python.ttd_databricks.handlers.third_party as tp_handler -from ttd_data.models import AdvertiserDataItem, OfflineConversionDataItem, PartnerDsrDataItem, ThirdPartyDataItem -from ttd_data.types import UNSET +from ttd_databricks_python.ttd_databricks.id_types import normalize_id_type # UNSET is not a singleton — the SDK creates fresh Unset() instances per field. # Use isinstance check rather than identity (is). @@ -27,7 +26,7 @@ # An array column reaches build_items as a list via the adhoc path # (collect + asDict) and as a numpy array via the batch path (mapInPandas). # build_items must handle both, so array-column tests run against each shape. -def _build_array_column(array_type: type, items: list[dict]) -> Union[list, np.ndarray]: +def _build_array_column(array_type: type, items: list[dict]) -> list | np.ndarray: return items if array_type is list else np.array(items, dtype=object) @@ -43,7 +42,7 @@ def test_builds_advertiser_data_item_with_correct_fields(self): # Handler maps id_type → AdvertiserDataItem field dynamically: {d["id_type"]: d["id_value"]} item = adv_handler.build_items([self._MINIMAL])[0] assert isinstance(item, AdvertiserDataItem) - assert getattr(item, "tdid") == "test-tdid-value" + assert item.tdid == "test-tdid-value" assert item.data[0].name == "test-segment-name" def test_none_optional_fields_are_not_sent_to_api(self): @@ -76,7 +75,7 @@ class TestThirdPartyBuildItems: def test_builds_third_party_data_item_with_correct_fields(self): item = tp_handler.build_items([self._MINIMAL])[0] assert isinstance(item, ThirdPartyDataItem) - assert getattr(item, "tdid") == "test-tdid-value" + assert item.tdid == "test-tdid-value" assert item.data[0].name == "test-segment-name" def test_none_optional_fields_are_not_sent_to_api(self): @@ -99,19 +98,19 @@ def test_optional_fields_are_passed_through_when_provided(self): def test_deletion_optout_advertiser_returns_partner_dsr_item_with_correct_id(): item = del_adv_handler.build_items([{"id_type": "TDID", "id_value": "test-advertiser-tdid"}])[0] assert isinstance(item, PartnerDsrDataItem) - assert getattr(item, "tdid") == "test-advertiser-tdid" + assert item.tdid == "test-advertiser-tdid" def test_deletion_optout_thirdparty_returns_partner_dsr_item_with_correct_id(): item = del_tp_handler.build_items([{"id_type": "UID2", "id_value": "test-thirdparty-uid2"}])[0] assert isinstance(item, PartnerDsrDataItem) - assert getattr(item, "uid2") == "test-thirdparty-uid2" + assert item.uid2 == "test-thirdparty-uid2" def test_deletion_optout_merchant_returns_partner_dsr_item_with_correct_id(): item = del_merch_handler.build_items([{"id_type": "TDID", "id_value": "test-merchant-tdid"}])[0] assert isinstance(item, PartnerDsrDataItem) - assert getattr(item, "tdid") == "test-merchant-tdid" + assert item.tdid == "test-merchant-tdid" # --------------------------------------------------------------------------- # @@ -143,8 +142,13 @@ def test_user_ids_converted_to_user_id_array_with_type_codes(self, array_type): def test_all_user_id_types_map_to_correct_codes(self): type_map = { - "TDID": "0", "DAID": "1", "UID2": "2", "UID2Token": "3", - "EUID": "4", "EUIDToken": "5", "RampID": "6", + "TDID": "0", + "DAID": "1", + "UID2": "2", + "UID2Token": "3", + "EUID": "4", + "EUIDToken": "5", + "RampID": "6", } for id_type, expected_code in type_map.items(): row = {**self._MINIMAL, "user_ids": [{"type": id_type, "id": f"test-{id_type}-value"}]} @@ -209,4 +213,4 @@ def test_collect_raw_pii_ids_keeps_only_pii_types(self, array_type): assert oc_handler.collect_raw_pii_ids_per_row(rows) == [["a@example.com"]] def test_collect_raw_pii_ids_handles_missing_user_ids(self): - assert oc_handler.collect_raw_pii_ids_per_row([self._MINIMAL]) == [[]] \ No newline at end of file + assert oc_handler.collect_raw_pii_ids_per_row([self._MINIMAL]) == [[]] diff --git a/tests/unit/test_process_partitions.py b/tests/unit/test_process_partitions.py index 3805292..5dea809 100644 --- a/tests/unit/test_process_partitions.py +++ b/tests/unit/test_process_partitions.py @@ -37,7 +37,9 @@ server_url=None, retry_config=None, timeout_ms=10_000, + ttd_auth="not-a-real-token", uid2_config=None, + graphql_server_url=None, ) @@ -46,6 +48,7 @@ class _StubHandler(BaseHTTPRequestHandler): status_code = 500 request_count = 0 + auth_headers: list[str | None] = [] # ThreadingHTTPServer handles each request on its own thread; `+= 1` is a # non-atomic read-modify-write, so guard it rather than relying on the spark # fixture staying single-threaded. @@ -56,10 +59,12 @@ def configure(cls, status_code: int) -> None: with cls.counter_lock: cls.status_code = status_code cls.request_count = 0 + cls.auth_headers = [] def do_POST(self) -> None: # noqa: N802 — required by stdlib BaseHTTPRequestHandler with type(self).counter_lock: type(self).request_count += 1 + type(self).auth_headers.append(self.headers.get("TTD-Auth")) body = b'{"Message":"forced error for test"}' self.send_response(type(self).status_code) self.send_header("Content-Type", "application/json") @@ -112,7 +117,6 @@ def test_mapinpandas_wires_up_and_round_trips(spark: SparkSession, stub_server: df=input_df, batch_size=3, output_schema=output_schema, - api_token="not-a-real-token", context=context, parallelism=2, client_config=_NO_RETRY_CLIENT_CONFIG, @@ -127,6 +131,9 @@ def test_mapinpandas_wires_up_and_round_trips(spark: SparkSession, stub_server: assert result_df.schema.fieldNames() == output_schema.fieldNames() # 4. Input column values survive Arrow → pandas → dict → pandas → Arrow round-trip. assert {row["id_value"] for row in result_rows} == set(input_ids) + # 5. The worker's rebuilt DataClient authenticates: ttd_auth travels in the client_config + # snapshot, not as a separate per-call argument. + assert set(_StubHandler.auth_headers) == {"not-a-real-token"} @pytest.mark.parametrize( @@ -136,7 +143,9 @@ def test_mapinpandas_wires_up_and_round_trips(spark: SparkSession, stub_server: 403, # token valid but not entitled to this advertiser or data provider ], ) -def test_401_and_403_stop_partition_without_failing_job(spark: SparkSession, stub_server: str, status_code: int) -> None: +def test_401_and_403_stop_partition_without_failing_job( + spark: SparkSession, stub_server: str, status_code: int +) -> None: """401/403 stop the partition without raising. The batch that was sent keeps the server's own status; every row after it is ABORTED, meaning it was never submitted.""" rows = [("TDID", f"id-{i}", "seg-a", None, None) for i in range(7)] @@ -149,7 +158,6 @@ def test_401_and_403_stop_partition_without_failing_job(spark: SparkSession, stu df=input_df, batch_size=3, output_schema=output_schema, - api_token="not-a-real-token", context=context, parallelism=1, client_config=_NO_RETRY_CLIENT_CONFIG, @@ -182,7 +190,6 @@ def test_other_4xx_fails_only_its_own_batch(spark: SparkSession, stub_server: st df=input_df, batch_size=3, output_schema=output_schema, - api_token="not-a-real-token", context=context, parallelism=1, client_config=_NO_RETRY_CLIENT_CONFIG, diff --git a/tests/unit/test_push_data.py b/tests/unit/test_push_data.py index 58233bf..21befaf 100644 --- a/tests/unit/test_push_data.py +++ b/tests/unit/test_push_data.py @@ -32,7 +32,6 @@ def _make_client(spark: SparkSession) -> TtdDatabricksClient: return TtdDatabricksClient( data_api_client=MagicMock(spec=DataClient), - api_token="test-token", spark=spark, ) diff --git a/tests/unit/test_uid2_resolutions.py b/tests/unit/test_uid2_resolutions.py index 295f973..7f9ead5 100644 --- a/tests/unit/test_uid2_resolutions.py +++ b/tests/unit/test_uid2_resolutions.py @@ -16,11 +16,11 @@ from unittest.mock import MagicMock, patch import pytest +from ttd_data import DataClient +from ttd_data.uid2 import UID2Resolution import ttd_databricks_python.ttd_databricks.handlers.advertiser as adv_handler import ttd_databricks_python.ttd_databricks.handlers.offline_conversion as oc_handler -from ttd_data import DataClient -from ttd_data.uid2 import UID2Resolution from ttd_databricks_python.ttd_databricks.contexts import AdvertiserContext, OfflineConversionContext from ttd_databricks_python.ttd_databricks.endpoints import TTDEndpoint from ttd_databricks_python.ttd_databricks.id_types import is_raw_pii_id_type @@ -31,7 +31,6 @@ from ttd_databricks_python.ttd_databricks.ttd_client import TtdDatabricksClient from ttd_databricks_python.ttd_databricks.utils import attach_resolutions - # --------------------------------------------------------------------------- # # id_types normalization # # --------------------------------------------------------------------------- # @@ -306,7 +305,7 @@ def test_raises_with_alter_table_hint_when_column_missing(self) -> None: def _make_client() -> TtdDatabricksClient: - return TtdDatabricksClient(data_api_client=MagicMock(spec=DataClient), api_token="test-token") + return TtdDatabricksClient(data_api_client=MagicMock(spec=DataClient)) def _make_rows(*dicts: dict) -> list[MagicMock]: @@ -355,9 +354,7 @@ def test_call_api_attaches_uid2_resolutions_array_for_offline_conversion() -> No ) with patch("importlib.import_module", return_value=mock_handler): - results = client._call_api( - OfflineConversionContext(data_provider_id="dp"), rows, batch_index=0 - ) + results = client._call_api(OfflineConversionContext(data_provider_id="dp"), rows, batch_index=0) assert len(results[0][UID2_RESOLUTIONS_COLUMN]) == 1 assert results[0][UID2_RESOLUTIONS_COLUMN][0]["current_uid2"] == "uid2-y" @@ -420,4 +417,4 @@ def test_batch_process_config_is_derived_from_data_api_client() -> None: assert client._data_api_client.config.uid2_config is uid2_cfg assert client._data_api_client.config.retry_config is retry_cfg - + assert client._data_api_client.config.ttd_auth == "tok" diff --git a/ttd_databricks_python/ttd_databricks/batching.py b/ttd_databricks_python/ttd_databricks/batching.py index dfdedc0..e9716e6 100644 --- a/ttd_databricks_python/ttd_databricks/batching.py +++ b/ttd_databricks_python/ttd_databricks/batching.py @@ -39,11 +39,10 @@ def process_partitions( df: DataFrame, batch_size: int, output_schema: StructType, - api_token: str, context: TTDContext, + client_config: ClientConfig, parallelism: Optional[int] = None, data_load_trace_id: Optional[str] = None, - client_config: Optional[ClientConfig] = None, ) -> DataFrame: """Process all rows through the API using a single mapInPandas pass. @@ -57,7 +56,7 @@ def process_partitions( on server responses. Falls back to _DEFAULT_PARALLELISM on serverless / Spark Connect where sparkContext is unavailable. - client_config is a snapshot of the driver DataClient's settings (server_url, + client_config is a snapshot of the driver DataClient's settings (ttd_auth, server_url, retry_config, timeout_ms, uid2_config), used to rebuild an equivalent DataClient per worker. @@ -81,7 +80,7 @@ def partition_to_results(pandas_df_iter: Iterable[pd.DataFrame]) -> Iterator[pd. import pandas as pd from ttd_data import DataClient - from ttd_databricks_python.ttd_databricks.constants import ABORTED_ERROR_CODE, DEFAULT_RETRY_CONFIG + from ttd_databricks_python.ttd_databricks.constants import ABORTED_ERROR_CODE from ttd_databricks_python.ttd_databricks.utils import ( attach_resolutions, classify_failure, @@ -92,11 +91,9 @@ def partition_to_results(pandas_df_iter: Iterable[pd.DataFrame]) -> Iterator[pd. global _worker_client if _worker_client is None: # Workers rebuild the client from the picklable client_config snapshot; - # DataClient itself can't be cloudpickled. - if client_config is None: - _worker_client = DataClient(timeout_ms=10_000, retry_config=DEFAULT_RETRY_CONFIG) - else: - _worker_client = DataClient.from_config(client_config) + # DataClient itself can't be cloudpickled. The snapshot carries ttd_auth, so the + # rebuilt client authenticates exactly as the driver's client does. + _worker_client = DataClient.from_config(client_config) client = _worker_client handler = importlib.import_module(handler_module) @@ -139,9 +136,7 @@ def abort(error_code: str, error_message: str) -> pd.DataFrame: try: items = handler.build_items(batch_rows) raw_pii_ids_per_row = handler.collect_raw_pii_ids_per_row(batch_rows) - failed_lines, identity_resolutions = handler.call_api( - client, context, items, api_token, data_load_trace_id - ) + failed_lines, identity_resolutions = handler.call_api(client, context, items, data_load_trace_id) row_results = parse_failed_lines(failed_lines, len(batch_rows)) attach_resolutions(row_results, raw_pii_ids_per_row, identity_resolutions) except Exception as exc: diff --git a/ttd_databricks_python/ttd_databricks/handlers/advertiser.py b/ttd_databricks_python/ttd_databricks/handlers/advertiser.py index bc8fe1d..8375592 100644 --- a/ttd_databricks_python/ttd_databricks/handlers/advertiser.py +++ b/ttd_databricks_python/ttd_databricks/handlers/advertiser.py @@ -53,7 +53,6 @@ def call_api( client: DataClient, context: AdvertiserContext, items: list[AdvertiserDataItem], - api_token: str, data_load_trace_id: Optional[str] = None, ) -> tuple[list[Any], dict[str, UID2Resolution]]: """Call ingest_advertiser_data. Returns (failed_lines, identity_resolutions). @@ -72,7 +71,6 @@ def call_api( try: response = client.advertiser.ingest_advertiser_data( advertiser_id=context.advertiser_id, - ttd_auth=api_token, data_provider_id=context.data_provider_id if context.data_provider_id is not None else UNSET, items=items, data_load_trace_id=data_load_trace_id if data_load_trace_id is not None else UNSET, diff --git a/ttd_databricks_python/ttd_databricks/handlers/deletion_optout_advertiser.py b/ttd_databricks_python/ttd_databricks/handlers/deletion_optout_advertiser.py index 74a1aa2..adc870b 100644 --- a/ttd_databricks_python/ttd_databricks/handlers/deletion_optout_advertiser.py +++ b/ttd_databricks_python/ttd_databricks/handlers/deletion_optout_advertiser.py @@ -37,7 +37,6 @@ def call_api( client: DataClient, context: DeletionOptOutAdvertiserContext, items: list[PartnerDsrDataItem], - api_token: str, data_load_trace_id: Optional[str] = None, ) -> tuple[list[Any], dict[str, UID2Resolution]]: """Call data_subject_request_advertiser_data. @@ -49,7 +48,6 @@ def call_api( try: response = client.deletion_opt_out.data_subject_request_advertiser_data( - ttd_auth=api_token, advertiser_id=context.advertiser_id, data_provider_id=context.data_provider_id if context.data_provider_id is not None else UNSET, items=items, diff --git a/ttd_databricks_python/ttd_databricks/handlers/deletion_optout_merchant.py b/ttd_databricks_python/ttd_databricks/handlers/deletion_optout_merchant.py index d96ca69..345769d 100644 --- a/ttd_databricks_python/ttd_databricks/handlers/deletion_optout_merchant.py +++ b/ttd_databricks_python/ttd_databricks/handlers/deletion_optout_merchant.py @@ -37,7 +37,6 @@ def call_api( client: DataClient, context: DeletionOptOutMerchantContext, items: list[PartnerDsrDataItem], - api_token: str, data_load_trace_id: Optional[str] = None, ) -> tuple[list[Any], dict[str, UID2Resolution]]: """Call data_subject_request_merchant_data. @@ -49,7 +48,6 @@ def call_api( try: response = client.deletion_opt_out.data_subject_request_merchant_data( - ttd_auth=api_token, merchant_id=context.merchant_id, items=items, data_load_trace_id=data_load_trace_id if data_load_trace_id is not None else UNSET, diff --git a/ttd_databricks_python/ttd_databricks/handlers/deletion_optout_thirdparty.py b/ttd_databricks_python/ttd_databricks/handlers/deletion_optout_thirdparty.py index 25aa0ae..97e544c 100644 --- a/ttd_databricks_python/ttd_databricks/handlers/deletion_optout_thirdparty.py +++ b/ttd_databricks_python/ttd_databricks/handlers/deletion_optout_thirdparty.py @@ -37,7 +37,6 @@ def call_api( client: DataClient, context: DeletionOptOutThirdPartyContext, items: list[PartnerDsrDataItem], - api_token: str, data_load_trace_id: Optional[str] = None, ) -> tuple[list[Any], dict[str, UID2Resolution]]: """Call data_subject_request_third_party_data. @@ -49,7 +48,6 @@ def call_api( try: response = client.deletion_opt_out.data_subject_request_third_party_data( - ttd_auth=api_token, data_provider_id=context.data_provider_id, brand_id=context.brand_id if context.brand_id is not None else UNSET, items=items, diff --git a/ttd_databricks_python/ttd_databricks/handlers/offline_conversion.py b/ttd_databricks_python/ttd_databricks/handlers/offline_conversion.py index 759afb9..c6779ec 100644 --- a/ttd_databricks_python/ttd_databricks/handlers/offline_conversion.py +++ b/ttd_databricks_python/ttd_databricks/handlers/offline_conversion.py @@ -116,7 +116,6 @@ def call_api( client: DataClient, context: OfflineConversionContext, items: list[OfflineConversionDataItem], - api_token: str, data_load_trace_id: Optional[str] = None, ) -> tuple[list[Any], dict[str, UID2Resolution]]: """Call ingest_offline_conversion_data. Returns (failed_lines, identity_resolutions). @@ -136,7 +135,6 @@ def call_api( try: response = client.offline_conversion.ingest_offline_conversion_data( - ttd_auth=api_token, data_provider_id=context.data_provider_id, user_id_array_metadata_format=["type", "id"] if has_user_id_array else UNSET, items=items, diff --git a/ttd_databricks_python/ttd_databricks/handlers/third_party.py b/ttd_databricks_python/ttd_databricks/handlers/third_party.py index a3fefbb..c679311 100644 --- a/ttd_databricks_python/ttd_databricks/handlers/third_party.py +++ b/ttd_databricks_python/ttd_databricks/handlers/third_party.py @@ -53,7 +53,6 @@ def call_api( client: DataClient, context: ThirdPartyContext, items: list[ThirdPartyDataItem], - api_token: str, data_load_trace_id: Optional[str] = None, ) -> tuple[list[Any], dict[str, UID2Resolution]]: """Call ingest_third_party_data. Returns (failed_lines, identity_resolutions). @@ -71,7 +70,6 @@ def call_api( try: response = client.third_party.ingest_third_party_data( - ttd_auth=api_token, data_provider_id=context.data_provider_id, items=items, is_user_id_already_hashed=context.is_user_id_already_hashed, diff --git a/ttd_databricks_python/ttd_databricks/schemas/advertiser.py b/ttd_databricks_python/ttd_databricks/schemas/advertiser.py index f689f05..03d7c4d 100644 --- a/ttd_databricks_python/ttd_databricks/schemas/advertiser.py +++ b/ttd_databricks_python/ttd_databricks/schemas/advertiser.py @@ -56,7 +56,7 @@ def input_schema() -> StructType: """ from pyspark.sql.types import ( DoubleType, - IntegerType, + LongType, StringType, StructField, StructType, @@ -72,7 +72,7 @@ def input_schema() -> StructType: # Optional StructField("cookie_mapping_partner_id", StringType(), True), StructField("timestamp_utc", TimestampType(), True), - StructField("ttl_in_minutes", IntegerType(), True), + StructField("ttl_in_minutes", LongType(), True), StructField("base_bid_cpm", DoubleType(), True), StructField("base_bid_cpm_metadata", StringType(), True), StructField("bid_factor", DoubleType(), True), diff --git a/ttd_databricks_python/ttd_databricks/schemas/third_party.py b/ttd_databricks_python/ttd_databricks/schemas/third_party.py index 29546dd..9572bd7 100644 --- a/ttd_databricks_python/ttd_databricks/schemas/third_party.py +++ b/ttd_databricks_python/ttd_databricks/schemas/third_party.py @@ -49,7 +49,7 @@ def input_schema() -> StructType: ttl_in_minutes → ThirdPartyData.TtlInMinutes """ from pyspark.sql.types import ( - IntegerType, + LongType, StringType, StructField, StructType, @@ -65,6 +65,6 @@ def input_schema() -> StructType: # Optional StructField("cookie_mapping_partner_id", StringType(), True), StructField("timestamp_utc", TimestampType(), True), - StructField("ttl_in_minutes", IntegerType(), True), + StructField("ttl_in_minutes", LongType(), True), ] ) diff --git a/ttd_databricks_python/ttd_databricks/ttd_client.py b/ttd_databricks_python/ttd_databricks/ttd_client.py index e24e086..cb8bc1a 100644 --- a/ttd_databricks_python/ttd_databricks/ttd_client.py +++ b/ttd_databricks_python/ttd_databricks/ttd_client.py @@ -7,7 +7,8 @@ # ttd-data is the external SDK for the TTD Data API. # Install via: pip install ttd-data -# DataClient is the main HTTP client. The TTD-Auth token is passed per API call. +# DataClient is the main HTTP client. It is given the TTD-Auth token when it is constructed, +# and authenticates every request the SDK makes with it. from ttd_data import DataClient from ttd_data.types import OptionalNullable from ttd_data.uid2 import UID2Config @@ -31,7 +32,7 @@ class TtdDatabricksClient: Supports two usage patterns: 1. Dependency Injection (recommended for testing): - client = TtdDatabricksClient(data_api_client=DataClient(), api_token="...") + client = TtdDatabricksClient(data_api_client=DataClient(ttd_auth="...")) 2. Factory method (convenience for notebooks): client = TtdDatabricksClient.from_params(api_token="...") @@ -40,7 +41,6 @@ class TtdDatabricksClient: def __init__( self, data_api_client: DataClient, - api_token: str, spark: Optional[SparkSession] = None, ) -> None: """ @@ -49,16 +49,23 @@ def __init__( Auto-detect spark session unless explicitly provided. Dependency Injection pattern: - - data_api_client: Required. Injected DataClient instance (from ttd-data package). - `batch_process` rebuilds an equivalent DataClient per worker from - `data_api_client.config` — a single place to configure uid2/retry settings. - - api_token: Required. TTD-Auth token passed to each API call for authentication. + - data_api_client: Required. Injected DataClient instance (from ttd-data package), + built with `ttd_auth` set to your TTD API token. `batch_process` rebuilds an + equivalent DataClient per worker from `data_api_client.config` — a single place + to configure the token and the uid2/retry settings. Note: DataClient is from the external ttd-data package. For factory pattern (creating clients from tokens), use `from_params()` class method. """ + from ttd_databricks_python.ttd_databricks.exceptions import TTDConfigurationError + + if data_api_client.config.ttd_auth is None: + raise TTDConfigurationError( + "data_api_client was created without a TTD API token. " + 'Build it as DataClient(ttd_auth=""), or use TtdDatabricksClient.from_params().' + ) + self._data_api_client = data_api_client - self._api_token = api_token self._spark = spark @classmethod @@ -88,12 +95,13 @@ def from_params( Returns: TtdDatabricksClient instance with internally created DataClient. """ data_api_client = DataClient( + ttd_auth=api_token, uid2_config=uid2_config, retry_config=retry_config, server_url=server_url, timeout_ms=timeout_ms, ) - return cls(data_api_client=data_api_client, api_token=api_token, spark=spark) + return cls(data_api_client=data_api_client, spark=spark) # ------------------------------------------------------------------ # Ad hoc mode @@ -248,14 +256,13 @@ def batch_process( output_schema = get_output_schema(df.schema) self._validate_output_table_schema(spark, output_table, output_schema) output_df = process_partitions( - df, - batch_size, - output_schema, - self._api_token, - context, - parallelism, - data_load_trace_id, - self._data_api_client.config, + df=df, + batch_size=batch_size, + output_schema=output_schema, + context=context, + client_config=self._data_api_client.config, + parallelism=parallelism, + data_load_trace_id=data_load_trace_id, ) output_df.write.format("delta").mode("append").saveAsTable(output_table) @@ -466,7 +473,7 @@ def fail_all(error_code: str, error_message: str) -> list[dict[str, Any]]: items = handler.build_items(rows_data) raw_pii_ids_per_row = handler.collect_raw_pii_ids_per_row(rows_data) failed_lines, identity_resolutions = handler.call_api( - self._data_api_client, context, items, self._api_token, data_load_trace_id + self._data_api_client, context, items, data_load_trace_id ) results = parse_failed_lines(failed_lines, len(rows)) attach_resolutions(results, raw_pii_ids_per_row, identity_resolutions)