Team Ai
Apppublic

drewgenai/protocol-sync

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
app.py359 linesDownload Raw Back to root
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)