cryogenic22/doc_knowledge_base
0
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()