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
2 changes: 2 additions & 0 deletions sdk/python/feast/data_source.py
Original file line number Diff line number Diff line change
Expand Up @@ -539,6 +539,7 @@ def __eq__(self, other):
or self.kafka_options.topic != other.kafka_options.topic
or self.kafka_options.watermark_delay_threshold
!= other.kafka_options.watermark_delay_threshold
or self.batch_source != other.batch_source
):
return False

Expand Down Expand Up @@ -831,6 +832,7 @@ def __eq__(self, other):
self.kinesis_options.record_format != other.kinesis_options.record_format
or self.kinesis_options.region != other.kinesis_options.region
or self.kinesis_options.stream_name != other.kinesis_options.stream_name
or self.batch_source != other.batch_source
):
return False

Expand Down
48 changes: 48 additions & 0 deletions sdk/python/tests/unit/test_data_sources.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,8 @@
from feast.infra.offline_stores.file_source import FileSource
from feast.infra.offline_stores.redshift_source import RedshiftSource
from feast.infra.offline_stores.snowflake_source import SnowflakeSource
from feast.infra.registry.registry import Registry
from feast.repo_config import RegistryConfig
from feast.types import Bool, Float32, Int64


Expand Down Expand Up @@ -379,3 +381,49 @@ def kafka_source(watermark):
KafkaSource.from_proto(unset_proto).kafka_options.watermark_delay_threshold
is None
)


def _stream_source_with_batch_path(source_class, path):
batch_source = FileSource(name="batch", path=path, timestamp_field="ts")
if source_class is KafkaSource:
return KafkaSource(
name="stream",
timestamp_field="ts",
kafka_bootstrap_servers="test_servers",
message_format=ProtoFormat("class_path"),
topic="test_topic",
batch_source=batch_source,
)
return KinesisSource(
name="stream",
timestamp_field="ts",
record_format=ProtoFormat("class_path"),
region="us-east-1",
stream_name="test_stream",
batch_source=batch_source,
)


@pytest.mark.parametrize("source_class", [KafkaSource, KinesisSource])
def test_stream_source_eq_compares_batch_source(source_class, tmp_path):
"""A changed batch source makes two stream sources unequal.

``apply_data_source`` skips the write when the stored source ``==`` the new
one, so ignoring ``batch_source`` in ``__eq__`` silently dropped the change.
"""
old = _stream_source_with_batch_path(source_class, "/tmp/old.parquet")
new = _stream_source_with_batch_path(source_class, "/tmp/new.parquet")

assert old != new
assert old == _stream_source_with_batch_path(source_class, "/tmp/old.parquet")

registry = Registry(
"project",
RegistryConfig(path=str(tmp_path / "registry.db"), cache_ttl_seconds=0),
None,
)
registry.apply_data_source(old, "project")
registry.apply_data_source(new, "project")

stored = registry.get_data_source("stream", "project", allow_cache=False)
assert stored.batch_source.path == "/tmp/new.parquet"
Loading