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
add test
Signed-off-by: HaoXuAI <sduxuhao@gmail.com>
  • Loading branch information
haoxu0 committed Apr 9, 2025
commit 25af94e6ef099924bc60f3822475316eaa6e1f9a
31 changes: 31 additions & 0 deletions sdk/python/feast/infra/compute_engines/dag/context.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,37 @@

@dataclass
class ExecutionContext:
"""
ExecutionContext holds all runtime information required to execute a DAG plan
within a ComputeEngine. It is passed into each DAGNode during execution and
contains shared context such as configuration, registry-backed entities, runtime
data (e.g. entity_df), and DAG evaluation state.

Attributes:
project: Feast project name (namespace for features, entities, views).

repo_config: Resolved RepoConfig containing provider and store configuration.

offline_store: Reference to the configured OfflineStore implementation.
Used for loading raw feature data during materialization or retrieval.

online_store: Reference to the OnlineStore implementation.
Used during materialization to write online features.

entity_defs: List of Entity definitions fetched from the registry.
Used for resolving join keys, inferring timestamp columns, and
validating FeatureViews against schema.

entity_df: A runtime DataFrame of entity rows used during historical
retrieval (e.g. for point-in-time join). Includes entity keys and
event timestamps. This is not part of the registry and is user-supplied
for training dataset generation.

node_outputs: Internal cache of DAGValue outputs keyed by DAGNode name.
Automatically populated during ExecutionPlan execution to avoid redundant
computation. Used by downstream nodes to access their input data.
"""

project: str
repo_config: RepoConfig
offline_store: OfflineStore
Expand Down
25 changes: 23 additions & 2 deletions sdk/python/feast/infra/compute_engines/dag/node.py
Original file line number Diff line number Diff line change
@@ -1,9 +1,8 @@
from abc import ABC, abstractmethod
from typing import List

from infra.compute_engines.dag.value import DAGValue

from feast.infra.compute_engines.dag.context import ExecutionContext
from feast.infra.compute_engines.dag.value import DAGValue


class DAGNode(ABC):
Expand All @@ -22,5 +21,27 @@ def add_input(self, node: "DAGNode"):
self.inputs.append(node)
node.outputs.append(self)

def get_input_values(self, context: ExecutionContext) -> List[DAGValue]:
input_values = []
for input_node in self.inputs:
if input_node.name not in context.node_outputs:
raise KeyError(
f"Missing output for input node '{input_node.name}' in context."
)
input_values.append(context.node_outputs[input_node.name])
return input_values

def get_single_input_value(self, context: ExecutionContext) -> DAGValue:
if len(self.inputs) != 1:
raise RuntimeError(
f"DAGNode '{self.name}' expected exactly 1 input, but got {len(self.inputs)}."
)
input_node = self.inputs[0]
if input_node.name not in context.node_outputs:
raise KeyError(
f"Missing output for input node '{input_node.name}' in context."
)
return context.node_outputs[input_node.name]

@abstractmethod
def execute(self, context: ExecutionContext) -> DAGValue: ...
38 changes: 29 additions & 9 deletions sdk/python/feast/infra/compute_engines/spark/node.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,13 +2,13 @@
from datetime import datetime
from typing import Dict, List, Optional, Union, cast

from infra.compute_engines.dag.context import ExecutionContext
from pyspark.sql import DataFrame, SparkSession, Window
from pyspark.sql import functions as F

from feast import BatchFeatureView, StreamFeatureView
from feast.aggregation import Aggregation
from feast.infra.compute_engines.base import HistoricalRetrievalTask
from feast.infra.compute_engines.dag.context import ExecutionContext
from feast.infra.compute_engines.dag.model import DAGFormat
from feast.infra.compute_engines.dag.node import DAGNode
from feast.infra.compute_engines.dag.value import DAGValue
Expand Down Expand Up @@ -170,7 +170,7 @@ def __init__(
self.timestamp_col = timestamp_col

def execute(self, context: ExecutionContext) -> DAGValue:
input_value = context.node_outputs[self.inputs[0].name]
input_value = self.get_single_input_value(context)
input_value.assert_format(DAGFormat.SPARK)
input_df: DataFrame = input_value.data

Expand Down Expand Up @@ -213,32 +213,52 @@ def __init__(
feature_node: DAGNode,
join_keys: List[str],
feature_view: Union[BatchFeatureView, StreamFeatureView],
spark_session: SparkSession,
):
super().__init__(name)
self.join_keys = join_keys
self.add_input(feature_node)
self.feature_view = feature_view
self.spark_session = spark_session

def execute(self, context: ExecutionContext) -> DAGValue:
feature_value = context.node_outputs[self.inputs[1].name]
feature_value = self.get_single_input_value(context)
feature_value.assert_format(DAGFormat.SPARK)
feature_df = feature_value.data

entity_df = context.entity_df
feature_df = feature_value.data
assert entity_df is not None, "entity_df must be set in ExecutionContext"

# Get timestamp fields from feature view
join_keys, feature_cols, ts_col, created_ts_col = _get_column_names(
self.feature_view, context.entity_defs
)

entity_event_ts_col = "event_timestamp" # Standardized by SparkEntityLoadNode
# Rename entity_df event_timestamp_col to match feature_df
entity_schema = _get_entity_schema(
spark_session=self.spark_session,
entity_df=entity_df,
)
event_timestamp_col = infer_event_timestamp_from_entity_df(
entity_schema=entity_schema,
)
entity_ts_alias = "__entity_event_timestamp"
entity_df = entity_df.withColumnRenamed(event_timestamp_col, entity_ts_alias)

# Perform left join + event timestamp filtering
joined = feature_df.join(entity_df, on=join_keys, how="left")
joined = joined.filter(F.col(ts_col) <= F.col(entity_event_ts_col))
joined = joined.filter(F.col(ts_col) <= F.col(entity_ts_alias))

# Optional TTL filter: feature.ts >= entity.event_timestamp - ttl
if self.feature_view.ttl:
ttl_seconds = int(self.feature_view.ttl.total_seconds())
lower_bound = F.col(entity_ts_alias) - F.expr(
f"INTERVAL {ttl_seconds} seconds"
)
joined = joined.filter(F.col(ts_col) >= lower_bound)

# Dedup with row_number
partition_cols = join_keys + [entity_event_ts_col]
partition_cols = join_keys + [entity_ts_alias]
ordering = [F.col(ts_col).desc()]
if created_ts_col:
ordering.append(F.col(created_ts_col).desc())
Expand Down Expand Up @@ -267,7 +287,7 @@ def __init__(
self.feature_view = feature_view

def execute(self, context: ExecutionContext) -> DAGValue:
spark_df: DataFrame = context.node_outputs[self.inputs[0].name].data
spark_df: DataFrame = self.get_single_input_value(context).data

# ✅ 1. Write to offline store (if enabled)
if self.feature_view.online:
Expand Down Expand Up @@ -305,7 +325,7 @@ def __init__(self, name: str, input_node: DAGNode, udf):
self.udf = udf

def execute(self, context: ExecutionContext) -> DAGValue:
input_val = context.node_outputs[self.inputs[0].name]
input_val = self.get_single_input_value(context)
input_val.assert_format(DAGFormat.SPARK)

transformed_df = self.udf(input_val.data)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -49,7 +49,9 @@ def build_aggregation_node(self, input_node):

def build_join_node(self, input_node):
join_keys = self.feature_view.entities
node = SparkJoinNode("join", input_node, join_keys, self.feature_view)
node = SparkJoinNode(
"join", input_node, join_keys, self.feature_view, self.spark_session
)
self.nodes.append(node)
return node

Expand Down
Loading