import json import security import boto3 from langchain.prompts import PromptTemplate from langchain.embeddings.bedrock import BedrockEmbeddings from langchain.llms.bedrock import Bedrock from langchain_community.vectorstores import Chroma from constants import CHROMA_SETTINGS, EMBEDDING_MODEL_NAME, PERSIST_DIRECTORY, MODEL_ID, MODEL_BASENAME, SOURCE_DIRECTORY, USER_LIST from langchain.chains import RetrievalQA import streamlit as st from streamlit_extras.add_vertical_space import add_vertical_space from streamlit_pdf_viewer import pdf_viewer # needs to be installed on the server import os import pandas as pd import numpy as np import util import anthropic from pydantic import BaseModel from typing import List import re import base64 from sf_conn import get_snowflake_conn from sf_conn import get_client_names, get_secret, save_to_sf import io (REDIRECT_URI, create_batch_url, doczy_pipeline) = util.load_page_details(2) user_list = USER_LIST st.set_page_config(layout = "wide") # Sidebar contents # with st.sidebar: # st.title("Doczy.AI ™") # st.markdown( # """ # ## About # This app extracts data from contracts # """ # ) # add_vertical_space(15) # # st.write("Doczy") # AARETE LOGO x,y,z = st.columns([15,2,15]) with y: st.image('aaretelogo.png') hide_img_fs = ''' ''' st.markdown(hide_img_fs, unsafe_allow_html=True) _,c1= st.columns([5,1]) try: util.setup_page(REDIRECT_URI) except: st.write("SSO Failed") st.session_state['user_info'] = {'mail': 'maamseek@aarete.com', 'displayName': 'Mayank Aamseek'} try: c1.write(f"User: **{st.session_state.user_info['displayName']}**") user_mail = st.session_state.user_info['mail'] except KeyError as e: st.write("Session Expired.") #st.write("Please sign-in to use this app.") auth_url = security.get_auth_url(REDIRECT_URI) st.markdown(f"Sign In", unsafe_allow_html=True) st.stop() # remove below try except statement if comparison with actual vales is not required try: conn = get_snowflake_conn('STG') cur = conn.cursor() query = 'select * from "TRAINING_DATA_RAW"' cur.execute(query) field_values = pd.DataFrame.from_records(iter(cur), columns=[x[0] for x in cur.description]) field_values['Document_Name'] = field_values['DOCUMENT_NAME'] except: # field_values = pd.read_csv('contract_field_values.csv', encoding='unicode_escape', skipinitialspace=True) # field_values = field_values.loc[:, ~field_values.columns.str.contains('Unnamed:')] st.write("Conn failed, unable to fetch data from training data table in Snowflake") try: query = 'select * from "PROMPT_CONFIG"' cur.execute(query) fields = pd.DataFrame.from_records(iter(cur), columns=[x[0] for x in cur.description]) fields.rename(columns={'FIELD_DESC': 'Field Name'}, inplace = True) fields.rename(columns={'PROMPT': 'Interrogation Question?'}, inplace = True) fields.rename(columns={'GROUP_ID': 'PRIORITY'}, inplace = True) fields.rename(columns={'FIELD_NAME': 'SF_DB_COL_NAME'}, inplace = True) fields.rename(columns={'FM_MODEL_ID': 'llm_selected'}, inplace = True) except Exception as e: st.write("Unable to fetch data from Snowflake: ",e) # fields = pd.read_csv('contract_fields.csv', encoding='unicode_escape', skipinitialspace=True) # fields = fields[~fields['SF_COL_NAME'].str.endswith('_PG', na=None)] # change the code below if contract list is fetched from snowflake s3_client = boto3.client('s3', region_name="us-east-2" ) client_list, s3_paths = get_client_names() client_s3_paths = dict(zip(client_list, s3_paths)) client_row = st.columns([0.2, 0.7, 0.1]) with client_row[0]: st.write("**Client Name**") with client_row[1]: client = st.selectbox('Client Name',(client_list), label_visibility = "collapsed", index= None) # client_bucket = client_s3_paths.get(client) client_bucket = 'doczy-dev-infra-textract' batch_objects = s3_client.list_objects_v2(Bucket=client_bucket , Prefix="contracts-landing-zone/", Delimiter='/') batch_list = [] for prefix in batch_objects['CommonPrefixes']: batch_list.append(prefix['Prefix'][:-1].split('/')[-1]) path_row = st.columns([0.1, 0.8]) with path_row[0]: st.write("**Batch ID**") with path_row[1]: batch_id = st.selectbox('**Batch ID**', batch_list, label_visibility = "collapsed", index = None) if batch_id: objects = s3_client.list_objects_v2(Bucket=client_bucket, Prefix="contract-text-file/"+batch_id+"/") file_list = [] for obj in objects['Contents']: if not obj['Key'].endswith('/'): file_list.append(obj['Key']) contract_list = sorted(file_list) file_row = st.columns([0.2, 0.7, 0.1]) with file_row[0]: st.write("**Contract Name**") with file_row[1]: file_name = st.selectbox('Select a file', ['All'] + contract_list, label_visibility = "collapsed", index= None) field_row = st.columns([0.2, 0.7, 0.1]) with field_row[0]: st.write("**Field Group**") with field_row[1]: field_group = st.selectbox('Field Group',('Unique Key', 'Contract Related', 'Pricing Before Carveouts - I' , 'Pricing Before Carveouts - II', 'Carveout Indicator, Code Type and Code #s - I' , 'Carveout Indicator, Code Type and Code #s - II', 'Carveout Indicator, Code Type and Code #s - III' , 'Optimize Carving Indic.', 'Carveout Method - I', 'Carveout Method - II', 'Provider' , 'Timeline'), label_visibility = "collapsed", index = None) if field_group == 'Unique Key': fields = fields[fields['PRIORITY'] == 'A'] elif field_group == 'Contract Related': fields = fields[fields['PRIORITY'] == 'C'] elif field_group == 'Pricing Before Carveouts - I': fields = fields[fields['PRIORITY'] == 'B'] fields = np.array_split(fields, 2)[0] elif field_group == 'Pricing Before Carveouts - II': fields = fields[fields['PRIORITY'] == 'B'] fields = np.array_split(fields, 2)[1] elif field_group == 'Carveout Indicator, Code Type and Code #s - I': fields = fields[fields['PRIORITY'] == 'F'] fields = np.array_split(fields, 3)[0] elif field_group == 'Carveout Indicator, Code Type and Code #s - II': fields = fields[fields['PRIORITY'] == 'F'] fields = np.array_split(fields, 3)[1] elif field_group == 'Carveout Indicator, Code Type and Code #s - III': fields = fields[fields['PRIORITY'] == 'F'] fields = np.array_split(fields, 3)[2] elif field_group == 'Carveout Methodology - I': fields = fields[fields['PRIORITY'] == 'G'] fields = np.array_split(fields, 4)[0] elif field_group == 'Carveout Methodology - II': fields = fields[fields['PRIORITY'] == 'G'] fields = np.array_split(fields, 4)[1] elif field_group == 'Carveout Method - III': fields = fields[fields['PRIORITY'] == 'G'] fields = np.array_split(fields, 4)[2] elif field_group == 'Carveout Method - IV': fields = fields[fields['PRIORITY'] == 'G'] fields = np.array_split(fields, 4)[3] elif field_group == 'Provider': fields = fields[fields['PRIORITY'] == 'D'] elif field_group == 'Timeline': fields = fields[fields['PRIORITY'] == 'E'] if st.button("Show Results"): query = 'select * from "DOCZY_PIPELINE_RAW_OUTPUT"' cur.execute(query) df2 = pd.DataFrame.from_records(iter(cur), columns=[x[0] for x in cur.description]) # get this dataframe from snowflake table # df2 = pd.DataFrame(columns=['Contract Name','Field Name', 'SF_DB_COL_NAME', 'Snippet','Page Number' # , 'Field Extracted Value', 'Actual Value','Imputed Value']) df2.to_csv('temp2.csv', index=False) if st.button("Show PDF"): if file_name == None or file_name == "All": st.error("Choose one specific file.") else: with st.sidebar: st.markdown( """ """, unsafe_allow_html=True, ) s3_obj = s3_client.get_object(Bucket = bucket, Key = 'textract-receiver-processed-pdfs/batch_2/2000-01-01 UCSD Medical Center PPA 14007097.PDF') data=s3_obj['Body'].read() pdf_viewer(data, width=1500) # if st.button("Show PDF"): # if file_name == None or file_name == "All": # st.error("Choose one specific file.") # else: # with st.sidebar: # with open(file_name, "rb") as f: # base64_pdf = base64.b64encode(f.read()).decode('utf-8') # # Embedding PDF in HTML # pdf_display = F'' # # Displaying File # st.markdown( # """ # # """, # unsafe_allow_html=True, # ) # st.markdown(pdf_display, unsafe_allow_html=True) df2 = pd.read_csv('temp2.csv') df2['Imputed Value'] = '' edited_df = st.data_editor(df2) @st.cache_data def convert_df(df): return df.to_csv(index=False).encode('utf-8') csv = convert_df(edited_df) buttons = st.columns(3) with buttons[0]: # st.button("Save All Imputations") st.download_button("Download Table", csv, "file.csv", "text/csv", key='download-csv') with buttons[1]: # st.download_button("Download Table", csv, "file.csv", "text/csv", key='download-csv') st.write("") with buttons[2]: if st.button("Kickoff Database Integration"): st.write("Stored in DB")