Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Prev Previous commit
Next Next commit
fix: Address Flink review blockers
Signed-off-by: Le Xuan An <anlx@viettel.com.vn>
  • Loading branch information
XuananLe authored and Le Xuan An committed Jun 11, 2026
commit 3af9cd7b67d911f48574bd8ca8a9eb8117eb9711
Original file line number Diff line number Diff line change
Expand Up @@ -39,10 +39,7 @@ def __init__(

def _should_join_entity_df(self) -> bool:
return isinstance(self.task, HistoricalRetrievalTask) and (
(
isinstance(self.task.entity_df, pd.DataFrame)
and not self.task.entity_df.empty
)
isinstance(self.task.entity_df, pd.DataFrame)
or (
isinstance(self.task.entity_df, str)
and bool(self.task.entity_df.strip())
Expand Down
68 changes: 51 additions & 17 deletions sdk/python/feast/infra/compute_engines/flink/nodes.py
Original file line number Diff line number Diff line change
Expand Up @@ -147,11 +147,22 @@ def _entity_value_from_dataframe(
table_env: Any,
entity_df: pd.DataFrame,
split_num: int,
join_keys: List[str],
) -> tuple[Any, List[str], str]:
entity_df = entity_df.copy()
if entity_df.empty:
for join_key in join_keys:
if join_key not in entity_df.columns:
entity_df[join_key] = pd.Series(dtype="object")
entity_ts_col = find_entity_timestamp_column(list(entity_df.columns))
if entity_ts_col is None:
entity_ts_col = ENTITY_TS_ALIAS
entity_df[entity_ts_col] = pd.Series(dtype="datetime64[ns]")
else:
entity_schema = dict(zip(entity_df.columns, entity_df.dtypes))
entity_ts_col = infer_entity_timestamp_column(entity_schema)

entity_df[ENTITY_ROW_ID] = range(len(entity_df))
entity_schema = dict(zip(entity_df.columns, entity_df.dtypes))
entity_ts_col = infer_entity_timestamp_column(entity_schema)
if entity_ts_col != ENTITY_TS_ALIAS:
entity_df = entity_df.rename(columns={entity_ts_col: ENTITY_TS_ALIAS})
return (
Expand Down Expand Up @@ -219,14 +230,48 @@ def _entity_value_from_context(
join_keys: List[str],
) -> tuple[Any, List[str], str]:
if isinstance(context.entity_df, pd.DataFrame):
return _entity_value_from_dataframe(table_env, context.entity_df, split_num)
return _entity_value_from_dataframe(
table_env, context.entity_df, split_num, join_keys
)
if isinstance(context.entity_df, str):
return _entity_value_from_sql(table_env, context.entity_df, join_keys)
raise TypeError(
"FlinkComputeEngine entity_df must be a pandas DataFrame, SQL string, or None."
)


def _retrieval_job_to_flink_table(
retrieval_job: Any,
table_env: Any,
split_num: int,
) -> tuple[Any, List[str]]:
to_flink_table = getattr(retrieval_job, "to_flink_table", None)
if callable(to_flink_table):
flink_table = to_flink_table(table_env)
columns = _get_columns_from_schema(flink_table)
if columns is None:
raise ValueError(
"Could not infer columns for source Flink table returned by "
"RetrievalJob.to_flink_table(table_env)."
)
return flink_table, columns

if not hasattr(retrieval_job, "to_arrow"):
raise TypeError(
"FlinkComputeEngine source reads require a RetrievalJob with either "
"to_flink_table(table_env) or to_arrow()."
)

arrow_table = retrieval_job.to_arrow()
if not isinstance(arrow_table, pa.Table):
raise TypeError(
"RetrievalJob.to_arrow() must return a pyarrow.Table for "
"FlinkComputeEngine source reads."
)
columns = list(arrow_table.column_names)
return pandas_to_flink_table(table_env, arrow_table.to_pandas(), split_num), columns


class FlinkSourceReadNode(DAGNode):
def __init__(
self,
Expand Down Expand Up @@ -254,20 +299,9 @@ def execute(self, context: ExecutionContext) -> DAGValue:
start_time=self.start_time,
end_time=self.end_time,
)
if not hasattr(retrieval_job, "to_flink_table"):
raise TypeError(
"FlinkComputeEngine source reads require RetrievalJob.to_flink_table("
"table_env). Configure an offline store retrieval job that returns "
"native PyFlink tables instead of Arrow/pandas results."
)

flink_table = retrieval_job.to_flink_table(self.table_env)
columns = _get_columns_from_schema(flink_table)
if columns is None:
raise ValueError(
"Could not infer columns for source Flink table returned by "
"RetrievalJob.to_flink_table(table_env)."
)
flink_table, columns = _retrieval_job_to_flink_table(
retrieval_job, self.table_env, self.split_num
)

if self.column_info.field_mapping:
view_name = _register_table(self.table_env, flink_table, "source_read")
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -549,7 +549,9 @@ def test_repo_config_loads_flink_batch_engine_config(tmp_path: Path) -> None:
assert config.batch_engine.pandas_split_num == 2


def test_flink_source_read_node_rejects_arrow_retrieval_jobs(tmp_path: Path) -> None:
def test_flink_source_read_node_converts_arrow_retrieval_jobs(
tmp_path: Path,
) -> None:
offline_store = MagicMock()
offline_store.pull_all_from_table_or_query.return_value = FakeRetrievalJob(
pa.Table.from_pandas(_feature_data())
Expand All @@ -564,8 +566,11 @@ def test_flink_source_read_node_rejects_arrow_retrieval_jobs(tmp_path: Path) ->
split_num=1,
)

with pytest.raises(TypeError, match="to_flink_table"):
node.execute(context)
result = node.execute(context)

assert result.format == DAGFormat.FLINK
assert result.metadata["columns"] == list(_feature_data().columns)
assert result.data.to_pandas().equals(_feature_data())


def test_flink_historical_retrieval_executes_dag_with_transformation(
Expand Down Expand Up @@ -600,7 +605,7 @@ def double_conv_rate(table: FakeFlinkTable) -> FakeFlinkTable:
)
task = HistoricalRetrievalTask(
project=config.project,
entity_df=pd.DataFrame(),
entity_df=None,
feature_view=feature_view,
full_feature_name=False,
registry=_registry(entity),
Expand All @@ -616,6 +621,42 @@ def double_conv_rate(table: FakeFlinkTable) -> FakeFlinkTable:
assert table_env.views == {}


def test_flink_historical_retrieval_with_empty_entity_df_returns_empty_result(
tmp_path: Path,
) -> None:
entity = _driver()
source = _source()
feature_view = _feature_view(source, online=False, offline=False)
config = _repo_config(tmp_path, {"type": "flink.engine", "pandas_split_num": 4})
table_env = FakeTableEnvironment()
engine = FlinkComputeEngine(
repo_config=config,
offline_store=_offline_store(_feature_data()),
online_store=MagicMock(),
table_environment=table_env,
)
task = HistoricalRetrievalTask(
project=config.project,
entity_df=pd.DataFrame(
{
"driver_id": pd.Series(dtype="int64"),
"event_timestamp": pd.Series(dtype="datetime64[ns]"),
}
),
feature_view=feature_view,
full_feature_name=False,
registry=_registry(entity),
)

job = engine.get_historical_features(_registry(entity), task)
result = job.to_df()

assert job.error() is None
assert result.empty
assert "conv_rate" in result.columns
assert table_env.created_tables[-1].empty


def test_flink_historical_retrieval_is_read_only_and_dedupes_per_entity_row(
tmp_path: Path,
) -> None:
Expand Down