mirror of
https://github.com/crewAIInc/crewAI.git
synced 2026-08-01 20:19:21 +00:00
Compare commits
1 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c8f441cffa |
@@ -271,7 +271,8 @@
|
||||
"edge/en/tools/database-data/qdrantvectorsearchtool",
|
||||
"edge/en/tools/database-data/weaviatevectorsearchtool",
|
||||
"edge/en/tools/database-data/mongodbvectorsearchtool",
|
||||
"edge/en/tools/database-data/singlestoresearchtool"
|
||||
"edge/en/tools/database-data/singlestoresearchtool",
|
||||
"edge/en/tools/database-data/db2searchtool"
|
||||
]
|
||||
},
|
||||
{
|
||||
|
||||
211
docs/edge/en/tools/database-data/db2searchtool.mdx
Normal file
211
docs/edge/en/tools/database-data/db2searchtool.mdx
Normal file
@@ -0,0 +1,211 @@
|
||||
---
|
||||
title: Db2 Vector Search Tool
|
||||
description: Semantic vector search for CrewAI agents using IBM Db2 native VECTOR_DISTANCE capabilities.
|
||||
icon: database
|
||||
mode: "wide"
|
||||
---
|
||||
|
||||
# `DB2VectorSearchTool`
|
||||
|
||||
## Description
|
||||
|
||||
Perform semantic vector similarity searches against IBM Db2 tables using the native `VECTOR_DISTANCE` function.
|
||||
Supports configurable distance metrics, OpenAI or custom embeddings, metadata filtering, and result shaping.
|
||||
|
||||
## Installation
|
||||
|
||||
```bash
|
||||
pip install ibm_db openai
|
||||
```
|
||||
|
||||
Or with uv:
|
||||
|
||||
```bash
|
||||
uv add ibm_db openai
|
||||
```
|
||||
|
||||
## Environment Variables
|
||||
|
||||
```bash
|
||||
OPENAI_API_KEY=your_openai_key # Required when using default OpenAI embeddings
|
||||
DB2_CONNECTION_STRING=DATABASE=TESTDB;HOSTNAME=localhost;PORT=50000;PROTOCOL=TCPIP;UID=db2user;PWD=password;
|
||||
```
|
||||
|
||||
## Basic Usage
|
||||
|
||||
```python
|
||||
from crewai import Agent
|
||||
from crewai_tools import DB2VectorSearchTool
|
||||
|
||||
tool = DB2VectorSearchTool(
|
||||
connection_string="DATABASE=TESTDB;HOSTNAME=localhost;PORT=50000;PROTOCOL=TCPIP;UID=db2user;PWD=password;",
|
||||
table_name="documents",
|
||||
vector_column="embedding",
|
||||
)
|
||||
|
||||
agent = Agent(
|
||||
role="Research Assistant",
|
||||
goal="Find relevant information in documents",
|
||||
tools=[tool],
|
||||
)
|
||||
```
|
||||
|
||||
## Full Semantic Search Workflow
|
||||
|
||||
```python
|
||||
import os
|
||||
from dotenv import load_dotenv
|
||||
from crewai import Agent, Task, Crew, Process
|
||||
from crewai_tools import DB2VectorSearchTool
|
||||
|
||||
load_dotenv()
|
||||
|
||||
db2_tool = DB2VectorSearchTool(
|
||||
connection_string=os.getenv("DB2_CONNECTION_STRING"),
|
||||
table_name="documents",
|
||||
vector_column="embedding",
|
||||
return_columns=["content", "category"],
|
||||
limit=3,
|
||||
distance_metric="COSINE",
|
||||
max_distance=0.35,
|
||||
)
|
||||
|
||||
search_agent = Agent(
|
||||
role="Senior Semantic Search Agent",
|
||||
goal="Find and analyse documents based on semantic search",
|
||||
backstory="You are an expert research assistant who can find relevant information using semantic search in a Db2 database.",
|
||||
tools=[db2_tool],
|
||||
verbose=True,
|
||||
)
|
||||
|
||||
answer_agent = Agent(
|
||||
role="Senior Answer Assistant",
|
||||
goal="Generate answers based on retrieved context",
|
||||
backstory="You are an expert assistant who generates answers from provided context.",
|
||||
tools=[db2_tool],
|
||||
verbose=True,
|
||||
)
|
||||
|
||||
search_task = Task(
|
||||
description="""Search for relevant documents about {query}.
|
||||
Include the relevant information found, vector distances, and returned fields.""",
|
||||
agent=search_agent,
|
||||
)
|
||||
|
||||
answer_task = Task(
|
||||
description="Given the retrieved Db2 context, generate a final answer.",
|
||||
agent=answer_agent,
|
||||
)
|
||||
|
||||
crew = Crew(
|
||||
agents=[search_agent, answer_agent],
|
||||
tasks=[search_task, answer_task],
|
||||
process=Process.sequential,
|
||||
verbose=True,
|
||||
)
|
||||
|
||||
result = crew.kickoff(inputs={"query": "What is the role of X in the document?"})
|
||||
print(result)
|
||||
```
|
||||
|
||||
## Tool Parameters
|
||||
|
||||
| Parameter | Type | Default | Description |
|
||||
|---|---|---|---|
|
||||
| `connection_string` | `str` | required | Db2 connection string. Format: `DATABASE=x;HOSTNAME=x;PORT=50000;PROTOCOL=TCPIP;UID=x;PWD=x;` |
|
||||
| `table_name` | `str` | `"documents"` | Table to search. Supports `schema.table` notation. |
|
||||
| `vector_column` | `str` | `"embedding"` | Column storing the vector embeddings. |
|
||||
| `embedding_model` | `str` | `"text-embedding-3-large"` | OpenAI model used when no custom embedding function is provided. |
|
||||
| `return_columns` | `list[str]` | `["content"]` | Columns to include in each result. Must contain at least one entry. |
|
||||
| `limit` | `int` | `3` | Maximum number of results (1–100). |
|
||||
| `distance_metric` | `str` | `"COSINE"` | Db2 distance metric. See supported values below. |
|
||||
| `max_distance` | `float \| None` | `None` | Drop results whose distance exceeds this value. |
|
||||
| `custom_embedding_fn` | `Callable[[str], list[float]] \| None` | `None` | Custom embedding function. Overrides OpenAI when provided. |
|
||||
|
||||
## Supported Distance Metrics
|
||||
|
||||
The following values map directly to the Db2 `VECTOR_DISTANCE` function:
|
||||
|
||||
- `COSINE`
|
||||
- `EUCLIDEAN`
|
||||
- `EUCLIDEAN_SQUARED`
|
||||
- `DOT`
|
||||
- `HAMMING`
|
||||
- `MANHATTAN`
|
||||
|
||||
Reference: [IBM Db2 VECTOR_DISTANCE documentation](https://www.ibm.com/docs/en/db2/12.1.x?topic=functions-vector-distance)
|
||||
|
||||
## Schema Parameters (per query)
|
||||
|
||||
| Parameter | Type | Required | Description |
|
||||
|---|---|---|---|
|
||||
| `query` | `str` | ✅ | The search query. |
|
||||
| `filter_by` | `str \| None` | ❌ | Column name for metadata filtering. Must be paired with `filter_value`. |
|
||||
| `filter_value` | `Any \| None` | ❌ | Value to filter on. Must be paired with `filter_by`. |
|
||||
|
||||
## Return Format
|
||||
|
||||
```json
|
||||
{
|
||||
"success": true,
|
||||
"results": [
|
||||
{
|
||||
"distance": 0.1401,
|
||||
"data": {
|
||||
"content": "Document content here",
|
||||
"category": "research"
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
On error:
|
||||
|
||||
```json
|
||||
{
|
||||
"success": false,
|
||||
"error": "Description of what went wrong",
|
||||
"error_type": "ExceptionClassName"
|
||||
}
|
||||
```
|
||||
|
||||
## Metadata Filtering
|
||||
|
||||
```python
|
||||
result = db2_tool.run(
|
||||
query="machine learning",
|
||||
filter_by="category",
|
||||
filter_value="research",
|
||||
)
|
||||
```
|
||||
|
||||
`filter_by` and `filter_value` must always be provided together. Providing only one raises a validation error.
|
||||
|
||||
## Custom Embeddings
|
||||
|
||||
Use any embedding model by supplying a `custom_embedding_fn`:
|
||||
|
||||
```python
|
||||
from sentence_transformers import SentenceTransformer
|
||||
from crewai_tools import DB2VectorSearchTool
|
||||
|
||||
model = SentenceTransformer("sentence-transformers/all-MiniLM-L6-v2")
|
||||
|
||||
def custom_embeddings(text: str) -> list[float]:
|
||||
return model.encode(text).tolist()
|
||||
|
||||
tool = DB2VectorSearchTool(
|
||||
connection_string="DATABASE=TESTDB;HOSTNAME=localhost;PORT=50000;PROTOCOL=TCPIP;UID=db2user;PWD=password;",
|
||||
table_name="documents",
|
||||
custom_embedding_fn=custom_embeddings,
|
||||
)
|
||||
```
|
||||
|
||||
When `custom_embedding_fn` is provided, `OPENAI_API_KEY` is not required.
|
||||
|
||||
## Security Features
|
||||
|
||||
- SQL identifier validation (table, column names must match `^[A-Za-z][A-Za-z0-9_]*(\.[A-Za-z][A-Za-z0-9_]*)?$`)
|
||||
- Parameterised SQL queries — values never interpolated into SQL strings
|
||||
- Distance metric whitelist — only valid Db2 metric names accepted
|
||||
@@ -59,6 +59,10 @@ from crewai_tools.tools.dalle_tool.dalle_tool import DallETool
|
||||
from crewai_tools.tools.databricks_query_tool.databricks_query_tool import (
|
||||
DatabricksQueryTool,
|
||||
)
|
||||
from crewai_tools.tools.db2_search_tool import (
|
||||
DB2ToolSchema,
|
||||
DB2VectorSearchTool,
|
||||
)
|
||||
from crewai_tools.tools.daytona_sandbox_tool import (
|
||||
DaytonaExecTool,
|
||||
DaytonaFileTool,
|
||||
@@ -248,6 +252,8 @@ __all__ = [
|
||||
"CrewaiPlatformTools",
|
||||
"DOCXSearchTool",
|
||||
"DallETool",
|
||||
"DB2ToolSchema",
|
||||
"DB2VectorSearchTool",
|
||||
"DatabricksQueryTool",
|
||||
"DaytonaExecTool",
|
||||
"DaytonaFileTool",
|
||||
|
||||
@@ -53,6 +53,10 @@ from crewai_tools.tools.daytona_sandbox_tool import (
|
||||
DaytonaFileTool,
|
||||
DaytonaPythonTool,
|
||||
)
|
||||
from crewai_tools.tools.db2_search_tool import (
|
||||
DB2ToolSchema,
|
||||
DB2VectorSearchTool,
|
||||
)
|
||||
from crewai_tools.tools.directory_read_tool.directory_read_tool import (
|
||||
DirectoryReadTool,
|
||||
)
|
||||
@@ -234,6 +238,8 @@ __all__ = [
|
||||
"DOCXSearchTool",
|
||||
"DallETool",
|
||||
"DatabricksQueryTool",
|
||||
"DB2ToolSchema",
|
||||
"DB2VectorSearchTool",
|
||||
"DaytonaExecTool",
|
||||
"DaytonaFileTool",
|
||||
"DaytonaPythonTool",
|
||||
|
||||
@@ -0,0 +1,91 @@
|
||||
# DB2 Vector Search Tool
|
||||
|
||||
IBM DB2 Vector Search Tool for CrewAI.
|
||||
|
||||
Supports:
|
||||
|
||||
- IBM DB2 native VECTOR search
|
||||
- OpenAI embeddings
|
||||
- Custom embedding functions
|
||||
- Metadata filtering
|
||||
- Runtime dynamic imports
|
||||
- Standardized CrewAI tool architecture
|
||||
|
||||
---
|
||||
|
||||
# Installation
|
||||
|
||||
```bash
|
||||
uv add ibm_db openai
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
# Environment Variables
|
||||
|
||||
```env
|
||||
OPENAI_API_KEY=your_openai_key
|
||||
|
||||
DB2_CONNECTION_STRING=DATABASE=TESTDB;HOSTNAME=localhost;PORT=50000;PROTOCOL=TCPIP;UID=db2user;PWD=password;
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
# Example Usage
|
||||
|
||||
```python
|
||||
from crewai_tools import DB2VectorSearchTool
|
||||
|
||||
tool = DB2VectorSearchTool(
|
||||
connection_string="DATABASE=TESTDB;HOSTNAME=localhost;PORT=50000;PROTOCOL=TCPIP;UID=db2user;PWD=password;",
|
||||
table_name="documents",
|
||||
)
|
||||
|
||||
result = tool.run(
|
||||
query="What is machine learning?",
|
||||
)
|
||||
|
||||
print(result)
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
# Example With Metadata Filtering
|
||||
|
||||
```python
|
||||
result = tool.run(
|
||||
query="AI papers",
|
||||
filter_by="category",
|
||||
filter_value="AI",
|
||||
)
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
# Supported Features
|
||||
|
||||
- DB2 VECTOR datatype
|
||||
- VECTOR_DISTANCE search
|
||||
- COSINE similarity
|
||||
- Metadata filtering
|
||||
- Uses a custom embedding function if supplied, otherwise OpenAI embeddings
|
||||
|
||||
---
|
||||
|
||||
# Architecture
|
||||
|
||||
This tool follows the same architecture as:
|
||||
|
||||
- QdrantVectorSearchTool
|
||||
- WeaviateVectorSearchTool
|
||||
|
||||
Responsibilities:
|
||||
|
||||
- Generate query embeddings
|
||||
- Perform vector similarity search
|
||||
- Apply optional metadata filters
|
||||
- Return normalized JSON results
|
||||
|
||||
This tool is retrieval-only.
|
||||
|
||||
Document ingestion should be handled separately.
|
||||
@@ -0,0 +1,10 @@
|
||||
from crewai_tools.tools.db2_search_tool.db2_search_tool import (
|
||||
DB2ToolSchema,
|
||||
DB2VectorSearchTool,
|
||||
)
|
||||
|
||||
|
||||
__all__ = [
|
||||
"DB2ToolSchema",
|
||||
"DB2VectorSearchTool",
|
||||
]
|
||||
@@ -0,0 +1,365 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
import datetime
|
||||
import decimal
|
||||
import importlib
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
from typing import Any, ClassVar
|
||||
|
||||
from crewai.tools import BaseTool, EnvVar
|
||||
from pydantic import BaseModel, ConfigDict, Field, model_validator
|
||||
from pydantic.types import ImportString
|
||||
|
||||
|
||||
class DB2JSONEncoder(json.JSONEncoder):
|
||||
"""Safely handles Decimal, Timestamps, and Bytes from DB2."""
|
||||
|
||||
def default(self, obj: object) -> object:
|
||||
if isinstance(obj, decimal.Decimal):
|
||||
return float(obj)
|
||||
if isinstance(obj, (datetime.date, datetime.datetime)):
|
||||
return obj.isoformat()
|
||||
if isinstance(obj, bytes):
|
||||
return "<binary_data>"
|
||||
return super().default(obj)
|
||||
|
||||
|
||||
class DB2ToolSchema(BaseModel):
|
||||
"""Input schema for DB2 vector search."""
|
||||
|
||||
query: str = Field(
|
||||
...,
|
||||
description="Query to search in IBM DB2 vector database - always required.",
|
||||
)
|
||||
|
||||
filter_by: str | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Column name used for metadata filtering. "
|
||||
"Must be used together with filter_value."
|
||||
),
|
||||
)
|
||||
|
||||
filter_value: Any | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Value used for metadata filtering. Must be used together with filter_by."
|
||||
),
|
||||
)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _validate_filter_pair(self) -> DB2ToolSchema:
|
||||
if self.filter_by is not None and not self.filter_by.strip():
|
||||
raise ValueError("filter_by must be a non-empty column name.")
|
||||
if (self.filter_by is None) ^ (self.filter_value is None):
|
||||
raise ValueError("filter_by and filter_value must be provided together.")
|
||||
return self
|
||||
|
||||
|
||||
class DB2VectorSearchTool(BaseTool):
|
||||
"""
|
||||
Fortified IBM DB2 Vector Search Tool.
|
||||
Includes SQL injection protection, dynamic relational support, and type-safe serialization.
|
||||
"""
|
||||
|
||||
model_config = ConfigDict(arbitrary_types_allowed=True)
|
||||
|
||||
name: str = "DB2VectorSearchTool"
|
||||
description: str = "Search IBM DB2 vector database for relevant documents. Uses a custom embedding function if supplied, otherwise OpenAI embeddings."
|
||||
args_schema: type[BaseModel] = DB2ToolSchema
|
||||
|
||||
# Internal Whitelist for distance metrics to prevent SQL injection
|
||||
# Aligned with Db2 VECTOR_DISTANCE API:
|
||||
# https://www.ibm.com/docs/en/db2/12.1.x?topic=functions-vector-distance
|
||||
_ALLOWED_METRICS: ClassVar[set[str]] = {
|
||||
"COSINE",
|
||||
"EUCLIDEAN",
|
||||
"EUCLIDEAN_SQUARED",
|
||||
"DOT",
|
||||
"HAMMING",
|
||||
"MANHATTAN",
|
||||
}
|
||||
|
||||
package_dependencies: list[str] = Field(
|
||||
default_factory=lambda: [
|
||||
"ibm_db",
|
||||
"openai", # Optional openai is used for embeddings
|
||||
]
|
||||
)
|
||||
|
||||
env_vars: list[EnvVar] = Field(
|
||||
default_factory=lambda: [
|
||||
EnvVar(
|
||||
name="OPENAI_API_KEY",
|
||||
description="OpenAI API key for embeddings.",
|
||||
required=False,
|
||||
),
|
||||
EnvVar(
|
||||
name="DB2_CONNECTION_STRING",
|
||||
description="IBM DB2 connection string (e.g. 'DATABASE=mydb;HOSTNAME=localhost;PORT=50000;PROTOCOL=TCPIP;UID=user;PWD=pass;').",
|
||||
required=False,
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
connection_string: str = Field(
|
||||
description=(
|
||||
"IBM DB2 connection string. "
|
||||
"Format: 'DATABASE=mydb;HOSTNAME=localhost;PORT=50000;PROTOCOL=TCPIP;UID=user;PWD=pass;' "
|
||||
"or just the database name for a local connection."
|
||||
)
|
||||
)
|
||||
|
||||
# Search settings
|
||||
table_name: str = "documents"
|
||||
vector_column: str = "embedding"
|
||||
embedding_model: str = "text-embedding-3-large"
|
||||
|
||||
return_columns: list[str] = Field(default_factory=lambda: ["content"])
|
||||
|
||||
limit: int = Field(
|
||||
default=3,
|
||||
ge=1,
|
||||
le=100,
|
||||
description="Number of documents to return. Must be between 1 and 100.",
|
||||
)
|
||||
|
||||
distance_metric: str = "COSINE"
|
||||
|
||||
max_distance: float | None = Field(
|
||||
default=None,
|
||||
ge=0.0,
|
||||
description="Maximum allowed distance for results. Cannot be negative.",
|
||||
)
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _validate_return_columns(self) -> DB2VectorSearchTool:
|
||||
if not self.return_columns:
|
||||
raise ValueError(
|
||||
"return_columns cannot be empty. At least one column must be specified "
|
||||
"for the SELECT query to be valid."
|
||||
)
|
||||
return self
|
||||
|
||||
db2_package: Any = Field(default=None, description="IBM DB2 base package.")
|
||||
db2_dbi_package: Any = Field(default=None, description="IBM DB2 DBI package.")
|
||||
|
||||
custom_embedding_fn: ImportString[Callable[[str], list[float]]] | None = Field(
|
||||
default=None,
|
||||
description="Optional custom embedding function.",
|
||||
)
|
||||
|
||||
connection: Any | None = None
|
||||
dbi_connection: Any | None = None
|
||||
cursor: Any | None = None
|
||||
_openai_client: Any | None = None
|
||||
|
||||
def _resolve_db2_packages(self) -> None:
|
||||
"""Lazily resolve IBM DB2 packages on first use.
|
||||
|
||||
Handles both default None values and explicit string inputs
|
||||
(e.g. db2_package="ibm_db") so the field always ends up as
|
||||
the real module object before _connect() uses it.
|
||||
"""
|
||||
if self.db2_package is None or isinstance(self.db2_package, str):
|
||||
pkg_name = self.db2_package or "ibm_db"
|
||||
self.db2_package = importlib.import_module(pkg_name)
|
||||
if self.db2_dbi_package is None or isinstance(self.db2_dbi_package, str):
|
||||
pkg_name = self.db2_dbi_package or "ibm_db_dbi"
|
||||
self.db2_dbi_package = importlib.import_module(pkg_name)
|
||||
|
||||
def _connect(self) -> None:
|
||||
self._resolve_db2_packages()
|
||||
self.connection = self.db2_package.connect(self.connection_string, "", "")
|
||||
self.dbi_connection = self.db2_dbi_package.Connection(self.connection)
|
||||
self.cursor = self.dbi_connection.cursor()
|
||||
|
||||
def _disconnect(self) -> None:
|
||||
try:
|
||||
if self.cursor:
|
||||
self.cursor.close()
|
||||
if self.dbi_connection:
|
||||
self.dbi_connection.close()
|
||||
if self.connection:
|
||||
self.db2_package.close(self.connection)
|
||||
finally:
|
||||
self.connection = None
|
||||
self.dbi_connection = None
|
||||
self.cursor = None
|
||||
|
||||
def _validate_identifier(self, name: str, allow_period: bool = False) -> str:
|
||||
"""
|
||||
Validates table and column names to prevent SQL injection.
|
||||
Simple identifiers must start with a letter and contain only letters, digits,
|
||||
or underscores. Schema-qualified names (allow_period=True) allow exactly one
|
||||
period separating two valid simple identifiers (e.g. myschema.mytable).
|
||||
"""
|
||||
pattern = (
|
||||
r"^[A-Za-z][A-Za-z0-9_]*(\.[A-Za-z][A-Za-z0-9_]*)?$"
|
||||
if allow_period
|
||||
else r"^[A-Za-z][A-Za-z0-9_]*$"
|
||||
)
|
||||
if not re.match(pattern, name):
|
||||
raise ValueError(
|
||||
f"Security Alert: Invalid database identifier detected: {name}"
|
||||
)
|
||||
return name
|
||||
|
||||
def _get_openai_client(self) -> Any:
|
||||
if self._openai_client is None:
|
||||
api_key = os.getenv("OPENAI_API_KEY")
|
||||
if not api_key:
|
||||
raise ValueError(
|
||||
"OPENAI_API_KEY environment variable is missing. Required for default embeddings."
|
||||
)
|
||||
openai = importlib.import_module("openai")
|
||||
self._openai_client = openai.OpenAI(api_key=api_key)
|
||||
return self._openai_client
|
||||
|
||||
def _generate_embedding(self, text: str) -> list[float]:
|
||||
if self.custom_embedding_fn:
|
||||
return self.custom_embedding_fn(text)
|
||||
|
||||
result = (
|
||||
self._get_openai_client()
|
||||
.embeddings.create(
|
||||
input=[text],
|
||||
model=self.embedding_model,
|
||||
)
|
||||
.data[0]
|
||||
.embedding
|
||||
)
|
||||
return list(result)
|
||||
|
||||
def _build_sql(
|
||||
self,
|
||||
column_query: str,
|
||||
v_col: str,
|
||||
vector_dimension: int,
|
||||
metric: str,
|
||||
table: str,
|
||||
filter_clause: str,
|
||||
) -> str:
|
||||
parts = [
|
||||
"SELECT " + column_query + ",",
|
||||
" VECTOR_DISTANCE("
|
||||
+ v_col
|
||||
+ ", VECTOR(CAST(? AS CLOB), "
|
||||
+ str(vector_dimension)
|
||||
+ ", FLOAT32), "
|
||||
+ metric
|
||||
+ ") AS distance",
|
||||
" FROM " + table,
|
||||
" " + filter_clause if filter_clause else "",
|
||||
" ORDER BY distance ASC",
|
||||
" FETCH FIRST " + str(self.limit) + " ROWS ONLY",
|
||||
]
|
||||
return "".join(parts)
|
||||
|
||||
def _run(
|
||||
self,
|
||||
query: str,
|
||||
filter_by: str | None = None,
|
||||
filter_value: Any | None = None,
|
||||
) -> str:
|
||||
# Validate query is not blank or whitespace-only
|
||||
if query is None or query.strip() == "":
|
||||
return json.dumps(
|
||||
{
|
||||
"success": False,
|
||||
"error": "Query cannot be empty or contain only whitespace.",
|
||||
},
|
||||
indent=2,
|
||||
)
|
||||
|
||||
try:
|
||||
query_vector = self._generate_embedding(query)
|
||||
|
||||
# Explicit Connection Handling
|
||||
try:
|
||||
self._connect()
|
||||
except Exception as e:
|
||||
self._disconnect() # Clean up any partial connection
|
||||
return json.dumps(
|
||||
{"success": False, "error": f"Failed to connect to DB2: {e!s}"}
|
||||
)
|
||||
|
||||
# Validate Metric
|
||||
metric = self.distance_metric.upper()
|
||||
if metric not in self._ALLOWED_METRICS:
|
||||
raise ValueError(f"Invalid distance metric: {metric}")
|
||||
|
||||
# Validate Identifiers
|
||||
table = self._validate_identifier(self.table_name, allow_period=True)
|
||||
v_col = self._validate_identifier(self.vector_column)
|
||||
ret_cols = [self._validate_identifier(c) for c in self.return_columns]
|
||||
|
||||
vector_dimension = len(query_vector)
|
||||
vector_string = str(query_vector)
|
||||
|
||||
filter_clause = ""
|
||||
params = [vector_string] # The vector string for the CLOB cast
|
||||
|
||||
if filter_by and filter_value is not None:
|
||||
f_col = self._validate_identifier(filter_by)
|
||||
filter_clause = f"WHERE {f_col} = ?"
|
||||
params.append(filter_value)
|
||||
|
||||
# DYNAMIC COLUMN SELECTION
|
||||
column_query = ", ".join(ret_cols)
|
||||
|
||||
sql = self._build_sql(
|
||||
column_query, v_col, vector_dimension, metric, table, filter_clause
|
||||
)
|
||||
|
||||
assert self.cursor is not None # noqa: S101
|
||||
self.cursor.execute(sql, tuple(params))
|
||||
rows = self.cursor.fetchall()
|
||||
|
||||
normalized_results = []
|
||||
|
||||
for row in rows:
|
||||
# The 'distance' is always the LAST column in our dynamic SELECT
|
||||
distance = float(row[-1])
|
||||
|
||||
if self.max_distance is not None and distance > self.max_distance:
|
||||
continue
|
||||
|
||||
# Automatically map the requested columns to their row values
|
||||
row_data = dict(zip(self.return_columns, row[:-1], strict=False))
|
||||
|
||||
normalized_results.append(
|
||||
{
|
||||
"distance": distance,
|
||||
"data": row_data,
|
||||
}
|
||||
)
|
||||
|
||||
# Explicit cleanup
|
||||
self._disconnect()
|
||||
|
||||
return json.dumps(
|
||||
{
|
||||
"success": True,
|
||||
"results": normalized_results,
|
||||
},
|
||||
indent=2,
|
||||
cls=DB2JSONEncoder,
|
||||
)
|
||||
|
||||
except Exception as error:
|
||||
self._disconnect()
|
||||
return json.dumps(
|
||||
{
|
||||
"success": False,
|
||||
"error": str(error),
|
||||
"error_type": type(error).__name__,
|
||||
},
|
||||
indent=2,
|
||||
)
|
||||
|
||||
def __del__(self) -> None:
|
||||
self._disconnect()
|
||||
707
lib/crewai-tools/tests/tools/test_db2_search_tool.py
Normal file
707
lib/crewai-tools/tests/tools/test_db2_search_tool.py
Normal file
@@ -0,0 +1,707 @@
|
||||
"""Tests for DB2VectorSearchTool.
|
||||
|
||||
All tests are fully unit-tested — no real IBM DB2 instance is required.
|
||||
ibm_db and ibm_db_dbi are mocked at import time so the suite runs without
|
||||
those optional packages installed.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import decimal
|
||||
import datetime
|
||||
import json
|
||||
import sys
|
||||
from types import ModuleType
|
||||
from unittest.mock import MagicMock, patch, call
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Stub ibm_db / ibm_db_dbi before any crewai_tools import, so the
|
||||
# ImportString validator on DB2VectorSearchTool does not fail.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def _make_ibm_db_stub() -> ModuleType:
|
||||
mod = ModuleType("ibm_db")
|
||||
mod.connect = MagicMock()
|
||||
mod.close = MagicMock()
|
||||
return mod
|
||||
|
||||
|
||||
def _make_ibm_db_dbi_stub() -> ModuleType:
|
||||
mod = ModuleType("ibm_db_dbi")
|
||||
|
||||
class FakeConnection:
|
||||
def __init__(self, conn):
|
||||
self._conn = conn
|
||||
self.cursor = MagicMock(return_value=MagicMock())
|
||||
|
||||
def close(self):
|
||||
pass
|
||||
|
||||
mod.Connection = FakeConnection
|
||||
return mod
|
||||
|
||||
|
||||
# Inject stubs before importing tool module
|
||||
_ibm_db_stub = _make_ibm_db_stub()
|
||||
_ibm_db_dbi_stub = _make_ibm_db_dbi_stub()
|
||||
sys.modules.setdefault("ibm_db", _ibm_db_stub)
|
||||
sys.modules.setdefault("ibm_db_dbi", _ibm_db_dbi_stub)
|
||||
|
||||
from crewai_tools.tools.db2_search_tool.db2_search_tool import ( # noqa: E402
|
||||
DB2JSONEncoder,
|
||||
DB2ToolSchema,
|
||||
DB2VectorSearchTool,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def _make_tool(
|
||||
*,
|
||||
table_name: str = "documents",
|
||||
vector_column: str = "embedding",
|
||||
return_columns: list[str] | None = None,
|
||||
limit: int = 3,
|
||||
distance_metric: str = "COSINE",
|
||||
max_distance: float | None = None,
|
||||
embedding_model: str = "text-embedding-3-large",
|
||||
custom_embedding_fn=None,
|
||||
) -> DB2VectorSearchTool:
|
||||
"""Return a DB2VectorSearchTool with mocked ibm_db packages."""
|
||||
return DB2VectorSearchTool(
|
||||
connection_string="DATABASE=TESTDB;HOSTNAME=localhost;PORT=50000;PROTOCOL=TCPIP;UID=user;PWD=pass;",
|
||||
table_name=table_name,
|
||||
vector_column=vector_column,
|
||||
return_columns=return_columns or ["content"],
|
||||
limit=limit,
|
||||
distance_metric=distance_metric,
|
||||
max_distance=max_distance,
|
||||
embedding_model=embedding_model,
|
||||
db2_package=_ibm_db_stub,
|
||||
db2_dbi_package=_ibm_db_dbi_stub,
|
||||
custom_embedding_fn=custom_embedding_fn,
|
||||
)
|
||||
|
||||
|
||||
def _fake_embedding(text: str) -> list[float]:
|
||||
return [0.1, 0.2, 0.3]
|
||||
|
||||
|
||||
def _make_cursor_with_rows(rows: list[tuple]) -> MagicMock:
|
||||
cursor = MagicMock()
|
||||
cursor.fetchall.return_value = rows
|
||||
return cursor
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# DB2ToolSchema validation
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestDB2ToolSchema:
|
||||
def test_valid_query_only(self):
|
||||
schema = DB2ToolSchema(query="find documents about AI")
|
||||
assert schema.query == "find documents about AI"
|
||||
assert schema.filter_by is None
|
||||
assert schema.filter_value is None
|
||||
|
||||
def test_valid_query_with_filter_pair(self):
|
||||
schema = DB2ToolSchema(query="search", filter_by="category", filter_value="tech")
|
||||
assert schema.filter_by == "category"
|
||||
assert schema.filter_value == "tech"
|
||||
|
||||
def test_filter_by_without_filter_value_raises(self):
|
||||
with pytest.raises(ValueError, match="filter_by and filter_value must be provided together"):
|
||||
DB2ToolSchema(query="search", filter_by="category")
|
||||
|
||||
def test_filter_value_without_filter_by_raises(self):
|
||||
with pytest.raises(ValueError, match="filter_by and filter_value must be provided together"):
|
||||
DB2ToolSchema(query="search", filter_value="tech")
|
||||
|
||||
def test_blank_filter_by_raises(self):
|
||||
with pytest.raises(ValueError, match="filter_by must be a non-empty column name"):
|
||||
DB2ToolSchema(query="search", filter_by=" ", filter_value="tech")
|
||||
|
||||
def test_none_filter_by_and_none_filter_value_is_valid(self):
|
||||
schema = DB2ToolSchema(query="hello", filter_by=None, filter_value=None)
|
||||
assert schema.filter_by is None
|
||||
assert schema.filter_value is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# DB2VectorSearchTool field validation
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestDB2VectorSearchToolConfig:
|
||||
_conn = "DATABASE=MYDB;HOSTNAME=localhost;PORT=50000;PROTOCOL=TCPIP;UID=u;PWD=p;"
|
||||
|
||||
def test_default_values(self):
|
||||
tool = DB2VectorSearchTool(
|
||||
connection_string=self._conn,
|
||||
db2_package=_ibm_db_stub,
|
||||
db2_dbi_package=_ibm_db_dbi_stub,
|
||||
)
|
||||
assert tool.return_columns == ["content"]
|
||||
assert tool.limit == 3
|
||||
assert tool.distance_metric == "COSINE"
|
||||
assert tool.max_distance is None
|
||||
|
||||
def test_empty_return_columns_raises(self):
|
||||
with pytest.raises(ValueError, match="return_columns cannot be empty"):
|
||||
DB2VectorSearchTool(
|
||||
connection_string=self._conn,
|
||||
return_columns=[],
|
||||
db2_package=_ibm_db_stub,
|
||||
db2_dbi_package=_ibm_db_dbi_stub,
|
||||
)
|
||||
|
||||
def test_limit_out_of_range_raises(self):
|
||||
with pytest.raises(ValueError):
|
||||
DB2VectorSearchTool(
|
||||
connection_string=self._conn,
|
||||
limit=0,
|
||||
db2_package=_ibm_db_stub,
|
||||
db2_dbi_package=_ibm_db_dbi_stub,
|
||||
)
|
||||
with pytest.raises(ValueError):
|
||||
DB2VectorSearchTool(
|
||||
connection_string=self._conn,
|
||||
limit=101,
|
||||
db2_package=_ibm_db_stub,
|
||||
db2_dbi_package=_ibm_db_dbi_stub,
|
||||
)
|
||||
|
||||
def test_negative_max_distance_raises(self):
|
||||
with pytest.raises(ValueError):
|
||||
DB2VectorSearchTool(
|
||||
connection_string=self._conn,
|
||||
max_distance=-1.0,
|
||||
db2_package=_ibm_db_stub,
|
||||
db2_dbi_package=_ibm_db_dbi_stub,
|
||||
)
|
||||
|
||||
def test_multiple_return_columns(self):
|
||||
tool = DB2VectorSearchTool(
|
||||
connection_string=self._conn,
|
||||
return_columns=["title", "body", "author"],
|
||||
db2_package=_ibm_db_stub,
|
||||
db2_dbi_package=_ibm_db_dbi_stub,
|
||||
)
|
||||
assert tool.return_columns == ["title", "body", "author"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# DB2JSONEncoder
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestDB2JSONEncoder:
|
||||
def test_encodes_decimal(self):
|
||||
result = json.dumps(decimal.Decimal("3.14"), cls=DB2JSONEncoder)
|
||||
assert result == "3.14"
|
||||
|
||||
def test_encodes_datetime(self):
|
||||
dt = datetime.datetime(2024, 1, 15, 12, 0, 0)
|
||||
result = json.dumps(dt, cls=DB2JSONEncoder)
|
||||
assert "2024-01-15" in result
|
||||
|
||||
def test_encodes_date(self):
|
||||
d = datetime.date(2024, 6, 1)
|
||||
result = json.dumps(d, cls=DB2JSONEncoder)
|
||||
assert "2024-06-01" in result
|
||||
|
||||
def test_encodes_bytes(self):
|
||||
result = json.dumps(b"\x00\xff", cls=DB2JSONEncoder)
|
||||
assert "<binary_data>" in result
|
||||
|
||||
def test_raises_for_unknown_type(self):
|
||||
class Unknown:
|
||||
pass
|
||||
with pytest.raises(TypeError):
|
||||
json.dumps(Unknown(), cls=DB2JSONEncoder)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _validate_identifier (SQL injection guard)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestValidateIdentifier:
|
||||
def test_valid_simple_name(self):
|
||||
tool = _make_tool()
|
||||
assert tool._validate_identifier("documents") == "documents"
|
||||
assert tool._validate_identifier("my_table_1") == "my_table_1"
|
||||
|
||||
def test_valid_schema_qualified_with_period(self):
|
||||
tool = _make_tool()
|
||||
assert tool._validate_identifier("myschema.documents", allow_period=True) == "myschema.documents"
|
||||
|
||||
def test_period_without_allow_period_raises(self):
|
||||
tool = _make_tool()
|
||||
with pytest.raises(ValueError, match="Security Alert"):
|
||||
tool._validate_identifier("schema.table", allow_period=False)
|
||||
|
||||
@pytest.mark.parametrize("bad_name", [
|
||||
"'; DROP TABLE documents; --",
|
||||
"table--",
|
||||
"col name",
|
||||
"col;name",
|
||||
"col OR 1=1",
|
||||
"",
|
||||
"1table", # must start with a letter
|
||||
"123", # must start with a letter
|
||||
".documents", # leading period
|
||||
"schema..table", # double period
|
||||
"schema.table.extra", # more than one period
|
||||
".....", # only dots — previously passed old regex
|
||||
])
|
||||
def test_injection_strings_raise(self, bad_name: str):
|
||||
tool = _make_tool()
|
||||
with pytest.raises(ValueError, match="Security Alert"):
|
||||
tool._validate_identifier(bad_name)
|
||||
|
||||
def test_allow_period_rejects_digit_led_schema(self):
|
||||
tool = _make_tool()
|
||||
with pytest.raises(ValueError, match="Security Alert"):
|
||||
tool._validate_identifier("1schema.table", allow_period=True)
|
||||
|
||||
def test_allow_period_rejects_digit_led_table(self):
|
||||
tool = _make_tool()
|
||||
with pytest.raises(ValueError, match="Security Alert"):
|
||||
tool._validate_identifier("schema.1table", allow_period=True)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _generate_embedding
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestGenerateEmbedding:
|
||||
def test_uses_custom_embedding_fn(self):
|
||||
called_with = []
|
||||
|
||||
def my_embed(text: str) -> list[float]:
|
||||
called_with.append(text)
|
||||
return [0.5, 0.5]
|
||||
|
||||
tool = _make_tool(custom_embedding_fn=my_embed)
|
||||
result = tool._generate_embedding("hello world")
|
||||
assert result == [0.5, 0.5]
|
||||
assert called_with == ["hello world"]
|
||||
|
||||
def test_falls_back_to_openai_with_api_key(self):
|
||||
tool = _make_tool()
|
||||
tool._openai_client = None # ensure cache is clear
|
||||
mock_openai = MagicMock()
|
||||
mock_openai.OpenAI.return_value.embeddings.create.return_value.data = [
|
||||
MagicMock(embedding=[0.1, 0.2])
|
||||
]
|
||||
|
||||
with patch.dict("os.environ", {"OPENAI_API_KEY": "test-key"}):
|
||||
with patch.dict("sys.modules", {"openai": mock_openai}):
|
||||
result = tool._generate_embedding("test query")
|
||||
|
||||
assert result == [0.1, 0.2]
|
||||
|
||||
def test_openai_client_is_reused_across_calls(self):
|
||||
tool = _make_tool()
|
||||
tool._openai_client = None # ensure cache is clear
|
||||
mock_openai = MagicMock()
|
||||
mock_client = mock_openai.OpenAI.return_value
|
||||
mock_client.embeddings.create.return_value.data = [MagicMock(embedding=[0.1, 0.2])]
|
||||
|
||||
with patch.dict("os.environ", {"OPENAI_API_KEY": "test-key"}):
|
||||
with patch.dict("sys.modules", {"openai": mock_openai}):
|
||||
tool._generate_embedding("first query")
|
||||
tool._generate_embedding("second query")
|
||||
|
||||
# OpenAI() constructor called only once — client was reused
|
||||
mock_openai.OpenAI.assert_called_once()
|
||||
|
||||
def test_raises_when_no_openai_key_and_no_custom_fn(self):
|
||||
tool = _make_tool()
|
||||
tool._openai_client = None # ensure cache is clear
|
||||
|
||||
import os
|
||||
env_without_key = {k: v for k, v in os.environ.items() if k != "OPENAI_API_KEY"}
|
||||
with patch.dict("os.environ", env_without_key, clear=True):
|
||||
with pytest.raises(ValueError, match="OPENAI_API_KEY"):
|
||||
tool._generate_embedding("test")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _run — empty / whitespace query guard
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestRunQueryValidation:
|
||||
def test_empty_query_returns_error_json(self):
|
||||
tool = _make_tool(custom_embedding_fn=_fake_embedding)
|
||||
result = json.loads(tool._run(query=""))
|
||||
assert result["success"] is False
|
||||
assert "empty" in result["error"].lower()
|
||||
|
||||
def test_whitespace_only_query_returns_error_json(self):
|
||||
tool = _make_tool(custom_embedding_fn=_fake_embedding)
|
||||
result = json.loads(tool._run(query=" "))
|
||||
assert result["success"] is False
|
||||
|
||||
def test_none_query_returns_error_json(self):
|
||||
tool = _make_tool(custom_embedding_fn=_fake_embedding)
|
||||
result = json.loads(tool._run(query=None))
|
||||
assert result["success"] is False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _run — connection failure
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestRunConnectionFailure:
|
||||
def test_connection_error_returns_error_json(self):
|
||||
tool = _make_tool(custom_embedding_fn=_fake_embedding)
|
||||
with patch.object(tool, "_connect", side_effect=Exception("Connection refused")):
|
||||
result = json.loads(tool._run(query="find AI docs"))
|
||||
assert result["success"] is False
|
||||
assert "Failed to connect to DB2" in result["error"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _run — invalid distance metric
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestRunInvalidMetric:
|
||||
def test_invalid_metric_returns_error_json(self):
|
||||
tool = _make_tool(
|
||||
custom_embedding_fn=_fake_embedding,
|
||||
distance_metric="INVALID_METRIC",
|
||||
)
|
||||
mock_cursor = _make_cursor_with_rows([])
|
||||
with patch.object(tool, "_connect"):
|
||||
with patch.object(tool, "_disconnect"):
|
||||
tool.cursor = mock_cursor
|
||||
result = json.loads(tool._run(query="test"))
|
||||
assert result["success"] is False
|
||||
assert "Invalid distance metric" in result["error"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _run — successful search (core happy path)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestRunSuccessful:
|
||||
def _setup_connected_tool(self, rows: list[tuple], **kwargs) -> DB2VectorSearchTool:
|
||||
tool = _make_tool(custom_embedding_fn=_fake_embedding, **kwargs)
|
||||
mock_cursor = _make_cursor_with_rows(rows)
|
||||
tool.cursor = mock_cursor
|
||||
return tool, mock_cursor
|
||||
|
||||
def test_returns_results_as_json(self):
|
||||
rows = [("Some document text", 0.12)]
|
||||
tool, cursor = self._setup_connected_tool(rows)
|
||||
|
||||
with patch.object(tool, "_connect"):
|
||||
with patch.object(tool, "_disconnect"):
|
||||
result = json.loads(tool._run(query="find documents about AI"))
|
||||
|
||||
assert result["success"] is True
|
||||
assert len(result["results"]) == 1
|
||||
assert result["results"][0]["distance"] == pytest.approx(0.12)
|
||||
assert result["results"][0]["data"]["content"] == "Some document text"
|
||||
|
||||
def test_multiple_return_columns_mapped_correctly(self):
|
||||
rows = [("Title A", "Body text A", 0.05)]
|
||||
tool, cursor = self._setup_connected_tool(
|
||||
rows, return_columns=["title", "body"]
|
||||
)
|
||||
|
||||
with patch.object(tool, "_connect"):
|
||||
with patch.object(tool, "_disconnect"):
|
||||
result = json.loads(tool._run(query="search"))
|
||||
|
||||
data = result["results"][0]["data"]
|
||||
assert data["title"] == "Title A"
|
||||
assert data["body"] == "Body text A"
|
||||
|
||||
def test_empty_db_result_returns_empty_list(self):
|
||||
tool, _ = self._setup_connected_tool([])
|
||||
|
||||
with patch.object(tool, "_connect"):
|
||||
with patch.object(tool, "_disconnect"):
|
||||
result = json.loads(tool._run(query="nothing"))
|
||||
|
||||
assert result["success"] is True
|
||||
assert result["results"] == []
|
||||
|
||||
def test_max_distance_filters_far_results(self):
|
||||
# Row 0 is close (0.2), Row 1 is too far (0.9)
|
||||
rows = [("Close doc", 0.2), ("Far doc", 0.9)]
|
||||
tool, _ = self._setup_connected_tool(rows, max_distance=0.5)
|
||||
|
||||
with patch.object(tool, "_connect"):
|
||||
with patch.object(tool, "_disconnect"):
|
||||
result = json.loads(tool._run(query="test"))
|
||||
|
||||
assert result["success"] is True
|
||||
assert len(result["results"]) == 1
|
||||
assert result["results"][0]["data"]["content"] == "Close doc"
|
||||
|
||||
def test_filter_by_and_filter_value_added_to_params(self):
|
||||
rows = [("Filtered doc", 0.1)]
|
||||
tool, cursor = self._setup_connected_tool(rows)
|
||||
|
||||
with patch.object(tool, "_connect"):
|
||||
with patch.object(tool, "_disconnect"):
|
||||
result = json.loads(
|
||||
tool._run(query="test", filter_by="category", filter_value="AI")
|
||||
)
|
||||
|
||||
assert result["success"] is True
|
||||
# The second param in the execute call must be the filter value
|
||||
execute_args = cursor.execute.call_args
|
||||
params_tuple = execute_args[0][1]
|
||||
assert "AI" in params_tuple
|
||||
|
||||
def test_sql_contains_correct_metric(self):
|
||||
rows = [("doc", 0.1)]
|
||||
tool, cursor = self._setup_connected_tool(rows, distance_metric="EUCLIDEAN")
|
||||
|
||||
with patch.object(tool, "_connect"):
|
||||
with patch.object(tool, "_disconnect"):
|
||||
tool._run(query="test")
|
||||
|
||||
executed_sql = cursor.execute.call_args[0][0]
|
||||
assert "EUCLIDEAN" in executed_sql
|
||||
|
||||
def test_sql_contains_correct_limit(self):
|
||||
rows = []
|
||||
tool, cursor = self._setup_connected_tool(rows, limit=7)
|
||||
|
||||
with patch.object(tool, "_connect"):
|
||||
with patch.object(tool, "_disconnect"):
|
||||
tool._run(query="test")
|
||||
|
||||
executed_sql = cursor.execute.call_args[0][0]
|
||||
assert "7" in executed_sql
|
||||
|
||||
def test_sql_contains_where_clause_when_filter_provided(self):
|
||||
rows = []
|
||||
tool, cursor = self._setup_connected_tool(rows)
|
||||
|
||||
with patch.object(tool, "_connect"):
|
||||
with patch.object(tool, "_disconnect"):
|
||||
tool._run(query="test", filter_by="dept", filter_value="HR")
|
||||
|
||||
executed_sql = cursor.execute.call_args[0][0]
|
||||
assert "WHERE dept = ?" in executed_sql
|
||||
|
||||
def test_sql_has_no_where_clause_without_filter(self):
|
||||
rows = []
|
||||
tool, cursor = self._setup_connected_tool(rows)
|
||||
|
||||
with patch.object(tool, "_connect"):
|
||||
with patch.object(tool, "_disconnect"):
|
||||
tool._run(query="test")
|
||||
|
||||
executed_sql = cursor.execute.call_args[0][0]
|
||||
assert "WHERE" not in executed_sql
|
||||
|
||||
def test_json_encoder_handles_decimal_in_results(self):
|
||||
rows = [(decimal.Decimal("42.50"), 0.1)]
|
||||
tool, _ = self._setup_connected_tool(rows, return_columns=["price"])
|
||||
|
||||
with patch.object(tool, "_connect"):
|
||||
with patch.object(tool, "_disconnect"):
|
||||
result = json.loads(tool._run(query="test"))
|
||||
|
||||
assert result["success"] is True
|
||||
assert result["results"][0]["data"]["price"] == pytest.approx(42.5)
|
||||
|
||||
def test_disconnect_called_after_successful_run(self):
|
||||
rows = [("doc", 0.1)]
|
||||
tool, cursor = self._setup_connected_tool(rows)
|
||||
|
||||
with patch.object(tool, "_connect"):
|
||||
with patch.object(tool, "_disconnect") as mock_disconnect:
|
||||
tool._run(query="test")
|
||||
|
||||
mock_disconnect.assert_called_once()
|
||||
|
||||
def test_disconnect_called_on_unexpected_error(self):
|
||||
tool = _make_tool(custom_embedding_fn=_fake_embedding)
|
||||
|
||||
with patch.object(tool, "_connect"):
|
||||
with patch.object(tool, "_disconnect") as mock_disconnect:
|
||||
# cursor is None → will raise AttributeError inside _run
|
||||
tool.cursor = None
|
||||
# Override _connect to set cursor to a raising mock
|
||||
def bad_cursor_setup():
|
||||
c = MagicMock()
|
||||
c.execute.side_effect = RuntimeError("Unexpected DB error")
|
||||
tool.cursor = c
|
||||
|
||||
tool._connect = bad_cursor_setup
|
||||
result = json.loads(tool._run(query="test"))
|
||||
|
||||
assert result["success"] is False
|
||||
mock_disconnect.assert_called()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _run — SQL injection via filter_by rejected
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestRunSQLInjectionPrevention:
|
||||
@pytest.mark.parametrize("bad_col", [
|
||||
"col; DROP TABLE documents; --",
|
||||
"col OR 1=1",
|
||||
"col name",
|
||||
# NOTE: empty string is falsy — _run skips the WHERE clause entirely
|
||||
# so it does NOT trigger _validate_identifier. The schema-level guard
|
||||
# (DB2ToolSchema._validate_filter_pair) catches the empty string case.
|
||||
])
|
||||
def test_injection_in_filter_by_returns_error(self, bad_col: str):
|
||||
tool = _make_tool(custom_embedding_fn=_fake_embedding)
|
||||
mock_cursor = _make_cursor_with_rows([])
|
||||
tool.cursor = mock_cursor
|
||||
|
||||
with patch.object(tool, "_connect"):
|
||||
with patch.object(tool, "_disconnect"):
|
||||
result = json.loads(
|
||||
tool._run(query="test", filter_by=bad_col, filter_value="val")
|
||||
)
|
||||
|
||||
assert result["success"] is False
|
||||
|
||||
def test_empty_filter_by_bypasses_where_clause(self):
|
||||
"""Empty string is falsy in Python — _run skips WHERE rather than injecting.
|
||||
The actual guard lives in DB2ToolSchema (schema-level validation).
|
||||
"""
|
||||
tool = _make_tool(custom_embedding_fn=_fake_embedding)
|
||||
mock_cursor = _make_cursor_with_rows([])
|
||||
tool.cursor = mock_cursor
|
||||
|
||||
with patch.object(tool, "_connect"):
|
||||
with patch.object(tool, "_disconnect"):
|
||||
result = json.loads(
|
||||
tool._run(query="test", filter_by="", filter_value="val")
|
||||
)
|
||||
|
||||
# The query succeeds (no WHERE clause injected) — success is True
|
||||
assert result["success"] is True
|
||||
executed_sql = mock_cursor.execute.call_args[0][0]
|
||||
assert "WHERE" not in executed_sql
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# _connect / _disconnect lifecycle
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestConnectDisconnect:
|
||||
def test_connect_builds_connection_objects(self):
|
||||
tool = _make_tool(custom_embedding_fn=_fake_embedding)
|
||||
mock_conn = MagicMock()
|
||||
_ibm_db_stub.connect.return_value = mock_conn
|
||||
|
||||
tool._connect()
|
||||
|
||||
assert tool.connection is mock_conn
|
||||
assert tool.cursor is not None
|
||||
|
||||
def test_connect_opens_fresh_connection_each_call(self):
|
||||
tool = _make_tool(custom_embedding_fn=_fake_embedding)
|
||||
mock_conn = MagicMock()
|
||||
local_connect = MagicMock(return_value=mock_conn)
|
||||
tool.db2_package = MagicMock()
|
||||
tool.db2_package.connect = local_connect
|
||||
tool.db2_package.close = MagicMock()
|
||||
|
||||
tool._connect()
|
||||
tool._connect() # connect-per-call: each invocation opens a new connection
|
||||
|
||||
assert local_connect.call_count == 2
|
||||
|
||||
def test_disconnect_resets_all_handles(self):
|
||||
tool = _make_tool(custom_embedding_fn=_fake_embedding)
|
||||
mock_conn = MagicMock()
|
||||
_ibm_db_stub.connect.return_value = mock_conn
|
||||
|
||||
tool._connect()
|
||||
tool._disconnect()
|
||||
|
||||
assert tool.connection is None
|
||||
assert tool.dbi_connection is None
|
||||
assert tool.cursor is None
|
||||
|
||||
def test_disconnect_is_safe_when_already_disconnected(self):
|
||||
tool = _make_tool(custom_embedding_fn=_fake_embedding)
|
||||
# Should not raise even with no open connection
|
||||
tool._disconnect()
|
||||
|
||||
def test_del_calls_disconnect(self):
|
||||
tool = _make_tool(custom_embedding_fn=_fake_embedding)
|
||||
with patch.object(tool, "_disconnect") as mock_disconnect:
|
||||
tool.__del__()
|
||||
mock_disconnect.assert_called_once()
|
||||
|
||||
def test_connect_resolves_packages_without_injection(self):
|
||||
"""Constructs the tool WITHOUT injecting db2_package / db2_dbi_package.
|
||||
|
||||
Verifies that _connect() automatically resolves package fields from sys.modules
|
||||
when left at their default of None, and successfully establishes a connection.
|
||||
"""
|
||||
tool = DB2VectorSearchTool(
|
||||
connection_string="DATABASE=TESTDB;HOSTNAME=localhost;PORT=50000;PROTOCOL=TCPIP;UID=user;PWD=pass;",
|
||||
custom_embedding_fn=_fake_embedding,
|
||||
)
|
||||
# Both fields start as None
|
||||
assert tool.db2_package is None
|
||||
assert tool.db2_dbi_package is None
|
||||
|
||||
# Exercise behavior via _connect()
|
||||
_ibm_db_stub.connect.return_value = MagicMock()
|
||||
tool._connect()
|
||||
|
||||
# Verify side-effects: packages were resolved and connections established
|
||||
assert tool.db2_package is _ibm_db_stub
|
||||
assert tool.db2_dbi_package is _ibm_db_dbi_stub
|
||||
assert tool.connection is not None
|
||||
assert tool.cursor is not None
|
||||
|
||||
tool._disconnect()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tool metadata
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TestToolMetadata:
|
||||
def test_tool_name(self):
|
||||
tool = _make_tool()
|
||||
assert tool.name == "DB2VectorSearchTool"
|
||||
|
||||
def test_tool_description(self):
|
||||
tool = _make_tool()
|
||||
assert "DB2" in tool.description
|
||||
assert "custom embedding function" in tool.description
|
||||
assert "OpenAI embeddings" in tool.description
|
||||
|
||||
def test_args_schema_is_db2_tool_schema(self):
|
||||
tool = _make_tool()
|
||||
assert tool.args_schema is DB2ToolSchema
|
||||
|
||||
def test_package_dependencies_listed(self):
|
||||
tool = _make_tool()
|
||||
assert "ibm_db" in tool.package_dependencies
|
||||
|
||||
def test_env_vars_declared(self):
|
||||
tool = _make_tool()
|
||||
env_var_names = {ev.name for ev in tool.env_vars}
|
||||
assert "OPENAI_API_KEY" in env_var_names
|
||||
assert "DB2_CONNECTION_STRING" in env_var_names
|
||||
|
||||
def test_public_import_from_crewai_tools(self):
|
||||
"""from crewai_tools import DB2VectorSearchTool must work at package level."""
|
||||
from crewai_tools import DB2ToolSchema # noqa: PLC0415
|
||||
from crewai_tools import DB2VectorSearchTool # noqa: PLC0415
|
||||
|
||||
assert DB2VectorSearchTool is not None
|
||||
assert DB2ToolSchema is not None
|
||||
Reference in New Issue
Block a user