Skip to content

Commit 37d6fce

Browse files
authored
style: apply ruff lint fixes and code formatting (#26)
- Fix import sorting (I001) across all modules - Remove unused imports (F401) and variables (F841) - Fix bare except clauses (E722) in url_fetcher.py - Remove duplicate DataType import (F811) in milvus.py - Remove trailing whitespace in docstrings (W293) - Fix module-level import ordering (E402) - Remove unused num_select parameter from LLMReranker.rerank() - Apply consistent code formatting with ruff format Signed-off-by: Cheney Zhang <chen.zhang@zilliz.com>
1 parent 575b8a5 commit 37d6fce

21 files changed

Lines changed: 417 additions & 397 deletions

src/vector_graph_rag/__init__.py

Lines changed: 7 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -6,17 +6,17 @@
66
"""
77

88
from vector_graph_rag.config import Settings
9-
from vector_graph_rag.models import Document, Triplet, Entity, Relation, Passage
10-
from vector_graph_rag.llm.extractor import TripletExtractor
11-
from vector_graph_rag.storage.embeddings import EmbeddingModel
12-
from vector_graph_rag.storage.milvus import MilvusStore
139
from vector_graph_rag.graph.builder import GraphBuilder
14-
from vector_graph_rag.graph.retriever import GraphRetriever
15-
from vector_graph_rag.graph.knowledge_graph import SubGraph
1610
from vector_graph_rag.graph.graph import Graph
11+
from vector_graph_rag.graph.knowledge_graph import SubGraph
12+
from vector_graph_rag.graph.retriever import GraphRetriever
13+
from vector_graph_rag.llm.cache import LLMCache, get_llm_cache
14+
from vector_graph_rag.llm.extractor import TripletExtractor
1715
from vector_graph_rag.llm.reranker import LLMReranker
16+
from vector_graph_rag.models import Document, Entity, Passage, Relation, Triplet
1817
from vector_graph_rag.rag import VectorGraphRAG, create_rag
19-
from vector_graph_rag.llm.cache import LLMCache, get_llm_cache
18+
from vector_graph_rag.storage.embeddings import EmbeddingModel
19+
from vector_graph_rag.storage.milvus import MilvusStore
2020

2121
__version__ = "0.1.3"
2222

src/vector_graph_rag/api/app.py

Lines changed: 130 additions & 68 deletions
Large diffs are not rendered by default.

src/vector_graph_rag/config.py

Lines changed: 14 additions & 27 deletions
Original file line numberDiff line numberDiff line change
@@ -3,9 +3,10 @@
33
"""
44

55
import os
6-
from typing import Optional, Dict, Any
7-
from pydantic_settings import BaseSettings
6+
from typing import Any, Dict, Optional
7+
88
from pydantic import Field
9+
from pydantic_settings import BaseSettings
910

1011

1112
class Settings(BaseSettings):
@@ -21,9 +22,7 @@ class Settings(BaseSettings):
2122
default_factory=lambda: os.getenv("OPENAI_API_KEY"),
2223
description="OpenAI API key for LLM and embeddings",
2324
)
24-
openai_base_url: Optional[str] = Field(
25-
default=None, description="Custom OpenAI API base URL"
26-
)
25+
openai_base_url: Optional[str] = Field(default=None, description="Custom OpenAI API base URL")
2726

2827
# Model Settings
2928
llm_model: str = Field(
@@ -80,12 +79,8 @@ class Settings(BaseSettings):
8079
)
8180

8281
# Retrieval Settings
83-
entity_top_k: int = Field(
84-
default=20, description="Number of top entities to retrieve"
85-
)
86-
relation_top_k: int = Field(
87-
default=20, description="Number of top relations to retrieve"
88-
)
82+
entity_top_k: int = Field(default=20, description="Number of top entities to retrieve")
83+
relation_top_k: int = Field(default=20, description="Number of top relations to retrieve")
8984
entity_similarity_threshold: float = Field(
9085
default=0.9,
9186
description="Similarity threshold for entity retrieval (keep if score > threshold)",
@@ -101,31 +96,23 @@ class Settings(BaseSettings):
10196
default=1000,
10297
description="Maximum number of expanded relations. If exceeded, use eviction strategy to filter by similarity.",
10398
)
104-
final_top_k: int = Field(
105-
default=3, description="Number of final passages to return"
106-
)
99+
final_top_k: int = Field(default=3, description="Number of final passages to return")
107100

108101
# LLM Settings
109-
llm_temperature: float = Field(
110-
default=0.0, description="Temperature for LLM generation"
111-
)
112-
llm_max_retries: int = Field(
113-
default=3, description="Maximum retries for LLM API calls"
114-
)
115-
use_llm_cache: bool = Field(
116-
default=True, description="Whether to use LLM response caching"
117-
)
102+
llm_temperature: float = Field(default=0.0, description="Temperature for LLM generation")
103+
llm_max_retries: int = Field(default=3, description="Maximum retries for LLM API calls")
104+
use_llm_cache: bool = Field(default=True, description="Whether to use LLM response caching")
118105

119106
# Processing Settings
120-
batch_size: int = Field(
121-
default=32, description="Batch size for embedding and insertion"
122-
)
107+
batch_size: int = Field(default=32, description="Batch size for embedding and insertion")
123108

124109
# NER Cache Settings
125110
ner_cache_dir: Optional[str] = Field(
126111
default_factory=lambda: os.path.join(
127112
os.path.dirname(os.path.dirname(os.path.dirname(__file__))), # -> vector-graph-rag/
128-
"evaluation", "data", "ner_cache"
113+
"evaluation",
114+
"data",
115+
"ner_cache",
129116
),
130117
description="Directory containing NER cache TSV files (HippoRAG format). "
131118
"Files should be named {dataset}_queries.named_entity_output.tsv",

src/vector_graph_rag/graph/__init__.py

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -3,14 +3,14 @@
33
"""
44

55
from vector_graph_rag.graph.builder import GraphBuilder
6-
from vector_graph_rag.graph.retriever import GraphRetriever, RetrievalResult
6+
from vector_graph_rag.graph.graph import Graph
77
from vector_graph_rag.graph.knowledge_graph import (
8-
SubGraph,
98
GraphEntity,
10-
GraphRelation,
119
GraphPassage,
10+
GraphRelation,
11+
SubGraph,
1212
)
13-
from vector_graph_rag.graph.graph import Graph
13+
from vector_graph_rag.graph.retriever import GraphRetriever, RetrievalResult
1414

1515
__all__ = [
1616
"Graph",

src/vector_graph_rag/graph/builder.py

Lines changed: 5 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -4,17 +4,17 @@
44

55
import uuid
66
from collections import defaultdict
7-
from typing import List, Dict, Optional
7+
from typing import Dict, List, Optional
88

9+
from vector_graph_rag.config import Settings, get_settings
10+
from vector_graph_rag.llm.extractor import processing_phrases
911
from vector_graph_rag.models import (
1012
Document,
1113
Entity,
14+
ExtractionResult,
1215
Relation,
1316
Triplet,
14-
ExtractionResult,
1517
)
16-
from vector_graph_rag.config import Settings, get_settings
17-
from vector_graph_rag.llm.extractor import processing_phrases
1818

1919

2020
def generate_id() -> str:
@@ -171,10 +171,7 @@ def build_from_documents(self, documents: List[Document]) -> ExtractionResult:
171171
self._process_documents(documents)
172172

173173
# Build result
174-
entities = [
175-
Entity(id=eid, name=self.entities[eid])
176-
for eid in self.entity_ids
177-
]
174+
entities = [Entity(id=eid, name=self.entities[eid]) for eid in self.entity_ids]
178175

179176
relations = []
180177
for rid in self.relation_ids:

src/vector_graph_rag/graph/graph.py

Lines changed: 45 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -11,14 +11,14 @@
1111
- Internal operations automatically handle Entity/Relation creation and linking
1212
"""
1313

14-
from typing import List, Optional, Dict, Any
14+
from typing import Any, Dict, List, Optional
1515

1616
from vector_graph_rag.config import Settings, get_settings
17-
from vector_graph_rag.storage.milvus import MilvusStore, generate_id
18-
from vector_graph_rag.storage.embeddings import EmbeddingModel
19-
from vector_graph_rag.models import Entity, Relation, Passage, Triplet
2017
from vector_graph_rag.graph.knowledge_graph import SubGraph
2118
from vector_graph_rag.llm.extractor import processing_phrases
19+
from vector_graph_rag.models import Entity, Passage, Relation, Triplet
20+
from vector_graph_rag.storage.embeddings import EmbeddingModel
21+
from vector_graph_rag.storage.milvus import MilvusStore, generate_id
2222

2323

2424
class Graph:
@@ -128,8 +128,12 @@ def _create_entity(
128128
existing = self._store._get_entities_by_ids([existing_id])
129129
if existing:
130130
current = existing[0]
131-
new_relation_ids = list(set(current.get("relation_ids", []) + (relation_ids or [])))
132-
new_passage_ids = list(set(current.get("passage_ids", []) + (passage_ids or [])))
131+
new_relation_ids = list(
132+
set(current.get("relation_ids", []) + (relation_ids or []))
133+
)
134+
new_passage_ids = list(
135+
set(current.get("passage_ids", []) + (passage_ids or []))
136+
)
133137
self._store._update_entity(
134138
existing_id,
135139
relation_ids=new_relation_ids,
@@ -326,8 +330,12 @@ def _create_relation(
326330
relation_id = id or generate_id()
327331

328332
# Create entities (if not exist) and get their IDs
329-
subject_id = self._create_entity(subject, relation_ids=[relation_id], passage_ids=passage_ids)
330-
object_id = self._create_entity(object_, relation_ids=[relation_id], passage_ids=passage_ids)
333+
subject_id = self._create_entity(
334+
subject, relation_ids=[relation_id], passage_ids=passage_ids
335+
)
336+
object_id = self._create_entity(
337+
object_, relation_ids=[relation_id], passage_ids=passage_ids
338+
)
331339

332340
# Generate embedding
333341
embedding = self._embedding_model.embed(relation_text)
@@ -413,8 +421,12 @@ def _update_relation(
413421
# Build new text if any triplet field changed
414422
new_text = None
415423
if subject or predicate or object_:
416-
new_subject = self._normalize_entity_name(subject) if subject else data.get("subject", "")
417-
new_predicate = processing_phrases(predicate) if predicate else data.get("predicate", "")
424+
new_subject = (
425+
self._normalize_entity_name(subject) if subject else data.get("subject", "")
426+
)
427+
new_predicate = (
428+
processing_phrases(predicate) if predicate else data.get("predicate", "")
429+
)
418430
new_object = self._normalize_entity_name(object_) if object_ else data.get("object", "")
419431
new_text = f"{new_subject} {new_predicate} {new_object}"
420432

@@ -453,15 +465,19 @@ def _delete_relation(self, relation_id: str) -> bool:
453465
entities = self._store._get_entities_by_ids([eid])
454466
if entities:
455467
e_data = entities[0]
456-
new_relation_ids = [rid for rid in e_data.get("relation_ids", []) if rid != relation_id]
468+
new_relation_ids = [
469+
rid for rid in e_data.get("relation_ids", []) if rid != relation_id
470+
]
457471
self._store._update_entity(eid, relation_ids=new_relation_ids)
458472

459473
# Update related passages (remove this relation from relation_ids)
460474
for pid in passage_ids:
461475
passages = self._store.get_passages_by_ids([pid])
462476
if passages:
463477
p_data = passages[0]
464-
new_relation_ids = [rid for rid in p_data.get("relation_ids", []) if rid != relation_id]
478+
new_relation_ids = [
479+
rid for rid in p_data.get("relation_ids", []) if rid != relation_id
480+
]
465481
self._store.update_passage(pid, relation_ids=new_relation_ids)
466482

467483
# Delete the relation
@@ -528,7 +544,9 @@ def create_passage(
528544
relation_ids.append(relation_id)
529545

530546
# Get entity IDs for this triplet
531-
subject_id = self._entity_name_to_id.get(self._normalize_entity_name(triplet.subject))
547+
subject_id = self._entity_name_to_id.get(
548+
self._normalize_entity_name(triplet.subject)
549+
)
532550
object_id = self._entity_name_to_id.get(self._normalize_entity_name(triplet.object))
533551

534552
if subject_id and subject_id not in entity_ids:
@@ -596,12 +614,14 @@ def search_passages(
596614
passages = []
597615
for r in results:
598616
data = r["entity"]
599-
passages.append(Passage(
600-
id=data["id"],
601-
text=data["text"],
602-
entity_ids=data.get("entity_ids", []),
603-
relation_ids=data.get("relation_ids", []),
604-
))
617+
passages.append(
618+
Passage(
619+
id=data["id"],
620+
text=data["text"],
621+
entity_ids=data.get("entity_ids", []),
622+
relation_ids=data.get("relation_ids", []),
623+
)
624+
)
605625

606626
return passages
607627

@@ -657,15 +677,19 @@ def delete_passage(self, passage_id: str) -> bool:
657677
entities = self._store._get_entities_by_ids([eid])
658678
if entities:
659679
e_data = entities[0]
660-
new_passage_ids = [pid for pid in e_data.get("passage_ids", []) if pid != passage_id]
680+
new_passage_ids = [
681+
pid for pid in e_data.get("passage_ids", []) if pid != passage_id
682+
]
661683
self._store._update_entity(eid, passage_ids=new_passage_ids)
662684

663685
# Update related relations (remove this passage from passage_ids)
664686
for rid in relation_ids:
665687
relations = self._store._get_relations_by_ids([rid])
666688
if relations:
667689
r_data = relations[0]
668-
new_passage_ids = [pid for pid in r_data.get("passage_ids", []) if pid != passage_id]
690+
new_passage_ids = [
691+
pid for pid in r_data.get("passage_ids", []) if pid != passage_id
692+
]
669693
self._store._update_relation(rid, passage_ids=new_passage_ids)
670694

671695
# Delete the passage

src/vector_graph_rag/graph/knowledge_graph.py

Lines changed: 14 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -6,8 +6,9 @@
66
"""
77

88
from __future__ import annotations
9-
from typing import List, Dict, Set, Optional, Any, TYPE_CHECKING
9+
1010
from dataclasses import dataclass, field
11+
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Set
1112

1213
if TYPE_CHECKING:
1314
from vector_graph_rag.storage.milvus import MilvusStore
@@ -401,7 +402,15 @@ def _fetch_relations(self, relation_ids: List[str]) -> None:
401402
results = self._store.client.query(
402403
collection_name=self._store.relation_collection,
403404
filter=filter_expr,
404-
output_fields=["id", "text", "entity_ids", "passage_ids", "subject", "predicate", "object"],
405+
output_fields=[
406+
"id",
407+
"text",
408+
"entity_ids",
409+
"passage_ids",
410+
"subject",
411+
"predicate",
412+
"object",
413+
],
405414
)
406415

407416
for r in results:
@@ -476,23 +485,17 @@ def passage_ids(self) -> Set[str]:
476485
@property
477486
def entities(self) -> List[GraphEntity]:
478487
"""Get all entity objects in this subgraph."""
479-
return [
480-
self._entities[eid] for eid in self._entity_ids if eid in self._entities
481-
]
488+
return [self._entities[eid] for eid in self._entity_ids if eid in self._entities]
482489

483490
@property
484491
def relations(self) -> List[GraphRelation]:
485492
"""Get all relation objects in this subgraph."""
486-
return [
487-
self._relations[rid] for rid in self._relation_ids if rid in self._relations
488-
]
493+
return [self._relations[rid] for rid in self._relation_ids if rid in self._relations]
489494

490495
@property
491496
def passages(self) -> List[GraphPassage]:
492497
"""Get all passage objects in this subgraph."""
493-
return [
494-
self._passages[pid] for pid in self._passage_ids if pid in self._passages
495-
]
498+
return [self._passages[pid] for pid in self._passage_ids if pid in self._passages]
496499

497500
@property
498501
def entity_names(self) -> List[str]:

src/vector_graph_rag/graph/retriever.py

Lines changed: 7 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -3,17 +3,17 @@
33
"""
44

55
import logging
6-
from typing import List, Dict, Any, Optional, Tuple
76
from dataclasses import dataclass, field
8-
9-
logger = logging.getLogger(__name__)
7+
from typing import List, Optional, Tuple
108

119
from vector_graph_rag.config import Settings, get_settings
12-
from vector_graph_rag.storage.embeddings import EmbeddingModel
13-
from vector_graph_rag.storage.milvus import MilvusStore
1410
from vector_graph_rag.graph.builder import GraphBuilder
1511
from vector_graph_rag.graph.knowledge_graph import SubGraph
1612
from vector_graph_rag.llm.extractor import EntityExtractor
13+
from vector_graph_rag.storage.embeddings import EmbeddingModel
14+
from vector_graph_rag.storage.milvus import MilvusStore
15+
16+
logger = logging.getLogger(__name__)
1717

1818

1919
@dataclass
@@ -93,9 +93,7 @@ def __init__(
9393
self.graph_builder = graph_builder
9494

9595
self.embedding_model = embedding_model or EmbeddingModel(settings=self.settings)
96-
self.entity_extractor = entity_extractor or EntityExtractor(
97-
settings=self.settings
98-
)
96+
self.entity_extractor = entity_extractor or EntityExtractor(settings=self.settings)
9997

10098
def _extract_query_entities(self, query: str) -> List[str]:
10199
"""Extract named entities from the query."""
@@ -338,9 +336,7 @@ def retrieve(
338336
)
339337

340338
# Expand subgraph
341-
subgraph = self._expand_subgraph(
342-
entity_ids, relation_ids, degree=expansion_degree
343-
)
339+
subgraph = self._expand_subgraph(entity_ids, relation_ids, degree=expansion_degree)
344340

345341
# Apply eviction strategy if needed
346342
threshold = relation_number_threshold or self.settings.relation_number_threshold

0 commit comments

Comments
 (0)