diff --git a/streamlit/interface_3.py b/streamlit/interface_3.py index a5f4a0a..ec49dca 100644 --- a/streamlit/interface_3.py +++ b/streamlit/interface_3.py @@ -22,26 +22,6 @@ import re import snowflake.connector from sf_conn import get_secret -try: - sf_secrets = get_secret() - conn = snowflake.connector.connect( - user=sf_secrets['user'], - password=sf_secrets['password'], - account="aarete-doczyai", - role = sf_secrets['ROLE'], - warehouse=sf_secrets['warehouse'], - database=sf_secrets['database'], - schema="STG" - ) - cur = conn.cursor() - - query = 'select * from "TRAINING_DATA_RAW"' - cur.execute(query) - - df = pd.DataFrame(cur.fetchall()) -except: - print("conn failed") - REDIRECT_URI = 'https://doczy.aarete.com:8503' user_list = ['maamseek@aarete.com', 'smahdavian@aarete.com', 'ahinge@aarete.com', 'akadam@aarete.com' @@ -70,6 +50,45 @@ try: except: user_mail = 'maamseek@aarete.com' +try: + sf_secrets = get_secret() + conn = snowflake.connector.connect( + user=sf_secrets['user'], + password=sf_secrets['password'], + account="aarete-doczyai", + role = sf_secrets['ROLE'], + warehouse=sf_secrets['warehouse'], + database=sf_secrets['database'], + schema="STG" + ) + cur = conn.cursor() + + query = 'select * from "TRAINING_DATA_RAW"' + cur.execute(query) + field_values = pd.DataFrame(cur.fetchall()) + field_values['Contract ID'] = field_values['Document_Name'] + +except: + field_values = pd.read_csv('contract_field_values.csv', encoding='unicode_escape', skipinitialspace=True) + field_values.rename(columns={'(internal) Document Name': 'Document_Name'}, inplace = True) + field_values.rename(columns={'(Internal) Carveout ID': 'Contract ID'}, inplace = True) + st.write("conn failed") + +try: + query = 'select * from "PROMPT_CONFIG"' + cur.execute(query) + fields = pd.DataFrame(cur.fetchall()) + field_values.rename(columns={'FIELD_DESC': 'Field Name'}, inplace = True) + field_values.rename(columns={'PROMPT': 'Interrogation Question?'}, inplace = True) + field_values.rename(columns={'GROUP_ID': 'PRIORITY'}, inplace = True) + field_values.rename(columns={'FIELD_NAME': 'SF_DB_COL_NAME'}, inplace = True) + field_values.rename(columns={'FM_MODEL_ID': 'llm_selected'}, inplace = True) + error('table is empty') +except: + fields = pd.read_csv('contract_fields.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()] + if user_mail in user_list: @@ -83,11 +102,8 @@ if user_mail in user_list: , 'Optimize Carving Indic.', 'Carveout Method - I', 'Carveout Method - II', 'Provider' , 'Timeline'), label_visibility = "collapsed") - fields = pd.read_csv('contract_fields.csv', encoding='unicode_escape', skipinitialspace=True) - 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()] + # priorty column will be relaced by group_id in snowflake db if field_group == 'Unique Key': fields = fields[fields['PRIORITY'] == 'A'] elif field_group == 'Contract Related': @@ -145,7 +161,7 @@ if user_mail in user_list: 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'])] + contract_list = [contract for contract in contract_list if contract.replace(' MU','').replace('_MU','').replace('.txt','') in list(field_values['Document_Name'])] with seed_row[0]: if contract_count in ['10', '20', '30', '50']: @@ -183,12 +199,12 @@ if user_mail in user_list: column_name = fields.loc[fields['Field Name'] == field, 'SF_DB_COL_NAME'].iloc[0] - column_list = ['(internal) Document Name', '(Internal) Carveout ID', column_name] + column_list = ['Document_Name', 'Contract 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.rename(columns={'Document_Name': 'Contract Name', column_name: 'Actual Value Stored' + , 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