File size: 5,900 Bytes
fd49a93 ae11f2d 309744b ae11f2d fd49a93 ae11f2d fd49a93 74c0f7a fd49a93 9c3c25d f76c26d 9c3c25d f76c26d fd49a93 ae11f2d fd49a93 ae11f2d fd49a93 9c3c25d f76c26d 9c3c25d f76c26d 9c3c25d ae11f2d 9c3c25d ae11f2d 9c3c25d fd49a93 ae11f2d fd49a93 ae11f2d f76c26d | 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 | 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", []) |