Merged DEV into dev_umistry

This commit is contained in:
Umang Mistry
2024-03-11 22:28:27 +00:00
6 changed files with 636 additions and 507 deletions
+5 -2
View File
@@ -50,12 +50,15 @@ Thumbs.db
# Data Files
streamlit/history.csv
streamlit/RESULTS/
streamlit/RESULTS
streamlit/DB/
streamlit/RAW_DOCUMENTS/
streamlit/SOURCE_DOCUMENTS/
streamlit/contract_field_values.csv
streamlit/contract_fields.csv
streamlit/sample.csv
streamlit/temp1.csv
streamlit/temp2.csv
# env
streamlit/venv
+89 -73
View File
@@ -5,9 +5,13 @@ import streamlit as st
import pandas as pd
from io import StringIO
from datetime import datetime
import boto3
import util
REDIRECT_URI = 'http://172.29.20.126:8501'
user_list = ['maamseek@aarete.com', 'smahdavian@aarete.com', 'ahinge@aarete.com', 'akadam@aarete.com'
, 'piragavarapu@aarete.com', 'umistry@aarete.com']
st.set_page_config(layout = "wide")
# # Sidebar contents
# with st.sidebar:
# st.title("Doczy.AI ™")
@@ -21,92 +25,104 @@ st.set_page_config(layout = "wide")
# add_vertical_space(15)
# # st.write("Doczy")
util.setup_page(REDIRECT_URI)
if st.session_state.user_info['mail'] in user_list:
client_row = st.columns([0.1, 0.8])
with client_row[0]:
st.write("**Client Name**")
with client_row[1]:
client = st.selectbox('Client Name',('Delaware First Health, Inc.', 'Community Health Choice, Inc','CareSource Network Partners LLC',
'HealthNet of Cali', 'Oklahoma Complete Health, Inc', 'HealthFirst', 'Molina Healthcare of TX', 'AvMed', 'Arizona Care1st',
'WellCare New Jersey'), label_visibility = "collapsed")
s3_client = boto3.client('s3',
region_name="us-east-1",
)
objects = s3_client.list_objects_v2(Bucket='doczy-dev-infra-textract'
, Prefix="batches/batch_1/", Delimiter='/')
def file_selector(folder_path='.'):
filenames = os.listdir(folder_path)
selected_filename = st.selectbox('**Path to folder**', filenames, label_visibility = "collapsed")
return os.path.join(folder_path, selected_filename)
# return selected_filename
folder_list = []
for prefix in objects['CommonPrefixes']:
folder_list.append(prefix['Prefix'][:-1].split('/')[-1])
path_row = st.columns([0.1, 0.8])
with path_row[0]:
st.write("**Path to folder**")
with path_row[1]:
# Directory = st.text_input("**Path to folder**", label_visibility = "collapsed")
Directory = file_selector()
client_row = st.columns([0.1, 0.8])
with client_row[0]:
st.write("**Client Name**")
with client_row[1]:
client = st.selectbox('Client Name',(folder_list), index=7, label_visibility = "collapsed")
checks = st.columns([0.1, 0.12, 0.12, 0.12, 0.12, 0.12, 0.12, 0.12])
with checks[0]:
st.write("**Group No.**")
with checks[1]:
a = st.checkbox('Unique Key', key = str(1))
with checks[2]:
b = st.checkbox('Pricing Before Carveouts', key = str(2))
with checks[3]:
c = st.checkbox('Contract Related', key = str(3))
with checks[4]:
d = st.checkbox('Provider', key = str(4))
with checks[5]:
e = st.checkbox('Timeline', key = str(5))
with checks[6]:
f = st.checkbox('Carveout Indicator', key = str(6))
with checks[7]:
g = st.checkbox('Carveout Methodology', key = str(7))
folder_objects = s3_client.list_objects_v2(Bucket='doczy-dev-infra-textract'
, Prefix="batches/batch_1/"+client+"/", Delimiter='/')
add_vertical_space(1)
folder_list_2 = []
for prefix in folder_objects['CommonPrefixes']:
folder_list_2.append(prefix['Prefix'][:-1].split('/')[-1])
df = pd.DataFrame(columns=['Request ID','Contract ID','Contract Name','Unique Key','Pricing Before Carveouts'
, 'Contract Related', 'Provider', 'Timeline', 'Carveout Indicator', 'Carveout Methodology'])
file_list = []
path_row = st.columns([0.1, 0.8])
with path_row[0]:
st.write("**Path to folder**")
with path_row[1]:
Directory = st.selectbox('**Path to folder**', folder_list_2, label_visibility = "collapsed")
if st.button("Read the contracts from Path"):
for filename in os.listdir(Directory):
# with open(os.path.join(Directory, filename), encoding="utf8") as f:
# context = f.read()
file_list.append(filename)
checks = st.columns([0.1, 0.12, 0.12, 0.12, 0.12, 0.12, 0.12, 0.12])
with checks[0]:
st.write("**Group No.**")
with checks[1]:
a = st.checkbox('Unique Key', key = str(1))
with checks[2]:
b = st.checkbox('Pricing Before Carveouts', key = str(2))
with checks[3]:
c = st.checkbox('Contract Related', key = str(3))
with checks[4]:
d = st.checkbox('Provider', key = str(4))
with checks[5]:
e = st.checkbox('Timeline', key = str(5))
with checks[6]:
f = st.checkbox('Carveout Indicator', key = str(6))
with checks[7]:
g = st.checkbox('Carveout Methodology', key = str(7))
df['Contract Name'] = file_list
df['Request ID'] = range(len(file_list))
df['Contract ID'] = file_list
# df['Folder Name'] = Directory
# df['Updated Group Number'] = updated_group
df['Unique Key'] = a
df['Pricing Before Carveouts'] = b
df['Contract Related'] = c
df['Provider'] = d
df['Timeline'] = e
df['Carveout Indicator'] = f
df['Carveout Methodology'] = g
df.to_csv('temp1.csv', index=False)
add_vertical_space(1)
add_vertical_space(1)
df = pd.DataFrame(columns=['Request ID','Contract ID','Contract Name','Unique Key','Pricing Before Carveouts'
, 'Contract Related', 'Provider', 'Timeline', 'Carveout Indicator', 'Carveout Methodology'])
file_list = []
file_objects = s3_client.list_objects_v2(Bucket='doczy-dev-infra-textract'
, Prefix="batches/batch_1/"+client+"/"+Directory+"/", Delimiter='/')
# df_copy = df.set_index(df.columns[0]).copy()
df2 = pd.read_csv('temp1.csv')
edited_df = st.data_editor(df2)
if st.button("Read the contracts from Path"):
for obj in file_objects.get('Contents',[]):
if not obj['Key'].endswith('/'):
file_list.append(obj['Key'].split('/')[-1])
df['Contract Name'] = file_list
df['Request ID'] = range(len(file_list))
df['Contract ID'] = file_list
df['Unique Key'] = a
df['Pricing Before Carveouts'] = b
df['Contract Related'] = c
df['Provider'] = d
df['Timeline'] = e
df['Carveout Indicator'] = f
df['Carveout Methodology'] = g
dir_path = os.path.dirname(os.path.realpath(__file__))
print(f'DEBUGGING: PWD= {dir_path}')
df.to_csv('temp1.csv', index=False)
add_vertical_space(1)
# df_copy = df.set_index(df.columns[0]).copy()
df2 = pd.read_csv('temp1.csv')
edited_df = st.data_editor(df2)
@st.cache_data
def convert_df(df):
return df.to_csv(index=False).encode('utf-8')
@st.cache_data
def convert_df(df):
return df.to_csv(index=False).encode('utf-8')
csv = convert_df(edited_df)
csv = convert_df(edited_df)
buttons = st.columns([0.45, 0.35, 0.2])
with buttons[0]:
st.button("Save All Edits")
with buttons[1]:
st.download_button("Download Table", csv, "file.csv", "text/csv", key='download-csv')
with buttons[2]:
st.button("Run Doczy.AI Pipeline")
buttons = st.columns([0.8, 0.2])
with buttons[0]:
st.download_button("Download Table", csv, "file.csv", "text/csv", key='download-csv')
with buttons[1]:
st.button("Run Doczy.AI Pipeline")
else:
st.write("Access Denied")
+163 -154
View File
@@ -12,9 +12,14 @@ import streamlit as st
from streamlit_extras.add_vertical_space import add_vertical_space
import os
import pandas as pd
import util
REDIRECT_URI = 'http://172.29.20.126:8502'
user_list = ['maamseek@aarete.com', 'smahdavian@aarete.com', 'ahinge@aarete.com', 'akadam@aarete.com'
, 'piragavarapu@aarete.com', 'umistry@aarete.com', 'ahutchison@aarete.com', 'bgrunst@aarete.com', 'ddimeglio@aarete.com'
, 'vnair@aarete.com']
st.set_page_config(layout = "wide")
# Sidebar contents
with st.sidebar:
st.title("Doczy.AI ™")
@@ -28,170 +33,174 @@ with st.sidebar:
add_vertical_space(15)
# st.write("Doczy")
sample = pd.read_csv('sample.csv')
sample = sample[['Filename', 'Attribute', 'Query', 'Answer']]
attribute_list = ['Agreement Name', 'Agreement Type', 'Contract State', 'Contract Type', 'Create Date', 'Effective Date', 'Gold Carded',
'Modify Date', 'National Contract', 'Provider State', 'Summary', 'Termination Date', 'TIN', 'Value-Based Contract']
sample = sample[sample['Attribute'].isin(attribute_list)]
field_prompt_mapping = sample[['Attribute', 'Query']].drop_duplicates().dropna()
field_prompt_mapping = dict(zip(field_prompt_mapping.Attribute, field_prompt_mapping.Query))
util.setup_page(REDIRECT_URI)
if st.session_state.user_info['mail'] in user_list:
fields = pd.read_csv('contract_fields.csv', encoding='unicode_escape', skipinitialspace=True)
fields = fields[fields['PRIORITY'] == 'A']
fields = fields.drop_duplicates(subset='Field Name', keep="first").sort_values('Field Name')
fields['Interrogation Question?'] = fields['Interrogation Question?'].fillna(' ')
field_prompt_mapping = dict(zip(fields['Field Name'], fields['Interrogation Question?']))
def file_selector(folder_path='RAW_DOCUMENTS'):
filenames = os.listdir(folder_path)
selected_filename = st.selectbox('Select a file', filenames, label_visibility = "collapsed")
# return os.path.join(folder_path, selected_filename)
return selected_filename
def file_selector(folder_path=SOURCE_DIRECTORY):
filenames = os.listdir(folder_path)
selected_filename = st.selectbox('Select a file', filenames, label_visibility = "collapsed")
# return os.path.join(folder_path, selected_filename)
return selected_filename
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.text_input("**Contract Name**", label_visibility = "collapsed")
file_name = file_selector()
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.text_input("**Contract Name**", label_visibility = "collapsed")
file_name = file_selector()
lob_row = st.columns([0.2, 0.7, 0.1])
with lob_row[0]:
st.write("**LOB**")
with lob_row[1]:
lob = st.selectbox('LOB',('Medicare', 'Medicaid'), label_visibility = "collapsed")
# lob_row = st.columns([0.2, 0.7, 0.1])
# with lob_row[0]:
# st.write("**LOB**")
# with lob_row[1]:
# lob = st.selectbox('LOB',('Medicare', 'Medicaid'), label_visibility = "collapsed")
llm_row = st.columns([0.2, 0.7, 0.1])
with llm_row[0]:
st.write("**Langauge Model**")
with llm_row[1]:
llm_selected = st.selectbox('Langauge Model',('Llama 2 Chat 13B', 'Llama 2 Chat 70B', 'Titan Text Express'), label_visibility = "collapsed")
llm_row = st.columns([0.2, 0.7, 0.1])
with llm_row[0]:
st.write("**Langauge Model**")
with llm_row[1]:
llm_selected = st.selectbox('Langauge Model',('Llama 2 Chat 13B', 'Llama 2 Chat 70B', 'Titan Text Express'), label_visibility = "collapsed")
page_list = []
with open(os.path.join(SOURCE_DIRECTORY, file_name[:-4]+'.txt'), 'r') as infile:
text = infile.read()
page_count = text.count('Start of Page No. = ')
for page in range(page_count+1):
file_path = "SOURCE_DOCUMENTS\\" + f'{file_name[:-4]}_page{page}.txt'
dict_with_pages = { 'source': { '$eq': file_path }}
page_list.append(dict_with_pages)
page_list = []
with open(os.path.join(SOURCE_DIRECTORY, file_name), 'r') as infile:
text = infile.read()
page_count = text.count('Start of Page No. = ')
for page in range(page_count+1):
file_path = "SOURCE_DOCUMENTS\\" + f'{file_name[:-4]}_page{page}.txt'
dict_with_pages = { 'source': { '$eq': file_path }}
page_list.append(dict_with_pages)
# AWS_ACCESS_KEY_ID = os.getenv('AWS_ACCESS_KEY_ID')
# AWS_SECRET_ACCESS_KEY = os.getenv('AWS_SECRET_ACCESS_KEY')
# AWS_SESSION_TOKEN=os.getenv('AWS_SESSION_TOKEN')
# AWS_ACCESS_KEY_ID = os.getenv('AWS_ACCESS_KEY_ID')
# AWS_SECRET_ACCESS_KEY = os.getenv('AWS_SECRET_ACCESS_KEY')
# AWS_SESSION_TOKEN=os.getenv('AWS_SESSION_TOKEN')
# Setup bedrock
bedrock_runtime = boto3.client(
service_name="bedrock-runtime",
region_name="us-east-1"
)
embeddings = BedrockEmbeddings(
client=bedrock_runtime,
model_id="amazon.titan-embed-text-v1",
)
DB = Chroma(
persist_directory=PERSIST_DIRECTORY,
embedding_function=embeddings,
client_settings=CHROMA_SETTINGS,
)
RETRIEVER = DB.as_retriever(search_kwargs={"filter": {"$or": page_list}, "k": 4})
if llm_selected == 'Titan Text Express':
LLM = Bedrock(
model_id="amazon.titan-text-express-v1",
client=bedrock_runtime,
model_kwargs={
"maxTokenCount": 4096,
"stopSequences": [],
"temperature": 0,
"topP": 1,
}
# Setup bedrock
bedrock_runtime = boto3.client(
service_name="bedrock-runtime",
region_name="us-east-1"
)
elif llm_selected == 'Llama 2 Chat 70B':
LLM = Bedrock(
model_id="meta.llama2-70b-chat-v1",
embeddings = BedrockEmbeddings(
client=bedrock_runtime,
model_kwargs={
"max_gen_len": 512,
"temperature": 0,
# "topP": 0.9,
}
model_id="amazon.titan-embed-text-v1",
)
DB = Chroma(
persist_directory=PERSIST_DIRECTORY,
embedding_function=embeddings,
client_settings=CHROMA_SETTINGS,
)
RETRIEVER = DB.as_retriever(search_kwargs={"filter": {"$or": page_list}, "k": 4})
if llm_selected == 'Titan Text Express':
LLM = Bedrock(
model_id="amazon.titan-text-express-v1",
client=bedrock_runtime,
model_kwargs={
"maxTokenCount": 4096,
"stopSequences": [],
"temperature": 0,
"topP": 1,
}
)
elif llm_selected == 'Llama 2 Chat 70B':
LLM = Bedrock(
model_id="meta.llama2-70b-chat-v1",
client=bedrock_runtime,
model_kwargs={
"max_gen_len": 512,
"temperature": 0,
# "topP": 0.9,
}
)
else:
LLM = Bedrock(
model_id="meta.llama2-13b-chat-v1",
client=bedrock_runtime,
model_kwargs={
"max_gen_len": 512,
"temperature": 0,
# "topP": 0.9,
}
)
template = """
Use the following pieces of context to answer the question at the end. If you don't know the answer,\
just say that you don't know, don't try to make up an answer.
{context}
Question: {question}
Answer:"""
prompt = PromptTemplate(input_variables=["context", "question"], template=template)
QA = RetrievalQA.from_chain_type(
llm=LLM,
chain_type="stuff",
retriever=RETRIEVER,
return_source_documents=True,
chain_type_kwargs={"prompt": prompt},
)
# query = "In which state or states is the Contract applicable? Answer in one or two words. State name: "
# response = QA({"query":query})
# st.write(query)
# st.write(response['result'])
# st.write("-----------")
# st.write(response)
# clicked = st.button("Show Results")
df = pd.DataFrame(columns=['Contract Name','Field Name','Snippet','Page Number','Confidence Level',
'Field Extracted Value','Imputed Value'])
field_list = list(field_prompt_mapping.keys())
query_list = [field_prompt_mapping[x] for x in field_list]
score_list = [DB.similarity_search_with_relevance_scores(query, k=4, filter={"$or": page_list}) for query in query_list]
confidence_list = []
for score in score_list:
confidence_list.append(max(d[1] for d in score))
# st.write(confidence_list)
if st.button("Show Results"):
response_list = [QA({"query":query}) for query in query_list]
answer_list = [response['result'] for response in response_list]
doc_list = [response['source_documents'] for response in response_list]
snippet_list = [str(doc[0].page_content) for doc in doc_list]
page_no_list = [int(str(doc[0].metadata["source"]).rsplit('_page')[1].replace('.txt',''))+1 for doc in doc_list]
df['Field Name'] = field_list
df['Contract Name'] = file_name
df['Snippet'] = snippet_list
df['Page Number'] = page_no_list
df['Confidence Level'] = confidence_list
df['Field Extracted Value'] = answer_list
df.to_csv('temp2.csv', index=False)
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")
with buttons[1]:
st.download_button("Download Table", csv, "file.csv", "text/csv", key='download-csv')
with buttons[2]:
st.button("Kickoff Database Integration")
else:
LLM = Bedrock(
model_id="meta.llama2-13b-chat-v1",
client=bedrock_runtime,
model_kwargs={
"max_gen_len": 512,
"temperature": 0,
# "topP": 0.9,
}
)
template = """
Use the following pieces of context to answer the question at the end. If you don't know the answer,\
just say that you don't know, don't try to make up an answer.
{context}
Question: {question}
Answer:"""
prompt = PromptTemplate(input_variables=["context", "question"], template=template)
QA = RetrievalQA.from_chain_type(
llm=LLM,
chain_type="stuff",
retriever=RETRIEVER,
return_source_documents=True,
chain_type_kwargs={"prompt": prompt},
)
# query = "In which state or states is the Contract applicable? Answer in one or two words. State name: "
# response = QA({"query":query})
# st.write(query)
# st.write(response['result'])
# st.write("-----------")
# st.write(response)
# clicked = st.button("Show Results")
df = pd.DataFrame(columns=['Contract Name','Field Name','Snippet','Page Number','Confidence Level',
'Field Extracted Value','Imputed Value'])
field_list = list(field_prompt_mapping.keys())
query_list = [field_prompt_mapping[x] for x in field_list]
score_list = [DB.similarity_search_with_relevance_scores(query, k=4, filter={"$or": page_list}) for query in query_list]
confidence_list = []
for score in score_list:
confidence_list.append(max(d[1] for d in score))
# st.write(confidence_list)
if st.button("Show Results"):
response_list = [QA({"query":query}) for query in query_list]
answer_list = [response['result'] for response in response_list]
doc_list = [response['source_documents'] for response in response_list]
snippet_list = [str(doc[0].page_content) for doc in doc_list]
page_no_list = [int(str(doc[0].metadata["source"]).rsplit('_page')[1].replace('.txt',''))+1 for doc in doc_list]
df['Field Name'] = field_list
df['Contract Name'] = file_name
df['Snippet'] = snippet_list
df['Page Number'] = page_no_list
df['Confidence Level'] = confidence_list
df['Field Extracted Value'] = answer_list
df.to_csv('temp2.csv', index=False)
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")
with buttons[1]:
st.download_button("Download Table", csv, "file.csv", "text/csv", key='download-csv')
with buttons[2]:
st.button("Kickoff Database Integration")
st.write("Access Denied")
+314 -278
View File
@@ -15,9 +15,14 @@ from datetime import datetime
import random
import os
import dateutil
import util
REDIRECT_URI = 'http://172.29.20.126:8503'
user_list = ['maamseek@aarete.com', 'smahdavian@aarete.com', 'ahinge@aarete.com', 'akadam@aarete.com'
, 'piragavarapu@aarete.com', 'umistry@aarete.com', 'ahutchison@aarete.com', 'bgrunst@aarete.com', 'ddimeglio@aarete.com'
, 'vnair@aarete.com']
st.set_page_config(layout = "wide")
# Sidebar contents
with st.sidebar:
st.title("Doczy.AI ™")
@@ -31,304 +36,335 @@ with st.sidebar:
add_vertical_space(15)
# st.write("Doczy")
fields = pd.read_csv('contract_fields.csv', encoding='unicode_escape', skipinitialspace=True)
fields = fields[fields['PRIORITY'].isin(['A','C'])]
field_values = pd.read_csv('contract_field_values.csv', encoding='unicode_escape', skipinitialspace=True)
fields = fields.drop_duplicates(subset='Field Name', keep="first").sort_values('Field Name')
fields = fields[~fields['Field Name'].isnull()]
field_prompt_mapping = dict(zip(fields['Field Name'], fields['Interrogation Question?']))
util.setup_page(REDIRECT_URI)
if st.session_state.user_info['mail'] in user_list:
field_row = st.columns([0.15, 0.45, 0.4])
with field_row[0]:
st.write("**Field Name**")
with field_row[1]:
field = st.selectbox('Field Name',sorted(set(field_prompt_mapping.keys())), index=0, label_visibility = "collapsed")
fields = pd.read_csv('contract_fields.csv', encoding='unicode_escape', skipinitialspace=True)
fields = fields[fields['PRIORITY'].isin(['A','C'])]
field_values = pd.read_csv('contract_field_values.csv', encoding='unicode_escape', skipinitialspace=True)
fields = fields.drop_duplicates(subset='Field Name', keep="first").sort_values('Field Name')
fields = fields[~fields['Field Name'].isnull()]
field_prompt_mapping = dict(zip(fields['Field Name'], fields['Interrogation Question?']))
contract_count_row = st.columns([0.15, 0.45, 0.4])
with contract_count_row[0]:
st.write("**# of Contracts**")
with contract_count_row[1]:
contract_count = st.selectbox('Contract count',('10', '20', '30', '50', 'All'), index=0, label_visibility = "collapsed")
field_row = st.columns([0.15, 0.45, 0.4])
with field_row[0]:
st.write("**Field Name**")
with field_row[1]:
field = st.selectbox('Field Name',sorted(set(field_prompt_mapping.keys())), index=0, label_visibility = "collapsed")
seed_row = st.columns([0.15, 0.45, 0.4])
with seed_row[0]:
st.write("**Seed Value**")
with seed_row[1]:
seed_value = st.text_input("**Seed Value**", value = 20, label_visibility = "collapsed")
contract_count_row = st.columns([0.15, 0.45, 0.4])
with contract_count_row[0]:
st.write("**# of Contracts**")
with contract_count_row[1]:
contract_count = st.selectbox('Contract count',('1', '10', '20', '30', '50', 'All'), index=1, label_visibility = "collapsed")
llm_row = st.columns([0.15, 0.45, 0.4])
with llm_row[0]:
st.write("**Langauge Model**")
with llm_row[1]:
llm_selected = st.selectbox('Langauge Model',('Claude 2', 'Claude Instant', 'Llama 2 Chat 13B', 'Llama 2 Chat 70B'
, 'Titan Text Express'), label_visibility = "collapsed")
seed_row = st.columns([0.15, 0.45, 0.4])
st.write("**Prompt**")
sequence_input = field_prompt_mapping.get(field)
prompt_row = st.columns([0.8, 0.2])
with prompt_row[1]:
if st.button("Clear Prompt"):
sequence_input = ''
if st.button("Back to default"):
prompt = sequence_input
st.button("Save Prompt")
with prompt_row[0]:
prompt = st.text_area("**Prompt**", sequence_input, height = 150, label_visibility = "collapsed")
random.seed(seed_value)
try:
contract_list = sorted(random.choices(os.listdir(SOURCE_DIRECTORY), k=int(contract_count)))
except:
contract_list = sorted(os.listdir(SOURCE_DIRECTORY))
# to be deleted later
contract_list = [contract for contract in contract_list if contract.replace(' MU','').replace('_MU','').replace('.txt','') in list(field_values['(internal) Document Name'])]
page_list_all = []
for contract in contract_list:
page_list = []
with open(os.path.join(SOURCE_DIRECTORY, contract[:-4]+'.txt'), 'r') as infile:
text = infile.read()
page_count = text.count('Start of Page No. = ')
for page in range(page_count+1):
file_path = "SOURCE_DOCUMENTS\\" + f'{contract[:-4]}_page{page}.txt'
dict_with_pages = { 'source': { '$eq': file_path }}
page_list.append(dict_with_pages)
page_list_all.append(page_list)
contract_txt_mapping = dict(zip(contract_list, page_list_all))
column_name = fields.loc[fields['Field Name'] == field, 'SF_DB_COL_NAME'].iloc[0]
column_list = ['(internal) Document Name', '(Internal) Carveout ID', column_name]
if column_name+'_PG' in list(field_values.columns):
column_list.append(column_name+'_PG')
field_values = field_values[column_list]
field_values.rename(columns={'(internal) Document Name': 'Contract Name', column_name: 'Actual Value Stored'
, '(Internal) Carveout ID': 'Contract ID', column_name+'_PG': 'Original Page Number'}, inplace=True)
field_values = field_values.drop_duplicates(subset='Contract Name', keep="first").sort_values('Contract Name')
# sample = dict(zip(sample.Filename, sample.Answer))
# actual_value_list = [sample.get(contract.rsplit('.',1)[0]+'.txt', ' ') for contract in contract_list]
# AWS_ACCESS_KEY_ID = os.getenv('AWS_ACCESS_KEY_ID')
# AWS_SECRET_ACCESS_KEY = os.getenv('AWS_SECRET_ACCESS_KEY')
# AWS_SESSION_TOKEN=os.getenv('AWS_SESSION_TOKEN')
# Setup bedrock
bedrock_runtime = boto3.client(
service_name="bedrock-runtime",
region_name="us-east-1",
)
# Define the retreiver
# load the vectorstore
if "EMBEDDINGS" not in st.session_state:
EMBEDDINGS = BedrockEmbeddings(
client=bedrock_runtime,
model_id="amazon.titan-embed-text-v1",
)
st.session_state.EMBEDDINGS = EMBEDDINGS
if "DB" not in st.session_state:
DB = Chroma(
persist_directory=PERSIST_DIRECTORY,
embedding_function=st.session_state.EMBEDDINGS,
client_settings=CHROMA_SETTINGS,
)
st.session_state.DB = DB
# if "RETRIEVER" not in st.session_state:
# # { "source": { '$eq': "SOURCE_DOCUMENTS\\A.1_UH_Health_System_eff_2_1_08 (1)_page0.txt"} }
# RETRIEVER = DB.as_retriever(search_kwargs={"filter": { "source": { '$eq': "SOURCE_DOCUMENTS\\A.1_UH_Health_System_eff_2_1_08 (1)_page0.txt"} }, "k": 2})
# st.session_state.RETRIEVER = RETRIEVER
# if "LLM" not in st.session_state:
if llm_selected == 'Titan Text Express':
LLM = Bedrock(
model_id="amazon.titan-text-express-v1",
client=bedrock_runtime,
model_kwargs={
"maxTokenCount": 512,
"stopSequences": [],
"temperature": 0,
"topP": 1,
}
)
elif llm_selected == 'Llama 2 Chat 70B':
LLM = Bedrock(
model_id="meta.llama2-70b-chat-v1",
client=bedrock_runtime,
model_kwargs={
"max_gen_len": 512,
"temperature": 0,
# "topP": 0.9,
}
)
elif llm_selected == 'Llama 2 Chat 13B':
LLM = Bedrock(
model_id="meta.llama2-13b-chat-v1",
client=bedrock_runtime,
model_kwargs={
"max_gen_len": 512,
"temperature": 0,
# "topP": 0.9,
}
)
elif llm_selected == 'Claude Instant':
LLM = Bedrock(
model_id="anthropic.claude-instant-v1",
client=bedrock_runtime,
model_kwargs={
# "max_tokens_to_sample": 512,
"temperature": 0,
# "topP": 0.9,
}
)
elif llm_selected == 'Claude 2':
LLM = Bedrock(
model_id="anthropic.claude-v2:1",
client=bedrock_runtime,
model_kwargs={
# "max_tokens_to_sample": 512,
"temperature": 0,
# "topP": 0.9,
}
)
st.session_state["LLM"] = LLM
# if "QA" not in st.session_state:
# prompt, memory = model_memory()
# QA = RetrievalQA.from_chain_type(
# llm=LLM,
# chain_type="stuff",
# retriever=RETRIEVER,
# return_source_documents=True,
# chain_type_kwargs={"prompt": prompt, "memory": memory},
# )
# st.session_state["QA"] = QA
# df = pd.DataFrame(columns=['Contract Name','Contract ID','Actual Value Stored','New Extracted value','Confidence Level','Snippet',
# 'Original Page Number', 'New Page Number', 'Revised Prompt', 'Result'])
df = pd.DataFrame(columns=['Contract Name','Raw value','New Extracted value','Confidence Level','Snippet','New Page Number'
, 'Revised Prompt', 'Result'])
try:
history = pd.read_csv('history.csv')
except:
history = pd.DataFrame(columns=['Field Name','# Contracts Tested', 'Username', 'Date/Time', 'Accuracy', 'Attempt #'])
attempt = 0
if llm_selected in ['Llama 2 Chat 13B', 'Llama 2 Chat 70B']:
k_value = 14
else:
k_value = 25
if st.button("Test Configuration"):
answer_list = []
doc_list = []
response_list = []
score_list = []
attempt = attempt + 1
for page_list in page_list_all:
RETRIEVER = st.session_state.DB.as_retriever(search_kwargs={"filter": {"$or": page_list}, "k": k_value})
QA = RetrievalQA.from_chain_type(
llm=st.session_state["LLM"],
chain_type="stuff",
retriever=RETRIEVER,
return_source_documents=True,
# chain_type_kwargs={"prompt": prompt, "memory": None},
)
score = st.session_state.DB.similarity_search_with_relevance_scores(prompt, k=4, filter={"$or": page_list})
score_list.append(max(d[1] for d in score))
response = QA(prompt)
answer, docs = response["result"], response["source_documents"]
answer_list.append(answer)
doc_list.append(docs)
response_list.append(response)
df['Raw value'] = answer_list
# post-processing
if 'Date' in field:
date_list = []
for answer in answer_list:
try:
extracted_date = dateutil.parser.parse(str(answer).replace('"',''), fuzzy=True).date()
except:
extracted_date = " "
date_list.append(extracted_date)
answer_list = date_list
elif llm_selected in ['Llama 2 Chat 13B', 'Llama 2 Chat 70B']:
answer_list = [answer.rstrip(".") for answer in answer_list]
answer_list = [answer if "I don't know" not in str(answer) else " " for answer in answer_list]
answer_list = [answer if "N/A" not in str(answer) else " " for answer in answer_list]
answer_list = [answer if "does not contain" not in str(answer) else " " for answer in answer_list]
answer_list = [answer if "None" not in str(answer) else " " for answer in answer_list]
answer_list = [answer if "Not specified in the contract" not in str(answer) else " " for answer in answer_list]
answer_list = [answer if "Not applicable" not in str(answer) else " " for answer in answer_list]
elif llm_selected in ['Claude 2', 'Claude Instant']:
answer_list = [answer if "Unfortunately, I do not have enough context" not in str(answer) else " " for answer in answer_list]
answer_list = [answer.rstrip(".") for answer in answer_list]
else:
answer_list = [answer.rstrip(".") for answer in answer_list]
# answer_list = [str(x).rsplit(':',1)[0] if len(str(x).rsplit(':',1)) < 2 else str(x).rsplit(':',1)[1] for x in answer_list]
df['Contract Name'] = contract_list
# to be deleted later
df['Contract Name'] = [contract.replace(' MU','').replace('_MU','').replace('.txt','') for contract in contract_list]
contract_list = [contract for contract in contract_list if contract.replace(' MU','').replace('_MU','').replace('.txt','') in list(field_values['(internal) Document Name'])]
df['New Extracted value'] = answer_list
df['Confidence Level'] = [round(score, 2) for score in score_list]
df['Snippet'] = [str(doc[0].page_content) for doc in doc_list]
df['New Page Number'] = [int(str(doc[0].metadata["source"]).rsplit('_page')[1].replace('.txt',''))+1 for doc in doc_list]
df['Revised Prompt'] = [prompt] * len(contract_list)
with seed_row[0]:
if contract_count in ['10', '20', '30', '50']:
st.write("**Seed Value**")
elif contract_count == '1':
st.write("**Contract Name**")
with seed_row[1]:
if contract_count in ['10', '20', '30', '50']:
seed_value = st.text_input("**Seed Value**", value = 20, label_visibility = "collapsed")
random.seed(seed_value)
contract_list = sorted(random.choices(os.listdir(SOURCE_DIRECTORY), k=int(contract_count)))
elif contract_count == '1':
contract_name = st.selectbox('Contract Name', (contract_list), label_visibility = "collapsed")
contract_list = [contract_name]
df = pd.merge(df, field_values, how ='left', on ='Contract Name')
answer_list = list(df['New Extracted value'])
actual_value_list = list(df['Actual Value Stored'])
result_list = [i==j for i, j in zip(actual_value_list, answer_list)]
df['Result'] = [str(x) for x in result_list]
df = df[~df['Contract ID'].isnull()]
if 'Original Page Number' in df.columns:
df = df[['Contract Name','Contract ID','Actual Value Stored','Raw value','New Extracted value','Confidence Level','Snippet'
,'Original Page Number', 'New Page Number', 'Revised Prompt', 'Result']]
else:
df = df[['Contract Name','Contract ID','Actual Value Stored','Raw value','New Extracted value','Confidence Level','Snippet'
, 'New Page Number', 'Revised Prompt', 'Result']]
llm_row = st.columns([0.15, 0.45, 0.4])
with llm_row[0]:
st.write("**Langauge Model**")
with llm_row[1]:
llm_selected = st.selectbox('Langauge Model',('Claude 2', 'Claude Instant', 'Llama 2 Chat 13B', 'Llama 2 Chat 70B'
, 'Titan Text Express'), label_visibility = "collapsed")
st.write("**Prompt**")
sequence_input = field_prompt_mapping.get(field)
prompt_row = st.columns([0.8, 0.2])
with prompt_row[1]:
if st.button("Clear Prompt"):
sequence_input = ''
if st.button("Back to default"):
prompt = sequence_input
st.button("Save Prompt")
with prompt_row[0]:
prompt = st.text_area("**Prompt**", sequence_input, height = 150, label_visibility = "collapsed")
page_list_all = []
for contract in contract_list:
page_list = []
with open(os.path.join(SOURCE_DIRECTORY, contract[:-4]+'.txt'), 'r') as infile:
text = infile.read()
page_count = text.count('Start of Page No. = ')
for page in range(page_count+1):
file_path = "SOURCE_DOCUMENTS\\" + f'{contract[:-4]}_page{page}.txt'
dict_with_pages = { 'source': { '$eq': file_path }}
page_list.append(dict_with_pages)
page_list_all.append(page_list)
contract_txt_mapping = dict(zip(contract_list, page_list_all))
column_name = fields.loc[fields['Field Name'] == field, 'SF_DB_COL_NAME'].iloc[0]
column_list = ['(internal) Document Name', '(Internal) Carveout ID', column_name]
if column_name+'_PG' in list(field_values.columns):
column_list.append(column_name+'_PG')
field_values = field_values[column_list]
field_values.rename(columns={'(internal) Document Name': 'Contract Name', column_name: 'Actual Value Stored'
, '(Internal) Carveout ID': 'Contract ID', column_name+'_PG': 'Original Page Number'}, inplace=True)
field_values = field_values.drop_duplicates(subset='Contract Name', keep="first").sort_values('Contract Name')
# Setup bedrock
bedrock_runtime = boto3.client(
service_name="bedrock-runtime",
region_name="us-east-1",
)
# Define the retreiver
# load the vectorstore
if "EMBEDDINGS" not in st.session_state:
EMBEDDINGS = BedrockEmbeddings(
client=bedrock_runtime,
model_id="amazon.titan-embed-text-v1",
)
st.session_state.EMBEDDINGS = EMBEDDINGS
if "DB" not in st.session_state:
DB = Chroma(
persist_directory=PERSIST_DIRECTORY,
embedding_function=st.session_state.EMBEDDINGS,
client_settings=CHROMA_SETTINGS,
)
st.session_state.DB = DB
# if "RETRIEVER" not in st.session_state:
# # { "source": { '$eq': "SOURCE_DOCUMENTS\\A.1_UH_Health_System_eff_2_1_08 (1)_page0.txt"} }
# RETRIEVER = DB.as_retriever(search_kwargs={"filter": { "source": { '$eq': "SOURCE_DOCUMENTS\\A.1_UH_Health_System_eff_2_1_08 (1)_page0.txt"} }, "k": 2})
# st.session_state.RETRIEVER = RETRIEVER
# if "LLM" not in st.session_state:
if llm_selected == 'Titan Text Express':
LLM = Bedrock(
model_id="amazon.titan-text-express-v1",
client=bedrock_runtime,
model_kwargs={
"maxTokenCount": 512,
"stopSequences": [],
"temperature": 0,
"topP": 1,
}
)
elif llm_selected == 'Llama 2 Chat 70B':
LLM = Bedrock(
model_id="meta.llama2-70b-chat-v1",
client=bedrock_runtime,
model_kwargs={
"max_gen_len": 512,
"temperature": 0,
# "topP": 0.9,
}
)
elif llm_selected == 'Llama 2 Chat 13B':
LLM = Bedrock(
model_id="meta.llama2-13b-chat-v1",
client=bedrock_runtime,
model_kwargs={
"max_gen_len": 512,
"temperature": 0,
# "topP": 0.9,
}
)
elif llm_selected == 'Claude Instant':
LLM = Bedrock(
model_id="anthropic.claude-instant-v1",
client=bedrock_runtime,
model_kwargs={
# "max_tokens_to_sample": 512,
"temperature": 0,
# "topP": 0.9,
}
)
elif llm_selected == 'Claude 2':
LLM = Bedrock(
model_id="anthropic.claude-v2:1",
client=bedrock_runtime,
model_kwargs={
# "max_tokens_to_sample": 512,
"temperature": 0,
# "topP": 0.9,
}
)
st.session_state["LLM"] = LLM
# if "QA" not in st.session_state:
# prompt, memory = model_memory()
# QA = RetrievalQA.from_chain_type(
# llm=LLM,
# chain_type="stuff",
# retriever=RETRIEVER,
# return_source_documents=True,
# chain_type_kwargs={"prompt": prompt, "memory": memory},
# )
# st.session_state["QA"] = QA
# df = pd.DataFrame(columns=['Contract Name','Contract ID','Actual Value Stored','New Extracted value','Confidence Level','Snippet',
# 'Original Page Number', 'New Page Number', 'Revised Prompt', 'Result'])
df = pd.DataFrame(columns=['Contract Name','Raw value','New Extracted value','Confidence Level','Snippet1','Snippet2','Snippet3'
,'Snippet4','Snippet5','Snippet6','Snippet7','Snippet8','Snippet9','Snippet10','New Page Number'
, 'Revised Prompt', 'Result'])
try:
accuracy = round(sum(bool(x) for x in result_list) * 100 / len(list(df['Result'])), 2)
history = pd.read_csv('history.csv')
except:
accuracy = 'NA'
history = pd.DataFrame(columns=['Field Name','# Contracts Tested', 'Username', 'Date/Time', 'Accuracy', 'Attempt #'])
attempt = 0
history.loc[len(history.index)] = [field, str(contract_count), None, datetime.now().strftime("%Y-%m-%d %H:%M:%S"), accuracy, attempt]
df.to_csv("RESULTS\\"+field.replace("?","").replace("/","_")+'-'+llm_selected+'.csv', index=False)
history.to_csv('history.csv', index=False)
# df_copy = df.set_index(df.columns[0]).copy()
# df_2_copy = history.set_index(history.columns[0]).copy()
st.dataframe(df)
st.dataframe(history)
if llm_selected in ['Llama 2 Chat 13B', 'Llama 2 Chat 70B']:
k_value = 10
else:
k_value = 20
# @st.cache_data
# def convert_df(df):
# return df.to_csv(index=False).encode('utf-8')
if st.button("Test Configuration"):
# csv = convert_df(edited_df)
answer_list = []
doc_list = []
response_list = []
score_list = []
attempt = attempt + 1
# buttons = st.columns(3)
# with buttons[0]:
# st.button("Save All Imputations")
# with buttons[1]:
# st.download_button("Download Table", csv, "file.csv", "text/csv", key='download-csv')
# with buttons[2]:
# st.button("Kickoff Database Integration")
for page_list in page_list_all:
RETRIEVER = st.session_state.DB.as_retriever(search_kwargs={"filter": {"$or": page_list}, "k": k_value})
QA = RetrievalQA.from_chain_type(
llm=st.session_state["LLM"],
chain_type="stuff",
retriever=RETRIEVER,
return_source_documents=True,
# chain_type_kwargs={"prompt": prompt, "memory": None},
)
score = st.session_state.DB.similarity_search_with_relevance_scores(prompt, k=4, filter={"$or": page_list})
score_list.append(max(d[1] for d in score))
response = QA(prompt)
answer, docs = response["result"], response["source_documents"]
answer_list.append(answer)
doc_list.append(docs)
response_list.append(response)
df['Raw value'] = answer_list
# post-processing
if 'Date' in field:
date_list = []
for answer in answer_list:
try:
extracted_date = dateutil.parser.parse(str(answer).replace('"',''), fuzzy=True).date()
except:
extracted_date = " "
date_list.append(extracted_date)
answer_list = date_list
elif llm_selected in ['Llama 2 Chat 13B', 'Llama 2 Chat 70B']:
answer_list = [answer.rstrip(".") for answer in answer_list]
answer_list = [answer if "I don't know" not in str(answer) else " " for answer in answer_list]
answer_list = [answer if "N/A" not in str(answer) else " " for answer in answer_list]
answer_list = [answer if "does not contain" not in str(answer) else " " for answer in answer_list]
answer_list = [answer if "None" not in str(answer) else " " for answer in answer_list]
answer_list = [answer if "Not specified in the contract" not in str(answer) else " " for answer in answer_list]
answer_list = [answer if "Not applicable" not in str(answer) else " " for answer in answer_list]
elif llm_selected in ['Claude 2', 'Claude Instant']:
answer_list = [answer if "Unfortunately, I do not have enough context" not in str(answer) else " " for answer in answer_list]
answer_list = [answer.rstrip(".") for answer in answer_list]
else:
answer_list = [answer.rstrip(".") for answer in answer_list]
# answer_list = [str(x).rsplit(':',1)[0] if len(str(x).rsplit(':',1)) < 2 else str(x).rsplit(':',1)[1] for x in answer_list]
df['Contract Name'] = contract_list
# to be deleted later
df['Contract Name'] = [contract.replace(' MU','').replace('_MU','').replace('.txt','') for contract in contract_list]
df['New Extracted value'] = answer_list
df['Confidence Level'] = [round(score, 2) for score in score_list]
Snippet = []
count = 0
for i in range(int(contract_count)):
for j in range(10):
try:
content = str(doc_list[i][j].page_content)
except:
content = " "
Snippet.append(content)
# df['Snippet1'] = [str(doc[0].page_content) for doc in doc_list]
df['Snippet1'] = Snippet[:int(contract_count)]
df['Snippet2'] = Snippet[int(contract_count):2*int(contract_count)]
df['Snippet3'] = Snippet[2*int(contract_count):3*int(contract_count)]
df['Snippet4'] = Snippet[3*int(contract_count):4*int(contract_count)]
df['Snippet5'] = Snippet[4*int(contract_count):5*int(contract_count)]
df['Snippet6'] = Snippet[5*int(contract_count):6*int(contract_count)]
df['Snippet7'] = Snippet[6*int(contract_count):7*int(contract_count)]
df['Snippet8'] = Snippet[7*int(contract_count):8*int(contract_count)]
df['Snippet9'] = Snippet[8*int(contract_count):9*int(contract_count)]
df['Snippet10'] = Snippet[9*int(contract_count):]
df['New Page Number'] = [int(str(doc[0].metadata["source"]).rsplit('_page')[1].replace('.txt',''))+1 for doc in doc_list]
df['Revised Prompt'] = [prompt] * len(contract_list)
df = pd.merge(df, field_values, how ='left', on ='Contract Name')
answer_list = list(df['New Extracted value'])
df['Actual Value Stored'] = pd.to_datetime(df['Actual Value Stored'],errors='coerce').dt.date
df.fillna(" ", inplace=True)
actual_value_list = list(df['Actual Value Stored'])
result_list = [i==j for i, j in zip(actual_value_list, answer_list)]
df['Result'] = [str(x) for x in result_list]
df = df[~df['Contract ID'].isnull()]
if 'Original Page Number' in df.columns:
df = df[['Contract Name','Contract ID','Actual Value Stored','Raw value','New Extracted value','Confidence Level'
,'Snippet1','Snippet2','Snippet3','Snippet4','Snippet5','Snippet6','Snippet7','Snippet8','Snippet9','Snippet10'
,'Original Page Number', 'New Page Number', 'Revised Prompt', 'Result']]
else:
df = df[['Contract Name','Contract ID','Actual Value Stored','Raw value','New Extracted value','Confidence Level'
,'Snippet1','Snippet2','Snippet3','Snippet4','Snippet5','Snippet6','Snippet7','Snippet8','Snippet9','Snippet10'
, 'New Page Number', 'Revised Prompt', 'Result']]
try:
accuracy = round(sum(bool(x) for x in result_list) * 100 / len(list(df['Result'])), 2)
except:
accuracy = 'NA'
history.loc[len(history.index)] = [field, str(contract_count), None, datetime.now().strftime("%Y-%m-%d %H:%M:%S"), accuracy, attempt]
# df.to_csv("RESULTS\\"+field.replace("?","").replace("/","_")+'-'+llm_selected+'.csv', index=False)
history.to_csv('history.csv', index=False)
# df_copy = df.set_index(df.columns[0]).copy()
# df_2_copy = history.set_index(history.columns[0]).copy()
st.dataframe(df)
st.dataframe(history)
st.write(column_name)
st.write(len(contract_list))
# @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")
# with buttons[1]:
# st.download_button("Download Table", csv, "file.csv", "text/csv", key='download-csv')
# with buttons[2]:
# st.button("Kickoff Database Integration")
st.write(column_name)
st.write(len(contract_list))
else:
st.write("Access Denied")
+41
View File
@@ -0,0 +1,41 @@
import streamlit as st
import msal
import requests
# Replace with your own values
CLIENT_ID = 'effafe90-7ed7-43a3-ab03-19a0be2f1758'
CLIENT_SECRET = 'bjQ8Q~lpR2uBcGI34VDu16t73doz8Crj0YY_~dgD'
# TENANT_ID = ''
AUTHORITY = 'https://login.microsoftonline.com/organizations/'
SCOPE = ['User.Read']
# REDIRECT_URI = 'http://localhost:8501'
app = msal.ConfidentialClientApplication(CLIENT_ID, authority=AUTHORITY, client_credential=CLIENT_SECRET)
def get_auth_url(REDIRECT_URI):
auth_url = app.get_authorization_request_url(SCOPE, redirect_uri=REDIRECT_URI)
return auth_url
def get_token_from_code(auth_code, REDIRECT_URI):
app = msal.ConfidentialClientApplication(CLIENT_ID, authority=AUTHORITY, client_credential=CLIENT_SECRET)
result = app.acquire_token_by_authorization_code(auth_code, scopes=SCOPE, redirect_uri=REDIRECT_URI)
return result['access_token']
def get_user_info(access_token):
headers = {'Authorization': f'Bearer {access_token}'}
response = requests.get('https://graph.microsoft.com/v1.0/me', headers=headers)
return response.json()
def handle_redirect(REDIRECT_URI):
if not st.session_state.get('access_token'):
code = st.query_params.get('code')
if code:
access_token = get_token_from_code(code, REDIRECT_URI)
st.session_state['access_token'] = access_token
st.session_state
+24
View File
@@ -0,0 +1,24 @@
import streamlit as st
import security
def setup_page(REDIRECT_URI):
# st.set_page_config(
# page_title=page_title,
# page_icon="👋",
# )
if st.query_params.get('code'):
security.handle_redirect(REDIRECT_URI)
access_token = st.session_state.get('access_token')
if access_token:
user_info = security.get_user_info(access_token)
st.session_state['user_info'] = user_info
return True
else:
st.write("Please sign-in to use this app.")
auth_url = security.get_auth_url(REDIRECT_URI)
st.markdown(f"<a href='{auth_url}' target='_self'>Sign In</a>", unsafe_allow_html=True)
st.stop()