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
Next Next commit
feat: Optional input_schema for ODFV (#6308)
Signed-off-by: Nick Quinn <nicholas_quinn@apple.com>
  • Loading branch information
nickquinn408 committed Apr 22, 2026
commit 56b0620e9b9c852f9dbe8d14bccde82dd0e202de
6 changes: 4 additions & 2 deletions sdk/python/feast/aggregation/__init__.py
Comment thread
nquinn408 marked this conversation as resolved.
Original file line number Diff line number Diff line change
Expand Up @@ -37,15 +37,17 @@ def __init__(
time_window: Optional[timedelta] = None,
slide_interval: Optional[timedelta] = None,
name: Optional[str] = None,
output: Optional[str] = None,
window: Optional[timedelta] = None,
):
self.column = column or ""
self.function = function or ""
self.time_window = time_window
self.time_window = window if window is not None else time_window
if not slide_interval:
self.slide_interval = self.time_window
else:
self.slide_interval = slide_interval
self.name = name or ""
self.name = output or name or ""

def to_proto(self) -> AggregationProto:
window_duration = None
Expand Down
75 changes: 64 additions & 11 deletions sdk/python/feast/on_demand_feature_view.py
Original file line number Diff line number Diff line change
Expand Up @@ -134,6 +134,7 @@ class OnDemandFeatureView(BaseFeatureView):
"""

_TRACK_METRICS_TAG = "feast:track_metrics"
_INPUT_SCHEMA_SOURCE_PREFIX = "__input_schema__"

name: str
entities: Optional[List[str]]
Expand All @@ -158,7 +159,8 @@ def __init__( # noqa: C901
name: str,
entities: Optional[List[Entity]] = None,
schema: Optional[List[Field]] = None,
sources: List[OnDemandSourceType],
sources: Optional[List[OnDemandSourceType]] = None,
input_schema: Optional[List[Field]] = None,
udf: Optional[FunctionType] = None,
udf_string: Optional[str] = "",
feature_transformation: Optional[Transformation] = None,
Expand All @@ -183,6 +185,9 @@ def __init__( # noqa: C901
sources: A map from input source names to the actual input sources, which may be
feature views, or request data sources. These sources serve as inputs to the udf,
which will refer to them by name.
input_schema (optional): A list of Fields describing the schema of the input data
for aggregation-based views. When provided, sources is not required — an
internal RequestSource will be created automatically.
udf: The user defined transformation function, which must take pandas
dataframes as inputs.
udf_string: The source code version of the udf (for diffing and displaying in Web UI)
Expand Down Expand Up @@ -214,15 +219,39 @@ def __init__( # noqa: C901
self.version = version
schema = schema or []
self.entities = [e.name for e in entities] if entities else [DUMMY_ENTITY_NAME]
self.sources = sources
self.input_schema = input_schema
self.mode = mode.lower()
self.udf = udf
self.udf_string = udf_string
self.source_feature_view_projections: dict[str, FeatureViewProjection] = {}
self.source_request_sources: dict[str, RequestSource] = {}

# Strip any existing sentinel from sources (handles __copy__ round-trip)
effective_sources: List[OnDemandSourceType] = [
s
for s in (sources or [])
if not (
isinstance(s, RequestSource)
and s.name.startswith(self._INPUT_SCHEMA_SOURCE_PREFIX)
)
]

if input_schema is not None:
# Automatically create an internal RequestSource from input_schema
sentinel = RequestSource(
name=f"{self._INPUT_SCHEMA_SOURCE_PREFIX}{name}",
schema=input_schema,
)
Comment thread
nquinn408 marked this conversation as resolved.
self.source_request_sources[sentinel.name] = sentinel
elif not effective_sources:
raise ValueError(
"Either 'sources' or 'input_schema' must be provided for OnDemandFeatureView."
)

self.sources = effective_sources

# Process each source with explicit type handling
for odfv_source in sources:
for odfv_source in effective_sources:
self._add_source_to_collections(odfv_source)

features: List[Field] = []
Expand Down Expand Up @@ -328,6 +357,7 @@ def __copy__(self):
schema=self.features,
sources=list(self.source_feature_view_projections.values())
+ list(self.source_request_sources.values()),
input_schema=self.input_schema,
feature_transformation=self.feature_transformation,
mode=self.mode,
description=self.description,
Expand All @@ -337,6 +367,7 @@ def __copy__(self):
singleton=self.singleton,
version=self.version,
track_metrics=self.track_metrics,
aggregations=self.aggregations,
)
fv.entities = self.entities
fv.features = self.features
Expand Down Expand Up @@ -559,7 +590,7 @@ def to_proto(self) -> OnDemandFeatureViewProto:
owner=self.owner,
write_to_online_store=self.write_to_online_store,
singleton=self.singleton or False,
aggregations=self.aggregations,
aggregations=[agg.to_proto() for agg in self.aggregations],
version=self.version,
)
return OnDemandFeatureViewProto(spec=spec, meta=meta)
Expand All @@ -585,6 +616,19 @@ def from_proto(
on_demand_feature_view_proto, skip_udf=skip_udf
)

# Detect and strip input_schema sentinel from sources
input_schema: Optional[List[Field]] = None
sources_without_sentinel: List[OnDemandSourceType] = []
for source in sources:
if (
isinstance(source, RequestSource)
and source.name.startswith(cls._INPUT_SCHEMA_SOURCE_PREFIX)
):
input_schema = source.schema
else:
sources_without_sentinel.append(source)
sources = sources_without_sentinel

# Parse transformation from proto (skip UDF deserialization if requested)
transformation = cls._parse_transformation_from_proto(
on_demand_feature_view_proto, skip_udf=skip_udf
Expand All @@ -607,6 +651,7 @@ def from_proto(
name=on_demand_feature_view_proto.spec.name,
schema=cls._parse_features_from_proto(on_demand_feature_view_proto),
sources=cast(List[OnDemandSourceType], sources),
input_schema=input_schema,
feature_transformation=transformation,
mode=on_demand_feature_view_proto.spec.mode or "pandas",
description=on_demand_feature_view_proto.spec.description,
Expand Down Expand Up @@ -1092,7 +1137,7 @@ def _is_array_type(self, dtype) -> bool:
"""Check if the dtype represents an array type."""
# Use proper type checking instead of string comparison
dtype_str = str(dtype)
return "Array" in dtype_str or "List" in dtype_str
return "Array" in dtype_str or "List" in dtype_str or "Set" in dtype_str

def _construct_random_input(
self, singleton: bool = False
Expand Down Expand Up @@ -1224,13 +1269,17 @@ def on_demand_feature_view(
name: Optional[str] = None,
entities: Optional[List[Entity]] = None,
schema: list[Field],
sources: list[
Union[
FeatureView,
RequestSource,
FeatureViewProjection,
sources: Optional[
list[
Union[
FeatureView,
RequestSource,
FeatureViewProjection,
]
]
],
] = None,
input_schema: Optional[list[Field]] = None,
aggregations: Optional[List[Aggregation]] = None,
mode: str = "pandas",
description: str = "",
tags: Optional[dict[str, str]] = None,
Expand All @@ -1252,6 +1301,8 @@ def on_demand_feature_view(
sources: A map from input source names to the actual input sources, which may be
feature views, or request data sources. These sources serve as inputs to the udf,
which will refer to them by name.
input_schema (optional): A list of Fields describing the schema of the input data
for aggregation-based views. When provided, sources is not required.
mode: The mode of execution (e.g,. Pandas or Python Native)
description (optional): A human-readable description.
tags (optional): A dictionary of key-value pairs to store arbitrary metadata.
Expand Down Expand Up @@ -1279,6 +1330,7 @@ def decorator(user_function):
on_demand_feature_view_obj = OnDemandFeatureView(
name=name if name is not None else user_function.__name__,
sources=sources,
input_schema=input_schema,
schema=schema,
mode=mode,
description=description,
Expand All @@ -1288,6 +1340,7 @@ def decorator(user_function):
entities=entities,
singleton=singleton,
track_metrics=track_metrics,
aggregations=aggregations,
udf=user_function,
udf_string=udf_string,
version=version,
Expand Down
165 changes: 165 additions & 0 deletions sdk/python/tests/unit/test_on_demand_feature_view_input_schema.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,165 @@
# Copyright 2025 The Feast Authors
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# https://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.

"""Tests for OnDemandFeatureView input_schema support."""

import copy
from datetime import timedelta

import pandas as pd
import pytest

from feast import Entity, Field
from feast.aggregation import Aggregation
from feast.on_demand_feature_view import OnDemandFeatureView, on_demand_feature_view
from feast.types import Float64, Int64
from feast.value_type import ValueType

user = Entity(name="user", join_keys=["user_id"], value_type=ValueType.INT64)


def test_decorator_with_input_schema():
"""The @on_demand_feature_view decorator supports input_schema without sources."""

@on_demand_feature_view(
input_schema=[
Field(name="txn_amount", dtype=Float64),
],
schema=[
Field(name="txn_count", dtype=Int64),
Field(name="total_txn_amount", dtype=Float64),
Field(name="avg_txn_amount", dtype=Float64),
],
aggregations=[
Aggregation(
column="txn_amount",
function="count",
output="txn_count",
window=timedelta(days=30),
),
Aggregation(
column="txn_amount",
function="sum",
output="total_txn_amount",
window=timedelta(days=30),
),
Aggregation(
column="txn_amount",
function="mean",
output="avg_txn_amount",
window=timedelta(days=30),
),
],
entities=[user],
)
def compute_txn_stats(df: pd.DataFrame) -> pd.DataFrame:
return df

assert isinstance(compute_txn_stats, OnDemandFeatureView)
assert compute_txn_stats.name == "compute_txn_stats"
assert compute_txn_stats.input_schema == [Field(name="txn_amount", dtype=Float64)]
assert len(compute_txn_stats.aggregations) == 3
assert len(compute_txn_stats.features) == 3

# The internal sentinel RequestSource should be present
sentinel_name = f"{OnDemandFeatureView._INPUT_SCHEMA_SOURCE_PREFIX}compute_txn_stats"
assert sentinel_name in compute_txn_stats.source_request_sources

# sources (user-visible) should be empty
assert compute_txn_stats.sources == []


def test_aggregation_aliases():
"""Aggregation accepts 'output' alias for 'name' and 'window' alias for 'time_window'."""
agg = Aggregation(
column="txn_amount",
function="sum",
output="total_txn_amount",
window=timedelta(days=30),
)
assert agg.name == "total_txn_amount"
assert agg.time_window == timedelta(days=30)


def test_input_schema_proto_roundtrip():
"""An ODFV with input_schema survives a to_proto / from_proto round-trip."""

@on_demand_feature_view(
input_schema=[
Field(name="txn_amount", dtype=Float64),
],
schema=[
Field(name="total_txn_amount", dtype=Float64),
],
aggregations=[
Aggregation(
column="txn_amount",
function="sum",
output="total_txn_amount",
window=timedelta(days=30),
),
],
entities=[user],
)
def txn_view(df: pd.DataFrame) -> pd.DataFrame:
return df

proto = txn_view.to_proto()
restored = OnDemandFeatureView.from_proto(proto)

assert restored.input_schema == txn_view.input_schema
assert restored.aggregations == txn_view.aggregations
sentinel_name = f"{OnDemandFeatureView._INPUT_SCHEMA_SOURCE_PREFIX}txn_view"
assert sentinel_name in restored.source_request_sources


def test_input_schema_copy():
"""__copy__ preserves input_schema and aggregations."""

@on_demand_feature_view(
input_schema=[
Field(name="txn_amount", dtype=Float64),
],
schema=[
Field(name="total_txn_amount", dtype=Float64),
],
aggregations=[
Aggregation(column="txn_amount", function="sum", output="total_txn_amount"),
],
entities=[user],
)
def copy_view(df: pd.DataFrame) -> pd.DataFrame:
return df

cloned = copy.copy(copy_view)
assert cloned.input_schema == copy_view.input_schema
assert cloned.aggregations == copy_view.aggregations
sentinel_name = f"{OnDemandFeatureView._INPUT_SCHEMA_SOURCE_PREFIX}copy_view"
assert sentinel_name in cloned.source_request_sources


def test_sources_required_without_input_schema():
"""Constructor raises if neither sources nor input_schema is provided."""
with pytest.raises(
(ValueError, TypeError),
Comment thread
nquinn408 marked this conversation as resolved.
Outdated
):

def dummy(df):
return df

OnDemandFeatureView(
name="bad_view",
schema=[Field(name="out", dtype=Float64)],
udf=dummy,
)