File size: 4,693 Bytes
f813ba1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Tests for the generation layer: LLM client, RAG chain, and cross-doc chain."""

from __future__ import annotations

from unittest.mock import AsyncMock, MagicMock, patch

import pytest


def test_llm_client_routes_model_correctly() -> None:
    """Test that LLM client routes to correct model based on question type."""
    with patch("app.core.generation.llm_client.Groq"):
        from app.core.generation.llm_client import LLMClient
        import app.core.generation.llm_client as llm_module
        llm_module._llm_client = None

        client = LLMClient()
        assert client.route_model("standard") == "llama-3.3-70b-versatile"
        assert client.route_model("reasoning") == "deepseek-r1-distill-llama-70b"
        assert client.route_model("compare") == "deepseek-r1-distill-llama-70b"


def test_llm_client_json_extraction() -> None:
    """Test JSON extraction from various response formats."""
    from app.core.generation.llm_client import _extract_json

    result = _extract_json('{"key": "value"}')
    assert result == {"key": "value"}

    result = _extract_json('```json\n{"key": "value"}\n```')
    assert result == {"key": "value"}

    result = _extract_json('Here is the result: {"score": 0.85}')
    assert result["score"] == 0.85


def test_llm_client_json_extraction_fails_gracefully() -> None:
    """Test that JSON extraction raises ValueError for invalid input."""
    from app.core.generation.llm_client import _extract_json
    with pytest.raises(ValueError, match="Could not extract valid JSON"):
        _extract_json("This is not JSON at all")


@pytest.mark.asyncio
@patch("app.core.generation.rag_chain.get_llm_client")
@patch("app.core.generation.rag_chain.retrieve_with_self_correction")
@patch("app.core.generation.rag_chain.get_vector_store")
async def test_standard_rag_chain_returns_answer(
    mock_vs: MagicMock, mock_retrieve: MagicMock, mock_llm: MagicMock,
) -> None:
    """Test that RAG chain returns a properly structured response."""
    mock_retrieve.return_value = (
        [{"id": "1", "score": 0.9, "payload": {
            "text": "ML is a subset of AI.", "doc_id": "doc-1",
            "filename": "paper.pdf", "chunk_index": 0,
        }}], 0.85, "What is ML?",
    )
    mock_client = MagicMock()
    mock_client.route_model.return_value = "llama-3.3-70b-versatile"
    mock_client.generate = AsyncMock(return_value="ML is a subset of AI.")
    mock_client.primary_model = "llama-3.3-70b-versatile"
    mock_client.reasoning_model = "deepseek-r1-distill-llama-70b"
    mock_llm.return_value = mock_client
    mock_store = MagicMock()
    mock_store.list_documents = AsyncMock(return_value=[])
    mock_vs.return_value = mock_store

    from app.core.generation.rag_chain import run_rag_chain
    from app.models.schemas import QueryRequest
    request = QueryRequest(question="What is ML?")
    response = await run_rag_chain(request)

    assert response.answer is not None
    assert response.model_used == "llama-3.3-70b-versatile"
    assert len(response.sources) == 1


@pytest.mark.asyncio
@patch("app.core.generation.cross_doc_chain.get_llm_client")
@patch("app.core.generation.cross_doc_chain.get_vector_store")
@patch("app.core.generation.cross_doc_chain.embed_query")
async def test_cross_doc_chain_returns_structured_output(
    mock_embed: MagicMock, mock_vs: MagicMock, mock_llm: MagicMock,
) -> None:
    """Test that cross-doc chain returns comparison with agreements."""
    mock_embed.return_value = [0.1] * 3072
    mock_store = MagicMock()
    mock_store.search = AsyncMock(side_effect=[
        [{"id": "1", "score": 0.9, "payload": {
            "text": "Paper A uses CNN.", "doc_id": "doc-1",
            "filename": "a.pdf", "chunk_index": 0}}],
        [{"id": "2", "score": 0.85, "payload": {
            "text": "Paper B uses RNN.", "doc_id": "doc-2",
            "filename": "b.pdf", "chunk_index": 0}}],
    ])
    mock_vs.return_value = mock_store
    mock_client = MagicMock()
    mock_client.reasoning_model = "deepseek-r1-distill-llama-70b"
    mock_client.generate = AsyncMock(return_value=(
        "Comparison text\n"
        'AGREEMENTS_JSON: ["Both classify"]\n'
        'CONTRADICTIONS_JSON: ["Different arch"]'
    ))
    mock_llm.return_value = mock_client

    from app.core.generation.cross_doc_chain import run_cross_doc_chain
    from app.models.schemas import CompareRequest
    request = CompareRequest(question="Compare", doc_ids=["doc-1", "doc-2"], aspect="methodology")
    response = await run_cross_doc_chain(request)

    assert response.comparison is not None
    assert response.model_used == "deepseek-r1-distill-llama-70b"
    assert mock_store.search.call_count == 2