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
16 changes: 8 additions & 8 deletions haystack/components/query/query_expander.py
Original file line number Diff line number Diff line change
Expand Up @@ -198,11 +198,11 @@ def run(self, query: str, n_expansions: int | None = None) -> dict[str, list[str

self.warm_up()

response = {"queries": [query] if self.include_original_query else []}

if not query.strip():
if not isinstance(query, str) or not query.strip():
logger.warning("Empty query provided to QueryExpander")
return response
return {"queries": [query] if self.include_original_query and isinstance(query, str) else []}

response = {"queries": [query] if self.include_original_query else []}

expansion_count = n_expansions if n_expansions is not None else self.n_expansions
if expansion_count <= 0:
Expand Down Expand Up @@ -268,11 +268,11 @@ async def run_async(self, query: str, n_expansions: int | None = None) -> dict[s

await self.warm_up_async()

response = {"queries": [query] if self.include_original_query else []}

if not query.strip():
if not isinstance(query, str) or not query.strip():
logger.warning("Empty query provided to QueryExpander")
return response
return {"queries": [query] if self.include_original_query and isinstance(query, str) else []}

response = {"queries": [query] if self.include_original_query else []}

expansion_count = n_expansions if n_expansions is not None else self.n_expansions
if expansion_count <= 0:
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
---
fixes:
- |
``QueryExpander`` now returns an empty query list when ``query`` is ``None`` or not a string,
instead of raising ``AttributeError`` on ``str.strip``.
14 changes: 14 additions & 0 deletions test/components/query/test_query_expander.py
Original file line number Diff line number Diff line change
Expand Up @@ -134,6 +134,20 @@ def test_run_whitespace_only_query(self, monkeypatch, caplog):
assert result["queries"] == ["\t\n \r"]
assert "Empty query provided" in caplog.text

def test_run_none_query(self, monkeypatch, caplog):
monkeypatch.setenv("OPENAI_API_KEY", "test-key-12345")
expander = QueryExpander()
with caplog.at_level(logging.WARNING):
result = expander.run(None) # type: ignore[arg-type]
assert result["queries"] == []
assert "Empty query provided" in caplog.text

def test_run_none_query_include_original(self, monkeypatch):
monkeypatch.setenv("OPENAI_API_KEY", "test-key-12345")
expander = QueryExpander(include_original_query=False)
result = expander.run(None) # type: ignore[arg-type]
assert result["queries"] == []

def test_run_generator_no_replies(self, mock_chat_generator):
mock_chat_generator.run.return_value = {"replies": []}
expander = QueryExpander(chat_generator=mock_chat_generator)
Expand Down
Loading