from typing import TypedDict, List from langchain_core.documents import Document from langgraph.graph import StateGraph, END from langchain_google_genai import ChatGoogleGenerativeAI from langchain_core.prompts import ChatPromptTemplate from langchain.retrievers import EnsembleRetriever from models import Question, FinalAnswer, GeneratedQueries from config import ANSWER_LLM_MODEL, QUERY_LLM_MODEL, GOOGLE_API_KEY class GraphState(TypedDict): original_questions: List[Question] decomposed_questions: GeneratedQueries retriever: EnsembleRetriever documents: List[List[Document]] answers: List[FinalAnswer] class RAGWorkflow: def __init__(self): self.generation_llm = ChatGoogleGenerativeAI(model=ANSWER_LLM_MODEL, google_api_key=GOOGLE_API_KEY, temperature=0.0) self.decomposition_llm = ChatGoogleGenerativeAI(model=QUERY_LLM_MODEL, google_api_key=GOOGLE_API_KEY, temperature=0.0) self.graph = self._build_graph() def _query_decomposition_node(self, state: GraphState): prompt = ChatPromptTemplate.from_template( """### ROLE: You are a world-class expert in information retrieval. ### GOAL: 🎯 For each user question, generate exactly 3 distinct, self-contained search queries. ### 📜 INSTRUCTIONS: - **Multi-Angle Approach:** Generate queries covering: 1. Definitional, 2. Procedural/Conditional, and 3. Quantitative/Limit. - **Use Policy Jargon:** Use precise insurance terminology. - **Be Self-Contained:** Each query must be understandable on its own. ### OUTPUT FORMAT: You MUST return a Pydantic `GeneratedQueries` object. ### USER QUESTIONS: {questions}""" ) questions_str = "\n".join(f"{i+1}. {q.question}" for i, q in enumerate(state["original_questions"])) structured_llm = self.decomposition_llm.with_structured_output(GeneratedQueries) decomposition_chain = prompt | structured_llm generated_lists: GeneratedQueries = decomposition_chain.invoke({"questions": questions_str}) for i, el in enumerate(generated_lists.lst): el.queries.append(state["original_questions"][i].question) return {"decomposed_questions": generated_lists} def _retrieval_node(self, state: GraphState): all_queries = [q for query_list in state["decomposed_questions"].lst for q in query_list.queries] retrieved_docs_lists = state["retriever"].batch(all_queries) queries_per_question = len(state["decomposed_questions"].lst[0].queries) final_documents: List[List[Document]] = [] for i in range(len(state["original_questions"])): start_index, end_index = i * queries_per_question, (i + 1) * queries_per_question single_question_docs = [doc for docs_list in retrieved_docs_lists[start_index:end_index] for doc in docs_list] unique_docs = {doc.page_content: doc for doc in single_question_docs} final_documents.append(list(unique_docs.values())) return {"documents": final_documents} def _generation_node(self, state: GraphState): prompt_template = """### ROLE: You are a meticulous, AI-powered Insurance Claims Adjudicator. ### 📜 CRITICAL MANDATES: 1. **Context is Absolute:** Your answer MUST be derived *only* from the provided CONTEXT. 2. **Exhaustive Factual Extraction:** Extract **all** relevant details: conditions, waiting periods, monetary limits, sub-limits, percentages, and exclusions. 3. **Handle Missing Information:** If context is insufficient, state: "The provided context does not contain sufficient information to answer this question." 4. **Be Factual:** Stick to the facts. No pleasantries. ### OUTPUT FORMAT: You MUST respond with a Pydantic `FinalAnswer` object. ### CONTEXT: {context} ### QUESTION: {question} """ final_answers = [] questions = state["original_questions"] documents = state["documents"] prompt = ChatPromptTemplate.from_template(prompt_template) structured_llm = self.generation_llm.with_structured_output(FinalAnswer) generation_chain = prompt | structured_llm for i, question in enumerate(questions): context_str = "\n\n---\n\n".join([doc.page_content for doc in documents[i]]) try: single_answer = generation_chain.invoke({"context": context_str, "question": question.question}) if single_answer is None: final_answers.append(FinalAnswer(answer="The answer could not be generated, likely due to a safety policy violation for this specific question.")) else: final_answers.append(single_answer) except Exception as e: print(f"Error processing question '{question.question}': {e}") final_answers.append(FinalAnswer(answer=f"An error occurred while processing this question.")) return {"answers": final_answers} def _build_graph(self): workflow = StateGraph(GraphState) workflow.add_node("decompose_query", self._query_decomposition_node) workflow.add_node("retrieve", self._retrieval_node) workflow.add_node("generate", self._generation_node) workflow.set_entry_point("decompose_query") workflow.add_edge("decompose_query", "retrieve") workflow.add_edge("retrieve", "generate") workflow.add_edge("generate", END) return workflow.compile() def invoke(self, questions: List[Question], retriever: EnsembleRetriever) -> List[FinalAnswer]: initial_state = {"original_questions": questions, "retriever": retriever} final_state = self.graph.invoke(initial_state) return final_state.get("answers", [])