diff --git a/pyproject.toml b/pyproject.toml
index 4cf5ab0..faf3341 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -106,9 +106,10 @@ convention = "google"
[tool.ruff.lint.per-file-ignores]
"tests/*" = [
- # Allow use of assert statements in tests
+ # Allow use of assert statements, magic numbers and unused variables in tests
"S101",
"PLR2004",
+ "F841"
]
[tool.ruff.format]
diff --git a/tests/test_indexers/readme.md b/tests/test_indexers/readme.md
new file mode 100644
index 0000000..f428b89
--- /dev/null
+++ b/tests/test_indexers/readme.md
@@ -0,0 +1,183 @@
+# Unit Testing Plan for Indexers Module
+## Test Structure Overview
+1. Dataclass Tests (test_indexers_dataclasses.py)
+* VectorStoreSearchInput:
+ * Valid dict/DataFrame converts and validates correctly
+ * Schema validation enforces column types
+ * Missing required columns raises validation error
+ * Type coercion works (strings, etc.)
+ * Property accessors work (id, query) / return correct series
+ * Empty inputs handled correctly
+* VectorStoreSearchOutput:
+ * Valid construction from dict/DataFrame
+ * Schema validation enforces column types
+ * Missing required columns raises validation error
+ * Rank column must be non-negative
+ * Score column accepts floats
+ * Property accessors work / return correct series
+ * Column ordering is preserved (queries broadcast down consecutive rows)
+ * Empty inputs handled correctly
+* VectorStoreEmbedInput/Output:
+ * Valid dict/DataFrame construction and validation
+ * Type coercion for id/text
+ * Embedding column accepts numpy arrays
+ * Empty Inputs/Output handled correctly
+* VectorStoreReverseSearchInput/Output:
+ * Valid construction from dicts/DataFrames
+ * Empty Input/Output DataFrame handles correctly
+ * Schema validation works
+ * Property accessors function properly
+ * Missing required columns raises validation error
+
+2. VectorStore Initialization Tests (test_vectorstore_init.py)
+* Input validation (DataValidationError):
+ * file_name must be non-empty string
+ * data_type validation (only "csv" supported)
+ * vectoriser must be VectoriserBase instance
+ * batch_size must be positive integer
+ * meta_data must be dict or None
+ * hooks must be dict or None
+ * output_dir must be string or None
+* File system handling (ConfigurationError):
+ * Input file must exist
+ * Output directory creation works
+ * overwrite flag prevents accidental overwrites
+ * gs:// paths require gcsfs (helpful error message)
+ * Invalid fsspec paths raise ConfigurationError
+* Index building (IndexBuildError):
+ * CSV file reads correctly
+ * UUID generation works
+ * Batch processing of embeddings
+ * Vectoriser failures wrapped appropriately
+ * Embeddings count matches batch size
+ * Metadata serialization to JSON
+ * Parquet file writing
+* skip_save flag:
+ * When True, no files written to disk
+ * When False, metadata.json and vectors.parquet created
+ * warning logged when output_dir set but skip_save=True
+
+3. VectorStore Search Tests (test_vectorstore_search.py)
+* Input validation (DataValidationError):
+ * query must be VectorStoreSearchInput
+ * n_results must be int >= 1
+ * batch_size must be int >= 1 or None
+ * Empty query raises error
+ * Vector store not initialized raises ConfigurationError
+* Search operation:
+ * Single query processes correctly
+ * Multiple queries in batch
+ * Similarity scores computed (dot-product)
+ * Top n_results returned per query
+ * Results ranked by score (descending)
+ * Output shape matches expected (n_queries * n_results rows)
+ * Metadata columns included in output
+ * Query batching with custom batch_size works
+* Error handling (VectorisationError/ClassifaiError):
+ * Query embedding failure
+ * Vectoriser.transform() exceptions wrapped
+ * Error context includes vectoriser class, batch info
+* Hooks integration:
+ * search_preprocess hook called before search
+ * search_postprocess hook called after search
+ * Hook failures raise HookError
+ * Multiple hooks in list processed in order
+ * Single hook converted to list automatically
+
+4. VectorStore Reverse Search Tests (test_vectorstore_reverse_search.py)
+* Input validation (DataValidationError):
+ * query must be VectorStoreReverseSearchInput
+ * max_n_results must be int >= 1 or -1
+ * Empty query raises error
+* Reverse search operation:
+ * Exact label matching works (default)
+ * Partial matching (prefix) when enabled
+ * max_n_results limits results per query
+ * max_n_results=-1 returns all matches
+ * Results include metadata columns
+ * Empty result sets handled (returns empty DataFrame with correct schema)
+ * Sorting by id and label works
+* Error handling:
+ * Vectoriser-independent (no embeddings needed)
+ * DataFrame join failures wrapped
+ * Error context includes max_n_results, query count
+* Hooks integration:
+ * reverse_search_preprocess hook calls before reverse search
+ * reverse_search_postprocess hook calls after reverse search
+ * Same checks and error handling as search
+
+5. VectorStore Embed Tests (test_vectorstore_embed.py)
+* Input validation (DataValidationError):
+ * query must be VectorStoreEmbedInput
+ * Invalid input type raises error
+* Embedding operation:
+ * Single text embeds correctly
+ * Multiple texts process correctly
+ * Output includes id, text, and embedding
+ * Embeddings are numpy arrays
+ * Output shape matches input count
+ * Vectoriser.transform() called with correct texts
+* Error handling (VectorisationError/ClassifaiError):
+ * Vectoriser failures wrapped with context
+ * Error includes vectoriser class, text count
+* Hooks integration:
+ * embed_preprocess hook called before embedding
+ * embed_postprocess hook called after embedding
+ * Same checks and error handling as search
+
+6. VectorStore Metadata Tests (test_vectorstore_metadata.py)
+* Metadata serialization (_save_metadata):
+ * JSON file created at correct path
+ * Contains all required fields (vectoriser_class, vector_shape, num_vectors, batch_size, created_at, meta_data)
+ * Type information preserved (str types → string names)
+ * Valid JSON format
+ * fsspec paths work (gs://, etc.)
+* Metadata loading (from_filespace):
+ * Metadata file read and parsed correctly
+ * Required keys validated
+ * Type deserialization works
+ * Backwards compatibility with v1.0.0 (missing batch_size)
+ * Default batch_size used when missing
+ * Warning logged for missing batch_size
+
+7. VectorStore from_filespace Tests (test_vectorstore_from_filespace.py)
+* Input validation (DataValidationError):
+ * folder_path must be non-empty string
+ * folder_path must be existing directory
+ * batch_size override must be int >= 1 or None
+ * hooks must be dict or None
+* File loading (IndexBuildError):
+ * metadata.json exists and valid
+ * vectors.parquet exists and valid
+ * Required columns present in parquet
+ * Parquet not empty
+ * Metadata can be deserialized
+* Configuration validation (ConfigurationError):
+ * Vectoriser class name matches metadata
+ * vectoriser must have callable .transform() method / inherit from base class
+ * fsspec paths (gs://) work with gcsfs
+ * Helpful error message when gcsfs missing
+* Instance construction:
+ * Instance created without calling init
+ * All attributes set correctly
+ * batch_size override works
+ * metadata.meta_data deserialized and set
+ * Vectoriser instance attached
+ * hooks parameter applied
+ * quiet_mode applied
+ * Instance is functional (can search/embed/reverse_search)
+
+
+## Key Testing Considerations
+| **Aspect** | **Strategy** |
+|--------------------------|-----------------------------------------------------------------------------------------------------------------------------------------------------------|
+| Vectoriser Mocking | Mock `VectoriserBase` to return predictable embeddings; separate vectoriser testing from vectorstore testing |
+| File System | Mock `fsspec` for local/remote paths; test real local paths in integration tests; `gs://` tests optional/skipped without `gcsfs` |
+| Dataclass Validation | Test both valid and invalid inputs; verify `pandera` schema enforcement; test type coercion |
+| Large Datasets | Use small synthetic CSVs (< 100 rows); mock large searches with synthetic embeddings to avoid slow tests |
+| Similarity Computation | Verify dot-product calculations; test edge cases (zero embeddings, identical embeddings, single query) |
+| Hook System | Mock hooks that modify input/output; test hook chains; verify error propagation; test that single hooks auto-convert to lists |
+| Error Context | Verify all exceptions include relevant context (vectoriser class, batch info, file paths) without exposing secrets |
+| Save/Load Cycle | Test round-trip (create → save → load); verify metadata preservation; test backwards compatibility with old metadata format |
+| Empty/Edge Cases | Empty query results, single document, single query, all identical embeddings, `max_n_results > available docs` |
+| Quiet Mode | Verify progress bars suppressed; verify logging levels adjusted; test both `True` and `False` paths |
diff --git a/tests/test_indexers/test_dataclasses.py b/tests/test_indexers/test_dataclasses.py
new file mode 100644
index 0000000..3ccca49
--- /dev/null
+++ b/tests/test_indexers/test_dataclasses.py
@@ -0,0 +1,503 @@
+"""Unit tests for VectorStore dataclasses."""
+
+from __future__ import annotations
+
+import numpy as np
+import pandas as pd
+import pandera as pa
+import pytest
+
+from classifai.indexers import (
+ VectorStoreEmbedInput,
+ VectorStoreEmbedOutput,
+ VectorStoreReverseSearchInput,
+ VectorStoreReverseSearchOutput,
+ VectorStoreSearchInput,
+ VectorStoreSearchOutput,
+)
+
+
+class TestVectorStoreSearchInput:
+ """Tests for VectorStoreSearchInput dataclass."""
+
+ def test_init_from_dict_valid_data(self):
+ """Valid dict with id and query columns should construct successfully."""
+ data = {"id": ["1", "2"], "query": ["hello", "world"]}
+ result = VectorStoreSearchInput(data)
+
+ assert isinstance(result, pd.DataFrame)
+ assert list(result.columns) == ["id", "query"]
+ assert len(result) == 2
+
+ def test_init_from_dataframe_valid_data(self):
+ """Valid DataFrame with id and query columns should construct successfully."""
+ df = pd.DataFrame({"id": ["1", "2"], "query": ["hello", "world"]})
+ result = VectorStoreSearchInput(df)
+
+ assert isinstance(result, pd.DataFrame)
+ assert list(result.columns) == ["id", "query"]
+ assert len(result) == 2
+
+ def test_init_missing_required_column_raises_schema_error(self):
+ """Missing 'query' column should raise SchemaError."""
+ data = {"id": ["1", "2"]}
+
+ with pytest.raises(pa.errors.SchemaError):
+ VectorStoreSearchInput(data)
+
+ def test_init_coerces_int_id_to_string(self):
+ """Integer id should be coerced to string due to coerce=True."""
+ data = {"id": [1, 2], "query": ["hello", "world"]}
+ result = VectorStoreSearchInput(data)
+
+ # Verify the values were coerced to strings
+ assert all(isinstance(x, str) for x in result["id"])
+ assert list(result["id"]) == ["1", "2"]
+
+ def test_property_id_returns_correct_series(self):
+ """The id property should return the 'id' column as a Series."""
+ data = {"id": ["1", "2"], "query": ["hello", "world"]}
+ result = VectorStoreSearchInput(data)
+
+ id_series = result.id
+ assert isinstance(id_series, pd.Series)
+ assert list(id_series) == ["1", "2"]
+
+ def test_property_query_returns_correct_series(self):
+ """The query property should return the 'query' column as a Series."""
+ data = {"id": ["1", "2"], "query": ["hello", "world"]}
+ result = VectorStoreSearchInput(data)
+
+ query_series = result.query
+ assert isinstance(query_series, pd.Series)
+ assert list(query_series) == ["hello", "world"]
+
+ def test_from_data_classmethod_valid_data(self):
+ """from_data classmethod should construct from dict or DataFrame."""
+ data = {"id": ["1", "2"], "query": ["hello", "world"]}
+ result = VectorStoreSearchInput.from_data(data)
+
+ assert isinstance(result, VectorStoreSearchInput)
+ assert len(result) == 2
+
+ def test_validate_classmethod_on_valid_data_returns_instance(self):
+ """Validate classmethod should return a VectorStoreSearchInput instance."""
+ df = pd.DataFrame({"id": ["1", "2"], "query": ["hello", "world"]})
+ result = VectorStoreSearchInput.validate(df)
+
+ assert isinstance(result, VectorStoreSearchInput)
+ assert len(result) == 2
+
+
+class TestVectorStoreSearchOutput:
+ """Tests for VectorStoreSearchOutput dataclass."""
+
+ def test_init_from_dict_valid_data(self):
+ """Valid dict with all 6 required columns should construct successfully."""
+ data = {
+ "query_id": ["1", "1"],
+ "query_text": ["what is AI?", "what is AI?"],
+ "doc_label": ["doc1", "doc2"],
+ "doc_text": [456, "Machine learning..."],
+ "rank": [0, 1],
+ "score": [0.95, 0.87],
+ }
+ result = VectorStoreSearchOutput(data)
+
+ assert isinstance(result, pd.DataFrame)
+ assert len(result) == 2
+ assert list(result.columns) == ["query_id", "query_text", "doc_label", "doc_text", "rank", "score"]
+
+ def test_init_missing_required_column_raises_schema_error(self):
+ """Missing 'score' column should raise SchemaError."""
+ data = {
+ "query_id": ["1"],
+ "query_text": ["query"],
+ "doc_label": ["doc1"],
+ "doc_text": ["text"],
+ "rank": [0],
+ }
+
+ with pytest.raises(pa.errors.SchemaError):
+ VectorStoreSearchOutput(data)
+
+ def test_init_rank_less_than_zero_raises_schema_error(self):
+ """Negative rank should raise SchemaError (rank >= 0 required)."""
+ data = {
+ "query_id": ["1"],
+ "query_text": ["query"],
+ "doc_label": ["doc1"],
+ "doc_text": ["text"],
+ "rank": [-1],
+ "score": [0.9],
+ }
+
+ with pytest.raises(pa.errors.SchemaError):
+ VectorStoreSearchOutput(data)
+
+ def test_init_columns_ordered_preserved_for_same_query(self):
+ """Multiple results for same query should be grouped consecutively."""
+ data = {
+ "query_id": ["1", "1", "2", "2"],
+ "query_text": ["q1", "q1", "q2", "q2"],
+ "doc_label": ["a", "b", "c", "d"],
+ "doc_text": ["t1", "t2", "t3", "t4"],
+ "rank": [0, 1, 0, 1],
+ "score": [0.9, 0.8, 0.85, 0.75],
+ }
+ result = VectorStoreSearchOutput(data)
+
+ # Verify grouping: all query_id="1" appear before query_id="2"
+ query_ids = result.query_id.tolist()
+ assert query_ids == ["1", "1", "2", "2"]
+
+ def test_property_query_id_returns_correct_series(self):
+ """The query_id property should return the 'query_id' column."""
+ data = {
+ "query_id": ["1", "2"],
+ "query_text": ["q1", "q2"],
+ "doc_label": ["a", "b"],
+ "doc_text": ["text1", "text2"],
+ "rank": [0, 0],
+ "score": [0.9, 0.8],
+ }
+ result = VectorStoreSearchOutput(data)
+
+ query_id_series = result.query_id
+ assert isinstance(query_id_series, pd.Series)
+ assert list(query_id_series) == ["1", "2"]
+
+ def test_property_score_returns_correct_series(self):
+ """The score property should return the 'score' column."""
+ data = {
+ "query_id": ["1"],
+ "query_text": ["q1"],
+ "doc_label": ["a"],
+ "doc_text": ["text"],
+ "rank": [0],
+ "score": [0.95],
+ }
+ result = VectorStoreSearchOutput(data)
+
+ score_series = result.score
+ assert isinstance(score_series, pd.Series)
+ assert list(score_series) == [0.95]
+
+ def test_from_data_classmethod_valid_data(self):
+ """from_data classmethod should construct from dict or DataFrame."""
+ data = {
+ "query_id": ["1"],
+ "query_text": ["q1"],
+ "doc_label": ["d1"],
+ "doc_text": ["text"],
+ "rank": [0],
+ "score": [0.9],
+ }
+ result = VectorStoreSearchOutput.from_data(data)
+
+ assert isinstance(result, VectorStoreSearchOutput)
+ assert len(result) == 1
+
+ def test_validate_classmethod_returns_instance(self):
+ """Validate classmethod should return a VectorStoreSearchOutput instance."""
+ df = pd.DataFrame(
+ {
+ "query_id": ["1"],
+ "query_text": ["q1"],
+ "doc_label": ["d1"],
+ "doc_text": ["text"],
+ "rank": [0],
+ "score": [0.9],
+ }
+ )
+ result = VectorStoreSearchOutput.validate(df)
+
+ assert isinstance(result, VectorStoreSearchOutput)
+
+
+class TestVectorStoreEmbedInput:
+ """Tests for VectorStoreEmbedInput dataclass."""
+
+ def test_init_from_dict_valid_data(self):
+ """Valid dict with id and text columns should construct successfully."""
+ data = {"id": ["1", "2"], "text": ["hello", "world"]}
+ result = VectorStoreEmbedInput(data)
+
+ assert isinstance(result, pd.DataFrame)
+ assert list(result.columns) == ["id", "text"]
+ assert len(result) == 2
+
+ def test_init_missing_required_column_raises_schema_error(self):
+ """Missing 'text' column should raise SchemaError."""
+ data = {"id": ["1", "2"]}
+
+ with pytest.raises(pa.errors.SchemaError):
+ VectorStoreEmbedInput(data)
+
+ def test_property_id_returns_correct_series(self):
+ """The id property should return the 'id' column as a Series."""
+ data = {"id": ["1", "2"], "text": ["hello", "world"]}
+ result = VectorStoreEmbedInput(data)
+
+ id_series = result.id
+ assert isinstance(id_series, pd.Series)
+ assert list(id_series) == ["1", "2"]
+
+ def test_property_text_returns_correct_series(self):
+ """The text property should return the 'text' column as a Series."""
+ data = {"id": ["1", "2"], "text": ["hello", "world"]}
+ result = VectorStoreEmbedInput(data)
+
+ text_series = result.text
+ assert isinstance(text_series, pd.Series)
+ assert list(text_series) == ["hello", "world"]
+
+ def test_from_data_classmethod_valid_data(self):
+ """from_data classmethod should construct from dict or DataFrame."""
+ data = {"id": ["1", "2"], "text": ["hello", "world"]}
+ result = VectorStoreEmbedInput.from_data(data)
+
+ assert isinstance(result, VectorStoreEmbedInput)
+ assert len(result) == 2
+
+ def test_validate_classmethod_returns_instance(self):
+ """Validate classmethod should return a VectorStoreEmbedInput instance."""
+ df = pd.DataFrame({"id": ["1"], "text": ["hello"]})
+ result = VectorStoreEmbedInput.validate(df)
+
+ assert isinstance(result, VectorStoreEmbedInput)
+
+
+class TestVectorStoreEmbedOutput:
+ """Tests for VectorStoreEmbedOutput dataclass."""
+
+ def test_init_from_dict_valid_data(self):
+ """Valid dict with id, text, and embedding columns should construct successfully."""
+ data = {
+ "id": ["1", "2"],
+ "text": ["hello", "world"],
+ "embedding": [np.array([0.1, 0.2, 0.3]), np.array([0.4, 0.5, 0.6])],
+ }
+ result = VectorStoreEmbedOutput(data)
+
+ assert isinstance(result, pd.DataFrame)
+ assert len(result) == 2
+ assert list(result.columns) == ["id", "text", "embedding"]
+
+ def test_init_missing_required_column_raises_schema_error(self):
+ """Missing 'text' column should raise SchemaError."""
+ data = {
+ "id": ["1"],
+ "embedding": [np.array([0.1, 0.2])],
+ }
+
+ with pytest.raises(pa.errors.SchemaError):
+ VectorStoreEmbedOutput(data)
+
+ def test_init_non_array_embedding_raises_schema_error(self):
+ """Non-numpy-array embedding should raise SchemaError."""
+ data = {
+ "id": ["1"],
+ "text": ["hello"],
+ "embedding": [[0.1, 0.2]], # list, not numpy array
+ }
+
+ with pytest.raises(pa.errors.SchemaError):
+ VectorStoreEmbedOutput(data)
+
+ def test_property_id_returns_correct_series(self):
+ """The id property should return the 'id' column as a Series."""
+ data = {
+ "id": ["1", "2"],
+ "text": ["hello", "world"],
+ "embedding": [np.array([0.1, 0.2]), np.array([0.3, 0.4])],
+ }
+ result = VectorStoreEmbedOutput(data)
+
+ id_series = result.id
+ assert isinstance(id_series, pd.Series)
+ assert list(id_series) == ["1", "2"]
+
+ def test_property_embedding_returns_correct_series(self):
+ """The embedding property should return the 'embedding' column."""
+ embedding1 = np.array([0.1, 0.2])
+ embedding2 = np.array([0.3, 0.4])
+ data = {
+ "id": ["1", "2"],
+ "text": ["hello", "world"],
+ "embedding": [embedding1, embedding2],
+ }
+ result = VectorStoreEmbedOutput(data)
+
+ embedding_series = result.embedding
+ assert isinstance(embedding_series, pd.Series)
+ assert isinstance(embedding_series.iloc[0], np.ndarray)
+
+ def test_from_data_classmethod_valid_data(self):
+ """from_data classmethod should construct from dict or DataFrame."""
+ data = {
+ "id": ["1"],
+ "text": ["hello"],
+ "embedding": [np.array([0.1, 0.2])],
+ }
+ result = VectorStoreEmbedOutput.from_data(data)
+
+ assert isinstance(result, VectorStoreEmbedOutput)
+ assert len(result) == 1
+
+
+class TestVectorStoreReverseSearchInput:
+ """Tests for VectorStoreReverseSearchInput dataclass."""
+
+ def test_init_from_dict_valid_data(self):
+ """Valid dict with id and doc_label columns should construct successfully."""
+ data = {"id": ["1", "2"], "doc_label": ["label1", "label2"]}
+ result = VectorStoreReverseSearchInput(data)
+
+ assert isinstance(result, pd.DataFrame)
+ assert list(result.columns) == ["id", "doc_label"]
+ assert len(result) == 2
+
+ def test_init_missing_required_column_raises_schema_error(self):
+ """Missing 'doc_label' column should raise SchemaError."""
+ data = {"id": ["1", "2"]}
+
+ with pytest.raises(pa.errors.SchemaError):
+ VectorStoreReverseSearchInput(data)
+
+ def test_property_id_returns_correct_series(self):
+ """The id property should return the 'id' column as a Series."""
+ data = {"id": ["1", "2"], "doc_label": ["label1", "label2"]}
+ result = VectorStoreReverseSearchInput(data)
+
+ id_series = result.id
+ assert isinstance(id_series, pd.Series)
+ assert list(id_series) == ["1", "2"]
+
+ def test_property_doc_label_returns_correct_series(self):
+ """The doc_label property should return the 'doc_label' column as a Series."""
+ data = {"id": ["1", "2"], "doc_label": ["label1", "label2"]}
+ result = VectorStoreReverseSearchInput(data)
+
+ doc_label_series = result.doc_label
+ assert isinstance(doc_label_series, pd.Series)
+ assert list(doc_label_series) == ["label1", "label2"]
+
+ def test_from_data_classmethod_valid_data(self):
+ """from_data classmethod should construct from dict or DataFrame."""
+ data = {"id": ["1", "2"], "doc_label": ["label1", "label2"]}
+ result = VectorStoreReverseSearchInput.from_data(data)
+
+ assert isinstance(result, VectorStoreReverseSearchInput)
+ assert len(result) == 2
+
+ # TODO: Uncomment after unique=True constraints merged to main - possibly needed for other dataclass tests as well depending on final implementation of ticket-167
+ # def test_init_duplicate_ids_raises_schema_error(self):
+ # """Duplicate 'id' values should raise SchemaError."""
+ # data = {"id": ["1", "1"], "doc_label": ["label1", "label2"]}
+ #
+ # with pytest.raises(pa.errors.SchemaError):
+ # VectorStoreReverseSearchInput(data)
+
+
+class TestVectorStoreReverseSearchOutput:
+ """Tests for VectorStoreReverseSearchOutput dataclass."""
+
+ def test_init_from_dict_valid_data(self):
+ """Valid dict with all 4 required columns should construct successfully."""
+ data = {
+ "id": ["1", "1"],
+ "searched_doc_label": ["label1", "label1"],
+ "doc_label": ["label1", "label1"],
+ "doc_text": ["text1", "text2"],
+ }
+ result = VectorStoreReverseSearchOutput(data)
+
+ assert isinstance(result, pd.DataFrame)
+ assert len(result) == 2
+ assert list(result.columns) == ["id", "searched_doc_label", "doc_label", "doc_text"]
+
+ def test_init_empty_dict_creates_valid_structure_with_columns(self):
+ """Empty dict should create a DataFrame with correct columns and schema."""
+ result = VectorStoreReverseSearchOutput({})
+
+ assert isinstance(result, pd.DataFrame)
+ assert len(result) == 0
+ # Verify all expected columns exist even when empty
+ expected_cols = ["id", "searched_doc_label", "doc_label", "doc_text"]
+ assert all(col in result.columns for col in expected_cols)
+
+ def test_init_missing_required_column_raises_schema_error(self):
+ """Missing 'doc_text' column should raise SchemaError."""
+ data = {
+ "id": ["1"],
+ "searched_doc_label": ["label"],
+ "doc_label": ["label"],
+ }
+
+ with pytest.raises(pa.errors.SchemaError):
+ VectorStoreReverseSearchOutput(data)
+
+ def test_property_id_returns_correct_series(self):
+ """The id property should return the 'id' column as a Series."""
+ data = {
+ "id": ["1", "2"],
+ "searched_doc_label": ["label1", "label2"],
+ "doc_label": ["label1", "label2"],
+ "doc_text": ["text1", "text2"],
+ }
+ result = VectorStoreReverseSearchOutput(data)
+
+ id_series = result.id
+ assert isinstance(id_series, pd.Series)
+ assert list(id_series) == ["1", "2"]
+
+ def test_property_searched_doc_label_returns_correct_series(self):
+ """The searched_doc_label property should return the 'searched_doc_label' column."""
+ data = {
+ "id": ["1", "1"],
+ "searched_doc_label": ["label1", "label1"],
+ "doc_label": ["label1", "label1"],
+ "doc_text": ["text1", "text2"],
+ }
+ result = VectorStoreReverseSearchOutput(data)
+
+ searched_doc_label_series = result.searched_doc_label
+ assert isinstance(searched_doc_label_series, pd.Series)
+ assert list(searched_doc_label_series) == ["label1", "label1"]
+
+ def test_from_data_classmethod_empty_data(self):
+ """from_data should handle empty data correctly and create valid columns."""
+ result = VectorStoreReverseSearchOutput.from_data({})
+
+ assert isinstance(result, VectorStoreReverseSearchOutput)
+ assert len(result) == 0
+ expected_cols = ["id", "searched_doc_label", "doc_label", "doc_text"]
+ assert all(col in result.columns for col in expected_cols)
+
+ def test_from_data_classmethod_valid_data(self):
+ """from_data classmethod should construct from dict or DataFrame."""
+ data = {
+ "id": ["1"],
+ "searched_doc_label": ["label"],
+ "doc_label": ["label"],
+ "doc_text": ["text"],
+ }
+ result = VectorStoreReverseSearchOutput.from_data(data)
+
+ assert isinstance(result, VectorStoreReverseSearchOutput)
+ assert len(result) == 1
+
+ def test_validate_classmethod_returns_instance(self):
+ """Validate classmethod should return a VectorStoreReverseSearchOutput instance."""
+ df = pd.DataFrame(
+ {
+ "id": ["1"],
+ "searched_doc_label": ["label"],
+ "doc_label": ["label"],
+ "doc_text": ["text"],
+ }
+ )
+ result = VectorStoreReverseSearchOutput.validate(df)
+
+ assert isinstance(result, VectorStoreReverseSearchOutput)
diff --git a/tests/test_indexers/test_vectorstore_embed.py b/tests/test_indexers/test_vectorstore_embed.py
new file mode 100644
index 0000000..5a1a23a
--- /dev/null
+++ b/tests/test_indexers/test_vectorstore_embed.py
@@ -0,0 +1,616 @@
+"""Unit tests for VectorStore.embed() method."""
+
+from __future__ import annotations
+
+from unittest.mock import Mock
+
+import numpy as np
+import pytest
+
+from classifai.exceptions import (
+ ClassifaiError,
+ DataValidationError,
+ HookError,
+)
+from classifai.indexers import VectorStore
+from classifai.indexers.dataclasses import VectorStoreEmbedInput, VectorStoreEmbedOutput
+from classifai.indexers.hooks import HookBase
+from classifai.vectorisers import VectoriserBase
+
+# ============================================================================
+# FIXTURES
+# ============================================================================
+
+
+@pytest.fixture
+def mock_vectoriser():
+ """Mock VectoriserBase that returns predictable embeddings."""
+ mock = Mock(spec=VectoriserBase)
+
+ def transform_side_effect(texts):
+ """Return one embedding per text, all with shape (3,)."""
+ num_texts = len(texts)
+ return np.array([np.linspace(0.1, 0.3, 3) + (i * 0.3) for i in range(num_texts)])
+
+ mock.transform.side_effect = transform_side_effect
+ mock.__class__.__name__ = "MockVectoriser"
+ return mock
+
+
+@pytest.fixture
+def initialized_vectorstore(mock_vectoriser, tmp_path):
+ """Create a fully initialized VectorStore with test data."""
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text("label,text\ndoc1,hello world\ndoc2,goodbye world\ndoc3,test document\n")
+
+ vs = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ output_dir=str(tmp_path / "output"),
+ skip_save=True,
+ )
+
+ return vs
+
+
+# ============================================================================
+# INPUT VALIDATION TESTS
+# ============================================================================
+
+
+class TestVectorStoreEmbedInputValidation:
+ """Tests for input validation in embed() method."""
+
+ def test_embed_query_must_be_vectorstore_embed_input(self, initialized_vectorstore):
+ """Query must be VectorStoreEmbedInput object."""
+ with pytest.raises(DataValidationError) as exc_info:
+ initialized_vectorstore.embed(query="not a VectorStoreEmbedInput")
+
+ assert "VectorStoreEmbedInput" in str(exc_info.value)
+
+ def test_embed_query_none_raises_error(self, initialized_vectorstore):
+ """query=None raises DataValidationError."""
+ with pytest.raises(DataValidationError) as exc_info:
+ initialized_vectorstore.embed(query=None)
+
+ assert "VectorStoreEmbedInput" in str(exc_info.value)
+
+ def test_embed_query_dict_raises_error(self, initialized_vectorstore):
+ """Query as dict (not VectorStoreEmbedInput) raises DataValidationError."""
+ with pytest.raises(DataValidationError) as exc_info:
+ initialized_vectorstore.embed(query={"id": ["1"], "text": ["hello"]})
+
+ assert "VectorStoreEmbedInput" in str(exc_info.value)
+
+ def test_embed_query_list_raises_error(self, initialized_vectorstore):
+ """Query as list raises DataValidationError."""
+ with pytest.raises(DataValidationError) as exc_info:
+ initialized_vectorstore.embed(query=["hello", "world"])
+
+ assert "VectorStoreEmbedInput" in str(exc_info.value)
+
+
+# ============================================================================
+# EMBEDDING OPERATION TESTS
+# ============================================================================
+
+
+class TestVectorStoreEmbedOperation:
+ """Tests for the core embedding operation."""
+
+ def test_embed_single_text_embeds_correctly(self, initialized_vectorstore):
+ """Single text embeds correctly and returns result."""
+ query = VectorStoreEmbedInput.from_data({"id": ["1"], "text": ["hello world"]})
+
+ result = initialized_vectorstore.embed(query=query)
+
+ assert isinstance(result, VectorStoreEmbedOutput)
+ assert len(result) == 1
+ assert result.id[0] == "1"
+ assert result.text[0] == "hello world"
+
+ def test_embed_multiple_texts_process_correctly(self, initialized_vectorstore):
+ """Multiple texts process correctly and all are embedded."""
+ query = VectorStoreEmbedInput.from_data(
+ {
+ "id": ["1", "2", "3"],
+ "text": ["hello", "world", "test"],
+ }
+ )
+
+ result = initialized_vectorstore.embed(query=query)
+
+ assert isinstance(result, VectorStoreEmbedOutput)
+ assert len(result) == 3
+ assert list(result.id) == ["1", "2", "3"]
+ assert list(result.text) == ["hello", "world", "test"]
+
+ def test_embed_output_includes_id_text_embedding(self, initialized_vectorstore):
+ """Output includes id, text, and embedding columns."""
+ query = VectorStoreEmbedInput.from_data({"id": ["1"], "text": ["test"]})
+
+ result = initialized_vectorstore.embed(query=query)
+
+ assert "id" in result.columns
+ assert "text" in result.columns
+ assert "embedding" in result.columns
+
+ def test_embed_embeddings_are_numpy_arrays(self, initialized_vectorstore):
+ """Embeddings in output are numpy arrays."""
+ query = VectorStoreEmbedInput.from_data(
+ {
+ "id": ["1", "2"],
+ "text": ["hello", "world"],
+ }
+ )
+
+ result = initialized_vectorstore.embed(query=query)
+
+ for embedding in result.embedding:
+ assert isinstance(embedding, np.ndarray)
+
+ def test_embed_output_shape_matches_input_count(self, initialized_vectorstore):
+ """Output has same number of rows as input texts."""
+ query = VectorStoreEmbedInput.from_data(
+ {
+ "id": ["1", "2", "3", "4", "5"],
+ "text": ["a", "b", "c", "d", "e"],
+ }
+ )
+
+ result = initialized_vectorstore.embed(query=query)
+
+ assert len(result) == 5
+
+ def test_embed_vectoriser_transform_called_with_correct_texts(self, initialized_vectorstore):
+ """Vectoriser.transform() called with correct text list."""
+ texts = ["hello", "world", "test"]
+ query = VectorStoreEmbedInput.from_data(
+ {
+ "id": ["1", "2", "3"],
+ "text": texts,
+ }
+ )
+
+ result = initialized_vectorstore.embed(query=query)
+
+ # Verify transform was called
+ initialized_vectorstore.vectoriser.transform.assert_called()
+
+ # Get the call arguments
+ call_args = initialized_vectorstore.vectoriser.transform.call_args[0][0]
+ assert call_args == texts
+
+ def test_embed_embedding_dimensions_correct(self, initialized_vectorstore):
+ """Embeddings have correct dimensions (shape)."""
+ query = VectorStoreEmbedInput.from_data({"id": ["1"], "text": ["test"]})
+
+ result = initialized_vectorstore.embed(query=query)
+
+ embedding = result.embedding[0]
+ # Mock vectoriser returns shape (3,)
+ assert embedding.shape == (3,)
+
+ def test_embed_preserves_id_order(self, initialized_vectorstore):
+ """IDs in output match input order."""
+ ids = ["id_a", "id_b", "id_c"]
+ query = VectorStoreEmbedInput.from_data(
+ {
+ "id": ids,
+ "text": ["text_a", "text_b", "text_c"],
+ }
+ )
+
+ result = initialized_vectorstore.embed(query=query)
+
+ assert list(result.id) == ids
+
+ def test_embed_preserves_text_content(self, initialized_vectorstore):
+ """Text content in output matches input exactly."""
+ texts = ["hello world", "foo bar", "test case"]
+ query = VectorStoreEmbedInput.from_data(
+ {
+ "id": ["1", "2", "3"],
+ "text": texts,
+ }
+ )
+
+ result = initialized_vectorstore.embed(query=query)
+
+ assert list(result.text) == texts
+
+
+# ============================================================================
+# ERROR HANDLING TESTS
+# ============================================================================
+
+
+class TestVectorStoreEmbedErrorHandling:
+ """Tests for error handling during embedding."""
+
+ def test_embed_vectoriser_failure_raises_classifai_error(self, initialized_vectorstore):
+ """Vectoriser failure raises ClassifaiError."""
+ initialized_vectorstore.vectoriser.transform.side_effect = Exception("Transform failed")
+
+ query = VectorStoreEmbedInput.from_data({"id": ["1"], "text": ["hello"]})
+
+ with pytest.raises(ClassifaiError) as exc_info:
+ initialized_vectorstore.embed(query=query)
+
+ assert "Embedding failed" in str(exc_info.value)
+
+ def test_embed_vectoriser_exceptions_wrapped(self, initialized_vectorstore):
+ """Vectoriser exceptions include context."""
+ initialized_vectorstore.vectoriser.transform.side_effect = RuntimeError("Bad embedding")
+
+ query = VectorStoreEmbedInput.from_data({"id": ["1"], "text": ["test"]})
+
+ with pytest.raises(ClassifaiError) as exc_info:
+ initialized_vectorstore.embed(query=query)
+
+ error_str = str(exc_info.value)
+ assert "MockVectoriser" in error_str or "vectoriser" in error_str.lower()
+
+ def test_embed_vectoriser_error_includes_text_count(self, initialized_vectorstore):
+ """Error context includes number of texts being embedded."""
+ initialized_vectorstore.vectoriser.transform.side_effect = ValueError("Embedding error")
+
+ query = VectorStoreEmbedInput.from_data(
+ {
+ "id": ["1", "2", "3"],
+ "text": ["a", "b", "c"],
+ }
+ )
+
+ with pytest.raises(ClassifaiError) as exc_info:
+ initialized_vectorstore.embed(query=query)
+
+ error_str = str(exc_info.value)
+ assert "3" in error_str or "n_texts" in error_str.lower()
+
+ def test_embed_vectoriser_error_includes_vectoriser_class(self, initialized_vectorstore):
+ """Error context includes vectoriser class name."""
+ initialized_vectorstore.vectoriser.transform.side_effect = Exception("Failed")
+
+ query = VectorStoreEmbedInput.from_data({"id": ["1"], "text": ["test"]})
+
+ with pytest.raises(ClassifaiError) as exc_info:
+ initialized_vectorstore.embed(query=query)
+
+ assert "MockVectoriser" in str(exc_info.value)
+
+ def test_embed_correct_number_of_embeddings(self, initialized_vectorstore):
+ """Embeddings with correct count are processed correctly."""
+ initialized_vectorstore.vectoriser.transform.side_effect = None
+ # Return exactly 3 embeddings for 3 texts
+ initialized_vectorstore.vectoriser.transform.return_value = np.array(
+ [
+ [0.1, 0.2, 0.3],
+ [0.4, 0.5, 0.6],
+ [0.7, 0.8, 0.9],
+ ]
+ )
+
+ query = VectorStoreEmbedInput.from_data(
+ {
+ "id": ["1", "2", "3"],
+ "text": ["hello", "world", "test"],
+ }
+ )
+
+ result = initialized_vectorstore.embed(query=query)
+
+ assert len(result) == 3
+ assert all(isinstance(emb, np.ndarray) for emb in result.embedding)
+
+
+# ============================================================================
+# HOOKS INTEGRATION TESTS
+# ============================================================================
+
+
+class TestVectorStoreEmbedHooksIntegration:
+ """Tests for hook integration in embed()."""
+
+ def test_embed_preprocess_hook_called_before_embedding(self, initialized_vectorstore):
+ """embed_preprocess hook called before embedding."""
+ mock_hook = Mock(spec=HookBase)
+ mock_hook.return_value = VectorStoreEmbedInput.from_data(
+ {
+ "id": ["1"],
+ "text": ["modified text"],
+ }
+ )
+
+ initialized_vectorstore.hooks = {"embed_preprocess": mock_hook}
+
+ query = VectorStoreEmbedInput.from_data({"id": ["1"], "text": ["hello"]})
+ result = initialized_vectorstore.embed(query=query)
+
+ mock_hook.assert_called_once()
+
+ def test_embed_postprocess_hook_called_after_embedding(self, initialized_vectorstore):
+ """embed_postprocess hook called after embedding."""
+ mock_result = VectorStoreEmbedOutput.from_data(
+ {
+ "id": ["1"],
+ "text": ["hello"],
+ "embedding": [np.array([0.1, 0.2, 0.3])],
+ }
+ )
+
+ mock_hook = Mock(spec=HookBase)
+ mock_hook.return_value = mock_result
+
+ initialized_vectorstore.hooks = {"embed_postprocess": mock_hook}
+
+ query = VectorStoreEmbedInput.from_data({"id": ["1"], "text": ["hello"]})
+ result = initialized_vectorstore.embed(query=query)
+
+ mock_hook.assert_called_once()
+
+ def test_embed_preprocess_hook_failure_raises_hook_error(self, initialized_vectorstore):
+ """Preprocess hook failure raises HookError."""
+ bad_hook = Mock(spec=HookBase)
+ bad_hook.side_effect = Exception("Preprocess failed")
+
+ initialized_vectorstore.hooks = {"embed_preprocess": bad_hook}
+
+ query = VectorStoreEmbedInput.from_data({"id": ["1"], "text": ["hello"]})
+
+ with pytest.raises(HookError) as exc_info:
+ initialized_vectorstore.embed(query=query)
+
+ assert "embed_preprocess" in str(exc_info.value)
+
+ def test_embed_postprocess_hook_failure_raises_hook_error(self, initialized_vectorstore):
+ """Postprocess hook failure raises HookError."""
+ bad_hook = Mock(spec=HookBase)
+ bad_hook.side_effect = Exception("Postprocess failed")
+
+ initialized_vectorstore.hooks = {"embed_postprocess": bad_hook}
+
+ query = VectorStoreEmbedInput.from_data({"id": ["1"], "text": ["hello"]})
+
+ with pytest.raises(HookError) as exc_info:
+ initialized_vectorstore.embed(query=query)
+
+ assert "embed_postprocess" in str(exc_info.value)
+
+ def test_embed_multiple_preprocess_hooks_processed_in_order(self, initialized_vectorstore):
+ """Multiple preprocess hooks processed in order."""
+ hook1 = Mock(spec=HookBase)
+ hook1.return_value = VectorStoreEmbedInput.from_data(
+ {
+ "id": ["1"],
+ "text": ["step1"],
+ }
+ )
+
+ hook2 = Mock(spec=HookBase)
+ hook2.return_value = VectorStoreEmbedInput.from_data(
+ {
+ "id": ["1"],
+ "text": ["step2"],
+ }
+ )
+
+ initialized_vectorstore.hooks = {"embed_preprocess": [hook1, hook2]}
+
+ query = VectorStoreEmbedInput.from_data({"id": ["1"], "text": ["original"]})
+ result = initialized_vectorstore.embed(query=query)
+
+ assert hook1.call_count == 1
+ assert hook2.call_count == 1
+
+ def test_embed_multiple_postprocess_hooks_processed_in_order(self, initialized_vectorstore):
+ """Multiple postprocess hooks processed in order."""
+ mock_output = VectorStoreEmbedOutput.from_data(
+ {
+ "id": ["1"],
+ "text": ["hello"],
+ "embedding": [np.array([0.1, 0.2, 0.3])],
+ }
+ )
+
+ hook1 = Mock(spec=HookBase)
+ hook1.return_value = mock_output
+
+ hook2 = Mock(spec=HookBase)
+ hook2.return_value = mock_output
+
+ initialized_vectorstore.hooks = {"embed_postprocess": [hook1, hook2]}
+
+ query = VectorStoreEmbedInput.from_data({"id": ["1"], "text": ["hello"]})
+ result = initialized_vectorstore.embed(query=query)
+
+ assert hook1.call_count == 1
+ assert hook2.call_count == 1
+
+ def test_embed_single_preprocess_hook_converted_to_list(self, initialized_vectorstore):
+ """Single preprocess hook automatically converted to list."""
+ mock_hook = Mock(spec=HookBase)
+ mock_hook.return_value = VectorStoreEmbedInput.from_data(
+ {
+ "id": ["1"],
+ "text": ["hello"],
+ }
+ )
+
+ initialized_vectorstore.hooks = {"embed_preprocess": mock_hook}
+
+ query = VectorStoreEmbedInput.from_data({"id": ["1"], "text": ["hello"]})
+ result = initialized_vectorstore.embed(query=query)
+
+ mock_hook.assert_called_once()
+
+ def test_embed_single_postprocess_hook_converted_to_list(self, initialized_vectorstore):
+ """Single postprocess hook automatically converted to list."""
+ mock_output = VectorStoreEmbedOutput.from_data(
+ {
+ "id": ["1"],
+ "text": ["hello"],
+ "embedding": [np.array([0.1, 0.2, 0.3])],
+ }
+ )
+
+ mock_hook = Mock(spec=HookBase)
+ mock_hook.return_value = mock_output
+
+ initialized_vectorstore.hooks = {"embed_postprocess": mock_hook}
+
+ query = VectorStoreEmbedInput.from_data({"id": ["1"], "text": ["hello"]})
+ result = initialized_vectorstore.embed(query=query)
+
+ mock_hook.assert_called_once()
+
+ def test_embed_preprocess_hook_modifies_input(self, initialized_vectorstore):
+ """Preprocess hook can modify input before embedding."""
+
+ def modify_hook(input_data):
+ # Add a prefix to all text
+ modified = VectorStoreEmbedInput.from_data(
+ {
+ "id": input_data.id,
+ "text": ["PREFIX: " + t for t in input_data.text],
+ }
+ )
+ return modified
+
+ initialized_vectorstore.hooks = {"embed_preprocess": modify_hook}
+
+ query = VectorStoreEmbedInput.from_data({"id": ["1"], "text": ["hello"]})
+ result = initialized_vectorstore.embed(query=query)
+
+ # Verify vectoriser was called with modified text
+ call_args = initialized_vectorstore.vectoriser.transform.call_args[0][0]
+ assert call_args[0] == "PREFIX: hello"
+
+ def test_embed_postprocess_hook_modifies_output(self, initialized_vectorstore):
+ """Postprocess hook can modify output after embedding."""
+
+ def modify_hook(output_data):
+ # Return same data (in a real scenario might filter/transform)
+ return output_data
+
+ initialized_vectorstore.hooks = {"embed_postprocess": modify_hook}
+
+ query = VectorStoreEmbedInput.from_data({"id": ["1"], "text": ["hello"]})
+ result = initialized_vectorstore.embed(query=query)
+
+ assert isinstance(result, VectorStoreEmbedOutput)
+ assert len(result) == 1
+
+
+# ============================================================================
+# EDGE CASE TESTS
+# ============================================================================
+
+
+class TestVectorStoreEmbedEdgeCases:
+ """Tests for edge cases in embed()."""
+
+ def test_embed_single_character_text(self, initialized_vectorstore):
+ """Embed single character text."""
+ query = VectorStoreEmbedInput.from_data({"id": ["1"], "text": ["a"]})
+
+ result = initialized_vectorstore.embed(query=query)
+
+ assert len(result) == 1
+ assert result.text[0] == "a"
+
+ def test_embed_very_long_text(self, initialized_vectorstore):
+ """Embed very long text."""
+ long_text = "word " * 1000 # 5000 characters
+ query = VectorStoreEmbedInput.from_data({"id": ["1"], "text": [long_text]})
+
+ result = initialized_vectorstore.embed(query=query)
+
+ assert len(result) == 1
+ assert result.text[0] == long_text
+
+ def test_embed_special_characters_in_text(self, initialized_vectorstore):
+ """Embed text with special characters."""
+ special_text = "Hello!@#$%^&*()_+-=[]{}|;:',.<>?/~`"
+ query = VectorStoreEmbedInput.from_data({"id": ["1"], "text": [special_text]})
+
+ result = initialized_vectorstore.embed(query=query)
+
+ assert result.text[0] == special_text
+
+ def test_embed_unicode_text(self, initialized_vectorstore):
+ """Embed text with unicode characters."""
+ unicode_text = "Hello 世界 🌍"
+ query = VectorStoreEmbedInput.from_data({"id": ["1"], "text": [unicode_text]})
+
+ result = initialized_vectorstore.embed(query=query)
+
+ assert result.text[0] == unicode_text
+
+ def test_embed_whitespace_only_text(self, initialized_vectorstore):
+ """Embed text containing only whitespace."""
+ query = VectorStoreEmbedInput.from_data({"id": ["1"], "text": [" "]})
+
+ result = initialized_vectorstore.embed(query=query)
+
+ assert len(result) == 1
+
+ def test_embed_empty_string_text(self, initialized_vectorstore):
+ """Embed empty string text."""
+ query = VectorStoreEmbedInput.from_data({"id": ["1"], "text": [""]})
+
+ result = initialized_vectorstore.embed(query=query)
+
+ assert len(result) == 1
+ assert result.text[0] == ""
+
+ def test_embed_newlines_in_text(self, initialized_vectorstore):
+ """Embed text containing newlines."""
+ multiline_text = "line1\nline2\nline3"
+ query = VectorStoreEmbedInput.from_data({"id": ["1"], "text": [multiline_text]})
+
+ result = initialized_vectorstore.embed(query=query)
+
+ assert result.text[0] == multiline_text
+
+ def test_embed_tabs_in_text(self, initialized_vectorstore):
+ """Embed text containing tabs."""
+ tab_text = "col1\tcol2\tcol3"
+ query = VectorStoreEmbedInput.from_data({"id": ["1"], "text": [tab_text]})
+
+ result = initialized_vectorstore.embed(query=query)
+
+ assert result.text[0] == tab_text
+
+ def test_embed_duplicate_ids_raises_validation_error(self, initialized_vectorstore):
+ """Duplicate IDs in input raises validation error."""
+ import pandera
+
+ with pytest.raises(pandera.errors.SchemaError):
+ VectorStoreEmbedInput.from_data(
+ {
+ "id": ["1", "1"], # duplicate
+ "text": ["hello", "world"],
+ }
+ )
+
+ def test_embed_large_batch(self, initialized_vectorstore):
+ """Embed large number of texts."""
+ n_texts = 100
+ ids = [str(i) for i in range(n_texts)]
+ texts = [f"text {i}" for i in range(n_texts)]
+
+ query = VectorStoreEmbedInput.from_data({"id": ids, "text": texts})
+
+ result = initialized_vectorstore.embed(query=query)
+
+ assert len(result) == n_texts
+
+ def test_embed_returns_correct_type(self, initialized_vectorstore):
+ """embed() returns VectorStoreEmbedOutput."""
+ query = VectorStoreEmbedInput.from_data({"id": ["1"], "text": ["hello"]})
+
+ result = initialized_vectorstore.embed(query=query)
+
+ assert isinstance(result, VectorStoreEmbedOutput)
diff --git a/tests/test_indexers/test_vectorstore_from_filespace.py b/tests/test_indexers/test_vectorstore_from_filespace.py
new file mode 100644
index 0000000..0ebf51b
--- /dev/null
+++ b/tests/test_indexers/test_vectorstore_from_filespace.py
@@ -0,0 +1,981 @@
+"""Unit tests for VectorStore.from_filespace() class method."""
+
+from __future__ import annotations
+
+import json
+import logging
+from unittest.mock import Mock
+
+import numpy as np
+import polars as pl
+import pytest
+
+from classifai.exceptions import (
+ ConfigurationError,
+ DataValidationError,
+ IndexBuildError,
+)
+from classifai.indexers import VectorStore
+from classifai.vectorisers import VectoriserBase
+
+# ============================================================================
+# FIXTURES
+# ============================================================================
+
+
+@pytest.fixture
+def mock_vectoriser():
+ """Mock VectoriserBase that returns predictable embeddings."""
+ mock = Mock(spec=VectoriserBase)
+
+ def transform_side_effect(texts):
+ """Return one embedding per text, all with shape (3,)."""
+ num_texts = len(texts)
+ return np.array([np.linspace(0.1, 0.3, 3) + (i * 0.3) for i in range(num_texts)])
+
+ mock.transform.side_effect = transform_side_effect
+ mock.__class__.__name__ = "MockVectoriser"
+ return mock
+
+
+@pytest.fixture
+def saved_vectorstore(mock_vectoriser, tmp_path):
+ """Create and save a VectorStore to disk for loading tests."""
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text(
+ "label,text,source\ncat_a,hello world,src1\ncat_a,goodbye world,src1\ncat_b,test document,src2\n"
+ )
+
+ output_dir = tmp_path / "vectorstore_output"
+
+ vs = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ meta_data={"source": str},
+ batch_size=64,
+ output_dir=str(output_dir),
+ skip_save=False,
+ )
+
+ return output_dir, mock_vectoriser
+
+
+# ============================================================================
+# INPUT VALIDATION TESTS
+# ============================================================================
+
+
+class TestVectorStoreFromFilespaceInputValidation:
+ """Tests for input validation in from_filespace()."""
+
+ def test_from_filespace_folder_path_must_be_string(self, mock_vectoriser):
+ """folder_path must be string."""
+ with pytest.raises(DataValidationError) as exc_info:
+ VectorStore.from_filespace(folder_path=123, vectoriser=mock_vectoriser)
+
+ assert "folder_path" in str(exc_info.value).lower()
+
+ def test_from_filespace_folder_path_must_be_non_empty(self, mock_vectoriser):
+ """folder_path must be non-empty string."""
+ with pytest.raises(DataValidationError) as exc_info:
+ VectorStore.from_filespace(folder_path="", vectoriser=mock_vectoriser)
+
+ assert "folder_path" in str(exc_info.value).lower()
+
+ def test_from_filespace_folder_path_must_exist(self, mock_vectoriser):
+ """folder_path must be existing directory."""
+ with pytest.raises(DataValidationError) as exc_info:
+ VectorStore.from_filespace(
+ folder_path="/nonexistent/path/to/vectorstore",
+ vectoriser=mock_vectoriser,
+ )
+
+ assert "folder_path" in str(exc_info.value).lower() or "exist" in str(exc_info.value).lower()
+
+ def test_from_filespace_batch_size_must_be_positive_int_or_none(self, saved_vectorstore):
+ """batch_size must be int >= 1 or None."""
+ output_dir, mock_vectoriser = saved_vectorstore
+
+ with pytest.raises(DataValidationError) as exc_info:
+ VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ batch_size=0,
+ )
+
+ assert "batch_size" in str(exc_info.value).lower()
+
+ def test_from_filespace_batch_size_negative_raises_error(self, saved_vectorstore):
+ """batch_size < 1 raises DataValidationError."""
+ output_dir, mock_vectoriser = saved_vectorstore
+
+ with pytest.raises(DataValidationError) as exc_info:
+ VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ batch_size=-5,
+ )
+
+ assert "batch_size" in str(exc_info.value).lower()
+
+ def test_from_filespace_batch_size_none_is_valid(self, saved_vectorstore):
+ """batch_size=None is valid."""
+ output_dir, mock_vectoriser = saved_vectorstore
+
+ # Should not raise
+ vs = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ batch_size=None,
+ )
+ assert vs is not None
+
+ def test_from_filespace_batch_size_non_int_raises_error(self, saved_vectorstore):
+ """batch_size as non-int raises DataValidationError."""
+ output_dir, mock_vectoriser = saved_vectorstore
+
+ with pytest.raises(DataValidationError) as exc_info:
+ VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ batch_size="64",
+ )
+
+ assert "batch_size" in str(exc_info.value).lower()
+
+ def test_from_filespace_hooks_must_be_dict_or_none(self, saved_vectorstore):
+ """Hooks must be dict or None."""
+ output_dir, mock_vectoriser = saved_vectorstore
+
+ with pytest.raises(DataValidationError) as exc_info:
+ VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ hooks="not a dict",
+ )
+
+ assert "hooks" in str(exc_info.value).lower()
+
+ def test_from_filespace_hooks_list_raises_error(self, saved_vectorstore):
+ """Hooks as list raises DataValidationError."""
+ output_dir, mock_vectoriser = saved_vectorstore
+
+ with pytest.raises(DataValidationError) as exc_info:
+ VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ hooks=[],
+ )
+
+ assert "hooks" in str(exc_info.value).lower()
+
+ def test_from_filespace_vectoriser_must_have_transform_method(self, saved_vectorstore):
+ """Vectoriser must have callable .transform() method."""
+ output_dir, _ = saved_vectorstore
+
+ bad_vectoriser = Mock()
+ bad_vectoriser.transform = None # Not callable
+ bad_vectoriser.__class__.__name__ = "BadVectoriser"
+
+ with pytest.raises(ConfigurationError) as exc_info:
+ VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=bad_vectoriser,
+ )
+
+ assert "transform" in str(exc_info.value).lower()
+
+ def test_from_filespace_vectoriser_without_transform_attribute(self, saved_vectorstore):
+ """Vectoriser without .transform attribute raises ConfigurationError."""
+ output_dir, _ = saved_vectorstore
+
+ bad_vectoriser = Mock(spec=[]) # No attributes
+ bad_vectoriser.__class__.__name__ = "BadVectoriser"
+
+ with pytest.raises(ConfigurationError) as exc_info:
+ VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=bad_vectoriser,
+ )
+
+ assert "transform" in str(exc_info.value).lower()
+
+
+# ============================================================================
+# FILE LOADING TESTS
+# ============================================================================
+
+
+class TestVectorStoreFromFilespaceFileLoading:
+ """Tests for file loading and validation."""
+
+ def test_from_filespace_metadata_json_must_exist(self, mock_vectoriser, tmp_path):
+ """metadata.json must exist in folder_path."""
+ output_dir = tmp_path / "empty_folder"
+ output_dir.mkdir()
+
+ with pytest.raises(DataValidationError) as exc_info:
+ VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ assert "metadata" in str(exc_info.value).lower()
+
+ def test_from_filespace_vectors_parquet_must_exist(self, mock_vectoriser, tmp_path):
+ """vectors.parquet must exist in folder_path."""
+ output_dir = tmp_path / "no_vectors"
+ output_dir.mkdir()
+
+ # Create only metadata.json
+ metadata = {
+ "vectoriser_class": "MockVectoriser",
+ "vector_shape": 3,
+ "num_vectors": 1,
+ "batch_size": 128,
+ "created_at": 1000.0,
+ "meta_data": {},
+ }
+
+ with open(output_dir / "metadata.json", "w") as f:
+ json.dump(metadata, f)
+
+ with pytest.raises(DataValidationError) as exc_info:
+ VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ assert "parquet" in str(exc_info.value).lower() or "vectors" in str(exc_info.value).lower()
+
+ def test_from_filespace_vectors_parquet_must_not_be_empty(self, mock_vectoriser, tmp_path):
+ """vectors.parquet must not be empty."""
+ output_dir = tmp_path / "empty_parquet"
+ output_dir.mkdir()
+
+ # Create metadata.json
+ metadata = {
+ "vectoriser_class": "MockVectoriser",
+ "vector_shape": 3,
+ "num_vectors": 0,
+ "batch_size": 128,
+ "created_at": 1000.0,
+ "meta_data": {},
+ }
+
+ with open(output_dir / "metadata.json", "w") as f:
+ json.dump(metadata, f)
+
+ # Create empty parquet file
+ empty_df = pl.DataFrame(
+ schema={
+ "label": pl.Utf8,
+ "text": pl.Utf8,
+ "embeddings": pl.List(pl.Float32),
+ "uuid": pl.Utf8,
+ }
+ )
+ empty_df.write_parquet(output_dir / "vectors.parquet")
+
+ with pytest.raises(DataValidationError) as exc_info:
+ VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ assert "empty" in str(exc_info.value).lower()
+
+ def test_from_filespace_malformed_metadata_json_raises_error(self, mock_vectoriser, tmp_path):
+ """Malformed metadata.json raises IndexBuildError."""
+ output_dir = tmp_path / "bad_metadata"
+ output_dir.mkdir()
+
+ # Write invalid JSON
+ with open(output_dir / "metadata.json", "w") as f:
+ f.write("{invalid json")
+
+ with pytest.raises(IndexBuildError) as exc_info:
+ VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ assert "metadata" in str(exc_info.value).lower()
+
+ def test_from_filespace_metadata_not_dict_raises_error(self, mock_vectoriser, tmp_path):
+ """metadata.json not containing dict raises DataValidationError."""
+ output_dir = tmp_path / "metadata_list"
+ output_dir.mkdir()
+
+ # Write JSON array instead of object
+ with open(output_dir / "metadata.json", "w") as f:
+ json.dump([], f)
+
+ with pytest.raises(DataValidationError) as exc_info:
+ VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ assert "object" in str(exc_info.value).lower() or "dict" in str(exc_info.value).lower()
+
+ def test_from_filespace_missing_required_metadata_keys(self, mock_vectoriser, tmp_path):
+ """Missing required metadata keys raises DataValidationError."""
+ output_dir = tmp_path / "incomplete_metadata"
+ output_dir.mkdir()
+
+ # Create incomplete metadata (missing vector_shape)
+ metadata = {
+ "vectoriser_class": "MockVectoriser",
+ "num_vectors": 1,
+ "created_at": 1000.0,
+ "meta_data": {},
+ }
+
+ with open(output_dir / "metadata.json", "w") as f:
+ json.dump(metadata, f)
+
+ with pytest.raises(DataValidationError) as exc_info:
+ VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ assert "missing" in str(exc_info.value).lower() or "required" in str(exc_info.value).lower()
+
+ def test_from_filespace_required_columns_in_parquet(self, saved_vectorstore):
+ """Required columns must be present in parquet."""
+ output_dir, mock_vectoriser = saved_vectorstore
+
+ # Load existing parquet and remove a required column
+ df = pl.read_parquet(output_dir / "vectors.parquet")
+ df_missing = df.drop("uuid")
+ df_missing.write_parquet(output_dir / "vectors.parquet")
+
+ with pytest.raises((DataValidationError, IndexBuildError)) as exc_info:
+ VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ error_message = str(exc_info.value).lower()
+ assert "missing" in error_message or "column" in error_message or "unable to find" in error_message
+
+
+# ============================================================================
+# CONFIGURATION VALIDATION TESTS
+# ============================================================================
+
+
+class TestVectorStoreFromFilespaceConfigurationValidation:
+ """Tests for configuration validation."""
+
+ def test_from_filespace_vectoriser_class_must_match_metadata(self, mock_vectoriser, tmp_path):
+ """Vectoriser class name must match metadata."""
+ output_dir = tmp_path / "class_mismatch"
+ output_dir.mkdir()
+
+ # Create metadata with different vectoriser class
+ metadata = {
+ "vectoriser_class": "DifferentVectoriser",
+ "vector_shape": 3,
+ "num_vectors": 1,
+ "batch_size": 128,
+ "created_at": 1000.0,
+ "meta_data": {},
+ }
+
+ with open(output_dir / "metadata.json", "w") as f:
+ json.dump(metadata, f)
+
+ # Create parquet file
+ df = pl.DataFrame(
+ {
+ "label": ["test"],
+ "text": ["hello"],
+ "embeddings": [np.array([0.1, 0.2, 0.3])],
+ "uuid": ["uuid1"],
+ }
+ )
+ df.write_parquet(output_dir / "vectors.parquet")
+
+ with pytest.raises(ConfigurationError) as exc_info:
+ VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ assert "vectoriser" in str(exc_info.value).lower() or "class" in str(exc_info.value).lower()
+
+ def test_from_filespace_meta_data_must_be_dict(self, mock_vectoriser, tmp_path):
+ """metadata.meta_data must be dict."""
+ output_dir = tmp_path / "bad_meta_data"
+ output_dir.mkdir()
+
+ # Create metadata with non-dict meta_data
+ metadata = {
+ "vectoriser_class": "MockVectoriser",
+ "vector_shape": 3,
+ "num_vectors": 1,
+ "batch_size": 128,
+ "created_at": 1000.0,
+ "meta_data": "not a dict",
+ }
+
+ with open(output_dir / "metadata.json", "w") as f:
+ json.dump(metadata, f)
+
+ with pytest.raises(DataValidationError) as exc_info:
+ VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ assert "meta_data" in str(exc_info.value).lower()
+
+
+# ============================================================================
+# METADATA DESERIALIZATION TESTS
+# ============================================================================
+
+
+class TestVectorStoreFromFilespaceMetadataDeserialization:
+ """Tests for metadata deserialization."""
+
+ def test_from_filespace_type_deserialization_str(self, mock_vectoriser, tmp_path):
+ """Type deserialization works for str."""
+ output_dir = tmp_path / "deserialize_str"
+ output_dir.mkdir()
+
+ metadata = {
+ "vectoriser_class": "MockVectoriser",
+ "vector_shape": 3,
+ "num_vectors": 1,
+ "batch_size": 128,
+ "created_at": 1000.0,
+ "meta_data": {"source": "str"},
+ }
+
+ with open(output_dir / "metadata.json", "w") as f:
+ json.dump(metadata, f)
+
+ df = pl.DataFrame(
+ {
+ "label": ["test"],
+ "text": ["hello"],
+ "embeddings": [np.array([0.1, 0.2, 0.3])],
+ "uuid": ["uuid1"],
+ "source": ["src1"],
+ }
+ )
+ df.write_parquet(output_dir / "vectors.parquet")
+
+ vs = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ assert vs.meta_data["source"] == "str"
+
+ def test_from_filespace_type_deserialization_int(self, mock_vectoriser, tmp_path):
+ """Type deserialization works for int."""
+ output_dir = tmp_path / "deserialize_int"
+ output_dir.mkdir()
+
+ metadata = {
+ "vectoriser_class": "MockVectoriser",
+ "vector_shape": 3,
+ "num_vectors": 1,
+ "batch_size": 128,
+ "created_at": 1000.0,
+ "meta_data": {"count": "int"},
+ }
+
+ with open(output_dir / "metadata.json", "w") as f:
+ json.dump(metadata, f)
+
+ df = pl.DataFrame(
+ {
+ "label": ["test"],
+ "text": ["hello"],
+ "embeddings": [np.array([0.1, 0.2, 0.3])],
+ "uuid": ["uuid1"],
+ "count": [5],
+ }
+ )
+ df.write_parquet(output_dir / "vectors.parquet")
+
+ vs = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ assert vs.meta_data["count"] == "int"
+
+ def test_from_filespace_type_deserialization_float(self, mock_vectoriser, tmp_path):
+ """Type deserialization works for float."""
+ output_dir = tmp_path / "deserialize_float"
+ output_dir.mkdir()
+
+ metadata = {
+ "vectoriser_class": "MockVectoriser",
+ "vector_shape": 3,
+ "num_vectors": 1,
+ "batch_size": 128,
+ "created_at": 1000.0,
+ "meta_data": {"score": "float"},
+ }
+
+ with open(output_dir / "metadata.json", "w") as f:
+ json.dump(metadata, f)
+
+ df = pl.DataFrame(
+ {
+ "label": ["test"],
+ "text": ["hello"],
+ "embeddings": [np.array([0.1, 0.2, 0.3])],
+ "uuid": ["uuid1"],
+ "score": [0.95],
+ }
+ )
+ df.write_parquet(output_dir / "vectors.parquet")
+
+ vs = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ assert vs.meta_data["score"] == "float"
+
+ def test_from_filespace_empty_meta_data_dict(self, mock_vectoriser, tmp_path):
+ """Empty meta_data dict handled correctly."""
+ output_dir = tmp_path / "empty_meta_data"
+ output_dir.mkdir()
+
+ metadata = {
+ "vectoriser_class": "MockVectoriser",
+ "vector_shape": 3,
+ "num_vectors": 1,
+ "batch_size": 128,
+ "created_at": 1000.0,
+ "meta_data": {},
+ }
+
+ with open(output_dir / "metadata.json", "w") as f:
+ json.dump(metadata, f)
+
+ df = pl.DataFrame(
+ {
+ "label": ["test"],
+ "text": ["hello"],
+ "embeddings": [np.array([0.1, 0.2, 0.3])],
+ "uuid": ["uuid1"],
+ }
+ )
+ df.write_parquet(output_dir / "vectors.parquet")
+
+ vs = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ assert vs.meta_data == {}
+
+
+# ============================================================================
+# BACKWARDS COMPATIBILITY TESTS
+# ============================================================================
+
+
+class TestVectorStoreFromFilespaceBackwardsCompatibility:
+ """Tests for backwards compatibility with v1.0.0."""
+
+ def test_from_filespace_missing_batch_size_uses_provided(self, mock_vectoriser, tmp_path):
+ """Missing batch_size uses provided value."""
+ output_dir = tmp_path / "missing_batch_size"
+ output_dir.mkdir()
+
+ # Create metadata without batch_size (v1.0.0)
+ metadata = {
+ "vectoriser_class": "MockVectoriser",
+ "vector_shape": 3,
+ "num_vectors": 1,
+ "created_at": 1000.0,
+ "meta_data": {},
+ }
+
+ with open(output_dir / "metadata.json", "w") as f:
+ json.dump(metadata, f)
+
+ df = pl.DataFrame(
+ {
+ "label": ["test"],
+ "text": ["hello"],
+ "embeddings": [np.array([0.1, 0.2, 0.3])],
+ "uuid": ["uuid1"],
+ }
+ )
+ df.write_parquet(output_dir / "vectors.parquet")
+
+ vs = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ batch_size=256,
+ )
+
+ assert vs.batch_size == 256
+
+ def test_from_filespace_missing_batch_size_uses_default(self, mock_vectoriser, tmp_path):
+ """Missing batch_size uses default when not provided."""
+ output_dir = tmp_path / "missing_batch_size_default"
+ output_dir.mkdir()
+
+ # Create metadata without batch_size (v1.0.0)
+ metadata = {
+ "vectoriser_class": "MockVectoriser",
+ "vector_shape": 3,
+ "num_vectors": 1,
+ "created_at": 1000.0,
+ "meta_data": {},
+ }
+
+ with open(output_dir / "metadata.json", "w") as f:
+ json.dump(metadata, f)
+
+ df = pl.DataFrame(
+ {
+ "label": ["test"],
+ "text": ["hello"],
+ "embeddings": [np.array([0.1, 0.2, 0.3])],
+ "uuid": ["uuid1"],
+ }
+ )
+ df.write_parquet(output_dir / "vectors.parquet")
+
+ vs = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ # Default should be 128
+ assert vs.batch_size == 128
+
+ def test_from_filespace_warning_logged_for_missing_batch_size(self, mock_vectoriser, tmp_path, caplog):
+ """Warning logged when batch_size missing from metadata."""
+ output_dir = tmp_path / "warning_batch_size"
+ output_dir.mkdir()
+
+ # Create metadata without batch_size (v1.0.0)
+ metadata = {
+ "vectoriser_class": "MockVectoriser",
+ "vector_shape": 3,
+ "num_vectors": 1,
+ "created_at": 1000.0,
+ "meta_data": {},
+ }
+
+ with open(output_dir / "metadata.json", "w") as f:
+ json.dump(metadata, f)
+
+ df = pl.DataFrame(
+ {
+ "label": ["test"],
+ "text": ["hello"],
+ "embeddings": [np.array([0.1, 0.2, 0.3])],
+ "uuid": ["uuid1"],
+ }
+ )
+ df.write_parquet(output_dir / "vectors.parquet")
+
+ with caplog.at_level(logging.WARNING):
+ vs = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ warning_messages = [record.message.lower() for record in caplog.records if record.levelname == "WARNING"]
+ assert any("outdated" in msg or "batch_size" in msg for msg in warning_messages)
+
+
+# ============================================================================
+# INSTANCE CONSTRUCTION TESTS
+# ============================================================================
+
+
+class TestVectorStoreFromFilespaceInstanceConstruction:
+ """Tests for instance construction and attribute assignment."""
+
+ def test_from_filespace_instance_created_without_init(self, saved_vectorstore):
+ """Instance created via object.__new__() (not calling __init__)."""
+ output_dir, mock_vectoriser = saved_vectorstore
+
+ vs = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ # If __init__ was called, it would fail because file_name/data_type are None
+ assert vs is not None
+ assert isinstance(vs, VectorStore)
+
+ def test_from_filespace_file_name_is_none(self, saved_vectorstore):
+ """file_name attribute is None."""
+ output_dir, mock_vectoriser = saved_vectorstore
+
+ vs = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ assert vs.file_name is None
+
+ def test_from_filespace_data_type_is_none(self, saved_vectorstore):
+ """data_type attribute is None."""
+ output_dir, mock_vectoriser = saved_vectorstore
+
+ vs = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ assert vs.data_type is None
+
+ def test_from_filespace_vectoriser_attached(self, saved_vectorstore):
+ """Vectoriser instance attached correctly."""
+ output_dir, mock_vectoriser = saved_vectorstore
+
+ vs = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ assert vs.vectoriser is mock_vectoriser
+
+ def test_from_filespace_batch_size_override_priority(self, saved_vectorstore):
+ """batch_size override has priority over metadata."""
+ output_dir, mock_vectoriser = saved_vectorstore
+
+ vs = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ batch_size=256,
+ )
+
+ # Override should take precedence
+ assert vs.batch_size == 256
+
+ def test_from_filespace_meta_data_deserialized(self, saved_vectorstore):
+ """meta_data deserialized and set correctly."""
+ output_dir, mock_vectoriser = saved_vectorstore
+
+ vs = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ assert isinstance(vs.meta_data, dict)
+ assert "source" in vs.meta_data
+ assert vs.meta_data["source"] == "str"
+
+ def test_from_filespace_vectors_loaded(self, saved_vectorstore):
+ """Vectors DataFrame loaded from parquet."""
+ output_dir, mock_vectoriser = saved_vectorstore
+
+ vs = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ assert vs.vectors is not None
+ assert isinstance(vs.vectors, pl.DataFrame)
+ assert len(vs.vectors) == 3
+
+ def test_from_filespace_vector_shape_set(self, saved_vectorstore):
+ """vector_shape attribute set from metadata."""
+ output_dir, mock_vectoriser = saved_vectorstore
+
+ vs = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ assert vs.vector_shape == 3
+
+ def test_from_filespace_num_vectors_set(self, saved_vectorstore):
+ """num_vectors attribute set from metadata."""
+ output_dir, mock_vectoriser = saved_vectorstore
+
+ vs = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ assert vs.num_vectors == 3
+
+ def test_from_filespace_vectoriser_class_set(self, saved_vectorstore):
+ """vectoriser_class attribute set from metadata."""
+ output_dir, mock_vectoriser = saved_vectorstore
+
+ vs = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ assert vs.vectoriser_class == "MockVectoriser"
+
+ def test_from_filespace_hooks_applied(self, saved_vectorstore):
+ """Hooks parameter applied correctly."""
+ output_dir, mock_vectoriser = saved_vectorstore
+
+ hooks = {"custom_hook": lambda x: x}
+
+ vs = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ hooks=hooks,
+ )
+
+ assert vs.hooks == hooks
+
+ def test_from_filespace_hooks_default_empty_dict(self, saved_vectorstore):
+ """Hooks defaults to empty dict when None."""
+ output_dir, mock_vectoriser = saved_vectorstore
+
+ vs = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ hooks=None,
+ )
+
+ assert vs.hooks == {}
+
+ def test_from_filespace_quiet_mode_applied(self, saved_vectorstore):
+ """quiet_mode parameter applied correctly."""
+ output_dir, mock_vectoriser = saved_vectorstore
+
+ vs = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ quiet_mode=True,
+ )
+
+ assert vs.quiet_mode is True
+
+
+# ============================================================================
+# EDGE CASE TESTS
+# ============================================================================
+
+
+class TestVectorStoreFromFilespaceEdgeCases:
+ """Tests for edge cases in from_filespace()."""
+
+ def test_from_filespace_large_batch_size_override(self, saved_vectorstore):
+ """Very large batch_size override handled correctly."""
+ output_dir, mock_vectoriser = saved_vectorstore
+
+ vs = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ batch_size=10000,
+ )
+
+ assert vs.batch_size == 10000
+
+ def test_from_filespace_batch_size_one(self, saved_vectorstore):
+ """batch_size=1 is valid."""
+ output_dir, mock_vectoriser = saved_vectorstore
+
+ vs = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ batch_size=1,
+ )
+
+ assert vs.batch_size == 1
+
+ def test_from_filespace_multiple_metadata_columns(self, mock_vectoriser, tmp_path):
+ """Multiple metadata columns deserialized correctly."""
+ output_dir = tmp_path / "multi_meta"
+ output_dir.mkdir()
+
+ metadata = {
+ "vectoriser_class": "MockVectoriser",
+ "vector_shape": 3,
+ "num_vectors": 1,
+ "batch_size": 128,
+ "created_at": 1000.0,
+ "meta_data": {"source": "str", "count": "int", "score": "float"},
+ }
+
+ with open(output_dir / "metadata.json", "w") as f:
+ json.dump(metadata, f)
+
+ df = pl.DataFrame(
+ {
+ "label": ["test"],
+ "text": ["hello"],
+ "embeddings": [np.array([0.1, 0.2, 0.3])],
+ "uuid": ["uuid1"],
+ "source": ["src1"],
+ "count": [5],
+ "score": [0.95],
+ }
+ )
+ df.write_parquet(output_dir / "vectors.parquet")
+
+ vs = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ assert vs.meta_data["source"] == "str"
+ assert vs.meta_data["count"] == "int"
+ assert vs.meta_data["score"] == "float"
+
+ def test_from_filespace_many_documents(self, mock_vectoriser, tmp_path):
+ """Large number of documents loaded correctly."""
+ output_dir = tmp_path / "many_docs"
+ output_dir.mkdir()
+
+ metadata = {
+ "vectoriser_class": "MockVectoriser",
+ "vector_shape": 3,
+ "num_vectors": 100,
+ "batch_size": 128,
+ "created_at": 1000.0,
+ "meta_data": {},
+ }
+
+ with open(output_dir / "metadata.json", "w") as f:
+ json.dump(metadata, f)
+
+ # Create 100 documents
+ data = {
+ "label": [f"cat_{i % 3}" for i in range(100)],
+ "text": [f"document {i}" for i in range(100)],
+ "embeddings": [np.random.rand(3) for _ in range(100)],
+ "uuid": [f"uuid_{i}" for i in range(100)],
+ }
+ df = pl.DataFrame(data)
+ df.write_parquet(output_dir / "vectors.parquet")
+
+ vs = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ assert vs.num_vectors == 100
+ assert len(vs.vectors) == 100
+
+ def test_from_filespace_returns_vectorstore_instance(self, saved_vectorstore):
+ """from_filespace() returns VectorStore instance."""
+ output_dir, mock_vectoriser = saved_vectorstore
+
+ vs = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ assert isinstance(vs, VectorStore)
diff --git a/tests/test_indexers/test_vectorstore_init.py b/tests/test_indexers/test_vectorstore_init.py
new file mode 100644
index 0000000..f80e3a8
--- /dev/null
+++ b/tests/test_indexers/test_vectorstore_init.py
@@ -0,0 +1,845 @@
+"""Unit tests for VectorStore initialization."""
+
+from __future__ import annotations
+
+import json
+import os
+from unittest.mock import MagicMock, Mock, patch
+
+import numpy as np
+import polars as pl
+import pytest
+
+from classifai._optional import OptionalDependencyError
+from classifai.exceptions import (
+ ConfigurationError,
+ DataValidationError,
+ IndexBuildError,
+)
+from classifai.indexers import VectorStore
+from classifai.vectorisers import VectoriserBase
+
+# ============================================================================
+# FIXTURES
+# ============================================================================
+
+
+@pytest.fixture
+def mock_vectoriser():
+ """Mock VectoriserBase that returns predictable embeddings."""
+ mock = Mock(spec=VectoriserBase)
+
+ # Make transform() return embeddings matching input size
+ def transform_side_effect(texts):
+ """Return one embedding per text, all with shape (3,)."""
+ num_texts = len(texts)
+ return np.array([np.linspace(0.1, 0.3, 3) + (i * 0.3) for i in range(num_texts)])
+
+ mock.transform.side_effect = transform_side_effect
+ mock.__class__.__name__ = "MockVectoriser"
+ return mock
+
+
+@pytest.fixture
+def temp_csv_file(tmp_path):
+ """Create a real temp CSV with id, label, text columns."""
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text("label,text\ndoc1,hello world\ndoc2,goodbye world\n")
+ return str(csv_path)
+
+
+@pytest.fixture
+def temp_csv_with_metadata(tmp_path):
+ """Create a real temp CSV with additional metadata columns."""
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text("label,text,source\ndoc1,hello,web\ndoc2,goodbye,file\n")
+ return str(csv_path)
+
+
+@pytest.fixture
+def temp_output_dir(tmp_path):
+ """Create a temporary output directory."""
+ return str(tmp_path / "output")
+
+
+# ============================================================================
+# INPUT VALIDATION TESTS (DataValidationError)
+# ============================================================================
+
+
+class TestVectorStoreInitValidation:
+ """Tests for input parameter validation in VectorStore.__init__."""
+
+ @patch("classifai.indexers.main.fsspec.core.url_to_fs")
+ def test_init_file_name_empty_string_raises_error(self, mock_url_to_fs, mock_vectoriser):
+ """file_name must be a non-empty string."""
+ with pytest.raises(DataValidationError) as exc_info:
+ VectorStore(
+ file_name="",
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ )
+ assert "file_name must be a non-empty string" in str(exc_info.value)
+
+ @patch("classifai.indexers.main.fsspec.core.url_to_fs")
+ def test_init_file_name_not_string_raises_error(self, mock_url_to_fs, mock_vectoriser):
+ """file_name must be a string, not other types."""
+ with pytest.raises(DataValidationError) as exc_info:
+ VectorStore(
+ file_name=123,
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ )
+ assert "file_name must be a non-empty string" in str(exc_info.value)
+
+ @patch("classifai.indexers.main.fsspec.core.url_to_fs")
+ def test_init_data_type_unsupported_raises_error(self, mock_url_to_fs, mock_vectoriser):
+ """data_type must be 'csv', others raise DataValidationError."""
+ mock_fs = Mock()
+ mock_fs.exists.return_value = True
+ mock_url_to_fs.return_value = (mock_fs, "/path/to/file.parquet")
+
+ with pytest.raises(DataValidationError) as exc_info:
+ VectorStore(
+ file_name="file.parquet",
+ data_type="parquet",
+ vectoriser=mock_vectoriser,
+ )
+ assert "Unsupported data_type" in str(exc_info.value)
+
+ @patch("classifai.indexers.main.fsspec.core.url_to_fs")
+ def test_init_vectoriser_not_base_instance_raises_error(self, mock_url_to_fs):
+ """Vectoriser must be VectoriserBase instance."""
+ mock_fs = Mock()
+ mock_fs.exists.return_value = True
+ mock_url_to_fs.return_value = (mock_fs, "/path/to/file.csv")
+
+ with pytest.raises(ConfigurationError) as exc_info:
+ VectorStore(
+ file_name="file.csv",
+ data_type="csv",
+ vectoriser="not_a_vectoriser", # type: ignore
+ )
+ assert "Vectoriser must be an instance" in str(exc_info.value)
+
+ @patch("classifai.indexers.main.fsspec.core.url_to_fs")
+ def test_init_batch_size_negative_raises_error(self, mock_url_to_fs, mock_vectoriser):
+ """batch_size must be >= 1."""
+ mock_fs = Mock()
+ mock_fs.exists.return_value = True
+ mock_url_to_fs.return_value = (mock_fs, "/path/to/file.csv")
+
+ with pytest.raises(DataValidationError) as exc_info:
+ VectorStore(
+ file_name="file.csv",
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ batch_size=-1,
+ )
+ assert "batch_size must be an integer >= 1" in str(exc_info.value)
+
+ @patch("classifai.indexers.main.fsspec.core.url_to_fs")
+ def test_init_batch_size_zero_raises_error(self, mock_url_to_fs, mock_vectoriser):
+ """batch_size must be >= 1."""
+ mock_fs = Mock()
+ mock_fs.exists.return_value = True
+ mock_url_to_fs.return_value = (mock_fs, "/path/to/file.csv")
+
+ with pytest.raises(DataValidationError) as exc_info:
+ VectorStore(
+ file_name="file.csv",
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ batch_size=0,
+ )
+ assert "batch_size must be an integer >= 1" in str(exc_info.value)
+
+ @patch("classifai.indexers.main.fsspec.core.url_to_fs")
+ def test_init_batch_size_not_int_raises_error(self, mock_url_to_fs, mock_vectoriser):
+ """batch_size must be an integer."""
+ mock_fs = Mock()
+ mock_fs.exists.return_value = True
+ mock_url_to_fs.return_value = (mock_fs, "/path/to/file.csv")
+
+ with pytest.raises(DataValidationError) as exc_info:
+ VectorStore(
+ file_name="file.csv",
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ batch_size="32", # type: ignore
+ )
+ assert "batch_size must be an integer >= 1" in str(exc_info.value)
+
+ @patch("classifai.indexers.main.fsspec.core.url_to_fs")
+ def test_init_meta_data_not_dict_raises_error(self, mock_url_to_fs, mock_vectoriser):
+ """meta_data must be dict or None."""
+ mock_fs = Mock()
+ mock_fs.exists.return_value = True
+ mock_url_to_fs.return_value = (mock_fs, "/path/to/file.csv")
+
+ with pytest.raises(DataValidationError) as exc_info:
+ VectorStore(
+ file_name="file.csv",
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ meta_data="not_a_dict", # type: ignore
+ )
+ assert "meta_data must be a dict or None" in str(exc_info.value)
+
+ @patch("classifai.indexers.main.fsspec.core.url_to_fs")
+ def test_init_hooks_not_dict_raises_error(self, mock_url_to_fs, mock_vectoriser):
+ """Hooks must be dict or None."""
+ mock_fs = Mock()
+ mock_fs.exists.return_value = True
+ mock_url_to_fs.return_value = (mock_fs, "/path/to/file.csv")
+
+ with pytest.raises(DataValidationError) as exc_info:
+ VectorStore(
+ file_name="file.csv",
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ hooks="not_a_dict", # type: ignore
+ )
+ assert "hooks must be a dict or None" in str(exc_info.value)
+
+ @patch("classifai.indexers.main.fsspec.core.url_to_fs")
+ def test_init_output_dir_not_string_raises_error(self, mock_url_to_fs, mock_vectoriser):
+ """output_dir must be string or None."""
+ mock_fs = Mock()
+ mock_fs.exists.return_value = True
+ mock_url_to_fs.return_value = (mock_fs, "/path/to/file.csv")
+
+ with pytest.raises(DataValidationError) as exc_info:
+ VectorStore(
+ file_name="file.csv",
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ output_dir=123, # type: ignore
+ skip_save=True,
+ )
+ assert "output_dir must be a string or None" in str(exc_info.value)
+
+
+# # ============================================================================
+# # FILE SYSTEM HANDLING TESTS (ConfigurationError / OptionalDependencyError)
+# # ============================================================================
+
+
+class TestVectorStoreInitFileSystem:
+ """Tests for file system handling in VectorStore.__init__."""
+
+ @patch("classifai.indexers.main.fsspec.core.url_to_fs")
+ def test_init_input_file_not_exist_raises_error(self, mock_url_to_fs, mock_vectoriser):
+ """Input file must exist."""
+ mock_fs = Mock()
+ mock_fs.exists.return_value = False
+ mock_url_to_fs.return_value = (mock_fs, "/path/to/missing.csv")
+
+ with pytest.raises(DataValidationError) as exc_info:
+ VectorStore(
+ file_name="/path/to/missing.csv",
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ )
+ assert "Input file does not exist" in str(exc_info.value)
+
+ @patch("classifai.indexers.main.fsspec.core.url_to_fs")
+ def test_init_input_fsspec_error_raises_configuration_error(self, mock_url_to_fs, mock_vectoriser):
+ """Fsspec resolution failure → ConfigurationError."""
+ mock_url_to_fs.side_effect = Exception("Invalid fsspec path")
+
+ with pytest.raises(ConfigurationError) as exc_info:
+ VectorStore(
+ file_name="invalid://path/file.csv",
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ )
+ assert "Failed to read input directory with file loader" in str(exc_info.value)
+
+ @patch("classifai.indexers.main.check_deps")
+ @patch("classifai.indexers.main.fsspec.core.url_to_fs")
+ def test_init_gs_path_without_gcsfs_raises_helpful_error(self, mock_url_to_fs, mock_check_deps, mock_vectoriser):
+ """gs:// path without gcsfs → OptionalDependencyError with helpful message."""
+ mock_url_to_fs.side_effect = ImportError("gcsfs not installed")
+ mock_check_deps.side_effect = OptionalDependencyError("gcsfs required")
+
+ with pytest.raises(OptionalDependencyError) as exc_info:
+ VectorStore(
+ file_name="gs://bucket/file.csv",
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ )
+ assert "gcsfs" in str(exc_info.value).lower()
+ assert "pip install" in str(exc_info.value).lower()
+
+
+# # ============================================================================
+# # OUTPUT DIRECTORY HANDLING TESTS (skip_save=False path)
+# # ============================================================================
+
+
+class TestVectorStoreInitOutputDirectory:
+ """Tests for output directory handling when skip_save=False."""
+
+ @patch("classifai.indexers.main.pl.read_csv")
+ @patch("classifai.indexers.main.fsspec.core.url_to_fs")
+ def test_init_output_dir_derived_from_file_name(self, mock_url_to_fs, mock_read_csv, mock_vectoriser):
+ """When output_dir=None, derive from file_name."""
+ # Mock input file
+ mock_fs = Mock()
+ mock_fs.exists.return_value = True
+ mock_fs.makedirs = Mock()
+
+ # Mock file object that supports context manager protocol
+ mock_file = MagicMock()
+ mock_fs.open.return_value.__enter__ = Mock(return_value=mock_file)
+ mock_fs.open.return_value.__exit__ = Mock(return_value=None)
+
+ mock_url_to_fs.side_effect = [
+ (mock_fs, "/path/to/test.csv"), # input
+ (mock_fs, "test"), # output (derived)
+ (mock_fs, "test/metadata.json"), # metadata save
+ ]
+
+ # Mock CSV read
+ mock_read_csv.return_value = pl.DataFrame(
+ {
+ "label": ["doc1", "doc2"],
+ "text": ["hello", "world"],
+ }
+ )
+
+ # Mock the parquet write to prevent actual file I/O
+ with patch("classifai.indexers.main.pl.DataFrame.write_parquet"):
+ vs = VectorStore(
+ file_name="/path/to/test.csv",
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ skip_save=False,
+ overwrite=True,
+ )
+
+ assert vs.output_dir == "test"
+ mock_fs.makedirs.assert_called()
+
+ @patch("classifai.indexers.main.pl.read_csv")
+ @patch("classifai.indexers.main.fsspec.core.url_to_fs")
+ def test_init_output_dir_exists_without_overwrite_raises_error(
+ self, mock_url_to_fs, mock_read_csv, mock_vectoriser
+ ):
+ """Existing output_dir without overwrite=True → ConfigurationError."""
+ mock_fs = Mock()
+ mock_fs.exists.side_effect = [True, True] # input exists, output exists
+ mock_url_to_fs.side_effect = [
+ (mock_fs, "/path/to/test.csv"),
+ (mock_fs, "/path/to/output"),
+ ]
+
+ with pytest.raises(ConfigurationError) as exc_info:
+ VectorStore(
+ file_name="/path/to/test.csv",
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ output_dir="/path/to/output",
+ overwrite=False,
+ skip_save=False,
+ )
+ assert "already exists" in str(exc_info.value).lower()
+
+ @patch("classifai.indexers.main.pl.read_csv")
+ @patch("classifai.indexers.main.fsspec.core.url_to_fs")
+ def test_init_output_dir_exists_with_overwrite_removes_directory(
+ self, mock_url_to_fs, mock_read_csv, mock_vectoriser
+ ):
+ """overwrite=True removes and recreates directory."""
+ mock_fs = Mock()
+ mock_fs.exists.side_effect = [True, True] # input, output
+ mock_fs.rm = Mock()
+ mock_fs.makedirs = Mock()
+
+ # Mock file object that supports context manager protocol
+ mock_file = MagicMock()
+ mock_fs.open.return_value.__enter__ = Mock(return_value=mock_file)
+ mock_fs.open.return_value.__exit__ = Mock(return_value=None)
+
+ mock_url_to_fs.side_effect = [
+ (mock_fs, "/path/to/test.csv"),
+ (mock_fs, "/path/to/output"),
+ (mock_fs, "/path/to/output/metadata.json"), # ← Add 3rd call for metadata
+ ]
+
+ mock_read_csv.return_value = pl.DataFrame(
+ {
+ "label": ["doc1"],
+ "text": ["hello"],
+ }
+ )
+
+ # Mock the parquet write to prevent actual file I/O
+ with patch("classifai.indexers.main.pl.DataFrame.write_parquet"):
+ VectorStore(
+ file_name="/path/to/test.csv",
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ output_dir="/path/to/output",
+ overwrite=True,
+ skip_save=False,
+ )
+
+ mock_fs.rm.assert_called_once()
+ mock_fs.makedirs.assert_called()
+
+ @patch("classifai.indexers.main.check_deps")
+ @patch("classifai.indexers.main.fsspec.core.url_to_fs")
+ def test_init_output_dir_gs_path_without_gcsfs_raises_helpful_error(
+ self, mock_url_to_fs, mock_check_deps, mock_vectoriser
+ ):
+ """gs:// output_dir without gcsfs → OptionalDependencyError."""
+ mock_fs = Mock()
+ mock_fs.exists.return_value = True
+ mock_url_to_fs.side_effect = [
+ (mock_fs, "/path/to/test.csv"), # input OK
+ ImportError("gcsfs not installed"), # output fails
+ ]
+ mock_check_deps.side_effect = OptionalDependencyError("gcsfs required")
+
+ with pytest.raises(OptionalDependencyError) as exc_info:
+ VectorStore(
+ file_name="/path/to/test.csv",
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ output_dir="gs://bucket/output",
+ skip_save=False,
+ )
+ assert "gcsfs" in str(exc_info.value).lower()
+
+ # # ============================================================================
+ # # INDEX BUILDING TESTS (IndexBuildError, _create_vector_store_index)
+ # # ============================================================================
+
+ # class TestVectorStoreInitIndexBuilding:
+ # """Tests for index building during initialization."""
+
+ @patch("classifai.indexers.main.pl.read_csv")
+ @patch("classifai.indexers.main.fsspec.core.url_to_fs")
+ def test_init_csv_reads_correctly(self, mock_url_to_fs, mock_read_csv, mock_vectoriser):
+ """CSV reads, UUIDs assigned, vectoriser called."""
+ mock_fs = Mock()
+ mock_fs.exists.return_value = True
+ mock_url_to_fs.return_value = (mock_fs, "/path/to/test.csv")
+
+ # Mock CSV with 2 rows
+ mock_read_csv.return_value = pl.DataFrame(
+ {
+ "label": ["doc1", "doc2"],
+ "text": ["hello world", "goodbye world"],
+ }
+ )
+
+ vs = VectorStore(
+ file_name="/path/to/test.csv",
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ skip_save=True,
+ )
+
+ # Verify vectoriser was called
+ mock_vectoriser.transform.assert_called()
+ # Verify UUIDs were assigned
+ assert "uuid" in vs.vectors.columns
+ assert len(vs.vectors) == 2
+
+ @patch("classifai.indexers.main.pl.read_csv")
+ @patch("classifai.indexers.main.fsspec.core.url_to_fs")
+ def test_init_csv_reads_with_metadata(self, mock_url_to_fs, mock_read_csv, mock_vectoriser):
+ """CSV reads with metadata columns included."""
+ mock_fs = Mock()
+ mock_fs.exists.return_value = True
+ mock_url_to_fs.return_value = (mock_fs, "/path/to/test.csv")
+
+ mock_read_csv.return_value = pl.DataFrame(
+ {
+ "label": ["doc1", "doc2"],
+ "text": ["hello", "world"],
+ "source": ["web", "file"],
+ }
+ )
+
+ vs = VectorStore(
+ file_name="/path/to/test.csv",
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ meta_data={"source": str},
+ skip_save=True,
+ )
+
+ assert "source" in vs.vectors.columns
+ assert list(vs.vectors["source"]) == ["web", "file"]
+
+ @patch("classifai.indexers.main.pl.read_csv")
+ @patch("classifai.indexers.main.fsspec.core.url_to_fs")
+ def test_init_vectoriser_failure_wrapped_appropriately(self, mock_url_to_fs, mock_read_csv, mock_vectoriser):
+ """vectoriser.transform() exception → IndexBuildError with context."""
+ mock_fs = Mock()
+ mock_fs.exists.return_value = True
+ mock_url_to_fs.return_value = (mock_fs, "/path/to/test.csv")
+
+ mock_read_csv.return_value = pl.DataFrame(
+ {
+ "label": ["doc1"],
+ "text": ["hello"],
+ }
+ )
+
+ mock_vectoriser.transform.side_effect = RuntimeError("Vectoriser crashed")
+
+ with pytest.raises(IndexBuildError) as exc_info:
+ VectorStore(
+ file_name="/path/to/test.csv",
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ skip_save=True,
+ )
+ assert "Vectoriser.transform failed" in str(exc_info.value)
+
+ @patch("classifai.indexers.main.pl.read_csv")
+ @patch("classifai.indexers.main.fsspec.core.url_to_fs")
+ def test_init_embeddings_count_mismatch_raises_error(self, mock_url_to_fs, mock_read_csv, mock_vectoriser):
+ """Vectoriser returns wrong # embeddings → IndexBuildError."""
+ mock_fs = Mock()
+ mock_fs.exists.return_value = True
+ mock_url_to_fs.return_value = (mock_fs, "/path/to/test.csv")
+
+ mock_read_csv.return_value = pl.DataFrame(
+ {
+ "label": ["doc1", "doc2"],
+ "text": ["hello", "world"],
+ }
+ )
+
+ # Clear side_effect and set return_value
+ mock_vectoriser.transform.side_effect = None
+ mock_vectoriser.transform.return_value = np.array([[0.1, 0.2, 0.3]])
+
+ with pytest.raises(IndexBuildError) as exc_info:
+ VectorStore(
+ file_name="/path/to/test.csv",
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ skip_save=True,
+ )
+ assert "wrong number of embeddings" in str(exc_info.value).lower()
+
+ @patch("classifai.indexers.main.pl.read_csv")
+ @patch("classifai.indexers.main.fsspec.core.url_to_fs")
+ def test_init_empty_csv_raises_error(self, mock_url_to_fs, mock_read_csv, mock_vectoriser):
+ """Empty CSV (no documents) raises DataValidationError."""
+ mock_fs = Mock()
+ mock_fs.exists.return_value = True
+ mock_url_to_fs.return_value = (mock_fs, "/path/to/test.csv")
+
+ mock_read_csv.return_value = pl.DataFrame(
+ {
+ "label": [],
+ "text": [],
+ }
+ )
+
+ with pytest.raises(DataValidationError) as exc_info:
+ VectorStore(
+ file_name="/path/to/test.csv",
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ skip_save=True,
+ )
+ assert "no documents" in str(exc_info.value).lower()
+
+ @patch("classifai.indexers.main.pl.read_csv")
+ @patch("classifai.indexers.main.fsspec.core.url_to_fs")
+ def test_init_batch_processing_works(self, mock_url_to_fs, mock_read_csv, mock_vectoriser):
+ """Batch processing embeds texts in chunks."""
+ mock_fs = Mock()
+ mock_fs.exists.return_value = True
+ mock_url_to_fs.return_value = (mock_fs, "/path/to/test.csv")
+
+ # 5 documents with batch_size=2 → 3 batches
+ mock_read_csv.return_value = pl.DataFrame(
+ {
+ "label": ["doc1", "doc2", "doc3", "doc4", "doc5"],
+ "text": ["a", "b", "c", "d", "e"],
+ }
+ )
+
+ # Mock vectoriser to return correct number of embeddings per call
+ mock_vectoriser.transform.side_effect = [
+ np.array([[0.1, 0.2, 0.3], [0.4, 0.5, 0.6]]), # batch 1: 2 embeddings
+ np.array([[0.7, 0.8, 0.9], [1.0, 1.1, 1.2]]), # batch 2: 2 embeddings
+ np.array([[1.3, 1.4, 1.5]]), # batch 3: 1 embedding
+ ]
+
+ vs = VectorStore(
+ file_name="/path/to/test.csv",
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ batch_size=2,
+ skip_save=True,
+ )
+
+ # Verify vectoriser was called 3 times
+ assert mock_vectoriser.transform.call_count == 3
+ # Verify all embeddings were collected
+ assert len(vs.vectors) == 5
+
+
+# # ============================================================================
+# # SAVE/METADATA TESTS (skip_save flag)
+# # ============================================================================
+
+
+class TestVectorStoreInitSaveHandling:
+ """Tests for save/metadata handling and skip_save flag."""
+
+ @patch("classifai.indexers.main.pl.read_csv")
+ @patch("classifai.indexers.main.fsspec.core.url_to_fs")
+ def test_init_skip_save_true_no_files_written(self, mock_url_to_fs, mock_read_csv, mock_vectoriser):
+ """skip_save=True → no parquet/JSON written."""
+ mock_fs = Mock()
+ mock_fs.exists.return_value = True
+ mock_url_to_fs.return_value = (mock_fs, "/path/to/test.csv")
+
+ mock_read_csv.return_value = pl.DataFrame(
+ {
+ "label": ["doc1"],
+ "text": ["hello"],
+ }
+ )
+
+ vs = VectorStore(
+ file_name="/path/to/test.csv",
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ output_dir="/path/to/output",
+ skip_save=True,
+ )
+
+ # Verify no write operations on the filesystem
+ assert not mock_fs.open.called
+
+ @patch("classifai.indexers.main.pl.read_csv")
+ @patch("classifai.indexers.main.fsspec.core.url_to_fs")
+ def test_init_skip_save_false_files_written(self, mock_url_to_fs, mock_read_csv, mock_vectoriser):
+ """skip_save=False → parquet + metadata.json created."""
+ mock_fs = Mock()
+ mock_fs.exists.return_value = True
+ mock_fs.makedirs = Mock()
+ mock_fs.open = MagicMock()
+
+ # Add 3 calls: input file, output dir, metadata JSON write
+ mock_url_to_fs.side_effect = [
+ (mock_fs, "/path/to/test.csv"), # input
+ (mock_fs, "/path/to/output"), # output
+ (mock_fs, "/path/to/output/metadata.json"), # metadata save
+ ]
+
+ mock_read_csv.return_value = pl.DataFrame(
+ {
+ "label": ["doc1"],
+ "text": ["hello"],
+ }
+ )
+
+ with patch("classifai.indexers.main.pl.DataFrame.write_parquet") as mock_write_parquet:
+ vs = VectorStore(
+ file_name="/path/to/test.csv",
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ output_dir="/path/to/output",
+ overwrite=True,
+ skip_save=False,
+ )
+
+ # Verify parquet write was called
+ mock_write_parquet.assert_called_once()
+ # Verify JSON metadata write was attempted
+ mock_fs.open.assert_called()
+
+ # @patch("classifai.indexers.main.pl.read_csv")
+ # @patch("classifai.indexers.main.fsspec.core.url_to_fs")
+ # def test_init_metadata_json_contains_required_fields(self, mock_url_to_fs, mock_read_csv, mock_vectoriser):
+ # """Saved metadata.json contains all required fields."""
+ # mock_fs = Mock()
+ # mock_fs.exists.return_value = True
+ # mock_fs.makedirs = Mock()
+
+ # # Capture written data
+ # written_data = {}
+
+ # def mock_open_func(path, mode="w", encoding=None):
+ # """Mock file that captures write calls."""
+ # class MockFile:
+ # def __enter__(self):
+ # return self
+ # def __exit__(self, *args):
+ # pass
+ # def write(self, data):
+ # nonlocal written_data
+ # written_data = data
+ # return MockFile()
+
+ # mock_fs.open = mock_open_func
+
+ # # Add 4 calls: input file, output dir, metadata JSON write (path), metadata JSON write (actual)
+ # mock_url_to_fs.side_effect = [
+ # (mock_fs, "/path/to/test.csv"), # input
+ # (mock_fs, "/path/to/output"), # output
+ # (mock_fs, "/path/to/output/metadata.json"), # metadata save path resolve
+ # (mock_fs, "/path/to/output/metadata.json"), # metadata save actual write
+ # ]
+
+ # mock_read_csv.return_value = pl.DataFrame({
+ # "label": ["doc1"],
+ # "text": ["hello"],
+ # })
+
+ # with patch("classifai.indexers.main.pl.DataFrame.write_parquet"):
+ # vs = VectorStore(
+ # file_name="/path/to/test.csv",
+ # data_type="csv",
+ # vectoriser=mock_vectoriser,
+ # output_dir="/path/to/output",
+ # overwrite=True,
+ # skip_save=False,
+ # )
+
+ # # Verify metadata was written
+ # assert written_data, "No data was written to metadata file"
+ # metadata = json.loads(written_data)
+ # assert "vectoriser_class" in metadata
+ # assert "vector_shape" in metadata
+ # assert "num_vectors" in metadata
+ # assert "batch_size" in metadata
+ # assert "created_at" in metadata
+ # assert "meta_data" in metadata
+ # assert metadata["vectoriser_class"] == "MockVectoriser"
+ # assert metadata["vector_shape"] == 3
+ # assert metadata["num_vectors"] == 1
+
+
+# # ============================================================================
+# # QUIET MODE TESTS
+# # ============================================================================
+
+
+class TestVectorStoreInitQuietMode:
+ """Tests for quiet mode behavior."""
+
+ @patch("classifai.indexers.main.pl.read_csv")
+ @patch("classifai.indexers.main.fsspec.core.url_to_fs")
+ def test_init_quiet_mode_true_suppresses_progress(self, mock_url_to_fs, mock_read_csv, mock_vectoriser):
+ """quiet_mode=True → progress bars suppressed."""
+ mock_fs = Mock()
+ mock_fs.exists.return_value = True
+ mock_url_to_fs.return_value = (mock_fs, "/path/to/test.csv")
+
+ mock_read_csv.return_value = pl.DataFrame(
+ {
+ "label": ["doc1", "doc2", "doc3"],
+ "text": ["a", "b", "c"],
+ }
+ )
+
+ vs = VectorStore(
+ file_name="/path/to/test.csv",
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ quiet_mode=True,
+ skip_save=True,
+ )
+
+ # When quiet_mode=True, classifai_tqdm should be identity function (no wrapping)
+ assert vs.classifai_tqdm([1, 2, 3]) == [1, 2, 3]
+
+ @patch("classifai.indexers.main.pl.read_csv")
+ @patch("classifai.indexers.main.fsspec.core.url_to_fs")
+ def test_init_quiet_mode_false_shows_progress(self, mock_url_to_fs, mock_read_csv, mock_vectoriser):
+ """quiet_mode=False → progress bars shown (tqdm enabled)."""
+ mock_fs = Mock()
+ mock_fs.exists.return_value = True
+ mock_url_to_fs.return_value = (mock_fs, "/path/to/test.csv")
+
+ mock_read_csv.return_value = pl.DataFrame(
+ {
+ "label": ["doc1"],
+ "text": ["hello"],
+ }
+ )
+
+ vs = VectorStore(
+ file_name="/path/to/test.csv",
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ quiet_mode=False,
+ skip_save=True,
+ )
+
+ # When quiet_mode=False, classifai_tqdm should be tqdm
+ from tqdm.autonotebook import tqdm
+
+ assert vs.classifai_tqdm == tqdm
+
+
+# ============================================================================
+# INTEGRATION TESTS (Real temp files)
+# ============================================================================
+
+
+class TestVectorStoreInitIntegration:
+ """Integration tests using real temporary files."""
+
+ def test_init_with_real_csv_file(self, temp_csv_file, mock_vectoriser, temp_output_dir):
+ """Full initialization with real CSV file and temp output directory."""
+ vs = VectorStore(
+ file_name=temp_csv_file,
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ output_dir=temp_output_dir,
+ skip_save=False,
+ )
+
+ # Verify VectorStore was created successfully
+ assert vs.file_name == temp_csv_file
+ assert vs.vectors is not None
+ assert len(vs.vectors) == 2
+ assert vs.vector_shape == 3 # embeddings have 3 dimensions
+ assert vs.num_vectors == 2
+ assert vs.vectoriser_class == "MockVectoriser"
+
+ # Verify files were saved
+ assert os.path.exists(os.path.join(temp_output_dir, "vectors.parquet"))
+ assert os.path.exists(os.path.join(temp_output_dir, "metadata.json"))
+
+ # Verify metadata.json content
+ with open(os.path.join(temp_output_dir, "metadata.json")) as f:
+ metadata = json.load(f)
+ assert metadata["vectoriser_class"] == "MockVectoriser"
+ assert metadata["vector_shape"] == 3
+ assert metadata["num_vectors"] == 2
+
+ def test_init_with_metadata_columns(self, temp_csv_with_metadata, mock_vectoriser, temp_output_dir):
+ """Initialization with metadata columns."""
+ vs = VectorStore(
+ file_name=temp_csv_with_metadata,
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ meta_data={"source": str},
+ output_dir=temp_output_dir,
+ skip_save=False,
+ )
+
+ assert "source" in vs.vectors.columns
+ assert vs.meta_data == {"source": str}
+
+ # Verify metadata is saved correctly
+ with open(os.path.join(temp_output_dir, "metadata.json")) as f:
+ metadata = json.load(f)
+ assert "source" in metadata["meta_data"]
diff --git a/tests/test_indexers/test_vectorstore_metadata.py b/tests/test_indexers/test_vectorstore_metadata.py
new file mode 100644
index 0000000..cd73251
--- /dev/null
+++ b/tests/test_indexers/test_vectorstore_metadata.py
@@ -0,0 +1,858 @@
+"""Unit tests for VectorStore metadata serialization and deserialization."""
+
+from __future__ import annotations
+
+import json
+import time
+from unittest.mock import Mock
+
+import numpy as np
+import pytest
+
+from classifai.exceptions import (
+ DataValidationError,
+ IndexBuildError,
+)
+from classifai.indexers import VectorStore
+from classifai.vectorisers import VectoriserBase
+
+# ============================================================================
+# FIXTURES
+# ============================================================================
+
+
+@pytest.fixture
+def mock_vectoriser():
+ """Mock VectoriserBase that returns predictable embeddings."""
+ mock = Mock(spec=VectoriserBase)
+
+ def transform_side_effect(texts):
+ """Return one embedding per text, all with shape (3,)."""
+ num_texts = len(texts)
+ return np.array([np.linspace(0.1, 0.3, 3) + (i * 0.3) for i in range(num_texts)])
+
+ mock.transform.side_effect = transform_side_effect
+ mock.__class__.__name__ = "MockVectoriser"
+ return mock
+
+
+@pytest.fixture
+def initialized_vectorstore(mock_vectoriser, tmp_path):
+ """Create a fully initialized VectorStore with test data."""
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text("label,text\ndoc1,hello world\ndoc2,goodbye world\n")
+
+ vs = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ output_dir=str(tmp_path / "output"),
+ skip_save=True,
+ )
+
+ return vs
+
+
+# ============================================================================
+# METADATA SERIALIZATION TESTS (_save_metadata)
+# ============================================================================
+
+
+class TestVectorStoreMetadataSerialization:
+ """Tests for metadata serialization (_save_metadata)."""
+
+ def test_save_metadata_json_file_created_at_correct_path(self, mock_vectoriser, tmp_path):
+ """JSON file created at correct path."""
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text("label,text\ndoc1,hello\ndoc2,world\n")
+
+ output_dir = tmp_path / "output"
+
+ vs = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ output_dir=str(output_dir),
+ skip_save=False,
+ )
+
+ metadata_path = output_dir / "metadata.json"
+ assert metadata_path.exists()
+
+ def test_save_metadata_contains_vectoriser_class(self, mock_vectoriser, tmp_path):
+ """Metadata contains vectoriser_class field."""
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text("label,text\ndoc1,hello\ndoc2,world\n")
+
+ output_dir = tmp_path / "output"
+
+ vs = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ output_dir=str(output_dir),
+ skip_save=False,
+ )
+
+ metadata_path = output_dir / "metadata.json"
+ with open(metadata_path) as f:
+ metadata = json.load(f)
+
+ assert "vectoriser_class" in metadata
+ assert metadata["vectoriser_class"] == "MockVectoriser"
+
+ def test_save_metadata_contains_vector_shape(self, mock_vectoriser, tmp_path):
+ """Metadata contains vector_shape field."""
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text("label,text\ndoc1,hello\ndoc2,world\n")
+
+ output_dir = tmp_path / "output"
+
+ vs = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ output_dir=str(output_dir),
+ skip_save=False,
+ )
+
+ metadata_path = output_dir / "metadata.json"
+ with open(metadata_path) as f:
+ metadata = json.load(f)
+
+ assert "vector_shape" in metadata
+ assert metadata["vector_shape"] == 3
+
+ def test_save_metadata_contains_num_vectors(self, mock_vectoriser, tmp_path):
+ """Metadata contains num_vectors field."""
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text("label,text\ndoc1,hello\ndoc2,world\n")
+
+ output_dir = tmp_path / "output"
+
+ vs = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ output_dir=str(output_dir),
+ skip_save=False,
+ )
+
+ metadata_path = output_dir / "metadata.json"
+ with open(metadata_path) as f:
+ metadata = json.load(f)
+
+ assert "num_vectors" in metadata
+ assert metadata["num_vectors"] == 2
+
+ def test_save_metadata_contains_batch_size(self, mock_vectoriser, tmp_path):
+ """Metadata contains batch_size field."""
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text("label,text\ndoc1,hello\ndoc2,world\n")
+
+ output_dir = tmp_path / "output"
+
+ vs = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ batch_size=64,
+ output_dir=str(output_dir),
+ skip_save=False,
+ )
+
+ metadata_path = output_dir / "metadata.json"
+ with open(metadata_path) as f:
+ metadata = json.load(f)
+
+ assert "batch_size" in metadata
+ assert metadata["batch_size"] == 64
+
+ def test_save_metadata_contains_created_at(self, mock_vectoriser, tmp_path):
+ """Metadata contains created_at field."""
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text("label,text\ndoc1,hello\ndoc2,world\n")
+
+ output_dir = tmp_path / "output"
+
+ before_creation = time.time()
+ vs = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ output_dir=str(output_dir),
+ skip_save=False,
+ )
+ after_creation = time.time()
+
+ metadata_path = output_dir / "metadata.json"
+ with open(metadata_path) as f:
+ metadata = json.load(f)
+
+ assert "created_at" in metadata
+ assert isinstance(metadata["created_at"], (int, float))
+ assert before_creation <= metadata["created_at"] <= after_creation
+
+ def test_save_metadata_contains_meta_data_field(self, mock_vectoriser, tmp_path):
+ """Metadata contains meta_data field."""
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text("label,text\ndoc1,hello\ndoc2,world\n")
+
+ output_dir = tmp_path / "output"
+
+ vs = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ output_dir=str(output_dir),
+ skip_save=False,
+ )
+
+ metadata_path = output_dir / "metadata.json"
+ with open(metadata_path) as f:
+ metadata = json.load(f)
+
+ assert "meta_data" in metadata
+ assert isinstance(metadata["meta_data"], dict)
+
+ def test_save_metadata_type_information_preserved_str(self, mock_vectoriser, tmp_path):
+ """Type information preserved (str types → string names)."""
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text("label,text,source\ndoc1,hello,src1\ndoc2,world,src2\n")
+
+ output_dir = tmp_path / "output"
+
+ vs = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ meta_data={"source": str},
+ output_dir=str(output_dir),
+ skip_save=False,
+ )
+
+ metadata_path = output_dir / "metadata.json"
+ with open(metadata_path) as f:
+ metadata = json.load(f)
+
+ assert "meta_data" in metadata
+ assert "source" in metadata["meta_data"]
+ assert metadata["meta_data"]["source"] == "str"
+
+ def test_save_metadata_type_information_preserved_int(self, mock_vectoriser, tmp_path):
+ """Type information preserved for int types."""
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text("label,text,count\ndoc1,hello,5\ndoc2,world,10\n")
+
+ output_dir = tmp_path / "output"
+
+ vs = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ meta_data={"count": int},
+ output_dir=str(output_dir),
+ skip_save=False,
+ )
+
+ metadata_path = output_dir / "metadata.json"
+ with open(metadata_path) as f:
+ metadata = json.load(f)
+
+ assert metadata["meta_data"]["count"] == "int"
+
+ def test_save_metadata_type_information_preserved_float(self, mock_vectoriser, tmp_path):
+ """Type information preserved for float types."""
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text("label,text,score\ndoc1,hello,0.5\ndoc2,world,0.8\n")
+
+ output_dir = tmp_path / "output"
+
+ vs = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ meta_data={"score": float},
+ output_dir=str(output_dir),
+ skip_save=False,
+ )
+
+ metadata_path = output_dir / "metadata.json"
+ with open(metadata_path) as f:
+ metadata = json.load(f)
+
+ assert metadata["meta_data"]["score"] == "float"
+
+ def test_save_metadata_multiple_types_preserved(self, mock_vectoriser, tmp_path):
+ """Type information preserved for multiple metadata columns."""
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text("label,text,source,count,score\ndoc1,hello,src1,5,0.5\ndoc2,world,src2,10,0.8\n")
+
+ output_dir = tmp_path / "output"
+
+ vs = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ meta_data={"source": str, "count": int, "score": float},
+ output_dir=str(output_dir),
+ skip_save=False,
+ )
+
+ metadata_path = output_dir / "metadata.json"
+ with open(metadata_path) as f:
+ metadata = json.load(f)
+
+ assert metadata["meta_data"]["source"] == "str"
+ assert metadata["meta_data"]["count"] == "int"
+ assert metadata["meta_data"]["score"] == "float"
+
+ def test_save_metadata_valid_json_format(self, mock_vectoriser, tmp_path):
+ """Metadata is valid JSON format."""
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text("label,text\ndoc1,hello\ndoc2,world\n")
+
+ output_dir = tmp_path / "output"
+
+ vs = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ output_dir=str(output_dir),
+ skip_save=False,
+ )
+
+ metadata_path = output_dir / "metadata.json"
+ with open(metadata_path) as f:
+ # Should not raise JSONDecodeError
+ metadata = json.load(f)
+
+ assert isinstance(metadata, dict)
+
+ def test_save_metadata_empty_meta_data(self, mock_vectoriser, tmp_path):
+ """Empty meta_data dict serialized correctly."""
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text("label,text\ndoc1,hello\ndoc2,world\n")
+
+ output_dir = tmp_path / "output"
+
+ vs = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ meta_data=None,
+ output_dir=str(output_dir),
+ skip_save=False,
+ )
+
+ metadata_path = output_dir / "metadata.json"
+ with open(metadata_path) as f:
+ metadata = json.load(f)
+
+ assert metadata["meta_data"] == {}
+
+ def test_save_metadata_path_must_be_string(self, initialized_vectorstore):
+ """Path argument must be string."""
+ with pytest.raises(DataValidationError) as exc_info:
+ initialized_vectorstore._save_metadata(path=None)
+
+ assert "path" in str(exc_info.value).lower()
+
+ def test_save_metadata_path_must_be_non_empty(self, initialized_vectorstore):
+ """Path argument must be non-empty string."""
+ with pytest.raises(DataValidationError) as exc_info:
+ initialized_vectorstore._save_metadata(path="")
+
+ assert "path" in str(exc_info.value).lower()
+
+ def test_save_metadata_invalid_path_type(self, initialized_vectorstore):
+ """Path as non-string type raises error."""
+ with pytest.raises(DataValidationError) as exc_info:
+ initialized_vectorstore._save_metadata(path=123)
+
+ assert "path" in str(exc_info.value).lower()
+
+
+# ============================================================================
+# METADATA LOADING TESTS (from_filespace related)
+# ============================================================================
+
+
+class TestVectorStoreMetadataLoading:
+ """Tests for metadata loading and deserialization."""
+
+ def test_load_metadata_file_read_and_parsed(self, mock_vectoriser, tmp_path):
+ """Metadata file read and parsed correctly from saved vectorstore."""
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text("label,text\ndoc1,hello\ndoc2,world\n")
+
+ output_dir = tmp_path / "output"
+
+ vs_saved = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ output_dir=str(output_dir),
+ skip_save=False,
+ )
+
+ metadata_path = output_dir / "metadata.json"
+ with open(metadata_path) as f:
+ metadata = json.load(f)
+
+ assert isinstance(metadata, dict)
+ assert "vectoriser_class" in metadata
+
+ def test_load_metadata_required_keys_validated(self, mock_vectoriser, tmp_path):
+ """Required keys validated when loading metadata."""
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text("label,text\ndoc1,hello\ndoc2,world\n")
+
+ output_dir = tmp_path / "output"
+
+ vs_saved = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ output_dir=str(output_dir),
+ skip_save=False,
+ )
+
+ metadata_path = output_dir / "metadata.json"
+ with open(metadata_path) as f:
+ metadata = json.load(f)
+
+ required_keys = ["vectoriser_class", "vector_shape", "num_vectors", "created_at", "meta_data"]
+ for key in required_keys:
+ assert key in metadata
+
+ def test_load_metadata_type_deserialization_str(self, mock_vectoriser, tmp_path):
+ """Type deserialization works for str."""
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text("label,text,source\ndoc1,hello,src1\ndoc2,world,src2\n")
+
+ output_dir = tmp_path / "output"
+
+ vs_saved = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ meta_data={"source": str},
+ output_dir=str(output_dir),
+ skip_save=False,
+ )
+
+ # Load it back
+ vs_loaded = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ # Type should be deserialized back to str
+ assert vs_loaded.meta_data["source"] == "str"
+
+ def test_load_metadata_type_deserialization_int(self, mock_vectoriser, tmp_path):
+ """Type deserialization works for int."""
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text("label,text,count\ndoc1,hello,5\ndoc2,world,10\n")
+
+ output_dir = tmp_path / "output"
+
+ vs_saved = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ meta_data={"count": int},
+ output_dir=str(output_dir),
+ skip_save=False,
+ )
+
+ # Load it back
+ vs_loaded = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ # Type should be deserialized back to int
+ assert vs_loaded.meta_data["count"] == "int"
+
+ def test_load_metadata_type_deserialization_float(self, mock_vectoriser, tmp_path):
+ """Type deserialization works for float."""
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text("label,text,score\ndoc1,hello,0.5\ndoc2,world,0.8\n")
+
+ output_dir = tmp_path / "output"
+
+ vs_saved = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ meta_data={"score": float},
+ output_dir=str(output_dir),
+ skip_save=False,
+ )
+
+ # Load it back
+ vs_loaded = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ # Type should be deserialized back to float
+ assert vs_loaded.meta_data["score"] == "float"
+
+ def test_load_metadata_backwards_compatibility_missing_batch_size(self, mock_vectoriser, tmp_path):
+ """Backwards compatibility with v1.0.0 (missing batch_size)."""
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text("label,text\ndoc1,hello\ndoc2,world\n")
+
+ output_dir = tmp_path / "output"
+
+ vs_saved = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ output_dir=str(output_dir),
+ skip_save=False,
+ )
+
+ # Remove batch_size from metadata to simulate v1.0.0
+ metadata_path = output_dir / "metadata.json"
+ with open(metadata_path) as f:
+ metadata = json.load(f)
+
+ del metadata["batch_size"]
+
+ with open(metadata_path, "w") as f:
+ json.dump(metadata, f)
+
+ # Should load without error
+ vs_loaded = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ assert vs_loaded.batch_size == 128 # default
+
+ def test_load_metadata_batch_size_override_works(self, mock_vectoriser, tmp_path):
+ """batch_size override works when loading."""
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text("label,text\ndoc1,hello\ndoc2,world\n")
+
+ output_dir = tmp_path / "output"
+
+ vs_saved = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ batch_size=64,
+ output_dir=str(output_dir),
+ skip_save=False,
+ )
+
+ # Load with different batch_size
+ vs_loaded = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ batch_size=256,
+ )
+
+ assert vs_loaded.batch_size == 256
+
+ def test_load_metadata_default_batch_size_when_missing(self, mock_vectoriser, tmp_path):
+ """Default batch_size used when missing from metadata."""
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text("label,text\ndoc1,hello\ndoc2,world\n")
+
+ output_dir = tmp_path / "output"
+
+ vs_saved = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ output_dir=str(output_dir),
+ skip_save=False,
+ )
+
+ # Remove batch_size from metadata
+ metadata_path = output_dir / "metadata.json"
+ with open(metadata_path) as f:
+ metadata = json.load(f)
+
+ del metadata["batch_size"]
+
+ with open(metadata_path, "w") as f:
+ json.dump(metadata, f)
+
+ # Load without specifying batch_size
+ vs_loaded = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ # Should use default
+ assert vs_loaded.batch_size == 128
+
+ def test_load_metadata_warning_logged_for_missing_batch_size(self, mock_vectoriser, tmp_path):
+ """Warning logged when batch_size missing from metadata."""
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text("label,text\ndoc1,hello\ndoc2,world\n")
+
+ output_dir = tmp_path / "output"
+
+ vs_saved = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ output_dir=str(output_dir),
+ skip_save=False,
+ )
+
+ # Remove batch_size from metadata
+ metadata_path = output_dir / "metadata.json"
+ with open(metadata_path) as f:
+ metadata = json.load(f)
+
+ del metadata["batch_size"]
+
+ with open(metadata_path, "w") as f:
+ json.dump(metadata, f)
+
+ # Load - should succeed even with missing batch_size
+ vs_loaded = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ # Verify it uses the default
+ assert vs_loaded.batch_size == 128
+ # Verify the vectorstore is still functional
+ assert vs_loaded.vector_shape == 3
+ assert vs_loaded.num_vectors == 2
+
+ def test_load_metadata_preserves_vectoriser_class(self, mock_vectoriser, tmp_path):
+ """Vectoriser class name preserved through save/load cycle."""
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text("label,text\ndoc1,hello\ndoc2,world\n")
+
+ output_dir = tmp_path / "output"
+
+ vs_saved = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ output_dir=str(output_dir),
+ skip_save=False,
+ )
+
+ vs_loaded = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ assert vs_loaded.vectoriser_class == "MockVectoriser"
+
+ def test_load_metadata_preserves_num_vectors(self, mock_vectoriser, tmp_path):
+ """num_vectors preserved through save/load cycle."""
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text("label,text\ndoc1,hello\ndoc2,world\ndoc3,test\n")
+
+ output_dir = tmp_path / "output"
+
+ vs_saved = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ output_dir=str(output_dir),
+ skip_save=False,
+ )
+
+ vs_loaded = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ assert vs_loaded.num_vectors == 3
+
+ def test_load_metadata_preserves_vector_shape(self, mock_vectoriser, tmp_path):
+ """vector_shape preserved through save/load cycle."""
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text("label,text\ndoc1,hello\ndoc2,world\n")
+
+ output_dir = tmp_path / "output"
+
+ vs_saved = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ output_dir=str(output_dir),
+ skip_save=False,
+ )
+
+ vs_loaded = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ assert vs_loaded.vector_shape == 3
+
+
+# ============================================================================
+# EDGE CASE TESTS
+# ============================================================================
+
+
+class TestVectorStoreMetadataEdgeCases:
+ """Tests for edge cases in metadata handling."""
+
+ def test_save_metadata_with_large_num_vectors(self, mock_vectoriser, tmp_path):
+ """Metadata serialization with large num_vectors."""
+ csv_lines = ["label,text"]
+ for i in range(100):
+ csv_lines.append(f"doc{i},text {i}")
+
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text("\n".join(csv_lines))
+
+ output_dir = tmp_path / "output"
+
+ vs = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ output_dir=str(output_dir),
+ skip_save=False,
+ )
+
+ metadata_path = output_dir / "metadata.json"
+ with open(metadata_path) as f:
+ metadata = json.load(f)
+
+ assert metadata["num_vectors"] == 100
+
+ def test_save_metadata_with_large_vector_shape(self, mock_vectoriser, tmp_path):
+ """Metadata serialization with large vector_shape."""
+ # Create vectoriser that returns larger embeddings
+ mock_large = Mock(spec=VectoriserBase)
+
+ def transform_large(texts):
+ return np.array([np.random.rand(512) for _ in texts])
+
+ mock_large.transform.side_effect = transform_large
+ mock_large.__class__.__name__ = "LargeVectoriser"
+
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text("label,text\ndoc1,hello\ndoc2,world\n")
+
+ output_dir = tmp_path / "output"
+
+ vs = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_large,
+ output_dir=str(output_dir),
+ skip_save=False,
+ )
+
+ metadata_path = output_dir / "metadata.json"
+ with open(metadata_path) as f:
+ metadata = json.load(f)
+
+ assert metadata["vector_shape"] == 512
+
+ def test_save_metadata_indentation_readable(self, mock_vectoriser, tmp_path):
+ """Metadata JSON is indented and readable."""
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text("label,text\ndoc1,hello\ndoc2,world\n")
+
+ output_dir = tmp_path / "output"
+
+ vs = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ output_dir=str(output_dir),
+ skip_save=False,
+ )
+
+ metadata_path = output_dir / "metadata.json"
+ with open(metadata_path) as f:
+ content = f.read()
+
+ # Check for indentation (newlines and spaces)
+ assert "\n" in content
+ assert " " in content or "\t" in content
+
+ def test_load_metadata_malformed_json_raises_error(self, mock_vectoriser, tmp_path):
+ """Malformed JSON raises IndexBuildError."""
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text("label,text\ndoc1,hello\ndoc2,world\n")
+
+ output_dir = tmp_path / "output"
+
+ vs_saved = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ output_dir=str(output_dir),
+ skip_save=False,
+ )
+
+ # Corrupt metadata.json
+ metadata_path = output_dir / "metadata.json"
+ with open(metadata_path, "w") as f:
+ f.write("{invalid json")
+
+ # Should raise error when loading
+ with pytest.raises((IndexBuildError, DataValidationError)):
+ VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ def test_save_metadata_special_characters_in_meta_data_keys(self, mock_vectoriser, tmp_path):
+ """Metadata keys with special characters handled correctly."""
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text("label,text,col_with_underscore\ndoc1,hello,val1\ndoc2,world,val2\n")
+
+ output_dir = tmp_path / "output"
+
+ vs = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ meta_data={"col_with_underscore": str},
+ output_dir=str(output_dir),
+ skip_save=False,
+ )
+
+ metadata_path = output_dir / "metadata.json"
+ with open(metadata_path) as f:
+ metadata = json.load(f)
+
+ assert "col_with_underscore" in metadata["meta_data"]
+
+ def test_load_metadata_with_multiple_meta_data_types(self, mock_vectoriser, tmp_path):
+ """Metadata with multiple column types round-trips correctly."""
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text("label,text,source,count,score\ndoc1,hello,src1,5,0.5\ndoc2,world,src2,10,0.8\n")
+
+ output_dir = tmp_path / "output"
+
+ vs_saved = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ meta_data={"source": str, "count": int, "score": float},
+ output_dir=str(output_dir),
+ skip_save=False,
+ )
+
+ vs_loaded = VectorStore.from_filespace(
+ folder_path=str(output_dir),
+ vectoriser=mock_vectoriser,
+ )
+
+ assert vs_loaded.meta_data["source"] == "str"
+ assert vs_loaded.meta_data["count"] == "int"
+ assert vs_loaded.meta_data["score"] == "float"
diff --git a/tests/test_indexers/test_vectorstore_reverse.py b/tests/test_indexers/test_vectorstore_reverse.py
new file mode 100644
index 0000000..dfd94f0
--- /dev/null
+++ b/tests/test_indexers/test_vectorstore_reverse.py
@@ -0,0 +1,878 @@
+"""Unit tests for VectorStore.reverse_search() method."""
+
+from __future__ import annotations
+
+from unittest.mock import Mock
+
+import numpy as np
+import pytest
+
+from classifai.exceptions import (
+ ClassifaiError,
+ DataValidationError,
+ HookError,
+)
+from classifai.indexers import VectorStore
+from classifai.indexers.dataclasses import (
+ VectorStoreReverseSearchInput,
+ VectorStoreReverseSearchOutput,
+)
+from classifai.indexers.hooks import HookBase
+from classifai.vectorisers import VectoriserBase
+
+# ============================================================================
+# FIXTURES
+# ============================================================================
+
+
+@pytest.fixture
+def mock_vectoriser():
+ """Mock VectoriserBase that returns predictable embeddings."""
+ mock = Mock(spec=VectoriserBase)
+
+ def transform_side_effect(texts):
+ """Return one embedding per text, all with shape (3,)."""
+ num_texts = len(texts)
+ return np.array([np.linspace(0.1, 0.3, 3) + (i * 0.3) for i in range(num_texts)])
+
+ mock.transform.side_effect = transform_side_effect
+ mock.__class__.__name__ = "MockVectoriser"
+ return mock
+
+
+@pytest.fixture
+def initialized_vectorstore(mock_vectoriser, tmp_path):
+ """Create a fully initialized VectorStore with test data."""
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text(
+ "label,text\n"
+ "category_a,document 1\n"
+ "category_a,document 2\n"
+ "category_b,document 3\n"
+ "category_b,document 4\n"
+ "category_c,document 5\n"
+ )
+
+ vs = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ output_dir=str(tmp_path / "output"),
+ skip_save=True,
+ )
+
+ return vs
+
+
+# ============================================================================
+# INPUT VALIDATION TESTS
+# ============================================================================
+
+
+class TestVectorStoreReverseSearchInputValidation:
+ """Tests for input validation in reverse_search() method."""
+
+ def test_reverse_search_query_must_be_vectorstore_reverse_search_input(self, initialized_vectorstore):
+ """Query must be VectorStoreReverseSearchInput object."""
+ with pytest.raises(DataValidationError) as exc_info:
+ initialized_vectorstore.reverse_search(query="not a VectorStoreReverseSearchInput")
+
+ assert "VectorStoreReverseSearchInput" in str(exc_info.value)
+
+ def test_reverse_search_query_none_raises_error(self, initialized_vectorstore):
+ """query=None raises DataValidationError."""
+ with pytest.raises(DataValidationError) as exc_info:
+ initialized_vectorstore.reverse_search(query=None)
+
+ assert "VectorStoreReverseSearchInput" in str(exc_info.value)
+
+ def test_reverse_search_query_dict_raises_error(self, initialized_vectorstore):
+ """Query as dict (not VectorStoreReverseSearchInput) raises DataValidationError."""
+ with pytest.raises(DataValidationError) as exc_info:
+ initialized_vectorstore.reverse_search(query={"id": ["1"], "doc_label": ["category_a"]})
+
+ assert "VectorStoreReverseSearchInput" in str(exc_info.value)
+
+ def test_reverse_search_query_list_raises_error(self, initialized_vectorstore):
+ """Query as list raises DataValidationError."""
+ with pytest.raises(DataValidationError) as exc_info:
+ initialized_vectorstore.reverse_search(query=["category_a", "category_b"])
+
+ assert "VectorStoreReverseSearchInput" in str(exc_info.value)
+
+ def test_reverse_search_max_n_results_must_be_positive_int(self, initialized_vectorstore):
+ """max_n_results must be int >= 1 or -1."""
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["1"],
+ "doc_label": ["category_a"],
+ }
+ )
+
+ with pytest.raises(DataValidationError) as exc_info:
+ initialized_vectorstore.reverse_search(query=query, max_n_results=0)
+
+ assert "max_n_results" in str(exc_info.value).lower()
+
+ def test_reverse_search_max_n_results_negative_not_minus_one_raises_error(self, initialized_vectorstore):
+ """max_n_results < -1 raises DataValidationError."""
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["1"],
+ "doc_label": ["category_a"],
+ }
+ )
+
+ with pytest.raises(DataValidationError) as exc_info:
+ initialized_vectorstore.reverse_search(query=query, max_n_results=-5)
+
+ assert "max_n_results" in str(exc_info.value).lower()
+
+ def test_reverse_search_max_n_results_non_int_raises_error(self, initialized_vectorstore):
+ """max_n_results as string raises DataValidationError."""
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["1"],
+ "doc_label": ["category_a"],
+ }
+ )
+
+ with pytest.raises(DataValidationError) as exc_info:
+ initialized_vectorstore.reverse_search(query=query, max_n_results="10")
+
+ assert "max_n_results" in str(exc_info.value).lower()
+
+ def test_reverse_search_max_n_results_minus_one_is_valid(self, initialized_vectorstore):
+ """max_n_results=-1 is valid (means return all)."""
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["1"],
+ "doc_label": ["category_a"],
+ }
+ )
+
+ # Should not raise
+ result = initialized_vectorstore.reverse_search(query=query, max_n_results=-1)
+ assert isinstance(result, VectorStoreReverseSearchOutput)
+
+ def test_reverse_search_empty_query_raises_error(self, initialized_vectorstore):
+ """Empty query DataFrame raises DataValidationError."""
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": [],
+ "doc_label": [],
+ }
+ )
+
+ with pytest.raises(DataValidationError) as exc_info:
+ initialized_vectorstore.reverse_search(query=query)
+
+ assert "empty" in str(exc_info.value).lower()
+
+ # TODO: implement if/when partial_match parameter checks are added
+ # def test_reverse_search_partial_match_must_be_bool(self, initialized_vectorstore):
+ # """partial_match must be boolean."""
+ # query = VectorStoreReverseSearchInput.from_data({
+ # "id": ["1"],
+ # "doc_label": ["category_a"],
+ # })
+
+ with pytest.raises((DataValidationError, TypeError)):
+ initialized_vectorstore.reverse_search(query=query, partial_match="yes")
+
+ def test_reverse_search_partial_match_default_false(self, initialized_vectorstore):
+ """partial_match defaults to False."""
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["1"],
+ "doc_label": ["category_a"],
+ }
+ )
+
+ # Should not raise - defaults to False
+ result = initialized_vectorstore.reverse_search(query=query)
+ assert isinstance(result, VectorStoreReverseSearchOutput)
+
+
+# ============================================================================
+# REVERSE SEARCH OPERATION TESTS
+# ============================================================================
+
+
+class TestVectorStoreReverseSearchOperation:
+ """Tests for the core reverse search operation."""
+
+ def test_reverse_search_exact_label_matching_works(self, initialized_vectorstore):
+ """Exact label matching works (default)."""
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "doc_label": ["category_a"],
+ }
+ )
+
+ result = initialized_vectorstore.reverse_search(query=query, max_n_results=10)
+
+ assert isinstance(result, VectorStoreReverseSearchOutput)
+ # Should find 2 documents with exact label "category_a"
+ assert len(result) == 2
+ assert all(label == "category_a" for label in result.doc_label)
+
+ def test_reverse_search_partial_matching_when_enabled(self, initialized_vectorstore):
+ """Partial matching (prefix) works when partial_match=True."""
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "doc_label": ["category"],
+ }
+ )
+
+ result = initialized_vectorstore.reverse_search(
+ query=query,
+ max_n_results=10,
+ partial_match=True,
+ )
+
+ # Should find all documents with labels starting with "category"
+ assert len(result) >= 5 # All 5 documents start with "category"
+
+ def test_reverse_search_exact_matching_excludes_partial(self, initialized_vectorstore):
+ """Exact matching excludes partial matches (default behavior)."""
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "doc_label": ["category"],
+ }
+ )
+
+ result = initialized_vectorstore.reverse_search(
+ query=query,
+ max_n_results=10,
+ partial_match=False,
+ )
+
+ # Should find 0 documents (no exact match for "category")
+ assert len(result) == 0
+
+ def test_reverse_search_max_n_results_limits_results(self, initialized_vectorstore):
+ """max_n_results limits results per query."""
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "doc_label": ["category_a"],
+ }
+ )
+
+ result = initialized_vectorstore.reverse_search(query=query, max_n_results=1)
+
+ # Should return only 1 result even though there are 2 matches
+ assert len(result) == 1
+
+ def test_reverse_search_max_n_results_minus_one_returns_all(self, initialized_vectorstore):
+ """max_n_results=-1 returns all matches."""
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "doc_label": ["category_a"],
+ }
+ )
+
+ result = initialized_vectorstore.reverse_search(query=query, max_n_results=-1)
+
+ # Should return all 2 matches
+ assert len(result) == 2
+
+ def test_reverse_search_multiple_queries_process_correctly(self, initialized_vectorstore):
+ """Multiple queries process correctly."""
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1", "q2"],
+ "doc_label": ["category_a", "category_b"],
+ }
+ )
+
+ result = initialized_vectorstore.reverse_search(query=query, max_n_results=10)
+
+ # Should find 2 + 2 = 4 documents total
+ assert len(result) == 4
+ assert list(result.id[:2]) == ["q1", "q1"]
+ assert list(result.id[2:]) == ["q2", "q2"]
+
+ def test_reverse_search_includes_required_columns(self, initialized_vectorstore):
+ """Output includes all required columns."""
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "doc_label": ["category_a"],
+ }
+ )
+
+ result = initialized_vectorstore.reverse_search(query=query, max_n_results=10)
+
+ required_columns = ["id", "searched_doc_label", "doc_label", "doc_text"]
+ for col in required_columns:
+ assert col in result.columns
+
+ def test_reverse_search_empty_result_set_returns_empty_dataframe(self, initialized_vectorstore):
+ """Empty result set returns empty DataFrame with correct schema."""
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "doc_label": ["nonexistent_label"],
+ }
+ )
+
+ result = initialized_vectorstore.reverse_search(query=query, max_n_results=10)
+
+ assert isinstance(result, VectorStoreReverseSearchOutput)
+ assert len(result) == 0
+ # Verify schema is still correct
+ required_columns = ["id", "searched_doc_label", "doc_label", "doc_text"]
+ for col in required_columns:
+ assert col in result.columns
+
+ def test_reverse_search_results_sorted_by_id_and_label(self, initialized_vectorstore):
+ """Results are sorted by id and searched_doc_label."""
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q2", "q1", "q3"],
+ "doc_label": ["category_b", "category_a", "category_c"],
+ }
+ )
+
+ result = initialized_vectorstore.reverse_search(query=query, max_n_results=10)
+
+ # Results should be sorted
+ ids = list(result.id)
+ # Verify ids are in order
+ assert isinstance(ids, list)
+
+ def test_reverse_search_case_sensitive_matching(self, initialized_vectorstore):
+ """Label matching is case-sensitive by default."""
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "doc_label": ["CATEGORY_A"], # uppercase
+ }
+ )
+
+ result = initialized_vectorstore.reverse_search(query=query, max_n_results=10)
+
+ # Should not match "category_a" (lowercase)
+ assert len(result) == 0
+
+ def test_reverse_search_preserves_query_id(self, initialized_vectorstore):
+ """Query ID is preserved in output."""
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["custom_id_123"],
+ "doc_label": ["category_a"],
+ }
+ )
+
+ result = initialized_vectorstore.reverse_search(query=query, max_n_results=10)
+
+ assert all(qid == "custom_id_123" for qid in result.id)
+
+ def test_reverse_search_preserves_query_label_in_searched_doc_label(self, initialized_vectorstore):
+ """Query label is preserved in searched_doc_label column."""
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "doc_label": ["category_a"],
+ }
+ )
+
+ result = initialized_vectorstore.reverse_search(query=query, max_n_results=10)
+
+ assert all(label == "category_a" for label in result.searched_doc_label)
+
+ def test_reverse_search_doc_label_matches_stored_labels(self, initialized_vectorstore):
+ """doc_label in output matches stored document labels."""
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "doc_label": ["category_a"],
+ }
+ )
+
+ result = initialized_vectorstore.reverse_search(query=query, max_n_results=10)
+
+ # doc_label should match the stored labels
+ assert all(label == "category_a" for label in result.doc_label)
+
+ def test_reverse_search_doc_text_contains_document_content(self, initialized_vectorstore):
+ """doc_text contains the actual document text."""
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "doc_label": ["category_a"],
+ }
+ )
+
+ result = initialized_vectorstore.reverse_search(query=query, max_n_results=10)
+
+ # doc_text should contain the original text from the CSV
+ assert all(isinstance(text, str) for text in result.doc_text)
+ assert len(result.doc_text) > 0
+
+
+# ============================================================================
+# ERROR HANDLING TESTS
+# ============================================================================
+
+
+class TestVectorStoreReverseSearchErrorHandling:
+ """Tests for error handling during reverse search."""
+
+ def test_reverse_search_vectoriser_independent(self, initialized_vectorstore):
+ """Reverse search doesn't require vectoriser (no embeddings used)."""
+ # Disable vectoriser
+ initialized_vectorstore.vectoriser = None
+
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "doc_label": ["category_a"],
+ }
+ )
+
+ result = initialized_vectorstore.reverse_search(query=query, max_n_results=10)
+
+ # Should still work fine
+ assert len(result) == 2
+
+ def test_reverse_search_dataframe_join_failure_wrapped(self, initialized_vectorstore):
+ """DataFrame join failures wrapped in ClassifaiError."""
+ # Corrupt the internal vectors to break join
+ initialized_vectorstore.vectors = None
+
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "doc_label": ["category_a"],
+ }
+ )
+
+ with pytest.raises((ClassifaiError, AttributeError)):
+ initialized_vectorstore.reverse_search(query=query, max_n_results=10)
+
+ def test_reverse_search_error_includes_context(self, initialized_vectorstore):
+ """Error context includes relevant information."""
+ # Force an error by corrupting data
+ initialized_vectorstore.vectors = None
+
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "doc_label": ["category_a"],
+ }
+ )
+
+ with pytest.raises((ClassifaiError, AttributeError)):
+ initialized_vectorstore.reverse_search(query=query, max_n_results=10)
+
+
+# ============================================================================
+# HOOKS INTEGRATION TESTS
+# ============================================================================
+
+
+class TestVectorStoreReverseSearchHooksIntegration:
+ """Tests for hook integration in reverse_search()."""
+
+ def test_reverse_search_preprocess_hook_called_before_search(self, initialized_vectorstore):
+ """reverse_search_preprocess hook called before search."""
+ mock_hook = Mock(spec=HookBase)
+ mock_hook.return_value = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "doc_label": ["category_a"],
+ }
+ )
+
+ initialized_vectorstore.hooks = {"reverse_search_preprocess": mock_hook}
+
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "doc_label": ["category_a"],
+ }
+ )
+ result = initialized_vectorstore.reverse_search(query=query, max_n_results=10)
+
+ mock_hook.assert_called_once()
+
+ def test_reverse_search_postprocess_hook_called_after_search(self, initialized_vectorstore):
+ """reverse_search_postprocess hook called after search."""
+ mock_result = VectorStoreReverseSearchOutput.from_data(
+ {
+ "id": ["q1"],
+ "searched_doc_label": ["category_a"],
+ "doc_label": ["category_a"],
+ "doc_text": ["document 1"],
+ }
+ )
+
+ mock_hook = Mock(spec=HookBase)
+ mock_hook.return_value = mock_result
+
+ initialized_vectorstore.hooks = {"reverse_search_postprocess": mock_hook}
+
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "doc_label": ["category_a"],
+ }
+ )
+ result = initialized_vectorstore.reverse_search(query=query, max_n_results=10)
+
+ mock_hook.assert_called_once()
+
+ def test_reverse_search_preprocess_hook_failure_raises_hook_error(self, initialized_vectorstore):
+ """Preprocess hook failure raises HookError."""
+ bad_hook = Mock(spec=HookBase)
+ bad_hook.side_effect = Exception("Preprocess failed")
+
+ initialized_vectorstore.hooks = {"reverse_search_preprocess": bad_hook}
+
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "doc_label": ["category_a"],
+ }
+ )
+
+ with pytest.raises(HookError) as exc_info:
+ initialized_vectorstore.reverse_search(query=query, max_n_results=10)
+
+ assert "reverse_search_preprocess" in str(exc_info.value)
+
+ def test_reverse_search_postprocess_hook_failure_raises_hook_error(self, initialized_vectorstore):
+ """Postprocess hook failure raises HookError."""
+ bad_hook = Mock(spec=HookBase)
+ bad_hook.side_effect = Exception("Postprocess failed")
+
+ initialized_vectorstore.hooks = {"reverse_search_postprocess": bad_hook}
+
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "doc_label": ["category_a"],
+ }
+ )
+
+ with pytest.raises(HookError) as exc_info:
+ initialized_vectorstore.reverse_search(query=query, max_n_results=10)
+
+ assert "reverse_search_postprocess" in str(exc_info.value)
+
+ def test_reverse_search_multiple_preprocess_hooks_processed_in_order(self, initialized_vectorstore):
+ """Multiple preprocess hooks processed in order."""
+ hook1 = Mock(spec=HookBase)
+ hook1.return_value = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "doc_label": ["category_a"],
+ }
+ )
+
+ hook2 = Mock(spec=HookBase)
+ hook2.return_value = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "doc_label": ["category_a"],
+ }
+ )
+
+ initialized_vectorstore.hooks = {"reverse_search_preprocess": [hook1, hook2]}
+
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "doc_label": ["category_a"],
+ }
+ )
+ result = initialized_vectorstore.reverse_search(query=query, max_n_results=10)
+
+ assert hook1.call_count == 1
+ assert hook2.call_count == 1
+
+ def test_reverse_search_multiple_postprocess_hooks_processed_in_order(self, initialized_vectorstore):
+ """Multiple postprocess hooks processed in order."""
+ mock_output = VectorStoreReverseSearchOutput.from_data(
+ {
+ "id": ["q1"],
+ "searched_doc_label": ["category_a"],
+ "doc_label": ["category_a"],
+ "doc_text": ["document 1"],
+ }
+ )
+
+ hook1 = Mock(spec=HookBase)
+ hook1.return_value = mock_output
+
+ hook2 = Mock(spec=HookBase)
+ hook2.return_value = mock_output
+
+ initialized_vectorstore.hooks = {"reverse_search_postprocess": [hook1, hook2]}
+
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "doc_label": ["category_a"],
+ }
+ )
+ result = initialized_vectorstore.reverse_search(query=query, max_n_results=10)
+
+ assert hook1.call_count == 1
+ assert hook2.call_count == 1
+
+ def test_reverse_search_single_preprocess_hook_converted_to_list(self, initialized_vectorstore):
+ """Single preprocess hook automatically converted to list."""
+ mock_hook = Mock(spec=HookBase)
+ mock_hook.return_value = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "doc_label": ["category_a"],
+ }
+ )
+
+ initialized_vectorstore.hooks = {"reverse_search_preprocess": mock_hook}
+
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "doc_label": ["category_a"],
+ }
+ )
+ result = initialized_vectorstore.reverse_search(query=query, max_n_results=10)
+
+ mock_hook.assert_called_once()
+
+ def test_reverse_search_single_postprocess_hook_converted_to_list(self, initialized_vectorstore):
+ """Single postprocess hook automatically converted to list."""
+ mock_output = VectorStoreReverseSearchOutput.from_data(
+ {
+ "id": ["q1"],
+ "searched_doc_label": ["category_a"],
+ "doc_label": ["category_a"],
+ "doc_text": ["document 1"],
+ }
+ )
+
+ mock_hook = Mock(spec=HookBase)
+ mock_hook.return_value = mock_output
+
+ initialized_vectorstore.hooks = {"reverse_search_postprocess": mock_hook}
+
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "doc_label": ["category_a"],
+ }
+ )
+ result = initialized_vectorstore.reverse_search(query=query, max_n_results=10)
+
+ mock_hook.assert_called_once()
+
+
+# ============================================================================
+# EDGE CASE TESTS
+# ============================================================================
+
+
+class TestVectorStoreReverseSearchEdgeCases:
+ """Tests for edge cases in reverse_search()."""
+
+ def test_reverse_search_single_label_match(self, initialized_vectorstore):
+ """Single label match returns one result."""
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "doc_label": ["category_c"],
+ }
+ )
+
+ result = initialized_vectorstore.reverse_search(query=query, max_n_results=10)
+
+ assert len(result) == 1
+
+ def test_reverse_search_all_queries_return_empty(self, initialized_vectorstore):
+ """All queries returning empty results."""
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1", "q2", "q3"],
+ "doc_label": ["nonexistent_a", "nonexistent_b", "nonexistent_c"],
+ }
+ )
+
+ result = initialized_vectorstore.reverse_search(query=query, max_n_results=10)
+
+ assert len(result) == 0
+
+ def test_reverse_search_mixed_empty_and_nonempty_results(self, initialized_vectorstore):
+ """Mix of queries with and without results."""
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1", "q2"],
+ "doc_label": ["category_a", "nonexistent"],
+ }
+ )
+
+ result = initialized_vectorstore.reverse_search(query=query, max_n_results=10)
+
+ # Only q1 should have results
+ assert len(result) == 2
+
+ def test_reverse_search_special_characters_in_label(self, initialized_vectorstore):
+ """Labels with special characters."""
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "doc_label": ["category_a"],
+ }
+ )
+
+ result = initialized_vectorstore.reverse_search(query=query, max_n_results=10)
+
+ assert isinstance(result, VectorStoreReverseSearchOutput)
+
+ def test_reverse_search_very_long_label(self, initialized_vectorstore):
+ """Very long label string."""
+ long_label = "a" * 1000
+
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "doc_label": [long_label],
+ }
+ )
+
+ result = initialized_vectorstore.reverse_search(query=query, max_n_results=10)
+
+ # Should find no matches (label not in store)
+ assert len(result) == 0
+
+ def test_reverse_search_max_n_results_exceeds_available(self, initialized_vectorstore):
+ """max_n_results larger than available matches."""
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "doc_label": ["category_a"],
+ }
+ )
+
+ result = initialized_vectorstore.reverse_search(query=query, max_n_results=1000)
+
+ # Should return only 2 matches (all available)
+ assert len(result) == 2
+
+ def test_reverse_search_whitespace_in_label(self, initialized_vectorstore):
+ """Label with leading/trailing whitespace."""
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "doc_label": [" category_a "], # with whitespace
+ }
+ )
+
+ result = initialized_vectorstore.reverse_search(query=query, max_n_results=10)
+
+ # Should not match (exact matching)
+ assert len(result) == 0
+
+ def test_reverse_search_unicode_label(self, initialized_vectorstore):
+ """Unicode characters in label."""
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "doc_label": ["类别_a"], # Chinese characters
+ }
+ )
+
+ result = initialized_vectorstore.reverse_search(query=query, max_n_results=10)
+
+ # Should find no matches
+ assert len(result) == 0
+
+ def test_reverse_search_empty_string_label(self, initialized_vectorstore):
+ """Empty string as label."""
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "doc_label": [""],
+ }
+ )
+
+ result = initialized_vectorstore.reverse_search(query=query, max_n_results=10)
+
+ # Should find no matches
+ assert len(result) == 0
+
+ def test_reverse_search_large_number_of_queries(self, initialized_vectorstore):
+ """Large number of queries processed correctly."""
+ n_queries = 50
+ ids = [str(i) for i in range(n_queries)]
+ labels = ["category_a" if i % 2 == 0 else "category_b" for i in range(n_queries)]
+
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ids,
+ "doc_label": labels,
+ }
+ )
+
+ result = initialized_vectorstore.reverse_search(query=query, max_n_results=10)
+
+ # Should process all queries (25 * 2 + 25 * 2 = 100 results)
+ assert len(result) == 100
+
+ def test_reverse_search_returns_correct_type(self, initialized_vectorstore):
+ """reverse_search() returns VectorStoreReverseSearchOutput."""
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "doc_label": ["category_a"],
+ }
+ )
+
+ result = initialized_vectorstore.reverse_search(query=query, max_n_results=10)
+
+ assert isinstance(result, VectorStoreReverseSearchOutput)
+
+ def test_reverse_search_partial_match_with_prefix(self, initialized_vectorstore):
+ """Partial matching with various prefix patterns."""
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "doc_label": ["cat"],
+ }
+ )
+
+ result = initialized_vectorstore.reverse_search(
+ query=query,
+ max_n_results=10,
+ partial_match=True,
+ )
+
+ # Should find documents starting with "cat"
+ assert len(result) >= 5
+
+ def test_reverse_search_partial_match_single_character(self, initialized_vectorstore):
+ """Partial matching with single character prefix."""
+ query = VectorStoreReverseSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "doc_label": ["c"],
+ }
+ )
+
+ result = initialized_vectorstore.reverse_search(
+ query=query,
+ max_n_results=10,
+ partial_match=True,
+ )
+
+ # Should find all documents starting with "c"
+ assert len(result) >= 5
diff --git a/tests/test_indexers/test_vectorstore_search.py b/tests/test_indexers/test_vectorstore_search.py
new file mode 100644
index 0000000..5e58ab0
--- /dev/null
+++ b/tests/test_indexers/test_vectorstore_search.py
@@ -0,0 +1,512 @@
+"""Unit tests for VectorStore.search() method."""
+
+from __future__ import annotations
+
+from unittest.mock import Mock
+
+import numpy as np
+import pytest
+
+from classifai.exceptions import (
+ ConfigurationError,
+ DataValidationError,
+ HookError,
+ VectorisationError,
+)
+from classifai.indexers import VectorStore
+from classifai.indexers.dataclasses import VectorStoreSearchInput, VectorStoreSearchOutput
+from classifai.indexers.hooks import HookBase
+from classifai.vectorisers import VectoriserBase
+
+# ============================================================================
+# FIXTURES
+# ============================================================================
+
+
+@pytest.fixture
+def mock_vectoriser():
+ """Mock VectoriserBase that returns predictable embeddings."""
+ mock = Mock(spec=VectoriserBase)
+
+ def transform_side_effect(texts):
+ """Return one embedding per text, all with shape (3,)."""
+ num_texts = len(texts)
+ return np.array([np.linspace(0.1, 0.3, 3) + (i * 0.3) for i in range(num_texts)])
+
+ mock.transform.side_effect = transform_side_effect
+ mock.__class__.__name__ = "MockVectoriser"
+ return mock
+
+
+@pytest.fixture
+def initialized_vectorstore(mock_vectoriser, tmp_path):
+ """Create a fully initialized VectorStore with test data."""
+ # Create a test CSV
+ csv_path = tmp_path / "test.csv"
+ csv_path.write_text("label,text\ndoc1,hello world\ndoc2,goodbye world\ndoc3,test document\n")
+
+ # Create VectorStore
+ vs = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ output_dir=str(tmp_path / "output"),
+ skip_save=True,
+ )
+
+ return vs
+
+
+# ============================================================================
+# INPUT VALIDATION TESTS
+# ============================================================================
+
+
+class TestVectorStoreSearchInputValidation:
+ """Tests for input validation in search() method."""
+
+ def test_search_query_must_be_vectorstore_search_input(self, initialized_vectorstore):
+ """Query must be VectorStoreSearchInput object."""
+ with pytest.raises(DataValidationError) as exc_info:
+ initialized_vectorstore.search(query="not a VectorStoreSearchInput")
+
+ assert "VectorStoreSearchInput" in str(exc_info.value)
+
+ def test_search_query_none_raises_error(self, initialized_vectorstore):
+ """query=None raises DataValidationError."""
+ with pytest.raises(DataValidationError) as exc_info:
+ initialized_vectorstore.search(query=None)
+
+ assert "VectorStoreSearchInput" in str(exc_info.value)
+
+ def test_search_n_results_must_be_positive_int(self, initialized_vectorstore):
+ """n_results must be int >= 1."""
+ query = VectorStoreSearchInput.from_data({"id": ["1"], "query": ["test"]})
+
+ with pytest.raises(DataValidationError) as exc_info:
+ initialized_vectorstore.search(query=query, n_results=0)
+
+ assert "n_results" in str(exc_info.value).lower()
+
+ def test_search_n_results_negative_raises_error(self, initialized_vectorstore):
+ """n_results < 1 raises DataValidationError."""
+ query = VectorStoreSearchInput.from_data({"id": ["1"], "query": ["test"]})
+
+ with pytest.raises(DataValidationError) as exc_info:
+ initialized_vectorstore.search(query=query, n_results=-1)
+
+ assert "n_results" in str(exc_info.value).lower()
+
+ def test_search_n_results_non_int_raises_error(self, initialized_vectorstore):
+ """n_results as float/string raises DataValidationError."""
+ query = VectorStoreSearchInput.from_data({"id": ["1"], "query": ["test"]})
+
+ with pytest.raises(DataValidationError) as exc_info:
+ initialized_vectorstore.search(query=query, n_results="10")
+
+ assert "n_results" in str(exc_info.value).lower()
+
+ def test_search_batch_size_must_be_positive_int(self, initialized_vectorstore):
+ """batch_size must be int >= 1 or None."""
+ query = VectorStoreSearchInput.from_data({"id": ["1"], "query": ["test"]})
+
+ with pytest.raises(DataValidationError) as exc_info:
+ initialized_vectorstore.search(query=query, batch_size=0)
+
+ assert "batch_size" in str(exc_info.value).lower()
+
+ def test_search_batch_size_negative_raises_error(self, initialized_vectorstore):
+ """batch_size < 1 raises DataValidationError."""
+ query = VectorStoreSearchInput.from_data({"id": ["1"], "query": ["test"]})
+
+ with pytest.raises(DataValidationError) as exc_info:
+ initialized_vectorstore.search(query=query, batch_size=-5)
+
+ assert "batch_size" in str(exc_info.value).lower()
+
+ def test_search_batch_size_non_int_raises_error(self, initialized_vectorstore):
+ """batch_size as string raises DataValidationError."""
+ query = VectorStoreSearchInput.from_data({"id": ["1"], "query": ["test"]})
+
+ with pytest.raises(DataValidationError) as exc_info:
+ initialized_vectorstore.search(query=query, batch_size="10")
+
+ assert "batch_size" in str(exc_info.value).lower()
+
+ def test_search_batch_size_none_is_valid(self, initialized_vectorstore):
+ """batch_size=None is valid and uses default."""
+ query = VectorStoreSearchInput.from_data({"id": ["1"], "query": ["test"]})
+
+ # Should not raise
+ result = initialized_vectorstore.search(query=query, n_results=3, batch_size=None)
+ assert isinstance(result, VectorStoreSearchOutput)
+
+ def test_search_empty_query_raises_error(self, initialized_vectorstore):
+ """Empty query DataFrame raises DataValidationError."""
+ query = VectorStoreSearchInput.from_data({"id": [], "query": []})
+
+ with pytest.raises(DataValidationError) as exc_info:
+ initialized_vectorstore.search(query=query)
+
+ assert "empty" in str(exc_info.value).lower()
+
+ def test_search_uninitialized_vectorstore_raises_error(self, mock_vectoriser):
+ """Vector store not initialized raises ConfigurationError."""
+ vs = Mock(spec=VectorStore)
+ vs.vectors = None
+ vs.batch_size = 128
+ vs.meta_data = {}
+
+ # Create a real instance but with vectors=None
+ vs_real = VectorStore.__new__(VectorStore)
+ vs_real.vectors = None
+ vs_real.batch_size = 128
+ vs_real.meta_data = {}
+ vs_real.vectoriser = mock_vectoriser
+ vs_real.vectoriser_class = "MockVectoriser"
+ vs_real.hooks = {}
+ vs_real.quiet_mode = False
+ vs_real.classifai_tqdm = lambda iterable, *args, **kwargs: iterable
+
+ query = VectorStoreSearchInput.from_data({"id": ["1"], "query": ["test"]})
+
+ with pytest.raises(ConfigurationError) as exc_info:
+ vs_real.search(query=query)
+
+ assert "not initialised" in str(exc_info.value).lower()
+
+
+# ============================================================================
+# SEARCH OPERATION TESTS
+# ============================================================================
+
+
+class TestVectorStoreSearchOperation:
+ """Tests for the core search operation."""
+
+ def test_search_single_query_returns_results(self, initialized_vectorstore):
+ """Single query processes correctly and returns results."""
+ query = VectorStoreSearchInput.from_data({"id": ["q1"], "query": ["hello"]})
+
+ result = initialized_vectorstore.search(query=query, n_results=2)
+
+ assert isinstance(result, VectorStoreSearchOutput)
+ assert len(result) == 2 # n_results=2
+ assert result.query_id[0] == "q1"
+
+ def test_search_multiple_queries_processes_correctly(self, initialized_vectorstore):
+ """Multiple queries in batch process correctly."""
+ query = VectorStoreSearchInput.from_data(
+ {
+ "id": ["q1", "q2"],
+ "query": ["hello", "test"],
+ }
+ )
+
+ result = initialized_vectorstore.search(query=query, n_results=2)
+
+ assert isinstance(result, VectorStoreSearchOutput)
+ assert len(result) == 4 # 2 queries * 2 results each
+ assert list(result.query_id[:2]) == ["q1", "q1"]
+ assert list(result.query_id[2:]) == ["q2", "q2"]
+
+ def test_search_similarity_scores_computed(self, initialized_vectorstore):
+ """Similarity scores are computed via dot-product."""
+ query = VectorStoreSearchInput.from_data({"id": ["q1"], "query": ["hello"]})
+
+ result = initialized_vectorstore.search(query=query, n_results=3)
+
+ # Check that scores are floats and have values
+ assert all(isinstance(score, (float, np.floating)) for score in result.score)
+ assert len(result.score) == 3
+
+ def test_search_top_n_results_returned(self, initialized_vectorstore):
+ """Top n_results returned per query."""
+ query = VectorStoreSearchInput.from_data(
+ {
+ "id": ["q1"],
+ "query": ["hello"],
+ }
+ )
+
+ result = initialized_vectorstore.search(query=query, n_results=2)
+
+ # Should have exactly 2 results (n_results=2)
+ assert len(result) == 2
+
+ def test_search_results_ranked_by_score_descending(self, initialized_vectorstore):
+ """Results ranked by score in descending order."""
+ query = VectorStoreSearchInput.from_data({"id": ["q1"], "query": ["hello"]})
+
+ result = initialized_vectorstore.search(query=query, n_results=3)
+
+ scores = list(result.score)
+ # Verify scores are in descending order
+ assert scores == sorted(scores, reverse=True)
+
+ def test_search_output_shape_matches_expected(self, initialized_vectorstore):
+ """Output shape is (n_queries * n_results) rows."""
+ query = VectorStoreSearchInput.from_data(
+ {
+ "id": ["q1", "q2", "q3"],
+ "query": ["hello", "test", "world"],
+ }
+ )
+
+ n_results = 2
+ result = initialized_vectorstore.search(query=query, n_results=n_results)
+
+ # Expected: 3 queries * 2 results = 6 rows
+ assert len(result) == 3 * n_results
+
+ def test_search_metadata_columns_included_in_output(self, mock_vectoriser, tmp_path):
+ """Metadata columns included in output when specified."""
+ csv_path = tmp_path / "test_meta.csv"
+ csv_path.write_text("label,text,source\ndoc1,hello,source_a\ndoc2,world,source_b\n")
+
+ vs = VectorStore(
+ file_name=str(csv_path),
+ data_type="csv",
+ vectoriser=mock_vectoriser,
+ meta_data={"source": str},
+ output_dir=str(tmp_path / "output"),
+ skip_save=True,
+ )
+
+ query = VectorStoreSearchInput.from_data({"id": ["q1"], "query": ["hello"]})
+ result = vs.search(query=query, n_results=2)
+
+ # Verify metadata column is in output
+ assert "source" in result.columns
+
+ def test_search_query_batching_with_custom_batch_size(self, initialized_vectorstore, mock_vectoriser):
+ """Query batching with custom batch_size works."""
+ query = VectorStoreSearchInput.from_data(
+ {
+ "id": ["q1", "q2", "q3"],
+ "query": ["hello", "test", "world"],
+ }
+ )
+
+ # Use batch_size=1 to force multiple batches
+ result = initialized_vectorstore.search(query=query, n_results=2, batch_size=1)
+
+ assert isinstance(result, VectorStoreSearchOutput)
+ assert len(result) == 6 # 3 queries * 2 results
+
+ def test_search_returns_correct_columns(self, initialized_vectorstore):
+ """Output contains all required columns."""
+ query = VectorStoreSearchInput.from_data({"id": ["q1"], "query": ["hello"]})
+
+ result = initialized_vectorstore.search(query=query, n_results=2)
+
+ required_columns = ["query_id", "query_text", "doc_label", "doc_text", "rank", "score"]
+ for col in required_columns:
+ assert col in result.columns
+
+
+# # ============================================================================
+# # ERROR HANDLING TESTS
+# # ============================================================================
+
+
+class TestVectorStoreSearchErrorHandling:
+ """Tests for error handling during search."""
+
+ def test_search_query_embedding_failure_raises_vectorisation_error(self, initialized_vectorstore):
+ """Query embedding failure raises VectorisationError."""
+ # Make vectoriser fail on query embedding
+ initialized_vectorstore.vectoriser.transform.side_effect = Exception("Vectoriser failed")
+
+ query = VectorStoreSearchInput.from_data({"id": ["q1"], "query": ["hello"]})
+
+ with pytest.raises(VectorisationError) as exc_info:
+ initialized_vectorstore.search(query=query)
+
+ assert "Failed to embed query batch" in str(exc_info.value)
+
+ def test_search_vectoriser_exceptions_wrapped(self, initialized_vectorstore):
+ """Vectoriser.transform() exceptions include context."""
+ initialized_vectorstore.vectoriser.transform.side_effect = RuntimeError("Transform failed")
+
+ query = VectorStoreSearchInput.from_data({"id": ["q1"], "query": ["hello"]})
+
+ with pytest.raises(VectorisationError) as exc_info:
+ initialized_vectorstore.search(query=query)
+
+ error_str = str(exc_info.value)
+ assert "MockVectoriser" in error_str # vectoriser class in context
+ assert "batch_size" in error_str or "batch" in error_str.lower()
+
+ def test_search_vectoriser_error_includes_batch_info(self, initialized_vectorstore):
+ """Error context includes vectoriser class and batch info."""
+ initialized_vectorstore.vectoriser.transform.side_effect = ValueError("Bad embedding")
+
+ query = VectorStoreSearchInput.from_data({"id": ["q1"], "query": ["test"]})
+
+ with pytest.raises(VectorisationError) as exc_info:
+ initialized_vectorstore.search(query=query)
+
+ context_str = str(exc_info.value)
+ assert "vectoriser" in context_str.lower() or "MockVectoriser" in context_str
+
+
+# # ============================================================================
+# # HOOKS INTEGRATION TESTS
+# # ============================================================================
+
+
+class TestVectorStoreSearchHooksIntegration:
+ """Tests for hook integration in search."""
+
+ def test_search_preprocess_hook_called_before_search(self, initialized_vectorstore):
+ """search_preprocess hook called before search."""
+ mock_hook = Mock(spec=HookBase)
+ mock_hook.return_value = VectorStoreSearchInput.from_data({"id": ["q1"], "query": ["hello"]})
+
+ initialized_vectorstore.hooks = {"search_preprocess": mock_hook}
+
+ query = VectorStoreSearchInput.from_data({"id": ["q1"], "query": ["hello"]})
+ result = initialized_vectorstore.search(query=query, n_results=3)
+
+ mock_hook.assert_called_once()
+
+ def test_search_postprocess_hook_called_after_search(self, initialized_vectorstore):
+ """search_postprocess hook called after search."""
+ # Create a real VectorStoreSearchOutput to return from the hook
+ mock_result = VectorStoreSearchOutput.from_data(
+ {
+ "query_id": ["q1"],
+ "query_text": ["hello"],
+ "doc_label": ["doc1"],
+ "doc_text": ["hello world"],
+ "rank": [1],
+ "score": [0.95],
+ }
+ )
+
+ mock_hook = Mock(spec=HookBase)
+ mock_hook.return_value = mock_result
+
+ initialized_vectorstore.hooks = {"search_postprocess": mock_hook}
+
+ query = VectorStoreSearchInput.from_data({"id": ["q1"], "query": ["hello"]})
+ result = initialized_vectorstore.search(query=query, n_results=3)
+
+ mock_hook.assert_called_once()
+ # Optionally we could verify the result is what we expect
+ assert isinstance(result, VectorStoreSearchOutput)
+
+ def test_search_hook_failure_raises_hook_error(self, initialized_vectorstore):
+ """Hook failure raises HookError."""
+ bad_hook = Mock(spec=HookBase)
+ bad_hook.side_effect = Exception("Hook failed")
+
+ initialized_vectorstore.hooks = {"search_preprocess": bad_hook}
+
+ query = VectorStoreSearchInput.from_data({"id": ["q1"], "query": ["hello"]})
+
+ with pytest.raises(HookError) as exc_info:
+ initialized_vectorstore.search(query=query, n_results=3)
+
+ assert "search_preprocess" in str(exc_info.value)
+
+ def test_search_multiple_hooks_processed_in_order(self, initialized_vectorstore):
+ """Multiple hooks in list processed in order."""
+ hook1 = Mock(spec=HookBase)
+ hook1.return_value = VectorStoreSearchInput.from_data({"id": ["q1"], "query": ["hello"]})
+
+ hook2 = Mock(spec=HookBase)
+ hook2.return_value = VectorStoreSearchInput.from_data({"id": ["q1"], "query": ["hello"]})
+
+ initialized_vectorstore.hooks = {"search_preprocess": [hook1, hook2]}
+
+ query = VectorStoreSearchInput.from_data({"id": ["q1"], "query": ["hello"]})
+ result = initialized_vectorstore.search(query=query, n_results=3)
+
+ assert hook1.call_count == 1
+ assert hook2.call_count == 1
+
+ def test_search_single_hook_converted_to_list(self, initialized_vectorstore):
+ """Single hook automatically converted to list."""
+ mock_hook = Mock(spec=HookBase)
+ mock_hook.return_value = VectorStoreSearchInput.from_data({"id": ["q1"], "query": ["hello"]})
+
+ initialized_vectorstore.hooks = {"search_preprocess": mock_hook}
+
+ query = VectorStoreSearchInput.from_data({"id": ["q1"], "query": ["hello"]})
+ result = initialized_vectorstore.search(query=query, n_results=3)
+
+ # Verify hook was called (auto-converted to list)
+ mock_hook.assert_called_once()
+
+ def test_search_postprocess_hook_failure_raises_hook_error(self, initialized_vectorstore):
+ """Postprocess hook failure raises HookError."""
+ bad_hook = Mock(spec=HookBase)
+ bad_hook.side_effect = Exception("Postprocess failed")
+
+ initialized_vectorstore.hooks = {"search_postprocess": bad_hook}
+
+ query = VectorStoreSearchInput.from_data({"id": ["q1"], "query": ["hello"]})
+
+ with pytest.raises(HookError) as exc_info:
+ initialized_vectorstore.search(query=query, n_results=3)
+
+ assert "search_postprocess" in str(exc_info.value)
+
+
+# # ============================================================================
+# # EDGE CASE TESTS
+# # ============================================================================
+
+# TODO: All of these tests rely on querying to retrieve large values of results, however currently there is a bug
+# that if n_results is greater than the number of documents in the store then an out of bounds error is thrown. This needs to be fixed in the VectorStore.search() method before these tests can be run.
+
+# class TestVectorStoreSearchEdgeCases:
+# """Tests for edge cases in search."""
+
+# def test_search_n_results_greater_than_available_documents(self, initialized_vectorstore):
+# """n_results > available documents returns all documents."""
+# # We have 3 documents in the store
+# query = VectorStoreSearchInput.from_data({"id": ["q1"], "query": ["hello"]})
+
+# result = initialized_vectorstore.search(query=query, n_results=100)
+
+# # Should return only 3 documents (all available)
+# assert len(result) == 3
+
+# def test_search_single_query_single_document_store(self, mock_vectoriser, tmp_path):
+# """Search with single document in store."""
+# csv_path = tmp_path / "single.csv"
+# csv_path.write_text("label,text\ndoc1,only document\n")
+
+# vs = VectorStore(
+# file_name=str(csv_path),
+# data_type="csv",
+# vectoriser=mock_vectoriser,
+# output_dir=str(tmp_path / "output"),
+# skip_save=True,
+# )
+
+# query = VectorStoreSearchInput.from_data({"id": ["q1"], "query": ["test"]})
+# result = vs.search(query=query, n_results=5)
+
+# assert len(result) == 1
+
+# def test_search_identical_embeddings_still_returns_results(self, initialized_vectorstore):
+# """Search works even with identical embeddings."""
+# query = VectorStoreSearchInput.from_data({"id": ["q1"], "query": ["hello"]})
+
+# result = initialized_vectorstore.search(query=query, n_results=2)
+
+# assert len(result) == 2
+# assert isinstance(result, VectorStoreSearchOutput)
+
+# def test_search_large_n_results_value(self, initialized_vectorstore):
+# """Search with very large n_results parameter."""
+# query = VectorStoreSearchInput.from_data({"id": ["q1"], "query": ["hello"]})
+
+# result = initialized_vectorstore.search(query=query, n_results=1000)
+
+# # Should return only 3 documents (all available)
+# assert len(result) == 3