Skip to content

Commit bee6d07

Browse files
committed
Add unit tests for key gaps: SparseValuesFactory, error classes, and utility functions
- Add comprehensive tests for SparseValuesFactory (REST version) covering all input types, error cases, and edge cases - Add tests for error classes in db_data/errors.py to ensure clear error messages - Add tests for utility functions: check_kwargs, validate_and_convert_errors, and fix_tuple_length - All tests pass and follow project conventions for readability
1 parent 745c7c7 commit bee6d07

5 files changed

Lines changed: 534 additions & 0 deletions

File tree

‎tests/unit/data/test_errors.py‎

Lines changed: 136 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,136 @@
1+
from pinecone.db_data.errors import (
2+
VectorDictionaryMissingKeysError,
3+
VectorDictionaryExcessKeysError,
4+
VectorTupleLengthError,
5+
SparseValuesTypeError,
6+
SparseValuesMissingKeysError,
7+
SparseValuesDictionaryExpectedError,
8+
MetadataDictionaryExpectedError,
9+
)
10+
11+
12+
class TestVectorDictionaryMissingKeysError:
13+
"""Test VectorDictionaryMissingKeysError exception."""
14+
15+
def test_error_message_includes_missing_keys(self):
16+
"""Test that error message lists missing required fields."""
17+
item = {"values": [0.1, 0.2]}
18+
error = VectorDictionaryMissingKeysError(item)
19+
assert "missing required fields" in str(error).lower()
20+
assert "id" in str(error)
21+
22+
def test_error_message_with_multiple_missing_keys(self):
23+
"""Test error message when multiple keys are missing."""
24+
item = {}
25+
error = VectorDictionaryMissingKeysError(item)
26+
assert "missing required fields" in str(error).lower()
27+
28+
29+
class TestVectorDictionaryExcessKeysError:
30+
"""Test VectorDictionaryExcessKeysError exception."""
31+
32+
def test_error_message_includes_excess_keys(self):
33+
"""Test that error message lists excess keys."""
34+
item = {"id": "1", "values": [0.1, 0.2], "extra_field": "value", "another_extra": 123}
35+
error = VectorDictionaryExcessKeysError(item)
36+
assert "excess keys" in str(error).lower()
37+
assert "extra_field" in str(error) or "another_extra" in str(error)
38+
39+
def test_error_message_includes_allowed_keys(self):
40+
"""Test that error message includes list of allowed keys."""
41+
item = {"id": "1", "values": [0.1, 0.2], "invalid": "key"}
42+
error = VectorDictionaryExcessKeysError(item)
43+
assert "allowed keys" in str(error).lower()
44+
45+
46+
class TestVectorTupleLengthError:
47+
"""Test VectorTupleLengthError exception."""
48+
49+
def test_error_message_includes_tuple_length(self):
50+
"""Test that error message includes the tuple length."""
51+
item = ("id", "values", "metadata", "extra")
52+
error = VectorTupleLengthError(item)
53+
assert str(len(item)) in str(error)
54+
assert "tuple" in str(error).lower()
55+
56+
def test_error_message_with_length_one(self):
57+
"""Test error message for tuple of length 1."""
58+
item = ("id",)
59+
error = VectorTupleLengthError(item)
60+
assert "1" in str(error)
61+
62+
def test_error_message_with_length_four(self):
63+
"""Test error message for tuple of length 4."""
64+
item = ("id", "values", "metadata", "extra")
65+
error = VectorTupleLengthError(item)
66+
assert "4" in str(error)
67+
68+
69+
class TestSparseValuesTypeError:
70+
"""Test SparseValuesTypeError exception."""
71+
72+
def test_error_message_mentions_sparse_values(self):
73+
"""Test that error message mentions sparse_values."""
74+
error = SparseValuesTypeError()
75+
assert "sparse_values" in str(error).lower()
76+
77+
def test_error_is_both_value_and_type_error(self):
78+
"""Test that SparseValuesTypeError is both ValueError and TypeError."""
79+
error = SparseValuesTypeError()
80+
assert isinstance(error, ValueError)
81+
assert isinstance(error, TypeError)
82+
83+
84+
class TestSparseValuesMissingKeysError:
85+
"""Test SparseValuesMissingKeysError exception."""
86+
87+
def test_error_message_includes_found_keys(self):
88+
"""Test that error message includes the keys that were found."""
89+
sparse_values_dict = {"indices": [0, 2]}
90+
error = SparseValuesMissingKeysError(sparse_values_dict)
91+
assert "missing required keys" in str(error).lower()
92+
assert "indices" in str(error) or "values" in str(error)
93+
94+
def test_error_message_with_empty_dict(self):
95+
"""Test error message when dictionary is empty."""
96+
sparse_values_dict = {}
97+
error = SparseValuesMissingKeysError(sparse_values_dict)
98+
assert "missing required keys" in str(error).lower()
99+
100+
101+
class TestSparseValuesDictionaryExpectedError:
102+
"""Test SparseValuesDictionaryExpectedError exception."""
103+
104+
def test_error_message_includes_actual_type(self):
105+
"""Test that error message includes the actual type found."""
106+
sparse_values_dict = "not a dict"
107+
error = SparseValuesDictionaryExpectedError(sparse_values_dict)
108+
assert "dictionary" in str(error).lower()
109+
assert "str" in str(error) or type(sparse_values_dict).__name__ in str(error)
110+
111+
def test_error_message_with_integer(self):
112+
"""Test error message when integer is provided."""
113+
sparse_values_dict = 123
114+
error = SparseValuesDictionaryExpectedError(sparse_values_dict)
115+
assert "dictionary" in str(error).lower()
116+
assert isinstance(error, ValueError)
117+
assert isinstance(error, TypeError)
118+
119+
120+
class TestMetadataDictionaryExpectedError:
121+
"""Test MetadataDictionaryExpectedError exception."""
122+
123+
def test_error_message_includes_actual_type(self):
124+
"""Test that error message includes the actual type found."""
125+
item = {"metadata": "not a dict"}
126+
error = MetadataDictionaryExpectedError(item)
127+
assert "dictionary" in str(error).lower()
128+
assert "metadata" in str(error).lower()
129+
130+
def test_error_message_with_list(self):
131+
"""Test error message when list is provided as metadata."""
132+
item = {"metadata": [1, 2, 3]}
133+
error = MetadataDictionaryExpectedError(item)
134+
assert "dictionary" in str(error).lower()
135+
assert isinstance(error, ValueError)
136+
assert isinstance(error, TypeError)
Lines changed: 152 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,152 @@
1+
import numpy as np
2+
import pandas as pd
3+
import pytest
4+
5+
from pinecone.db_data.sparse_values_factory import SparseValuesFactory
6+
from pinecone import SparseValues
7+
from pinecone.core.openapi.db_data.models import SparseValues as OpenApiSparseValues
8+
from pinecone.db_data.errors import (
9+
SparseValuesTypeError,
10+
SparseValuesMissingKeysError,
11+
SparseValuesDictionaryExpectedError,
12+
)
13+
14+
15+
class TestSparseValuesFactory:
16+
"""Test SparseValuesFactory for REST API (db_data module)."""
17+
18+
def test_build_when_none_returns_none(self):
19+
"""Test that None input returns None."""
20+
assert SparseValuesFactory.build(None) is None
21+
22+
def test_build_when_passed_openapi_sparse_values(self):
23+
"""Test that OpenApiSparseValues are returned unchanged."""
24+
sv = OpenApiSparseValues(indices=[0, 2], values=[0.1, 0.3])
25+
actual = SparseValuesFactory.build(sv)
26+
assert actual == sv
27+
assert actual is sv
28+
29+
def test_build_when_given_sparse_values_dataclass(self):
30+
"""Test conversion from SparseValues dataclass to OpenApiSparseValues."""
31+
sv = SparseValues(indices=[0, 2], values=[0.1, 0.3])
32+
actual = SparseValuesFactory.build(sv)
33+
expected = OpenApiSparseValues(indices=[0, 2], values=[0.1, 0.3])
34+
assert isinstance(actual, OpenApiSparseValues)
35+
assert actual.indices == expected.indices
36+
assert actual.values == expected.values
37+
38+
@pytest.mark.parametrize(
39+
"input_dict",
40+
[
41+
{"indices": [2], "values": [0.3]},
42+
{"indices": [88, 102], "values": [-0.1, 0.3]},
43+
{"indices": [0, 2, 4], "values": [0.1, 0.3, 0.5]},
44+
{"indices": [0, 2, 4, 6], "values": [0.1, 0.3, 0.5, 0.7]},
45+
],
46+
)
47+
def test_build_when_valid_dictionary(self, input_dict):
48+
"""Test building from valid dictionary input."""
49+
actual = SparseValuesFactory.build(input_dict)
50+
expected = OpenApiSparseValues(indices=input_dict["indices"], values=input_dict["values"])
51+
assert actual.indices == expected.indices
52+
assert actual.values == expected.values
53+
54+
@pytest.mark.parametrize(
55+
"input_dict",
56+
[
57+
{"indices": np.array([0, 2]), "values": [0.1, 0.3]},
58+
{"indices": [0, 2], "values": np.array([0.1, 0.3])},
59+
{"indices": np.array([0, 2]), "values": np.array([0.1, 0.3])},
60+
{"indices": pd.array([0, 2]), "values": [0.1, 0.3]},
61+
{"indices": [0, 2], "values": pd.array([0.1, 0.3])},
62+
{"indices": pd.array([0, 2]), "values": pd.array([0.1, 0.3])},
63+
],
64+
)
65+
def test_build_when_special_data_types(self, input_dict):
66+
"""Test that the factory handles numpy/pandas arrays correctly."""
67+
actual = SparseValuesFactory.build(input_dict)
68+
expected = OpenApiSparseValues(indices=[0, 2], values=[0.1, 0.3])
69+
assert actual.indices == expected.indices
70+
assert actual.values == expected.values
71+
72+
@pytest.mark.parametrize(
73+
"input_dict",
74+
[{"indices": [2], "values": [0.3, 0.3]}, {"indices": [88, 102], "values": [-0.1]}],
75+
)
76+
def test_build_when_list_sizes_dont_match(self, input_dict):
77+
"""Test that mismatched indices and values lengths raise ValueError."""
78+
with pytest.raises(
79+
ValueError, match="Sparse values indices and values must have the same length"
80+
):
81+
SparseValuesFactory.build(input_dict)
82+
83+
@pytest.mark.parametrize(
84+
"input_dict",
85+
[
86+
{"indices": [2.0], "values": [0.3]},
87+
{"indices": ["2"], "values": [0.3]},
88+
{"indices": np.array([2.0]), "values": [0.3]},
89+
{"indices": pd.array([2.0]), "values": [0.3]},
90+
],
91+
)
92+
def test_build_when_non_integer_indices(self, input_dict):
93+
"""Test that non-integer indices raise SparseValuesTypeError."""
94+
with pytest.raises(SparseValuesTypeError):
95+
SparseValuesFactory.build(input_dict)
96+
97+
@pytest.mark.parametrize(
98+
"input_dict", [{"indices": [2], "values": ["3.2"]}, {"indices": [2], "values": [True]}]
99+
)
100+
def test_build_when_non_float_values(self, input_dict):
101+
"""Test that non-float values raise SparseValuesTypeError."""
102+
with pytest.raises(SparseValuesTypeError):
103+
SparseValuesFactory.build(input_dict)
104+
105+
def test_build_when_missing_indices_key(self):
106+
"""Test that missing 'indices' key raises SparseValuesMissingKeysError."""
107+
input_dict = {"values": [0.1, 0.3]}
108+
with pytest.raises(SparseValuesMissingKeysError) as exc_info:
109+
SparseValuesFactory.build(input_dict)
110+
assert "indices" in str(exc_info.value)
111+
112+
def test_build_when_missing_values_key(self):
113+
"""Test that missing 'values' key raises SparseValuesMissingKeysError."""
114+
input_dict = {"indices": [0, 2]}
115+
with pytest.raises(SparseValuesMissingKeysError) as exc_info:
116+
SparseValuesFactory.build(input_dict)
117+
assert "values" in str(exc_info.value)
118+
119+
def test_build_when_missing_both_keys(self):
120+
"""Test that missing both keys raises SparseValuesMissingKeysError."""
121+
input_dict = {}
122+
with pytest.raises(SparseValuesMissingKeysError) as exc_info:
123+
SparseValuesFactory.build(input_dict)
124+
assert "indices" in str(exc_info.value) or "values" in str(exc_info.value)
125+
126+
def test_build_when_not_a_dictionary(self):
127+
"""Test that non-dictionary input raises SparseValuesDictionaryExpectedError."""
128+
with pytest.raises(SparseValuesDictionaryExpectedError) as exc_info:
129+
SparseValuesFactory.build("not a dict")
130+
assert "dictionary" in str(exc_info.value).lower()
131+
132+
with pytest.raises(SparseValuesDictionaryExpectedError):
133+
SparseValuesFactory.build(123)
134+
135+
with pytest.raises(SparseValuesDictionaryExpectedError):
136+
SparseValuesFactory.build([1, 2, 3])
137+
138+
def test_build_when_empty_indices_list(self):
139+
"""Test that empty indices list is handled correctly."""
140+
input_dict = {"indices": [], "values": []}
141+
actual = SparseValuesFactory.build(input_dict)
142+
expected = OpenApiSparseValues(indices=[], values=[])
143+
assert actual.indices == expected.indices
144+
assert actual.values == expected.values
145+
146+
def test_build_when_empty_values_list(self):
147+
"""Test that empty values list is handled correctly."""
148+
input_dict = {"indices": [], "values": []}
149+
actual = SparseValuesFactory.build(input_dict)
150+
expected = OpenApiSparseValues(indices=[], values=[])
151+
assert actual.indices == expected.indices
152+
assert actual.values == expected.values
Lines changed: 85 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,85 @@
1+
from unittest.mock import patch
2+
3+
from pinecone.utils import check_kwargs
4+
5+
6+
def example_function(arg1, arg2, arg3=None):
7+
"""Example function for testing."""
8+
pass
9+
10+
11+
class TestCheckKwargs:
12+
"""Test check_kwargs utility function."""
13+
14+
def test_no_unexpected_kwargs_no_logging(self):
15+
"""Test that no logging occurs when all kwargs are valid."""
16+
with patch("logging.exception") as mock_log:
17+
check_kwargs(example_function, {"arg1", "arg2", "arg3"})
18+
mock_log.assert_not_called()
19+
20+
def test_unexpected_kwargs_logs_warning(self):
21+
"""Test that unexpected kwargs trigger logging."""
22+
with patch("logging.exception") as mock_log:
23+
check_kwargs(example_function, {"arg1", "arg2", "arg3", "unexpected_arg"})
24+
mock_log.assert_called_once()
25+
call_args = mock_log.call_args[0][0]
26+
assert "unexpected keyword argument" in call_args.lower()
27+
assert "unexpected_arg" in call_args
28+
29+
def test_multiple_unexpected_kwargs_logs_all(self):
30+
"""Test that multiple unexpected kwargs are all logged."""
31+
with patch("logging.exception") as mock_log:
32+
check_kwargs(example_function, {"arg1", "arg2", "arg3", "unexpected1", "unexpected2"})
33+
mock_log.assert_called_once()
34+
call_args = mock_log.call_args[0][0]
35+
assert "unexpected1" in call_args or "unexpected2" in call_args
36+
37+
def test_only_unexpected_kwargs(self):
38+
"""Test when only unexpected kwargs are provided."""
39+
with patch("logging.exception") as mock_log:
40+
check_kwargs(example_function, {"unexpected_arg"})
41+
mock_log.assert_called_once()
42+
call_args = mock_log.call_args[0][0]
43+
assert "unexpected keyword argument" in call_args.lower()
44+
45+
def test_empty_kwargs_set(self):
46+
"""Test with empty kwargs set."""
47+
with patch("logging.exception") as mock_log:
48+
check_kwargs(example_function, set())
49+
mock_log.assert_not_called()
50+
51+
def test_function_with_no_args(self):
52+
"""Test with function that has no arguments."""
53+
54+
def no_args_function():
55+
pass
56+
57+
with patch("logging.exception") as mock_log:
58+
check_kwargs(no_args_function, {"any_arg"})
59+
mock_log.assert_called_once()
60+
61+
def test_function_with_varargs(self):
62+
"""Test with function that has *args."""
63+
64+
def varargs_function(*args):
65+
pass
66+
67+
with patch("logging.exception") as mock_log:
68+
check_kwargs(varargs_function, {"any_arg"})
69+
mock_log.assert_called_once()
70+
71+
def test_function_with_kwargs(self):
72+
"""Test with function that has **kwargs.
73+
74+
Note: check_kwargs only checks explicit args, not **kwargs,
75+
so it will still log unexpected args even for functions with **kwargs.
76+
This is the current behavior of the function.
77+
"""
78+
79+
def kwargs_function(**kwargs):
80+
pass
81+
82+
with patch("logging.exception") as mock_log:
83+
check_kwargs(kwargs_function, {"any_arg"})
84+
# check_kwargs doesn't check for **kwargs, so it will log
85+
mock_log.assert_called_once()

0 commit comments

Comments
 (0)