mirror of
https://github.com/encounter/adk-python.git
synced 2026-07-09 18:19:28 -07:00
feat: add Spanner similarity_search tool
Similarity search tool supports similarity search on Spanner data by embedding a text query to a vector and run vector search with the embedded vector. PiperOrigin-RevId: 806502499
This commit is contained in:
committed by
Copybara-Service
parent
0c1f1fadeb
commit
c29d41f0d0
@@ -0,0 +1,446 @@
|
|||||||
|
# Copyright 2025 Google LLC
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
from typing import Any
|
||||||
|
from typing import Dict
|
||||||
|
from typing import List
|
||||||
|
from typing import Optional
|
||||||
|
|
||||||
|
from google.adk.tools.spanner import client
|
||||||
|
from google.adk.tools.spanner.settings import SpannerToolSettings
|
||||||
|
from google.adk.tools.tool_context import ToolContext
|
||||||
|
from google.auth.credentials import Credentials
|
||||||
|
from google.cloud.spanner_admin_database_v1.types import DatabaseDialect
|
||||||
|
from google.cloud.spanner_v1.database import Database
|
||||||
|
|
||||||
|
# Embedding options
|
||||||
|
_SPANNER_EMBEDDING_MODEL_NAME = "spanner_embedding_model_name"
|
||||||
|
_VERTEX_AI_EMBEDDING_MODEL_ENDPOINT = "vertex_ai_embedding_model_endpoint"
|
||||||
|
|
||||||
|
# Search options
|
||||||
|
_TOP_K = "top_k"
|
||||||
|
_DISTANCE_TYPE = "distance_type"
|
||||||
|
_NEAREST_NEIGHBORS_ALGORITHM = "nearest_neighbors_algorithm"
|
||||||
|
_EXACT_NEAREST_NEIGHBORS = "EXACT_NEAREST_NEIGHBORS"
|
||||||
|
_APPROXIMATE_NEAREST_NEIGHBORS = "APPROXIMATE_NEAREST_NEIGHBORS"
|
||||||
|
_NUM_LEAVES_TO_SEARCH = "num_leaves_to_search"
|
||||||
|
|
||||||
|
# Constants
|
||||||
|
_DISTANCE_ALIAS = "distance"
|
||||||
|
_GOOGLESQL_PARAMETER_TEXT_QUERY = "query"
|
||||||
|
_POSTGRESQL_PARAMETER_TEXT_QUERY = "1"
|
||||||
|
_GOOGLESQL_PARAMETER_QUERY_EMBEDDING = "embedding"
|
||||||
|
_POSTGRESQL_PARAMETER_QUERY_EMBEDDING = "1"
|
||||||
|
|
||||||
|
|
||||||
|
def _generate_googlesql_for_embedding_query(
|
||||||
|
spanner_embedding_model_name: str,
|
||||||
|
) -> str:
|
||||||
|
return f"""
|
||||||
|
SELECT embeddings.values
|
||||||
|
FROM ML.PREDICT(
|
||||||
|
MODEL {spanner_embedding_model_name},
|
||||||
|
(SELECT CAST(@{_GOOGLESQL_PARAMETER_TEXT_QUERY} AS STRING) as content)
|
||||||
|
)
|
||||||
|
"""
|
||||||
|
|
||||||
|
|
||||||
|
def _generate_postgresql_for_embedding_query(
|
||||||
|
vertex_ai_embedding_model_endpoint: str,
|
||||||
|
) -> str:
|
||||||
|
return f"""
|
||||||
|
SELECT spanner.FLOAT32_ARRAY( spanner.ML_PREDICT_ROW(
|
||||||
|
'{vertex_ai_embedding_model_endpoint}',
|
||||||
|
JSONB_BUILD_OBJECT(
|
||||||
|
'instances',
|
||||||
|
JSONB_BUILD_ARRAY( JSONB_BUILD_OBJECT(
|
||||||
|
'content',
|
||||||
|
${_POSTGRESQL_PARAMETER_TEXT_QUERY}::TEXT
|
||||||
|
))
|
||||||
|
)
|
||||||
|
) -> 'predictions'->0->'embeddings'->'values' )
|
||||||
|
"""
|
||||||
|
|
||||||
|
|
||||||
|
def _get_embedding_for_query(
|
||||||
|
database: Database,
|
||||||
|
dialect: DatabaseDialect,
|
||||||
|
spanner_embedding_model_name: Optional[str],
|
||||||
|
vertex_ai_embedding_model_endpoint: Optional[str],
|
||||||
|
query: str,
|
||||||
|
) -> List[float]:
|
||||||
|
"""Gets the embedding for the query."""
|
||||||
|
if dialect == DatabaseDialect.POSTGRESQL:
|
||||||
|
embedding_query = _generate_postgresql_for_embedding_query(
|
||||||
|
vertex_ai_embedding_model_endpoint
|
||||||
|
)
|
||||||
|
params = {f"p{_POSTGRESQL_PARAMETER_TEXT_QUERY}": query}
|
||||||
|
else:
|
||||||
|
embedding_query = _generate_googlesql_for_embedding_query(
|
||||||
|
spanner_embedding_model_name
|
||||||
|
)
|
||||||
|
params = {_GOOGLESQL_PARAMETER_TEXT_QUERY: query}
|
||||||
|
with database.snapshot() as snapshot:
|
||||||
|
result_set = snapshot.execute_sql(embedding_query, params=params)
|
||||||
|
return result_set.one()[0]
|
||||||
|
|
||||||
|
|
||||||
|
def _get_postgresql_distance_function(distance_type: str) -> str:
|
||||||
|
return {
|
||||||
|
"COSINE_DISTANCE": "spanner.cosine_distance",
|
||||||
|
"EUCLIDEAN_DISTANCE": "spanner.euclidean_distance",
|
||||||
|
"DOT_PRODUCT": "spanner.dot_product",
|
||||||
|
}[distance_type]
|
||||||
|
|
||||||
|
|
||||||
|
def _get_googlesql_distance_function(distance_type: str, ann: bool) -> str:
|
||||||
|
if ann:
|
||||||
|
return {
|
||||||
|
"COSINE_DISTANCE": "APPROX_COSINE_DISTANCE",
|
||||||
|
"EUCLIDEAN_DISTANCE": "APPROX_EUCLIDEAN_DISTANCE",
|
||||||
|
"DOT_PRODUCT": "APPROX_DOT_PRODUCT",
|
||||||
|
}[distance_type]
|
||||||
|
return {
|
||||||
|
"COSINE_DISTANCE": "COSINE_DISTANCE",
|
||||||
|
"EUCLIDEAN_DISTANCE": "EUCLIDEAN_DISTANCE",
|
||||||
|
"DOT_PRODUCT": "DOT_PRODUCT",
|
||||||
|
}[distance_type]
|
||||||
|
|
||||||
|
|
||||||
|
def _generate_sql_for_knn(
|
||||||
|
dialect: DatabaseDialect,
|
||||||
|
table_name: str,
|
||||||
|
embedding_column_to_search: str,
|
||||||
|
columns,
|
||||||
|
additional_filter: Optional[str],
|
||||||
|
distance_type: str,
|
||||||
|
top_k: int,
|
||||||
|
) -> str:
|
||||||
|
"""Generates a SQL query for kNN search."""
|
||||||
|
if dialect == DatabaseDialect.POSTGRESQL:
|
||||||
|
distance_function = _get_postgresql_distance_function(distance_type)
|
||||||
|
embedding_parameter = f"${_POSTGRESQL_PARAMETER_QUERY_EMBEDDING}"
|
||||||
|
else:
|
||||||
|
distance_function = _get_googlesql_distance_function(
|
||||||
|
distance_type, ann=False
|
||||||
|
)
|
||||||
|
embedding_parameter = f"@{_GOOGLESQL_PARAMETER_QUERY_EMBEDDING}"
|
||||||
|
columns += [f"""{distance_function}(
|
||||||
|
{embedding_column_to_search},
|
||||||
|
{embedding_parameter}) AS {_DISTANCE_ALIAS}
|
||||||
|
"""]
|
||||||
|
columns = ", ".join(columns)
|
||||||
|
if additional_filter is None:
|
||||||
|
additional_filter = "1=1"
|
||||||
|
|
||||||
|
optional_limit_clause = ""
|
||||||
|
if top_k > 0:
|
||||||
|
optional_limit_clause = f"""LIMIT {top_k}"""
|
||||||
|
return f"""
|
||||||
|
SELECT {columns}
|
||||||
|
FROM {table_name}
|
||||||
|
WHERE {additional_filter}
|
||||||
|
ORDER BY {_DISTANCE_ALIAS}
|
||||||
|
{optional_limit_clause}
|
||||||
|
"""
|
||||||
|
|
||||||
|
|
||||||
|
def _generate_sql_for_ann(
|
||||||
|
dialect: DatabaseDialect,
|
||||||
|
table_name: str,
|
||||||
|
embedding_column_to_search: str,
|
||||||
|
columns,
|
||||||
|
additional_filter: Optional[str],
|
||||||
|
distance_type: str,
|
||||||
|
top_k: int,
|
||||||
|
num_leaves_to_search: int,
|
||||||
|
):
|
||||||
|
"""Generates a SQL query for ANN search."""
|
||||||
|
if dialect == DatabaseDialect.POSTGRESQL:
|
||||||
|
raise NotImplementedError(
|
||||||
|
f"{_APPROXIMATE_NEAREST_NEIGHBORS} is not supported for PostgreSQL"
|
||||||
|
" dialect."
|
||||||
|
)
|
||||||
|
distance_function = _get_googlesql_distance_function(distance_type, ann=True)
|
||||||
|
columns = columns + [f"""{distance_function}(
|
||||||
|
{embedding_column_to_search},
|
||||||
|
@{_GOOGLESQL_PARAMETER_QUERY_EMBEDDING},
|
||||||
|
options => JSON '{{"num_leaves_to_search": {num_leaves_to_search}}}'
|
||||||
|
) AS {_DISTANCE_ALIAS}
|
||||||
|
"""]
|
||||||
|
columns = ", ".join(columns)
|
||||||
|
query_filter = f"{embedding_column_to_search} IS NOT NULL"
|
||||||
|
if additional_filter is not None:
|
||||||
|
query_filter = f"{query_filter} AND {additional_filter}"
|
||||||
|
|
||||||
|
return f"""
|
||||||
|
SELECT {columns}
|
||||||
|
FROM {table_name}
|
||||||
|
WHERE {query_filter}
|
||||||
|
ORDER BY {_DISTANCE_ALIAS}
|
||||||
|
LIMIT {top_k}
|
||||||
|
"""
|
||||||
|
|
||||||
|
|
||||||
|
def similarity_search(
|
||||||
|
project_id: str,
|
||||||
|
instance_id: str,
|
||||||
|
database_id: str,
|
||||||
|
table_name: str,
|
||||||
|
query: str,
|
||||||
|
embedding_column_to_search: str,
|
||||||
|
columns: List[str],
|
||||||
|
embedding_options: Dict[str, str],
|
||||||
|
credentials: Credentials,
|
||||||
|
settings: SpannerToolSettings,
|
||||||
|
tool_context: ToolContext,
|
||||||
|
additional_filter: Optional[str] = None,
|
||||||
|
search_options: Optional[Dict[str, Any]] = None,
|
||||||
|
) -> Dict[str, Any]:
|
||||||
|
# fmt: off
|
||||||
|
"""Similarity search in Spanner using a text query.
|
||||||
|
|
||||||
|
The function will use embedding service (provided from options) to embed
|
||||||
|
the text query automatically, then use the embedding vector to do similarity
|
||||||
|
search and to return requested data. This is suitable when the Spanner table
|
||||||
|
contains a column that stores the embeddings of the data that we want to
|
||||||
|
search the `query` against.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
project_id (str): The GCP project id in which the spanner database
|
||||||
|
resides.
|
||||||
|
instance_id (str): The instance id of the spanner database.
|
||||||
|
database_id (str): The database id of the spanner database.
|
||||||
|
table_name (str): The name of the table used for vector search.
|
||||||
|
query (str): The user query for which the tool will find the top similar
|
||||||
|
content. The query will be embedded and used for vector search.
|
||||||
|
embedding_column_to_search (str): The name of the column that contains the
|
||||||
|
embeddings of the documents. The tool will do similarity search on this
|
||||||
|
column.
|
||||||
|
columns (List[str]): A list of column names, representing the additional
|
||||||
|
columns to return in the search results.
|
||||||
|
embedding_options (Dict[str, str]): A dictionary of options to use for
|
||||||
|
the embedding service. The following options are supported:
|
||||||
|
- spanner_embedding_model_name: (For GoogleSQL dialect) The
|
||||||
|
name of the embedding model that is registered in Spanner via a
|
||||||
|
`CREATE MODEL` statement. For more details, see
|
||||||
|
https://cloud.google.com/spanner/docs/ml-tutorial-embeddings#generate_and_store_text_embeddings
|
||||||
|
- vertex_ai_embedding_model_endpoint: (For PostgreSQL dialect)
|
||||||
|
The fully qualified endpoint of the Vertex AI embedding model,
|
||||||
|
in the format of
|
||||||
|
`projects/$project/locations/$location/publishers/google/models/$model_name`,
|
||||||
|
where $project is the project hosting the Vertex AI endpoint,
|
||||||
|
$location is the location of the endpoint, and $model_name is
|
||||||
|
the name of the text embedding model.
|
||||||
|
credentials (Credentials): The credentials to use for the request.
|
||||||
|
settings (SpannerToolSettings): The configuration for the tool.
|
||||||
|
tool_context (ToolContext): The context for the tool.
|
||||||
|
additional_filter (Optional[str]): An optional filter to apply to the
|
||||||
|
search query. If provided, this will be added to the WHERE clause of the
|
||||||
|
final query.
|
||||||
|
search_options (Optional[Dict[str, Any]]): A dictionary of options to use
|
||||||
|
for the similarity search. The following options are supported:
|
||||||
|
- top_k: The number of most similar documents to return. The
|
||||||
|
default value is 4.
|
||||||
|
- distance_type: The distance type to use to perform the
|
||||||
|
similarity search. Valid values include "COSINE_DISTANCE",
|
||||||
|
"EUCLIDEAN_DISTANCE", and "DOT_PRODUCT". Default value is
|
||||||
|
"COSINE_DISTANCE".
|
||||||
|
- nearest_neighbors_algorithm: The nearest neighbors search
|
||||||
|
algorithm to use. Valid values include "EXACT_NEAREST_NEIGHBORS"
|
||||||
|
and "APPROXIMATE_NEAREST_NEIGHBORS". Default value is
|
||||||
|
"EXACT_NEAREST_NEIGHBORS".
|
||||||
|
- num_leaves_to_search: (Only applies when the
|
||||||
|
nearest_neighbors_algorithm is APPROXIMATE_NEAREST_NEIGHBORS.)
|
||||||
|
The number of leaves to search in the vector index.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dict[str, Any]: A dictionary representing the result of the search.
|
||||||
|
On success, it contains {"status": "SUCCESS", "rows": [...]}. The last
|
||||||
|
column of each row is the distance between the query and the column
|
||||||
|
embedding (i.e. the embedding_column_to_search).
|
||||||
|
On error, it contains {"status": "ERROR", "error_details": "..."}.
|
||||||
|
|
||||||
|
Examples:
|
||||||
|
Search for relevant products given a user's text description and a filter
|
||||||
|
on the price:
|
||||||
|
>>> similarity_search(
|
||||||
|
... project_id="my-project",
|
||||||
|
... instance_id="my-instance",
|
||||||
|
... database_id="my-database",
|
||||||
|
... table_name="my-product-table",
|
||||||
|
... query="Tools that can help me clean my house.",
|
||||||
|
... embedding_column_to_search="product_description_embedding",
|
||||||
|
... columns=["product_name", "product_description", "price_in_cents"],
|
||||||
|
... credentials=credentials,
|
||||||
|
... settings=settings,
|
||||||
|
... tool_context=tool_context,
|
||||||
|
... additional_filter="price_in_cents < 100000",
|
||||||
|
... embedding_options={
|
||||||
|
... "spanner_embedding_model_name": "my_embedding_model"
|
||||||
|
... },
|
||||||
|
... search_options={
|
||||||
|
... "top_k": 2,
|
||||||
|
... "distance_type": "COSINE_DISTANCE"
|
||||||
|
... }
|
||||||
|
... )
|
||||||
|
{
|
||||||
|
"status": "SUCCESS",
|
||||||
|
"rows": [
|
||||||
|
(
|
||||||
|
"Powerful Robot Vacuum",
|
||||||
|
"This is a powerful robot vacuum that can clean carpets and wood floors.",
|
||||||
|
99999,
|
||||||
|
0.31,
|
||||||
|
),
|
||||||
|
(
|
||||||
|
"Nice Mop",
|
||||||
|
"Great for cleaning different surfaces.",
|
||||||
|
5099,
|
||||||
|
0.45,
|
||||||
|
),
|
||||||
|
],
|
||||||
|
}
|
||||||
|
"""
|
||||||
|
# fmt: on
|
||||||
|
try:
|
||||||
|
# Get Spanner client
|
||||||
|
spanner_client = client.get_spanner_client(
|
||||||
|
project=project_id, credentials=credentials
|
||||||
|
)
|
||||||
|
instance = spanner_client.instance(instance_id)
|
||||||
|
database = instance.database(database_id)
|
||||||
|
|
||||||
|
assert database.database_dialect in [
|
||||||
|
DatabaseDialect.GOOGLE_STANDARD_SQL,
|
||||||
|
DatabaseDialect.POSTGRESQL,
|
||||||
|
], (
|
||||||
|
"Unsupported database dialect: %s" % database.database_dialect
|
||||||
|
)
|
||||||
|
|
||||||
|
if embedding_options is None:
|
||||||
|
embedding_options = {}
|
||||||
|
if search_options is None:
|
||||||
|
search_options = {}
|
||||||
|
spanner_embedding_model_name = embedding_options.get(
|
||||||
|
_SPANNER_EMBEDDING_MODEL_NAME
|
||||||
|
)
|
||||||
|
if (
|
||||||
|
database.database_dialect == DatabaseDialect.GOOGLE_STANDARD_SQL
|
||||||
|
and spanner_embedding_model_name is None
|
||||||
|
):
|
||||||
|
raise ValueError(
|
||||||
|
f"embedding_options['{_SPANNER_EMBEDDING_MODEL_NAME}']"
|
||||||
|
" must be specified for GoogleSQL dialect."
|
||||||
|
)
|
||||||
|
vertex_ai_embedding_model_endpoint = embedding_options.get(
|
||||||
|
_VERTEX_AI_EMBEDDING_MODEL_ENDPOINT
|
||||||
|
)
|
||||||
|
if (
|
||||||
|
database.database_dialect == DatabaseDialect.POSTGRESQL
|
||||||
|
and vertex_ai_embedding_model_endpoint is None
|
||||||
|
):
|
||||||
|
raise ValueError(
|
||||||
|
f"embedding_options['{_VERTEX_AI_EMBEDDING_MODEL_ENDPOINT}']"
|
||||||
|
" must be specified for PostgreSQL dialect."
|
||||||
|
)
|
||||||
|
|
||||||
|
# Use cosine distance by default.
|
||||||
|
distance_type = search_options.get(_DISTANCE_TYPE)
|
||||||
|
if distance_type is None:
|
||||||
|
distance_type = "COSINE_DISTANCE"
|
||||||
|
|
||||||
|
top_k = search_options.get(_TOP_K)
|
||||||
|
if top_k is None:
|
||||||
|
top_k = 4
|
||||||
|
|
||||||
|
# Use EXACT_NEAREST_NEIGHBORS (i.e. kNN) by default.
|
||||||
|
nearest_neighbors_algorithm = search_options.get(
|
||||||
|
_NEAREST_NEIGHBORS_ALGORITHM, _EXACT_NEAREST_NEIGHBORS
|
||||||
|
)
|
||||||
|
if nearest_neighbors_algorithm not in (
|
||||||
|
_EXACT_NEAREST_NEIGHBORS,
|
||||||
|
_APPROXIMATE_NEAREST_NEIGHBORS,
|
||||||
|
):
|
||||||
|
raise NotImplementedError(
|
||||||
|
f"Unsupported search_options['{_NEAREST_NEIGHBORS_ALGORITHM}']:"
|
||||||
|
f" {nearest_neighbors_algorithm}"
|
||||||
|
)
|
||||||
|
|
||||||
|
# copy_bara:strip_begin(internal comment)
|
||||||
|
# TODO: Once Spanner supports ML.PREDICT in WITH CTE, we can execute
|
||||||
|
# the embedding and the search in one query instead of two separate queries.
|
||||||
|
# copy_bara:strip_end
|
||||||
|
embedding = _get_embedding_for_query(
|
||||||
|
database,
|
||||||
|
database.database_dialect,
|
||||||
|
spanner_embedding_model_name,
|
||||||
|
vertex_ai_embedding_model_endpoint,
|
||||||
|
query,
|
||||||
|
)
|
||||||
|
|
||||||
|
if nearest_neighbors_algorithm == _EXACT_NEAREST_NEIGHBORS:
|
||||||
|
sql = _generate_sql_for_knn(
|
||||||
|
database.database_dialect,
|
||||||
|
table_name,
|
||||||
|
embedding_column_to_search,
|
||||||
|
columns,
|
||||||
|
additional_filter,
|
||||||
|
distance_type,
|
||||||
|
top_k,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
num_leaves_to_search = search_options.get(_NUM_LEAVES_TO_SEARCH)
|
||||||
|
if num_leaves_to_search is None:
|
||||||
|
num_leaves_to_search = 1000
|
||||||
|
sql = _generate_sql_for_ann(
|
||||||
|
database.database_dialect,
|
||||||
|
table_name,
|
||||||
|
embedding_column_to_search,
|
||||||
|
columns,
|
||||||
|
additional_filter,
|
||||||
|
distance_type,
|
||||||
|
top_k,
|
||||||
|
num_leaves_to_search,
|
||||||
|
)
|
||||||
|
|
||||||
|
if database.database_dialect == DatabaseDialect.POSTGRESQL:
|
||||||
|
params = {f"p{_POSTGRESQL_PARAMETER_QUERY_EMBEDDING}": embedding}
|
||||||
|
else:
|
||||||
|
params = {_GOOGLESQL_PARAMETER_QUERY_EMBEDDING: embedding}
|
||||||
|
|
||||||
|
with database.snapshot() as snapshot:
|
||||||
|
result_set = snapshot.execute_sql(sql, params=params)
|
||||||
|
rows = []
|
||||||
|
result = {}
|
||||||
|
for row in result_set:
|
||||||
|
try:
|
||||||
|
# if the json serialization of the row succeeds, use it as is
|
||||||
|
json.dumps(row)
|
||||||
|
except:
|
||||||
|
row = str(row)
|
||||||
|
|
||||||
|
rows.append(row)
|
||||||
|
|
||||||
|
result["status"] = "SUCCESS"
|
||||||
|
result["rows"] = rows
|
||||||
|
return result
|
||||||
|
except Exception as ex:
|
||||||
|
return {
|
||||||
|
"status": "ERROR",
|
||||||
|
"error_details": str(ex),
|
||||||
|
}
|
||||||
@@ -19,10 +19,11 @@ from typing import Optional
|
|||||||
from typing import Union
|
from typing import Union
|
||||||
|
|
||||||
from google.adk.agents.readonly_context import ReadonlyContext
|
from google.adk.agents.readonly_context import ReadonlyContext
|
||||||
|
from google.adk.tools.spanner import metadata_tool
|
||||||
|
from google.adk.tools.spanner import query_tool
|
||||||
|
from google.adk.tools.spanner import search_tool
|
||||||
from typing_extensions import override
|
from typing_extensions import override
|
||||||
|
|
||||||
from . import metadata_tool
|
|
||||||
from . import query_tool
|
|
||||||
from ...tools.base_tool import BaseTool
|
from ...tools.base_tool import BaseTool
|
||||||
from ...tools.base_toolset import BaseToolset
|
from ...tools.base_toolset import BaseToolset
|
||||||
from ...tools.base_toolset import ToolPredicate
|
from ...tools.base_toolset import ToolPredicate
|
||||||
@@ -113,6 +114,13 @@ class SpannerToolset(BaseToolset):
|
|||||||
tool_settings=self._tool_settings,
|
tool_settings=self._tool_settings,
|
||||||
)
|
)
|
||||||
)
|
)
|
||||||
|
all_tools.append(
|
||||||
|
GoogleTool(
|
||||||
|
func=search_tool.similarity_search,
|
||||||
|
credentials_config=self._credentials_config,
|
||||||
|
tool_settings=self._tool_settings,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
|
||||||
return [
|
return [
|
||||||
tool
|
tool
|
||||||
|
|||||||
@@ -0,0 +1,301 @@
|
|||||||
|
# Copyright 2025 Google LLC
|
||||||
|
#
|
||||||
|
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||||
|
# you may not use this file except in compliance with the License.
|
||||||
|
# You may obtain a copy of the License at
|
||||||
|
#
|
||||||
|
# http://www.apache.org/licenses/LICENSE-2.0
|
||||||
|
#
|
||||||
|
# Unless required by applicable law or agreed to in writing, software
|
||||||
|
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||||
|
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||||
|
# See the License for the specific language governing permissions and
|
||||||
|
# limitations under the License.
|
||||||
|
|
||||||
|
from unittest.mock import MagicMock
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
from google.adk.tools.spanner import search_tool
|
||||||
|
from google.cloud.spanner_admin_database_v1.types import DatabaseDialect
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def mock_credentials():
|
||||||
|
return MagicMock()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def mock_spanner_ids():
|
||||||
|
return {
|
||||||
|
"project_id": "test-project",
|
||||||
|
"instance_id": "test-instance",
|
||||||
|
"database_id": "test-database",
|
||||||
|
"table_name": "test-table",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@patch("google.adk.tools.spanner.client.get_spanner_client")
|
||||||
|
def test_similarity_search_knn_success(
|
||||||
|
mock_get_spanner_client, mock_spanner_ids, mock_credentials
|
||||||
|
):
|
||||||
|
"""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
|
||||||
|
|
||||||
|
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_embedding_model_name": "test_model"},
|
||||||
|
credentials=mock_credentials,
|
||||||
|
settings=MagicMock(),
|
||||||
|
tool_context=MagicMock(),
|
||||||
|
)
|
||||||
|
assert result["status"] == "SUCCESS", result
|
||||||
|
assert result["rows"] == [("result1",), ("result2",)]
|
||||||
|
|
||||||
|
# Check the generated SQL for kNN search
|
||||||
|
call_args = mock_snapshot.execute_sql.call_args
|
||||||
|
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]}}
|
||||||
|
|
||||||
|
|
||||||
|
@patch("google.adk.tools.spanner.client.get_spanner_client")
|
||||||
|
def test_similarity_search_ann_success(
|
||||||
|
mock_get_spanner_client, mock_spanner_ids, mock_credentials
|
||||||
|
):
|
||||||
|
"""Test similarity_search function with ANN 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 ANN search
|
||||||
|
mock_snapshot.execute_sql.side_effect = [
|
||||||
|
mock_embedding_result,
|
||||||
|
iter([("ann_result1",), ("ann_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
|
||||||
|
|
||||||
|
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_embedding_model_name": "test_model"},
|
||||||
|
credentials=mock_credentials,
|
||||||
|
settings=MagicMock(),
|
||||||
|
tool_context=MagicMock(),
|
||||||
|
search_options={
|
||||||
|
"nearest_neighbors_algorithm": "APPROXIMATE_NEAREST_NEIGHBORS"
|
||||||
|
},
|
||||||
|
)
|
||||||
|
assert result["status"] == "SUCCESS", result
|
||||||
|
assert result["rows"] == [("ann_result1",), ("ann_result2",)]
|
||||||
|
call_args = mock_snapshot.execute_sql.call_args
|
||||||
|
sql = call_args.args[0]
|
||||||
|
assert "APPROX_COSINE_DISTANCE" in sql
|
||||||
|
assert "@embedding" in sql
|
||||||
|
assert call_args.kwargs == {"params": {"embedding": [0.1, 0.2, 0.3]}}
|
||||||
|
|
||||||
|
|
||||||
|
@patch("google.adk.tools.spanner.client.get_spanner_client")
|
||||||
|
def test_similarity_search_error(
|
||||||
|
mock_get_spanner_client, mock_spanner_ids, mock_credentials
|
||||||
|
):
|
||||||
|
"""Test similarity_search function with a generic error."""
|
||||||
|
mock_get_spanner_client.side_effect = Exception("Test Exception")
|
||||||
|
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",
|
||||||
|
embedding_options={"spanner_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"
|
||||||
|
|
||||||
|
|
||||||
|
@patch("google.adk.tools.spanner.client.get_spanner_client")
|
||||||
|
def test_similarity_search_postgresql_knn_success(
|
||||||
|
mock_get_spanner_client, mock_spanner_ids, mock_credentials
|
||||||
|
):
|
||||||
|
"""Test similarity_search with PostgreSQL dialect for kNN."""
|
||||||
|
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],)
|
||||||
|
mock_snapshot.execute_sql.side_effect = [
|
||||||
|
mock_embedding_result,
|
||||||
|
iter([("pg_result",)]),
|
||||||
|
]
|
||||||
|
mock_database.snapshot.return_value.__enter__.return_value = mock_snapshot
|
||||||
|
mock_database.database_dialect = DatabaseDialect.POSTGRESQL
|
||||||
|
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_endpoint": "test_endpoint"},
|
||||||
|
credentials=mock_credentials,
|
||||||
|
settings=MagicMock(),
|
||||||
|
tool_context=MagicMock(),
|
||||||
|
)
|
||||||
|
assert result["status"] == "SUCCESS", result
|
||||||
|
assert result["rows"] == [("pg_result",)]
|
||||||
|
call_args = mock_snapshot.execute_sql.call_args
|
||||||
|
sql = call_args.args[0]
|
||||||
|
assert "spanner.cosine_distance" in sql
|
||||||
|
assert "$1" in sql
|
||||||
|
assert call_args.kwargs == {"params": {"p1": [0.1, 0.2, 0.3]}}
|
||||||
|
|
||||||
|
|
||||||
|
@patch("google.adk.tools.spanner.client.get_spanner_client")
|
||||||
|
def test_similarity_search_postgresql_ann_unsupported(
|
||||||
|
mock_get_spanner_client, mock_spanner_ids, mock_credentials
|
||||||
|
):
|
||||||
|
"""Test similarity_search with unsupported ANN for PostgreSQL dialect."""
|
||||||
|
mock_spanner_client = MagicMock()
|
||||||
|
mock_instance = MagicMock()
|
||||||
|
mock_database = MagicMock()
|
||||||
|
mock_database.database_dialect = DatabaseDialect.POSTGRESQL
|
||||||
|
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_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."
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@patch("google.adk.tools.spanner.client.get_spanner_client")
|
||||||
|
def test_similarity_search_missing_spanner_embedding_model_name_error(
|
||||||
|
mock_get_spanner_client, mock_spanner_ids, mock_credentials
|
||||||
|
):
|
||||||
|
"""Test similarity_search with missing spanner_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={},
|
||||||
|
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."
|
||||||
|
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_get_spanner_client, mock_spanner_ids, mock_credentials
|
||||||
|
):
|
||||||
|
"""Test similarity_search with missing vertex_ai_embedding_model_endpoint."""
|
||||||
|
mock_spanner_client = MagicMock()
|
||||||
|
mock_instance = MagicMock()
|
||||||
|
mock_database = MagicMock()
|
||||||
|
mock_database.database_dialect = DatabaseDialect.POSTGRESQL
|
||||||
|
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={},
|
||||||
|
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."
|
||||||
|
in result["error_details"]
|
||||||
|
)
|
||||||
@@ -37,7 +37,7 @@ async def test_spanner_toolset_tools_default():
|
|||||||
tools = await toolset.get_tools()
|
tools = await toolset.get_tools()
|
||||||
assert tools is not None
|
assert tools is not None
|
||||||
|
|
||||||
assert len(tools) == 6
|
assert len(tools) == 7
|
||||||
assert all([isinstance(tool, GoogleTool) for tool in tools])
|
assert all([isinstance(tool, GoogleTool) for tool in tools])
|
||||||
|
|
||||||
expected_tool_names = set([
|
expected_tool_names = set([
|
||||||
@@ -47,6 +47,7 @@ async def test_spanner_toolset_tools_default():
|
|||||||
"list_named_schemas",
|
"list_named_schemas",
|
||||||
"get_table_schema",
|
"get_table_schema",
|
||||||
"execute_sql",
|
"execute_sql",
|
||||||
|
"similarity_search",
|
||||||
])
|
])
|
||||||
actual_tool_names = set([tool.name for tool in tools])
|
actual_tool_names = set([tool.name for tool in tools])
|
||||||
assert actual_tool_names == expected_tool_names
|
assert actual_tool_names == expected_tool_names
|
||||||
|
|||||||
Reference in New Issue
Block a user