diff --git a/pyobvector/client/ob_client.py b/pyobvector/client/ob_client.py index 8a0c010..15f266b 100644 --- a/pyobvector/client/ob_client.py +++ b/pyobvector/client/ob_client.py @@ -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. @@ -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") @@ -327,6 +327,7 @@ def upsert( ) upsert_stmt = upsert_stmt.values(data) conn.execute(upsert_stmt) + self._flush_seekdb_index() def update( self, diff --git a/tests/test_ob_client.py b/tests/test_ob_client.py new file mode 100644 index 0000000..454ebb7 --- /dev/null +++ b/tests/test_ob_client.py @@ -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()