Skip to content

Commit 85d09aa

Browse files
authored
feat: add auth_header to TransformContext and support ctx in lighweight (#100)
1 parent f325241 commit 85d09aa

2 files changed

Lines changed: 63 additions & 21 deletions

File tree

libs/transforms/src/transforms/api/_transform.py

Lines changed: 16 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -113,7 +113,7 @@ def _compute_spark(
113113
"""
114114
kwargs = {name: i.init_input(context).dataframe() for name, i in self.inputs.items()}
115115
if self._use_context:
116-
kwargs["ctx"] = TransformContext()
116+
kwargs["ctx"] = TransformContext(context)
117117

118118
output_df = self(**kwargs)
119119

@@ -136,7 +136,7 @@ def _compute_pandas(
136136
"""
137137
kwargs = {name: i.init_input(context).dataframe().toPandas() for name, i in self.inputs.items()}
138138
if self._use_context:
139-
kwargs["ctx"] = TransformContext()
139+
kwargs["ctx"] = TransformContext(context)
140140

141141
output_df = self(**kwargs)
142142
from foundry_dev_tools._optional.pandas import pd
@@ -159,7 +159,7 @@ def _compute_transform( # noqa: ANN202
159159
kwargs = {**inputs, **outputs}
160160

161161
if self._use_context:
162-
kwargs["ctx"] = TransformContext()
162+
kwargs["ctx"] = TransformContext(context)
163163

164164
self(**kwargs)
165165

@@ -169,17 +169,16 @@ def _compute_lightweight( # noqa: ANN202
169169
self,
170170
context: FoundryContext,
171171
):
172-
if self._use_context:
173-
msg = "Lightweight transforms do not support the context argument."
174-
raise ValueError(msg)
175-
176172
inputs = {argument_name: LightweightTransformInput(i, context) for argument_name, i in self.inputs.items()}
177173
outputs = {
178174
argument_name: LightweightTransformOutput(o, argument_name, context)
179175
for argument_name, o in self.outputs.items()
180176
}
181177
kwargs = {**inputs, **outputs}
182178

179+
if self._use_context:
180+
kwargs["ctx"] = TransformContext(context)
181+
183182
self(**kwargs)
184183

185184
return {name: i.df for name, i in outputs.items()}
@@ -188,34 +187,30 @@ def _compute_lightweight_pandas(
188187
self,
189188
context: FoundryContext,
190189
) -> pd.core.frame.DataFrame:
191-
if self._use_context:
192-
msg = "Lightweight transforms do not support the context argument."
193-
raise ValueError(msg)
194-
195190
inputs = {
196191
argument_name: LightweightTransformInput(i, context).pandas() for argument_name, i in self.inputs.items()
197192
}
193+
if self._use_context:
194+
inputs["ctx"] = TransformContext(context)
198195
return self(**inputs)
199196

200197
def _compute_lightweight_polars(
201198
self,
202199
context: FoundryContext,
203200
) -> pl.DataFrame:
204-
if self._use_context:
205-
msg = "Lightweight transforms do not support the context argument."
206-
raise ValueError(msg)
207-
208201
inputs = {
209202
argument_name: LightweightTransformInput(i, context).polars() for argument_name, i in self.inputs.items()
210203
}
204+
if self._use_context:
205+
inputs["ctx"] = TransformContext(context)
211206
return self(**inputs)
212207

213208

214209
class TransformContext:
215210
"""The TransformContext is passed to the transform function if ctx is the first argument."""
216211

217-
def __init__(self):
218-
pass
212+
def __init__(self, foundry_ctx: FoundryContext):
213+
self._foundry_ctx = foundry_ctx
219214

220215
@property
221216
def spark_session(self) -> pyspark.sql.SparkSession:
@@ -241,6 +236,10 @@ def is_incremental(self) -> bool:
241236
warnings.warn("is_incremental functionality not implemented in Foundry DevTools")
242237
return False
243238

239+
@property
240+
def auth_header(self) -> str:
241+
return f"Bearer {self._foundry_ctx.token_provider.token}"
242+
244243

245244
class TransformInput:
246245
"""TransformInput class, passed when using @transform decorator."""

tests/unit/test_transforms.py

Lines changed: 47 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -275,16 +275,17 @@ def transform_me(input1, input2, out):
275275

276276

277277
def test_transform_polars_transform_two_inputs(spark_df_return_data_one, spark_df_return_data_two, transforms_context):
278-
from transforms.api import Input, Output, transform_polars
278+
from transforms.api import Input, Output, TransformContext, transform_polars
279279

280280
@transform_polars(
281281
Output("/output/to/dataset"),
282282
input1=Input("/input1"),
283283
input2=Input("/input2"),
284284
)
285-
def transform_me(input1, input2) -> DataFrame:
285+
def transform_me(ctx, input1, input2) -> DataFrame:
286286
assert_frame_equal(input1.to_pandas(), spark_df_return_data_one.toPandas())
287287
assert_frame_equal(input2.to_pandas(), spark_df_return_data_two.toPandas())
288+
assert isinstance(ctx, TransformContext) # ctx is our TransformContext
288289
return input1.extend(input2)
289290

290291
df = transform_me.compute(transforms_context)
@@ -358,20 +359,34 @@ def transform_me(input1: pd.DataFrame) -> pd.DataFrame:
358359

359360

360361
def test_transform_pandas_one_input_with_ctx(spark_df_return_data_one, transforms_context):
361-
from transforms.api import Input, Output, TransformContext, transform_pandas
362+
from transforms.api import Input, Output, TransformContext, lightweight, transform_pandas
363+
364+
expected_df = spark_df_return_data_one.toPandas()
362365

363366
@transform_pandas(Output("/output/to/dataset"), input1=Input("/input1"))
364367
def transform_me(ctx, input1: pd.DataFrame) -> pd.DataFrame:
365368
assert isinstance(input1, pd.DataFrame)
366369
assert isinstance(ctx, TransformContext) # ctx is our TransformContext
367370
assert isinstance(ctx.spark_session, SparkSession) # ctx.spark_session is a SparkSession
368-
assert_frame_equal(input1, spark_df_return_data_one.toPandas())
371+
assert_frame_equal(input1, expected_df)
369372
return input1
370373

371374
df = transform_me.compute(transforms_context)
372375
assert isinstance(df, pd.DataFrame)
373376
assert_frame_equal(df, spark_df_return_data_one.toPandas())
374377

378+
@lightweight
379+
@transform_pandas(Output("/output/to/dataset"), input1=Input("/input1"))
380+
def transform_me(ctx, input1: pd.DataFrame) -> pd.DataFrame:
381+
assert isinstance(input1, pd.DataFrame)
382+
assert isinstance(ctx, TransformContext) # ctx is our TransformContext
383+
assert_frame_equal(input1, expected_df)
384+
return input1
385+
386+
df = transform_me.compute(transforms_context)
387+
assert isinstance(df, pd.DataFrame)
388+
assert_frame_equal(df, expected_df)
389+
375390

376391
def test_transform_pandas_two_inputs(spark_df_return_data_one, spark_df_return_data_two, transforms_context):
377392
from transforms.api import Input, Output, TransformContext, transform_pandas
@@ -502,6 +517,20 @@ def transform_me(output1, input1):
502517
assert "will have no effect" in record_message
503518

504519

520+
def test_lightweight_transforms_with_context(transforms_context):
521+
from transforms.api import Input, Output, TransformContext, lightweight, transform
522+
523+
@lightweight()
524+
@transform(
525+
output1=Output("/output/to/dataset"),
526+
input1=Input("/input1"),
527+
)
528+
def transform_me(ctx, output1, input1):
529+
assert isinstance(ctx, TransformContext) # ctx is our TransformContext
530+
531+
transform_me.compute(transforms_context)
532+
533+
505534
def test_transforms_with_incremental():
506535
from transforms.api import Input, Output, incremental, transform
507536

@@ -587,3 +616,17 @@ def transform_me(input1):
587616
return input1
588617

589618
transform_me.compute(transforms_context)
619+
620+
621+
def test_transform_context_auth_header(transforms_context):
622+
from transforms.api import Input, Output, transform_df
623+
624+
@transform_df(
625+
Output("output1"),
626+
input1=Input("/input1"),
627+
)
628+
def transform_me(ctx, input1):
629+
assert ctx.auth_header == "Bearer test_transform_context_auth_header_token"
630+
return input1
631+
632+
transform_me.compute(transforms_context)

0 commit comments

Comments
 (0)