diff --git a/docs/docs.json b/docs/docs.json index a764cdef4..656e2d270 100644 --- a/docs/docs.json +++ b/docs/docs.json @@ -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" ] }, { diff --git a/docs/edge/en/tools/database-data/db2searchtool.mdx b/docs/edge/en/tools/database-data/db2searchtool.mdx new file mode 100644 index 000000000..5d1c06f60 --- /dev/null +++ b/docs/edge/en/tools/database-data/db2searchtool.mdx @@ -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 diff --git a/lib/crewai-tools/src/crewai_tools/__init__.py b/lib/crewai-tools/src/crewai_tools/__init__.py index 4f272014a..98fe3fc2b 100644 --- a/lib/crewai-tools/src/crewai_tools/__init__.py +++ b/lib/crewai-tools/src/crewai_tools/__init__.py @@ -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", diff --git a/lib/crewai-tools/src/crewai_tools/tools/__init__.py b/lib/crewai-tools/src/crewai_tools/tools/__init__.py index dfab952b9..6337c88ec 100644 --- a/lib/crewai-tools/src/crewai_tools/tools/__init__.py +++ b/lib/crewai-tools/src/crewai_tools/tools/__init__.py @@ -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", diff --git a/lib/crewai-tools/src/crewai_tools/tools/db2_search_tool/README.md b/lib/crewai-tools/src/crewai_tools/tools/db2_search_tool/README.md new file mode 100644 index 000000000..cbd4ed907 --- /dev/null +++ b/lib/crewai-tools/src/crewai_tools/tools/db2_search_tool/README.md @@ -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. \ No newline at end of file diff --git a/lib/crewai-tools/src/crewai_tools/tools/db2_search_tool/__init__.py b/lib/crewai-tools/src/crewai_tools/tools/db2_search_tool/__init__.py new file mode 100644 index 000000000..81287f6ef --- /dev/null +++ b/lib/crewai-tools/src/crewai_tools/tools/db2_search_tool/__init__.py @@ -0,0 +1,10 @@ +from crewai_tools.tools.db2_search_tool.db2_search_tool import ( + DB2ToolSchema, + DB2VectorSearchTool, +) + + +__all__ = [ + "DB2ToolSchema", + "DB2VectorSearchTool", +] diff --git a/lib/crewai-tools/src/crewai_tools/tools/db2_search_tool/db2_search_tool.py b/lib/crewai-tools/src/crewai_tools/tools/db2_search_tool/db2_search_tool.py new file mode 100644 index 000000000..9374fb269 --- /dev/null +++ b/lib/crewai-tools/src/crewai_tools/tools/db2_search_tool/db2_search_tool.py @@ -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 "" + 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() diff --git a/lib/crewai-tools/tests/tools/test_db2_search_tool.py b/lib/crewai-tools/tests/tools/test_db2_search_tool.py new file mode 100644 index 000000000..4fcaae392 --- /dev/null +++ b/lib/crewai-tools/tests/tools/test_db2_search_tool.py @@ -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 "" 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