Team Ai
Apppublic

kumardatascience/Multi-Source-RAG-AI-System-with-Query-Routing

sourceHugging Faceupdated 3mo agoView on Hugging Face
1likes
router.py126 linesDownload Raw Back to app
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)