cryogenic22/doc_knowledge_base
0
1"""2Streamlit application for the Pharmaceutical R&D Knowledge Ecosystem.3"""4 5import streamlit as st6import os7import pandas as pd8import json9import tempfile10import time11from datetime import datetime12 13from pdf_processor import PDFProcessor14from knowledge_store import KnowledgeStore15from llm_interface import LLMInterface16from graph_builder import (17 init_handlers,18 build_document_extraction_graph,19 build_protocol_coach_graph,20 build_content_authoring_graph,21 build_traceability_graph22)23 24# =========================================================================25# App Setup and Configuration26# =========================================================================27 28# Page configuration29st.set_page_config(30 page_title="Pharma R&D Knowledge Ecosystem",31 page_icon="๐",32 layout="wide",33 initial_sidebar_state="expanded"34)35 36# Initialize session state variables if they don't exist37if "chat_history" not in st.session_state:38 st.session_state.chat_history = []39 40if "documents" not in st.session_state:41 st.session_state.documents = []42 43if "knowledge_base_stats" not in st.session_state:44 st.session_state.knowledge_base_stats = {45 "documents": 0,46 "studies": 0,47 "endpoints": 0,48 "objectives": 0,49 "vectors": 050 }51 52# Initialize our handlers and graphs53@st.cache_resource54def initialize_app():55 """Initialize app resources and LangGraph workflows."""56 # Get API key from environment or secrets57 api_key = os.environ.get("ANTHROPIC_API_KEY")58 if not api_key and hasattr(st, "secrets") and "ANTHROPIC_API_KEY" in st.secrets:59 api_key = st.secrets["ANTHROPIC_API_KEY"]60 61 # Initialize handlers62 pdf_processor, knowledge_store, llm_interface = init_handlers(api_key)63 64 # Build LangGraph workflows65 extraction_graph = build_document_extraction_graph()66 coach_graph = build_protocol_coach_graph()67 authoring_graph = build_content_authoring_graph()68 traceability_graph = build_traceability_graph()69 70 return {71 "pdf_processor": pdf_processor,72 "knowledge_store": knowledge_store,73 "llm_interface": llm_interface,74 "extraction_graph": extraction_graph,75 "coach_graph": coach_graph,76 "authoring_graph": authoring_graph,77 "traceability_graph": traceability_graph78 }79 80# Initialize app resources81app_resources = initialize_app()82pdf_processor = app_resources["pdf_processor"]83knowledge_store = app_resources["knowledge_store"]84llm_interface = app_resources["llm_interface"]85extraction_graph = app_resources["extraction_graph"]86coach_graph = app_resources["coach_graph"]87authoring_graph = app_resources["authoring_graph"]88traceability_graph = app_resources["traceability_graph"]89 90# =========================================================================91# Helper Functions92# =========================================================================93 94def update_knowledge_base_stats():95 """Update the knowledge base statistics in session state."""96 try:97 # Get counts of different entity types98 documents = knowledge_store.get_all_documents()99 document_count = len(documents)100 101 # Get unique protocol IDs102 protocol_ids = set()103 for doc in documents:104 if "protocol_id" in doc and doc["protocol_id"]:105 protocol_ids.add(doc["protocol_id"])106 107 # Get vector store stats108 vector_stats = knowledge_store.get_vector_store_stats()109 vector_count = vector_stats.get("document_count", 0)110 111 # Count objectives and endpoints across all protocols112 objective_count = 0113 endpoint_count = 0114 for protocol_id in protocol_ids:115 objectives = knowledge_store.get_objectives_by_protocol_id(protocol_id)116 endpoints = knowledge_store.get_endpoints_by_protocol_id(protocol_id)117 objective_count += len(objectives)118 endpoint_count += len(endpoints)119 120 # Update session state121 st.session_state.knowledge_base_stats = {122 "documents": document_count,123 "studies": len(protocol_ids),124 "objectives": objective_count,125 "endpoints": endpoint_count,126 "vectors": vector_count127 }128 except Exception as e:129 st.error(f"Error updating knowledge base stats: {e}")130 131def process_document(uploaded_file):132 """Process an uploaded document and store in knowledge base."""133 try:134 # Create a progress bar135 progress_bar = st.progress(0)136 status_text = st.empty()137 138 # Step 1: Save the uploaded file139 status_text.text("Saving uploaded file...")140 progress_bar.progress(10)141 142 # Save uploaded file temporarily143 with tempfile.NamedTemporaryFile(delete=False, suffix=".pdf") as tmp_file:144 tmp_file.write(uploaded_file.getbuffer())145 file_path = tmp_file.name146 147 # Step 2: Process through LangGraph extraction workflow148 status_text.text("Parsing document...")149 progress_bar.progress(20)150 151 # Initialize state for extraction152 initial_state = {153 "document_path": file_path,154 "status": "initialized"155 }156 157 # Run extraction workflow158 result_state = extraction_graph.invoke(initial_state)159 160 # Update progress based on status161 if result_state.get("status") == "error":162 status_text.text(f"Error: {result_state.get('error', 'Unknown error')}")163 progress_bar.progress(100)164 return {165 "status": "error",166 "error": result_state.get("error", "Unknown error"),167 "filename": uploaded_file.name168 }169 170 # Update progress171 status_text.text("Processing completed successfully!")172 progress_bar.progress(100)173 174 # Update knowledge base stats175 update_knowledge_base_stats()176 177 # Return result178 return {179 "status": "success",180 "filename": uploaded_file.name,181 "document_id": result_state.get("document_id"),182 "protocol_id": result_state.get("extracted_study", {}).get("protocol_id")183 }184 except Exception as e:185 st.error(f"Error processing document: {e}")186 return {187 "status": "error",188 "error": str(e),189 "filename": uploaded_file.name190 }191 finally:192 # Clean up temporary file193 if 'file_path' in locals():194 try:195 os.unlink(file_path)196 except:197 pass198 199def chat_with_protocol_coach(query):200 """Process a query through the Protocol Coach."""201 try:202 # Initialize state for Protocol Coach203 initial_state = {204 "query": query,205 "chat_history": st.session_state.chat_history206 }207 208 # Run Protocol Coach workflow209 result_state = coach_graph.invoke(initial_state)210 211 return {212 "status": "success",213 "response": result_state.get("response", "I couldn't generate a response."),214 "context": result_state.get("retrieved_context", [])215 }216 except Exception as e:217 return {218 "status": "error",219 "error": str(e)220 }221 222def generate_document_section(section_type, protocol_id=None, style_guide=None):223 """Generate a document section using the content authoring workflow."""224 try:225 # Initialize state for Content Authoring226 initial_state = {227 "section_type": section_type,228 "target_protocol_id": protocol_id,229 "style_guide": style_guide230 }231 232 # Run Content Authoring workflow233 result_state = authoring_graph.invoke(initial_state)234 235 return {236 "status": "success",237 "content": result_state.get("generated_content", "I couldn't generate the content."),238 "context": result_state.get("retrieved_context", [])239 }240 except Exception as e:241 return {242 "status": "error",243 "error": str(e)244 }245 246def analyze_document_traceability(source_id, target_id, entity_type):247 """Analyze traceability between two documents."""248 try:249 # Initialize state for Traceability Analysis250 initial_state = {251 "source_document_id": source_id,252 "target_document_id": target_id,253 "entity_type": entity_type254 }255 256 # Run Traceability Analysis workflow257 result_state = traceability_graph.invoke(initial_state)258 259 return {260 "status": "success",261 "analysis": result_state.get("analysis", "I couldn't perform the analysis."),262 "matched_pairs": result_state.get("matched_pairs", [])263 }264 except Exception as e:265 return {266 "status": "error",267 "error": str(e)268 }269 270# =========================================================================271# Sidebar: Document Upload and Management272# =========================================================================273 274def render_sidebar():275 """Render the sidebar for document management."""276 st.sidebar.title("Document Management")277 278 # Knowledge Base Stats279 st.sidebar.subheader("Knowledge Base Stats")280 stats = st.session_state.knowledge_base_stats281 col1, col2 = st.sidebar.columns(2)282 col1.metric("Documents", stats["documents"])283 col2.metric("Studies", stats["studies"])284 col1.metric("Objectives", stats["objectives"])285 col2.metric("Endpoints", stats["endpoints"])286 st.sidebar.metric("Vector Chunks", stats["vectors"])287 288 # Document Upload289 st.sidebar.subheader("Upload Documents")290 uploaded_files = st.sidebar.file_uploader(291 "Upload Protocol/SAP PDFs", 292 type="pdf", 293 accept_multiple_files=True,294 help="Upload clinical documents (Protocol, SAP, etc.) to add to the knowledge base."295 )296 297 # Process uploaded files if any298 if uploaded_files:299 if st.sidebar.button("Process Documents"):300 with st.sidebar.expander("Processing Results", expanded=True):301 for uploaded_file in uploaded_files:302 st.write(f"Processing: {uploaded_file.name}")303 result = process_document(uploaded_file)304 305 if result["status"] == "success":306 st.success(f"Successfully processed {result['filename']}")307 308 # Add to documents list if not already there309 doc_exists = False310 for doc in st.session_state.documents:311 if doc.get("filename") == result["filename"]:312 doc_exists = True313 break314 315 if not doc_exists:316 st.session_state.documents.append({317 "filename": result["filename"],318 "document_id": result.get("document_id"),319 "protocol_id": result.get("protocol_id"),320 "processed_date": datetime.now().strftime("%Y-%m-%d %H:%M")321 })322 else:323 st.error(f"Error processing {result['filename']}: {result.get('error', 'Unknown error')}")324 325 # Document list326 st.sidebar.subheader("Processed Documents")327 if not st.session_state.documents:328 st.sidebar.info("No documents processed yet.")329 else:330 for i, doc in enumerate(st.session_state.documents):331 with st.sidebar.expander(f"{doc['filename']}"):332 st.write(f"**Protocol ID:** {doc.get('protocol_id', 'Unknown')}")333 st.write(f"**Processed:** {doc.get('processed_date', 'Unknown')}")334 335 # Refresh Stats Button336 if st.sidebar.button("Refresh Stats"):337 update_knowledge_base_stats()338 st.sidebar.success("Stats refreshed!")339 340# =========================================================================341# Main Content Tabs342# =========================================================================343 344def render_protocol_coach_tab():345 """Render the Protocol Coach chatbot tab."""346 st.header("Protocol Coach Chatbot")347 st.info("Ask questions about the protocol documents in the knowledge base. The Protocol Coach will retrieve relevant information to answer your questions.")348 349 # Initialize or display chat history350 for message in st.session_state.chat_history:351 with st.chat_message(message["role"]):352 st.markdown(message["content"])353 354 # Chat input355 if query := st.chat_input("Ask about protocols..."):356 # Add user message to chat history and display357 st.session_state.chat_history.append({"role": "user", "content": query})358 with st.chat_message("user"):359 st.markdown(query)360 361 # Process query362 with st.chat_message("assistant"):363 with st.spinner("Thinking..."):364 result = chat_with_protocol_coach(query)365 if result["status"] == "success":366 st.markdown(result["response"])367 368 # Show context sources if debug mode enabled369 if st.session_state.get("debug_mode", False):370 with st.expander("Context Sources"):371 for i, ctx in enumerate(result.get("context", [])):372 st.write(f"**Source {i+1}:** {ctx.get('metadata', {}).get('source', 'Unknown')}")373 st.write(f"**Section:** {ctx.get('metadata', {}).get('section', 'Unknown')}")374 st.write("---")375 376 # Add assistant response to chat history377 st.session_state.chat_history.append({"role": "assistant", "content": result["response"]})378 else:379 st.error(f"Error: {result.get('error', 'Unknown error')}")380 st.session_state.chat_history.append({"role": "assistant", "content": f"Error: {result.get('error', 'Unknown error')}"})381 382def render_content_authoring_tab():383 """Render the Content Authoring tab."""384 st.header("Content Authoring Assistant")385 st.info("Generate document sections based on knowledge extracted from similar documents.")386 387 col1, col2 = st.columns([1, 1])388 389 with col1:390 st.subheader("Content Generation Settings")391 392 # Section Type Selection393 section_types = [394 "Introduction",395 "Objectives and Endpoints",396 "Study Design",397 "Study Population",398 "Statistical Considerations",399 "Inclusion Criteria",400 "Exclusion Criteria",401 "Safety Assessments",402 "Pharmacokinetic Assessments"403 ]404 section_type = st.selectbox("Select Section Type", section_types)405 406 # Protocol Selection for Context (Optional)407 protocol_options = ["--None--"]408 for doc in st.session_state.documents:409 if doc.get("protocol_id"):410 protocol_options.append(doc.get("protocol_id"))411 412 target_protocol = st.selectbox(413 "Target Protocol ID (Optional)",414 protocol_options415 )416 target_protocol = None if target_protocol == "--None--" else target_protocol417 418 # Style Guide (Optional)419 style_guide = st.text_area(420 "Style Guide (Optional)",421 placeholder="Enter any specific style guidelines or content requirements..."422 )423 424 # Generate Button425 generate_button = st.button("Generate Content")426 427 # Debug toggle428 st.session_state.debug_mode = st.checkbox("Show Context Sources", value=st.session_state.get("debug_mode", False))429 430 with col2:431 st.subheader("Generated Content")432 433 if generate_button:434 with st.spinner("Generating content..."):435 result = generate_document_section(436 section_type=section_type,437 protocol_id=target_protocol,438 style_guide=style_guide if style_guide else None439 )440 441 if result["status"] == "success":442 st.markdown(result["content"])443 444 # Show context sources if debug mode enabled445 if st.session_state.get("debug_mode", False):446 with st.expander("Context Sources"):447 for i, ctx in enumerate(result.get("context", [])):448 st.write(f"**Source {i+1}:** {ctx.get('metadata', {}).get('source', 'Unknown')}")449 st.write(f"**Section:** {ctx.get('metadata', {}).get('section', 'Unknown')}")450 st.write("---")451 else:452 st.error(f"Error: {result.get('error', 'Unknown error')}")453 454def render_traceability_tab():455 """Render the Document Traceability tab."""456 st.header("Cross-Document Traceability")457 st.info("Analyze relationships between related documents (e.g., Protocol and SAP).")458 459 col1, col2 = st.columns([1, 1])460 461 with col1:462 st.subheader("Traceability Analysis Settings")463 464 # Document Selection465 document_options = []466 for doc in st.session_state.documents:467 document_options.append({468 "id": doc.get("document_id", ""),469 "label": f"{doc['filename']} ({doc.get('protocol_id', 'Unknown')})"470 })471 472 # Source Document473 source_options = [{"id": "", "label": "--Select Source Document--"}] + document_options474 source_doc = st.selectbox(475 "Source Document",476 options=source_options,477 format_func=lambda x: x["label"]478 )479 480 # Target Document 481 target_options = [{"id": "", "label": "--Select Target Document--"}] + document_options482 target_doc = st.selectbox(483 "Target Document",484 options=target_options,485 format_func=lambda x: x["label"]486 )487 488 # Entity Type489 entity_types = [490 {"id": "objectives", "label": "Study Objectives"},491 {"id": "endpoints", "label": "Endpoints"},492 {"id": "population", "label": "Population Criteria"}493 ]494 entity_type = st.selectbox(495 "Entity Type to Compare",496 options=entity_types,497 format_func=lambda x: x["label"]498 )499 500 # Analyze Button501 analyze_button = st.button("Analyze Traceability")502 503 with col2:504 st.subheader("Analysis Results")505 506 if analyze_button:507 if not source_doc["id"] or not target_doc["id"]:508 st.error("Please select both source and target documents.")509 else:510 with st.spinner("Analyzing traceability..."):511 result = analyze_document_traceability(512 source_id=source_doc["id"],513 target_id=target_doc["id"],514 entity_type=entity_type["id"]515 )516 517 if result["status"] == "success":518 st.markdown(result["analysis"])519 520 # Show matched pairs if debug mode enabled521 if st.session_state.get("debug_mode", False) and result.get("matched_pairs"):522 with st.expander("Matched Entity Pairs"):523 for i, pair in enumerate(result["matched_pairs"]):524 st.write(f"**Pair {i+1}**")525 st.write(f"**Source:** {pair.get('source_text', 'Unknown')}")526 st.write(f"**Target:** {pair.get('target_text', 'Unknown')}")527 st.write("---")528 else:529 st.error(f"Error: {result.get('error', 'Unknown error')}")530 531def render_knowledge_explorer_tab():532 """Render the Knowledge Base Explorer tab."""533 st.header("Knowledge Base Explorer")534 st.info("Explore the structured data extracted from documents in the knowledge base.")535 536 # Entity Type Selection537 entity_types = [538 {"id": "studies", "label": "Studies"},539 {"id": "objectives", "label": "Study Objectives"},540 {"id": "endpoints", "label": "Endpoints"},541 {"id": "population", "label": "Population Criteria"},542 {"id": "documents", "label": "Documents"}543 ]544 entity_type = st.selectbox(545 "Select Entity Type",546 options=entity_types,547 format_func=lambda x: x["label"]548 )549 550 # Filter by Protocol ID (Optional)551 protocol_options = ["--All Protocols--"]552 for doc in st.session_state.documents:553 if doc.get("protocol_id") and doc.get("protocol_id") not in protocol_options:554 protocol_options.append(doc.get("protocol_id"))555 556 filter_protocol = st.selectbox(557 "Filter by Protocol ID",558 protocol_options559 )560 filter_protocol = None if filter_protocol == "--All Protocols--" else filter_protocol561 562 # Search Query (Optional)563 search_query = st.text_input(564 "Search Query (Optional)",565 placeholder="Enter text to search for..."566 )567 568 # Display Results569 st.subheader("Results")570 571 try:572 # Retrieve data based on entity type573 if entity_type["id"] == "studies":574 if filter_protocol:575 data = [knowledge_store.get_study_by_protocol_id(filter_protocol)]576 else:577 data = knowledge_store.get_all_studies()578 elif entity_type["id"] == "objectives":579 if filter_protocol:580 data = knowledge_store.get_objectives_by_protocol_id(filter_protocol)581 else:582 # Get all objectives across protocols583 data = []584 documents = knowledge_store.get_all_documents()585 protocol_ids = set()586 for doc in documents:587 if "protocol_id" in doc and doc["protocol_id"]:588 protocol_ids.add(doc["protocol_id"])589 590 for pid in protocol_ids:591 data.extend(knowledge_store.get_objectives_by_protocol_id(pid))592 elif entity_type["id"] == "endpoints":593 if filter_protocol:594 data = knowledge_store.get_endpoints_by_protocol_id(filter_protocol)595 else:596 # Get all endpoints across protocols597 data = []598 documents = knowledge_store.get_all_documents()599 protocol_ids = set()600 for doc in documents:601 if "protocol_id" in doc and doc["protocol_id"]:602 protocol_ids.add(doc["protocol_id"])603 604 for pid in protocol_ids:605 data.extend(knowledge_store.get_endpoints_by_protocol_id(pid))606 elif entity_type["id"] == "population":607 if filter_protocol:608 data = knowledge_store.get_population_criteria_by_protocol_id(filter_protocol)609 else:610 # Get all population criteria across protocols611 data = []612 documents = knowledge_store.get_all_documents()613 protocol_ids = set()614 for doc in documents:615 if "protocol_id" in doc and doc["protocol_id"]:616 protocol_ids.add(doc["protocol_id"])617 618 for pid in protocol_ids:619 data.extend(knowledge_store.get_population_criteria_by_protocol_id(pid))620 elif entity_type["id"] == "documents":621 if filter_protocol:622 data = knowledge_store.get_documents_by_protocol_id(filter_protocol)623 else:624 data = knowledge_store.get_all_documents()625 else:626 data = []627 628 # Filter by search query if provided629 if search_query:630 filtered_data = []631 search_lower = search_query.lower()632 for item in data:633 # Convert item to string for searching634 item_str = json.dumps(item).lower()635 if search_lower in item_str:636 filtered_data.append(item)637 data = filtered_data638 639 # Display results640 if not data:641 st.info("No data found.")642 else:643 st.write(f"{len(data)} items found")644 645 # Display as table if possible, otherwise as JSON646 try:647 df = pd.DataFrame(data)648 st.dataframe(df, use_container_width=True)649 except Exception as e:650 st.json(data)651 except Exception as e:652 st.error(f"Error retrieving data: {e}")653 654# =========================================================================655# Main App656# =========================================================================657 658def main():659 """Main application function."""660 st.title("๐ง Pharmaceutical R&D Knowledge Ecosystem")661 662 # Render the sidebar for document management663 render_sidebar()664 665 # Initialize knowledge base stats on first load666 if st.session_state.knowledge_base_stats["documents"] == 0:667 update_knowledge_base_stats()668 669 # Main content tabs670 tab1, tab2, tab3, tab4 = st.tabs([671 "๐ Content Authoring",672 "๐ค Protocol Coach", 673 "๐ Knowledge Explorer",674 "๐ Cross-Document Traceability"675 ])676 677 with tab1:678 render_content_authoring_tab()679 680 with tab2:681 render_protocol_coach_tab()682 683 with tab3:684 render_knowledge_explorer_tab()685 686 with tab4:687 render_traceability_tab()688 689 # Footer690 st.markdown("---")691 st.caption("Pharmaceutical R&D Knowledge Ecosystem | A demonstration of AI-assisted document processing and knowledge management")692 693if __name__ == "__main__":694 main()