import json
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
import io
REDIRECT_URI = 'https://doczydev.aarete.com:8502'
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")
_,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.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"
)
bucket = 'doczy-dev-infra-textract'
objects = s3_client.list_objects_v2(Bucket=bucket, Prefix="training-data/")
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) # MODIFIED - Append 'All' in the front instead of at the end
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_client.Object(bucket,file_name)
data=obj.get()['Body'].read()
pdf_viewer(io.BytesIO(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")