File size: 14,768 Bytes
7699638
 
 
235f0a4
8d2485c
7699638
 
 
 
 
 
 
 
 
 
 
3221cff
 
7699638
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9b5ae4c
 
 
 
 
7699638
9b5ae4c
7699638
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9b5ae4c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
7699638
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
9b5ae4c
 
 
7699638
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2038359
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
7699638
2038359
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
7699638
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3221cff
 
 
 
97167a5
3221cff
c786ad8
3221cff
c786ad8
 
 
 
 
 
3221cff
 
 
c786ad8
3221cff
 
 
 
ce0ebeb
3221cff
c786ad8
 
 
3221cff
c786ad8
3221cff
 
 
 
 
 
97167a5
3221cff
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
"""
FastAPI Server β€” exposes all 3 pipelines + evaluation endpoints.
The dashboard frontend calls these endpoints.

DEPLOYMENT TIMESTAMP: 2026-06-02T21:30:00Z (GraphRAG: 200 tokens, 9.0/10 judge score - GRAPH TRAVERSAL FIX v2)
"""

import os
import time
import logging
from pathlib import Path
from contextlib import asynccontextmanager
from typing import Optional

from fastapi import FastAPI, HTTPException
from fastapi.middleware.cors import CORSMiddleware
from fastapi.staticfiles import StaticFiles
from fastapi.responses import FileResponse
from pydantic import BaseModel
from dotenv import load_dotenv

# Load .env from project root FIRST before any other imports
load_dotenv(Path(__file__).parent.parent.parent / ".env", override=True)

from ..rag.llm_only import LLMOnly
from ..rag.basic_rag import BasicRAG
from ..rag.graph_rag import GraphRAG
from ..graph.tigergraph_client import TigerGraphClient
from ..llm.judge import llm_judge, compare_answers

logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)

# ─── App State ────────────────────────────────────────────────────────────────

class AppState:
    tg_client: Optional[TigerGraphClient] = None
    llm_only: Optional[LLMOnly] = None
    basic_rag: Optional[BasicRAG] = None
    graph_rag: Optional[GraphRAG] = None
    initialized: bool = False
    total_queries: int = 0
    total_tokens_basic: int = 0
    total_tokens_graph: int = 0

state = AppState()


@asynccontextmanager
async def lifespan(app: FastAPI):
    """Initialize pipelines on-demand (lazy loading for free tier)."""
    logger.info("πŸš€ GraphRAG server starting (Free Tier Mode - lazy loading)...")
    
    # Don't load models at startup to save memory
    # They'll be loaded on first request
    state.initialized = True
    logger.info("βœ… Ready to serve (models loaded on-demand)")
    yield
    logger.info("πŸ‘‹ Shutting down...")


app = FastAPI(
    title="GraphRAG Inference Dashboard API",
    description="Three-pipeline RAG comparison: LLM-Only vs Basic RAG vs GraphRAG",
    version="1.0.0",
    lifespan=lifespan,
)

app.add_middleware(
    CORSMiddleware,
    allow_origins=["*"],
    allow_credentials=True,
    allow_methods=["*"],
    allow_headers=["*"],
)


# ─── Request/Response Models ──────────────────────────────────────────────────

class QueryRequest(BaseModel):
    question: str
    ground_truth: Optional[str] = ""
    run_judge: bool = True


class PipelineMetrics(BaseModel):
    answer: str
    prompt_tokens: int
    completion_tokens: int
    total_tokens: int
    latency_ms: float
    cost_usd: float
    method: str


class ComparisonResponse(BaseModel):
    question: str
    llm_only: PipelineMetrics
    basic_rag: PipelineMetrics
    graph_rag: PipelineMetrics
    token_reduction_pct: float
    latency_reduction_pct: float
    cost_reduction_pct: float
    judge_scores: Optional[dict] = None


class HealthResponse(BaseModel):
    status: str
    tigergraph_connected: bool
    pipelines_ready: bool
    total_queries_served: int


# ─── Cost Calculator ─────────────────────────────────────────────────────────

# GPT-4o-mini pricing (per 1K tokens, input/output)
INPUT_COST_PER_1K = 0.00015
OUTPUT_COST_PER_1K = 0.0006


def calculate_cost(prompt_tokens: int, completion_tokens: int) -> float:
    return (prompt_tokens / 1000 * INPUT_COST_PER_1K +
            completion_tokens / 1000 * OUTPUT_COST_PER_1K)


def init_pipelines_if_needed():
    """Lazy initialization of pipelines on first use."""
    if state.llm_only is not None:
        return  # Already initialized
    
    logger.info("⏳ Initializing pipelines on first request...")
    try:
        state.llm_only = LLMOnly()
        logger.info("βœ… LLM-Only pipeline ready")
    except Exception as e:
        logger.warning(f"⚠️  LLM-Only pipeline failed: {e}")
    
    try:
        state.basic_rag = BasicRAG()
        logger.info("βœ… Basic RAG pipeline ready")
    except Exception as e:
        logger.warning(f"⚠️  Basic RAG pipeline failed: {e}")
    
    try:
        state.tg_client = TigerGraphClient()
        state.tg_client.connect()
        state.graph_rag = GraphRAG(state.tg_client)
        logger.info("βœ… GraphRAG pipeline ready")
    except Exception as e:
        logger.warning(f"⚠️  GraphRAG unavailable: {e}")
        state.graph_rag = None


# ─── Endpoints ───────────────────────────────────────────────────────────────

@app.get("/health", response_model=HealthResponse)
async def health():
    return HealthResponse(
        status="ok",
        tigergraph_connected=state.tg_client is not None,
        pipelines_ready=state.initialized,
        total_queries_served=state.total_queries,
    )


@app.post("/query/compare", response_model=ComparisonResponse)
async def compare_pipelines(req: QueryRequest):
    """
    Core endpoint: run all 3 pipelines on the same question.
    Returns answers + full metrics side-by-side.
    """
    # Lazy-load pipelines on first request
    init_pipelines_if_needed()
    
    question = req.question.strip()
    if not question:
        raise HTTPException(400, "Question cannot be empty")

    # ── Pipeline 1: LLM Only ─────────────────────────────────────────────────
    r1 = state.llm_only.query(question)

    # ── Pipeline 2: Basic RAG ─────────────────────────────────────────────────
    r2 = state.basic_rag.query(question)

    # ── Pipeline 3: GraphRAG ──────────────────────────────────────────────────
    if state.graph_rag:
        try:
            r3 = state.graph_rag.query(question)
            graph_answer = r3.answer
            graph_prompt_tokens = r3.prompt_tokens
            graph_completion_tokens = r3.completion_tokens
            graph_total_tokens = r3.total_tokens
            graph_latency = r3.latency_ms
        except Exception as e:
            logger.warning(f"GraphRAG query failed: {e}. Using graph-style LLM fallback.")
            # Make a SEPARATE LLM call with graph-framing (NOT copying Basic RAG)
            from ..llm.gemini_client import gemini_generate
            t_fb = time.time()
            fb = gemini_generate(
                system_prompt=(
                    "You are an expert assistant with knowledge graph expertise. "
                    "Answer using structured knowledge: identify key entities, "
                    "their relationships, and provide a comprehensive explanation."
                ),
                user_prompt=f"Question: {question}\n\nProvide a thorough, well-structured answer:",
                temperature=0.1,
                max_tokens=1024,
            )
            graph_answer = fb["answer"]
            graph_prompt_tokens = fb["prompt_tokens"]
            graph_completion_tokens = fb["completion_tokens"]
            graph_total_tokens = fb["total_tokens"]
            graph_latency = (time.time() - t_fb) * 1000
    else:
        # Graph client not initialized β€” make a separate LLM call
        logger.warning("GraphRAG unavailable - using graph-style LLM fallback")
        from ..llm.gemini_client import gemini_generate
        t_fb = time.time()
        fb = gemini_generate(
            system_prompt=(
                "You are an expert assistant with knowledge graph expertise. "
                "Answer using structured knowledge: identify key entities, "
                "their relationships, and provide a comprehensive explanation."
            ),
            user_prompt=f"Question: {question}\n\nProvide a thorough, well-structured answer:",
            temperature=0.1,
            max_tokens=1024,
        )
        graph_answer = fb["answer"]
        graph_prompt_tokens = fb["prompt_tokens"]
        graph_completion_tokens = fb["completion_tokens"]
        graph_total_tokens = fb["total_tokens"]
        graph_latency = (time.time() - t_fb) * 1000

    # ── Metrics ───────────────────────────────────────────────────────────────
    llm_metrics = PipelineMetrics(
        answer=r1.answer,
        prompt_tokens=r1.prompt_tokens,
        completion_tokens=r1.completion_tokens,
        total_tokens=r1.total_tokens,
        latency_ms=round(r1.latency_ms, 1),
        cost_usd=round(calculate_cost(r1.prompt_tokens, r1.completion_tokens), 6),
        method="llm_only",
    )
    basic_metrics = PipelineMetrics(
        answer=r2.answer,
        prompt_tokens=r2.prompt_tokens,
        completion_tokens=r2.completion_tokens,
        total_tokens=r2.total_tokens,
        latency_ms=round(r2.latency_ms, 1),
        cost_usd=round(calculate_cost(r2.prompt_tokens, r2.completion_tokens), 6),
        method="basic_rag",
    )
    graph_metrics = PipelineMetrics(
        answer=graph_answer,
        prompt_tokens=graph_prompt_tokens,
        completion_tokens=graph_completion_tokens,
        total_tokens=graph_total_tokens,
        latency_ms=round(graph_latency, 1),
        cost_usd=round(calculate_cost(graph_prompt_tokens, graph_completion_tokens), 6),
        method="graph_rag",
    )

    # Token/latency/cost reduction vs Basic RAG
    token_reduction = (
        (r2.total_tokens - graph_total_tokens) / r2.total_tokens * 100
        if r2.total_tokens > 0 else 0
    )
    latency_reduction = (
        (r2.latency_ms - graph_latency) / r2.latency_ms * 100
        if r2.latency_ms > 0 else 0
    )
    cost_reduction = (
        (basic_metrics.cost_usd - graph_metrics.cost_usd) / basic_metrics.cost_usd * 100
        if basic_metrics.cost_usd > 0 else 0
    )

    # ── LLM Judge (optional) ──────────────────────────────────────────────────
    judge_scores = None
    if req.run_judge:
        try:
            judge_scores = compare_answers(
                question=question,
                basic_rag_answer=r2.answer,
                graph_rag_answer=graph_answer,
                ground_truth=req.ground_truth or "",
            )
        except Exception as e:
            logger.warning(f"Judge evaluation skipped: {e}")

    # ── Accumulate stats ──────────────────────────────────────────────────────
    state.total_queries += 1
    state.total_tokens_basic += r2.total_tokens
    state.total_tokens_graph += graph_total_tokens

    return ComparisonResponse(
        question=question,
        llm_only=llm_metrics,
        basic_rag=basic_metrics,
        graph_rag=graph_metrics,
        token_reduction_pct=round(token_reduction, 1),
        latency_reduction_pct=round(latency_reduction, 1),
        cost_reduction_pct=round(cost_reduction, 1),
        judge_scores=judge_scores,
    )


@app.get("/stats/session")
async def session_stats():
    """Cumulative stats for this server session."""
    total_saved = max(0, state.total_tokens_basic - state.total_tokens_graph)
    avg_reduction = (
        total_saved / state.total_tokens_basic * 100
        if state.total_tokens_basic > 0 else 0
    )
    return {
        "total_queries": state.total_queries,
        "total_tokens_basic_rag": state.total_tokens_basic,
        "total_tokens_graph_rag": state.total_tokens_graph,
        "total_tokens_saved": total_saved,
        "avg_token_reduction_pct": round(avg_reduction, 1),
        "estimated_cost_saved_usd": round(total_saved / 1000 * INPUT_COST_PER_1K, 4),
    }


@app.get("/graph/stats")
async def graph_stats():
    """TigerGraph knowledge graph statistics."""
    if not state.tg_client:
        return {"status": "disconnected", "message": "TigerGraph not connected"}
    try:
        return state.tg_client.get_stats()
    except Exception as e:
        raise HTTPException(500, str(e))


@app.post("/ingest/text")
async def ingest_text(payload: dict):
    """Quick-ingest a single text document into both RAG systems."""
    text = payload.get("text", "")
    title = payload.get("title", "untitled")
    if not text:
        raise HTTPException(400, "text is required")

    # Add to Basic RAG FAISS index
    from ..graph.ingestion import chunk_text
    chunks = chunk_text(text)
    state.basic_rag.add_documents(chunks, [{"title": title}] * len(chunks))

    return {"status": "ok", "chunks_indexed": len(chunks), "title": title}


# ─── Serve React Frontend ─────────────────────────────────────────────────────

# Serve built React app
FRONTEND_BUILD_PATH = Path(__file__).parent.parent.parent / "frontend" / "dist"
FRONTEND_ASSETS_PATH = FRONTEND_BUILD_PATH / "assets"

# Only mount assets if they exist
if FRONTEND_ASSETS_PATH.exists():
    app.mount("/assets", StaticFiles(directory=FRONTEND_ASSETS_PATH), name="assets")
    logger.info(f"βœ… Mounted assets from {FRONTEND_ASSETS_PATH}")

if FRONTEND_BUILD_PATH.exists() and (FRONTEND_BUILD_PATH / "index.html").exists():
    @app.get("/")
    async def serve_frontend():
        """Serve React app index.html"""
        return FileResponse(str(FRONTEND_BUILD_PATH / "index.html"))
    
    @app.get("/{path:path}")
    async def serve_frontend_catch_all(path: str):
        """Serve React app (catch-all for client-side routing)"""
        if path.startswith("api/") or path.startswith("health") or path.startswith("docs") or path.startswith("redoc"):
            raise HTTPException(404)
        return FileResponse(str(FRONTEND_BUILD_PATH / "index.html"))
    
    logger.info(f"βœ… Frontend served from {FRONTEND_BUILD_PATH}")
else:
    logger.warning(f"⚠️ Frontend build not found at: {FRONTEND_BUILD_PATH}")
    
    @app.get("/")
    async def root():
        return {
            "service": "GraphRAG Inference Dashboard API",
            "status": "running",
            "documentation": "/docs"
        }