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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 3 additions & 3 deletions haystack/components/embedders/mock_document_embedder.py
Original file line number Diff line number Diff line change
Expand Up @@ -81,13 +81,13 @@ def __init__(
:param meta_fields_to_embed: List of metadata fields to embed along with the document text.
:param embedding_separator: Separator used to concatenate the metadata fields to the document text.
:param progress_bar: Accepted for interface compatibility with real Document Embedders and ignored.
:raises ValueError: If both `embedding` and `embedding_fn` are provided, if `dimension` is not positive, or
if `embedding` is an empty list.
:raises ValueError: If both `embedding` and `embedding_fn` are provided, if `embedding` is an empty list,
or if neither is provided and `dimension` is not positive.
:raises TypeError: If `embedding` is not a sequence of numbers.
"""
if embedding is not None and embedding_fn is not None:
raise ValueError("Pass either 'embedding' or 'embedding_fn', not both.")
if dimension <= 0:
if embedding is None and embedding_fn is None and dimension <= 0:
raise ValueError("'dimension' must be a positive integer.")

self.embedding = _coerce_embedding(embedding, name="'embedding'") if embedding is not None else None
Expand Down
6 changes: 3 additions & 3 deletions haystack/components/embedders/mock_text_embedder.py
Original file line number Diff line number Diff line change
Expand Up @@ -70,13 +70,13 @@ def __init__(
:param meta: Additional metadata merged into the output `meta`.
:param prefix: A string to add at the beginning of the text before embedding.
:param suffix: A string to add at the end of the text before embedding.
:raises ValueError: If both `embedding` and `embedding_fn` are provided, if `dimension` is not positive, or
if `embedding` is an empty list.
:raises ValueError: If both `embedding` and `embedding_fn` are provided, if `embedding` is an empty list,
or if neither is provided and `dimension` is not positive.
:raises TypeError: If `embedding` is not a sequence of numbers.
"""
if embedding is not None and embedding_fn is not None:
raise ValueError("Pass either 'embedding' or 'embedding_fn', not both.")
if dimension <= 0:
if embedding is None and embedding_fn is None and dimension <= 0:
raise ValueError("'dimension' must be a positive integer.")

self.embedding = _coerce_embedding(embedding, name="'embedding'") if embedding is not None else None
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,6 @@
fixes:
- |
``MockTextEmbedder`` and ``MockDocumentEmbedder`` now accept non-positive
``dimension`` values when ``embedding`` or ``embedding_fn`` is provided,
matching the documented behavior. The default deterministic embedding mode
still requires a positive ``dimension``.
13 changes: 9 additions & 4 deletions test/components/embedders/test_mock_document_embedder.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@ class TestMockDocumentEmbedder:
("args", "kwargs", "match"),
[
(([0.1],), {"embedding_fn": _ones}, "either 'embedding' or 'embedding_fn'"),
((), {"dimension": 0}, "must be a positive integer"),
((), {"dimension": -1}, "must be a positive integer"),
],
)
Expand All @@ -38,12 +39,16 @@ def test_consistent_with_text_embedder(self):
doc_embedding = MockDocumentEmbedder(dimension=8).run([Document(content="pizza")])["documents"][0].embedding
assert text_embedding == doc_embedding

def test_fixed_embedding(self):
result = MockDocumentEmbedder([0.5, 0.5]).run([Document(content="a"), Document(content="b")])
@pytest.mark.parametrize("dimension", [768, 0, -1])
def test_fixed_embedding(self, dimension):
embedder = MockDocumentEmbedder([0.5, 0.5], dimension=dimension)
result = embedder.run([Document(content="a"), Document(content="b")])
assert all(doc.embedding == [0.5, 0.5] for doc in result["documents"])

def test_embedding_fn(self):
result = MockDocumentEmbedder(embedding_fn=_ones).run([Document(content="a")])
@pytest.mark.parametrize("dimension", [768, 0, -1])
def test_embedding_fn(self, dimension):
embedder = MockDocumentEmbedder(embedding_fn=_ones, dimension=dimension)
result = embedder.run([Document(content="a")])
assert result["documents"][0].embedding == [1.0, 1.0, 1.0]

def test_meta_fields_to_embed_affect_embedding(self):
Expand Down
12 changes: 8 additions & 4 deletions test/components/embedders/test_mock_text_embedder.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,7 @@ class TestMockTextEmbedder:
[
(([0.1, 0.2],), {"embedding_fn": _ones}, ValueError, "either 'embedding' or 'embedding_fn'"),
((), {"dimension": 0}, ValueError, "must be a positive integer"),
((), {"dimension": -1}, ValueError, "must be a positive integer"),
(([],), {}, ValueError, "must not be empty"),
((["not", "numbers"],), {}, TypeError, "must be a sequence of numbers"),
],
Expand All @@ -51,13 +52,16 @@ def test_deterministic_distinguishes_texts(self):
MockTextEmbedder(dimension=8).run("x")["embedding"] == MockTextEmbedder(dimension=8).run("x")["embedding"]
)

def test_fixed_embedding(self):
embedder = MockTextEmbedder([0.1, 0.2, 0.3])
@pytest.mark.parametrize("dimension", [768, 0, -1])
def test_fixed_embedding(self, dimension):
embedder = MockTextEmbedder([0.1, 0.2, 0.3], dimension=dimension)
assert embedder.run("anything")["embedding"] == [0.1, 0.2, 0.3]
assert embedder.run("something else")["embedding"] == [0.1, 0.2, 0.3]

def test_embedding_fn(self):
assert MockTextEmbedder(embedding_fn=_ones).run("hello")["embedding"] == [1.0, 1.0, 1.0]
@pytest.mark.parametrize("dimension", [768, 0, -1])
def test_embedding_fn(self, dimension):
embedder = MockTextEmbedder(embedding_fn=_ones, dimension=dimension)
assert embedder.run("hello")["embedding"] == [1.0, 1.0, 1.0]

def test_embedding_fn_invalid_return_raises(self):
# embedding_fn deliberately returns a non-vector to exercise the runtime type check
Expand Down
Loading