diff --git a/.gitignore b/.gitignore index b3393a2..dd97710 100644 --- a/.gitignore +++ b/.gitignore @@ -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 diff --git a/streamlit/interface_1.py b/streamlit/interface_1.py index 400adb0..334959f 100644 --- a/streamlit/interface_1.py +++ b/streamlit/interface_1.py @@ -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") diff --git a/streamlit/interface_2.py b/streamlit/interface_2.py index 9a27c78..8ccbfa2 100644 --- a/streamlit/interface_2.py +++ b/streamlit/interface_2.py @@ -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") diff --git a/streamlit/interface_3.py b/streamlit/interface_3.py index 2cb2b06..aa49067 100644 --- a/streamlit/interface_3.py +++ b/streamlit/interface_3.py @@ -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") diff --git a/streamlit/security.py b/streamlit/security.py new file mode 100644 index 0000000..1337f58 --- /dev/null +++ b/streamlit/security.py @@ -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 \ No newline at end of file diff --git a/streamlit/util.py b/streamlit/util.py new file mode 100644 index 0000000..7e4565a --- /dev/null +++ b/streamlit/util.py @@ -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"Sign In", unsafe_allow_html=True) + st.stop() \ No newline at end of file