Skip to content
Merged
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
29 changes: 22 additions & 7 deletions sdk/python/feast/infra/online_stores/bigtable.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,7 +7,7 @@
import google
from google.cloud import bigtable
from google.cloud.bigtable import row_filters
from pydantic import StrictStr
from pydantic import PositiveInt, StrictStr

from feast import Entity, FeatureView, utils
from feast.feature_view import DUMMY_ENTITY_NAME
Expand All @@ -19,11 +19,13 @@

logger = logging.getLogger(__name__)

# Number of mutations per Bigtable write operation we're aiming for. The official max is
# 100K; we're being conservative.
# Default number of mutations per Bigtable write operation we're aiming for. The official
# max is 100K; we're being conservative. Overridable via
# ``BigtableOnlineStoreConfig.mutations_per_write``.
MUTATIONS_PER_OP = 50_000
# The Bigtable client library limits the connection pool size to 10. This imposes a
# limitation to the concurrency we can get using a thread pool in each worker.
# Default thread-pool size used to parallelize writes within a worker. The Bigtable client
# library limits the connection pool size to 10, which bounds the useful concurrency.
# Overridable via ``BigtableOnlineStoreConfig.write_concurrency``.
BIGTABLE_CLIENT_CONNECTION_POOL_SIZE = 10


Expand All @@ -42,6 +44,16 @@ class BigtableOnlineStoreConfig(FeastConfigBaseModel):
max_versions: int = 2
"""The number of historical versions of data that will be kept around."""

mutations_per_write: PositiveInt = MUTATIONS_PER_OP
"""(optional) Target number of cell mutations per Bigtable write request. Bigtable's
hard limit is 100K; lower values produce smaller requests, which reduces the write
burst a materialization places on the instance. Defaults to 50,000."""

write_concurrency: PositiveInt = BIGTABLE_CLIENT_CONNECTION_POOL_SIZE
"""(optional) Maximum number of concurrent write requests (thread-pool size) issued
per worker during ``online_write_batch``. Lower values reduce the write load placed
on the Bigtable instance at the cost of longer materialization. Defaults to 10."""


class BigtableOnlineStore(OnlineStore):
_client: Optional[bigtable.Client] = None
Expand Down Expand Up @@ -133,10 +145,13 @@ def online_write_batch(
# `columns_per_row` is used to calculate the number of rows we are allowed to
# mutate in one request.
columns_per_row = len(feature_view.features) + 1 # extra for event timestamp
rows_per_write = MUTATIONS_PER_OP // columns_per_row
# At least one row per request even when a feature view is very wide.
rows_per_write = max(
1, config.online_store.mutations_per_write // columns_per_row
)

with futures.ThreadPoolExecutor(
max_workers=BIGTABLE_CLIENT_CONNECTION_POOL_SIZE
max_workers=config.online_store.write_concurrency
) as executor:
fs = []
while data:
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,119 @@
from dataclasses import dataclass, field
from datetime import datetime, timezone
from typing import List
from unittest.mock import patch

import pytest

pytest.importorskip("google.cloud.bigtable")

from concurrent import futures # noqa: E402

import feast.infra.online_stores.bigtable as bigtable_module # noqa: E402
from feast.infra.offline_stores.dask import DaskOfflineStoreConfig # noqa: E402
from feast.infra.online_stores.bigtable import ( # noqa: E402
BigtableOnlineStore,
BigtableOnlineStoreConfig,
)
from feast.repo_config import RepoConfig # noqa: E402


@dataclass
class MockFeatureView:
name: str
features: List[object] = field(default_factory=list)


def _make_repo_config(**online_store_kwargs) -> RepoConfig:
return RepoConfig(
registry="s3://test_registry/registry.db",
project="test_bigtable",
provider="local",
online_store=BigtableOnlineStoreConfig(
instance="test-instance",
project_id="test-project",
**online_store_kwargs,
),
offline_store=DaskOfflineStoreConfig(),
entity_key_serialization_version=3,
)


def _rows(n: int):
# Contents are irrelevant: _write_rows_to_bt is mocked, only the slicing
# performed by online_write_batch is under test.
return [(None, {}, datetime.now(timezone.utc), None) for _ in range(n)]


def _run_write(store, config, feature_view, data, captured_chunk_sizes):
def _capture(rows_to_write, **_):
captured_chunk_sizes.append(len(rows_to_write))

with (
patch.object(BigtableOnlineStore, "_get_client"),
patch.object(BigtableOnlineStore, "_get_table_name", return_value="tbl"),
patch.object(BigtableOnlineStore, "_write_rows_to_bt", side_effect=_capture),
):
store.online_write_batch(config, feature_view, data, None)


def test_defaults_match_legacy_constants():
config = BigtableOnlineStoreConfig(instance="test-instance")
assert config.mutations_per_write == bigtable_module.MUTATIONS_PER_OP == 50_000
assert (
config.write_concurrency
== bigtable_module.BIGTABLE_CLIENT_CONNECTION_POOL_SIZE
== 10
)


@pytest.mark.parametrize("field_name", ["mutations_per_write", "write_concurrency"])
@pytest.mark.parametrize("bad_value", [0, -1])
def test_positive_int_validation(field_name, bad_value):
from pydantic import ValidationError

with pytest.raises(ValidationError):
BigtableOnlineStoreConfig(instance="test-instance", **{field_name: bad_value})


def test_online_write_batch_respects_mutations_per_write():
# 1 feature -> columns_per_row = 2 -> rows_per_write = max(1, 10 // 2) = 5
store = BigtableOnlineStore()
config = _make_repo_config(mutations_per_write=10)
feature_view = MockFeatureView(name="fv", features=[object()])

captured: List[int] = []
_run_write(store, config, feature_view, _rows(12), captured)

assert captured == [5, 5, 2]


def test_rows_per_write_floors_to_one_for_wide_feature_view():
# 10 features -> columns_per_row = 11 -> 2 // 11 = 0 -> floored to 1
store = BigtableOnlineStore()
config = _make_repo_config(mutations_per_write=2)
feature_view = MockFeatureView(name="fv", features=[object() for _ in range(10)])

captured: List[int] = []
_run_write(store, config, feature_view, _rows(3), captured)

assert captured == [1, 1, 1]


def test_write_concurrency_sets_thread_pool_size():
store = BigtableOnlineStore()
config = _make_repo_config(write_concurrency=3)
feature_view = MockFeatureView(name="fv", features=[object()])

with (
patch.object(BigtableOnlineStore, "_get_client"),
patch.object(BigtableOnlineStore, "_get_table_name", return_value="tbl"),
patch.object(BigtableOnlineStore, "_write_rows_to_bt"),
patch(
"feast.infra.online_stores.bigtable.futures.ThreadPoolExecutor",
wraps=futures.ThreadPoolExecutor,
) as mock_pool,
):
store.online_write_batch(config, feature_view, _rows(4), None)

mock_pool.assert_called_once_with(max_workers=3)
Loading