Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
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
Prev Previous commit
fix: Load Milvus collections once instead of on every query
online_read and retrieve_online_documents_v2 called load_collection on
every request, adding a round-trip to each read. The call also hid the
fact that newly created collections were never loaded on Milvus
servers, because indexes were created after the collection.

Collections are now created with their index params, which makes
Milvus load them straight away. Existing collections are loaded only
if their load state isn't Loaded, and only the first time a store
accesses them. The per-query load_collection calls are removed.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Signed-off-by: Simon Hearne <simon.hearne@gmail.com>
  • Loading branch information
2 people authored and ntkathole committed Sep 29, 2026
commit f75c4445eaa3074df1cede85f1423fe848dff601
7 changes: 7 additions & 0 deletions docs/reference/online-stores/milvus.md
Original file line number Diff line number Diff line change
Expand Up @@ -76,6 +76,13 @@ online_store:

The full set of configuration options is available in [MilvusOnlineStoreConfig](https://rtd.feast.dev/en/latest/#feast.infra.online_stores.milvus.MilvusOnlineStoreConfig).

## Collection loading

Feast creates collections together with their indexes, which makes Milvus load them straight away.
When Feast finds an existing collection it checks its load state and loads it only if needed.
Reads and searches never load collections, so a collection released outside Feast is only reloaded
the next time a Feast process first accesses it.

## Feature views without vectors

Milvus requires every collection to have a vector field. For feature views that have no vector
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
FieldSchema,
MilvusClient,
)
from pymilvus.client.types import LoadState

from feast import Entity
from feast.feature_view import FeatureView
Expand Down Expand Up @@ -367,13 +368,7 @@ def _get_or_create_collection(
collection_name=collection_name
)
if not collection_exists:
self.client.create_collection(
collection_name=collection_name,
dimension=config.online_store.embedding_dim,
schema=schema,
)
index_params = self.client.prepare_index_params()
indices_added = False
for vector_field in schema.fields:
if vector_field.dtype not in [
DataType.FLOAT_VECTOR,
Expand Down Expand Up @@ -405,19 +400,31 @@ def _get_or_create_collection(
index_type="FLAT",
index_name=f"vector_index_{vector_field.name}",
)
indices_added = True
if indices_added:
self.client.create_index(
collection_name=collection_name,
index_params=index_params,
)
# Every collection has at least one vector field, and every
# vector field is indexed, so passing the index params here
# makes Milvus create the indexes and load the collection.
self.client.create_collection(
collection_name=collection_name,
dimension=config.online_store.embedding_dim,
schema=schema,
index_params=index_params,
)
else:
self.client.load_collection(collection_name)
self._ensure_loaded(collection_name)
# Collections are only cached once loaded, so reads and searches
# don't need to load them again.
self._collections[collection_name] = self.client.describe_collection(
collection_name
)
return self._collections[collection_name]

def _ensure_loaded(self, collection_name: str) -> None:
"""Load an existing collection unless Milvus already has it loaded."""
assert self.client is not None, "Milvus client is not initialized"
load_state = self.client.get_load_state(collection_name).get("state")
if load_state != LoadState.Loaded:
self.client.load_collection(collection_name)

def online_write_batch(
self,
config: RepoConfig,
Expand Down Expand Up @@ -560,7 +567,6 @@ def online_read(
+ ", ".join([f"'{e}'" for e in composite_entities])
+ "]"
)
self.client.load_collection(collection_name)
results = self.client.query(
collection_name=collection_name,
filter=query_filter_for_entities,
Expand Down Expand Up @@ -773,8 +779,6 @@ def retrieve_online_documents_v2(
ann_search_field = field["name"]
break

self.client.load_collection(collection_name)

if filters and filters_contain_numeric_comparison(filters):
collection_field_types = {
f["name"]: f["type"] for f in collection["fields"]
Expand Down
108 changes: 107 additions & 1 deletion sdk/python/tests/integration/online_store/test_milvus_remote.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
from datetime import datetime, timedelta, timezone
from pathlib import Path
from typing import Any, Callable, Dict, Iterator, List, Optional, TypeVar
from unittest.mock import patch
from urllib.parse import urlparse

import pytest
Expand All @@ -25,7 +26,7 @@
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 RepoConfig
from feast.types import Float32, Int64, String
from feast.types import Array, Float32, Int64, String
from feast.value_type import ValueType

T = TypeVar("T")
Expand Down Expand Up @@ -171,3 +172,108 @@ def test_scalar_feature_view_round_trip(
)

assert rows[0] is not None and rows[0]["city"].string_val == "Paris"


def _vector_feature_view(name: str = "driver_embeddings") -> FeatureView:
return FeatureView(
name=name,
entities=[
Entity(
name="driver_id", join_keys=["driver_id"], value_type=ValueType.INT64
)
],
ttl=timedelta(days=1),
schema=[
Field(name="driver_id", dtype=Int64),
Field(
name="embedding",
dtype=Array(Float32),
vector_index=True,
vector_search_metric="COSINE",
),
Field(name="city", dtype=String),
],
)


def _vector_rows() -> Dict[int, Dict[str, ValueProto]]:
def embedding(x: float, y: float) -> ValueProto:
value = ValueProto()
value.float_list_val.val.extend([x, y])
return value

return {
1: {"embedding": embedding(1.0, 0.0), "city": ValueProto(string_val="Paris")},
2: {"embedding": embedding(0.0, 1.0), "city": ValueProto(string_val="Rome")},
}


def _search(
store: MilvusOnlineStore,
config: RepoConfig,
fv: FeatureView,
embedding: List[float],
top_k: int = 1,
**kwargs: Any,
) -> List[Dict[str, ValueProto]]:
results = store.retrieve_online_documents_v2(
config,
fv,
["embedding", "city"],
embedding=embedding,
top_k=top_k,
distance_metric="COSINE",
**kwargs,
)
return [values for _, _, values in results if values]


def test_load_collection_not_called_per_query(
tmp_path: Path, project: str, store: MilvusOnlineStore
) -> None:
config = _repo_config(tmp_path, project)
fv = _vector_feature_view()
store.update(config, [], [fv], [], [], partial=False)
_write_rows(store, config, fv, _vector_rows())

assert store.client is not None
with patch.object(
store.client, "load_collection", wraps=store.client.load_collection
) as load_spy:
hits = _eventually(
lambda: _search(store, config, fv, [1.0, 0.0]),
lambda hits: len(hits) == 1,
)
for _ in range(3):
_read(store, config, fv, [1, 2], ["city"])
_search(store, config, fv, [1.0, 0.0])

assert hits[0]["city"].string_val == "Paris"
assert load_spy.call_count == 0


def test_released_collection_is_loaded_once(
tmp_path: Path, project: str, store: MilvusOnlineStore
) -> None:
config = _repo_config(tmp_path, project)
fv = _vector_feature_view()
store.update(config, [], [fv], [], [], partial=False)
_write_rows(store, config, fv, _vector_rows())
assert store.client is not None
collection_name = f"{project}_{fv.name}"
store.client.release_collection(collection_name)

# A fresh store, e.g. a new feature server process, finds it unloaded.
fresh_store = MilvusOnlineStore()
fresh_store.client = store.client
with patch.object(
store.client, "load_collection", wraps=store.client.load_collection
) as load_spy:
rows = _eventually(
lambda: _read(fresh_store, config, fv, [1], ["city"]),
lambda rows: rows[0] is not None,
)
_read(fresh_store, config, fv, [1], ["city"])

assert rows[0] is not None and rows[0]["city"].string_val == "Paris"
assert load_spy.call_count == 1
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,9 @@
from typing import Any, Dict, List, Optional
from unittest.mock import MagicMock, patch

import pytest
from pymilvus import DataType, MilvusClient
from pymilvus.client.types import LoadState

from feast import Entity, FeatureView
from feast.field import Field
Expand Down Expand Up @@ -211,3 +213,91 @@ def test_placeholder_values_are_finite_and_match_collection_dim(

data = mock_client.upsert.call_args.kwargs["data"]
assert data[0][PLACEHOLDER_VECTOR_FIELD] == [0.0]


def _existing_collection_description() -> Dict[str, Any]:
return {
"collection_name": "test_milvus_driver_stats",
"fields": [
{"name": "driver_id_pk", "type": DataType.VARCHAR, "params": {}},
{"name": "driver_id", "type": DataType.VARCHAR, "params": {}},
{"name": "event_ts", "type": DataType.INT64, "params": {}},
{"name": "created_ts", "type": DataType.INT64, "params": {}},
{"name": "trips_today", "type": DataType.VARCHAR, "params": {}},
{"name": "city", "type": DataType.VARCHAR, "params": {}},
{
"name": PLACEHOLDER_VECTOR_FIELD,
"type": DataType.FLOAT_VECTOR,
"params": {"dim": PLACEHOLDER_VECTOR_DIM},
},
],
}


@pytest.mark.parametrize(
"load_state, expected_loads",
[(LoadState.Loaded, 0), (LoadState.NotLoad, 1)],
)
@patch(f"{MILVUS_MODULE}.MilvusClient")
def test_existing_collection_is_loaded_at_most_once(
mock_client_cls: MagicMock, load_state: LoadState, expected_loads: int
) -> None:
mock_client = _mock_client(mock_client_cls, has_collection=True)
mock_client.describe_collection.return_value = _existing_collection_description()
mock_client.get_load_state.return_value = {"state": load_state}
mock_client.query.return_value = []

store = MilvusOnlineStore()
config = _mock_config()
fv = _scalar_feature_view()
for _ in range(5):
store.online_read(config, fv, [_entity_key(1)], ["city"])

assert mock_client.load_collection.call_count == expected_loads
assert mock_client.query.call_count == 5


@patch(f"{MILVUS_MODULE}.MilvusClient")
def test_new_collection_is_created_with_indexes_so_it_loads(
mock_client_cls: MagicMock,
) -> None:
mock_client = _mock_client(mock_client_cls, has_collection=False)

store = MilvusOnlineStore()
store._get_or_create_collection(_mock_config(), _scalar_feature_view())

# MilvusClient.create_collection loads the collection when it is given
# index params; creating indexes separately would leave it unloaded.
assert mock_client.create_collection.call_args.kwargs["index_params"]
mock_client.create_index.assert_not_called()
mock_client.load_collection.assert_not_called()


def test_load_collection_not_called_per_query(tmp_path: Path) -> None:
config = _lite_config(tmp_path)
fv = _scalar_feature_view()
store = MilvusOnlineStore()
store.update(config, [], [fv], [], [], partial=False)
_write_rows(
store,
config,
fv,
{
1: {
"trips_today": ValueProto(float_val=1.0),
"city": ValueProto(string_val="Oslo"),
}
},
)

assert store.client is not None
with patch.object(
store.client, "load_collection", wraps=store.client.load_collection
) as load_spy:
for _ in range(3):
_read(store, config, fv, [1], ["city"])
store.retrieve_online_documents_v2(
config, fv, ["city"], embedding=None, top_k=1, query_string="Oslo"
)

assert load_spy.call_count == 0