drewgenai/protocol-sync
0
1import os2import shutil3import json4import pandas as pd5import chainlit as cl6from dotenv import load_dotenv7from langchain_core.documents import Document8from langchain_community.document_loaders import PyMuPDFLoader9from langchain_experimental.text_splitter import SemanticChunker10from langchain_community.vectorstores import Qdrant11from langchain_huggingface import HuggingFaceEmbeddings12from langchain_core.output_parsers import StrOutputParser13from langchain_openai import ChatOpenAI14from langchain_core.prompts import ChatPromptTemplate15from langchain.tools import tool16from langchain.schema import HumanMessage17from typing_extensions import List, TypedDict18from operator import itemgetter19from langchain.agents import AgentExecutor, create_openai_tools_agent20from langchain_core.prompts import MessagesPlaceholder21from qdrant_client import QdrantClient22from qdrant_client.models import VectorParams, Distance23 24load_dotenv()25 26 27UPLOAD_PATH = "upload/"28OUTPUT_PATH = "output/"29INITIAL_DATA_PATH = "./data/Instruments_Definitions.xlsx"30os.makedirs(UPLOAD_PATH, exist_ok=True)31os.makedirs(OUTPUT_PATH, exist_ok=True)32 33# Initialize embeddings model34model_id = "Snowflake/snowflake-arctic-embed-m"35embedding_model = HuggingFaceEmbeddings(model_name=model_id)36semantic_splitter = SemanticChunker(embedding_model, add_start_index=True, buffer_size=30)37llm = ChatOpenAI(model="gpt-4o-mini")38 39# Export comparison prompt40export_prompt = """41CONTEXT:42{context}43 44QUERY:45{question}46 47You are a helpful assistant. Use the available context to answer the question.48 49Between these two files containing protocols, identify and match **entire assessment sections** based on conceptual similarity. Do NOT match individual questions.50 51### **Output Format:**52Return the response in **valid JSON format** structured as a list of dictionaries, where each dictionary contains:53[54 {{55 "Derived Description": "A short name for the matched concept",56 "Protocol_1": "Protocol 1 - Matching Element",57 "Protocol_2": "Protocol 2 - Matching Element"58 }},59 ...60]61### **Example Output:**62[63 {{64 "Derived Description": "Pain Coping Strategies",65 "Protocol_1": "Pain Coping Strategy Scale (PCSS-9)",66 "Protocol_2": "Chronic Pain Adjustment Index (CPAI-10)"67 }},68 {{69 "Derived Description": "Work Stress and Fatigue",70 "Protocol_1": "Work-Related Stress Scale (WRSS-8)",71 "Protocol_2": "Occupational Fatigue Index (OFI-7)"72 }},73 ...74]75 76### Rules:771. Only output **valid JSON** with no explanations, summaries, or markdown formatting.782. Ensure each entry in the JSON list represents a single matched data element from the two protocols.793. If no matching element is found in a protocol, leave it empty ("").804. **Do NOT include headers, explanations, or additional formatting**—only return the raw JSON list.815. It should include all the elements in the two protocols.826. If it cannot match the element, create the row and include the protocol it did find and put "could not match" in the other protocol column.837. protocol should be the between84"""85 86compare_export_prompt = ChatPromptTemplate.from_template(export_prompt)87 88QUERY_PROMPT = """89You are a helpful assistant. Use the available context to answer the question concisely and informatively.90 91CONTEXT:92{context}93 94QUERY:95{question}96 97Provide a natural-language response using the given information. If you do not know the answer, say so.98"""99 100query_prompt = ChatPromptTemplate.from_template(QUERY_PROMPT)101 102 103@tool104def document_query_tool(question: str) -> str:105 """Retrieves relevant document sections and answers questions based on the uploaded documents."""106 107 retriever = cl.user_session.get("qdrant_retriever")108 if not retriever:109 return "Error: No documents available for retrieval. Please upload two PDF files first."110 retriever = retriever.with_config({"k": 10})111 112 # Use a RAG chain similar to the comparison tool113 rag_chain = (114 {"context": itemgetter("question") | retriever, "question": itemgetter("question")}115 | query_prompt | llm | StrOutputParser()116 )117 response_text = rag_chain.invoke({"question": question})118 119 # Get the retrieved docs for context120 retrieved_docs = retriever.invoke(question)121 122 return {123 "messages": [HumanMessage(content=response_text)],124 "context": retrieved_docs125 }126 127 128@tool129def document_comparison_tool(question: str) -> str:130 """Compares the two uploaded documents, identifies matched elements, exports them as JSON, formats into CSV, and provides a download link."""131 132 # Retrieve the vector database retriever133 retriever = cl.user_session.get("qdrant_retriever")134 if not retriever:135 return "Error: No documents available for retrieval. Please upload two PDF files first."136 137 # Process query using RAG138 rag_chain = (139 {"context": itemgetter("question") | retriever, "question": itemgetter("question")}140 | compare_export_prompt | llm | StrOutputParser()141 )142 response_text = rag_chain.invoke({"question": question})143 144 # Parse response and save as CSV145 try:146 structured_data = json.loads(response_text)147 if not structured_data:148 return "Error: No matched elements found."149 150 # Define output file path151 file_path = os.path.join(OUTPUT_PATH, "comparison_results.csv")152 153 # Save to CSV154 df = pd.DataFrame(structured_data, columns=["Derived Description", "Protocol_1", "Protocol_2"])155 df.to_csv(file_path, index=False)156 157 # Send the message with the file directly from the tool158 cl.run_sync(159 cl.Message(160 content="Comparison complete! Download the CSV below:",161 elements=[cl.File(name="comparison_results.csv", path=file_path, display="inline")],162 ).send()163 )164 165 # Return a simple confirmation message166 return "Comparison results have been generated and displayed."167 168 except json.JSONDecodeError:169 return "Error: Response is not valid JSON."170 171 172# Define tools for the agent173tools = [document_query_tool, document_comparison_tool]174 175# Set up the agent with a system prompt176system_prompt = """You are an intelligent document analysis assistant. You have access to two tools:177 1781. document_query_tool: Use this when a user wants information or has questions about the content of uploaded documents.1792. document_comparison_tool: Use this when a user wants to compare elements between two uploaded documents or export comparison results.180 181Analyze the user's request carefully to determine which tool is most appropriate.182"""183 184# Create the agent using OpenAI function calling185agent_prompt = ChatPromptTemplate.from_messages([186 ("system", system_prompt),187 MessagesPlaceholder(variable_name="chat_history"),188 ("human", "{input}"),189 MessagesPlaceholder(variable_name="agent_scratchpad"),190])191 192agent = create_openai_tools_agent(193 llm=ChatOpenAI(model="gpt-4o", temperature=0),194 tools=tools,195 prompt=agent_prompt196)197 198# Create the agent executor199agent_executor = AgentExecutor.from_agent_and_tools(200 agent=agent,201 tools=tools,202 verbose=True,203 handle_parsing_errors=True,204)205 206 207def initialize_vector_store():208 """Initialize an empty Qdrant vector store"""209 try:210 # Create a Qdrant client for in-memory storage211 client = QdrantClient(location=":memory:")212 213 # Create the collection with the appropriate vector size214 # Snowflake/snowflake-arctic-embed-m produces 768-dimensional vectors215 vector_size = 768 # Changed from 1536 to match your embedding model216 217 # Check if collection exists, if not create it218 collections = client.get_collections().collections219 collection_names = [collection.name for collection in collections]220 221 if "document_comparison" not in collection_names:222 client.create_collection(223 collection_name="document_comparison",224 vectors_config=VectorParams(size=vector_size, distance=Distance.COSINE)225 )226 print("Created new collection: document_comparison")227 228 # Create the vector store with the client229 vectorstore = Qdrant(230 client=client,231 collection_name="document_comparison",232 embeddings=embedding_model233 )234 print("Vector store initialized successfully")235 return vectorstore236 except Exception as e:237 print(f"Error initializing vector store: {str(e)}")238 return None239 240 241async def load_reference_data(vectorstore):242 """Load reference Excel data into the vector database"""243 if not os.path.exists(INITIAL_DATA_PATH):244 print(f"Warning: Initial data file {INITIAL_DATA_PATH} not found")245 return vectorstore246 247 try:248 # Load Excel file249 df = pd.read_excel(INITIAL_DATA_PATH)250 251 # Convert DataFrame to documents252 documents = []253 for _, row in df.iterrows():254 # Combine all columns into a single text255 content = " ".join([f"{col}: {str(val)}" for col, val in row.items()])256 doc = Document(page_content=content, metadata={"source": "Instruments_Definitions.xlsx"})257 documents.append(doc)258 259 # Add documents to vector store260 if documents:261 vectorstore.add_documents(documents)262 print(f"Successfully loaded {len(documents)} entries from {INITIAL_DATA_PATH}")263 264 return vectorstore265 except Exception as e:266 print(f"Error loading reference data: {str(e)}")267 return vectorstore268 269 270async def process_uploaded_files(files, vectorstore):271 """Process uploaded PDF files and add them to the vector store"""272 documents_with_metadata = []273 for file in files:274 file_path = os.path.join(UPLOAD_PATH, file.name)275 shutil.copyfile(file.path, file_path)276 277 loader = PyMuPDFLoader(file_path)278 documents = loader.load()279 280 for doc in documents:281 source_name = file.name282 chunks = semantic_splitter.split_text(doc.page_content)283 for chunk in chunks:284 doc_chunk = Document(page_content=chunk, metadata={"source": source_name})285 documents_with_metadata.append(doc_chunk)286 287 if documents_with_metadata:288 # Add documents to vector store289 vectorstore.add_documents(documents_with_metadata)290 print(f"Added {len(documents_with_metadata)} chunks from uploaded files")291 return True292 return False293 294 295@cl.on_chat_start296async def start():297 # Initialize chat history for the agent298 cl.user_session.set("chat_history", [])299 300 # Initialize vector store301 vectorstore = initialize_vector_store()302 if not vectorstore:303 await cl.Message("Error: Could not initialize vector store.").send()304 return305 306 # Load reference data307 with cl.Step("Loading reference data"):308 vectorstore = await load_reference_data(vectorstore)309 cl.user_session.set("qdrant_vectorstore", vectorstore)310 cl.user_session.set("qdrant_retriever", vectorstore.as_retriever())311 await cl.Message("Reference data loaded successfully!").send()312 313 # Ask for PDF uploads314 files = await cl.AskFileMessage(315 content="Please upload **two PDF files** for comparison:",316 accept=["application/pdf"],317 max_files=2318 ).send()319 320 if len(files) != 2:321 await cl.Message("Error: You must upload exactly two PDF files.").send()322 return323 324 # Process uploaded files325 with cl.Step("Processing uploaded files"):326 success = await process_uploaded_files(files, vectorstore)327 if success:328 # Update the retriever with the latest vector store329 cl.user_session.set("qdrant_retriever", vectorstore.as_retriever())330 await cl.Message("Files uploaded and processed successfully! You can now enter your query.").send()331 else:332 await cl.Message("Error: Unable to process files. Please try again.").send()333 334 335@cl.on_message336async def handle_message(message: cl.Message):337 # Get chat history338 chat_history = cl.user_session.get("chat_history", [])339 340 # Run the agent341 with cl.Step("Agent thinking"):342 response = await cl.make_async(agent_executor.invoke)(343 {"input": message.content, "chat_history": chat_history}344 )345 346 # Handle the response based on the tool that was called347 if isinstance(response["output"], dict) and "messages" in response["output"]:348 # This is from document_query_tool349 await cl.Message(response["output"]["messages"][0].content).send()350 else:351 # Generic response (including the confirmation from document_comparison_tool)352 await cl.Message(content=str(response["output"])).send()353 354 # Update chat history with the new exchange355 chat_history.extend([356 HumanMessage(content=message.content),357 HumanMessage(content=str(response["output"]))358 ])359 cl.user_session.set("chat_history", chat_history)