Skip to content
Open
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
2 changes: 2 additions & 0 deletions docs/index.rst
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,8 @@ API Reference
langchain_google_alloydb_pg/vectorstore
langchain_google_alloydb_pg/loader
langchain_google_alloydb_pg/history
langchain_google_alloydb_pg/toolkit
langchain_google_alloydb_pg/embeddings

How to Choose a Nearest-Neighbor Index Guide
--------------------------------------------
Expand Down
7 changes: 7 additions & 0 deletions docs/langchain_google_alloydb_pg/embeddings.rst
Original file line number Diff line number Diff line change
@@ -0,0 +1,7 @@
Embeddings
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~

.. automodule:: langchain_google_alloydb_pg.embeddings
:members:
:private-members:
:noindex:
7 changes: 7 additions & 0 deletions docs/langchain_google_alloydb_pg/toolkit.rst
Original file line number Diff line number Diff line change
@@ -0,0 +1,7 @@
Toolkit
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~

.. automodule:: langchain_google_alloydb_pg.toolkit
:members:
:private-members:
:noindex:
3 changes: 3 additions & 0 deletions src/langchain_google_alloydb_pg/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,7 @@
from .loader import AlloyDBDocumentSaver, AlloyDBLoader
from .model_manager import AlloyDBModel, AlloyDBModelManager
from .vectorstore import AlloyDBVectorStore
from .toolkit import AlloyDBNL2SQLTool, AlloyDBToolkit
from .version import __version__

__all__ = [
Expand All @@ -43,4 +44,6 @@
"reciprocal_rank_fusion",
"weighted_sum_ranking",
"__version__",
"AlloyDBNL2SQLTool",
"AlloyDBToolkit",
]
63 changes: 62 additions & 1 deletion src/langchain_google_alloydb_pg/embeddings.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
# Copyright 2024 Google LLC
# Copyright 2026 Google LLC
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
Expand Down Expand Up @@ -169,3 +169,64 @@ async def __aembed_query(self, query: str) -> list[float]:
result_map = result.mappings()
results = result_map.fetchall()
return json.loads(results[0]["embedding"])

async def aembed_image(self, image_uri: str) -> list[float]:
"""Asynchronous Embed image.

Args:
image_uri (str): Image URI to embed.

Returns:
list[float]: Embedding.
"""
embeddings = await self._engine._run_as_async(self.__aembed_image(image_uri))
return embeddings

def embed_image(self, image_uri: str) -> list[float]:
"""Embed image.

Args:
image_uri (str): Image URI to embed.

Returns:
list[float]: Embedding.
"""
return self._engine._run_as_sync(self.__aembed_image(image_uri))

async def __aembed_image(self, image_uri: str) -> list[float]:
"""Coroutine for generating embeddings for a given image.

Args:
image_uri (str): Image URI to embed.

Returns:
list[float]: Embedding.
"""
query = f" SELECT google_ml.image_embedding('{self.model_id}', '{image_uri}')::vector as embedding "
async with self._engine._pool.connect() as conn:
result = await conn.execute(text(query))
result_map = result.mappings()
results = result_map.fetchall()
return json.loads(results[0]["embedding"])

def embed_images(self, image_uris: list[str]) -> list[list[float]]:
"""Embed list of images.

Args:
image_uris (list[str]): List of Image URIs to embed.

Returns:
list[list[float]]: Embeddings.
"""
return [self.embed_image(uri) for uri in image_uris]

async def aembed_images(self, image_uris: list[str]) -> list[list[float]]:
"""Asynchronous Embed list of images.

Args:
image_uris (list[str]): List of Image URIs to embed.

Returns:
list[list[float]]: Embeddings.
"""
return [await self.aembed_image(uri) for uri in image_uris]
75 changes: 75 additions & 0 deletions src/langchain_google_alloydb_pg/toolkit.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,75 @@
# Copyright 2026 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 typing import List, Optional, Type

from langchain_core.callbacks import CallbackManagerForToolRun
from langchain_core.tools import BaseTool
from langchain_core.tools.base import BaseToolkit
from pydantic import BaseModel, ConfigDict, Field
from sqlalchemy import text

from .engine import AlloyDBEngine


class NL2SQLInput(BaseModel):
query: str = Field(description="The natural language query to translate to SQL.")


class AlloyDBNL2SQLTool(BaseTool):
"""Tool for translating natural language to SQL using AlloyDB's native NL2SQL."""

name: str = "alloydb_nl2sql_tool"
description: str = (
"Translate a natural language question into a SQL query using AlloyDB AI. "
"Pass the user's natural language question as the query."
)
args_schema: Type[BaseModel] = NL2SQLInput
engine: AlloyDBEngine

def _run(
self,
query: str,
run_manager: Optional[CallbackManagerForToolRun] = None,
) -> str:
"""Run the tool synchronously."""
return self.engine._run_as_sync(self._arun(query, run_manager))

async def _arun(
self,
query: str,
run_manager: Optional[CallbackManagerForToolRun] = None,
) -> str:
"""Run the tool asynchronously."""
# Using ai.generate_sql based on general AlloyDB AI syntax patterns
# or similar functions to generate SQL natively.
sql_query = "SELECT google_ml.generate_sql(:query)"
async with self.engine._pool.connect() as conn:
result = await conn.execute(text(sql_query), {"query": query})
return str(result.scalar())


class AlloyDBToolkit(BaseToolkit):
"""Toolkit for interacting with AlloyDB databases using native AI features."""

engine: AlloyDBEngine = Field(exclude=True)

model_config = ConfigDict(
arbitrary_types_allowed=True,
)

def get_tools(self) -> List[BaseTool]:
"""Get the tools in the toolkit."""
nl2sql_tool = AlloyDBNL2SQLTool(engine=self.engine)
return [nl2sql_tool]
15 changes: 14 additions & 1 deletion tests/test_embeddings.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
# Copyright 2024 Google LLC
# Copyright 2026 Google LLC
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
Expand Down Expand Up @@ -122,3 +122,16 @@ async def test_aembed_query(self, embeddings):
for embedding_field in embedding:
assert isinstance(embedding_field, float)
assert -1 <= embedding_field <= 1

async def test_embed_image(self, embeddings):
"""Test image embedding generation via synchronous wrapper."""
# We assume the image returns an embedding list similar to embed_query
# when running against a live integration.
embedding = embeddings.embed_image("gs://bucket/test_image.jpg")
assert isinstance(embedding, list)

async def test_aembed_image(self, embeddings):
"""Test image embedding generation asynchronously."""
embedding = await embeddings.aembed_image("gs://bucket/test_image.jpg")
assert isinstance(embedding, list)

62 changes: 62 additions & 0 deletions tests/test_toolkit.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,62 @@
# Copyright 2026 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.

import pytest
from unittest.mock import AsyncMock, MagicMock
from langchain_google_alloydb_pg.toolkit import AlloyDBToolkit, AlloyDBNL2SQLTool
from langchain_google_alloydb_pg.engine import AlloyDBEngine

@pytest.fixture
def mock_engine():
class DummyEngine(AlloyDBEngine):
def __init__(self):
pass

engine = DummyEngine()
engine._run_as_sync = MagicMock(side_effect=lambda coro: "mocked_sql_query")

# Mock pool and connection for async runs
pool_mock = MagicMock()
conn_mock = AsyncMock()
conn_mock.execute.return_value = MagicMock(scalar=MagicMock(return_value="mocked_sql_query"))
pool_mock.connect.return_value.__aenter__.return_value = conn_mock
engine._pool = pool_mock

return engine

class TestAlloyDBToolkit:

@pytest.mark.asyncio
async def test_nl2sql_tool_arun(self, mock_engine):
"""Test that AlloyDBNL2SQLTool._arun executes the NL2SQL generation asynchronously."""
tool = AlloyDBNL2SQLTool(engine=mock_engine)
result = await tool._arun("Show me all users")
assert result == "mocked_sql_query"
# Verify the exact SQL query
conn_mock = mock_engine._pool.connect.return_value.__aenter__.return_value
executed_query = conn_mock.execute.call_args[0][0].text
assert "SELECT google_ml.generate_sql" in executed_query

def test_nl2sql_tool_run(self, mock_engine):
"""Test that AlloyDBNL2SQLTool._run executes the NL2SQL generation synchronously."""
tool = AlloyDBNL2SQLTool(engine=mock_engine)
result = tool._run("Show me all users")
assert result == "mocked_sql_query"

def test_toolkit_get_tools(self, mock_engine):
"""Test that AlloyDBToolkit properly exposes its internal list of AI tools."""
toolkit = AlloyDBToolkit(engine=mock_engine)
tools = toolkit.get_tools()
assert len(tools) == 1
assert isinstance(tools[0], AlloyDBNL2SQLTool)
Loading