kumardatascience/Multi-Source-RAG-AI-System-with-Query-Routing
1
1"""LlamaIndex Workflow that routes user queries to the right knowledge source.2Emits Chainlit steps so the user sees reasoning in real time."""3 4import asyncio5import chainlit as cl6from llama_index.core.workflow import (7 Workflow,8 step,9 Event,10 StartEvent,11 StopEvent,12)13 14from llm import client, MODEL_NAME, stream_response15from rag import retrieve_context_with_session16from web_search import search_web17 18 19# --- Events ---20 21class RouteDecision(Event):22 decision: str23 question: str24 history: list[dict]25 session_id: str | None = None26 27 28class ContextReady(Event):29 context: str | None30 question: str31 history: list[dict]32 decision: str33 session_id: str | None = None34 35 36class ChatStopEvent(StopEvent):37 stream: object38 decision: str39 40 41# --- Router prompt ---42 43ROUTER_PROMPT = """You are a query router for an AI assistant. Classify the user's question into ONE of these categories:44 45- "direct": General knowledge, math, writing, coding, casual chat, definitions. Doesn't need external info.46- "rag": About uploaded documents (e.g., company policies, contracts, the user's PDFs).47- "web": Needs current/real-time information (news, weather, today's events, recent prices).48- "multi": Needs BOTH documents and current web info (e.g., "compare our pricing to current market rates").49 50Reply with EXACTLY one word: direct, rag, web, or multi. No punctuation, no explanation.51 52Question: {question}53"""54 55 56# --- The Workflow ---57 58class ChatRouter(Workflow):59 60 @step61 async def route_query(self, ev: StartEvent) -> RouteDecision:62 async with cl.Step(name="๐ฆ Routing the query", type="tool") as cl_step:63 cl_step.input = ev.question64 65 prompt = ROUTER_PROMPT.format(question=ev.question)66 response = await client.aio.models.generate_content(67 model=MODEL_NAME,68 contents=prompt,69 )70 71 raw = (response.text or "").strip().lower().split()72 decision = raw[0] if raw else "direct"73 if decision not in {"direct", "rag", "web", "multi"}:74 decision = "direct"75 76 cl_step.output = f"Decision: **{decision}**"77 78 return RouteDecision(79 decision=decision,80 question=ev.question,81 history=ev.history,82 session_id=getattr(ev, "session_id", None),83 )84 85 @step86 async def gather_context(self, ev: RouteDecision) -> ContextReady:87 context_parts = []88 89 if ev.decision in ("rag", "multi"):90 async with cl.Step(name="๐ Retrieving from documents", type="retrieval") as cl_step:91 cl_step.input = ev.question92 rag_chunks = retrieve_context_with_session(93 ev.question,94 session_id=getattr(ev, "session_id", None),95 top_k=396)97 if rag_chunks:98 context_parts.append("--- FROM YOUR DOCUMENTS ---\n" + rag_chunks)99 cl_step.output = f"Found {len(rag_chunks)} characters of relevant content."100 else:101 cl_step.output = "No relevant chunks found."102 103 if ev.decision in ("web", "multi"):104 async with cl.Step(name="๐ Searching the web", type="tool") as cl_step:105 cl_step.input = ev.question106 web_results = await asyncio.to_thread(search_web, ev.question, 3)107 if web_results:108 context_parts.append("--- FROM THE WEB ---\n" + web_results)109 cl_step.output = f"Found {len(web_results)} characters of web content."110 else:111 cl_step.output = "No web results found."112 113 context = "\n\n".join(context_parts) if context_parts else None114 115 return ContextReady(116 context=context,117 question=ev.question,118 history=ev.history,119 decision=ev.decision,120 session_id=ev.session_id,121 )122 123 @step124 async def generate_answer(self, ev: ContextReady) -> ChatStopEvent:125 generator = stream_response(ev.history, context=ev.context)126 return ChatStopEvent(stream=generator, decision=ev.decision)