diff --git a/pyproject.toml b/pyproject.toml
index 5fd3280..4cf5ab0 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -108,6 +108,7 @@ convention = "google"
"tests/*" = [
# Allow use of assert statements in tests
"S101",
+ "PLR2004",
]
[tool.ruff.format]
diff --git a/src/classifai/vectorisers/base.py b/src/classifai/vectorisers/base.py
index c92266f..7b26767 100644
--- a/src/classifai/vectorisers/base.py
+++ b/src/classifai/vectorisers/base.py
@@ -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
+ ...
diff --git a/tests/readme.md b/tests/readme.md
new file mode 100644
index 0000000..4ea6257
--- /dev/null
+++ b/tests/readme.md
@@ -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.
\ No newline at end of file
diff --git a/tests/test_vectorisers/readme.md b/tests/test_vectorisers/readme.md
new file mode 100644
index 0000000..937ddcc
--- /dev/null
+++ b/tests/test_vectorisers/readme.md
@@ -0,0 +1,72 @@
+# Unit Testing Plan for Vectorisers Module
+## Test Structure Overview
+
+1. Base Class Tests (test_vectoriser_base.py)
+* Verify VectoriserBase is abstract and cannot be instantiated
+* Verify transform method is abstract
+* Test that subclasses must implement transform
+* Test generic subclass behaviours
+
+
+2. HuggingFaceVectoriser Tests(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
+
+
+3. GcpVectoriser Tests (test_gcp_vectoriser.py)
+* 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
+
+
+4. OllamaVectoriser Tests (test_ollama_vectoriser.py)
+* 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 |
diff --git a/tests/test_vectorisers/test_vectoriser_base.py b/tests/test_vectorisers/test_vectoriser_base.py
new file mode 100644
index 0000000..faecd10
--- /dev/null
+++ b/tests/test_vectorisers/test_vectoriser_base.py
@@ -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
diff --git a/tests/test_vectorisers/test_vectoriser_gcp.py b/tests/test_vectorisers/test_vectoriser_gcp.py
new file mode 100644
index 0000000..66e2b20
--- /dev/null
+++ b/tests/test_vectorisers/test_vectoriser_gcp.py
@@ -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
diff --git a/tests/test_vectorisers/test_vectoriser_huggingface.py b/tests/test_vectorisers/test_vectoriser_huggingface.py
new file mode 100644
index 0000000..1ee89c8
--- /dev/null
+++ b/tests/test_vectorisers/test_vectoriser_huggingface.py
@@ -0,0 +1,256 @@
+"""Unit tests for HuggingFaceVectoriser."""
+
+from __future__ import annotations
+
+from unittest.mock import Mock, patch
+
+import pytest
+import torch
+
+from classifai.exceptions import ConfigurationError, ExternalServiceError, VectorisationError
+from classifai.vectorisers import HuggingFaceVectoriser
+
+
+def make_fake_tokenizer_output(input_ids, attention_mask):
+ """Build a fake tokenizer output supporting dict access and .to(device)."""
+
+ class FakeBatchEncoding(dict):
+ def to(self, device):
+ return self
+
+ return FakeBatchEncoding(
+ {
+ "input_ids": torch.tensor(input_ids),
+ "attention_mask": torch.tensor(attention_mask),
+ }
+ )
+
+
+class TestHuggingFaceVectoriserInitialization:
+ """Tests for HuggingFaceVectoriser initialization."""
+
+ @patch("classifai.vectorisers.huggingface.check_deps")
+ def test_init_missing_dependencies_raises_error(self, mock_check_deps):
+ """Missing torch/transformers should raise the error surfaced by check_deps."""
+ mock_check_deps.side_effect = ImportError("torch/transformers not installed")
+
+ with pytest.raises(ImportError):
+ HuggingFaceVectoriser(model_name="bert-base-uncased")
+
+ mock_check_deps.assert_called_once_with(["transformers", "torch"], extra="huggingface")
+
+ @patch("classifai.vectorisers.huggingface.check_deps")
+ @patch("transformers.AutoModel")
+ @patch("transformers.AutoTokenizer")
+ def test_init_valid_model_loads_successfully(self, mock_autotokenizer, mock_automodel, mock_check_deps):
+ """A valid model name should load tokenizer and model without raising."""
+ vectoriser = HuggingFaceVectoriser(model_name="bert-base-uncased")
+
+ assert vectoriser.model_name == "bert-base-uncased"
+ assert vectoriser.tokenizer is mock_autotokenizer.from_pretrained.return_value
+ assert vectoriser.model is mock_automodel.from_pretrained.return_value
+
+ @patch("classifai.vectorisers.huggingface.check_deps")
+ @patch("transformers.AutoModel")
+ @patch("transformers.AutoTokenizer")
+ def test_init_invalid_model_name_raises_external_service_error(
+ self, mock_autotokenizer, mock_automodel, mock_check_deps
+ ):
+ """A failure loading the tokenizer/model should raise ExternalServiceError."""
+ mock_autotokenizer.from_pretrained.side_effect = OSError("model not found")
+
+ with pytest.raises(ExternalServiceError):
+ HuggingFaceVectoriser(model_name="not-a-real-model")
+
+ @patch("classifai.vectorisers.huggingface.check_deps")
+ @patch("transformers.AutoModel")
+ @patch("transformers.AutoTokenizer")
+ def test_init_device_defaults_to_cpu_when_cuda_unavailable(
+ self, mock_autotokenizer, mock_automodel, mock_check_deps
+ ):
+ """When no device is specified and CUDA is unavailable, device should default to cpu."""
+ with patch("torch.cuda.is_available", return_value=False):
+ vectoriser = HuggingFaceVectoriser(model_name="bert-base-uncased")
+
+ assert vectoriser.device == torch.device("cpu")
+
+ @patch("classifai.vectorisers.huggingface.check_deps")
+ @patch("transformers.AutoModel")
+ @patch("transformers.AutoTokenizer")
+ def test_init_device_defaults_to_cuda_when_available(self, mock_autotokenizer, mock_automodel, mock_check_deps):
+ """When no device is specified and CUDA is available, device should default to cuda."""
+ with patch("torch.cuda.is_available", return_value=True):
+ vectoriser = HuggingFaceVectoriser(model_name="bert-base-uncased")
+
+ assert vectoriser.device == torch.device("cuda")
+
+ @patch("classifai.vectorisers.huggingface.check_deps")
+ @patch("transformers.AutoModel")
+ @patch("transformers.AutoTokenizer")
+ def test_init_explicit_device_is_respected(self, mock_autotokenizer, mock_automodel, mock_check_deps):
+ """An explicitly passed device should be used as-is."""
+ vectoriser = HuggingFaceVectoriser(model_name="bert-base-uncased", device="cpu")
+
+ assert vectoriser.device == "cpu"
+
+ @patch("classifai.vectorisers.huggingface.check_deps")
+ @patch("transformers.AutoModel")
+ @patch("transformers.AutoTokenizer")
+ def test_init_bad_device_raises_configuration_error(self, mock_autotokenizer, mock_automodel, mock_check_deps):
+ """If placing the model on the device fails, ConfigurationError should be raised."""
+ mock_automodel.from_pretrained.return_value.to.side_effect = RuntimeError("invalid device")
+
+ with pytest.raises(ConfigurationError):
+ HuggingFaceVectoriser(model_name="bert-base-uncased", device="bad-device")
+
+ @patch("classifai.vectorisers.huggingface.check_deps")
+ @patch("transformers.AutoModel")
+ @patch("transformers.AutoTokenizer")
+ def test_init_trust_remote_code_defaults_to_false(self, mock_autotokenizer, mock_automodel, mock_check_deps):
+ """trust_remote_code should default to False for both tokenizer and model kwargs."""
+ HuggingFaceVectoriser(model_name="bert-base-uncased")
+
+ _, tok_kwargs = mock_autotokenizer.from_pretrained.call_args
+ _, model_kwargs = mock_automodel.from_pretrained.call_args
+
+ assert tok_kwargs["trust_remote_code"] is False
+ assert model_kwargs["trust_remote_code"] is False
+
+ @patch("classifai.vectorisers.huggingface.check_deps")
+ @patch("transformers.AutoModel")
+ @patch("transformers.AutoTokenizer")
+ def test_init_custom_tokenizer_kwargs_passed_through(self, mock_autotokenizer, mock_automodel, mock_check_deps):
+ """Custom tokenizer_kwargs should be forwarded to AutoTokenizer.from_pretrained."""
+ HuggingFaceVectoriser(
+ model_name="bert-base-uncased",
+ tokenizer_kwargs={"use_fast": False},
+ )
+
+ _, call_kwargs = mock_autotokenizer.from_pretrained.call_args
+ assert call_kwargs["use_fast"] is False
+
+ @patch("classifai.vectorisers.huggingface.check_deps")
+ @patch("transformers.AutoModel")
+ @patch("transformers.AutoTokenizer")
+ def test_init_custom_model_kwargs_passed_through(self, mock_autotokenizer, mock_automodel, mock_check_deps):
+ """Custom model_kwargs should be forwarded to AutoModel.from_pretrained."""
+ HuggingFaceVectoriser(
+ model_name="bert-base-uncased",
+ model_kwargs={"trust_remote_code": True},
+ )
+
+ _, call_kwargs = mock_automodel.from_pretrained.call_args
+ assert call_kwargs["trust_remote_code"] is True
+
+ @patch("classifai.vectorisers.huggingface.check_deps")
+ @patch("transformers.AutoModel")
+ @patch("transformers.AutoTokenizer")
+ def test_init_model_revision_passed_through(self, mock_autotokenizer, mock_automodel, mock_check_deps):
+ """model_revision should be forwarded to both from_pretrained calls."""
+ HuggingFaceVectoriser(model_name="bert-base-uncased", model_revision="v2")
+
+ _, tok_kwargs = mock_autotokenizer.from_pretrained.call_args
+ _, model_kwargs = mock_automodel.from_pretrained.call_args
+
+ assert tok_kwargs["revision"] == "v2"
+ assert model_kwargs["revision"] == "v2"
+
+
+class TestHuggingFaceVectoriserTransform:
+ """Tests for HuggingFaceVectoriser transform method."""
+
+ @pytest.fixture
+ def mock_vectoriser(self):
+ """Return a HuggingFaceVectoriser instance with mocked tokenizer/model."""
+ with (
+ patch("classifai.vectorisers.huggingface.check_deps"),
+ patch("transformers.AutoModel"),
+ patch("transformers.AutoTokenizer"),
+ ):
+ vectoriser = HuggingFaceVectoriser(model_name="bert-base-uncased", device="cpu")
+
+ # Replace with controllable mocks for the transform tests.
+ vectoriser.tokenizer = Mock()
+ vectoriser.model = Mock()
+
+ yield vectoriser
+
+ def _configure_successful_forward_pass(self, vectoriser, batch_size, seq_len, hidden_dim):
+ """Wire up tokenizer + model mocks to produce a real small tensor output."""
+ vectoriser.tokenizer.return_value = make_fake_tokenizer_output(
+ input_ids=[[1] * seq_len] * batch_size,
+ attention_mask=[[1] * seq_len] * batch_size,
+ )
+
+ fake_model_output = Mock()
+ fake_model_output.last_hidden_state = torch.randn(batch_size, seq_len, hidden_dim)
+ vectoriser.model.return_value = fake_model_output
+
+ def test_transform_single_string_converts_to_list_and_processes(self, mock_vectoriser):
+ """A single string input should be wrapped in a list and produce one embedding row."""
+ self._configure_successful_forward_pass(mock_vectoriser, batch_size=1, seq_len=5, hidden_dim=8)
+
+ result = mock_vectoriser.transform("hello world")
+
+ assert result.shape[0] == 1
+ call_args, _ = mock_vectoriser.tokenizer.call_args
+ assert call_args[0] == ["hello world"]
+
+ def test_transform_list_of_strings_processes_correctly(self, mock_vectoriser):
+ """A list of strings should be passed through unchanged."""
+ self._configure_successful_forward_pass(mock_vectoriser, batch_size=3, seq_len=6, hidden_dim=8)
+
+ mock_vectoriser.transform(["a", "b", "c"])
+
+ call_args, _ = mock_vectoriser.tokenizer.call_args
+ assert call_args[0] == ["a", "b", "c"]
+
+ def test_transform_returns_2d_numpy_array(self, mock_vectoriser):
+ """Output should be a 2D numpy array."""
+ self._configure_successful_forward_pass(mock_vectoriser, batch_size=2, seq_len=6, hidden_dim=8)
+
+ result = mock_vectoriser.transform(["a", "b"])
+
+ 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."""
+ self._configure_successful_forward_pass(mock_vectoriser, batch_size=4, seq_len=6, hidden_dim=8)
+
+ result = mock_vectoriser.transform(["a", "b", "c", "d"])
+
+ assert result.shape[0] == 4
+
+ def test_transform_tokenisation_failure_raises_vectorisation_error(self, mock_vectoriser):
+ """If tokenization raises, it should be wrapped in VectorisationError."""
+ mock_vectoriser.tokenizer.side_effect = RuntimeError("bad tokenizer input")
+
+ with pytest.raises(VectorisationError):
+ mock_vectoriser.transform(["hello"])
+
+ def test_transform_model_inference_failure_raises_vectorisation_error(self, mock_vectoriser):
+ """If the model forward pass raises, it should be wrapped in VectorisationError."""
+ mock_vectoriser.tokenizer.return_value = make_fake_tokenizer_output(
+ input_ids=[[1, 2, 3]],
+ attention_mask=[[1, 1, 1]],
+ )
+ mock_vectoriser.model.side_effect = RuntimeError("model forward failed")
+
+ with pytest.raises(VectorisationError):
+ mock_vectoriser.transform(["hello"])
+
+ def test_transform_pooling_failure_raises_vectorisation_error(self, mock_vectoriser):
+ """If pooling fails (e.g. shape mismatch), it should be wrapped in VectorisationError."""
+ mock_vectoriser.tokenizer.return_value = make_fake_tokenizer_output(
+ input_ids=[[1, 2, 3]],
+ attention_mask=[[1, 1, 1]],
+ )
+
+ # Mismatched shape between hidden_state (seq_len=5) and attention_mask (seq_len=3)
+ # forces a broadcasting error inside the pooling try-block.
+ fake_model_output = Mock()
+ fake_model_output.last_hidden_state = torch.randn(1, 5, 8)
+ mock_vectoriser.model.return_value = fake_model_output
+
+ with pytest.raises(VectorisationError):
+ mock_vectoriser.transform(["hello"])
diff --git a/tests/test_vectorisers/test_vectoriser_ollama.py b/tests/test_vectorisers/test_vectoriser_ollama.py
new file mode 100644
index 0000000..21aa1c0
--- /dev/null
+++ b/tests/test_vectorisers/test_vectoriser_ollama.py
@@ -0,0 +1,116 @@
+"""Unit tests for OllamaVectoriser."""
+
+from __future__ import annotations
+
+from unittest.mock import Mock, patch
+
+import numpy as np
+import pytest
+
+from classifai.exceptions import ExternalServiceError, VectorisationError
+from classifai.vectorisers import OllamaVectoriser
+
+
+class TestOllamaVectoriserInitialization:
+ """Tests for OllamaVectoriser initialization."""
+
+ @patch("classifai.vectorisers.ollama.check_deps")
+ def test_init_missing_dependencies_raises_error(self, mock_check_deps):
+ """Missing ollama package should raise the error surfaced by check_deps."""
+ mock_check_deps.side_effect = ImportError("ollama not installed")
+
+ with pytest.raises(ImportError):
+ OllamaVectoriser(model_name="nomic-embed-text")
+
+ mock_check_deps.assert_called_once_with(["ollama"], extra="ollama")
+
+ @patch("classifai.vectorisers.ollama.check_deps")
+ def test_init_model_name_stored_correctly(self, mock_check_deps):
+ """model_name should be stored as an attribute."""
+ vectoriser = OllamaVectoriser(model_name="nomic-embed-text")
+
+ assert vectoriser.model_name == "nomic-embed-text"
+
+
+class TestOllamaVectoriserTransform:
+ """Tests for OllamaVectoriser transform method."""
+
+ @pytest.fixture
+ def vectoriser(self):
+ """Return an OllamaVectoriser instance with check_deps mocked out."""
+ with patch("classifai.vectorisers.ollama.check_deps"):
+ return OllamaVectoriser(model_name="nomic-embed-text")
+
+ def _configure_successful_embed_response(self, mock_embed, n_texts, dim=4):
+ """Wire up ollama.embed to return a fake embeddings response."""
+ fake_response = Mock()
+ fake_response.embeddings = [[0.1] * dim for _ in range(n_texts)]
+ mock_embed.return_value = fake_response
+
+ @patch("ollama.embed")
+ def test_transform_single_string_converts_to_list(self, mock_embed, vectoriser):
+ """A single string input should be wrapped in a list before the API call."""
+ self._configure_successful_embed_response(mock_embed, n_texts=1)
+
+ vectoriser.transform("hello world")
+
+ _, call_kwargs = mock_embed.call_args
+ assert call_kwargs["input"] == ["hello world"]
+
+ @patch("ollama.embed")
+ def test_transform_list_processes_correctly(self, mock_embed, vectoriser):
+ """A list of strings should be passed through unchanged."""
+ texts = ["text1", "text2", "text3"]
+ self._configure_successful_embed_response(mock_embed, n_texts=len(texts))
+
+ vectoriser.transform(texts)
+
+ _, call_kwargs = mock_embed.call_args
+ assert call_kwargs["input"] == texts
+
+ @patch("ollama.embed")
+ def test_transform_returns_2d_numpy_array(self, mock_embed, vectoriser):
+ """Output should be a 2D numpy array."""
+ self._configure_successful_embed_response(mock_embed, n_texts=2)
+
+ result = vectoriser.transform(["a", "b"])
+
+ assert isinstance(result, np.ndarray)
+ assert result.ndim == 2
+
+ @patch("ollama.embed")
+ def test_transform_output_shape_matches_input_count(self, mock_embed, vectoriser):
+ """Number of output rows should match number of input texts."""
+ texts = ["a", "b", "c", "d"]
+ self._configure_successful_embed_response(mock_embed, n_texts=len(texts))
+
+ result = vectoriser.transform(texts)
+
+ assert result.shape[0] == len(texts)
+
+ @patch("ollama.embed")
+ def test_transform_model_name_passed_to_embed_call(self, mock_embed, vectoriser):
+ """The stored model_name should be passed to ollama.embed."""
+ self._configure_successful_embed_response(mock_embed, n_texts=1)
+
+ vectoriser.transform(["hello"])
+
+ _, call_kwargs = mock_embed.call_args
+ assert call_kwargs["model"] == vectoriser.model_name
+
+ @patch("ollama.embed")
+ def test_transform_service_failure_raises_external_service_error(self, mock_embed, vectoriser):
+ """If ollama.embed itself raises, it should be wrapped in ExternalServiceError."""
+ mock_embed.side_effect = Exception("connection refused")
+
+ with pytest.raises(ExternalServiceError):
+ vectoriser.transform(["hello"])
+
+ @patch("ollama.embed")
+ def test_transform_response_parsing_failure_raises_vectorisation_error(self, mock_embed, vectoriser):
+ """If extracting/converting .embeddings fails, it should be wrapped in VectorisationError."""
+ fake_response = Mock(spec=[]) # no .embeddings attribute at all
+ mock_embed.return_value = fake_response
+
+ with pytest.raises(VectorisationError):
+ vectoriser.transform(["hello"])