Team Ai
Apppublic

cryogenic22/doc_knowledge_base

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
graph_builder.py774 linesDownload Raw Back to root
1"""2LangGraph agent orchestration for document processing, content authoring, and protocol coach.3"""4 5from langgraph.graph import StateGraph, END6from typing import TypedDict, Dict, List, Any, Optional, Literal, Annotated, cast7import operator8import uuid9 10from schemas import DocumentExtractionState, ProtocolCoachState, ContentAuthoringState, TraceabilityState11from pdf_processor import PDFProcessor12from knowledge_store import KnowledgeStore13from llm_interface import LLMInterface14 15# Initialize handlers16pdf_processor = None17knowledge_store = None18llm_interface = None19 20def init_handlers(api_key=None):21    """Initialize handlers for PDF processing, knowledge store, and LLM."""22    global pdf_processor, knowledge_store, llm_interface23    24    pdf_processor = PDFProcessor()25    knowledge_store = KnowledgeStore()26    llm_interface = LLMInterface(api_key=api_key)27    28    return pdf_processor, knowledge_store, llm_interface29 30# =========================================================================31# Document Extraction Workflow Nodes32# =========================================================================33 34def parse_document(state: DocumentExtractionState) -> DocumentExtractionState:35    """Parse PDF document and extract text."""36    try:37        document_path = state["document_path"]38        39        # Process document with PDFProcessor40        result = pdf_processor.process_complete_document(document_path)41        42        if result["status"] == "error":43            return {44                **state,45                "status": "error",46                "error": f"Failed to parse document: {result.get('error', 'Unknown error')}"47            }48        49        return {50            **state,51            "document_text": result.get("full_text", ""),52            "document_metadata": result.get("metadata", {}),53            "sections": result.get("sections", {}),54            "vector_chunks": result.get("chunks", []),55            "status": "parsed"56        }57    except Exception as e:58        return {59            **state,60            "status": "error",61            "error": f"Exception in parse_document: {str(e)}"62        }63 64def extract_study_info(state: DocumentExtractionState) -> DocumentExtractionState:65    """Extract study information using LLM."""66    if state.get("status") == "error":67        return state68    69    try:70        # Use synopsis or first few sections for study info extraction71        text_for_extraction = ""72        sections = state.get("sections", {})73        74        # Check if sections is a list (section names only) or a dict (section name -> content)75        if isinstance(sections, list):76            # Just use the document text since we don't have section content77            if "document_text" in state:78                text_for_extraction = state["document_text"][:20000]  # Use first 20k chars79        else:80            # Try to find synopsis or summary section first81            for section_name in ["synopsis", "summary", "overview"]:82                if section_name.lower() in [s.lower() for s in sections.keys()]:83                    section_key = next(k for k in sections.keys() if k.lower() == section_name.lower())84                    text_for_extraction = sections[section_key]85                    break86            87            # If no synopsis found, use the beginning of the document88            if not text_for_extraction and "document_text" in state:89                text_for_extraction = state["document_text"][:20000]  # Use first 20k chars90        91        if not text_for_extraction:92            return {93                **state,94                "status": "error",95                "error": "No text available for study info extraction"96            }97        98        # Extract study info using LLM99        study_info = llm_interface.extract_study_info(text_for_extraction)100        101        if not study_info:102            return {103                **state,104                "status": "error",105                "error": "Failed to extract study information"106            }107        108        # Ensure protocol_id is in study_info109        if "protocol_id" not in study_info and "document_metadata" in state:110            study_info["protocol_id"] = state["document_metadata"].get("protocol_id")111        112        return {113            **state,114            "extracted_study": study_info,115            "status": "study_extracted"116        }117    except Exception as e:118        return {119            **state,120            "status": "error",121            "error": f"Exception in extract_study_info: {str(e)}"122        }123 124def extract_objectives_endpoints(state: DocumentExtractionState) -> DocumentExtractionState:125    """Extract objectives and endpoints using LLM."""126    if state.get("status") == "error":127        return state128    129    try:130        sections = state.get("sections", {})131        protocol_id = state.get("extracted_study", {}).get("protocol_id")132        133        if not protocol_id:134            protocol_id = state.get("document_metadata", {}).get("protocol_id")135        136        if not protocol_id:137            return {138                **state,139                "status": "error",140                "error": "No protocol ID available for extraction"141            }142        143        # Find objectives/endpoints section144        text_for_extraction = ""145        for section_name in ["objectives", "objective", "endpoint", "endpoints"]:146            for key in sections.keys():147                if section_name.lower() in key.lower():148                    text_for_extraction = sections[key]149                    break150            if text_for_extraction:151                break152        153        if not text_for_extraction:154            return {155                **state,156                "status": "warning",157                "error": "No objectives/endpoints section found"158            }159        160        # Extract objectives and endpoints161        result = llm_interface.extract_objectives_and_endpoints(text_for_extraction, protocol_id)162        163        if not result:164            return {165                **state,166                "status": "warning",167                "error": "Failed to extract objectives and endpoints"168            }169        170        return {171            **state,172            "extracted_objectives": result.get("objectives", []),173            "extracted_endpoints": result.get("endpoints", []),174            "status": "objectives_endpoints_extracted"175        }176    except Exception as e:177        return {178            **state,179            "status": "error",180            "error": f"Exception in extract_objectives_endpoints: {str(e)}"181        }182 183def extract_population_criteria(state: DocumentExtractionState) -> DocumentExtractionState:184    """Extract inclusion and exclusion criteria using LLM."""185    if state.get("status") == "error":186        return state187    188    try:189        sections = state.get("sections", {})190        protocol_id = state.get("extracted_study", {}).get("protocol_id")191        192        if not protocol_id:193            protocol_id = state.get("document_metadata", {}).get("protocol_id")194        195        # Find criteria section196        text_for_extraction = ""197        for section_name in ["eligibility", "inclusion", "exclusion", "criteria", "population"]:198            for key in sections.keys():199                if section_name.lower() in key.lower():200                    text_for_extraction = sections[key]201                    break202            if text_for_extraction:203                break204        205        if not text_for_extraction:206            return {207                **state,208                "status": "warning",209                "error": "No population criteria section found"210            }211        212        # Extract criteria213        result = llm_interface.extract_population_criteria(text_for_extraction, protocol_id)214        215        if not result:216            return {217                **state,218                "status": "warning",219                "error": "Failed to extract population criteria"220            }221        222        return {223            **state,224            "extracted_population": result,225            "status": "population_extracted"226        }227    except Exception as e:228        return {229            **state,230            "status": "error",231            "error": f"Exception in extract_population_criteria: {str(e)}"232        }233 234def extract_study_design(state: DocumentExtractionState) -> DocumentExtractionState:235    """Extract study design information using LLM."""236    if state.get("status") == "error":237        return state238    239    try:240        sections = state.get("sections", {})241        protocol_id = state.get("extracted_study", {}).get("protocol_id")242        243        if not protocol_id:244            protocol_id = state.get("document_metadata", {}).get("protocol_id")245        246        # Find study design section247        text_for_extraction = ""248        for section_name in ["study design", "design", "methodology"]:249            for key in sections.keys():250                if section_name.lower() in key.lower():251                    text_for_extraction = sections[key]252                    break253            if text_for_extraction:254                break255        256        if not text_for_extraction:257            return {258                **state,259                "status": "warning",260                "error": "No study design section found"261            }262        263        # Extract study design264        result = llm_interface.extract_study_design(text_for_extraction, protocol_id)265        266        if not result:267            return {268                **state,269                "status": "warning",270                "error": "Failed to extract study design"271            }272        273        return {274            **state,275            "extracted_design": result,276            "status": "design_extracted"277        }278    except Exception as e:279        return {280            **state,281            "status": "error",282            "error": f"Exception in extract_study_design: {str(e)}"283        }284 285def store_in_knowledge_base(state: DocumentExtractionState) -> DocumentExtractionState:286    """Store extracted information in the knowledge base."""287    try:288        # Skip if there was a critical error289        if state.get("status") == "error":290            return state291        292        # Extract data from state293        document_metadata = state.get("document_metadata", {})294        study_info = state.get("extracted_study", {})295        objectives = state.get("extracted_objectives", [])296        endpoints = state.get("extracted_endpoints", [])297        population = state.get("extracted_population", {})298        design = state.get("extracted_design", {})299        vector_chunks = state.get("vector_chunks", [])300        301        # Ensure we have a protocol ID302        protocol_id = study_info.get("protocol_id")303        if not protocol_id:304            protocol_id = document_metadata.get("protocol_id")305        306        if not protocol_id:307            return {308                **state,309                "status": "error",310                "error": "No protocol ID available for knowledge base storage"311            }312        313        # Add protocol_id to document_metadata314        document_metadata["protocol_id"] = protocol_id315        316        # Store in NoSQL DB317        doc_id = knowledge_store.store_document_metadata(document_metadata)318        319        # Store study info if available320        if study_info:321            study_id = knowledge_store.store_study_info(study_info)322        323        # Store objectives if available324        if objectives:325            knowledge_store.store_objectives(protocol_id, objectives)326        327        # Store endpoints if available328        if endpoints:329            knowledge_store.store_endpoints(protocol_id, endpoints)330        331        # Store population criteria if available332        if population and "inclusion_criteria" in population:333            inclusion = population.get("inclusion_criteria", [])334            exclusion = population.get("exclusion_criteria", [])335            336            # Add criterion_type to each criterion337            for criterion in inclusion:338                criterion["criterion_type"] = "Inclusion"339                criterion["protocol_id"] = protocol_id340            341            for criterion in exclusion:342                criterion["criterion_type"] = "Exclusion"343                criterion["protocol_id"] = protocol_id344            345            # Store all criteria346            all_criteria = inclusion + exclusion347            knowledge_store.store_population_criteria(protocol_id, all_criteria)348        349        # Store in vector store if chunks available350        if vector_chunks:351            result = knowledge_store.add_documents(vector_chunks)352            353            if result.get("status") == "error":354                return {355                    **state,356                    "status": "warning",357                    "error": f"Warning: Failed to add to vector store: {result.get('message')}"358                }359        360        return {361            **state,362            "status": "completed",363            "document_id": doc_id,364        }365    except Exception as e:366        return {367            **state,368            "status": "error",369            "error": f"Exception in store_in_knowledge_base: {str(e)}"370        }371 372# =========================================================================373# Protocol Coach Workflow Nodes374# =========================================================================375 376def retrieve_context_for_query(state: ProtocolCoachState) -> ProtocolCoachState:377    """Retrieve relevant context for a user query."""378    try:379        query = state["query"]380        381        # Query vector store for context382        relevant_docs = knowledge_store.similarity_search(383            query=query,384            k=5  # Get top 5 most relevant chunks385        )386        387        if not relevant_docs:388            return {389                **state,390                "retrieved_context": [],391                "error": "No relevant context found"392            }393        394        # Format results for easy use395        context = [396            {397                "page_content": doc.page_content,398                "metadata": doc.metadata399            }400            for doc in relevant_docs401        ]402        403        return {404            **state,405            "retrieved_context": context406        }407    except Exception as e:408        return {409            **state,410            "error": f"Exception in retrieve_context_for_query: {str(e)}"411        }412 413def answer_query(state: ProtocolCoachState) -> ProtocolCoachState:414    """Generate answer to user query using retrieved context."""415    try:416        query = state["query"]417        context = state.get("retrieved_context", [])418        chat_history = state.get("chat_history", [])419        420        if not context:421            return {422                **state,423                "response": "I don't have enough context to answer that question about the protocol. Please try asking something else or upload relevant documents."424            }425        426        # Generate response using LLM427        response = llm_interface.answer_protocol_question(428            question=query,429            context=context,430            chat_history=chat_history431        )432        433        if not response:434            return {435                **state,436                "response": "I encountered an issue while generating a response. Please try again."437            }438        439        return {440            **state,441            "response": response442        }443    except Exception as e:444        return {445            **state,446            "response": f"Error: {str(e)}",447            "error": f"Exception in answer_query: {str(e)}"448        }449 450# =========================================================================451# Content Authoring Workflow Nodes452# =========================================================================453 454def retrieve_content_examples(state: ContentAuthoringState) -> ContentAuthoringState:455    """Retrieve examples of similar content for authoring."""456    try:457        section_type = state["section_type"]458        target_protocol_id = state.get("target_protocol_id")459        460        # Create a search query based on section type461        search_query = f"{section_type} section for clinical study protocol"462        463        # Set up potential filters464        filter_dict = None465        if target_protocol_id:466            # Exclude the target protocol from examples if specified467            filter_dict = {"protocol_id": {"$ne": target_protocol_id}}468        469        # Query vector store for examples470        relevant_docs = knowledge_store.similarity_search(471            query=search_query,472            k=3,473            filter_dict=filter_dict474        )475        476        if not relevant_docs:477            return {478                **state,479                "retrieved_context": [],480                "error": "No relevant examples found"481            }482        483        # Format results for easy use484        context = [485            {486                "page_content": doc.page_content,487                "metadata": doc.metadata488            }489            for doc in relevant_docs490        ]491        492        return {493            **state,494            "retrieved_context": context495        }496    except Exception as e:497        return {498            **state,499            "error": f"Exception in retrieve_content_examples: {str(e)}"500        }501 502def generate_content(state: ContentAuthoringState) -> ContentAuthoringState:503    """Generate content for authoring."""504    try:505        section_type = state["section_type"]506        context = state.get("retrieved_context", [])507        target_protocol_id = state.get("target_protocol_id")508        style_guide = state.get("style_guide")509        510        if not context:511            return {512                **state,513                "generated_content": "I don't have enough examples to generate a good section. Please upload more documents or try a different section type.",514                "error": "No context available for generation"515            }516        517        # Generate content using LLM518        content = llm_interface.generate_content_from_knowledge(519            section_type=section_type,520            context=context,521            protocol_id=target_protocol_id,522            style_guide=style_guide523        )524        525        if not content:526            return {527                **state,528                "generated_content": "I encountered an issue while generating content. Please try again.",529                "error": "Failed to generate content"530            }531        532        return {533            **state,534            "generated_content": content535        }536    except Exception as e:537        return {538            **state,539            "generated_content": f"Error: {str(e)}",540            "error": f"Exception in generate_content: {str(e)}"541        }542 543def critique_content(state: ContentAuthoringState) -> ContentAuthoringState:544    """Critique generated content for quality and consistency."""545    # This would normally use an LLM to critique content546    # For simplicity, we're returning the content unchanged547    return state548 549# =========================================================================550# Traceability Workflow Nodes551# =========================================================================552 553def retrieve_document_entities(state: TraceabilityState) -> TraceabilityState:554    """Retrieve entities from source and target documents."""555    try:556        source_doc_id = state["source_document_id"]557        target_doc_id = state["target_document_id"]558        entity_type = state["entity_type"]559        560        # Get document metadata561        source_doc = knowledge_store.get_document_by_id(source_doc_id)562        target_doc = knowledge_store.get_document_by_id(target_doc_id)563        564        if not source_doc or not target_doc:565            return {566                **state,567                "error": "One or both documents not found"568            }569        570        # Get protocol IDs571        source_protocol_id = source_doc.get("protocol_id")572        target_protocol_id = target_doc.get("protocol_id")573        574        if not source_protocol_id or not target_protocol_id:575            return {576                **state,577                "error": "Protocol ID missing from one or both documents"578            }579        580        # Retrieve entities based on entity type581        source_entities = []582        target_entities = []583        584        if entity_type == "objectives":585            source_entities = knowledge_store.get_objectives_by_protocol_id(source_protocol_id)586            target_entities = knowledge_store.get_objectives_by_protocol_id(target_protocol_id)587        elif entity_type == "endpoints":588            source_entities = knowledge_store.get_endpoints_by_protocol_id(source_protocol_id)589            target_entities = knowledge_store.get_endpoints_by_protocol_id(target_protocol_id)590        elif entity_type == "population":591            source_entities = knowledge_store.get_population_criteria_by_protocol_id(source_protocol_id)592            target_entities = knowledge_store.get_population_criteria_by_protocol_id(target_protocol_id)593        594        if not source_entities or not target_entities:595            return {596                **state,597                "error": f"No {entity_type} found in one or both documents"598            }599        600        return {601            **state,602            "source_entities": source_entities,603            "target_entities": target_entities604        }605    except Exception as e:606        return {607            **state,608            "error": f"Exception in retrieve_document_entities: {str(e)}"609        }610 611def match_entities(state: TraceabilityState) -> TraceabilityState:612    """Match entities between documents based on similarity."""613    try:614        if "error" in state:615            return state616        617        source_entities = state.get("source_entities", [])618        target_entities = state.get("target_entities", [])619        620        # Simple matching - in a real system this would use more sophisticated comparison621        matched_pairs = []622        623        for source_entity in source_entities:624            matches = []625            626            for target_entity in target_entities:627                # Compare based on description/text628                source_text = source_entity.get("description", source_entity.get("text", ""))629                target_text = target_entity.get("description", target_entity.get("text", ""))630                631                if not source_text or not target_text:632                    continue633                634                # Simple text comparison - LLM would do better comparison in real system635                if len(source_text) > 0 and len(target_text) > 0:636                    matches.append({637                        "source_entity": source_entity,638                        "target_entity": target_entity,639                        "source_text": source_text,640                        "target_text": target_text,641                        "entity_type": state["entity_type"]642                    })643            644            # If matches found, take the top one645            if matches:646                matched_pairs.append(matches[0])647        648        return {649            **state,650            "matched_pairs": matched_pairs651        }652    except Exception as e:653        return {654            **state,655            "error": f"Exception in match_entities: {str(e)}"656        }657 658def analyze_matches(state: TraceabilityState) -> TraceabilityState:659    """Analyze matches between documents to identify consistency issues."""660    try:661        if "error" in state:662            return state663        664        matched_pairs = state.get("matched_pairs", [])665        source_doc_id = state["source_document_id"]666        target_doc_id = state["target_document_id"]667        668        if not matched_pairs:669            return {670                **state,671                "analysis": "No matching entities found between the documents."672            }673        674        # Get document metadata675        source_doc = knowledge_store.get_document_by_id(source_doc_id)676        target_doc = knowledge_store.get_document_by_id(target_doc_id)677        678        # Use LLM to analyze matches679        analysis = llm_interface.find_document_connections(680            source_doc_info=source_doc,681            target_doc_info=target_doc,682            entity_pairs=matched_pairs683        )684        685        return {686            **state,687            "analysis": analysis688        }689    except Exception as e:690        return {691            **state,692            "error": f"Exception in analyze_matches: {str(e)}",693            "analysis": f"Error analyzing matches: {str(e)}"694        }695 696# =========================================================================697# Graph Building Functions698# =========================================================================699 700def build_document_extraction_graph():701    """Build and return document extraction workflow graph."""702    workflow = StateGraph(DocumentExtractionState)703    704    # Add nodes705    workflow.add_node("parse_document", parse_document)706    workflow.add_node("extract_study_info", extract_study_info)707    workflow.add_node("extract_objectives_endpoints", extract_objectives_endpoints)708    workflow.add_node("extract_population_criteria", extract_population_criteria)709    workflow.add_node("extract_study_design", extract_study_design)710    workflow.add_node("store_in_knowledge_base", store_in_knowledge_base)711    712    # Add edges - sequential process713    workflow.add_edge("parse_document", "extract_study_info")714    workflow.add_edge("extract_study_info", "extract_objectives_endpoints")715    workflow.add_edge("extract_objectives_endpoints", "extract_population_criteria")716    workflow.add_edge("extract_population_criteria", "extract_study_design")717    workflow.add_edge("extract_study_design", "store_in_knowledge_base")718    workflow.add_edge("store_in_knowledge_base", END)719    720    # Instead of using conditional edges for all nodes, 721    # let each function handle its own error status722    # This simplifies the graph structure and avoids the conditional edge issue723    724    workflow.set_entry_point("parse_document")725    return workflow.compile()726 727def build_protocol_coach_graph():728    """Build and return protocol coach workflow graph."""729    workflow = StateGraph(ProtocolCoachState)730    731    # Add nodes732    workflow.add_node("retrieve_context", retrieve_context_for_query)733    workflow.add_node("answer_query", answer_query)734    735    # Add edges736    workflow.add_edge("retrieve_context", "answer_query")737    workflow.add_edge("answer_query", END)738    739    workflow.set_entry_point("retrieve_context")740    return workflow.compile()741 742def build_content_authoring_graph():743    """Build and return content authoring workflow graph."""744    workflow = StateGraph(ContentAuthoringState)745    746    # Add nodes747    workflow.add_node("retrieve_examples", retrieve_content_examples)748    workflow.add_node("generate_content", generate_content)749    workflow.add_node("critique_content", critique_content)750    751    # Add edges752    workflow.add_edge("retrieve_examples", "generate_content")753    workflow.add_edge("generate_content", "critique_content")754    workflow.add_edge("critique_content", END)755    756    workflow.set_entry_point("retrieve_examples")757    return workflow.compile()758 759def build_traceability_graph():760    """Build and return traceability analysis workflow graph."""761    workflow = StateGraph(TraceabilityState)762    763    # Add nodes764    workflow.add_node("retrieve_entities", retrieve_document_entities)765    workflow.add_node("match_entities", match_entities)766    workflow.add_node("analyze_matches", analyze_matches)767    768    # Add edges769    workflow.add_edge("retrieve_entities", "match_entities")770    workflow.add_edge("match_entities", "analyze_matches")771    workflow.add_edge("analyze_matches", END)772    773    workflow.set_entry_point("retrieve_entities")774    return workflow.compile()