Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
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
5 changes: 1 addition & 4 deletions sdk/python/feast/data_source.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down
44 changes: 44 additions & 0 deletions sdk/python/tests/unit/test_data_sources.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
64 changes: 64 additions & 0 deletions sdk/python/tests/unit/test_unit_feature_store.py
Original file line number Diff line number Diff line change
@@ -1,13 +1,20 @@
from dataclasses import dataclass
from pathlib import Path
from typing import Dict, List
from unittest.mock import MagicMock, patch

import pandas as pd
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
Expand All @@ -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)],
Expand Down
Loading