Skip to content

Commit bae07fc

Browse files
committed
fix: Use do_exchange for offline server reads to support HPA
The remote offline server used a two-phase Arrow Flight protocol for read operations: do_put (stores entity data in an in-memory dict) followed by get_flight_info + do_get (retrieves results). With HPA or multiple replicas, these separate gRPC calls can be load-balanced to different pods, causing 'Flight not found' errors because the in-memory flights dict is per-pod. Replace the two-phase read path with Arrow Flight's do_exchange RPC, which handles both the upload and result download in a single bidirectional gRPC stream. This guarantees both phases hit the same pod, making the offline server compatible with horizontal scaling. Signed-off-by: ntkathole <nikhilkathole2683@gmail.com>
1 parent fa8f06b commit bae07fc

3 files changed

Lines changed: 329 additions & 21 deletions

File tree

‎sdk/python/feast/infra/offline_stores/remote.py‎

Lines changed: 44 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -65,6 +65,14 @@ def do_put(
6565
):
6666
return super().do_put(descriptor, schema, options)
6767

68+
@arrow_client_error_handling_decorator
69+
def do_exchange(
70+
self,
71+
descriptor: FlightDescriptor,
72+
options: FlightCallOptions = None,
73+
):
74+
return super().do_exchange(descriptor, options)
75+
6876
@arrow_client_error_handling_decorator
6977
def list_flights(self, criteria: bytes = b"", options: FlightCallOptions = None):
7078
return super().list_flights(criteria, options)
@@ -535,14 +543,42 @@ def _send_retrieve_remote(
535543
table: Optional[pa.Table],
536544
client: FeastFlightClient,
537545
):
538-
command_descriptor = _call_put(
539-
api,
540-
api_parameters,
541-
client,
542-
entity_df,
543-
table,
544-
)
545-
return _call_get(client, command_descriptor)
546+
return _call_exchange(api, api_parameters, client, entity_df, table)
547+
548+
549+
def _call_exchange(
550+
api: str,
551+
api_parameters: Dict[str, Any],
552+
client: FeastFlightClient,
553+
entity_df: Optional[Union[pd.DataFrame, str]],
554+
table: Optional[pa.Table],
555+
) -> pa.Table:
556+
"""Execute a read API via a single ``do_exchange`` bidirectional stream.
557+
558+
This replaces the legacy two-phase ``do_put`` → ``get_flight_info`` →
559+
``do_get`` flow that relied on in-memory state (``self.flights``) on the
560+
server. Because both the upload and the result download happen on the
561+
same gRPC stream, the request is always handled by the same server pod —
562+
making the remote offline server compatible with HPA / multiple replicas.
563+
"""
564+
command_id = str(uuid.uuid4())
565+
command = {"command_id": command_id, "api": api, **api_parameters}
566+
567+
descriptor = fl.FlightDescriptor.for_command(json.dumps(command))
568+
569+
upload_table: pa.Table
570+
if entity_df is not None and not isinstance(entity_df, str):
571+
upload_table = pa.Table.from_pandas(entity_df)
572+
elif table is not None:
573+
upload_table = table
574+
else:
575+
upload_table = _create_empty_table()
576+
577+
writer, reader = client.do_exchange(descriptor)
578+
writer.begin(upload_table.schema)
579+
writer.write_table(upload_table)
580+
writer.done_writing()
581+
return read_all(reader)
546582

547583

548584
def _call_get(

‎sdk/python/feast/offline_server.py‎

Lines changed: 111 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -285,19 +285,7 @@ def do_get(self, context: fl.ServerCallContext, ticket: fl.Ticket):
285285
logger.debug(f"get command is {command}")
286286
logger.debug(f"requested api is {api}")
287287
try:
288-
if api == OfflineServer.get_historical_features.__name__:
289-
table = self.get_historical_features(command, key).to_arrow()
290-
elif api == OfflineServer.pull_all_from_table_or_query.__name__:
291-
table = self.pull_all_from_table_or_query(command).to_arrow()
292-
elif api == OfflineServer.pull_latest_from_table_or_query.__name__:
293-
table = self.pull_latest_from_table_or_query(command).to_arrow()
294-
elif (
295-
api
296-
== OfflineServer.get_table_column_names_and_types_from_data_source.__name__
297-
):
298-
table = self.get_table_column_names_and_types_from_data_source(command)
299-
else:
300-
raise NotImplementedError
288+
table = self._execute_read_api(api, command, key)
301289
except Exception as e:
302290
logger.exception(e)
303291
traceback.print_exc()
@@ -307,6 +295,116 @@ def do_get(self, context: fl.ServerCallContext, ticket: fl.Ticket):
307295
del self.flights[key]
308296
return fl.RecordBatchStream(table)
309297

298+
@inject_user_details_decorator
299+
@arrow_server_error_handling_decorator
300+
def do_exchange(
301+
self,
302+
context: fl.ServerCallContext,
303+
descriptor: fl.FlightDescriptor,
304+
reader: fl.MetadataRecordBatchReader,
305+
writer: fl.MetadataRecordBatchWriter,
306+
):
307+
"""Handle read APIs in a single bidirectional stream.
308+
309+
Unlike the legacy do_put → get_flight_info → do_get flow (which stores
310+
intermediate state in ``self.flights`` between calls), ``do_exchange``
311+
receives the entity data **and** returns the query results within one
312+
gRPC stream. This makes the offline server compatible with multiple
313+
replicas / HPA because no cross-call in-memory state is required.
314+
"""
315+
key = OfflineServer.descriptor_to_key(descriptor)
316+
command = json.loads(key[1])
317+
self._validate_do_get_parameters(command)
318+
api = command["api"]
319+
320+
logger.debug(f"do_exchange: api={api}, command={command}")
321+
322+
# Read the entity data sent by the client on this stream.
323+
data = reader.read_all()
324+
325+
# For get_historical_features the entity table must be converted to a
326+
# pandas DataFrame and passed in via the ``key`` mechanism. Other read
327+
# APIs do not use entity data so we can call them directly.
328+
try:
329+
if api == OfflineServer.get_historical_features.__name__:
330+
entity_df = pa.Table.to_pandas(data)
331+
if len(entity_df.columns) == 1 and "key" in entity_df.columns:
332+
entity_df = None
333+
if entity_df is None and "entity_df_sql" in command:
334+
entity_df = command["entity_df_sql"]
335+
table = self._get_historical_features_direct(
336+
command, entity_df
337+
).to_arrow()
338+
else:
339+
table = self._execute_read_api(api, command, key=None)
340+
except Exception as e:
341+
logger.exception(e)
342+
traceback.print_exc()
343+
raise e
344+
345+
writer.begin(table.schema)
346+
writer.write_table(table)
347+
348+
def _execute_read_api(
349+
self, api: str, command: dict, key: Optional[str] = None
350+
) -> pa.Table:
351+
"""Dispatch a read API call and return the result as an Arrow table."""
352+
if api == OfflineServer.get_historical_features.__name__:
353+
return self.get_historical_features(command, key).to_arrow()
354+
elif api == OfflineServer.pull_all_from_table_or_query.__name__:
355+
return self.pull_all_from_table_or_query(command).to_arrow()
356+
elif api == OfflineServer.pull_latest_from_table_or_query.__name__:
357+
return self.pull_latest_from_table_or_query(command).to_arrow()
358+
elif (
359+
api
360+
== OfflineServer.get_table_column_names_and_types_from_data_source.__name__
361+
):
362+
return self.get_table_column_names_and_types_from_data_source(command)
363+
else:
364+
raise NotImplementedError(f"Unknown read API: {api}")
365+
366+
def _get_historical_features_direct(self, command: dict, entity_df):
367+
"""Run get_historical_features without relying on self.flights."""
368+
self._validate_get_historical_features_parameters(command, key=None)
369+
370+
feature_view_names = command["feature_view_names"]
371+
name_aliases = command["name_aliases"]
372+
feature_refs = command["feature_refs"]
373+
project = command["project"]
374+
full_feature_names = command["full_feature_names"]
375+
376+
feature_views = self.list_feature_views_by_name(
377+
feature_view_names=feature_view_names,
378+
name_aliases=name_aliases,
379+
project=project,
380+
)
381+
382+
for feature_view in feature_views:
383+
assert_permissions(
384+
resource=feature_view, actions=[AuthzedAction.READ_OFFLINE]
385+
)
386+
387+
kwargs = {}
388+
if "start_date" in command and command["start_date"] is not None:
389+
kwargs["start_date"] = utils.make_tzaware(
390+
datetime.fromisoformat(command["start_date"])
391+
)
392+
if "end_date" in command and command["end_date"] is not None:
393+
kwargs["end_date"] = utils.make_tzaware(
394+
datetime.fromisoformat(command["end_date"])
395+
)
396+
397+
return self.offline_store.get_historical_features(
398+
config=self.store.config,
399+
feature_views=feature_views,
400+
feature_refs=feature_refs,
401+
entity_df=entity_df,
402+
registry=self.store.registry,
403+
project=project,
404+
full_feature_names=full_feature_names,
405+
**kwargs,
406+
)
407+
310408
def _validate_offline_write_batch_parameters(self, command: dict):
311409
assert "feature_view_names" in command, (
312410
"feature_view_names is a mandatory parameter"

‎sdk/python/tests/unit/test_offline_server.py‎

Lines changed: 174 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,15 +1,18 @@
1+
import json
12
import os
23
import subprocess
34
import sys
45
import textwrap
56
from unittest.mock import MagicMock, mock_open, patch
67

78
import assertpy
9+
import pyarrow as pa
810
import pytest
911

1012
from feast.infra.offline_stores.remote import (
1113
RemoteOfflineStore,
1214
RemoteOfflineStoreConfig,
15+
_call_exchange,
1316
_create_retrieval_metadata,
1417
)
1518
from feast.offline_server import (
@@ -208,3 +211,174 @@ def tracking_import(name, *args, **kwargs):
208211
assert result.returncode == 0, (
209212
f"Subprocess failed:\nstdout: {result.stdout}\nstderr: {result.stderr}"
210213
)
214+
215+
216+
# ---------------------------------------------------------------------------
217+
# do_exchange tests — HPA-safe single-stream read path
218+
# ---------------------------------------------------------------------------
219+
220+
221+
def test_do_exchange_get_historical_features():
222+
"""do_exchange delegates to _get_historical_features_direct for
223+
get_historical_features and writes the result table back."""
224+
import pyarrow.flight as fl
225+
226+
result_table = pa.table({"col": [1, 2, 3]})
227+
mock_job = MagicMock()
228+
mock_job.to_arrow.return_value = result_table
229+
mock_offline_store = MagicMock()
230+
mock_offline_store.get_historical_features.return_value = mock_job
231+
232+
mock_store = MagicMock()
233+
mock_store.config.project = "test"
234+
235+
server = MagicMock(spec=OfflineServer)
236+
server.offline_store = mock_offline_store
237+
server.store = mock_store
238+
server.flights = {}
239+
server.list_feature_views_by_name.return_value = []
240+
server._validate_do_get_parameters = (
241+
OfflineServer._validate_do_get_parameters.__get__(server)
242+
)
243+
server._execute_read_api = OfflineServer._execute_read_api.__get__(server)
244+
server._get_historical_features_direct = (
245+
OfflineServer._get_historical_features_direct.__get__(server)
246+
)
247+
server._validate_get_historical_features_parameters = (
248+
OfflineServer._validate_get_historical_features_parameters.__get__(server)
249+
)
250+
server.get_historical_features = OfflineServer.get_historical_features.__get__(
251+
server
252+
)
253+
254+
command = {
255+
"api": "get_historical_features",
256+
"command_id": "test-123",
257+
"feature_view_names": [],
258+
"name_aliases": [],
259+
"feature_refs": ["driver_hourly_stats:conv_rate"],
260+
"project": "test",
261+
"full_feature_names": False,
262+
}
263+
descriptor = fl.FlightDescriptor.for_command(json.dumps(command))
264+
265+
entity_table = pa.table({"key": ["mock_key"]})
266+
mock_reader = MagicMock()
267+
mock_reader.read_all.return_value = entity_table
268+
269+
mock_writer = MagicMock()
270+
271+
OfflineServer.do_exchange.__wrapped__.__wrapped__(
272+
server, MagicMock(), descriptor, mock_reader, mock_writer
273+
)
274+
275+
mock_writer.begin.assert_called_once_with(result_table.schema)
276+
mock_writer.write_table.assert_called_once_with(result_table)
277+
278+
279+
def test_do_exchange_pull_all_from_table_or_query():
280+
"""do_exchange delegates to pull_all_from_table_or_query correctly."""
281+
import pyarrow.flight as fl
282+
283+
result_table = pa.table({"feature": [10, 20]})
284+
mock_job = MagicMock()
285+
mock_job.to_arrow.return_value = result_table
286+
287+
server = MagicMock(spec=OfflineServer)
288+
server.flights = {}
289+
server._validate_do_get_parameters = (
290+
OfflineServer._validate_do_get_parameters.__get__(server)
291+
)
292+
server._execute_read_api = OfflineServer._execute_read_api.__get__(server)
293+
server.pull_all_from_table_or_query.return_value = mock_job
294+
295+
command = {
296+
"api": "pull_all_from_table_or_query",
297+
"command_id": "test-456",
298+
"data_source_name": "ds",
299+
"join_key_columns": [],
300+
"feature_name_columns": [],
301+
"timestamp_field": "ts",
302+
"created_timestamp_column": "",
303+
"start_date": "2021-01-01T00:00:00",
304+
"end_date": "2021-12-31T00:00:00",
305+
}
306+
descriptor = fl.FlightDescriptor.for_command(json.dumps(command))
307+
308+
mock_reader = MagicMock()
309+
mock_reader.read_all.return_value = pa.table({"key": ["mock_key"]})
310+
mock_writer = MagicMock()
311+
312+
OfflineServer.do_exchange.__wrapped__.__wrapped__(
313+
server, MagicMock(), descriptor, mock_reader, mock_writer
314+
)
315+
316+
server.pull_all_from_table_or_query.assert_called_once()
317+
mock_writer.begin.assert_called_once_with(result_table.schema)
318+
mock_writer.write_table.assert_called_once_with(result_table)
319+
320+
321+
def test_call_exchange_sends_entity_df_and_reads_result():
322+
"""_call_exchange sends entity data via do_exchange and reads the result."""
323+
import pandas as pd
324+
325+
result_table = pa.table({"result": [1, 2, 3]})
326+
327+
mock_writer = MagicMock()
328+
mock_reader = MagicMock()
329+
mock_reader._connection_retries = 0
330+
mock_reader.read_all.return_value = result_table
331+
332+
mock_client = MagicMock()
333+
mock_client.do_exchange.return_value = (mock_writer, mock_reader)
334+
335+
entity_df = pd.DataFrame(
336+
{"driver_id": [1, 2], "event_timestamp": ["2021-01-01", "2021-01-02"]}
337+
)
338+
339+
table = _call_exchange(
340+
api="get_historical_features",
341+
api_parameters={
342+
"feature_refs": ["f1"],
343+
"project": "test",
344+
"full_feature_names": False,
345+
"feature_view_names": [],
346+
"name_aliases": [],
347+
},
348+
client=mock_client,
349+
entity_df=entity_df,
350+
table=None,
351+
)
352+
353+
mock_client.do_exchange.assert_called_once()
354+
mock_writer.begin.assert_called_once()
355+
mock_writer.write_table.assert_called_once()
356+
mock_writer.done_writing.assert_called_once()
357+
assertpy.assert_that(table).is_equal_to(result_table)
358+
359+
360+
def test_call_exchange_sends_empty_table_when_no_entity_df():
361+
"""_call_exchange sends a stub table when entity_df and table are both None."""
362+
result_table = pa.table({"result": [42]})
363+
364+
mock_writer = MagicMock()
365+
mock_reader = MagicMock()
366+
mock_reader._connection_retries = 0
367+
mock_reader.read_all.return_value = result_table
368+
369+
mock_client = MagicMock()
370+
mock_client.do_exchange.return_value = (mock_writer, mock_reader)
371+
372+
table = _call_exchange(
373+
api="pull_all_from_table_or_query",
374+
api_parameters={"data_source_name": "ds"},
375+
client=mock_client,
376+
entity_df=None,
377+
table=None,
378+
)
379+
380+
mock_client.do_exchange.assert_called_once()
381+
mock_writer.begin.assert_called_once()
382+
call_args = mock_writer.write_table.call_args[0][0]
383+
assertpy.assert_that(call_args.column_names).contains("key")
384+
assertpy.assert_that(table).is_equal_to(result_table)

0 commit comments

Comments
 (0)