Compare commits

..

1 Commits

Author SHA1 Message Date
PawanThakurIBM
c8f441cffa feat(crewai-tools): add IBM Db2 search tool (#5885)
Some checks failed
CodeQL Advanced / Analyze (actions) (push) Has been cancelled
CodeQL Advanced / Analyze (python) (push) Has been cancelled
Check Documentation Broken Links / Check broken links (push) Has been cancelled
Vulnerability Scan / pip-audit (push) Has been cancelled
* feat(crewai-tools): add db2 search tool

* refactor(crewai-tools): improve db2 search tool implementation

* feat(tools): improve DB2VectorSearchTool validation, security, and configurability

* docs: add DB2SearchTool documentation

* feat: add DB2 search tool

* docs: update DB2SearchTool documentation

* fix: address CodeRabbit review feedback

* fix: validate non-empty filter_by in DB2ToolSchema

* chore: trigger CodeRabbit re-review

* feat: fortify DB2 tool; fixed JSON response shape, added input guards and config validation

* refactor(db2): replace DB2Config with connection_string field

* refactor(db2): remove dead _setup_db2 validator and importlib import

* refactor(db2): remove dead guard in _connect as  _disconnect() is called at the end of every
_run, so self.connection is always None when _connect is called next.
The 'if not self.connection' guard was dead code.

* fix(db2): tighten _validate_identifier regex. Old regex allowed leading digits, multiple periods and dot-only strings (e.g. '.....' passed).

* fix(db2): replace __import__ with importlib.import_module in _generate_embedding as keeping openai as a lazy optional import since it is not always required.

* perf(db2): cache OpenAI client in _openai_client to avoid re-instantiation as OpenAI(api_key=...) was recreated on every _generate_embedding call. Extract into _get_openai_client() which lazily initialises and caches self._openai_client on first use, reusing it for all subsequent queries.

* docs(db2): clarify tool description to mention embedding fallback

* docs(db2): update README supported features to clarify embedding behaviour. 'OpenAI embedding fallback' implied it was
optional. Replaced with 'Uses a custom embedding function if supplied,
otherwise OpenAI embeddings.'

* updated both code examples to use the correct import path and public run() method.

* feat(crewai-tools): add db2 search tool

* refactor(crewai-tools): improve db2 search tool implementation

* feat(tools): improve DB2VectorSearchTool validation, security, and configurability

* docs: add DB2SearchTool documentation

* feat: add DB2 search tool

* docs: update DB2SearchTool documentation

* fix: address CodeRabbit review feedback

* fix: validate non-empty filter_by in DB2ToolSchema

* chore: trigger CodeRabbit re-review

* feat: fortify DB2 tool; fixed JSON response shape, added input guards and config validation

* fix(db2): address ruff and mypy linter errors

* style(db2): apply ruff format to db2_search_tool.py

* fix(db2-search-tool): address PR review comments

- Restore DirectoryReadTool export accidentally removed; add DB2VectorSearchTool
  and DB2ToolSchema to crewai_tools.tools __init__ and __all__
- Align _ALLOWED_METRICS whitelist with Db2 VECTOR_DISTANCE API:
  replace DOT_PRODUCT/L2_DISTANCE with EUCLIDEAN_SQUARED/DOT/HAMMING/MANHATTAN
- Replace ImportString fields for db2_package/db2_dbi_package with plain Any +
  lazy importlib.import_module in new _resolve_db2_packages() to avoid Pydantic
  default-validation gap where strings were never resolved at construction time
- Move docs from frozen docs/v1.13.0/ snapshot to docs/edge/en/tools/database-data/
  and register in docs/docs.json; update examples to match actual API
  (connection_string constructor, not DB2Config), correct return format, and
  align documented distance metrics with the whitelist

* fix(db2-search-tool): resolve default and string db2 package imports dynamically

* fix(db2-search-tool): export DB2VectorSearchTool and DB2ToolSchema from package-level crewai_tools

* docs(db2-search-tool): fix installation command and import path in README

---------

Co-authored-by: priyanshu-krishnan1 <priyanshu.krishnan1@ibm.com>
Co-authored-by: GeetikaChugh24 <geetika@ibm.com>
Co-authored-by: Lorenze Jay <63378463+lorenzejay@users.noreply.github.com>
Co-authored-by: Dhruv Chaturvedi <dhruv_insights@Dhruvs-MacBook-Pro.local>
2026-07-31 09:13:33 -07:00
8 changed files with 1398 additions and 1 deletions

View File

@@ -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"
]
},
{

View 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 (1100). |
| `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

View File

@@ -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",

View File

@@ -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",

View File

@@ -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.

View File

@@ -0,0 +1,10 @@
from crewai_tools.tools.db2_search_tool.db2_search_tool import (
DB2ToolSchema,
DB2VectorSearchTool,
)
__all__ = [
"DB2ToolSchema",
"DB2VectorSearchTool",
]

View File

@@ -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()

View 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