| 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", []) |