Skip to content

Commit 084a0d7

Browse files
SudipSinhaclaude
andcommitted
fix(drift): derive fitColumns from stored metadata when omitted
Apply the same fix as PR #180 to the streaming KS test endpoints: when fitColumns is not provided, derive it from the model's stored input schema metadata instead of rejecting with HTTP 400. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com> Signed-off-by: Sudip Sinha <Sudip.Sinha@RedHat.com>
1 parent 2fd8510 commit 084a0d7

3 files changed

Lines changed: 45 additions & 26 deletions

File tree

src/endpoints/metrics/drift/kolmogorov_smirnov_streaming.py

Lines changed: 17 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -97,9 +97,13 @@ async def compute_ksteststreaming(
9797
)
9898

9999
if not request.fit_columns:
100-
raise HTTPException(
101-
status_code=HTTPStatus.BAD_REQUEST,
102-
detail="fitColumns is required - specify which features to test for drift",
100+
data_source = get_data_source()
101+
metadata = await data_source.get_metadata(request.model_id)
102+
request.fit_columns = list(metadata.input_schema.items.keys())
103+
logger.info(
104+
"fitColumns not specified, using all input columns for model %s: %s",
105+
request.model_id,
106+
request.fit_columns,
103107
)
104108

105109
try:
@@ -215,6 +219,16 @@ async def schedule_ksteststreaming(
215219
request: ApproxKSTestMetricRequest,
216220
) -> dict[str, str]:
217221
"""Schedule a recurring computation of KS Test Streaming metric."""
222+
if not request.fit_columns:
223+
data_source = get_data_source()
224+
metadata = await data_source.get_metadata(request.model_id)
225+
request.fit_columns = list(metadata.input_schema.items.keys())
226+
logger.info(
227+
"fitColumns not specified, using all input columns for model %s: %s",
228+
request.model_id,
229+
request.fit_columns,
230+
)
231+
218232
# Get the scheduler and validate availability
219233
scheduler = get_prometheus_scheduler()
220234
if not scheduler:

tests/endpoints/metrics/drift/factory.py

Lines changed: 16 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -73,16 +73,17 @@ def make_compute_endpoint_test(
7373
def test_impl(_: object, mock_ds: MagicMock) -> None:
7474
"""Test compute endpoint returns valid response structure."""
7575
# Create sample dataframe (Pandas or Polars based on df_type)
76-
sample_df = _create_sample_dataframe(
77-
request_payload.get("fitColumns", ["feature1"]),
78-
df_type=df_type,
79-
)
76+
columns = request_payload.get("fitColumns", ["feature1"])
77+
sample_df = _create_sample_dataframe(columns, df_type=df_type)
8078

8179
# Mock data source
8280
mock_data_source = MagicMock()
8381
mock_data_source.get_dataframe_by_tag = AsyncMock(return_value=sample_df)
8482
mock_data_source.get_organic_dataframe = AsyncMock(return_value=sample_df)
8583
mock_data_source.get_dataframe = AsyncMock(return_value=sample_df)
84+
mock_metadata = MagicMock()
85+
mock_metadata.input_schema.items.keys.return_value = columns
86+
mock_data_source.get_metadata = AsyncMock(return_value=mock_metadata)
8687
mock_ds.return_value = mock_data_source
8788

8889
# Send request
@@ -184,7 +185,9 @@ def test_impl(_: object, mock_ds: MagicMock, mock_sched_fn: MagicMock) -> None:
184185

185186
# Mock data source (needed for KS test registration)
186187
mock_data_source = MagicMock()
187-
mock_data_source.get_metadata = AsyncMock(return_value={"feature1": "type1"})
188+
mock_sched_metadata = MagicMock()
189+
mock_sched_metadata.input_schema.items.keys.return_value = ["feature1"]
190+
mock_data_source.get_metadata = AsyncMock(return_value=mock_sched_metadata)
188191
mock_ds.return_value = mock_data_source
189192

190193
# Send request
@@ -338,17 +341,17 @@ def make_compute_endpoint_error_test(
338341
def test_impl(_: object, mock_ds: MagicMock) -> None:
339342
"""Test compute endpoint error handling."""
340343
if setup_mocks:
341-
# Create sample dataframe
342-
sample_df = _create_sample_dataframe(
343-
["feature1", "feature2"],
344-
df_type=df_type,
345-
)
344+
columns = ["feature1", "feature2"]
345+
sample_df = _create_sample_dataframe(columns, df_type=df_type)
346346

347347
# Mock data source
348348
mock_data_source = MagicMock()
349349
mock_data_source.get_dataframe_by_tag = AsyncMock(return_value=sample_df)
350350
mock_data_source.get_organic_dataframe = AsyncMock(return_value=sample_df)
351351
mock_data_source.get_dataframe = AsyncMock(return_value=sample_df)
352+
mock_metadata = MagicMock()
353+
mock_metadata.input_schema.items.keys.return_value = columns
354+
mock_data_source.get_metadata = AsyncMock(return_value=mock_metadata)
352355
mock_ds.return_value = mock_data_source
353356

354357
# Send request
@@ -430,7 +433,9 @@ def test_impl(_: object, mock_ds: MagicMock, mock_sched_fn: MagicMock) -> None:
430433

431434
# Mock data source
432435
mock_data_source = MagicMock()
433-
mock_data_source.get_metadata = AsyncMock(return_value={"feature1": "type1"})
436+
mock_sched_metadata = MagicMock()
437+
mock_sched_metadata.input_schema.items.keys.return_value = ["feature1"]
438+
mock_data_source.get_metadata = AsyncMock(return_value=mock_sched_metadata)
434439
mock_ds.return_value = mock_data_source
435440

436441
# Send request

tests/endpoints/metrics/drift/test_kolmogorov_smirnov_streaming.py

Lines changed: 12 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -126,18 +126,18 @@ class TestKSTestStreamingEndpoints:
126126
expected_error_substring="referenceTag is required",
127127
)
128128

129-
test_compute_missing_fit_columns = factory.make_compute_endpoint_error_test(
130-
metric_name="KSTestStreaming",
131-
module_path="src.endpoints.metrics.drift.kolmogorov_smirnov_streaming",
132-
endpoint_path="/metrics/drift/ksteststreaming",
133-
client=client,
134-
request_payload={
135-
"modelId": "test-model",
136-
"referenceTag": "baseline",
137-
# Missing fitColumns
138-
},
139-
expected_status_code=HTTPStatus.BAD_REQUEST,
140-
expected_error_substring="fitColumns is required",
129+
test_compute_missing_fit_columns_derives_from_metadata = (
130+
factory.make_compute_endpoint_test(
131+
metric_name="KSTestStreaming",
132+
module_path="src.endpoints.metrics.drift.kolmogorov_smirnov_streaming",
133+
endpoint_path="/metrics/drift/ksteststreaming",
134+
client=client,
135+
request_payload={
136+
"modelId": "test-model",
137+
"referenceTag": "baseline",
138+
},
139+
expected_response_keys=["status", "value", "drift_detected"],
140+
)
141141
)
142142

143143
test_compute_invalid_feature = factory.make_compute_endpoint_error_test(

0 commit comments

Comments
 (0)