| """ |
| GraphRAG Redis Worker (Upstash) |
| Background process to run community detection (Leiden Algorithm) on PostgreSQL Graph. |
| Consumes jobs from 'graph_reindex_jobs' Redis queue. |
| """ |
|
|
| import json |
| import logging |
| import os |
| import sys |
| import time |
| from typing import Any, Dict, List, Optional |
| import uuid |
| from sqlalchemy import text |
|
|
| |
| sys.path.append(os.getcwd()) |
|
|
| from core.database import SessionLocal, get_db_session |
| from core.models import CommunityMembership, GraphCommunity, GraphEdge, GraphNode |
| from core.service_factory import ServiceFactory |
| import asyncio |
|
|
| try: |
| import networkx as nx |
| except ImportError: |
| class MockGraph: |
| def __init__(self): |
| self._nodes = {} |
| self._edges = {} |
| @property |
| def nodes(self): return self._nodes |
| def add_node(self, id, **attr): self._nodes[id] = attr |
| def add_edge(self, u, v, **attr): self._edges[(u, v)] = attr |
| def number_of_nodes(self): return len(self._nodes) |
|
|
| class nx: |
| Graph = MockGraph |
| @staticmethod |
| def connected_components(G): |
| return [list(G.nodes.keys())] |
|
|
| |
| logging.basicConfig( |
| format='%(asctime)s - %(name)s - %(levelname)s - %(message)s', |
| level=logging.INFO |
| ) |
| logger = logging.getLogger(__name__) |
|
|
| class RedisWorker: |
| def __init__(self, redis_url: str = None): |
| self.max_ram_nodes = 50000 |
| self.queue_name = "graph_reindex_jobs" |
| self.redis_client = None |
| |
| |
| redis_url = redis_url or os.getenv("UPSTASH_REDIS_URL") or os.getenv("REDIS_URL") |
| if redis_url: |
| try: |
| import redis |
| self.redis_client = redis.from_url(redis_url) |
| logger.info(f"Connected to Redis: {redis_url.split('@')[-1]}") |
| except ImportError: |
| logger.warning("redis-py not installed.") |
| except Exception as e: |
| logger.error(f"Failed to connect to Redis: {e}") |
| else: |
| logger.warning("No REDIS_URL provided. Worker will not listen to queue.") |
|
|
| def fetch_graph(self, workspace_id: str) -> nx.Graph: |
| """Load entire workspace graph into NetworkX""" |
| session = SessionLocal() |
| G = nx.Graph() |
| try: |
| logger.info(f"Fetching nodes for workspace {workspace_id}...") |
| nodes = session.query(GraphNode.id, GraphNode.name).filter_by(workspace_id=workspace_id).all() |
| for n in nodes: |
| G.add_node(n.id, name=n.name) |
| |
| logger.info(f"Fetching edges for workspace {workspace_id}...") |
| edges = session.query(GraphEdge.source_node_id, GraphEdge.target_node_id, GraphEdge.weight).filter_by(workspace_id=workspace_id).all() |
| for e in edges: |
| G.add_edge(e.source_node_id, e.target_node_id, weight=e.weight) |
| |
| return G |
| finally: |
| session.close() |
|
|
| def detect_communities(self, G: nx.Graph) -> List[List[str]]: |
| """Run Louvain/Leiden algorithm""" |
| if G.number_of_nodes() == 0: |
| return [] |
| |
| try: |
| from networkx.algorithms.community import louvain_communities |
| logger.info(f"Running Louvain on {G.number_of_nodes()} nodes...") |
| communities = louvain_communities(G, seed=42) |
| return [list(c) for c in communities] |
| except ImportError: |
| logger.warning("Louvain not available, falling back to connected components") |
| return [list(c) for c in nx.connected_components(G)] |
|
|
| async def summarize_community(self, workspace_id: str, G: nx.Graph, community_nodes: List[str]) -> Dict[str, Any]: |
| """Generate LLM summary and keywords for a community using unified LLMService""" |
| llm = ServiceFactory.get_llm_service() |
| |
| |
| nodes_list = [f"- {G.nodes[n].get('name', 'Unknown')} ({G.nodes[n].get('type', 'entity')})" for n in community_nodes[:20]] |
| nodes_str = "\n".join(nodes_list) |
| |
| prompt = f"""Summarize this knowledge graph community of related entities. |
| Entities: |
| {nodes_str} |
| |
| Respond in valid JSON only with this structure: |
| {{ |
| "summary": "Short 1-2 sentence description emphasizing the common theme", |
| "keywords": ["keyword1", "keyword2", "keyword3"] |
| }}""" |
|
|
| try: |
| |
| result = await llm.generate_response( |
| prompt=prompt, |
| tenant_id=workspace_id, |
| system_prompt="You are a GraphRAG Community Analyst. Categorize and summarize groups of entities.", |
| json_mode=True |
| ) |
| |
| data = json.loads(result) |
| return { |
| "summary": data.get("summary", f"Community of {len(community_nodes)} entities."), |
| "keywords": data.get("keywords", []) |
| } |
| except Exception as e: |
| logger.error(f"Failed to summarize community via LLMService: {e}") |
| node_names = [G.nodes[n].get("name", "Unknown") for n in community_nodes[:3]] |
| return { |
| "summary": f"Group related to {', '.join(node_names)}.", |
| "keywords": node_names |
| } |
|
|
| def save_communities(self, workspace_id: str, communities: List[List[str]], G: nx.Graph): |
| """Persist results to Postgres""" |
| session = SessionLocal() |
| try: |
| session.execute(text("DELETE FROM graph_communities WHERE workspace_id = :ws_id"), {"ws_id": workspace_id}) |
| session.commit() |
| |
| logger.info(f"Summarizing and saving {len(communities)} communities...") |
| for i, members in enumerate(communities): |
| if len(members) < 2: continue |
| |
| |
| res = asyncio.run(self.summarize_community(workspace_id, G, members)) |
| |
| comm = GraphCommunity( |
| workspace_id=workspace_id, |
| level=0, |
| summary=res["summary"], |
| keywords=res["keywords"] |
| ) |
| session.add(comm) |
| session.flush() |
| |
| for node_id in members: |
| membership = CommunityMembership( |
| community_id=comm.id, |
| node_id=node_id |
| ) |
| session.add(membership) |
| |
| session.commit() |
| logger.info(f"Saved {len(communities)} communities for workspace {workspace_id}") |
| |
| except Exception as e: |
| session.rollback() |
| logger.error(f"Failed to save communities: {e}") |
| finally: |
| session.close() |
|
|
| def process_job(self, workspace_id: str): |
| logger.info(f"WORKER: Starting job for {workspace_id}") |
| G = self.fetch_graph(workspace_id) |
| if G.number_of_nodes() > self.max_ram_nodes: |
| logger.error(f"Graph too large ({G.number_of_nodes()} nodes).") |
| return |
| |
| communities = self.detect_communities(G) |
| self.save_communities(workspace_id, communities, G) |
| logger.info("WORKER: Job Finished.") |
|
|
| def run(self): |
| """Main listening loop""" |
| if not self.redis_client: |
| logger.error("Redis not connected. Exiting.") |
| return |
|
|
| logger.info(f"Listening on queue: {self.queue_name}...") |
| |
| while True: |
| |
| |
| try: |
| |
| job = self.redis_client.brpop(self.queue_name, timeout=30) |
| |
| if job: |
| _, workspace_id_bytes = job |
| workspace_id = workspace_id_bytes.decode('utf-8') |
| self.process_job(workspace_id) |
| else: |
| logger.info("Queue empty (timeout). Idle...") |
| |
| |
| except Exception as e: |
| logger.error(f"Worker Error: {e}") |
| time.sleep(5) |
|
|
| if __name__ == "__main__": |
| |
| if len(sys.argv) > 1: |
| workspace_id = sys.argv[1] |
| worker = RedisWorker(redis_url="mock://") |
| worker.process_job(workspace_id) |
| else: |
| |
| worker = RedisWorker() |
| worker.run() |
|
|