feat: add Spanner vector_store_similarity_search tool

The vector_store_similarity_search tool performs similarity search against data in a Spanner vector store table, using the provided Spanner tool settings for configuration.

PiperOrigin-RevId: 839352057
This commit is contained in:
Google Team Member
2025-12-02 11:18:57 -08:00
committed by Copybara-Service
parent 8da61be45a
commit 090711934f
9 changed files with 948 additions and 350 deletions
+230 -51
View File
@@ -12,10 +12,12 @@
# See the License for the specific language governing permissions and
# limitations under the License.
from unittest import mock
from unittest.mock import MagicMock
from unittest.mock import patch
from google.adk.tools.spanner import client
from google.adk.tools.spanner import search_tool
from google.adk.tools.spanner import utils
from google.cloud.spanner_admin_database_v1.types import DatabaseDialect
import pytest
@@ -35,29 +37,59 @@ def mock_spanner_ids():
}
@patch("google.adk.tools.spanner.client.get_spanner_client")
@pytest.mark.parametrize(
("embedding_option_key", "embedding_option_value", "expected_embedding"),
[
pytest.param(
"spanner_googlesql_embedding_model_name",
"EmbeddingsModel",
[0.1, 0.2, 0.3],
id="spanner_googlesql_embedding_model",
),
pytest.param(
"vertex_ai_embedding_model_name",
"text-embedding-005",
[0.4, 0.5, 0.6],
id="vertex_ai_embedding_model",
),
],
)
@mock.patch.object(utils, "embed_contents")
@mock.patch.object(client, "get_spanner_client")
def test_similarity_search_knn_success(
mock_get_spanner_client, mock_spanner_ids, mock_credentials
mock_get_spanner_client,
mock_embed_contents,
mock_spanner_ids,
mock_credentials,
embedding_option_key,
embedding_option_value,
expected_embedding,
):
"""Test similarity_search function with kNN success."""
mock_spanner_client = MagicMock()
mock_instance = MagicMock()
mock_database = MagicMock()
mock_snapshot = MagicMock()
mock_embedding_result = MagicMock()
mock_embedding_result.one.return_value = ([0.1, 0.2, 0.3],)
# First call to execute_sql is for getting the embedding
# Second call is for the kNN search
mock_snapshot.execute_sql.side_effect = [
mock_embedding_result,
iter([("result1",), ("result2",)]),
]
mock_database.snapshot.return_value.__enter__.return_value = mock_snapshot
mock_database.database_dialect = DatabaseDialect.GOOGLE_STANDARD_SQL
mock_instance.database.return_value = mock_database
mock_spanner_client.instance.return_value = mock_instance
mock_get_spanner_client.return_value = mock_spanner_client
if embedding_option_key == "vertex_ai_embedding_model_name":
mock_embed_contents.return_value = [expected_embedding]
# execute_sql is called once for the kNN search
mock_snapshot.execute_sql.return_value = iter([("result1",), ("result2",)])
else:
mock_embedding_result = MagicMock()
mock_embedding_result.one.return_value = (expected_embedding,)
# First call to execute_sql is for getting the embedding,
# second call is for the kNN search
mock_snapshot.execute_sql.side_effect = [
mock_embedding_result,
iter([("result1",), ("result2",)]),
]
result = search_tool.similarity_search(
project_id=mock_spanner_ids["project_id"],
instance_id=mock_spanner_ids["instance_id"],
@@ -66,10 +98,8 @@ def test_similarity_search_knn_success(
query="test query",
embedding_column_to_search="embedding_col",
columns=["col1"],
embedding_options={"spanner_embedding_model_name": "test_model"},
embedding_options={embedding_option_key: embedding_option_value},
credentials=mock_credentials,
settings=MagicMock(),
tool_context=MagicMock(),
)
assert result["status"] == "SUCCESS", result
assert result["rows"] == [("result1",), ("result2",)]
@@ -79,10 +109,14 @@ def test_similarity_search_knn_success(
sql = call_args.args[0]
assert "COSINE_DISTANCE" in sql
assert "@embedding" in sql
assert call_args.kwargs == {"params": {"embedding": [0.1, 0.2, 0.3]}}
assert call_args.kwargs == {"params": {"embedding": expected_embedding}}
if embedding_option_key == "vertex_ai_embedding_model_name":
mock_embed_contents.assert_called_once_with(
embedding_option_value, ["test query"], None
)
@patch("google.adk.tools.spanner.client.get_spanner_client")
@mock.patch.object(client, "get_spanner_client")
def test_similarity_search_ann_success(
mock_get_spanner_client, mock_spanner_ids, mock_credentials
):
@@ -113,10 +147,10 @@ def test_similarity_search_ann_success(
query="test query",
embedding_column_to_search="embedding_col",
columns=["col1"],
embedding_options={"spanner_embedding_model_name": "test_model"},
embedding_options={
"spanner_googlesql_embedding_model_name": "test_model"
},
credentials=mock_credentials,
settings=MagicMock(),
tool_context=MagicMock(),
search_options={
"nearest_neighbors_algorithm": "APPROXIMATE_NEAREST_NEIGHBORS"
},
@@ -130,7 +164,7 @@ def test_similarity_search_ann_success(
assert call_args.kwargs == {"params": {"embedding": [0.1, 0.2, 0.3]}}
@patch("google.adk.tools.spanner.client.get_spanner_client")
@mock.patch.object(client, "get_spanner_client")
def test_similarity_search_error(
mock_get_spanner_client, mock_spanner_ids, mock_credentials
):
@@ -143,17 +177,17 @@ def test_similarity_search_error(
table_name=mock_spanner_ids["table_name"],
query="test query",
embedding_column_to_search="embedding_col",
embedding_options={"spanner_embedding_model_name": "test_model"},
embedding_options={
"spanner_googlesql_embedding_model_name": "test_model"
},
columns=["col1"],
credentials=mock_credentials,
settings=MagicMock(),
tool_context=MagicMock(),
)
assert result["status"] == "ERROR"
assert result["error_details"] == "Test Exception"
assert "Test Exception" in result["error_details"]
@patch("google.adk.tools.spanner.client.get_spanner_client")
@mock.patch.object(client, "get_spanner_client")
def test_similarity_search_postgresql_knn_success(
mock_get_spanner_client, mock_spanner_ids, mock_credentials
):
@@ -182,10 +216,12 @@ def test_similarity_search_postgresql_knn_success(
query="test query",
embedding_column_to_search="embedding_col",
columns=["col1"],
embedding_options={"vertex_ai_embedding_model_endpoint": "test_endpoint"},
embedding_options={
"spanner_postgresql_vertex_ai_embedding_model_endpoint": (
"test_endpoint"
)
},
credentials=mock_credentials,
settings=MagicMock(),
tool_context=MagicMock(),
)
assert result["status"] == "SUCCESS", result
assert result["rows"] == [("pg_result",)]
@@ -196,7 +232,7 @@ def test_similarity_search_postgresql_knn_success(
assert call_args.kwargs == {"params": {"p1": [0.1, 0.2, 0.3]}}
@patch("google.adk.tools.spanner.client.get_spanner_client")
@mock.patch.object(client, "get_spanner_client")
def test_similarity_search_postgresql_ann_unsupported(
mock_get_spanner_client, mock_spanner_ids, mock_credentials
):
@@ -217,27 +253,28 @@ def test_similarity_search_postgresql_ann_unsupported(
query="test query",
embedding_column_to_search="embedding_col",
columns=["col1"],
embedding_options={"vertex_ai_embedding_model_endpoint": "test_endpoint"},
embedding_options={
"spanner_postgresql_vertex_ai_embedding_model_endpoint": (
"test_endpoint"
)
},
credentials=mock_credentials,
settings=MagicMock(),
tool_context=MagicMock(),
search_options={
"nearest_neighbors_algorithm": "APPROXIMATE_NEAREST_NEIGHBORS"
},
)
assert result["status"] == "ERROR"
assert (
result["error_details"]
== "APPROXIMATE_NEAREST_NEIGHBORS is not supported for PostgreSQL"
" dialect."
"APPROXIMATE_NEAREST_NEIGHBORS is not supported for PostgreSQL dialect."
in result["error_details"]
)
@patch("google.adk.tools.spanner.client.get_spanner_client")
def test_similarity_search_missing_spanner_embedding_model_name_error(
@mock.patch.object(client, "get_spanner_client")
def test_similarity_search_gsql_missing_embedding_model_error(
mock_get_spanner_client, mock_spanner_ids, mock_credentials
):
"""Test similarity_search with missing spanner_embedding_model_name."""
"""Test similarity_search with missing embedding_options for GoogleSQL dialect."""
mock_spanner_client = MagicMock()
mock_instance = MagicMock()
mock_database = MagicMock()
@@ -254,24 +291,27 @@ def test_similarity_search_missing_spanner_embedding_model_name_error(
query="test query",
embedding_column_to_search="embedding_col",
columns=["col1"],
embedding_options={},
embedding_options={
"spanner_postgresql_vertex_ai_embedding_model_endpoint": (
"test_endpoint"
)
},
credentials=mock_credentials,
settings=MagicMock(),
tool_context=MagicMock(),
)
assert result["status"] == "ERROR"
assert (
"embedding_options['spanner_embedding_model_name'] must be"
" specified for GoogleSQL dialect."
"embedding_options['vertex_ai_embedding_model_name'] or"
" embedding_options['spanner_googlesql_embedding_model_name'] must be"
" specified for GoogleSQL dialect Spanner database."
in result["error_details"]
)
@patch("google.adk.tools.spanner.client.get_spanner_client")
def test_similarity_search_missing_vertex_ai_embedding_model_endpoint_error(
@mock.patch.object(client, "get_spanner_client")
def test_similarity_search_pg_missing_embedding_model_error(
mock_get_spanner_client, mock_spanner_ids, mock_credentials
):
"""Test similarity_search with missing vertex_ai_embedding_model_endpoint."""
"""Test similarity_search with missing embedding_options for PostgreSQL dialect."""
mock_spanner_client = MagicMock()
mock_instance = MagicMock()
mock_database = MagicMock()
@@ -288,14 +328,153 @@ def test_similarity_search_missing_vertex_ai_embedding_model_endpoint_error(
query="test query",
embedding_column_to_search="embedding_col",
columns=["col1"],
embedding_options={},
embedding_options={
"spanner_googlesql_embedding_model_name": "EmbeddingsModel"
},
credentials=mock_credentials,
settings=MagicMock(),
tool_context=MagicMock(),
)
assert result["status"] == "ERROR"
assert (
"embedding_options['vertex_ai_embedding_model_endpoint'] must "
"be specified for PostgreSQL dialect."
"embedding_options['vertex_ai_embedding_model_name'] or"
" embedding_options['spanner_postgresql_vertex_ai_embedding_model_endpoint']"
" must be specified for PostgreSQL dialect Spanner database."
in result["error_details"]
)
@pytest.mark.parametrize(
"embedding_options",
[
pytest.param(
{
"vertex_ai_embedding_model_name": "test-model",
"spanner_googlesql_embedding_model_name": "test-model-2",
},
id="vertex_ai_and_googlesql",
),
pytest.param(
{
"vertex_ai_embedding_model_name": "test-model",
"spanner_postgresql_vertex_ai_embedding_model_endpoint": (
"test-endpoint"
),
},
id="vertex_ai_and_postgresql",
),
pytest.param(
{
"spanner_googlesql_embedding_model_name": "test-model",
"spanner_postgresql_vertex_ai_embedding_model_endpoint": (
"test-endpoint"
),
},
id="googlesql_and_postgresql",
),
pytest.param(
{
"vertex_ai_embedding_model_name": "test-model",
"spanner_googlesql_embedding_model_name": "test-model-2",
"spanner_postgresql_vertex_ai_embedding_model_endpoint": (
"test-endpoint"
),
},
id="all_three_models",
),
pytest.param(
{},
id="no_models",
),
],
)
@mock.patch.object(client, "get_spanner_client")
def test_similarity_search_multiple_embedding_options_error(
mock_get_spanner_client,
mock_spanner_ids,
mock_credentials,
embedding_options,
):
"""Test similarity_search with multiple embedding models."""
mock_spanner_client = MagicMock()
mock_instance = MagicMock()
mock_database = MagicMock()
mock_database.database_dialect = DatabaseDialect.GOOGLE_STANDARD_SQL
mock_instance.database.return_value = mock_database
mock_spanner_client.instance.return_value = mock_instance
mock_get_spanner_client.return_value = mock_spanner_client
result = search_tool.similarity_search(
project_id=mock_spanner_ids["project_id"],
instance_id=mock_spanner_ids["instance_id"],
database_id=mock_spanner_ids["database_id"],
table_name=mock_spanner_ids["table_name"],
query="test query",
embedding_column_to_search="embedding_col",
columns=["col1"],
embedding_options=embedding_options,
credentials=mock_credentials,
)
assert result["status"] == "ERROR"
assert (
"Exactly one embedding model option must be specified."
in result["error_details"]
)
@mock.patch.object(client, "get_spanner_client")
def test_similarity_search_output_dimensionality_gsql_error(
mock_get_spanner_client, mock_spanner_ids, mock_credentials
):
"""Test similarity_search with output_dimensionality and spanner_googlesql_embedding_model_name."""
mock_spanner_client = MagicMock()
mock_instance = MagicMock()
mock_database = MagicMock()
mock_database.database_dialect = DatabaseDialect.GOOGLE_STANDARD_SQL
mock_instance.database.return_value = mock_database
mock_spanner_client.instance.return_value = mock_instance
mock_get_spanner_client.return_value = mock_spanner_client
result = search_tool.similarity_search(
project_id=mock_spanner_ids["project_id"],
instance_id=mock_spanner_ids["instance_id"],
database_id=mock_spanner_ids["database_id"],
table_name=mock_spanner_ids["table_name"],
query="test query",
embedding_column_to_search="embedding_col",
columns=["col1"],
embedding_options={
"spanner_googlesql_embedding_model_name": "EmbeddingsModel",
"output_dimensionality": 128,
},
credentials=mock_credentials,
)
assert result["status"] == "ERROR"
assert "is not supported when" in result["error_details"]
@mock.patch.object(client, "get_spanner_client")
def test_similarity_search_unsupported_algorithm_error(
mock_get_spanner_client, mock_spanner_ids, mock_credentials
):
"""Test similarity_search with an unsupported nearest neighbors algorithm."""
mock_spanner_client = MagicMock()
mock_instance = MagicMock()
mock_database = MagicMock()
mock_database.database_dialect = DatabaseDialect.GOOGLE_STANDARD_SQL
mock_instance.database.return_value = mock_database
mock_spanner_client.instance.return_value = mock_instance
mock_get_spanner_client.return_value = mock_spanner_client
result = search_tool.similarity_search(
project_id=mock_spanner_ids["project_id"],
instance_id=mock_spanner_ids["instance_id"],
database_id=mock_spanner_ids["database_id"],
table_name=mock_spanner_ids["table_name"],
query="test query",
embedding_column_to_search="embedding_col",
columns=["col1"],
embedding_options={"vertex_ai_embedding_model_name": "test-model"},
credentials=mock_credentials,
search_options={"nearest_neighbors_algorithm": "INVALID_ALGORITHM"},
)
assert result["status"] == "ERROR"
assert "Unsupported search_options" in result["error_details"]
@@ -15,9 +15,23 @@
from __future__ import annotations
from google.adk.tools.spanner.settings import SpannerToolSettings
from google.adk.tools.spanner.settings import SpannerVectorStoreSettings
from pydantic import ValidationError
import pytest
def common_spanner_vector_store_settings(vector_length=None):
return {
"project_id": "test-project",
"instance_id": "test-instance",
"database_id": "test-database",
"table_name": "test-table",
"content_column": "test-content-column",
"embedding_column": "test-embedding-column",
"vector_length": 128 if vector_length is None else vector_length,
}
def test_spanner_tool_settings_experimental_warning():
"""Test SpannerToolSettings experimental warning."""
with pytest.warns(
@@ -25,3 +39,34 @@ def test_spanner_tool_settings_experimental_warning():
match="Tool settings defaults may have breaking change in the future.",
):
SpannerToolSettings()
def test_spanner_vector_store_settings_all_fields_present():
"""Test SpannerVectorStoreSettings with all required fields present."""
settings = SpannerVectorStoreSettings(
**common_spanner_vector_store_settings(),
vertex_ai_embedding_model_name="test-embedding-model",
)
assert settings is not None
assert settings.selected_columns == ["test-content-column"]
assert settings.vertex_ai_embedding_model_name == "test-embedding-model"
def test_spanner_vector_store_settings_missing_embedding_model_name():
"""Test SpannerVectorStoreSettings with missing vertex_ai_embedding_model_name."""
with pytest.raises(ValidationError) as excinfo:
SpannerVectorStoreSettings(**common_spanner_vector_store_settings())
assert "Field required" in str(excinfo.value)
assert "vertex_ai_embedding_model_name" in str(excinfo.value)
def test_spanner_vector_store_settings_invalid_vector_length():
"""Test SpannerVectorStoreSettings with invalid vector_length."""
with pytest.raises(ValidationError) as excinfo:
SpannerVectorStoreSettings(
**common_spanner_vector_store_settings(vector_length=0),
vertex_ai_embedding_model_name="test-embedding-model",
)
assert "Invalid vector length in the Spanner vector store settings." in str(
excinfo.value
)
@@ -18,6 +18,7 @@ from google.adk.tools.google_tool import GoogleTool
from google.adk.tools.spanner import SpannerCredentialsConfig
from google.adk.tools.spanner import SpannerToolset
from google.adk.tools.spanner.settings import SpannerToolSettings
from google.adk.tools.spanner.settings import SpannerVectorStoreSettings
import pytest
@@ -184,3 +185,50 @@ async def test_spanner_toolset_without_read_capability(
expected_tool_names = set(returned_tools)
actual_tool_names = set([tool.name for tool in tools])
assert actual_tool_names == expected_tool_names
@pytest.mark.asyncio
async def test_spanner_toolset_with_vector_store_search():
"""Test Spanner toolset with vector store search.
This test verifies the behavior of the Spanner toolset when vector store
settings is provided.
"""
credentials_config = SpannerCredentialsConfig(
client_id="abc", client_secret="def"
)
spanner_tool_settings = SpannerToolSettings(
vector_store_settings=SpannerVectorStoreSettings(
project_id="test-project",
instance_id="test-instance",
database_id="test-database",
table_name="test-table",
content_column="test-content-column",
embedding_column="test-embedding-column",
vector_length=128,
vertex_ai_embedding_model_name="test-embedding-model",
)
)
toolset = SpannerToolset(
credentials_config=credentials_config,
spanner_tool_settings=spanner_tool_settings,
)
tools = await toolset.get_tools()
assert tools is not None
assert len(tools) == 8
assert all([isinstance(tool, GoogleTool) for tool in tools])
expected_tool_names = set([
"list_table_names",
"list_table_indexes",
"list_table_index_columns",
"list_named_schemas",
"get_table_schema",
"execute_sql",
"similarity_search",
"vector_store_similarity_search",
])
actual_tool_names = set([tool.name for tool in tools])
assert actual_tool_names == expected_tool_names