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..a04f57a 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.0,<0.4.0", "pandas>=1.0.5", "pyarrow>=4.0.0", "setuptools>=63.4.1", @@ -54,6 +54,7 @@ select = [ ] ignore = [ "UP045", # prefer Optional[X] over X | None + "UP007", # prefer Union[X, Y] over X | Y ] [tool.ruff.lint.isort] 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..6b7ec5c 100644 --- a/tests/unit/test_call_api.py +++ b/tests/unit/test_call_api.py @@ -30,7 +30,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_process_partitions.py b/tests/unit/test_process_partitions.py index 3805292..ec79dea 100644 --- a/tests/unit/test_process_partitions.py +++ b/tests/unit/test_process_partitions.py @@ -17,11 +17,12 @@ from collections.abc import Iterator from http import HTTPStatus from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from typing import Optional import pytest from pyspark.sql import SparkSession from pyspark.sql.types import StringType, StructField, StructType, TimestampType -from ttd_data import ClientConfig +from ttd_data import DataClient from ttd_databricks_python.ttd_databricks.batching import process_partitions from ttd_databricks_python.ttd_databricks.contexts import AdvertiserContext @@ -29,16 +30,15 @@ pytestmark = pytest.mark.spark -# Pinned so request counts stay exact: retry_config=None leaves the SDK's retry wrapper -# off, so each batch makes exactly one call even when the stub returns a retryable 5xx. +_TOKEN = "not-a-real-token" + +# Snapshotted off a real DataClient rather than hand-built, so the test tracks whatever +# fields ClientConfig carries in the installed ttd-data. +# retry_config=None leaves the SDK's retry wrapper off, keeping request counts exact: each +# batch makes exactly one call even when the stub returns a retryable 5xx. # Shared by both tests — workers cache one DataClient per process, so a differing config # in a second test would be silently ignored. -_NO_RETRY_CLIENT_CONFIG = ClientConfig( - server_url=None, - retry_config=None, - timeout_ms=10_000, - uid2_config=None, -) +_NO_RETRY_CLIENT_CONFIG = DataClient(ttd_auth=_TOKEN, retry_config=None, timeout_ms=10_000).config class _StubHandler(BaseHTTPRequestHandler): @@ -46,6 +46,7 @@ class _StubHandler(BaseHTTPRequestHandler): status_code = 500 request_count = 0 + auth_headers: list[Optional[str]] = [] # 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 +57,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 +115,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 +129,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) == {_TOKEN} @pytest.mark.parametrize( @@ -149,7 +154,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 +186,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..10dbd8b 100644 --- a/tests/unit/test_uid2_resolutions.py +++ b/tests/unit/test_uid2_resolutions.py @@ -306,7 +306,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]: @@ -420,4 +420,5 @@ 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/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)