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
5 changes: 3 additions & 2 deletions pyobvector/client/ob_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -159,7 +159,7 @@ def _is_seekdb(self) -> bool:
return is_seekdb

def _flush_seekdb_index(self) -> None:
"""Flush async HNSW index builds in embedded seekdb after insert.
"""Flush async HNSW index builds in embedded seekdb after a write.

No-op when not using embedded seekdb or when the server does not expose
a ``refresh_index`` method.
Expand All @@ -169,7 +169,7 @@ def _flush_seekdb_index(self) -> None:
try:
server.refresh_index()
except Exception as e:
logger.warning("seekdb index refresh failed after insert: %s", e)
logger.warning("seekdb index refresh failed after write: %s", e)

def _insert_partition_hint_for_query_sql(self, sql: str, partition_hint: str):
from_index = sql.find("FROM")
Expand Down Expand Up @@ -327,6 +327,7 @@ def upsert(
)
upsert_stmt = upsert_stmt.values(data)
conn.execute(upsert_stmt)
self._flush_seekdb_index()

def update(
self,
Expand Down
27 changes: 27 additions & 0 deletions tests/test_ob_client.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,27 @@
from unittest import TestCase
from unittest.mock import MagicMock, patch

from pyobvector.client.ob_client import ObClient


class TestObClient(TestCase):
def test_upsert_refreshes_embedded_seekdb_index(self) -> None:
client = ObClient.__new__(ObClient)
client.engine = MagicMock()
client.metadata_obj = MagicMock()
client._flush_seekdb_index = MagicMock()

table = MagicMock()
upsert_statement = MagicMock()
upsert_statement.values.return_value = upsert_statement

with (
patch("pyobvector.client.ob_client.Table", return_value=table),
patch(
"pyobvector.client.ob_client.ReplaceStmt",
return_value=upsert_statement,
),
):
client.upsert("test_table", [{"id": "doc-1"}])

client._flush_seekdb_index.assert_called_once_with()
Loading