Files

781 lines
21 KiB
Python
Raw Permalink Normal View History

"""Phase 6 FastAPI application: Graph API + GraphQL + RAG Pipeline.
Features:
- REST API for graph operations (entity resolution, subgraph, patterns, analytics)
- GraphQL endpoint for flexible queries
- RAG pipeline integrating with LLM
"""
from fastapi import FastAPI, APIRouter, HTTPException, Query, Request
from fastapi.responses import JSONResponse
from typing import Optional, List, Dict, Any
import time
import asyncio
import logging
import json
from ont_platform.core.graph.neo4j_adapter import Neo4jAdapter, Neo4jConfig
from ont_platform.core.graph.entity_resolver import EntityResolver
from ont_platform.core.graph.subgraph_retriever import SubgraphRetriever
from ont_platform.core.graph.pattern_matcher import PatternMatcher
from ont_platform.core.graph.graph_analytics import GraphAnalytics
logger = logging.getLogger(__name__)
app = FastAPI(
title="Ontology Platform - Phase 6 GraphRAG",
description="Graph API + GraphQL + RAG Pipeline",
version="0.6.0",
)
# Routers
graph_router = APIRouter(prefix="/api/v1/graph", tags=["graph"])
rag_router = APIRouter(prefix="/api/v1/rag", tags=["rag"])
# Global instances
_neo4j_adapter: Optional[Neo4jAdapter] = None
_entity_resolver: Optional[EntityResolver] = None
_subgraph_retriever: Optional[SubgraphRetriever] = None
_pattern_matcher: Optional[PatternMatcher] = None
_graph_analytics: Optional[GraphAnalytics] = None
async def get_neo4j_adapter() -> Neo4jAdapter:
"""Get or create Neo4j adapter instance."""
global _neo4j_adapter
if _neo4j_adapter is None:
config = Neo4jConfig(
uri="bolt://localhost:7687",
username="neo4j",
password="ontology123",
)
_neo4j_adapter = Neo4jAdapter(config)
if not await _neo4j_adapter.connect():
logger.warning("Neo4j not available")
else:
try:
await _neo4j_adapter.initialize_embedder()
except Exception as e:
logger.warning(f"Failed to initialize embedder: {e}")
return _neo4j_adapter
async def get_components():
"""Initialize all graph components."""
global _entity_resolver, _subgraph_retriever, _pattern_matcher, _graph_analytics
adapter = await get_neo4j_adapter()
if _entity_resolver is None:
_entity_resolver = EntityResolver()
await _entity_resolver.initialize_embedder()
if _subgraph_retriever is None:
_subgraph_retriever = SubgraphRetriever(adapter)
if _pattern_matcher is None:
_pattern_matcher = PatternMatcher(adapter)
if _graph_analytics is None:
_graph_analytics = GraphAnalytics(adapter)
return {
"adapter": adapter,
"resolver": _entity_resolver,
"retriever": _subgraph_retriever,
"matcher": _pattern_matcher,
"analytics": _graph_analytics,
}
# ============================================================================
# Entity Resolution Endpoints
# ============================================================================
@graph_router.post("/resolve")
async def resolve_entities(
entities: List[Dict[str, Any]],
vector_threshold: float = Query(0.85),
text_threshold: float = Query(0.88),
):
"""
Detect and resolve duplicate entities.
Request:
```json
{
"entities": [
{"id": 1, "label": "Apple Inc.", "type": "Company"},
{"id": 2, "label": "Apple Inc", "type": "Company"}
]
}
```
Response:
```json
{
"clusters": [
{
"cluster_id": "C_1_2",
"canonical_id": 1,
"duplicates": [2],
"confidence": 0.92,
"reason": "combined"
}
]
}
```
"""
try:
components = await get_components()
resolver = components["resolver"]
resolver.vector_threshold = vector_threshold
resolver.text_threshold = text_threshold
clusters = await resolver.detect_duplicates(entities)
return {
"status": "success",
"clusters": [
{
"cluster_id": c.cluster_id,
"canonical_id": c.canonical_id,
"duplicates": c.duplicates,
"confidence": c.confidence,
"reason": c.reason,
}
for c in clusters
],
"total_clusters": len(clusters),
}
except Exception as e:
logger.error(f"Entity resolution failed: {e}")
raise HTTPException(status_code=500, detail=str(e))
# ============================================================================
# Subgraph Retrieval Endpoints
# ============================================================================
@graph_router.get("/subgraph/neighborhood/{entity_id}")
async def get_neighborhood(
entity_id: int,
hops: int = Query(2, ge=1, le=3),
limit: int = Query(500),
min_confidence: float = Query(0.0),
):
"""
Extract N-hop neighborhood around an entity.
Returns:
```json
{
"center_entity": {...},
"nodes": [{id, label, type, confidence}, ...],
"edges": [{source_id, target_id, predicate, confidence}, ...],
"node_count": 125,
"edge_count": 287
}
```
"""
try:
components = await get_components()
retriever = components["retriever"]
result = await retriever.retrieve_neighborhood(
entity_id=entity_id,
hops=hops,
limit=limit,
min_confidence=min_confidence,
)
return {
"status": "success",
"data": result,
}
except Exception as e:
logger.error(f"Subgraph retrieval failed: {e}")
raise HTTPException(status_code=500, detail=str(e))
@graph_router.post("/subgraph/context")
async def get_context(
entity_ids: List[int],
context_hops: int = Query(2, ge=1, le=3),
):
"""
Find common context between multiple entities.
Request:
```json
{
"entity_ids": [1, 2, 3]
}
```
"""
try:
components = await get_components()
retriever = components["retriever"]
result = await retriever.retrieve_context(
entity_ids=entity_ids,
context_hops=context_hops,
)
return {
"status": "success",
"data": result,
}
except Exception as e:
logger.error(f"Context retrieval failed: {e}")
raise HTTPException(status_code=500, detail=str(e))
# ============================================================================
# Pattern Matching Endpoints
# ============================================================================
@graph_router.post("/patterns/paths")
async def find_paths(
start_id: int,
end_id: int,
max_length: int = Query(5, ge=2, le=6),
):
"""
Find all paths between two entities.
Returns:
```json
{
"paths": [
{"path": [1, 2, 3, 5], "length": 3, "confidence": 0.87},
{"path": [1, 4, 5], "length": 2, "confidence": 0.91}
]
}
```
"""
try:
components = await get_components()
matcher = components["matcher"]
paths = await matcher.find_paths(
start_entity_id=start_id,
end_entity_id=end_id,
max_length=max_length,
)
return {
"status": "success",
"paths": paths,
"total_paths": len(paths),
}
except Exception as e:
logger.error(f"Path finding failed: {e}")
raise HTTPException(status_code=500, detail=str(e))
@graph_router.post("/patterns/cycles")
async def find_cycles(
min_length: int = Query(2, ge=2),
max_length: int = Query(5, ge=2, le=6),
):
"""
Detect cycles in the knowledge graph.
"""
try:
components = await get_components()
matcher = components["matcher"]
cycles = await matcher.find_cycles(
min_length=min_length,
max_length=max_length,
)
return {
"status": "success",
"cycles": cycles,
"total_cycles": len(cycles),
}
except Exception as e:
logger.error(f"Cycle detection failed: {e}")
raise HTTPException(status_code=500, detail=str(e))
@graph_router.post("/patterns/motifs")
async def find_motifs(
motif_type: str = Query("triangle"),
limit: int = Query(100),
):
"""
Detect graph motifs (triangle, chain, star).
"""
try:
components = await get_components()
matcher = components["matcher"]
motifs = await matcher.find_motifs(
motif_type=motif_type,
limit=limit,
)
return {
"status": "success",
"motif_type": motif_type,
"motifs": motifs,
"total_motifs": len(motifs),
}
except Exception as e:
logger.error(f"Motif detection failed: {e}")
raise HTTPException(status_code=500, detail=str(e))
# ============================================================================
# Graph Analytics Endpoints
# ============================================================================
@graph_router.post("/analytics/centrality")
async def calculate_centrality(
centrality_type: str = Query("pagerank"),
top_n: int = Query(100),
):
"""
Calculate entity centrality metrics.
Types: degree, pagerank, betweenness, closeness
"""
try:
components = await get_components()
analytics = components["analytics"]
entities = await analytics.calculate_centrality(
centrality_type=centrality_type,
top_n=top_n,
)
return {
"status": "success",
"centrality_type": centrality_type,
"entities": entities,
"total_entities": len(entities),
}
except Exception as e:
logger.error(f"Centrality calculation failed: {e}")
raise HTTPException(status_code=500, detail=str(e))
@graph_router.post("/analytics/communities")
async def detect_communities(
algorithm: str = Query("louvain"),
min_size: int = Query(2),
):
"""
Detect communities in the graph.
Algorithms: louvain, label_propagation
"""
try:
components = await get_components()
analytics = components["analytics"]
communities = await analytics.detect_communities(
algorithm=algorithm,
)
filtered = [c for c in communities if c["size"] >= min_size]
return {
"status": "success",
"algorithm": algorithm,
"communities": filtered,
"total_communities": len(filtered),
}
except Exception as e:
logger.error(f"Community detection failed: {e}")
raise HTTPException(status_code=500, detail=str(e))
@graph_router.get("/analytics/statistics")
async def get_graph_statistics():
"""
Get overall graph statistics.
Returns:
```json
{
"total_nodes": 1000,
"total_edges": 5000,
"density": 0.01,
"diameter": 7,
"is_connected": true
}
```
"""
try:
components = await get_components()
analytics = components["analytics"]
stats = await analytics.get_graph_statistics()
return {
"status": "success",
"statistics": stats,
}
except Exception as e:
logger.error(f"Statistics calculation failed: {e}")
raise HTTPException(status_code=500, detail=str(e))
@graph_router.get("/analytics/influential")
async def get_influential_entities(
top_n: int = Query(20),
):
"""
Get most influential entities (composite score).
"""
try:
components = await get_components()
analytics = components["analytics"]
entities = await analytics.find_influential_entities(top_n=top_n)
return {
"status": "success",
"entities": entities,
"total_entities": len(entities),
}
except Exception as e:
logger.error(f"Influential entity detection failed: {e}")
raise HTTPException(status_code=500, detail=str(e))
# ============================================================================
# RAG Pipeline Endpoints
# ============================================================================
@rag_router.post("/context-extraction")
async def extract_rag_context(
query_text: str,
entity_id: Optional[int] = None,
hops: int = Query(2, ge=1, le=3),
max_entities: int = Query(100),
):
"""
Extract RAG context from knowledge graph.
If entity_id provided: use neighborhood
If query_text provided: search and extract context
"""
try:
components = await get_components()
retriever = components["retriever"]
analytics = components["analytics"]
if entity_id:
# Extract from known entity
context = await retriever.retrieve_neighborhood(
entity_id=entity_id,
hops=hops,
limit=max_entities,
)
else:
# Search for query in entities (simple text match)
adapter = components["adapter"]
results = await adapter.vector_search(query_text, limit=5)
if not results:
return {
"status": "no_results",
"message": f"No entities found for: {query_text}",
"context": None,
}
# Use top result for context
top_entity = results[0]
context = await retriever.retrieve_neighborhood(
entity_id=top_entity["id"],
hops=hops,
limit=max_entities,
)
return {
"status": "success",
"query": query_text or f"entity_{entity_id}",
"context": context,
"context_size": len(context.get("nodes", [])),
}
except Exception as e:
logger.error(f"Context extraction failed: {e}")
raise HTTPException(status_code=500, detail=str(e))
@rag_router.post("/query")
async def rag_query(
query: str,
context_hops: int = Query(2),
use_graph_context: bool = Query(True),
):
"""
Process a RAG query with graph context.
Returns:
```json
{
"query": "What is Apple?",
"context": {...},
"llm_prompt": "...",
"ready_for_llm": true
}
```
Note: For LLM inference, send the llm_prompt to your LLM service.
"""
try:
components = await get_components()
retriever = components["retriever"]
adapter = components["adapter"]
# Step 1: Search for relevant entities
search_results = await adapter.vector_search(query, limit=3)
if not search_results:
return {
"status": "no_results",
"message": "No relevant entities found",
"query": query,
}
# Step 2: Extract context from top results
context_data = []
for result in search_results:
context = await retriever.retrieve_neighborhood(
entity_id=result["id"],
hops=context_hops,
limit=50,
)
context_data.append(
{
"entity": result,
"subgraph": context,
}
)
# Step 3: Build LLM prompt
llm_prompt = _build_rag_prompt(query, context_data)
return {
"status": "success",
"query": query,
"relevant_entities": [r["label"] for r in search_results],
"context_nodes": sum(
len(c["subgraph"].get("nodes", [])) for c in context_data
),
"llm_prompt": llm_prompt,
"ready_for_llm": True,
"context": context_data if use_graph_context else None,
}
except Exception as e:
logger.error(f"RAG query failed: {e}")
raise HTTPException(status_code=500, detail=str(e))
def _build_rag_prompt(query: str, context_data: List[Dict]) -> str:
"""
Build a structured prompt for LLM with graph context.
"""
prompt = f"""You are a helpful assistant with access to a knowledge graph.
KNOWLEDGE GRAPH CONTEXT:
"""
for i, ctx in enumerate(context_data, 1):
entity = ctx["entity"]
subgraph = ctx["subgraph"]
prompt += f"\n--- Source Entity {i}: {entity['label']} ---\n"
prompt += f"Type: {entity['type']}\n"
prompt += f"Confidence: {entity['similarity']:.3f}\n"
if subgraph.get("nodes"):
prompt += f"\nRelated Entities ({len(subgraph['nodes'])} total):\n"
for node in subgraph["nodes"][:10]: # Show top 10
prompt += f" - {node['label']} (type: {node['type']})\n"
if subgraph.get("edges"):
prompt += f"\nRelationships ({len(subgraph['edges'])} total):\n"
for edge in subgraph["edges"][:5]: # Show top 5
prompt += (
f" - {edge['source_id']} --{edge['predicate']}--> "
f"{edge['target_id']} (confidence: {edge['confidence']:.2f})\n"
)
prompt += f"\nUSER QUERY: {query}\n\n"
prompt += "Based on the knowledge graph context above, please answer the user's query comprehensively.\n"
prompt += "If information is found in the graph, cite it. If not found, say so clearly.\n"
return prompt
# ============================================================================
# GraphQL Endpoint (Simple Implementation)
# ============================================================================
@app.post("/graphql")
async def graphql_query(request: Request):
"""
Simple GraphQL endpoint for flexible graph queries.
Example query:
```graphql
{
entity(id: 1) {
id
label
type
neighbors(hops: 2) {
id
label
distance
}
}
}
```
"""
try:
body = await request.json()
query = body.get("query", "")
variables = body.get("variables", {})
# Simple GraphQL parser (in production, use graphene or similar)
result = await _process_graphql(query, variables)
return {
"data": result,
}
except Exception as e:
logger.error(f"GraphQL query failed: {e}")
return {
"errors": [{"message": str(e)}],
}
async def _process_graphql(query: str, variables: Dict) -> Dict:
"""
Process GraphQL query (simplified implementation).
Supports:
- entity(id): Get entity with neighbors
- entities: List all entities
- communities: List detected communities
"""
components = await get_components()
# Simple parsing (in production, use proper GraphQL parser)
if "entity(" in query:
# Extract entity ID from query
import re
match = re.search(r"entity\(id:\s*(\d+)", query)
if match:
entity_id = int(match.group(1))
retriever = components["retriever"]
context = await retriever.retrieve_neighborhood(entity_id=entity_id)
return {
"entity": {
"id": entity_id,
"data": context,
}
}
elif "communities" in query:
analytics = components["analytics"]
communities = await analytics.detect_communities()
return {"communities": communities}
return {"error": "Query not supported"}
# ============================================================================
# Health Check & Info Endpoints
# ============================================================================
@app.get("/health")
async def health_check():
"""Health check endpoint."""
try:
adapter = await get_neo4j_adapter()
neo4j_status = "connected" if adapter._driver else "disconnected"
except Exception as e:
neo4j_status = f"error: {str(e)}"
return {
"status": "ok",
"version": "0.6.0",
"neo4j": neo4j_status,
}
@app.get("/info")
async def info():
"""API information."""
return {
"name": "Ontology Platform - Phase 6",
"version": "0.6.0",
"phase": 6,
"features": [
"REST API for graph operations",
"GraphQL endpoint",
"RAG pipeline integration",
"Entity resolution",
"Subgraph retrieval",
"Pattern matching",
"Graph analytics",
],
"endpoints": {
"graph": "/api/v1/graph",
"rag": "/api/v1/rag",
"graphql": "/graphql",
},
}
# Register routers
app.include_router(graph_router)
app.include_router(rag_router)
@app.on_event("shutdown")
async def shutdown_event():
"""Cleanup on shutdown."""
global _neo4j_adapter
if _neo4j_adapter:
await _neo4j_adapter.close()
logger.info("Neo4j connection closed")
if __name__ == "__main__":
import uvicorn
uvicorn.run(app, host="0.0.0.0", port=8000, reload=True)