diff --git a/docs/how-to-guides/online-server-performance-tuning.md b/docs/how-to-guides/online-server-performance-tuning.md index 50e746a6f03..4807ab000ad 100644 --- a/docs/how-to-guides/online-server-performance-tuning.md +++ b/docs/how-to-guides/online-server-performance-tuning.md @@ -32,7 +32,9 @@ When the server processes a `get_online_features()` call, it groups the requeste **Guideline:** For features that share the same entity key and are frequently requested together, consolidate them into a **single feature view**. This reduces the number of store round-trips per request. Split feature views only when features have different entities, different materialization schedules, or different data source update frequencies. {% hint style="info" %} -**Redis exception:** The Redis online store overrides `get_online_features()` to batch all `HMGET` commands across every feature view into a **single pipeline execution**. Because all feature views for the same entity share one Redis hash key, the number of Redis round trips is always **1**, regardless of how many feature views the request touches. This means the "fewer feature views" guideline is less critical for Redis than for other stores — but consolidating feature views still reduces serialization and protobuf overhead at the application layer. +**Redis exception:** The Redis online store overrides `_read_features_per_fv()` to batch all `HMGET` commands across every feature view into a **single pipeline execution**. Because all feature views for the same entity share one Redis hash key, the number of Redis round trips is always **1**, regardless of how many feature views the request touches. This means the "fewer feature views" guideline is less critical for Redis than for other stores — but consolidating feature views still reduces serialization and protobuf overhead at the application layer. + +**PostgreSQL exception:** PostgreSQL stores each feature view in its own table and overrides `_read_features_per_fv()` to combine the per-view reads into a **single `UNION ALL` statement**, so a request touching any number of feature views costs **1** query on both the sync and async paths. This helps most where round trips dominate — a handful of entities against a remote database. Requests larger than `max_batched_result_rows` (entity rows × feature views, default 2048) fall back to one query per view, because at that size the saved round trips no longer offset holding every view's rows at once. A single-view request is never batched, since there is nothing to combine. Set `max_batched_result_rows: 0` in the online store config to turn batching off. As with Redis, consolidating feature views still reduces application-layer overhead. {% endhint %} ### Feature services are free (and can be faster) @@ -87,7 +89,7 @@ Requesting just `combined_score` triggers reads from **both** `driver_stats_fv` ## Pre-computed feature vectors -When a `get_online_features()` request touches multiple feature views, the server issues a separate store read per feature view. For services spanning 5–15+ feature views, this fan-out dominates latency — even with Redis pipeline batching, the protobuf deserialization and response-building overhead grows linearly with the number of views. +When a `get_online_features()` request touches multiple feature views, the server issues a separate store read per feature view. For services spanning 5–15+ feature views, this fan-out dominates latency — even where the store batches its reads (Redis, PostgreSQL), the protobuf deserialization and response-building overhead grows linearly with the number of views. **Pre-computed feature vectors** solve this by storing all of a feature service's features for each entity as a single serialized blob. At read time, the server fetches one blob per entity instead of N reads per feature view, reducing the operation to O(1). @@ -276,7 +278,7 @@ The online store is the single largest factor in `get_online_features()` latency | ----- | ------------------- | ---------- | -------- | ------------- | | **Redis / Dragonfly** | < 1 ms | No (threadpool) | Ultra-low latency, high throughput; all FV reads batched into 1 pipeline | Requires in-memory capacity for your dataset | | **DynamoDB** | 2–5 ms | Yes | Serverless, auto-scaling on AWS | Pay-per-request cost; batch API limits (100 items) | -| **PostgreSQL** | 3–10 ms | No (threadpool) | Teams with existing Postgres infra | Connection pooling needed at scale | +| **PostgreSQL** | 3–10 ms | No (threadpool) | Teams with existing Postgres infra; all FV reads batched into 1 query | Connection pooling needed at scale | | **MongoDB** | 2–5 ms | Yes | Flexible schema, async-native | Requires index tuning for large datasets | | **Aerospike** | < 1 ms | No (threadpool) | Ultra-low latency, hybrid memory (RAM + SSD), large datasets | Namespace must be pre-configured on the cluster | | **Bigtable** | 3–8 ms | No (threadpool) | Large-scale GCP workloads | Row-key design affects read performance | @@ -304,8 +306,8 @@ The feature server can read from the online store using either an **async** or * | ----- | ---------- | ----------- | ----- | | **DynamoDB** | Yes | Yes | Uses `aiobotocore` for non-blocking I/O | | **MongoDB** | Yes | Yes | Uses `motor` (async MongoDB driver) | -| **PostgreSQL** | Implemented | No | Has `online_read_async` but does not yet advertise via `async_supported`; uses sync/threadpool path | -| **Redis** | Implemented | **Yes** | `online_read_async` and `online_write_batch_async` both implemented; uses sync/threadpool path for `get_online_features` (overridden with batched single pipeline) | +| **PostgreSQL** | Implemented | No | Has `online_read_async` but does not yet advertise via `async_supported`; uses sync/threadpool path. `_read_features_per_fv` is overridden to batch all feature view reads into a single `UNION ALL` query | +| **Redis** | Implemented | **Yes** | `online_read_async` and `online_write_batch_async` both implemented; uses sync/threadpool path for `get_online_features` (`_read_features_per_fv` overridden with batched single pipeline) | | **Aerospike** | Implemented | No | Async methods wrap the blocking C client via `run_in_executor`; does not yet advertise via `async_supported`, so the server still uses the threadpool path | | All others | No | No | Fall back to sync with `run_in_threadpool()` | @@ -387,7 +389,7 @@ online_store: #### Batched multi-feature-view reads -The Redis online store overrides `get_online_features()` to issue all `HMGET` commands — across every feature view in the request — in a **single pipeline execution**. This reduces Redis round trips from `N` (one per feature view) to `1` regardless of request size. +The Redis online store overrides `_read_features_per_fv()` to issue all `HMGET` commands — across every feature view in the request — in a **single pipeline execution**. This reduces Redis round trips from `N` (one per feature view) to `1` regardless of request size. | Feature views | Round trips (other stores) | Round trips (Redis) | | :---: | :---: | :---: | diff --git a/docs/reference/online-stores/postgres.md b/docs/reference/online-stores/postgres.md index 10f3f871710..27e7e2ca71f 100644 --- a/docs/reference/online-stores/postgres.md +++ b/docs/reference/online-stores/postgres.md @@ -8,6 +8,8 @@ The PostgreSQL online store provides support for materializing feature values in * `sslmode` defaults to `require`, which encrypts the connection without certificate verification. To disable SSL (e.g. for local development), set `sslmode: disable`. For certificate verification, set `sslmode` to `verify-ca` or `verify-full` and provide the corresponding `sslrootcert_path` (and optionally `sslcert_path` and `sslkey_path` for mutual TLS) +* Reads that span several feature views are combined into one query. `max_batched_result_rows` (default 2048) caps the request size, in entity rows × feature views, that is batched this way; set it to `0` to always issue one query per feature view + ## Getting started In order to use this online store, you'll need to run `pip install 'feast[postgres]'`. You can get started by then running `feast init -t postgres`. diff --git a/sdk/python/feast/infra/online_stores/postgres_online_store/postgres.py b/sdk/python/feast/infra/online_stores/postgres_online_store/postgres.py index 09be88c991d..fc9cb0813d2 100644 --- a/sdk/python/feast/infra/online_stores/postgres_online_store/postgres.py +++ b/sdk/python/feast/infra/online_stores/postgres_online_store/postgres.py @@ -10,6 +10,7 @@ Generator, List, Literal, + NamedTuple, Optional, Sequence, Tuple, @@ -19,8 +20,9 @@ from psycopg import AsyncConnection, sql from psycopg.connection import Connection from psycopg_pool import AsyncConnectionPool, ConnectionPool +from pydantic import NonNegativeInt -from feast import Entity, FeatureView, ValueType +from feast import Entity, FeatureView, ValueType, utils from feast.filter_models import ( ComparisonFilter, CompoundFilter, @@ -55,6 +57,18 @@ "inner_product": "<#>", } +# Default for ``max_batched_result_rows``: the largest request, measured as entity +# rows x feature views, that is read in one batched query. Past that, batching stops +# paying for itself: the saved round trips are amortized away by the payload, while +# holding every feature view's rows at once makes peak memory grow with the number of +# views. The measure is a proxy taken from the request shape; the rows actually +# returned also scale with the number of requested features. Measured against +# Postgres 16, a 10-view request is ~3-5x faster at 1-100 entities and a wash at +# 5000, where the combined result costs tens of MB per in-flight request. Past the +# limit the generic one-query-per-view path is used instead, which processes and +# releases a single view at a time. +MAX_BATCHED_RESULT_ROWS = 2048 + _PG_COMPARISON_OPS: Dict[str, str] = { "eq": "=", "ne": "!=", @@ -65,6 +79,16 @@ } +class _BatchedRead(NamedTuple): + """One feature view's share of a batched online read.""" + + table: FeatureView + requested_features: List[str] + keys: List[bytes] + idxs: Tuple[List[int], ...] + output_len: int + + class PostgresFilterTranslator(FilterTranslator): """Translates Feast filters into Postgres SQL WHERE clause fragments.""" @@ -159,6 +183,9 @@ def _pg_filter_col_and_val(value: Any) -> Tuple[str, Any]: class PostgreSQLOnlineStoreConfig(PostgreSQLConfig, VectorStoreConfig): type: Literal["postgres"] = "postgres" enable_openai_compatible_store: Optional[bool] = False + max_batched_result_rows: NonNegativeInt = MAX_BATCHED_RESULT_ROWS + """Largest request, as entity rows x feature views, read in one batched query. + Set to 0 to always read one feature view per query.""" class PostgreSQLOnlineStore(OnlineStore): @@ -424,6 +451,13 @@ def _process_rows( row[0] if isinstance(row[0], bytes) else row[0].tobytes() ].append(row[1:]) + return PostgreSQLOnlineStore._result_from_values_dict(keys, values_dict) + + @staticmethod + def _result_from_values_dict( + keys: List[bytes], values_dict: Dict[bytes, List[Tuple]] + ) -> List[Tuple[Optional[datetime], Optional[Dict[str, ValueProto]]]]: + """Assemble per-entity rows in ``keys`` order from a feature-value mapping.""" result: List[Tuple[Optional[datetime], Optional[Dict[str, ValueProto]]]] = [] for key in keys: if key in values_dict: @@ -438,6 +472,225 @@ def _process_rows( result.append((None, None)) return result + def _read_features_per_fv( + self, + config: RepoConfig, + grouped_refs: List, + join_key_values: Dict, + entity_name_to_join_key_map: Dict, + online_features_response, + full_feature_names: bool, + include_feature_view_version_metadata: bool, + ) -> None: + """Read every requested feature view in a single round trip. + + The generic path issues one query per feature view. Each feature view lives + in its own table here, so the per-view reads can be combined with UNION ALL + and split apart again afterwards. + """ + if not self._should_batch(config, grouped_refs, join_key_values): + return super()._read_features_per_fv( + config, + grouped_refs, + join_key_values, + entity_name_to_join_key_map, + online_features_response, + full_feature_names, + include_feature_view_version_metadata, + ) + + reads = self._prepare_batched_reads( + config, grouped_refs, join_key_values, entity_name_to_join_key_map + ) + + query, params = self._construct_batched_query_and_params(config, reads) + with self._get_conn(config, autocommit=True) as conn, conn.cursor() as cur: + cur.execute(query, params) + buckets = self._bucket_rows(reads, cur) + + self._populate_from_buckets( + reads, + buckets, + online_features_response, + full_feature_names, + include_feature_view_version_metadata, + ) + + async def _read_features_per_fv_async( + self, + config: RepoConfig, + grouped_refs: List, + join_key_values: Dict, + entity_name_to_join_key_map: Dict, + online_features_response, + full_feature_names: bool, + include_feature_view_version_metadata: bool, + ) -> None: + """Async version of :meth:`_read_features_per_fv`. + + The generic async path issues the per-view queries concurrently. Combining + them still replaces those N queries with one. + """ + if not self._should_batch(config, grouped_refs, join_key_values): + return await super()._read_features_per_fv_async( + config, + grouped_refs, + join_key_values, + entity_name_to_join_key_map, + online_features_response, + full_feature_names, + include_feature_view_version_metadata, + ) + + reads = self._prepare_batched_reads( + config, grouped_refs, join_key_values, entity_name_to_join_key_map + ) + + query, params = self._construct_batched_query_and_params(config, reads) + async with self._get_conn_async(config, autocommit=True) as conn: + async with conn.cursor() as cur: + await cur.execute(query, params) + buckets = [defaultdict(list) for _ in reads] # type: List[Dict[bytes, List[Tuple]]] + async for row in cur: + self._bucket_row(buckets, row) + + self._populate_from_buckets( + reads, + buckets, + online_features_response, + full_feature_names, + include_feature_view_version_metadata, + ) + + def _prepare_batched_reads( + self, + config: RepoConfig, + grouped_refs: List, + join_key_values: Dict, + entity_name_to_join_key_map: Dict, + ) -> List["_BatchedRead"]: + """Resolve the entity keys each feature view needs before querying. + + Feature views can be keyed on different entities, so the serialized keys are + computed per view rather than shared. + """ + reads = [] + for table, requested_features in grouped_refs: + table_entity_values, idxs, output_len = utils._get_unique_entities( + table, join_key_values, entity_name_to_join_key_map + ) + entity_key_protos = utils._get_entity_key_protos(table_entity_values) + reads.append( + _BatchedRead( + table=table, + requested_features=requested_features, + keys=self._prepare_keys( + entity_key_protos, config.entity_key_serialization_version + ), + idxs=idxs, + output_len=output_len, + ) + ) + return reads + + @staticmethod + def _construct_batched_query_and_params( + config: RepoConfig, reads: List["_BatchedRead"] + ) -> Tuple[sql.Composed, List[Any]]: + """UNION ALL the per feature view reads into one statement. + + Each branch selects a constant tag so the combined result can be split back + apart. The table name itself is not in the result set, and two feature views + can return the same entity key, so the tag is what keeps them separable. + """ + versioning = config.registry.enable_online_feature_view_versioning + branches: List[sql.Composed] = [] + params: List[Any] = [] + for tag, read in enumerate(reads): + table_name = _table_id(config.project, read.table, versioning) + if read.requested_features: + branch = sql.SQL( + "SELECT {tag} AS fv_tag, entity_key, feature_name, value, event_ts " + "FROM {table} WHERE entity_key = ANY(%s) AND feature_name = ANY(%s)" + ).format(tag=sql.Literal(tag), table=sql.Identifier(table_name)) + params.extend([read.keys, list(read.requested_features)]) + else: + branch = sql.SQL( + "SELECT {tag} AS fv_tag, entity_key, feature_name, value, event_ts " + "FROM {table} WHERE entity_key = ANY(%s)" + ).format(tag=sql.Literal(tag), table=sql.Identifier(table_name)) + params.append(read.keys) + branches.append(branch) + + return sql.SQL(" UNION ALL ").join(branches), params + + @staticmethod + def _should_batch( + config: RepoConfig, grouped_refs: List, join_key_values: Dict + ) -> bool: + """Whether to read this request's feature views in one batched query. + + A single view has no round trips to save, so it keeps the generic path. The + size check is estimated from the request shape rather than from the resolved + keys, so that deciding against batching costs nothing. The number of entity + rows in the request is an upper bound on the unique keys any one view will + read. See :data:`MAX_BATCHED_RESULT_ROWS`. + """ + if len(grouped_refs) < 2: + return False + limit = getattr( + config.online_store, "max_batched_result_rows", MAX_BATCHED_RESULT_ROWS + ) + entity_rows = max( + (len(values) for values in join_key_values.values()), default=0 + ) + return entity_rows * len(grouped_refs) <= limit + + @staticmethod + def _bucket_row(buckets: List[Dict[bytes, List[Tuple]]], row: Tuple) -> None: + """File one row under its feature view's bucket, keyed by entity key. + + Only the three payload columns are kept. Slicing the row instead would copy + it a second time while the combined result is still referenced. + """ + entity_key = row[1] if isinstance(row[1], bytes) else row[1].tobytes() + buckets[row[0]][entity_key].append((row[2], row[3], row[4])) + + def _bucket_rows( + self, reads: List["_BatchedRead"], row_iter + ) -> List[Dict[bytes, List[Tuple]]]: + """Consume the combined result one row at a time. + + The result spans every requested feature view, so calling ``fetchall()`` and + then slicing each row would hold two or three copies of it at peak. Bucketing + as rows arrive keeps a single copy. + """ + buckets: List[Dict[bytes, List[Tuple]]] = [defaultdict(list) for _ in reads] + for row in row_iter: + self._bucket_row(buckets, row) + return buckets + + def _populate_from_buckets( + self, + reads: List["_BatchedRead"], + buckets: List[Dict[bytes, List[Tuple]]], + online_features_response, + full_feature_names: bool, + include_feature_view_version_metadata: bool, + ) -> None: + """Hand each feature view's bucket to the response in grouped_refs order.""" + for read, bucket in zip(reads, buckets): + utils._populate_response_from_feature_data( + read.requested_features, + self._result_from_values_dict(read.keys, bucket), + read.idxs, + online_features_response, + full_feature_names, + read.table, + read.output_len, + include_feature_view_version_metadata, + ) + def update( self, config: RepoConfig, diff --git a/sdk/python/tests/integration/online_store/test_postgres_batched_read.py b/sdk/python/tests/integration/online_store/test_postgres_batched_read.py new file mode 100644 index 00000000000..ff56b91f0e5 --- /dev/null +++ b/sdk/python/tests/integration/online_store/test_postgres_batched_read.py @@ -0,0 +1,233 @@ +"""Integration tests for batched PostgreSQL online reads. + +The batched path must return exactly what the generic per-feature-view path +returns. These run against a real database because the UNION ALL and its tag +column are the part worth checking, and neither is exercised by mocks. + +Run with: pytest --integration sdk/python/tests/integration/online_store/test_postgres_batched_read.py +""" + +import shutil +from datetime import datetime, timedelta, timezone +from unittest.mock import patch + +import psycopg +import pytest + +from feast import Entity, FeatureView +from feast.field import Field +from feast.infra.online_stores.online_store import OnlineStore +from feast.infra.online_stores.postgres_online_store.postgres import ( + PostgreSQLOnlineStore, +) +from feast.protos.feast.serving.ServingService_pb2 import ( + FieldStatus, + GetOnlineFeaturesResponse, +) +from feast.protos.feast.types.EntityKey_pb2 import EntityKey as EntityKeyProto +from feast.protos.feast.types.Value_pb2 import Value as ValueProto +from feast.repo_config import RegistryConfig, RepoConfig +from feast.types import Int64 +from feast.value_type import ValueType + +DRIVER = Entity(name="driver", join_keys=["driver_id"], value_type=ValueType.INT64) +JOIN_KEY_MAP = {"driver": "driver_id"} +DRIVER_IDS = [1001, 1002, 1003] + + +def _feature_view(name: str, feature: str) -> FeatureView: + return FeatureView( + name=name, + entities=[DRIVER], + ttl=timedelta(days=1), + schema=[Field(name="driver_id", dtype=Int64), Field(name=feature, dtype=Int64)], + ) + + +def _entity_key(join_key: str, value: int) -> EntityKeyProto: + key = EntityKeyProto() + key.join_keys.append(join_key) + key.entity_values.append(ValueProto(int64_val=value)) + return key + + +@pytest.mark.integration +@pytest.mark.skipif(not shutil.which("docker"), reason="Docker not available") +class TestPostgresBatchedRead: + @pytest.fixture(autouse=True) + def setup_postgres(self, tmp_path): + try: + from testcontainers.postgres import PostgresContainer + except ImportError: + pytest.skip("testcontainers[postgres] not installed") + + self.registry_path = str(tmp_path / "registry.pb") + self.container = PostgresContainer( + "postgres:16", + username="root", + password="testpass", # pragma: allowlist secret + dbname="test", + ).with_exposed_ports(5432) + self.container.start() + self.port = self.container.get_exposed_port(5432) + yield + self.container.stop() + + def _config(self) -> RepoConfig: + from feast.infra.online_stores.postgres_online_store.postgres import ( + PostgreSQLOnlineStoreConfig, + ) + + return RepoConfig( + project="batched", + provider="local", + online_store=PostgreSQLOnlineStoreConfig( + type="postgres", + host="localhost", + port=int(self.port), + user="root", + password="testpass", # pragma: allowlist secret + database="test", + sslmode="disable", + ), + registry=RegistryConfig(path=self.registry_path), + entity_key_serialization_version=3, + ) + + def _write(self, store, config, fv, feature, values, ids=None): + now = datetime.now(tz=timezone.utc) + ids = ids if ids is not None else DRIVER_IDS + store.online_write_batch( + config, + fv, + [ + ( + _entity_key("driver_id", i), + {feature: ValueProto(int64_val=v)}, + now, + now, + ) + for i, v in zip(ids, values) + ], + None, + ) + + def _args(self, config, grouped_refs, join_key_values, response): + return dict( + config=config, + grouped_refs=grouped_refs, + join_key_values=join_key_values, + entity_name_to_join_key_map=JOIN_KEY_MAP, + online_features_response=response, + full_feature_names=False, + include_feature_view_version_metadata=False, + ) + + def _read(self, store, config, grouped_refs, join_key_values, generic=False): + response = GetOnlineFeaturesResponse() + args = self._args(config, grouped_refs, join_key_values, response) + if generic: + OnlineStore._read_features_per_fv(store, **args) + else: + store._read_features_per_fv(**args) + return response + + async def _read_async( + self, store, config, grouped_refs, join_key_values, generic=False + ): + response = GetOnlineFeaturesResponse() + args = self._args(config, grouped_refs, join_key_values, response) + if generic: + await OnlineStore._read_features_per_fv_async(store, **args) + else: + await store._read_features_per_fv_async(**args) + return response + + def _drivers(self, ids=None): + return { + "driver_id": [ValueProto(int64_val=i) for i in (ids or DRIVER_IDS)], + } + + def _setup_three_views(self, store, config, feature_name=None): + """Three views; by default they share one feature name so a demux slip shows.""" + specs = [("fv_a", 10), ("fv_b", 20), ("fv_c", 30)] + views = [] + for name, base in specs: + feature = feature_name or f"f_{name}" + fv = _feature_view(name, feature) + views.append((fv, feature, base)) + store.update(config, [], [v[0] for v in views], [], [], False) + for fv, feature, base in views: + self._write(store, config, fv, feature, [base, base + 1, base + 2]) + return [(fv, [feature]) for fv, feature, _ in views] + + def test_batched_matches_generic_path_in_one_execute(self): + """Same results as the generic path, in one execute instead of three.""" + config = self._config() + store = PostgreSQLOnlineStore() + grouped_refs = self._setup_three_views(store, config, feature_name="feat") + + real_execute = psycopg.Cursor.execute + + def counted(executes): + def execute(self, query, params=None, **kwargs): + executes.append(query) + return real_execute(self, query, params, **kwargs) + + return execute + + batched_executes: list = [] + with patch.object(psycopg.Cursor, "execute", counted(batched_executes)): + batched = self._read(store, config, grouped_refs, self._drivers()) + + generic_executes: list = [] + with patch.object(psycopg.Cursor, "execute", counted(generic_executes)): + generic = self._read( + store, config, grouped_refs, self._drivers(), generic=True + ) + + assert batched == generic + assert [v.int64_val for v in batched.results[0].values] == [10, 11, 12] + assert [v.int64_val for v in batched.results[1].values] == [20, 21, 22] + assert [v.int64_val for v in batched.results[2].values] == [30, 31, 32] + assert len(batched_executes) == 1 + assert len(generic_executes) == 3 + + def test_missing_entities_report_not_found_like_generic(self): + """An entity absent from one view must not shift another view's values.""" + config = self._config() + store = PostgreSQLOnlineStore() + + fv_full = _feature_view("fv_full", "feat") + fv_sparse = _feature_view("fv_sparse", "feat") + store.update(config, [], [fv_full, fv_sparse], [], [], False) + self._write(store, config, fv_full, "feat", [100, 101, 102]) + # Only the middle driver exists in the sparse view. + self._write(store, config, fv_sparse, "feat", [7], ids=[DRIVER_IDS[1]]) + + grouped_refs = [(fv_full, ["feat"]), (fv_sparse, ["feat"])] + batched = self._read(store, config, grouped_refs, self._drivers()) + generic = self._read(store, config, grouped_refs, self._drivers(), generic=True) + + assert batched == generic + # Status, not value: a stored zero and a missing row both read as 0. + assert list(batched.results[1].statuses) == [ + FieldStatus.NOT_FOUND, + FieldStatus.PRESENT, + FieldStatus.NOT_FOUND, + ] + assert batched.results[1].values[1].int64_val == 7 + + async def test_async_batched_matches_generic(self): + config = self._config() + store = PostgreSQLOnlineStore() + grouped_refs = self._setup_three_views(store, config, feature_name="feat") + + batched = await self._read_async(store, config, grouped_refs, self._drivers()) + generic = await self._read_async( + store, config, grouped_refs, self._drivers(), generic=True + ) + + assert batched == generic + assert [v.int64_val for v in batched.results[0].values] == [10, 11, 12] + assert [v.int64_val for v in batched.results[1].values] == [20, 21, 22] diff --git a/sdk/python/tests/unit/infra/online_store/test_postgres_batched_read.py b/sdk/python/tests/unit/infra/online_store/test_postgres_batched_read.py new file mode 100644 index 00000000000..362206d0518 --- /dev/null +++ b/sdk/python/tests/unit/infra/online_store/test_postgres_batched_read.py @@ -0,0 +1,330 @@ +"""Unit tests for batching PostgreSQL online reads across feature views.""" + +import contextlib +from datetime import datetime, timedelta +from typing import List, Tuple +from unittest.mock import patch + +import pytest +from pydantic import ValidationError + +from feast import Entity, FeatureView +from feast.field import Field +from feast.infra.online_stores.postgres_online_store.postgres import ( + MAX_BATCHED_RESULT_ROWS, + PostgreSQLOnlineStore, +) +from feast.protos.feast.serving.ServingService_pb2 import GetOnlineFeaturesResponse +from feast.protos.feast.types.Value_pb2 import Value as ValueProto +from feast.repo_config import RegistryConfig, RepoConfig +from feast.types import Int64 +from feast.value_type import ValueType + +DRIVER = Entity(name="driver", join_keys=["driver_id"], value_type=ValueType.INT64) +CUSTOMER = Entity( + name="customer", join_keys=["customer_id"], value_type=ValueType.INT64 +) +JOIN_KEY_MAP = {"driver": "driver_id", "customer": "customer_id"} +TS = datetime(2026, 1, 1) + + +def _feature_view(name: str, feature: str, entity: Entity = DRIVER) -> FeatureView: + return FeatureView( + name=name, + entities=[entity], + ttl=timedelta(days=1), + schema=[Field(name=feature, dtype=Int64)], + ) + + +def _config(**online_store) -> RepoConfig: + return RepoConfig( + project="proj", + provider="local", + registry=RegistryConfig(path="registry.db"), + online_store={ + "type": "postgres", + "host": "localhost", + "port": 5432, + "database": "test", + "db_schema": "public", + "user": "root", + "password": "test", # pragma: allowlist secret + **online_store, + }, + entity_key_serialization_version=3, + ) + + +def _int(value: int) -> bytes: + return ValueProto(int64_val=value).SerializeToString() + + +class _FakeCursor: + """Records executed statements and replays canned rows by iteration.""" + + def __init__(self, rows: List[Tuple], executed: List): + self._rows = rows + self._executed = executed + + def execute(self, query, params=None): + self._executed.append((query, params)) + + def __iter__(self): + return iter(self._rows) + + def __enter__(self): + return self + + def __exit__(self, *args): + return False + + +class _FakeConnection: + def __init__(self, rows: List[Tuple], executed: List): + self._rows = rows + self._executed = executed + + def cursor(self): + return _FakeCursor(self._rows, self._executed) + + +@contextlib.contextmanager +def _fake_conn(rows: List[Tuple], executed: List): + yield _FakeConnection(rows, executed) + + +class _AsyncFakeCursor(_FakeCursor): + async def execute(self, query, params=None): # type: ignore[override] + self._executed.append((query, params)) + + def __aiter__(self): + async def gen(): + for row in self._rows: + yield row + + return gen() + + async def __aenter__(self): + return self + + async def __aexit__(self, *args): + return False + + +class _AsyncFakeConnection(_FakeConnection): + def cursor(self): # type: ignore[override] + return _AsyncFakeCursor(self._rows, self._executed) + + +@contextlib.asynccontextmanager +async def _fake_conn_async(rows: List[Tuple], executed: List): + yield _AsyncFakeConnection(rows, executed) + + +def _args(grouped_refs, join_key_values, response, config=None): + return dict( + config=config or _config(), + grouped_refs=grouped_refs, + join_key_values=join_key_values, + entity_name_to_join_key_map=JOIN_KEY_MAP, + online_features_response=response, + full_feature_names=False, + include_feature_view_version_metadata=False, + ) + + +def _drivers(*ids) -> dict: + return {"driver_id": [ValueProto(int64_val=i) for i in ids]} + + +def _read(store, grouped_refs, rows, join_key_values=None, config=None): + """Drive the sync override against a fake connection.""" + executed: List = [] + response = GetOnlineFeaturesResponse() + with patch.object( + PostgreSQLOnlineStore, + "_get_conn", + lambda self, cfg, autocommit=False: _fake_conn(rows, executed), + ): + store._read_features_per_fv( + **_args(grouped_refs, join_key_values or _drivers(1, 2), response, config) + ) + return executed, response + + +async def _read_async(store, grouped_refs, rows, join_key_values=None, config=None): + """Drive the async override against a fake connection.""" + executed: List = [] + response = GetOnlineFeaturesResponse() + with patch.object( + PostgreSQLOnlineStore, + "_get_conn_async", + lambda self, cfg, autocommit=False: _fake_conn_async(rows, executed), + ): + await store._read_features_per_fv_async( + **_args(grouped_refs, join_key_values or _drivers(1, 2), response, config) + ) + return executed, response + + +async def _read_sync(store, grouped_refs, rows, join_key_values=None, config=None): + """Await-compatible wrapper so both overrides share one test body.""" + return _read(store, grouped_refs, rows, join_key_values, config) + + +READERS = [pytest.param(_read_sync, id="sync"), pytest.param(_read_async, id="async")] + + +@pytest.mark.parametrize("read", READERS) +async def test_views_read_in_one_query_and_demultiplexed_by_tag(read): + """Two views sharing an entity key AND a feature name come back in one query. + + Distinct feature names make a demux failure invisible: _process_rows keys values + by feature name, so a cross-contaminated bucket still yields the right value. + """ + store = PostgreSQLOnlineStore() + grouped_refs = [ + (_feature_view("fv_a", "feat"), ["feat"]), + (_feature_view("fv_b", "feat"), ["feat"]), + ] + reads = store._prepare_batched_reads( + _config(), grouped_refs, _drivers(1, 2), JOIN_KEY_MAP + ) + shared_key = reads[0].keys[0] + assert shared_key == reads[1].keys[0] + rows = [ + (0, shared_key, "feat", _int(11), TS), + (1, shared_key, "feat", _int(22), TS), + ] + + executed, response = await read(store, grouped_refs, rows) + + assert len(executed) == 1 + statement = executed[0][0].as_string(None) + assert statement.count("UNION ALL") == 1 + for table in ("proj_fv_a", "proj_fv_b"): + assert table in statement + assert response.results[0].values[0].int64_val == 11 + assert response.results[1].values[0].int64_val == 22 + + +def test_per_view_entity_bookkeeping_is_not_shared(): + """Views on different entities must each use their own idxs and output_len. + + Every view here resolves a different number of unique entities, so reusing one + view's index mapping for another shows up as a wrong or misplaced value. + """ + store = PostgreSQLOnlineStore() + fv_driver = _feature_view("fv_driver", "d_feat", DRIVER) + fv_customer = _feature_view("fv_customer", "c_feat", CUSTOMER) + grouped_refs = [(fv_driver, ["d_feat"]), (fv_customer, ["c_feat"])] + # Three request rows; the driver view sees 3 distinct keys, the customer view 2. + join_key_values = { + "driver_id": [ValueProto(int64_val=i) for i in (1, 2, 3)], + "customer_id": [ValueProto(int64_val=i) for i in (7, 7, 8)], + } + + reads = store._prepare_batched_reads( + _config(), grouped_refs, join_key_values, JOIN_KEY_MAP + ) + assert len(reads[0].keys) == 3 + assert len(reads[1].keys) == 2 + + rows = [(0, reads[0].keys[i], "d_feat", _int(10 + i), TS) for i in range(3)] + [ + (1, reads[1].keys[i], "c_feat", _int(70 + i), TS) for i in range(2) + ] + + _, response = _read(store, grouped_refs, rows, join_key_values) + + assert [v.int64_val for v in response.results[0].values] == [10, 11, 12] + # customer 7 is requested twice and must fan back out to both output positions + assert [v.int64_val for v in response.results[1].values] == [70, 70, 71] + + +@pytest.mark.parametrize("read", READERS) +@pytest.mark.parametrize( + ("views", "entity_rows", "online_store", "exp_batched"), + [ + (2, 2, {}, True), + (2, MAX_BATCHED_RESULT_ROWS, {}, False), + (1, 2, {}, False), + (2, 2, {"max_batched_result_rows": 4}, True), + (2, 2, {"max_batched_result_rows": 3}, False), + (2, 2, {"max_batched_result_rows": 0}, False), + ], + ids=[ + "under_default_limit", + "over_default_limit", + "single_view", + "at_configured_limit", + "over_configured_limit", + "zero_disables_batching", + ], +) +async def test_when_batching_applies( + read, views, entity_rows, online_store, exp_batched +): + """Batching needs at least two views and a request within the row limit. + + A batched request issues one statement and never calls online_read; the generic + path issues none itself and reads each view through online_read instead. + """ + store = PostgreSQLOnlineStore() + names = [f"fv_{i}" for i in range(views)] + grouped_refs = [(_feature_view(n, f"feat_{n}"), [f"feat_{n}"]) for n in names] + per_view_reads: List[str] = [] + + def one_row_per_key(self, config, table, entity_keys, requested_features=None): + per_view_reads.append(table.name) + return [(None, None)] * len(entity_keys) + + async def one_row_per_key_async(self, *args, **kwargs): + return one_row_per_key(self, *args, **kwargs) + + with ( + patch.object(PostgreSQLOnlineStore, "online_read", one_row_per_key), + patch.object(PostgreSQLOnlineStore, "online_read_async", one_row_per_key_async), + ): + executed, _ = await read( + store, + grouped_refs, + [], + _drivers(*range(entity_rows)), + _config(**online_store), + ) + + assert len(executed) == (1 if exp_batched else 0) + assert sorted(per_view_reads) == ([] if exp_batched else names) + + +def test_max_batched_result_rows_config(): + """Defaults to 2048, and rejects negatives rather than silently never batching.""" + assert _config().online_store.max_batched_result_rows == 2048 + with pytest.raises(ValidationError): + _config(max_batched_result_rows=-1) + + +@pytest.mark.parametrize( + ("requested", "exp_feature_filter", "exp_params"), + [(["feat_a"], True, 2), ([], False, 1)], + ids=["with_requested_features", "without_requested_features"], +) +def test_query_shape_with_and_without_requested_features( + requested, exp_feature_filter, exp_params +): + """Omitting requested features must drop the filter, not pass an empty one.""" + store = PostgreSQLOnlineStore() + config = _config() + reads = store._prepare_batched_reads( + config, + [(_feature_view("fv_a", "feat_a"), requested)], + _drivers(1), + JOIN_KEY_MAP, + ) + + query, params = store._construct_batched_query_and_params(config, reads) + where_clause = query.as_string(None).split("WHERE")[1] + + assert ("feature_name" in where_clause) is exp_feature_filter + assert len(params) == exp_params