LightRAG源码:NetworkXStorage测试(2)
·
测试代码与逻辑解释
1. 导入依赖
import os
import shutil
import pytest
import networkx as nx
import numpy as np
import asyncio
import json
from lightrag import LightRAG
from lightrag.storage import NetworkXStorage
from lightrag.utils import wrap_embedding_func_with_attrs
- 导入了必要的Python库和模块:
os和shutil:用于文件和目录操作。pytest:用于编写和运行测试。networkx:用于图结构的创建和操作。numpy:用于数值计算。asyncio:用于异步编程。json:用于JSON数据的处理。LightRAG、NetworkXStorage和wrap_embedding_func_with_attrs:来自自定义模块lightrag,分别用于RAG(Retrieval-Augmented Generation)框架、图存储和嵌入函数的装饰器。
2. 定义工作目录
WORKING_DIR = "./tests/nano_graphrag_cache_networkx_storage_test"
- 定义了一个工作目录路径
WORKING_DIR,用于存储测试过程中生成的文件。
3. 测试环境的设置与清理
@pytest.fixture(scope="function")
def setup_teardown():
if os.path.exists(WORKING_DIR):
shutil.rmtree(WORKING_DIR)
os.mkdir(WORKING_DIR)
yield
shutil.rmtree(WORKING_DIR)
- 使用
pytest.fixture定义了一个测试环境的设置与清理函数setup_teardown:- 在每个测试函数运行前,检查并删除
WORKING_DIR目录(如果存在),然后重新创建该目录。 - 在测试函数运行后,再次删除
WORKING_DIR目录,确保测试环境的干净。
- 在每个测试函数运行前,检查并删除
4. 模拟嵌入函数
@wrap_embedding_func_with_attrs(embedding_dim=384, max_token_size=8192)
async def mock_embedding(texts: list[str]) -> np.ndarray:
return np.random.rand(len(texts), 384)
- 使用
wrap_embedding_func_with_attrs装饰器定义了一个模拟的嵌入函数mock_embedding:- 输入为文本列表
texts,输出为一个随机的numpy数组,形状为(len(texts), 384)。 - 该函数用于模拟文本嵌入的过程。
- 输入为文本列表
5. 初始化 NetworkXStorage
@pytest.fixture
def networkx_storage(setup_teardown):
rag = LightRAG(working_dir=WORKING_DIR, embedding_func=mock_embedding)
return NetworkXStorage(
namespace="test",
global_config=rag.__dict__,
)
- 使用
pytest.fixture定义了一个networkx_storage的初始化函数:- 创建了一个
LightRAG实例rag,并传入工作目录和模拟嵌入函数。 - 返回一个
NetworkXStorage实例,用于存储图数据。
- 创建了一个
6. 测试持久化功能
@pytest.mark.asyncio
async def test_persistence(setup_teardown):
rag = LightRAG(working_dir=WORKING_DIR, embedding_func=mock_embedding)
initial_storage = NetworkXStorage(
namespace="test_persistence",
global_config=rag.__dict__,
)
await initial_storage.upsert_node("node1", {"attr": "value"})
await initial_storage.upsert_node("node2", {"attr": "value"})
await initial_storage.upsert_edge("node1", "node2", {"weight": 1.0})
await initial_storage.index_done_callback()
new_storage = NetworkXStorage(
namespace="test_persistence",
global_config=rag.__dict__,
)
assert await new_storage.has_node("node1")
assert await new_storage.has_node("node2")
assert await new_storage.has_edge("node1", "node2")
node1_data = await new_storage.get_node("node1")
assert node1_data == {"attr": "value"}
edge_data = await new_storage.get_edge("node1", "node2")
assert edge_data == {"weight": 1.0}
- 测试
NetworkXStorage的持久化功能:- 创建初始存储实例
initial_storage,并插入两个节点和一条边。 - 调用
index_done_callback方法,确保数据被持久化。 - 创建新的存储实例
new_storage,验证节点和边是否被正确加载。 - 检查节点和边的属性是否正确。
- 创建初始存储实例
7. 测试节点嵌入功能
@pytest.mark.asyncio
async def test_embed_nodes(networkx_storage):
for i in range(5):
await networkx_storage.upsert_node(f"node{i}", {"id": f"node{i}"})
for i in range(4):
await networkx_storage.upsert_edge(f"node{i}", f"node{i+1}", {})
embeddings, node_ids = await networkx_storage.embed_nodes("node2vec")
assert embeddings.shape == (
5,
networkx_storage.global_config["node2vec_params"]["dimensions"],
)
assert len(node_ids) == 5
assert all(f"node{i}" in node_ids for i in range(5))
- 测试
NetworkXStorage的节点嵌入功能:- 插入5个节点和4条边。
- 使用
node2vec方法生成节点嵌入。 - 验证嵌入结果的形状和节点ID是否正确。
8. 测试稳定最大连通分量
@pytest.mark.asyncio
async def test_stable_largest_connected_component_equal_components():
G = nx.Graph()
G.add_edges_from([("A", "B"), ("C", "D"), ("E", "F")])
result = NetworkXStorage.stable_largest_connected_component(G)
assert sorted(result.nodes()) == ["A", "B"]
assert list(result.edges()) == [("A", "B")]
- 测试
stable_largest_connected_component方法:- 创建一个包含多个连通分量的图。
- 验证返回的最大连通分量是否正确。
9. 测试稳定最大连通分量的稳定性
@pytest.mark.asyncio
async def test_stable_largest_connected_component_stability():
G = nx.Graph()
G.add_edges_from([("A", "B"), ("B", "C"), ("C", "D"), ("E", "F")])
result1 = NetworkXStorage.stable_largest_connected_component(G)
result2 = NetworkXStorage.stable_largest_connected_component(G)
assert nx.is_isomorphic(result1, result2)
assert list(result1.nodes()) == list(result2.nodes())
assert list(result1.edges()) == list(result2.edges())
- 测试
stable_largest_connected_component方法的稳定性:- 验证多次调用该方法返回的结果是否一致。
10. 测试有向图的稳定最大连通分量
@pytest.mark.asyncio
async def test_stable_largest_connected_component_directed_graph():
G = nx.DiGraph()
G.add_edges_from([("A", "B"), ("B", "C"), ("C", "D"), ("E", "F")])
result = NetworkXStorage.stable_largest_connected_component(G)
assert sorted(result.nodes()) == ["A", "B", "C", "D"]
assert sorted(result.edges()) == [("A", "B"), ("B", "C"), ("C", "D")]
- 测试
stable_largest_connected_component方法在有向图中的表现:- 验证返回的最大连通分量是否正确。
11. 测试自环和平行边的稳定最大连通分量
@pytest.mark.asyncio
async def test_stable_largest_connected_component_self_loops_and_parallel_edges():
G = nx.Graph()
G.add_edges_from(
[("A", "B"), ("B", "C"), ("C", "A"), ("A", "A"), ("B", "B"), ("A", "B")]
)
result = NetworkXStorage.stable_largest_connected_component(G)
assert sorted(result.nodes()) == ["A", "B", "C"]
assert sorted(result.edges()) == [
("A", "A"),
("A", "B"),
("A", "C"),
("B", "B"),
("B", "C"),
]
- 测试
stable_largest_connected_component方法在包含自环和平行边的图中的表现:- 验证返回的最大连通分量是否正确。
总结
- 这些测试代码主要验证了
NetworkXStorage类的功能,包括持久化、节点嵌入、最大连通分量等。 - 通过
pytest和asyncio,测试代码以异步方式运行,确保代码的正确性和稳定性。 - 每个测试函数都专注于一个特定的功能点,并通过断言验证结果是否符合预期。
更多推荐


所有评论(0)