File size: 2,985 Bytes
0828c2c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Tests for retrieval base classes and factory."""

import pytest

from src.retrieval import RetrieverFactory
from src.retrieval.base import BaseRetrieverStrategy

# --- RetrieverFactory Tests ---


def test_available_strategies():
    """Test that default strategies are registered."""
    strategies = RetrieverFactory.available_strategies()
    assert "vector" in strategies
    assert "bm25" in strategies
    assert "bm25_vector" in strategies


def test_is_registered():
    """Test is_registered method."""
    assert RetrieverFactory.is_registered("vector")
    assert RetrieverFactory.is_registered("bm25")
    assert RetrieverFactory.is_registered("bm25_vector")
    assert not RetrieverFactory.is_registered("nonexistent")


def test_create_unknown_strategy_raises():
    """Test that creating unknown strategy raises ValueError."""
    with pytest.raises(ValueError, match="Unknown retrieval strategy"):
        RetrieverFactory.create("nonexistent_strategy", {})


def test_create_vector_strategy():
    """Test creating vector strategy."""
    config = {"retrieval": {"vector": {"search_type": "similarity", "k": 4}}}
    strategy = RetrieverFactory.create("vector", config)
    assert strategy.name == "vector"
    assert isinstance(strategy, BaseRetrieverStrategy)


def test_create_bm25_strategy():
    """Test creating BM25 strategy."""
    config = {
        "retrieval": {
            "bm25": {
                "k": 10,
                "persist_path": "./test_bm25_index",
                "tokenizer": "simple",
            }
        }
    }
    strategy = RetrieverFactory.create("bm25", config)
    assert strategy.name == "bm25"
    assert isinstance(strategy, BaseRetrieverStrategy)


def test_create_bm25_vector_strategy():
    """Test creating BM25+Vector strategy."""
    config = {
        "retrieval": {
            "final_k": 4,
            "vector": {"search_type": "similarity", "k": 10},
            "bm25": {"k": 10, "persist_path": "./test_bm25_index"},
            "fusion": {
                "algorithm": "rrf",
                "rrf_k": 60,
                "weights": {"vector": 0.7, "bm25": 0.3},
            },
        }
    }
    strategy = RetrieverFactory.create("bm25_vector", config)
    assert strategy.name == "bm25_vector"
    assert isinstance(strategy, BaseRetrieverStrategy)


def test_register_custom_strategy():
    """Test registering a custom strategy."""

    @RetrieverFactory.register("test_custom")
    class CustomStrategy(BaseRetrieverStrategy):
        @property
        def name(self):
            return "test_custom"

        def build_index(self, documents):
            pass

        def load_index(self):
            return True

        def retrieve(self, query, k=4):
            return []

        def as_retriever(self, **kwargs):
            return None

    assert RetrieverFactory.is_registered("test_custom")
    strategy = RetrieverFactory.create("test_custom", {})
    assert strategy.name == "test_custom"