Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -108,6 +108,7 @@ convention = "google"
"tests/*" = [
# Allow use of assert statements in tests
"S101",
"PLR2004",
]

[tool.ruff.format]
Expand Down
2 changes: 1 addition & 1 deletion src/classifai/vectorisers/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -80,4 +80,4 @@ def transform(self, texts: str | list[str]) -> np.ndarray:
text(s). Each row corresponds to the embedding of a single input
text.
"""
pass
...
9 changes: 9 additions & 0 deletions tests/readme.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,9 @@
A testing plan including:

- a plan for implementing concrete unit tests in each of ClassifAI's 4 modules - describing which componets to test,
- table of key concerns for testing each module,
- a corresponding first pass example of unit tests generated by Claude to test a specific class or feature detailed in the module unit testing plan.

The above list describes content that we can use to implement sets of unit tests for each module, testing each independently of the other modules. Additional later tests could include end-to-end integration tests where we use multiple modules together (vectoriiser + VectorStore for example).

Additionally, a `test_exports.py` test script within the parent folder. The purpose of this test is to ensure that each of the modules from the package that should be importable to a user are succsefully exported by the package. This ensure thats the the base API for the package is accessible and functioning.
72 changes: 72 additions & 0 deletions tests/test_vectorisers/readme.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,72 @@
# Unit Testing Plan for Vectorisers Module
## Test Structure Overview

<b>1. Base Class Tests (test_vectoriser_base.py)</b>
* Verify VectoriserBase is abstract and cannot be instantiated
* Verify transform method is abstract
* Test that subclasses must implement transform
* Test generic subclass behaviours


<b>2. HuggingFaceVectoriser Tests</b>(test_huggingface_vectoriser.py)
* Initialisation:
* Missing dependencies raise appropriate errors
* Valid model loads successfully
* Invalid model name raises ExternalServiceError
* Device selection (CPU/GPU) works correctly
* Bad device selection raises ConfigurationError
* trust_remote_code defaults to False
* Custom kwargs are passed through
* Transform method:
* Single string input converts to list and processes
* List of strings processes correctly
* Returns 2D numpy array
* Output shape matches input count
* Tokenisation failures raise VectorisationError
* Model inference failures raise VectorisationError
* Pooling failures raise VectorisationError


<b>3. GcpVectoriser Tests (test_gcp_vectoriser.py)</b>
* Initialization:
* Missing dependencies raise appropriate errors
* project_id + location authentication works
* api_key authentication works
* Missing both auth methods raises ConfigurationError
* Providing both auth methods raises ConfigurationError
* Client initialisation failures raise ConfigurationError
* Test model name and task type stored correctly
* Transform method:
* Single string input converts to list
* List processes correctly
* Returns 2D numpy array
* Output shape matches input count
* API request failures raise ExternalServiceError
* Unexpected response format raises VectorisationError
* Model named passed to call correctly


<b>4. OllamaVectoriser Tests (test_ollama_vectoriser.py)</b>
* Initialization:
* Missing dependencies raise appropriate errors
* Model name is stored correctly
* Transform method:
* Single string input converts to list
* List processes correctly
* Returns 2D numpy array
* Service failures raise ExternalServiceError
* Response parsing failures raise VectorisationError



## Key Testing Considerations


| **Aspect** | **Strategy** |
|------------------------|-----------------------------------------------------------------------------|
| External Dependencies | Use `pytest-mock` or `unittest.mock` to patch external libraries (torch, transformers, ollama, google.genai) |
| GPU/Device Testing | Mock `torch.cuda` to test both CPU and GPU branches if we are concerned with GPU compatibility |
| API Responses | Mock service responses with realistic embedding data |
| Error Cases | Test each exception path in try-except blocks |
| Input Validation | Test both string and list inputs |
| Output Validation | Verify numpy array shape, dtype, and content |
73 changes: 73 additions & 0 deletions tests/test_vectorisers/test_vectoriser_base.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,73 @@
"""Unit tests for VectoriserBase abstract class."""

from __future__ import annotations

import numpy as np
import pytest

from classifai.vectorisers.base import VectoriserBase


class IncompleteVectoriser(VectoriserBase):
"""Subclass that does NOT implement transform - should fail to instantiate."""

pass


class ConcreteVectoriser(VectoriserBase):
"""Subclass that DOES implement transform - should instantiate fine."""

def transform(self, texts):
if isinstance(texts, str):
texts = [texts]
return np.zeros((len(texts), 4))


class TestVectoriserBaseIsAbstract:
"""Tests that VectoriserBase enforces its abstract interface correctly."""

def test_cannot_instantiate_base_class_directly(self):
"""VectoriserBase should raise TypeError on direct instantiation."""
with pytest.raises(TypeError):
VectoriserBase()

def test_transform_is_registered_as_abstract_method(self):
"""Transform should be listed in __abstractmethods__."""
assert "transform" in VectoriserBase.__abstractmethods__

def test_subclass_missing_transform_cannot_instantiate(self):
"""A subclass that doesn't implement transform should also fail to instantiate."""
with pytest.raises(TypeError):
IncompleteVectoriser()

def test_concrete_subclass_can_be_instantiated(self):
"""A subclass implementing transform should instantiate successfully."""
vectoriser = ConcreteVectoriser()
assert isinstance(vectoriser, VectoriserBase)


class TestVectoriserBaseSubclassBehaviour:
"""Sanity checks on how a concrete subclass should behave."""

def test_concrete_subclass_transform_accepts_single_string(self):
"""Transform should accept a bare string input."""
vectoriser = ConcreteVectoriser()
result = vectoriser.transform("hello")

assert isinstance(result, np.ndarray)
assert result.shape[0] == 1

def test_concrete_subclass_transform_accepts_list_of_strings(self):
"""Transform should accept a list of strings input."""
vectoriser = ConcreteVectoriser()
result = vectoriser.transform(["hello", "world"])

assert isinstance(result, np.ndarray)
assert result.shape[0] == 2

def test_concrete_subclass_transform_returns_2d_array(self):
"""The returned array should always be 2-dimensional."""
vectoriser = ConcreteVectoriser()
result = vectoriser.transform(["a", "b", "c"])

assert result.ndim == 2
169 changes: 169 additions & 0 deletions tests/test_vectorisers/test_vectoriser_gcp.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,169 @@
"""Unit tests for GcpVectoriser."""

from __future__ import annotations

from unittest.mock import Mock, patch

import numpy as np
import pytest

from classifai.exceptions import ConfigurationError, ExternalServiceError, VectorisationError
from classifai.vectorisers import GcpVectoriser


class TestGcpVectoriserInitialization:
"""Tests for GcpVectoriser initialization."""

@patch("classifai.vectorisers.gcp.check_deps")
def test_init_missing_dependencies_raises_error(self, mock_check_deps):
"""Missing google-genai should raise the error surfaced by check_deps."""
mock_check_deps.side_effect = ImportError("google-genai not installed")

with pytest.raises(ImportError):
GcpVectoriser(project_id="my-project", location="europe-west2")

mock_check_deps.assert_called_once_with(["google-genai"], extra="gcp")

@patch("classifai.vectorisers.gcp.check_deps")
@patch("google.genai.Client")
def test_init_project_id_and_location_authentication_works(self, mock_client, mock_check_deps):
"""project_id + location should be forwarded to the client constructor."""
vectoriser = GcpVectoriser(project_id="my-project", location="europe-west2")

_, call_kwargs = mock_client.call_args
assert call_kwargs["project"] == "my-project"
assert call_kwargs["location"] == "europe-west2"
assert vectoriser.vectoriser is mock_client.return_value

@patch("classifai.vectorisers.gcp.check_deps")
@patch("google.genai.Client")
def test_init_api_key_authentication_works(self, mock_client, mock_check_deps):
"""api_key alone should be forwarded to the client constructor."""
vectoriser = GcpVectoriser(api_key="fake-api-key")

_, call_kwargs = mock_client.call_args
assert call_kwargs["api_key"] == "fake-api-key" # pragma: allowlist secret
assert vectoriser.vectoriser is mock_client.return_value

@patch("classifai.vectorisers.gcp.check_deps")
def test_init_missing_both_auth_methods_raises_configuration_error(self, mock_check_deps):
"""Providing neither project_id/location nor api_key should raise ConfigurationError."""
with pytest.raises(ConfigurationError):
GcpVectoriser()

@patch("classifai.vectorisers.gcp.check_deps")
def test_init_providing_both_auth_methods_raises_configuration_error(self, mock_check_deps):
"""Providing both project_id and api_key should raise ConfigurationError."""
with pytest.raises(ConfigurationError):
GcpVectoriser(project_id="my-project", api_key="fake-api-key")

@patch("classifai.vectorisers.gcp.check_deps")
@patch("google.genai.Client")
def test_init_client_initialisation_failure_raises_configuration_error(self, mock_client, mock_check_deps):
"""If the underlying client constructor raises, wrap it in ConfigurationError."""
mock_client.side_effect = Exception("bad credentials")

with pytest.raises(ConfigurationError):
GcpVectoriser(project_id="my-project", location="europe-west2")

@patch("classifai.vectorisers.gcp.check_deps")
@patch("google.genai.Client")
def test_init_model_name_stored_correctly(self, mock_client, mock_check_deps):
"""model_name should be stored as an attribute."""
vectoriser = GcpVectoriser(api_key="fake-api-key", model_name="text-embedding-005")

assert vectoriser.model_name == "text-embedding-005"

@patch("classifai.vectorisers.gcp.check_deps")
@patch("google.genai.Client")
def test_init_task_type_passed_to_embed_content_config(self, mock_client, mock_check_deps):
"""task_type should be forwarded to EmbedContentConfig."""
with patch("google.genai.types.EmbedContentConfig") as mock_config:
GcpVectoriser(api_key="fake-api-key", task_type="RETRIEVAL_QUERY") # pragma: allowlist secret

_, call_kwargs = mock_config.call_args
assert call_kwargs["task_type"] == "RETRIEVAL_QUERY"


class TestGcpVectoriserTransform:
"""Tests for GcpVectoriser transform method."""

@pytest.fixture
def mock_vectoriser(self):
"""Return a GcpVectoriser instance with a mocked client."""
with (
patch("classifai.vectorisers.gcp.check_deps"),
patch("google.genai.Client"),
):
vectoriser = GcpVectoriser(api_key="fake-api-key")

# Replace with a controllable mock for the transform tests.
vectoriser.vectoriser = Mock()

yield vectoriser

def _configure_successful_embed_response(self, vectoriser, n_texts, dim=4):
"""Wire up the client mock to return a fake embeddings response."""
fake_response = Mock()
fake_response.embeddings = [Mock(values=[0.1] * dim) for _ in range(n_texts)]
vectoriser.vectoriser.models.embed_content.return_value = fake_response

def test_transform_single_string_converts_to_list(self, mock_vectoriser):
"""A single string input should be wrapped in a list before the API call."""
self._configure_successful_embed_response(mock_vectoriser, n_texts=1)

mock_vectoriser.transform("hello world")

_, call_kwargs = mock_vectoriser.vectoriser.models.embed_content.call_args
assert call_kwargs["contents"] == ["hello world"]

def test_transform_list_processes_correctly(self, mock_vectoriser):
"""A list of strings should be passed through unchanged."""
texts = ["text1", "text2", "text3"]
self._configure_successful_embed_response(mock_vectoriser, n_texts=len(texts))

mock_vectoriser.transform(texts)

_, call_kwargs = mock_vectoriser.vectoriser.models.embed_content.call_args
assert call_kwargs["contents"] == texts

def test_transform_returns_2d_numpy_array(self, mock_vectoriser):
"""Output should be a 2D numpy array."""
self._configure_successful_embed_response(mock_vectoriser, n_texts=2)

result = mock_vectoriser.transform(["a", "b"])

assert isinstance(result, np.ndarray)
assert result.ndim == 2

def test_transform_output_shape_matches_input_count(self, mock_vectoriser):
"""Number of output rows should match number of input texts."""
texts = ["a", "b", "c", "d"]
self._configure_successful_embed_response(mock_vectoriser, n_texts=len(texts))

result = mock_vectoriser.transform(texts)

assert result.shape[0] == len(texts)

def test_transform_api_request_failure_raises_external_service_error(self, mock_vectoriser):
"""If the API call itself raises, it should be wrapped in ExternalServiceError."""
mock_vectoriser.vectoriser.models.embed_content.side_effect = Exception("network error")

with pytest.raises(ExternalServiceError):
mock_vectoriser.transform(["hello"])

def test_transform_unexpected_response_format_raises_vectorisation_error(self, mock_vectoriser):
"""If the response doesn't have the expected .embeddings attribute, raise VectorisationError."""
mock_vectoriser.vectoriser.models.embed_content.return_value = Mock(spec=[]) # no .embeddings attribute

with pytest.raises(VectorisationError):
mock_vectoriser.transform(["hello"])

def test_transform_model_name_passed_to_api_call(self, mock_vectoriser):
"""The stored model_name should be passed to embed_content."""
self._configure_successful_embed_response(mock_vectoriser, n_texts=1)

mock_vectoriser.transform(["hello"])

_, call_kwargs = mock_vectoriser.vectoriser.models.embed_content.call_args
assert call_kwargs["model"] == mock_vectoriser.model_name
Loading
Loading