mirror of
https://github.com/crewAIInc/crewAI.git
synced 2026-09-11 05:28:00 +00:00
Merge branch 'main' into feat/agents-md-amp-and-coding-agent-telemetry
This commit is contained in:
@@ -271,7 +271,8 @@
|
|||||||
"edge/en/tools/database-data/qdrantvectorsearchtool",
|
"edge/en/tools/database-data/qdrantvectorsearchtool",
|
||||||
"edge/en/tools/database-data/weaviatevectorsearchtool",
|
"edge/en/tools/database-data/weaviatevectorsearchtool",
|
||||||
"edge/en/tools/database-data/mongodbvectorsearchtool",
|
"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 (
|
from crewai_tools.tools.databricks_query_tool.databricks_query_tool import (
|
||||||
DatabricksQueryTool,
|
DatabricksQueryTool,
|
||||||
)
|
)
|
||||||
|
from crewai_tools.tools.db2_search_tool import (
|
||||||
|
DB2ToolSchema,
|
||||||
|
DB2VectorSearchTool,
|
||||||
|
)
|
||||||
from crewai_tools.tools.daytona_sandbox_tool import (
|
from crewai_tools.tools.daytona_sandbox_tool import (
|
||||||
DaytonaExecTool,
|
DaytonaExecTool,
|
||||||
DaytonaFileTool,
|
DaytonaFileTool,
|
||||||
@@ -248,6 +252,8 @@ __all__ = [
|
|||||||
"CrewaiPlatformTools",
|
"CrewaiPlatformTools",
|
||||||
"DOCXSearchTool",
|
"DOCXSearchTool",
|
||||||
"DallETool",
|
"DallETool",
|
||||||
|
"DB2ToolSchema",
|
||||||
|
"DB2VectorSearchTool",
|
||||||
"DatabricksQueryTool",
|
"DatabricksQueryTool",
|
||||||
"DaytonaExecTool",
|
"DaytonaExecTool",
|
||||||
"DaytonaFileTool",
|
"DaytonaFileTool",
|
||||||
|
|||||||
@@ -53,6 +53,10 @@ from crewai_tools.tools.daytona_sandbox_tool import (
|
|||||||
DaytonaFileTool,
|
DaytonaFileTool,
|
||||||
DaytonaPythonTool,
|
DaytonaPythonTool,
|
||||||
)
|
)
|
||||||
|
from crewai_tools.tools.db2_search_tool import (
|
||||||
|
DB2ToolSchema,
|
||||||
|
DB2VectorSearchTool,
|
||||||
|
)
|
||||||
from crewai_tools.tools.directory_read_tool.directory_read_tool import (
|
from crewai_tools.tools.directory_read_tool.directory_read_tool import (
|
||||||
DirectoryReadTool,
|
DirectoryReadTool,
|
||||||
)
|
)
|
||||||
@@ -234,6 +238,8 @@ __all__ = [
|
|||||||
"DOCXSearchTool",
|
"DOCXSearchTool",
|
||||||
"DallETool",
|
"DallETool",
|
||||||
"DatabricksQueryTool",
|
"DatabricksQueryTool",
|
||||||
|
"DB2ToolSchema",
|
||||||
|
"DB2VectorSearchTool",
|
||||||
"DaytonaExecTool",
|
"DaytonaExecTool",
|
||||||
"DaytonaFileTool",
|
"DaytonaFileTool",
|
||||||
"DaytonaPythonTool",
|
"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