diff --git a/docsgpt/graphrag/store.py b/docsgpt/graphrag/store.py index ceb0fc48..116d3611 100644 --- a/docsgpt/graphrag/store.py +++ b/docsgpt/graphrag/store.py @@ -175,19 +175,24 @@ class GraphStore: Returns: Whatever ``operation`` returns. + + Raises: + Exception: Anything ``operation`` raises that is not connection + loss, and anything the single retry raises. """ - for attempt in (1, 2): - conn = self._get_connection() - try: - return operation(conn) - except Exception as exc: - if attempt == 2 or not _is_connection_lost(exc): - raise - logging.warning( - "Graph write lost its connection (%s); reconnecting and retrying once.", - exc, - ) - self.close() + try: + return operation(self._get_connection()) + except Exception as exc: + if not _is_connection_lost(exc): + raise + logging.warning( + "Graph write lost its connection (%s); reconnecting and retrying once.", + exc, + ) + self.close() + # Second and final attempt, on a connection freshly checked out by + # ``_get_connection``. A failure here belongs to the caller. + return operation(self._get_connection()) def _register_pgvector_types(self, conn) -> None: """Register pgvector's adapters, tolerating a not-yet-created extension. diff --git a/docsgpt/retriever/graph_rag.py b/docsgpt/retriever/graph_rag.py index 5593a9b2..ef4b0ea5 100644 --- a/docsgpt/retriever/graph_rag.py +++ b/docsgpt/retriever/graph_rag.py @@ -103,7 +103,7 @@ def _personalized_pagerank( # Row-normalized transitions. An undirected edge is traversable from both # endpoints, so each node normalizes over its own incident weights. - transitions: Dict[Any, List[Any]] = {} + transitions: Dict[Any, List[tuple[Any, float]]] = {} for node in nodes: neighbors = [] total = 0.0 @@ -134,6 +134,13 @@ def _personalized_pagerank( ranks[node] += (leaked + 1.0 - alpha) * restart[node] if sum(abs(ranks[node] - previous[node]) for node in nodes) < node_count * tol: break + else: + logging.debug( + "Personalized PageRank hit its %s-iteration cap on a %s-node " + "subgraph; ranking with the last iterate.", + max_iter, + node_count, + ) return ranks