File size: 3,344 Bytes
8e874f5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Retrievers module - modular graph retrieval algorithms.

This module provides pluggable retrieval algorithms for knowledge graph traversal.
Similar to how neural network frameworks organize optimizers (Adam, SGD, etc.),
this module organizes graph traversal algorithms.

Available Retrievers:
    - FlowDiffusionRetriever: Query-Aware Weighted Flow Diffusion algorithm

Usage:
    # Direct import
    from src.retrievers import FlowDiffusionRetriever

    retriever = FlowDiffusionRetriever(
        graph=my_graph,
        source_node="node_a",
        target_node="node_b",
        confidence=0.5
    )
    result = retriever.retrieve()

    # Or use the registry
    from src.retrievers import get_retriever

    retriever = get_retriever(
        "flow_diffusion",
        graph=my_graph,
        source_node="node_a",
        target_node="node_b"
    )

Adding Custom Retrievers:
    from src.retrievers import BaseRetriever, register_retriever

    class MyCustomRetriever(BaseRetriever):
        def retrieve(self, source_node, target_node=None, **kwargs):
            # Implementation
            pass

        def find_path(self, source, target):
            # Implementation
            pass

    register_retriever("my_custom", MyCustomRetriever)
"""

from .base import BaseRetriever, RetrieverResult
from .flow_diffusion import FlowDiffusionRetriever, QueryAwareWeightedFlowDiffusion

# Registry for retriever classes
RETRIEVER_REGISTRY = {
    "flow_diffusion": FlowDiffusionRetriever,
    "qafd": FlowDiffusionRetriever,  # Alias
}


def get_retriever(name: str, **kwargs) -> BaseRetriever:
    """Factory function to get a retriever by name.

    Args:
        name: Retriever name (e.g., "flow_diffusion", "qafd")
        **kwargs: Arguments to pass to the retriever constructor

    Returns:
        Instantiated retriever

    Raises:
        ValueError: If retriever name is not found

    Example:
        retriever = get_retriever(
            "flow_diffusion",
            graph=my_graph,
            source_node="node_a",
            target_node="node_b",
            confidence=0.5
        )
    """
    if name not in RETRIEVER_REGISTRY:
        available = ", ".join(RETRIEVER_REGISTRY.keys())
        raise ValueError(f"Unknown retriever: '{name}'. Available retrievers: {available}")

    return RETRIEVER_REGISTRY[name](**kwargs)


def register_retriever(name: str, cls: type):
    """Register a custom retriever class.

    Args:
        name: Name to register the retriever under
        cls: Retriever class (should inherit from BaseRetriever)

    Example:
        register_retriever("my_retriever", MyRetrieverClass)
    """
    if not issubclass(cls, BaseRetriever):
        raise TypeError(f"Retriever class must inherit from BaseRetriever, got {cls}")
    RETRIEVER_REGISTRY[name] = cls


def list_retrievers() -> list:
    """List all available retriever names.

    Returns:
        List of registered retriever names
    """
    return list(RETRIEVER_REGISTRY.keys())


__all__ = [
    # Base classes
    "BaseRetriever",
    "RetrieverResult",
    # Retriever implementations
    "FlowDiffusionRetriever",
    "QueryAwareWeightedFlowDiffusion",  # Backward compatibility alias
    # Registry functions
    "get_retriever",
    "register_retriever",
    "list_retrievers",
    "RETRIEVER_REGISTRY",
]