From 685766349e3763880ab56bf3585d878c565d2243 Mon Sep 17 00:00:00 2001 From: Arnav Sharma <73969466+Arnavsharma2@users.noreply.github.com> Date: Wed, 9 Sep 2026 16:11:33 -0500 Subject: [PATCH] fix: Preserve request source schema additions and removals Signed-off-by: Arnav Sharma <73969466+Arnavsharma2@users.noreply.github.com> --- sdk/python/feast/data_source.py | 5 +- sdk/python/tests/unit/test_data_sources.py | 44 +++++++++++++ .../tests/unit/test_unit_feature_store.py | 64 +++++++++++++++++++ 3 files changed, 109 insertions(+), 4 deletions(-) diff --git a/sdk/python/feast/data_source.py b/sdk/python/feast/data_source.py index 93473fb732e..96c9200347f 100644 --- a/sdk/python/feast/data_source.py +++ b/sdk/python/feast/data_source.py @@ -673,10 +673,7 @@ def __eq__(self, other): return False if isinstance(self.schema, List) and isinstance(other.schema, List): - for field1, field2 in zip(self.schema, other.schema): - if field1 != field2: - return False - return True + return self.schema == other.schema else: return False diff --git a/sdk/python/tests/unit/test_data_sources.py b/sdk/python/tests/unit/test_data_sources.py index e2a928f7bb9..b15d10afb81 100644 --- a/sdk/python/tests/unit/test_data_sources.py +++ b/sdk/python/tests/unit/test_data_sources.py @@ -63,6 +63,50 @@ def test_request_source_primitive_type_to_proto(): assert deserialized_request_source == request_source +@pytest.mark.parametrize( + "schema_names,other_schema_names,equal", + [ + pytest.param(["f1"], ["f1", "f2"], False, id="different-lengths"), + pytest.param([], ["f1"], False, id="empty-and-nonempty"), + pytest.param([], [], True, id="both-empty"), + pytest.param(["f1", "f2"], ["f1", "f2"], True, id="equal-fields"), + pytest.param(["f1"], ["other"], False, id="changed-field-name"), + pytest.param(["f1", "f2"], ["f2", "f1"], False, id="changed-order"), + ], +) +def test_request_source_schema_equality( + schema_names: list[str], other_schema_names: list[str], equal: bool +) -> None: + source = RequestSource( + name="source", schema=[Field(name=name, dtype=Int64) for name in schema_names] + ) + other = RequestSource( + name="source", + schema=[Field(name=name, dtype=Int64) for name in other_schema_names], + ) + + assert (source == other) is equal + assert (other == source) is equal + + +@pytest.mark.parametrize( + "other_field", + [ + pytest.param(Field(name="f1", dtype=Bool), id="changed-dtype"), + pytest.param( + Field(name="f1", dtype=Int64, tags={"team": "ml"}), + id="changed-field-metadata", + ), + ], +) +def test_request_source_schema_field_changes(other_field: Field) -> None: + source = RequestSource(name="source", schema=[Field(name="f1", dtype=Int64)]) + other = RequestSource(name="source", schema=[other_field]) + + assert source != other + assert other != source + + def test_hash(): push_source_1 = PushSource( name="test", diff --git a/sdk/python/tests/unit/test_unit_feature_store.py b/sdk/python/tests/unit/test_unit_feature_store.py index 0d3df24c5ac..2ff441a594f 100644 --- a/sdk/python/tests/unit/test_unit_feature_store.py +++ b/sdk/python/tests/unit/test_unit_feature_store.py @@ -1,4 +1,5 @@ from dataclasses import dataclass +from pathlib import Path from typing import Dict, List from unittest.mock import MagicMock, patch @@ -6,8 +7,14 @@ import pytest from feast import utils +from feast.data_source import RequestSource from feast.feature_store import FeatureStore +from feast.field import Field +from feast.infra.registry.registry import Registry +from feast.infra.registry.sql import SqlRegistry from feast.protos.feast.types.Value_pb2 import Value +from feast.repo_config import RepoConfig +from feast.types import Int64 @dataclass @@ -22,6 +29,63 @@ class MockFeatureView: projection: MockFeatureViewProjection +@pytest.mark.parametrize("registry_type", ["file", "sql"]) +@pytest.mark.parametrize( + "initial_fields,updated_fields", + [ + (["amount"], ["amount", "multiplier"]), + (["amount", "multiplier"], ["amount"]), + ([], ["amount"]), + (["amount"], []), + ], + ids=["add-field", "remove-field", "add-to-empty", "remove-all"], +) +def test_apply_request_source_schema_changes( + tmp_path: Path, + registry_type: str, + initial_fields: List[str], + updated_fields: List[str], +) -> None: + """Request schema changes must survive applying and reopening the registry.""" + registry_path = tmp_path / "registry.db" + config = RepoConfig( + project="test_project", + provider="local", + registry={ + "registry_type": registry_type, + "path": ( + str(registry_path) + if registry_type == "file" + else f"sqlite:///{registry_path}" + ), + }, + online_store={"type": "sqlite", "path": str(tmp_path / "online.db")}, + offline_store={"type": "file"}, + ) + store = FeatureStore(config=config) + reloaded_store = None + try: + store.apply( + RequestSource( + name="request", + schema=[Field(name=name, dtype=Int64) for name in initial_fields], + ) + ) + updated_schema = [Field(name=name, dtype=Int64) for name in updated_fields] + store.apply(RequestSource(name="request", schema=updated_schema)) + + reloaded_store = FeatureStore(config=config) + persisted_source = reloaded_store.get_data_source("request") + assert isinstance(persisted_source, RequestSource) + assert persisted_source.schema == updated_schema + finally: + for current_store in (reloaded_store, store): + if current_store is not None: + registry = current_store.registry + assert isinstance(registry, (Registry, SqlRegistry)) + registry.teardown() + + def test_get_unique_entities_success(): entity_values = { "entity_1": [Value(int64_val=1), Value(int64_val=2), Value(int64_val=1)],