File size: 9,408 Bytes
7ab7df1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1521cd2
7ab7df1
1521cd2
7ab7df1
1521cd2
7ab7df1
1521cd2
7ab7df1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1521cd2
 
7ab7df1
1521cd2
7ab7df1
 
 
 
 
1521cd2
7ab7df1
 
 
 
 
 
1521cd2
7ab7df1
1521cd2
7ab7df1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
"""Tests for input validation and edge cases."""

from unittest.mock import MagicMock, Mock, patch

import pytest

from qa_chain import QAChainWrapper
from langchain_core.prompts import ChatPromptTemplate
from langchain_core.documents import Document
from ui.handlers import create_stream_chat_response, create_respond_handler


@pytest.fixture
def qa_chain_wrapper(mock_vectorstore):
    """Create a QAChainWrapper instance."""
    prompt = ChatPromptTemplate.from_template("Test: {question}")
    return QAChainWrapper(mock_vectorstore, prompt)


@patch("qa_chain.create_llm")
@patch("qa_chain.format_chat_history")
def test_very_long_query(mock_format_history, mock_create_llm, qa_chain_wrapper):
    """Test handling of extremely long queries (>10k characters)."""
    very_long_query = "What is " + "RAG? " * 2000  # ~10k+ characters

    mock_format_history.return_value = ""
    mock_llm = MagicMock()
    mock_chunk = MagicMock()
    mock_chunk.content = "Response"
    mock_llm.stream.return_value = [mock_chunk]
    mock_create_llm.return_value = mock_llm

    mock_retriever = MagicMock()
    mock_retriever.invoke.return_value = [Document(page_content="Test", metadata={})]
    qa_chain_wrapper._retriever = mock_retriever

    mock_chain = MagicMock()
    mock_chain.stream.return_value = [mock_chunk]
    qa_chain_wrapper._prompt.__or__ = MagicMock(return_value=mock_chain)

    inputs = {
        "question": very_long_query,
        "chat_history": [],
    }

    # Should handle long query without crashing
    results = list(qa_chain_wrapper.stream(inputs))
    assert len(results) > 0


def test_empty_whitespace_query():
    """Test handling of whitespace-only queries."""
    from ui.handlers import create_respond_handler

    stream_fn = MagicMock()
    respond_fn = create_respond_handler(stream_fn)

    # Empty string
    results = list(respond_fn("", [], False, "mmr", "All Documents", False, False, 70))
    assert results[0][0] == ""  # Should return empty immediately

    # Whitespace only - now also returns empty (treated as empty after strip)
    results = list(respond_fn("   ", [], False, "mmr", "All Documents", False, False, 70))
    assert results[0][0] == ""  # Should return empty immediately

    # Newlines only - also returns empty
    results = list(respond_fn("\n\n\n", [], False, "mmr", "All Documents", False, False, 70))
    assert results[0][0] == ""  # Should return empty immediately


@patch("qa_chain.create_llm")
@patch("qa_chain.format_chat_history")
def test_special_characters_in_query(mock_format_history, mock_create_llm, qa_chain_wrapper):
    """Test handling of special characters in queries."""
    special_chars_query = "What is RAG? @#$%^&*()[]{}|\\/<>?~`"

    mock_format_history.return_value = ""
    mock_llm = MagicMock()
    mock_chunk = MagicMock()
    mock_chunk.content = "Response"
    mock_llm.stream.return_value = [mock_chunk]
    mock_create_llm.return_value = mock_llm

    mock_retriever = MagicMock()
    mock_retriever.invoke.return_value = [Document(page_content="Test", metadata={})]
    qa_chain_wrapper._retriever = mock_retriever

    mock_chain = MagicMock()
    mock_chain.stream.return_value = [mock_chunk]
    qa_chain_wrapper._prompt.__or__ = MagicMock(return_value=mock_chain)

    inputs = {
        "question": special_chars_query,
        "chat_history": [],
    }

    # Should handle special characters without crashing
    results = list(qa_chain_wrapper.stream(inputs))
    assert len(results) > 0


@patch("qa_chain.create_llm")
@patch("qa_chain.format_chat_history")
def test_unicode_emoji_in_query(mock_format_history, mock_create_llm, qa_chain_wrapper):
    """Test handling of Unicode and emoji in queries."""
    unicode_query = "What is RAG? πŸš€ δ½ ε₯½ Ω…Ψ±Ψ­Ψ¨Ψ§"

    mock_format_history.return_value = ""
    mock_llm = MagicMock()
    mock_chunk = MagicMock()
    mock_chunk.content = "Response"
    mock_llm.stream.return_value = [mock_chunk]
    mock_create_llm.return_value = mock_llm

    mock_retriever = MagicMock()
    mock_retriever.invoke.return_value = [Document(page_content="Test", metadata={})]
    qa_chain_wrapper._retriever = mock_retriever

    mock_chain = MagicMock()
    mock_chain.stream.return_value = [mock_chunk]
    qa_chain_wrapper._prompt.__or__ = MagicMock(return_value=mock_chain)

    inputs = {
        "question": unicode_query,
        "chat_history": [],
    }

    # Should handle Unicode/emoji without crashing
    results = list(qa_chain_wrapper.stream(inputs))
    assert len(results) > 0


def test_invalid_document_filter():
    """Test handling of invalid document filter selection."""
    from ui.handlers import create_stream_chat_response

    mock_qa_chain = MagicMock()
    mock_qa_chain.stream.return_value = [
        {
            "chunk": "Response",
            "source_documents": [],
            "docs_with_scores": None,
            "rewritten_query": None,
            "hybrid_scores": None,
        }
    ]

    # Pass valid sources - nonexistent.pdf is NOT in the list
    stream_fn = create_stream_chat_response(mock_qa_chain, available_sources=["valid.pdf"])

    # Invalid filter (document that doesn't exist in available_sources)
    results = list(
        stream_fn(
            "test question",
            [],
            "RAG",
            doc_filter="nonexistent.pdf",  # Invalid filter - not in available_sources
            search_type="mmr",
        )
    )

    # Should handle invalid filter gracefully
    assert len(results) > 0
    # Invalid filter should be ignored (no filter passed to chain)
    call_args = mock_qa_chain.stream.call_args[0][0]
    assert "filter" not in call_args  # Filter is NOT passed for invalid sources


def test_malformed_chat_history():
    """Test handling of malformed chat history."""
    from ui.handlers import messages_to_tuples

    # Missing role - should handle gracefully
    try:
        malformed_history1 = [
            {"content": "Message without role"},
        ]
        tuples1 = messages_to_tuples(malformed_history1)
        assert isinstance(tuples1, list)
        # Should return empty (no valid user/assistant pairs)
        assert tuples1 == []
    except (KeyError, TypeError):
        # Exception is acceptable for malformed input
        pass

    # Missing content - should handle gracefully
    try:
        malformed_history2 = [
            {"role": "user"},
        ]
        tuples2 = messages_to_tuples(malformed_history2)
        assert isinstance(tuples2, list)
    except (KeyError, TypeError):
        # Exception is acceptable for malformed input
        pass

    # Invalid role - should handle gracefully
    malformed_history3 = [
        {"role": "invalid", "content": "Message"},
    ]
    tuples3 = messages_to_tuples(malformed_history3)
    # Should return empty (only processes user/assistant pairs)
    assert isinstance(tuples3, list)
    assert tuples3 == []


@patch("qa_chain.create_llm")
def test_query_rewriting_with_very_short_query(mock_create_llm, qa_chain_wrapper):
    """Test query rewriting with very short query."""
    very_short_query = "RAG?"

    mock_llm = MagicMock()
    mock_response = MagicMock()
    # Make rewritten query shorter than 30% of original
    # "RAG?" is 4 chars, 30% = 1.2, so "R" (1 char) should trigger fallback
    mock_response.content = "R"
    mock_llm.invoke.return_value = mock_response
    mock_create_llm.return_value = mock_llm

    result = qa_chain_wrapper.rewrite_query(very_short_query)

    # Should return original if rewritten is too short (< 30% of original length)
    assert result == very_short_query


@patch("qa_chain.create_llm")
@patch("qa_chain.format_chat_history")
def test_hybrid_search_with_extreme_alpha(mock_format_history, mock_create_llm, qa_chain_wrapper):
    """Test hybrid search with extreme alpha values."""
    from langchain_core.documents import Document

    mock_format_history.return_value = ""
    mock_llm = MagicMock()
    mock_chunk = MagicMock()
    mock_chunk.content = "Response"
    mock_llm.stream.return_value = [mock_chunk]
    mock_create_llm.return_value = mock_llm

    mock_retriever = MagicMock()
    mock_retriever.invoke.return_value = [Document(page_content="Test", metadata={})]
    qa_chain_wrapper._retriever = mock_retriever

    mock_chain = MagicMock()
    mock_chain.stream.return_value = [mock_chunk]
    qa_chain_wrapper._prompt.__or__ = MagicMock(return_value=mock_chain)

    # Test with alpha = 0.0 (pure keyword)
    inputs = {
        "question": "test",
        "chat_history": [],
        "search_type": "hybrid",
        "hybrid_alpha": 0.0,
    }
    results = list(qa_chain_wrapper.stream(inputs))
    assert len(results) > 0

    # Test with alpha = 1.0 (pure semantic)
    inputs["hybrid_alpha"] = 1.0
    results = list(qa_chain_wrapper.stream(inputs))
    assert len(results) > 0


def test_rerank_with_very_few_documents(qa_chain_wrapper, sample_documents):
    """Test re-ranking with very few documents."""
    # Re-rank with only 1 document
    result = qa_chain_wrapper.rerank_documents("query", sample_documents[:1], top_k=5)
    
    # Should return the single document
    assert len(result) == 1

    # Re-rank with more documents than top_k
    result = qa_chain_wrapper.rerank_documents("query", sample_documents[:10], top_k=3)
    
    # Should return only top_k documents
    assert len(result) == 3