diff --git a/Doczy.AI_Automation/config/aws_conn.py b/Doczy.AI_Automation/config/aws_conn.py index 6fe9589..23e4290 100644 --- a/Doczy.AI_Automation/config/aws_conn.py +++ b/Doczy.AI_Automation/config/aws_conn.py @@ -4,32 +4,31 @@ import snowflake.connector def get_files(s3_client, path, client_bucket, batch_id): - batch_objects = s3_client.list_objects_v2(Bucket=client_bucket, - Prefix=path+"/"+batch_id, - Delimiter='/') + batch_objects = s3_client.list_objects_v2( + Bucket=client_bucket, Prefix=path + "/" + batch_id, Delimiter="/" + ) file_list = [] - for prefix in batch_objects['CommonPrefixes']: - file_list.append(prefix['Prefix'][:-1].split('/')[-1]) + for prefix in batch_objects["CommonPrefixes"]: + file_list.append(prefix["Prefix"][:-1].split("/")[-1]) file_count = len(file_list) return file_count, file_list def get_s3_client(): region_name = "us-east-2" - s3_client = boto3.client('s3', - region_name=region_name) + s3_client = boto3.client("s3", region_name=region_name) return s3_client def get_snowflake_conn(schema): - secret = ''#get_secret() + secret = "" # get_secret() secret_dict = eval(secret) conn = snowflake.connector.connect( - user=secret_dict['user'], - password=secret_dict['password'], - account=secret_dict['account_alias'], - warehouse=secret_dict['warehouse'], - database=secret_dict['database'], - schema=schema + user=secret_dict["user"], + password=secret_dict["password"], + account=secret_dict["account_alias"], + warehouse=secret_dict["warehouse"], + database=secret_dict["database"], + schema=schema, ) - return conn \ No newline at end of file + return conn diff --git a/Doczy.AI_Automation/elements/interface_0_elements.py b/Doczy.AI_Automation/elements/interface_0_elements.py index e8ee273..433075a 100644 --- a/Doczy.AI_Automation/elements/interface_0_elements.py +++ b/Doczy.AI_Automation/elements/interface_0_elements.py @@ -1,19 +1,29 @@ from selenium.webdriver.common.by import By -from utils.element_related_methods import get_clickable_element, get_invisible_element, get_element +from utils.element_related_methods import ( + get_clickable_element, + get_invisible_element, + get_element, +) class Interface_0_elements: client_name_dropdown_xpath = '//div[@data-testid="stSelectbox"]' - client_list_xpath = '//li' + client_list_xpath = "//li" file_input_xpath = "//input[@type='file']" - file_input_progress_xpath = '//div[@role="progressbar"]'#'//div[@class="stFileUploaderFileName"]' - create_batch_btn_xpath = "//div[@data-testid='stButton']/button[@data-testid='baseButton-secondary']" + file_input_progress_xpath = ( + '//div[@role="progressbar"]' #'//div[@class="stFileUploaderFileName"]' + ) + create_batch_btn_xpath = ( + "//div[@data-testid='stButton']/button[@data-testid='baseButton-secondary']" + ) p_element_xpath = "//div[@data-testid='stCodeBlock']/pre/div/code/span" error_xpath = "//div[@data-testid='stNotificationContentError']/div/div/p" def get_client_name_dropdown(self, driver, config): - return get_clickable_element(driver, config, By.XPATH, self.client_name_dropdown_xpath) + return get_clickable_element( + driver, config, By.XPATH, self.client_name_dropdown_xpath + ) def get_client_list(self, driver): return driver.find_elements(By.XPATH, self.client_list_xpath) @@ -25,7 +35,9 @@ class Interface_0_elements: return driver.find_elements(By.XPATH, self.file_input_progress_xpath) def get_create_batch_btn(self, driver, config): - return get_clickable_element(driver, config, By.XPATH, self.create_batch_btn_xpath) + return get_clickable_element( + driver, config, By.XPATH, self.create_batch_btn_xpath + ) def get_p_element(self, driver, config): return get_element(driver, config, By.XPATH, self.p_element_xpath) diff --git a/Doczy.AI_Automation/elements/interface_1_elements.py b/Doczy.AI_Automation/elements/interface_1_elements.py index 8f58a9a..2624305 100644 --- a/Doczy.AI_Automation/elements/interface_1_elements.py +++ b/Doczy.AI_Automation/elements/interface_1_elements.py @@ -1,22 +1,30 @@ from selenium.webdriver.common.by import By -from utils.element_related_methods import get_clickable_element, get_invisible_element, get_element +from utils.element_related_methods import ( + get_clickable_element, + get_invisible_element, + get_element, +) class Interface_1_elements: client_name_dropdown_xpath = '//div[@data-testid="stSelectbox"]' - client_list_xpath = '//li' + client_list_xpath = "//li" batch_id_xpath = '(//div[@data-testid="stSelectbox"])[2]' - batch_id_list_xpath = '//li' + batch_id_list_xpath = "//li" checkbox_list_xpath = "//div[@data-testid='stCheckbox']" - read_the_contracts_from_path_btn_xpath = "//button[.//p[text()='Read the contracts from Path']]" + read_the_contracts_from_path_btn_xpath = ( + "//button[.//p[text()='Read the contracts from Path']]" + ) result_table_xpath = "//table[@aria-colcount='9']" table_rows_xpath = ".//tbody/tr" run_doczy_ai_pipeline_btn_xpath = "//button[.//p[text()='Run Doczy.AI Pipeline']]" error_xpath = "//div[@data-testid='stNotificationContentError']/div/div/p" def get_client_name_dropdown(self, driver, config): - return get_clickable_element(driver, config, By.XPATH, self.client_name_dropdown_xpath) + return get_clickable_element( + driver, config, By.XPATH, self.client_name_dropdown_xpath + ) def get_client_list(self, driver): return driver.find_elements(By.XPATH, self.client_list_xpath) @@ -31,7 +39,9 @@ class Interface_1_elements: return driver.find_elements(By.XPATH, self.checkbox_list_xpath) def get_read_the_contracts_from_path_btn(self, driver, config): - return get_clickable_element(driver, config, By.XPATH, self.read_the_contracts_from_path_btn_xpath) + return get_clickable_element( + driver, config, By.XPATH, self.read_the_contracts_from_path_btn_xpath + ) def get_result_table(self, driver, config): return get_element(driver, config, By.XPATH, self.result_table_xpath) @@ -40,7 +50,9 @@ class Interface_1_elements: return result_table.find_elements(By.XPATH, self.table_rows_xpath) def get_run_doczy_ai_pipeline_btn(self, driver, config): - return get_clickable_element(driver, config, By.XPATH, self.run_doczy_ai_pipeline_btn_xpath) + return get_clickable_element( + driver, config, By.XPATH, self.run_doczy_ai_pipeline_btn_xpath + ) def get_error(self, driver, config): return get_element(driver, config, By.XPATH, self.error_xpath) diff --git a/Doczy.AI_Automation/elements/interface_2_elements.py b/Doczy.AI_Automation/elements/interface_2_elements.py index 9ae454e..a1ad3d4 100644 --- a/Doczy.AI_Automation/elements/interface_2_elements.py +++ b/Doczy.AI_Automation/elements/interface_2_elements.py @@ -1,17 +1,21 @@ from selenium.webdriver.common.by import By -from utils.element_related_methods import get_clickable_element, get_invisible_element, get_element +from utils.element_related_methods import ( + get_clickable_element, + get_invisible_element, + get_element, +) class Interface_2_elements: client_name_dropdown_xpath = '//div[@data-testid="stSelectbox"]' - client_list_xpath = '//li' + client_list_xpath = "//li" batch_id_xpath = '(//div[@data-testid="stSelectbox"])[2]' - batch_id_list_xpath = '//li' + batch_id_list_xpath = "//li" contract_name_xpath = '(//div[@data-testid="stSelectbox"])[3]' - contract_name_list_xpath = '//li' + contract_name_list_xpath = "//li" field_group_xpath = '(//div[@data-testid="stSelectbox"])[4]' - field_group_list_xpath = '//li' + field_group_list_xpath = "//li" show_result_btn_xpath = "//button[.//p[text()='Show Results']]" result_table_xpath = "//table[@aria-colcount='9']" table_rows_xpath = ".//tbody/tr" @@ -19,7 +23,9 @@ class Interface_2_elements: error_xpath = "//div[@data-testid='stNotificationContentError']/div/div/p" def get_client_name_dropdown(self, driver, config): - return get_clickable_element(driver, config, By.XPATH, self.client_name_dropdown_xpath) + return get_clickable_element( + driver, config, By.XPATH, self.client_name_dropdown_xpath + ) def get_client_list(self, driver): return driver.find_elements(By.XPATH, self.client_list_xpath) @@ -43,7 +49,9 @@ class Interface_2_elements: return driver.find_elements(By.XPATH, self.field_group_list_xpath) def get_show_result_btn(self, driver, config): - return get_clickable_element(driver, config, By.XPATH, self.show_result_btn_xpath) + return get_clickable_element( + driver, config, By.XPATH, self.show_result_btn_xpath + ) def get_result_table(self, driver, config): return get_element(driver, config, By.XPATH, self.result_table_xpath) @@ -52,7 +60,9 @@ class Interface_2_elements: return result_table.find_elements(By.XPATH, self.table_rows_xpath) def get_kickoff_database_btn(self, driver, config): - return get_clickable_element(driver, config, By.XPATH, self.kickoff_database_btn_xpath) + return get_clickable_element( + driver, config, By.XPATH, self.kickoff_database_btn_xpath + ) def get_error_toast(self, driver, config): return get_element(driver, config, By.XPATH, self.error_xpath) diff --git a/Doczy.AI_Automation/pages/Email_sign_in_page.py b/Doczy.AI_Automation/pages/Email_sign_in_page.py index 27c6bb6..1dc571a 100644 --- a/Doczy.AI_Automation/pages/Email_sign_in_page.py +++ b/Doczy.AI_Automation/pages/Email_sign_in_page.py @@ -5,7 +5,7 @@ from utils.element_related_methods import get_element, get_clickable_element class Email_sign_in_page: def __init__(self, driver): self.driver = driver - self.email = get_element(self.driver, By.NAME, 'loginfmt') + self.email = get_element(self.driver, By.NAME, "loginfmt") self.next_btn = get_clickable_element(self.driver, By.ID, "idSIButton9") def set_email(self, email): @@ -16,5 +16,5 @@ class Email_sign_in_page: self.next_btn.click() def sign_in(self, config): - self.set_email(config['username']) - self.click_next_btn() \ No newline at end of file + self.set_email(config["username"]) + self.click_next_btn() diff --git a/Doczy.AI_Automation/pages/Password_sign_in_page.py b/Doczy.AI_Automation/pages/Password_sign_in_page.py index f92df4c..e7069cc 100644 --- a/Doczy.AI_Automation/pages/Password_sign_in_page.py +++ b/Doczy.AI_Automation/pages/Password_sign_in_page.py @@ -6,8 +6,8 @@ from utils.element_related_methods import get_clickable_element, get_element class Password_sign_in_page: def __init__(self, driver): self.driver = driver - self.password = get_element(self.driver, By.NAME, 'passwd') - self.sign_in_btn = get_clickable_element(self.driver, By.ID, 'idSIButton9') + self.password = get_element(self.driver, By.NAME, "passwd") + self.sign_in_btn = get_clickable_element(self.driver, By.ID, "idSIButton9") def set_password(self, password): self.password.clear() @@ -17,5 +17,5 @@ class Password_sign_in_page: self.sign_in_btn.click() def sign_in(self, config): - self.set_password(config['password']) - self.click_sign_in_btn() \ No newline at end of file + self.set_password(config["password"]) + self.click_sign_in_btn() diff --git a/Doczy.AI_Automation/pages/Pre_sign_in_page.py b/Doczy.AI_Automation/pages/Pre_sign_in_page.py index bc86910..53be9fa 100644 --- a/Doczy.AI_Automation/pages/Pre_sign_in_page.py +++ b/Doczy.AI_Automation/pages/Pre_sign_in_page.py @@ -8,7 +8,7 @@ from utils.element_related_methods import get_clickable_element class Pre_sign_in_page: def __init__(self, driver): self.driver = driver - self.sign_in = get_clickable_element(self.driver, By.TAG_NAME, 'a') + self.sign_in = get_clickable_element(self.driver, By.TAG_NAME, "a") def click_sign_in_link(self): - self.sign_in.click() \ No newline at end of file + self.sign_in.click() diff --git a/Doczy.AI_Automation/pages/Remember_sign_in_page.py b/Doczy.AI_Automation/pages/Remember_sign_in_page.py index d7c9fd2..dd3b1c0 100644 --- a/Doczy.AI_Automation/pages/Remember_sign_in_page.py +++ b/Doczy.AI_Automation/pages/Remember_sign_in_page.py @@ -6,8 +6,8 @@ from utils.element_related_methods import get_clickable_element class Remember_sign_in_page: def __init__(self, driver): self.driver = driver - self.no_btn = get_clickable_element(self.driver, By.ID, 'idBtn_Back') - self.yes_btn = get_clickable_element(self.driver, By.ID, 'idSIButton9') + self.no_btn = get_clickable_element(self.driver, By.ID, "idBtn_Back") + self.yes_btn = get_clickable_element(self.driver, By.ID, "idSIButton9") def click_no_btn(self): - self.no_btn.click() \ No newline at end of file + self.no_btn.click() diff --git a/Doczy.AI_Automation/pages/interface_0_page_objects.py b/Doczy.AI_Automation/pages/interface_0_page_objects.py index 9d28204..ed86edf 100644 --- a/Doczy.AI_Automation/pages/interface_0_page_objects.py +++ b/Doczy.AI_Automation/pages/interface_0_page_objects.py @@ -17,12 +17,14 @@ class Interface_0(Interface_0_elements): self.driver = driver self.config = config - def select_client(self, client_name=''): - client_name_dropdown = super().get_client_name_dropdown(self.driver, self.config) + def select_client(self, client_name=""): + client_name_dropdown = super().get_client_name_dropdown( + self.driver, self.config + ) client_name_dropdown.click() client_list = super().get_client_list(self.driver) assert len(client_list) > 0 - if client_name == '': + if client_name == "": random_number = random.randint(0, len(client_list) - 1) time.sleep(1) self.logger.info(f"Selected Client Name: {client_list[random_number].text}") @@ -41,8 +43,10 @@ class Interface_0(Interface_0_elements): file_list = "" try: file_input = super().get_file_input(self.driver, self.config) - file_list = read_files(self.config["contract_files_path"]+client_name+"\\") - files = '\n'.join(file_list) + file_list = read_files( + self.config["contract_files_path"] + client_name + "\\" + ) + files = "\n".join(file_list) print(files) file_input.send_keys(files) self.logger.info(f"Uploading files...") @@ -66,8 +70,7 @@ class Interface_0(Interface_0_elements): def verify_message(self, tc): time.sleep(3) - if tc == 'POS': + if tc == "POS": return super().get_p_element(self.driver, self.config) else: return super().get_error_toast(self.driver, self.config) - diff --git a/Doczy.AI_Automation/pages/interface_1_page_objects.py b/Doczy.AI_Automation/pages/interface_1_page_objects.py index a7363c8..f557f42 100644 --- a/Doczy.AI_Automation/pages/interface_1_page_objects.py +++ b/Doczy.AI_Automation/pages/interface_1_page_objects.py @@ -6,13 +6,17 @@ from selenium.webdriver.common.by import By from selenium.webdriver.support.ui import WebDriverWait from elements.interface_1_elements import Interface_1_elements -from utils.element_related_methods import get_clickable_element, get_element, get_invisible_element +from utils.element_related_methods import ( + get_clickable_element, + get_element, + get_invisible_element, +) from utils.logger import get_logger class Interface_1(Interface_1_elements): logger = get_logger() - selected_batch = '' + selected_batch = "" result_message = None p_element = None @@ -20,15 +24,19 @@ class Interface_1(Interface_1_elements): self.driver = driver self.config = config - def select_client(self, client_name=''): - client_name_dropdown = super().get_client_name_dropdown(self.driver, self.config) + def select_client(self, client_name=""): + client_name_dropdown = super().get_client_name_dropdown( + self.driver, self.config + ) client_name_dropdown.click() client_list = super().get_client_list(self.driver) assert len(client_list) > 0 - if client_name == '': + if client_name == "": random_number = random.randint(0, len(client_list) - 1) time.sleep(1) - self.logger.debug(f"Selected Client Name: {client_list[random_number].text}") + self.logger.debug( + f"Selected Client Name: {client_list[random_number].text}" + ) client_name = client_list[random_number].text client_list[random_number].click() else: @@ -40,12 +48,12 @@ class Interface_1(Interface_1_elements): time.sleep(20) return client_name - def select_batch(self, batch_id_name=''): + def select_batch(self, batch_id_name=""): batch_id = super().get_batch_id(self.driver, self.config) batch_id.click() batch_id_list = super().get_batch_id_list(self.driver) assert len(batch_id_list) > 0 - if batch_id_name == '': + if batch_id_name == "": random_number = random.randint(1, len(batch_id_list) - 1) self.selected_batch = batch_id_list[random_number].text time.sleep(1) @@ -68,27 +76,29 @@ class Interface_1(Interface_1_elements): group_checkbox = [] checkbox_list = super().get_checkbox_list(self.driver) assert len(checkbox_list) > 0 - if(len(group_checkbox)>0): - checked_groups = '' + if len(group_checkbox) > 0: + checked_groups = "" for checkbox in checkbox_list: if checkbox.text in group_checkbox: checkbox.click() - checked_groups = f'{checked_groups}{checkbox.text}, ' + checked_groups = f"{checked_groups}{checkbox.text}, " self.logger.debug(f"checked Group No.: {checked_groups[:-2]}") else: random_number = random.randint(0, len(checkbox_list) - 1) random_elements = random.sample(checkbox_list, random_number) - checked_groups = '' + checked_groups = "" for checkbox in random_elements: time.sleep(1) checkbox.click() time.sleep(1) - checked_groups = f'{checked_groups}{checkbox.text}, ' + checked_groups = f"{checked_groups}{checkbox.text}, " self.logger.debug(f"checked Group No.: {checked_groups[:-2]}") time.sleep(5) def click_read_the_contracts_from_path(self): - read_the_contracts_from_path_btn = super().get_read_the_contracts_from_path_btn(self.driver, self.config) + read_the_contracts_from_path_btn = super().get_read_the_contracts_from_path_btn( + self.driver, self.config + ) read_the_contracts_from_path_btn.click() time.sleep(5) @@ -98,14 +108,18 @@ class Interface_1(Interface_1_elements): for row in rows: print(row) if len(rows) > 0: - self.logger.debug(f"{len(rows)} contracts are read from batch: {self.selected_batch}") + self.logger.debug( + f"{len(rows)} contracts are read from batch: {self.selected_batch}" + ) assert True else: assert False time.sleep(5) def click_run_doczy_ai_pipeline(self): - run_doczy_ai_pipeline_btn = super().get_run_doczy_ai_pipeline_btn(self.driver, self.config) + run_doczy_ai_pipeline_btn = super().get_run_doczy_ai_pipeline_btn( + self.driver, self.config + ) run_doczy_ai_pipeline_btn.click() time.sleep(3) self.result_message = super().get_error(self.driver, self.config) diff --git a/Doczy.AI_Automation/pages/interface_2_page_objects.py b/Doczy.AI_Automation/pages/interface_2_page_objects.py index 4c9dcaa..59282e8 100644 --- a/Doczy.AI_Automation/pages/interface_2_page_objects.py +++ b/Doczy.AI_Automation/pages/interface_2_page_objects.py @@ -6,15 +6,19 @@ from selenium.webdriver.common.by import By from selenium.webdriver.support.ui import WebDriverWait from elements.interface_2_elements import Interface_2_elements -from utils.element_related_methods import get_clickable_element, get_element, get_invisible_element +from utils.element_related_methods import ( + get_clickable_element, + get_element, + get_invisible_element, +) from utils.logger import get_logger class Interface_2(Interface_2_elements): logger = get_logger() - selected_batch = '' - selected_contract = '' - selected_field_group = '' + selected_batch = "" + selected_contract = "" + selected_field_group = "" result_message = None p_element = None @@ -22,12 +26,14 @@ class Interface_2(Interface_2_elements): self.driver = driver self.config = config - def select_client(self, client_name=''): - client_name_dropdown = super().get_client_name_dropdown(self.driver, self.config) + def select_client(self, client_name=""): + client_name_dropdown = super().get_client_name_dropdown( + self.driver, self.config + ) client_name_dropdown.click() client_list = super().get_client_list(self.driver) assert len(client_list) > 0 - if client_name == '': + if client_name == "": random_number = random.randint(0, len(client_list) - 1) time.sleep(1) self.logger.info(f"Selected Client Name: {client_list[random_number].text}") @@ -43,12 +49,16 @@ class Interface_2(Interface_2_elements): time.sleep(20) return client_name - def select_batch(self, batch_id_name=''): - batch_id = super().get_batch_id(self.driver, self.config)#get_clickable_element(self.driver, By.XPATH, '/html/body/div/div[1]/div[1]/div/div/div/section[2]/div[1]/div/div/div/div[5]/div[2]/div/div/div/div/div') + def select_batch(self, batch_id_name=""): + batch_id = super().get_batch_id( + self.driver, self.config + ) # get_clickable_element(self.driver, By.XPATH, '/html/body/div/div[1]/div[1]/div/div/div/section[2]/div[1]/div/div/div/div[5]/div[2]/div/div/div/div/div') batch_id.click() - batch_id_list = super().get_batch_id_list(self.driver)#self.driver.find_elements(By.XPATH, "//li") + batch_id_list = super().get_batch_id_list( + self.driver + ) # self.driver.find_elements(By.XPATH, "//li") assert len(batch_id_list) > 0 - if batch_id_name == '': + if batch_id_name == "": random_number = random.randint(1, len(batch_id_list) - 1) self.selected_batch = batch_id_list[random_number].text time.sleep(1) @@ -88,7 +98,9 @@ class Interface_2(Interface_2_elements): time.sleep(1) self.selected_field_group = field_group_list[random_number].text time.sleep(1) - self.logger.info(f"Selected field froup: {field_group_list[random_number].text}") + self.logger.info( + f"Selected field froup: {field_group_list[random_number].text}" + ) time.sleep(1) field_group_list[random_number].click() time.sleep(5) @@ -102,14 +114,18 @@ class Interface_2(Interface_2_elements): result_table = super().get_result_table(self.driver, self.config) rows = super().get_table_rows(result_table) if len(rows) > 0: - self.logger.info(f"{len(rows)} fields are read from contract: {self.selected_contract}") + self.logger.info( + f"{len(rows)} fields are read from contract: {self.selected_contract}" + ) assert True else: assert False time.sleep(5) def click_kickoff_database_integration_btn(self): - kickoff_database_btn = super().get_kickoff_database_btn(self.driver, self.config) + kickoff_database_btn = super().get_kickoff_database_btn( + self.driver, self.config + ) kickoff_database_btn.click() time.sleep(15) diff --git a/Doczy.AI_Automation/tests/conftest.py b/Doczy.AI_Automation/tests/conftest.py index 3123fec..017f549 100644 --- a/Doczy.AI_Automation/tests/conftest.py +++ b/Doczy.AI_Automation/tests/conftest.py @@ -24,9 +24,9 @@ def pytest_html_report_title(report): def client_list(): - file_path = 'C:\\Ankit\\Code\\doczy.ai\\Doczy.AI_Automation\\data\\data.xlsx' - sheet_name = 'client_data' - return read_excel(file_path, sheet_name, ['client']) + file_path = "C:\\Ankit\\Code\\doczy.ai\\Doczy.AI_Automation\\data\\data.xlsx" + sheet_name = "client_data" + return read_excel(file_path, sheet_name, ["client"]) @pytest.fixture @@ -36,28 +36,28 @@ def client_list_fixture(): @pytest.fixture def groups(): - file_path = 'C:\\Ankit\\Code\\doczy.ai\\Doczy.AI_Automation\\data\\data.xlsx' - sheet_name = 'groups' - return read_excel(file_path, sheet_name, ['groups']) + file_path = "C:\\Ankit\\Code\\doczy.ai\\Doczy.AI_Automation\\data\\data.xlsx" + sheet_name = "groups" + return read_excel(file_path, sheet_name, ["groups"]) def batch_list(): - file_path = 'C:\\Ankit\\Code\\doczy.ai\\Doczy.AI_Automation\\data\\data.xlsx' - sheet_name = 'batch_data' - return read_excel(file_path, sheet_name, ['client', 'Batch']) + file_path = "C:\\Ankit\\Code\\doczy.ai\\Doczy.AI_Automation\\data\\data.xlsx" + sheet_name = "batch_data" + return read_excel(file_path, sheet_name, ["client", "Batch"]) @pytest.fixture @allure.title("Creating Web Driver") def driver(config, request): - browser = request.config.getoption('--browser') - if browser == 'remote': + browser = request.config.getoption("--browser") + if browser == "remote": driver = webdriver.Remote() yield driver - elif browser == 'firefox': + elif browser == "firefox": driver = webdriver.Firefox() yield driver - elif browser == 'edge': + elif browser == "edge": driver = webdriver.Edge() yield driver else: @@ -83,7 +83,7 @@ def driver(config, request): @pytest.fixture def config(request): - env = request.config.getoption('--env') + env = request.config.getoption("--env") return read_yaml_config(env) @@ -95,5 +95,9 @@ def read_yaml_config(env): def pytest_addoption(parser): - parser.addoption('--env', action='store', default='DEV', help='Environment to run tests against') - parser.addoption('--browser', action='store', default='firefox', help='browser to run test on') + parser.addoption( + "--env", action="store", default="DEV", help="Environment to run tests against" + ) + parser.addoption( + "--browser", action="store", default="firefox", help="browser to run test on" + ) diff --git a/Doczy.AI_Automation/tests/test_end_to_end.py b/Doczy.AI_Automation/tests/test_end_to_end.py index 9c2a735..cffe79c 100644 --- a/Doczy.AI_Automation/tests/test_end_to_end.py +++ b/Doczy.AI_Automation/tests/test_end_to_end.py @@ -24,12 +24,15 @@ class Test_end_to_end: @allure.suite("Tests for End to end") @allure.sub_suite("Tests for End to end functionality") @allure.title("Verifying End to end functionality") - @allure.feature('End to end') - @allure.story('Verifying End to end functionality') + @allure.feature("End to end") + @allure.story("Verifying End to end functionality") @allure.severity(allure.severity_level.CRITICAL) def test_end_to_end(self, client, driver, config, groups): self.logger.info("Verifying End to end functionality") - driver.get(config['url_interface_0'] + "?code=0.AVsA57e_zrnCv0-sKQYbnoaPz5D--u_XfqNDqwMZoL4vF1hbAN8.AgABBAIAAAApTwJmzXqdR4BN2miheQMYAgDs_wUA9P90Kn8ixTenU_xRich4oROh6FiAIvqe9IcmtiRPSuy7Cl3U70LQWDBnyGZF1wF2US5pmZAl3cYSs3CXilUakzjFaRQS0hZDH-aAb1gwPAtMLD8kA5R0r8ZSuFDfE5TL30BocqNvItA96_6L9DmagcfQmtFhimBmT9zuiGWjDb13rNFI97_UyUojskrzjPLW_PSP_ynOqk9bKCGlsEU4mqnUfHW0HRarnk4nrv-H7mp4LXmHxlpNIUJjIT4lp1vhFz0TdsPXkTypHOhgjPQ6LA0S6NFcdd3pbejPGiRibqwOdS9HI2nYk4g6-FEO7-jRiEchFn_CeXILamMPkZuh4HXvvITZxkYfGHUXqap2mABFXvBknDT0QXqyonjknLCl8-HBmpRt8xe3D5cm95P_147j63uKS3CK9HM7aTAJ9aHAcBwkiF4axlYMoG9WRqC8RbPKaYmBKhcjScCNorqQRKbElXi0O4tkJwkS-CxLElQtQYhwdd1sAxLRSNg4WPz4RN8eH2osmh7RlZz3NrLe465U0PMIFYVRmJSEmiBilYY1LC-Ydg521gxePSpK882qOxJe4URjD2PziYAZ0mFVDvWpsLNtbV3B67mPBlyJD98gkcUsaY-34bmqRxfzbCJRxqBmBrLke-Y4vV16r4m4aAgNIH_sz9nC61hoB46UDpjiwYvMrCRDMrRthsUYD3FpA2CCxTmgNBXYZggLuHrJ_U7pHL71owHdv94yKVmNZOm4D8qs3-0y3zlYzlOC8f-RQQ1WRr7J20L1iN5VOesrv6xwVYS9EQRPOahkmOdwz1LrAcaDoxpr9CI&session_state=7515cdc9-6fd7-4e93-b526-f0a3c23401fd#") + driver.get( + config["url_interface_0"] + + "?code=0.AVsA57e_zrnCv0-sKQYbnoaPz5D--u_XfqNDqwMZoL4vF1hbAN8.AgABBAIAAAApTwJmzXqdR4BN2miheQMYAgDs_wUA9P90Kn8ixTenU_xRich4oROh6FiAIvqe9IcmtiRPSuy7Cl3U70LQWDBnyGZF1wF2US5pmZAl3cYSs3CXilUakzjFaRQS0hZDH-aAb1gwPAtMLD8kA5R0r8ZSuFDfE5TL30BocqNvItA96_6L9DmagcfQmtFhimBmT9zuiGWjDb13rNFI97_UyUojskrzjPLW_PSP_ynOqk9bKCGlsEU4mqnUfHW0HRarnk4nrv-H7mp4LXmHxlpNIUJjIT4lp1vhFz0TdsPXkTypHOhgjPQ6LA0S6NFcdd3pbejPGiRibqwOdS9HI2nYk4g6-FEO7-jRiEchFn_CeXILamMPkZuh4HXvvITZxkYfGHUXqap2mABFXvBknDT0QXqyonjknLCl8-HBmpRt8xe3D5cm95P_147j63uKS3CK9HM7aTAJ9aHAcBwkiF4axlYMoG9WRqC8RbPKaYmBKhcjScCNorqQRKbElXi0O4tkJwkS-CxLElQtQYhwdd1sAxLRSNg4WPz4RN8eH2osmh7RlZz3NrLe465U0PMIFYVRmJSEmiBilYY1LC-Ydg521gxePSpK882qOxJe4URjD2PziYAZ0mFVDvWpsLNtbV3B67mPBlyJD98gkcUsaY-34bmqRxfzbCJRxqBmBrLke-Y4vV16r4m4aAgNIH_sz9nC61hoB46UDpjiwYvMrCRDMrRthsUYD3FpA2CCxTmgNBXYZggLuHrJ_U7pHL71owHdv94yKVmNZOm4D8qs3-0y3zlYzlOC8f-RQQ1WRr7J20L1iN5VOesrv6xwVYS9EQRPOahkmOdwz1LrAcaDoxpr9CI&session_state=7515cdc9-6fd7-4e93-b526-f0a3c23401fd#" + ) with allure.step("Signing in"): # sign_in(driver, config) self.logger.info("Signed in successful") @@ -41,26 +44,29 @@ class Test_end_to_end: with allure.step("Creating Batch"): interface_0_page.create_batch() with allure.step("Asserting Batch Creating"): - assert interface_0_page.verify_message('POS').startswith("batch_") - batch = interface_0_page.verify_message('POS') + assert interface_0_page.verify_message("POS").startswith("batch_") + batch = interface_0_page.verify_message("POS") self.logger.info(f"Created batch with Batch ID: {batch}") - with allure.step('Navigating to interface 1'): - driver.get(config['url_interface_1'] + "?code=0.AVsA57e_zrnCv0-sKQYbnoaPz5D--u_XfqNDqwMZoL4vF1hbAN8.AgABBAIAAAApTwJmzXqdR4BN2miheQMYAgDs_wUA9P90Kn8ixTenU_xRich4oROh6FiAIvqe9IcmtiRPSuy7Cl3U70LQWDBnyGZF1wF2US5pmZAl3cYSs3CXilUakzjFaRQS0hZDH-aAb1gwPAtMLD8kA5R0r8ZSuFDfE5TL30BocqNvItA96_6L9DmagcfQmtFhimBmT9zuiGWjDb13rNFI97_UyUojskrzjPLW_PSP_ynOqk9bKCGlsEU4mqnUfHW0HRarnk4nrv-H7mp4LXmHxlpNIUJjIT4lp1vhFz0TdsPXkTypHOhgjPQ6LA0S6NFcdd3pbejPGiRibqwOdS9HI2nYk4g6-FEO7-jRiEchFn_CeXILamMPkZuh4HXvvITZxkYfGHUXqap2mABFXvBknDT0QXqyonjknLCl8-HBmpRt8xe3D5cm95P_147j63uKS3CK9HM7aTAJ9aHAcBwkiF4axlYMoG9WRqC8RbPKaYmBKhcjScCNorqQRKbElXi0O4tkJwkS-CxLElQtQYhwdd1sAxLRSNg4WPz4RN8eH2osmh7RlZz3NrLe465U0PMIFYVRmJSEmiBilYY1LC-Ydg521gxePSpK882qOxJe4URjD2PziYAZ0mFVDvWpsLNtbV3B67mPBlyJD98gkcUsaY-34bmqRxfzbCJRxqBmBrLke-Y4vV16r4m4aAgNIH_sz9nC61hoB46UDpjiwYvMrCRDMrRthsUYD3FpA2CCxTmgNBXYZggLuHrJ_U7pHL71owHdv94yKVmNZOm4D8qs3-0y3zlYzlOC8f-RQQ1WRr7J20L1iN5VOesrv6xwVYS9EQRPOahkmOdwz1LrAcaDoxpr9CI&session_state=7515cdc9-6fd7-4e93-b526-f0a3c23401fd#") + with allure.step("Navigating to interface 1"): + driver.get( + config["url_interface_1"] + + "?code=0.AVsA57e_zrnCv0-sKQYbnoaPz5D--u_XfqNDqwMZoL4vF1hbAN8.AgABBAIAAAApTwJmzXqdR4BN2miheQMYAgDs_wUA9P90Kn8ixTenU_xRich4oROh6FiAIvqe9IcmtiRPSuy7Cl3U70LQWDBnyGZF1wF2US5pmZAl3cYSs3CXilUakzjFaRQS0hZDH-aAb1gwPAtMLD8kA5R0r8ZSuFDfE5TL30BocqNvItA96_6L9DmagcfQmtFhimBmT9zuiGWjDb13rNFI97_UyUojskrzjPLW_PSP_ynOqk9bKCGlsEU4mqnUfHW0HRarnk4nrv-H7mp4LXmHxlpNIUJjIT4lp1vhFz0TdsPXkTypHOhgjPQ6LA0S6NFcdd3pbejPGiRibqwOdS9HI2nYk4g6-FEO7-jRiEchFn_CeXILamMPkZuh4HXvvITZxkYfGHUXqap2mABFXvBknDT0QXqyonjknLCl8-HBmpRt8xe3D5cm95P_147j63uKS3CK9HM7aTAJ9aHAcBwkiF4axlYMoG9WRqC8RbPKaYmBKhcjScCNorqQRKbElXi0O4tkJwkS-CxLElQtQYhwdd1sAxLRSNg4WPz4RN8eH2osmh7RlZz3NrLe465U0PMIFYVRmJSEmiBilYY1LC-Ydg521gxePSpK882qOxJe4URjD2PziYAZ0mFVDvWpsLNtbV3B67mPBlyJD98gkcUsaY-34bmqRxfzbCJRxqBmBrLke-Y4vV16r4m4aAgNIH_sz9nC61hoB46UDpjiwYvMrCRDMrRthsUYD3FpA2CCxTmgNBXYZggLuHrJ_U7pHL71owHdv94yKVmNZOm4D8qs3-0y3zlYzlOC8f-RQQ1WRr7J20L1iN5VOesrv6xwVYS9EQRPOahkmOdwz1LrAcaDoxpr9CI&session_state=7515cdc9-6fd7-4e93-b526-f0a3c23401fd#" + ) # pre_sign_in_page = Pre_sign_in_page(driver) # pre_sign_in_page.click_sign_in_link() # self.logger.info("Signed in successful") get_title(driver, config, "interface_1 · Streamlit") - with allure.step('Selecting Client'): + with allure.step("Selecting Client"): interface_1 = Interface_1(driver, config) interface_1.select_client(client) - with allure.step('Selecting Batch'): + with allure.step("Selecting Batch"): interface_1.select_batch(batch) - with allure.step('Selecting Groups'): + with allure.step("Selecting Groups"): interface_1.check_group_no(groups) - with allure.step('Reading Contracts from S3 Bucket'): + with allure.step("Reading Contracts from S3 Bucket"): interface_1.click_read_the_contracts_from_path() interface_1.check_table_contain() - with allure.step('Running Doczy.AI Pipeline'): + with allure.step("Running Doczy.AI Pipeline"): interface_1.click_run_doczy_ai_pipeline() if interface_1.result_message.text == "Success": self.logger.info(f"Files are passed to Doczy.AI pipeline") @@ -70,8 +76,10 @@ class Test_end_to_end: assert False with allure.step("Signing in"): time.sleep(900) - driver.get(config[ - 'url_interface_2'] + "?code=0.AVsA57e_zrnCv0-sKQYbnoaPz5D--u_XfqNDqwMZoL4vF1hbAN8.AgABBAIAAAApTwJmzXqdR4BN2miheQMYAgDs_wUA9P90Kn8ixTenU_xRich4oROh6FiAIvqe9IcmtiRPSuy7Cl3U70LQWDBnyGZF1wF2US5pmZAl3cYSs3CXilUakzjFaRQS0hZDH-aAb1gwPAtMLD8kA5R0r8ZSuFDfE5TL30BocqNvItA96_6L9DmagcfQmtFhimBmT9zuiGWjDb13rNFI97_UyUojskrzjPLW_PSP_ynOqk9bKCGlsEU4mqnUfHW0HRarnk4nrv-H7mp4LXmHxlpNIUJjIT4lp1vhFz0TdsPXkTypHOhgjPQ6LA0S6NFcdd3pbejPGiRibqwOdS9HI2nYk4g6-FEO7-jRiEchFn_CeXILamMPkZuh4HXvvITZxkYfGHUXqap2mABFXvBknDT0QXqyonjknLCl8-HBmpRt8xe3D5cm95P_147j63uKS3CK9HM7aTAJ9aHAcBwkiF4axlYMoG9WRqC8RbPKaYmBKhcjScCNorqQRKbElXi0O4tkJwkS-CxLElQtQYhwdd1sAxLRSNg4WPz4RN8eH2osmh7RlZz3NrLe465U0PMIFYVRmJSEmiBilYY1LC-Ydg521gxePSpK882qOxJe4URjD2PziYAZ0mFVDvWpsLNtbV3B67mPBlyJD98gkcUsaY-34bmqRxfzbCJRxqBmBrLke-Y4vV16r4m4aAgNIH_sz9nC61hoB46UDpjiwYvMrCRDMrRthsUYD3FpA2CCxTmgNBXYZggLuHrJ_U7pHL71owHdv94yKVmNZOm4D8qs3-0y3zlYzlOC8f-RQQ1WRr7J20L1iN5VOesrv6xwVYS9EQRPOahkmOdwz1LrAcaDoxpr9CI&session_state=7515cdc9-6fd7-4e93-b526-f0a3c23401fd#") + driver.get( + config["url_interface_2"] + + "?code=0.AVsA57e_zrnCv0-sKQYbnoaPz5D--u_XfqNDqwMZoL4vF1hbAN8.AgABBAIAAAApTwJmzXqdR4BN2miheQMYAgDs_wUA9P90Kn8ixTenU_xRich4oROh6FiAIvqe9IcmtiRPSuy7Cl3U70LQWDBnyGZF1wF2US5pmZAl3cYSs3CXilUakzjFaRQS0hZDH-aAb1gwPAtMLD8kA5R0r8ZSuFDfE5TL30BocqNvItA96_6L9DmagcfQmtFhimBmT9zuiGWjDb13rNFI97_UyUojskrzjPLW_PSP_ynOqk9bKCGlsEU4mqnUfHW0HRarnk4nrv-H7mp4LXmHxlpNIUJjIT4lp1vhFz0TdsPXkTypHOhgjPQ6LA0S6NFcdd3pbejPGiRibqwOdS9HI2nYk4g6-FEO7-jRiEchFn_CeXILamMPkZuh4HXvvITZxkYfGHUXqap2mABFXvBknDT0QXqyonjknLCl8-HBmpRt8xe3D5cm95P_147j63uKS3CK9HM7aTAJ9aHAcBwkiF4axlYMoG9WRqC8RbPKaYmBKhcjScCNorqQRKbElXi0O4tkJwkS-CxLElQtQYhwdd1sAxLRSNg4WPz4RN8eH2osmh7RlZz3NrLe465U0PMIFYVRmJSEmiBilYY1LC-Ydg521gxePSpK882qOxJe4URjD2PziYAZ0mFVDvWpsLNtbV3B67mPBlyJD98gkcUsaY-34bmqRxfzbCJRxqBmBrLke-Y4vV16r4m4aAgNIH_sz9nC61hoB46UDpjiwYvMrCRDMrRthsUYD3FpA2CCxTmgNBXYZggLuHrJ_U7pHL71owHdv94yKVmNZOm4D8qs3-0y3zlYzlOC8f-RQQ1WRr7J20L1iN5VOesrv6xwVYS9EQRPOahkmOdwz1LrAcaDoxpr9CI&session_state=7515cdc9-6fd7-4e93-b526-f0a3c23401fd#" + ) # pre_sign_in_page = Pre_sign_in_page(driver) # pre_sign_in_page.click_sign_in_link() # self.logger.info("Signed in successful") diff --git a/Doczy.AI_Automation/tests/test_interface_0.py b/Doczy.AI_Automation/tests/test_interface_0.py index 22cf759..bad5946 100644 --- a/Doczy.AI_Automation/tests/test_interface_0.py +++ b/Doczy.AI_Automation/tests/test_interface_0.py @@ -18,14 +18,14 @@ class Test_interface_0: @allure.suite("Tests for Interface 0") @allure.sub_suite("Tests for Navigation") @allure.title("Verifying designated Interface 0 Page opens successfully") - @allure.feature('Interface 0') - @allure.story('Verifying designated Interface 0 Page opens successfully') + @allure.feature("Interface 0") + @allure.story("Verifying designated Interface 0 Page opens successfully") @allure.severity(allure.severity_level.CRITICAL) def test_validate_navigation(self, driver, config): self.logger.info("Verifying designated Interface 0 Page opens successfully") - driver.get(config['url_interface_0']) + driver.get(config["url_interface_0"]) get_title(driver, config, "interface_0 · Streamlit") - if driver.title == 'interface_0 · Streamlit': + if driver.title == "interface_0 · Streamlit": self.logger.info("Designated Interface 0 Page opened successfully") assert True else: @@ -37,16 +37,19 @@ class Test_interface_0: @allure.suite("Tests for Interface 0") @allure.sub_suite("Tests for Authentication") @allure.title("Verifying Single Sign-On (SSO) Login for Interface 0") - @allure.feature('Interface 0') - @allure.story('Verifying Single Sign-On (SSO) Login for Interface 0') + @allure.feature("Interface 0") + @allure.story("Verifying Single Sign-On (SSO) Login for Interface 0") @allure.severity(allure.severity_level.CRITICAL) def test_verify_sign_in(self, driver, config): self.logger.info("Verifying Single Sign-On (SSO) Login for Interface 0") - driver.get(config['url_interface_0'] + "?code=0.AVsA57e_zrnCv0-sKQYbnoaPz5D--u_XfqNDqwMZoL4vF1hbAN8.AgABBAIAAAApTwJmzXqdR4BN2miheQMYAgDs_wUA9P-Xu65_k1cptZUCoWDROksYjcCYq_DvUcJbXG9MczlC6LlsO80A1Y4mYcEajsGirxfVgaxYazfwYb5sXUbARKILzxaoqtvPuW3nJDMPcIjjOCG5bNQo2tu67oUGeIWJEQIMWH0jVXB4D8N1j3m3Rd1lPkV-d_Wr6LNfieAK1ralDynQXV8zd5y5TNJTkplLjazIKbhjMr9SrhfCs4m-_IrWIPcGIteLt4oly-XzxAxgX6mhs0WTgcXqUfuyTuYkPJrZs70exEdNQi48Ml_vZRXux7NXn_yuyRu51qZ34sG89Wj8nkIdmPyoOi-CTwl0QFSUzNKixWczhQssq0WyRfKgw4-afANx2mkHozy9gwL1JbGm6Je29H00ntHrQX0YvSWtEvTT4kdXRh2F6zaS12LZJ31-GUAOUNdZcz09FitSbdEVSF-Z9e4VPNEZVPoUhj3cbFAf02N5Wv3K4vapE24Aof18CHav9gvg5-BcZUzfPhq-zAva2WoC-DANUSPHOjCQlUBRt2WjnJ-9qqAWNV8A9MIf3-tF5kES48HpIVwnt05uCs0t6M6bDJIK6eml4lV9YfVtlmO9LOxhtoMO69Sj5RExcol8Yz_-5T1u0xlbdKlWetUXQTFNwfjcz1u_HYjXVy_fpLFG1M2MdhjtvQ85hxGBv32d8BDYuGeXpjoc-fYEOnm1tOlESrqbZBzFLRv1MmRlkxiAg3uW0Myb32wKjS2iTFp3pkQ-D8Yxv1h0_HRd9ZeCQE-JfLJ2HALswI9wxLFidsCvZ4xaOXpqLZX0o8U7SwHvLw9BRvPvPol5b7mHQ9_OS&session_state=7515cdc9-6fd7-4e93-b526-f0a3c23401fd#") + driver.get( + config["url_interface_0"] + + "?code=0.AVsA57e_zrnCv0-sKQYbnoaPz5D--u_XfqNDqwMZoL4vF1hbAN8.AgABBAIAAAApTwJmzXqdR4BN2miheQMYAgDs_wUA9P-Xu65_k1cptZUCoWDROksYjcCYq_DvUcJbXG9MczlC6LlsO80A1Y4mYcEajsGirxfVgaxYazfwYb5sXUbARKILzxaoqtvPuW3nJDMPcIjjOCG5bNQo2tu67oUGeIWJEQIMWH0jVXB4D8N1j3m3Rd1lPkV-d_Wr6LNfieAK1ralDynQXV8zd5y5TNJTkplLjazIKbhjMr9SrhfCs4m-_IrWIPcGIteLt4oly-XzxAxgX6mhs0WTgcXqUfuyTuYkPJrZs70exEdNQi48Ml_vZRXux7NXn_yuyRu51qZ34sG89Wj8nkIdmPyoOi-CTwl0QFSUzNKixWczhQssq0WyRfKgw4-afANx2mkHozy9gwL1JbGm6Je29H00ntHrQX0YvSWtEvTT4kdXRh2F6zaS12LZJ31-GUAOUNdZcz09FitSbdEVSF-Z9e4VPNEZVPoUhj3cbFAf02N5Wv3K4vapE24Aof18CHav9gvg5-BcZUzfPhq-zAva2WoC-DANUSPHOjCQlUBRt2WjnJ-9qqAWNV8A9MIf3-tF5kES48HpIVwnt05uCs0t6M6bDJIK6eml4lV9YfVtlmO9LOxhtoMO69Sj5RExcol8Yz_-5T1u0xlbdKlWetUXQTFNwfjcz1u_HYjXVy_fpLFG1M2MdhjtvQ85hxGBv32d8BDYuGeXpjoc-fYEOnm1tOlESrqbZBzFLRv1MmRlkxiAg3uW0Myb32wKjS2iTFp3pkQ-D8Yxv1h0_HRd9ZeCQE-JfLJ2HALswI9wxLFidsCvZ4xaOXpqLZX0o8U7SwHvLw9BRvPvPol5b7mHQ9_OS&session_state=7515cdc9-6fd7-4e93-b526-f0a3c23401fd#" + ) # sign_in(driver, config) self.logger.info("Signed in successful") get_title(driver, config, "interface_0 · Streamlit") - if driver.title == 'interface_0 · Streamlit': + if driver.title == "interface_0 · Streamlit": self.logger.info("Interface 0 page opened successfully") assert True else: @@ -58,13 +61,16 @@ class Test_interface_0: @allure.suite("Tests for Interface 0") @allure.sub_suite("Tests for Batch Creation") @allure.title("Verifying S3 bucket is created with batch ID") - @allure.feature('Interface 0') - @allure.story('Verifying S3 bucket is created with batch ID') + @allure.feature("Interface 0") + @allure.story("Verifying S3 bucket is created with batch ID") @allure.severity(allure.severity_level.CRITICAL) @allure.testcase("TC-003") def test_verify_batch_creation(self, driver, config): self.logger.info("Verifying S3 bucket is created with batch ID") - driver.get(config['url_interface_0'] + "?code=0.AVsA57e_zrnCv0-sKQYbnoaPz5D--u_XfqNDqwMZoL4vF1hbAN8.AgABBAIAAAApTwJmzXqdR4BN2miheQMYAgDs_wUA9P-Xu65_k1cptZUCoWDROksYjcCYq_DvUcJbXG9MczlC6LlsO80A1Y4mYcEajsGirxfVgaxYazfwYb5sXUbARKILzxaoqtvPuW3nJDMPcIjjOCG5bNQo2tu67oUGeIWJEQIMWH0jVXB4D8N1j3m3Rd1lPkV-d_Wr6LNfieAK1ralDynQXV8zd5y5TNJTkplLjazIKbhjMr9SrhfCs4m-_IrWIPcGIteLt4oly-XzxAxgX6mhs0WTgcXqUfuyTuYkPJrZs70exEdNQi48Ml_vZRXux7NXn_yuyRu51qZ34sG89Wj8nkIdmPyoOi-CTwl0QFSUzNKixWczhQssq0WyRfKgw4-afANx2mkHozy9gwL1JbGm6Je29H00ntHrQX0YvSWtEvTT4kdXRh2F6zaS12LZJ31-GUAOUNdZcz09FitSbdEVSF-Z9e4VPNEZVPoUhj3cbFAf02N5Wv3K4vapE24Aof18CHav9gvg5-BcZUzfPhq-zAva2WoC-DANUSPHOjCQlUBRt2WjnJ-9qqAWNV8A9MIf3-tF5kES48HpIVwnt05uCs0t6M6bDJIK6eml4lV9YfVtlmO9LOxhtoMO69Sj5RExcol8Yz_-5T1u0xlbdKlWetUXQTFNwfjcz1u_HYjXVy_fpLFG1M2MdhjtvQ85hxGBv32d8BDYuGeXpjoc-fYEOnm1tOlESrqbZBzFLRv1MmRlkxiAg3uW0Myb32wKjS2iTFp3pkQ-D8Yxv1h0_HRd9ZeCQE-JfLJ2HALswI9wxLFidsCvZ4xaOXpqLZX0o8U7SwHvLw9BRvPvPol5b7mHQ9_OS&session_state=7515cdc9-6fd7-4e93-b526-f0a3c23401fd#") + driver.get( + config["url_interface_0"] + + "?code=0.AVsA57e_zrnCv0-sKQYbnoaPz5D--u_XfqNDqwMZoL4vF1hbAN8.AgABBAIAAAApTwJmzXqdR4BN2miheQMYAgDs_wUA9P-Xu65_k1cptZUCoWDROksYjcCYq_DvUcJbXG9MczlC6LlsO80A1Y4mYcEajsGirxfVgaxYazfwYb5sXUbARKILzxaoqtvPuW3nJDMPcIjjOCG5bNQo2tu67oUGeIWJEQIMWH0jVXB4D8N1j3m3Rd1lPkV-d_Wr6LNfieAK1ralDynQXV8zd5y5TNJTkplLjazIKbhjMr9SrhfCs4m-_IrWIPcGIteLt4oly-XzxAxgX6mhs0WTgcXqUfuyTuYkPJrZs70exEdNQi48Ml_vZRXux7NXn_yuyRu51qZ34sG89Wj8nkIdmPyoOi-CTwl0QFSUzNKixWczhQssq0WyRfKgw4-afANx2mkHozy9gwL1JbGm6Je29H00ntHrQX0YvSWtEvTT4kdXRh2F6zaS12LZJ31-GUAOUNdZcz09FitSbdEVSF-Z9e4VPNEZVPoUhj3cbFAf02N5Wv3K4vapE24Aof18CHav9gvg5-BcZUzfPhq-zAva2WoC-DANUSPHOjCQlUBRt2WjnJ-9qqAWNV8A9MIf3-tF5kES48HpIVwnt05uCs0t6M6bDJIK6eml4lV9YfVtlmO9LOxhtoMO69Sj5RExcol8Yz_-5T1u0xlbdKlWetUXQTFNwfjcz1u_HYjXVy_fpLFG1M2MdhjtvQ85hxGBv32d8BDYuGeXpjoc-fYEOnm1tOlESrqbZBzFLRv1MmRlkxiAg3uW0Myb32wKjS2iTFp3pkQ-D8Yxv1h0_HRd9ZeCQE-JfLJ2HALswI9wxLFidsCvZ4xaOXpqLZX0o8U7SwHvLw9BRvPvPol5b7mHQ9_OS&session_state=7515cdc9-6fd7-4e93-b526-f0a3c23401fd#" + ) # sign_in(driver, config) with allure.step("Signing in"): self.logger.info("Signed in successful") @@ -78,43 +84,51 @@ class Test_interface_0: interface_0_page.create_batch() # self.logger.info(f"Created batch with Batch ID: {interface_0_page.p_element.text}") with allure.step("Asserting Batch Creating"): - assert interface_0_page.verify_message('POS').startswith("batch_") + assert interface_0_page.verify_message("POS").startswith("batch_") @allure.parent_suite("Tests for web interface") @allure.suite("Tests for Interface 0") @allure.sub_suite("Tests for Batch Creation without client") @allure.title("Verifying S3 bucket is created without selecting client") - @allure.feature('Interface 0') - @allure.story('Verifying S3 bucket is created without selecting client') + @allure.feature("Interface 0") + @allure.story("Verifying S3 bucket is created without selecting client") @allure.severity(allure.severity_level.CRITICAL) def test_verify_batch_creation_without_client(self, driver, config): self.logger.info("Verifying S3 bucket is created without selecting client") - driver.get(config[ - 'url_interface_0'] + "?code=0.AVsA57e_zrnCv0-sKQYbnoaPz5D--u_XfqNDqwMZoL4vF1hbAN8.AgABBAIAAAApTwJmzXqdR4BN2miheQMYAgDs_wUA9P-Xu65_k1cptZUCoWDROksYjcCYq_DvUcJbXG9MczlC6LlsO80A1Y4mYcEajsGirxfVgaxYazfwYb5sXUbARKILzxaoqtvPuW3nJDMPcIjjOCG5bNQo2tu67oUGeIWJEQIMWH0jVXB4D8N1j3m3Rd1lPkV-d_Wr6LNfieAK1ralDynQXV8zd5y5TNJTkplLjazIKbhjMr9SrhfCs4m-_IrWIPcGIteLt4oly-XzxAxgX6mhs0WTgcXqUfuyTuYkPJrZs70exEdNQi48Ml_vZRXux7NXn_yuyRu51qZ34sG89Wj8nkIdmPyoOi-CTwl0QFSUzNKixWczhQssq0WyRfKgw4-afANx2mkHozy9gwL1JbGm6Je29H00ntHrQX0YvSWtEvTT4kdXRh2F6zaS12LZJ31-GUAOUNdZcz09FitSbdEVSF-Z9e4VPNEZVPoUhj3cbFAf02N5Wv3K4vapE24Aof18CHav9gvg5-BcZUzfPhq-zAva2WoC-DANUSPHOjCQlUBRt2WjnJ-9qqAWNV8A9MIf3-tF5kES48HpIVwnt05uCs0t6M6bDJIK6eml4lV9YfVtlmO9LOxhtoMO69Sj5RExcol8Yz_-5T1u0xlbdKlWetUXQTFNwfjcz1u_HYjXVy_fpLFG1M2MdhjtvQ85hxGBv32d8BDYuGeXpjoc-fYEOnm1tOlESrqbZBzFLRv1MmRlkxiAg3uW0Myb32wKjS2iTFp3pkQ-D8Yxv1h0_HRd9ZeCQE-JfLJ2HALswI9wxLFidsCvZ4xaOXpqLZX0o8U7SwHvLw9BRvPvPol5b7mHQ9_OS&session_state=7515cdc9-6fd7-4e93-b526-f0a3c23401fd#") + driver.get( + config["url_interface_0"] + + "?code=0.AVsA57e_zrnCv0-sKQYbnoaPz5D--u_XfqNDqwMZoL4vF1hbAN8.AgABBAIAAAApTwJmzXqdR4BN2miheQMYAgDs_wUA9P-Xu65_k1cptZUCoWDROksYjcCYq_DvUcJbXG9MczlC6LlsO80A1Y4mYcEajsGirxfVgaxYazfwYb5sXUbARKILzxaoqtvPuW3nJDMPcIjjOCG5bNQo2tu67oUGeIWJEQIMWH0jVXB4D8N1j3m3Rd1lPkV-d_Wr6LNfieAK1ralDynQXV8zd5y5TNJTkplLjazIKbhjMr9SrhfCs4m-_IrWIPcGIteLt4oly-XzxAxgX6mhs0WTgcXqUfuyTuYkPJrZs70exEdNQi48Ml_vZRXux7NXn_yuyRu51qZ34sG89Wj8nkIdmPyoOi-CTwl0QFSUzNKixWczhQssq0WyRfKgw4-afANx2mkHozy9gwL1JbGm6Je29H00ntHrQX0YvSWtEvTT4kdXRh2F6zaS12LZJ31-GUAOUNdZcz09FitSbdEVSF-Z9e4VPNEZVPoUhj3cbFAf02N5Wv3K4vapE24Aof18CHav9gvg5-BcZUzfPhq-zAva2WoC-DANUSPHOjCQlUBRt2WjnJ-9qqAWNV8A9MIf3-tF5kES48HpIVwnt05uCs0t6M6bDJIK6eml4lV9YfVtlmO9LOxhtoMO69Sj5RExcol8Yz_-5T1u0xlbdKlWetUXQTFNwfjcz1u_HYjXVy_fpLFG1M2MdhjtvQ85hxGBv32d8BDYuGeXpjoc-fYEOnm1tOlESrqbZBzFLRv1MmRlkxiAg3uW0Myb32wKjS2iTFp3pkQ-D8Yxv1h0_HRd9ZeCQE-JfLJ2HALswI9wxLFidsCvZ4xaOXpqLZX0o8U7SwHvLw9BRvPvPol5b7mHQ9_OS&session_state=7515cdc9-6fd7-4e93-b526-f0a3c23401fd#" + ) # sign_in(driver, config) self.logger.info("Signed in successful") interface_0_page = Interface_0(driver, config) interface_0_page.create_batch() - assert interface_0_page.verify_message('NEG') == "No Client Name Selected.", "Batch ID is created instead of message" + assert ( + interface_0_page.verify_message("NEG") == "No Client Name Selected." + ), "Batch ID is created instead of message" # @allure.label("type", 'regression') @allure.parent_suite("Tests for web interface") @allure.suite("Tests for Interface 0") @allure.sub_suite("Tests for Batch Creation without files") @allure.title("Verifying S3 bucket is created with no files are uploaded") - @allure.feature('Interface 0') - @allure.story('Verifying S3 bucket is created with no files are uploaded') + @allure.feature("Interface 0") + @allure.story("Verifying S3 bucket is created with no files are uploaded") @allure.severity(allure.severity_level.CRITICAL) def test_verify_no_files_uploaded(self, driver, config): self.logger.info("Verifying S3 bucket is created with no files are uploaded") - driver.get(config['url_interface_0'] + "?code=0.AVsA57e_zrnCv0-sKQYbnoaPz5D--u_XfqNDqwMZoL4vF1hbAN8.AgABBAIAAAApTwJmzXqdR4BN2miheQMYAgDs_wUA9P9LX84t_sXygReogqiCvr7kFfaCaIcCEnNTGHigLFFW27KaCXXvBZ4b6UzF9L92ozdGa9JxHLXSNERV8wUc7DZRC5U_FC1U-hcSZjGFGA_rteC6-iOnqSqzp58hTnxJw5gBPxwqO0ZIBZy3zbjfe2piABPrTTvvftQ_ccdzbZoSoBfVbvpnWNpesYapFyc8zCnUdfKAS3sz65mtBouZqqgRQ6C31NhaoUMl5qFmqLs01vM50k3VrnRA_AqkzMTne0nNgt53FsdQ2WbubuIIZnwNFbDO8WwskHT2r_jVjnSXXs7pgufPMgV0DOlu2BRZKnabLSqiYa0OGsUepnB7dTV9QlduQw-_TYyXdiReapwXxkz9RtfgM8EUgt1ldTGB5Qd4GOtXDwACNcas53RxUoKJeI2vLVLSXKAtvhROs8NqvF222T1IWWDjjari6-M661OZd2kEmx9JeRilSlmIylonz8qhmL40dRWsDppiSeoPiCgIQsoIqZJiCD8TTMyUSo23qxDcQZamIR8LlteELMQdZHUpSfQG8vKCF02I7Bxc8zD8aksP-27O7g7JTCwWQ3OOFbTGH2hy9_JZFzWCj6W8vYbb9yjGJrAccevXg6T-FBa6LBqWZBmTP9q6Ic8XQOUFZaYMkceMuDg6-ARDxyLXWE60ajN3fQA0xkFrWXOyaCJY3fYnxpHwRk_jyC_yEAaZCciUHFyDPK_Uiudh4L4jbpy7wKM_GPR0faJE0OPJbTc32SId6-LC9e2ECaFYtyUNTkLw-DVhZ2MnFbUEj-Y0SAoRYchAroTF7EfByMFLqSXAKeGj&session_state=7515cdc9-6fd7-4e93-b526-f0a3c23401fd#") + driver.get( + config["url_interface_0"] + + "?code=0.AVsA57e_zrnCv0-sKQYbnoaPz5D--u_XfqNDqwMZoL4vF1hbAN8.AgABBAIAAAApTwJmzXqdR4BN2miheQMYAgDs_wUA9P9LX84t_sXygReogqiCvr7kFfaCaIcCEnNTGHigLFFW27KaCXXvBZ4b6UzF9L92ozdGa9JxHLXSNERV8wUc7DZRC5U_FC1U-hcSZjGFGA_rteC6-iOnqSqzp58hTnxJw5gBPxwqO0ZIBZy3zbjfe2piABPrTTvvftQ_ccdzbZoSoBfVbvpnWNpesYapFyc8zCnUdfKAS3sz65mtBouZqqgRQ6C31NhaoUMl5qFmqLs01vM50k3VrnRA_AqkzMTne0nNgt53FsdQ2WbubuIIZnwNFbDO8WwskHT2r_jVjnSXXs7pgufPMgV0DOlu2BRZKnabLSqiYa0OGsUepnB7dTV9QlduQw-_TYyXdiReapwXxkz9RtfgM8EUgt1ldTGB5Qd4GOtXDwACNcas53RxUoKJeI2vLVLSXKAtvhROs8NqvF222T1IWWDjjari6-M661OZd2kEmx9JeRilSlmIylonz8qhmL40dRWsDppiSeoPiCgIQsoIqZJiCD8TTMyUSo23qxDcQZamIR8LlteELMQdZHUpSfQG8vKCF02I7Bxc8zD8aksP-27O7g7JTCwWQ3OOFbTGH2hy9_JZFzWCj6W8vYbb9yjGJrAccevXg6T-FBa6LBqWZBmTP9q6Ic8XQOUFZaYMkceMuDg6-ARDxyLXWE60ajN3fQA0xkFrWXOyaCJY3fYnxpHwRk_jyC_yEAaZCciUHFyDPK_Uiudh4L4jbpy7wKM_GPR0faJE0OPJbTc32SId6-LC9e2ECaFYtyUNTkLw-DVhZ2MnFbUEj-Y0SAoRYchAroTF7EfByMFLqSXAKeGj&session_state=7515cdc9-6fd7-4e93-b526-f0a3c23401fd#" + ) # sign_in(driver, config) with allure.step("Sign In"): self.logger.info("Signed in successful") interface_0_page = Interface_0(driver, config) with allure.step("Selecting Client"): - interface_0_page.select_client('Centene') + interface_0_page.select_client("Centene") with allure.step("Creating Batch"): interface_0_page.create_batch() - assert interface_0_page.verify_message('NEG') == "No Files Selected.", "Batch ID is created instead of message" - + assert ( + interface_0_page.verify_message("NEG") == "No Files Selected." + ), "Batch ID is created instead of message" diff --git a/Doczy.AI_Automation/tests/test_interface_1.py b/Doczy.AI_Automation/tests/test_interface_1.py index d3eec15..3352ea8 100644 --- a/Doczy.AI_Automation/tests/test_interface_1.py +++ b/Doczy.AI_Automation/tests/test_interface_1.py @@ -16,15 +16,15 @@ class Test_interface_1: @allure.suite("Tests for Interface 1") @allure.sub_suite("Tests for Navigation") @allure.title("Verifying designated Interface 1 Page opens successfully") - @allure.feature('Interface 1') - @allure.story('Verifying designated Interface 1 Page opens successfully') + @allure.feature("Interface 1") + @allure.story("Verifying designated Interface 1 Page opens successfully") @allure.severity(allure.severity_level.CRITICAL) def test_validate_navigation(self, driver, config): self.logger.info("Verifying designated Interface 1 Page opens successfully") - driver.get(config['url_interface_1']) + driver.get(config["url_interface_1"]) self.logger.info(f"url_interface_1 {config['url_interface_1']} opened") get_title(driver, config, "interface_1 · Streamlit") - if driver.title == 'interface_1 · Streamlit': + if driver.title == "interface_1 · Streamlit": self.logger.info("Designated Interface 1 Page opened successfully") assert True else: @@ -35,16 +35,19 @@ class Test_interface_1: @allure.suite("Tests for Interface 1") @allure.sub_suite("Tests for Authentication") @allure.title("Verifying Single Sign-On (SSO) Login for Interface 1") - @allure.feature('Interface 1') - @allure.story('Verifying Single Sign-On (SSO) Login for Interface 1') + @allure.feature("Interface 1") + @allure.story("Verifying Single Sign-On (SSO) Login for Interface 1") @allure.severity(allure.severity_level.CRITICAL) def test_verify_sign_in(self, driver, config): self.logger.info("Verifying Single Sign-On (SSO) Login for Interface 1") - driver.get(config['url_interface_1'] + "?code=0.AVsA57e_zrnCv0-sKQYbnoaPz5D--u_XfqNDqwMZoL4vF1hbAN8.AgABBAIAAAApTwJmzXqdR4BN2miheQMYAgDs_wUA9P-OZmUc_k67ay91q5x8iGq-FCMcsGOuN-RILqMWR3WzTHoAUItcREFppJowN22APhWP0_3aO-H-ByYLq_4toR4SkrJ83TIfavORbh2T-CllhzINKejOopV6Oj2TGqnANs2vp_vgpufdLpQiYcT5Cntcx6Zfz4TZCFrmFNjxoLox5FcXrNiwv_yTFyoiHaIUJG-cIaOn_oTqRGo7O-hUouN9JVxhmnU7ufPFdqEyvj566q9FddhHWKDUAh15f8nxWwRKqI65-vBXTiygJZboznnFjs4kQaXi7sRrmoT1uXt73lxr6FQe6bzlIWxHRWXci50qp4J_MU5gfZ1vSzJS5zFzFy7W2spTG7hByAH9LVv3us5Q_dAvr2_J3rn6zLJaZBnp_xRJ6cJM9e_zgcxL5zD6vVPtFkuqxWpV_Axtm2xrqTS7q03WRoDDwhRlHC02whsP8HVovurA-knXTrIRHM4p1eRSfK99EuKkev0eOGL5cNSKT-JNLQoe23R-JstAqhlyTSTis2wNK8I5Y2kd5esJGnBoe1jHeyTD5a-yppMfmBd6nW90_nvOc0eipFxMuWnhvYs56ex5Mv7PM6iWvnrP2ewVPbRI7oFa8WM_YhJ84Je4JRVw-HGi8Ftfv3pZfn8kxfxAp2lvncXw9GmCjm3eig-NdJqB1gpNGXN-lPZgMzjOwX5REWe3Wvqn2wK9sNd9gSzGTUCtHCN9NCkuSqkUAlAxECMy7FVuUl8i3CqNCkfrd7EI-LwczhRaOnKT9VJF0E544fAK1W5GWKMBGcrcrp8fJwOdAHqf8R7Mcsl3XzrafYzJo1SO&session_state=7515cdc9-6fd7-4e93-b526-f0a3c23401fd#") + driver.get( + config["url_interface_1"] + + "?code=0.AVsA57e_zrnCv0-sKQYbnoaPz5D--u_XfqNDqwMZoL4vF1hbAN8.AgABBAIAAAApTwJmzXqdR4BN2miheQMYAgDs_wUA9P-OZmUc_k67ay91q5x8iGq-FCMcsGOuN-RILqMWR3WzTHoAUItcREFppJowN22APhWP0_3aO-H-ByYLq_4toR4SkrJ83TIfavORbh2T-CllhzINKejOopV6Oj2TGqnANs2vp_vgpufdLpQiYcT5Cntcx6Zfz4TZCFrmFNjxoLox5FcXrNiwv_yTFyoiHaIUJG-cIaOn_oTqRGo7O-hUouN9JVxhmnU7ufPFdqEyvj566q9FddhHWKDUAh15f8nxWwRKqI65-vBXTiygJZboznnFjs4kQaXi7sRrmoT1uXt73lxr6FQe6bzlIWxHRWXci50qp4J_MU5gfZ1vSzJS5zFzFy7W2spTG7hByAH9LVv3us5Q_dAvr2_J3rn6zLJaZBnp_xRJ6cJM9e_zgcxL5zD6vVPtFkuqxWpV_Axtm2xrqTS7q03WRoDDwhRlHC02whsP8HVovurA-knXTrIRHM4p1eRSfK99EuKkev0eOGL5cNSKT-JNLQoe23R-JstAqhlyTSTis2wNK8I5Y2kd5esJGnBoe1jHeyTD5a-yppMfmBd6nW90_nvOc0eipFxMuWnhvYs56ex5Mv7PM6iWvnrP2ewVPbRI7oFa8WM_YhJ84Je4JRVw-HGi8Ftfv3pZfn8kxfxAp2lvncXw9GmCjm3eig-NdJqB1gpNGXN-lPZgMzjOwX5REWe3Wvqn2wK9sNd9gSzGTUCtHCN9NCkuSqkUAlAxECMy7FVuUl8i3CqNCkfrd7EI-LwczhRaOnKT9VJF0E544fAK1W5GWKMBGcrcrp8fJwOdAHqf8R7Mcsl3XzrafYzJo1SO&session_state=7515cdc9-6fd7-4e93-b526-f0a3c23401fd#" + ) # sign_in(driver, config) self.logger.info("Signed in successful") get_title(driver, config, "interface_1 · Streamlit") - if driver.title == 'interface_1 · Streamlit': + if driver.title == "interface_1 · Streamlit": self.logger.info("Interface 1 page opened successfully") assert True else: @@ -55,24 +58,29 @@ class Test_interface_1: @allure.suite("Tests for Interface 1") @allure.sub_suite("Tests for Read the contracts from Path") @allure.title('Verify the "Read the contracts from Path" button functionality') - @allure.feature('Interface 1') + @allure.feature("Interface 1") @allure.story('Verify the "Read the contracts from Path" button functionality') @allure.severity(allure.severity_level.CRITICAL) def test_read_contract_from_path(self, driver, config, groups): - self.logger.info('Verify the "Read the contracts from Path" button functionality') - driver.get(config['url_interface_1'] + "?code=0.AVsA57e_zrnCv0-sKQYbnoaPz5D--u_XfqNDqwMZoL4vF1hbAN8.AgABBAIAAAApTwJmzXqdR4BN2miheQMYAgDs_wUA9P-OZmUc_k67ay91q5x8iGq-FCMcsGOuN-RILqMWR3WzTHoAUItcREFppJowN22APhWP0_3aO-H-ByYLq_4toR4SkrJ83TIfavORbh2T-CllhzINKejOopV6Oj2TGqnANs2vp_vgpufdLpQiYcT5Cntcx6Zfz4TZCFrmFNjxoLox5FcXrNiwv_yTFyoiHaIUJG-cIaOn_oTqRGo7O-hUouN9JVxhmnU7ufPFdqEyvj566q9FddhHWKDUAh15f8nxWwRKqI65-vBXTiygJZboznnFjs4kQaXi7sRrmoT1uXt73lxr6FQe6bzlIWxHRWXci50qp4J_MU5gfZ1vSzJS5zFzFy7W2spTG7hByAH9LVv3us5Q_dAvr2_J3rn6zLJaZBnp_xRJ6cJM9e_zgcxL5zD6vVPtFkuqxWpV_Axtm2xrqTS7q03WRoDDwhRlHC02whsP8HVovurA-knXTrIRHM4p1eRSfK99EuKkev0eOGL5cNSKT-JNLQoe23R-JstAqhlyTSTis2wNK8I5Y2kd5esJGnBoe1jHeyTD5a-yppMfmBd6nW90_nvOc0eipFxMuWnhvYs56ex5Mv7PM6iWvnrP2ewVPbRI7oFa8WM_YhJ84Je4JRVw-HGi8Ftfv3pZfn8kxfxAp2lvncXw9GmCjm3eig-NdJqB1gpNGXN-lPZgMzjOwX5REWe3Wvqn2wK9sNd9gSzGTUCtHCN9NCkuSqkUAlAxECMy7FVuUl8i3CqNCkfrd7EI-LwczhRaOnKT9VJF0E544fAK1W5GWKMBGcrcrp8fJwOdAHqf8R7Mcsl3XzrafYzJo1SO&session_state=7515cdc9-6fd7-4e93-b526-f0a3c23401fd#") + self.logger.info( + 'Verify the "Read the contracts from Path" button functionality' + ) + driver.get( + config["url_interface_1"] + + "?code=0.AVsA57e_zrnCv0-sKQYbnoaPz5D--u_XfqNDqwMZoL4vF1hbAN8.AgABBAIAAAApTwJmzXqdR4BN2miheQMYAgDs_wUA9P-OZmUc_k67ay91q5x8iGq-FCMcsGOuN-RILqMWR3WzTHoAUItcREFppJowN22APhWP0_3aO-H-ByYLq_4toR4SkrJ83TIfavORbh2T-CllhzINKejOopV6Oj2TGqnANs2vp_vgpufdLpQiYcT5Cntcx6Zfz4TZCFrmFNjxoLox5FcXrNiwv_yTFyoiHaIUJG-cIaOn_oTqRGo7O-hUouN9JVxhmnU7ufPFdqEyvj566q9FddhHWKDUAh15f8nxWwRKqI65-vBXTiygJZboznnFjs4kQaXi7sRrmoT1uXt73lxr6FQe6bzlIWxHRWXci50qp4J_MU5gfZ1vSzJS5zFzFy7W2spTG7hByAH9LVv3us5Q_dAvr2_J3rn6zLJaZBnp_xRJ6cJM9e_zgcxL5zD6vVPtFkuqxWpV_Axtm2xrqTS7q03WRoDDwhRlHC02whsP8HVovurA-knXTrIRHM4p1eRSfK99EuKkev0eOGL5cNSKT-JNLQoe23R-JstAqhlyTSTis2wNK8I5Y2kd5esJGnBoe1jHeyTD5a-yppMfmBd6nW90_nvOc0eipFxMuWnhvYs56ex5Mv7PM6iWvnrP2ewVPbRI7oFa8WM_YhJ84Je4JRVw-HGi8Ftfv3pZfn8kxfxAp2lvncXw9GmCjm3eig-NdJqB1gpNGXN-lPZgMzjOwX5REWe3Wvqn2wK9sNd9gSzGTUCtHCN9NCkuSqkUAlAxECMy7FVuUl8i3CqNCkfrd7EI-LwczhRaOnKT9VJF0E544fAK1W5GWKMBGcrcrp8fJwOdAHqf8R7Mcsl3XzrafYzJo1SO&session_state=7515cdc9-6fd7-4e93-b526-f0a3c23401fd#" + ) with allure.step("Signing in"): # sign_in(driver, config) self.logger.debug("Signed in successful") get_title(driver, config, "interface_1 · Streamlit") - with allure.step('Selecting Client'): + with allure.step("Selecting Client"): interface_1 = Interface_1(driver, config) interface_1.select_client("Centene") - with allure.step('Selecting Batch'): + with allure.step("Selecting Batch"): interface_1.select_batch("batch_190624022330") - with allure.step('Selecting Groups'): + with allure.step("Selecting Groups"): interface_1.check_group_no(groups) - with allure.step('Reading Contracts from S3 Bucket'): + with allure.step("Reading Contracts from S3 Bucket"): interface_1.click_read_the_contracts_from_path() interface_1.check_table_contain() @@ -81,12 +89,15 @@ class Test_interface_1: @allure.suite("Tests for Interface 1") @allure.sub_suite("Tests for Run Doczy.AI Pipeline") @allure.title('Verify "Run Doczy.AI Pipeline" button functionality') - @allure.feature('Interface 1') + @allure.feature("Interface 1") @allure.story('Verify "Run Doczy.AI Pipeline" button functionality') @allure.severity(allure.severity_level.CRITICAL) def test_run_doczy_ai_pipeline(self, driver, config, groups): self.logger.info('Verify "Run Doczy.AI Pipeline" button functionality') - driver.get(config['url_interface_1'] + "?code=0.AVsA57e_zrnCv0-sKQYbnoaPz5D--u_XfqNDqwMZoL4vF1hbAN8.AgABBAIAAAApTwJmzXqdR4BN2miheQMYAgDs_wUA9P9xUVp2PJ6-Ymm_h4pBRlXjhIYg3cEOpiosJGKkG8EF0WMNT10RyUeKXqrs3UDMkLlTCZy3qMPFePktnSIkLU7AmdUv_QXjPc221BMScHwImuvMxH5YhN8s8jeJuS9GWMnllIbKInqHn_DNzL5uBSMVCd3lx1YEtHTr2Bn_MY4Vp9TS7Z3byW3upWK77c6SE6Fsv4nBTTab_4ebD_cLAanFoY2BPZdKXPngveBjjxDPUcAXD3JvRwUGHPUR8Iv3sRip6lqa02SsXt0_UzhypRYRmP67P52D_xhGjfQ--FuVF_M87UTnFK4FeNVHSeU6SRKJA9VWG9QDnizP-xiRB_8cnfmaaBw2UWnpbXC1o_EnsWPbjyomHylkPTc69gygghWBOGwR8agZA3B8sWjz-pMJO8E7hWpOIBlsEeRLJieXgBZxdRG77JZ4iUtJZZ0Dd-1VFxQCh8bSUl9TBdHA8wzui-dqNd15xKDHtS7FA3BHGZOICrYcKeiaShjC2OhQl71oUOtF8o31zsyUvq9jeTCL64X1l_xL8ePDqzinfJeaN--mPea3QZHkd3P6jOGdT4cTD9u-RBTD-Qbu6dj5egQDgUSPCcfL7kVeSz0QSi8jpqUoxJi4RTOXZEFLX77n88kUJczMsUsSQSdzTaKBDXY5HQwFK79jE5uJ7hwdBcEtKBHv0zoE0pn8RogLwzKk6SG-6R94UAuvuYw-ivkW8HyCWbGHQq_QSj5YdMEyBO2C-woCH5qwFbN4qoTc-OabOvW4lHTfFyeoJB7ryUdHilMd1kdAPe6WNNRKCmLb0Ofn_Jv5eqFID6L7DtmEJCc&session_state=7515cdc9-6fd7-4e93-b526-f0a3c23401fd#") + driver.get( + config["url_interface_1"] + + "?code=0.AVsA57e_zrnCv0-sKQYbnoaPz5D--u_XfqNDqwMZoL4vF1hbAN8.AgABBAIAAAApTwJmzXqdR4BN2miheQMYAgDs_wUA9P9xUVp2PJ6-Ymm_h4pBRlXjhIYg3cEOpiosJGKkG8EF0WMNT10RyUeKXqrs3UDMkLlTCZy3qMPFePktnSIkLU7AmdUv_QXjPc221BMScHwImuvMxH5YhN8s8jeJuS9GWMnllIbKInqHn_DNzL5uBSMVCd3lx1YEtHTr2Bn_MY4Vp9TS7Z3byW3upWK77c6SE6Fsv4nBTTab_4ebD_cLAanFoY2BPZdKXPngveBjjxDPUcAXD3JvRwUGHPUR8Iv3sRip6lqa02SsXt0_UzhypRYRmP67P52D_xhGjfQ--FuVF_M87UTnFK4FeNVHSeU6SRKJA9VWG9QDnizP-xiRB_8cnfmaaBw2UWnpbXC1o_EnsWPbjyomHylkPTc69gygghWBOGwR8agZA3B8sWjz-pMJO8E7hWpOIBlsEeRLJieXgBZxdRG77JZ4iUtJZZ0Dd-1VFxQCh8bSUl9TBdHA8wzui-dqNd15xKDHtS7FA3BHGZOICrYcKeiaShjC2OhQl71oUOtF8o31zsyUvq9jeTCL64X1l_xL8ePDqzinfJeaN--mPea3QZHkd3P6jOGdT4cTD9u-RBTD-Qbu6dj5egQDgUSPCcfL7kVeSz0QSi8jpqUoxJi4RTOXZEFLX77n88kUJczMsUsSQSdzTaKBDXY5HQwFK79jE5uJ7hwdBcEtKBHv0zoE0pn8RogLwzKk6SG-6R94UAuvuYw-ivkW8HyCWbGHQq_QSj5YdMEyBO2C-woCH5qwFbN4qoTc-OabOvW4lHTfFyeoJB7ryUdHilMd1kdAPe6WNNRKCmLb0Ofn_Jv5eqFID6L7DtmEJCc&session_state=7515cdc9-6fd7-4e93-b526-f0a3c23401fd#" + ) # sign_in(driver, config) self.logger.info("Signed in successful") get_title(driver, config, "interface_1 · Streamlit") @@ -108,13 +119,22 @@ class Test_interface_1: @allure.parent_suite("Tests for web interface") @allure.suite("Tests for Interface 1") @allure.sub_suite("Tests for Run Doczy.AI Pipeline with no Groups") - @allure.title('Verify if none Group No is selected from Group No Checkbox or from table then "Run Doczy.AI Pipeline"should not be functional') - @allure.feature('Interface 1') - @allure.story('Verify if none Group No is selected from Group No Checkbox or from table then "Run Doczy.AI Pipeline"should not be functional') + @allure.title( + 'Verify if none Group No is selected from Group No Checkbox or from table then "Run Doczy.AI Pipeline"should not be functional' + ) + @allure.feature("Interface 1") + @allure.story( + 'Verify if none Group No is selected from Group No Checkbox or from table then "Run Doczy.AI Pipeline"should not be functional' + ) @allure.severity(allure.severity_level.CRITICAL) def test_run_doczy_ai_pipeline_without_selecting_group_no(self, driver, config): - self.logger.info('Verify if none Group No is selected from Group No Checkbox or from table then "Run Doczy.AI Pipeline"should not be functional') - driver.get(config['url_interface_1'] + "?code=0.AVsA57e_zrnCv0-sKQYbnoaPz5D--u_XfqNDqwMZoL4vF1hbAN8.AgABBAIAAAApTwJmzXqdR4BN2miheQMYAgDs_wUA9P-OZmUc_k67ay91q5x8iGq-FCMcsGOuN-RILqMWR3WzTHoAUItcREFppJowN22APhWP0_3aO-H-ByYLq_4toR4SkrJ83TIfavORbh2T-CllhzINKejOopV6Oj2TGqnANs2vp_vgpufdLpQiYcT5Cntcx6Zfz4TZCFrmFNjxoLox5FcXrNiwv_yTFyoiHaIUJG-cIaOn_oTqRGo7O-hUouN9JVxhmnU7ufPFdqEyvj566q9FddhHWKDUAh15f8nxWwRKqI65-vBXTiygJZboznnFjs4kQaXi7sRrmoT1uXt73lxr6FQe6bzlIWxHRWXci50qp4J_MU5gfZ1vSzJS5zFzFy7W2spTG7hByAH9LVv3us5Q_dAvr2_J3rn6zLJaZBnp_xRJ6cJM9e_zgcxL5zD6vVPtFkuqxWpV_Axtm2xrqTS7q03WRoDDwhRlHC02whsP8HVovurA-knXTrIRHM4p1eRSfK99EuKkev0eOGL5cNSKT-JNLQoe23R-JstAqhlyTSTis2wNK8I5Y2kd5esJGnBoe1jHeyTD5a-yppMfmBd6nW90_nvOc0eipFxMuWnhvYs56ex5Mv7PM6iWvnrP2ewVPbRI7oFa8WM_YhJ84Je4JRVw-HGi8Ftfv3pZfn8kxfxAp2lvncXw9GmCjm3eig-NdJqB1gpNGXN-lPZgMzjOwX5REWe3Wvqn2wK9sNd9gSzGTUCtHCN9NCkuSqkUAlAxECMy7FVuUl8i3CqNCkfrd7EI-LwczhRaOnKT9VJF0E544fAK1W5GWKMBGcrcrp8fJwOdAHqf8R7Mcsl3XzrafYzJo1SO&session_state=7515cdc9-6fd7-4e93-b526-f0a3c23401fd#") + self.logger.info( + 'Verify if none Group No is selected from Group No Checkbox or from table then "Run Doczy.AI Pipeline"should not be functional' + ) + driver.get( + config["url_interface_1"] + + "?code=0.AVsA57e_zrnCv0-sKQYbnoaPz5D--u_XfqNDqwMZoL4vF1hbAN8.AgABBAIAAAApTwJmzXqdR4BN2miheQMYAgDs_wUA9P-OZmUc_k67ay91q5x8iGq-FCMcsGOuN-RILqMWR3WzTHoAUItcREFppJowN22APhWP0_3aO-H-ByYLq_4toR4SkrJ83TIfavORbh2T-CllhzINKejOopV6Oj2TGqnANs2vp_vgpufdLpQiYcT5Cntcx6Zfz4TZCFrmFNjxoLox5FcXrNiwv_yTFyoiHaIUJG-cIaOn_oTqRGo7O-hUouN9JVxhmnU7ufPFdqEyvj566q9FddhHWKDUAh15f8nxWwRKqI65-vBXTiygJZboznnFjs4kQaXi7sRrmoT1uXt73lxr6FQe6bzlIWxHRWXci50qp4J_MU5gfZ1vSzJS5zFzFy7W2spTG7hByAH9LVv3us5Q_dAvr2_J3rn6zLJaZBnp_xRJ6cJM9e_zgcxL5zD6vVPtFkuqxWpV_Axtm2xrqTS7q03WRoDDwhRlHC02whsP8HVovurA-knXTrIRHM4p1eRSfK99EuKkev0eOGL5cNSKT-JNLQoe23R-JstAqhlyTSTis2wNK8I5Y2kd5esJGnBoe1jHeyTD5a-yppMfmBd6nW90_nvOc0eipFxMuWnhvYs56ex5Mv7PM6iWvnrP2ewVPbRI7oFa8WM_YhJ84Je4JRVw-HGi8Ftfv3pZfn8kxfxAp2lvncXw9GmCjm3eig-NdJqB1gpNGXN-lPZgMzjOwX5REWe3Wvqn2wK9sNd9gSzGTUCtHCN9NCkuSqkUAlAxECMy7FVuUl8i3CqNCkfrd7EI-LwczhRaOnKT9VJF0E544fAK1W5GWKMBGcrcrp8fJwOdAHqf8R7Mcsl3XzrafYzJo1SO&session_state=7515cdc9-6fd7-4e93-b526-f0a3c23401fd#" + ) # sign_in(driver, config) self.logger.info("Signed in successful") get_title(driver, config, "interface_1 · Streamlit") diff --git a/Doczy.AI_Automation/tests/test_interface_2.py b/Doczy.AI_Automation/tests/test_interface_2.py index 02738a1..8852b38 100644 --- a/Doczy.AI_Automation/tests/test_interface_2.py +++ b/Doczy.AI_Automation/tests/test_interface_2.py @@ -14,14 +14,14 @@ class Test_interface_2: @allure.suite("Tests for Interface 2") @allure.sub_suite("Tests for Navigation") @allure.title("Verifying designated Interface 2 Page opens successfully") - @allure.feature('Interface 2') - @allure.story('Verifying designated Interface 2 Page opens successfully') + @allure.feature("Interface 2") + @allure.story("Verifying designated Interface 2 Page opens successfully") @allure.severity(allure.severity_level.CRITICAL) def test_validate_navigation(self, driver, config): self.logger.info("Verifying designated Interface 2 Page opens successfully") - driver.get(config['url_interface_2']) + driver.get(config["url_interface_2"]) get_title(driver, config, "interface_2 · Streamlit") - if driver.title == 'interface_2 · Streamlit': + if driver.title == "interface_2 · Streamlit": self.logger.info("Designated Interface 2 Page opened successfully") assert True else: @@ -33,16 +33,19 @@ class Test_interface_2: @allure.suite("Tests for Interface 2") @allure.sub_suite("Tests for Authentication") @allure.title("Verifying Single Sign-On (SSO) Login for Interface 2") - @allure.feature('Interface 2') - @allure.story('Verifying Single Sign-On (SSO) Login for Interface 2') + @allure.feature("Interface 2") + @allure.story("Verifying Single Sign-On (SSO) Login for Interface 2") @allure.severity(allure.severity_level.CRITICAL) def test_verify_sign_in(self, driver, config): self.logger.info("Verifying Single Sign-On (SSO) Login for Interface 2") - driver.get(config['url_interface_2'] + "?code=0.AVsA57e_zrnCv0-sKQYbnoaPz5D--u_XfqNDqwMZoL4vF1hbAN8.AgABBAIAAAApTwJmzXqdR4BN2miheQMYAgDs_wUA9P91FtoYFvR20dqHG0dEKQhDe1BftWl3mQWsfh9s8UJttZdIVkp4cYIkmKKhTtdGruEzTbeZjcCH6mii-K8QH5QDmQ9chVv-P5lGlTOPig_9AZ_6ZZEty4XWPFPJzseHk4JsvCk36kgebUZvfT2X2MfMEsRi-TnOC2CWhhlurm0a8mCR7WzDWHLzPHpFAwcd-GS7V3N8NUri6LDskw8E1HiOtMSs0SwwZCGw2A7GZV59MrYvNxi8YyousqNRLhdOg56PrzlhLk5l-hgm9ZX9NrKYmlNmb3LYFrxd8zgWQ5fdulyColwquQHTGZ4kRwM49xCdTEsCDprDYqopmMoPXcwvjFz69f_a9S-qT8jNhFGN_WmdoUlq02J0BSLnD-IfKNAQe5w2D6SKCYH4LLZnSSKpQrJUFosqYmxOsk34Z80enm-f4kiJhzqYt9Qwa6n70J3bF-6vXCjZdu3l5iI5CLc7QcaicPmfiOQBGu_EyFploqsW7fmqfNhJ92_W0jL_yQvgZaBtMeZH23BmVpYQGdNeBhRauBj4qaHpTkPHSUPBw26atNtq5yWOs0j6gIL_hhHljLqqvMmI3vw0B3deVxqy2MvgU40wkYYfETwaBaIVTH5OtO60EBULwN7PBCYDVOT2ZskBZnQCp29IZqam1GwZHDBHEVbw9djRXl9ZwKOUHa9mv_dKHJAB_cMtpAzOF0H6fnyzLr1kZEd-hfHUY-ZF1mrEPgUBKC5swYTfgYpBpgWjnlxMKva9tC1iHAwKSyRcsRc7BnSzBJv9ZU8UhzrhIlJ1ILGjjJQcUwicfVIokN9X79fxllkkqQ&session_state=7515cdc9-6fd7-4e93-b526-f0a3c23401fd#") + driver.get( + config["url_interface_2"] + + "?code=0.AVsA57e_zrnCv0-sKQYbnoaPz5D--u_XfqNDqwMZoL4vF1hbAN8.AgABBAIAAAApTwJmzXqdR4BN2miheQMYAgDs_wUA9P91FtoYFvR20dqHG0dEKQhDe1BftWl3mQWsfh9s8UJttZdIVkp4cYIkmKKhTtdGruEzTbeZjcCH6mii-K8QH5QDmQ9chVv-P5lGlTOPig_9AZ_6ZZEty4XWPFPJzseHk4JsvCk36kgebUZvfT2X2MfMEsRi-TnOC2CWhhlurm0a8mCR7WzDWHLzPHpFAwcd-GS7V3N8NUri6LDskw8E1HiOtMSs0SwwZCGw2A7GZV59MrYvNxi8YyousqNRLhdOg56PrzlhLk5l-hgm9ZX9NrKYmlNmb3LYFrxd8zgWQ5fdulyColwquQHTGZ4kRwM49xCdTEsCDprDYqopmMoPXcwvjFz69f_a9S-qT8jNhFGN_WmdoUlq02J0BSLnD-IfKNAQe5w2D6SKCYH4LLZnSSKpQrJUFosqYmxOsk34Z80enm-f4kiJhzqYt9Qwa6n70J3bF-6vXCjZdu3l5iI5CLc7QcaicPmfiOQBGu_EyFploqsW7fmqfNhJ92_W0jL_yQvgZaBtMeZH23BmVpYQGdNeBhRauBj4qaHpTkPHSUPBw26atNtq5yWOs0j6gIL_hhHljLqqvMmI3vw0B3deVxqy2MvgU40wkYYfETwaBaIVTH5OtO60EBULwN7PBCYDVOT2ZskBZnQCp29IZqam1GwZHDBHEVbw9djRXl9ZwKOUHa9mv_dKHJAB_cMtpAzOF0H6fnyzLr1kZEd-hfHUY-ZF1mrEPgUBKC5swYTfgYpBpgWjnlxMKva9tC1iHAwKSyRcsRc7BnSzBJv9ZU8UhzrhIlJ1ILGjjJQcUwicfVIokN9X79fxllkkqQ&session_state=7515cdc9-6fd7-4e93-b526-f0a3c23401fd#" + ) # sign_in(driver, config) self.logger.info("Signed in successful") get_title(driver, config, "interface_2 · Streamlit") - if driver.title == 'interface_2 · Streamlit': + if driver.title == "interface_2 · Streamlit": self.logger.info("Interface 2 page opened successfully") assert True else: @@ -53,14 +56,16 @@ class Test_interface_2: @allure.suite("Tests for Interface 2") @allure.sub_suite("Tests for Show Result button") @allure.title("Verify the 'show result' button functionality") - @allure.feature('Interface 2') + @allure.feature("Interface 2") @allure.story("Verify the 'show result' button functionality") @allure.severity(allure.severity_level.CRITICAL) @allure.testcase("TC-008") def test_verify_show_results_button(self, driver, config): self.logger.info("Verify the 'show result' button functionality") - driver.get(config[ - 'url_interface_2'] + "?code=0.AVsA57e_zrnCv0-sKQYbnoaPz5D--u_XfqNDqwMZoL4vF1hbAN8.AgABBAIAAAApTwJmzXqdR4BN2miheQMYAgDs_wUA9P91FtoYFvR20dqHG0dEKQhDe1BftWl3mQWsfh9s8UJttZdIVkp4cYIkmKKhTtdGruEzTbeZjcCH6mii-K8QH5QDmQ9chVv-P5lGlTOPig_9AZ_6ZZEty4XWPFPJzseHk4JsvCk36kgebUZvfT2X2MfMEsRi-TnOC2CWhhlurm0a8mCR7WzDWHLzPHpFAwcd-GS7V3N8NUri6LDskw8E1HiOtMSs0SwwZCGw2A7GZV59MrYvNxi8YyousqNRLhdOg56PrzlhLk5l-hgm9ZX9NrKYmlNmb3LYFrxd8zgWQ5fdulyColwquQHTGZ4kRwM49xCdTEsCDprDYqopmMoPXcwvjFz69f_a9S-qT8jNhFGN_WmdoUlq02J0BSLnD-IfKNAQe5w2D6SKCYH4LLZnSSKpQrJUFosqYmxOsk34Z80enm-f4kiJhzqYt9Qwa6n70J3bF-6vXCjZdu3l5iI5CLc7QcaicPmfiOQBGu_EyFploqsW7fmqfNhJ92_W0jL_yQvgZaBtMeZH23BmVpYQGdNeBhRauBj4qaHpTkPHSUPBw26atNtq5yWOs0j6gIL_hhHljLqqvMmI3vw0B3deVxqy2MvgU40wkYYfETwaBaIVTH5OtO60EBULwN7PBCYDVOT2ZskBZnQCp29IZqam1GwZHDBHEVbw9djRXl9ZwKOUHa9mv_dKHJAB_cMtpAzOF0H6fnyzLr1kZEd-hfHUY-ZF1mrEPgUBKC5swYTfgYpBpgWjnlxMKva9tC1iHAwKSyRcsRc7BnSzBJv9ZU8UhzrhIlJ1ILGjjJQcUwicfVIokN9X79fxllkkqQ&session_state=7515cdc9-6fd7-4e93-b526-f0a3c23401fd#") + driver.get( + config["url_interface_2"] + + "?code=0.AVsA57e_zrnCv0-sKQYbnoaPz5D--u_XfqNDqwMZoL4vF1hbAN8.AgABBAIAAAApTwJmzXqdR4BN2miheQMYAgDs_wUA9P91FtoYFvR20dqHG0dEKQhDe1BftWl3mQWsfh9s8UJttZdIVkp4cYIkmKKhTtdGruEzTbeZjcCH6mii-K8QH5QDmQ9chVv-P5lGlTOPig_9AZ_6ZZEty4XWPFPJzseHk4JsvCk36kgebUZvfT2X2MfMEsRi-TnOC2CWhhlurm0a8mCR7WzDWHLzPHpFAwcd-GS7V3N8NUri6LDskw8E1HiOtMSs0SwwZCGw2A7GZV59MrYvNxi8YyousqNRLhdOg56PrzlhLk5l-hgm9ZX9NrKYmlNmb3LYFrxd8zgWQ5fdulyColwquQHTGZ4kRwM49xCdTEsCDprDYqopmMoPXcwvjFz69f_a9S-qT8jNhFGN_WmdoUlq02J0BSLnD-IfKNAQe5w2D6SKCYH4LLZnSSKpQrJUFosqYmxOsk34Z80enm-f4kiJhzqYt9Qwa6n70J3bF-6vXCjZdu3l5iI5CLc7QcaicPmfiOQBGu_EyFploqsW7fmqfNhJ92_W0jL_yQvgZaBtMeZH23BmVpYQGdNeBhRauBj4qaHpTkPHSUPBw26atNtq5yWOs0j6gIL_hhHljLqqvMmI3vw0B3deVxqy2MvgU40wkYYfETwaBaIVTH5OtO60EBULwN7PBCYDVOT2ZskBZnQCp29IZqam1GwZHDBHEVbw9djRXl9ZwKOUHa9mv_dKHJAB_cMtpAzOF0H6fnyzLr1kZEd-hfHUY-ZF1mrEPgUBKC5swYTfgYpBpgWjnlxMKva9tC1iHAwKSyRcsRc7BnSzBJv9ZU8UhzrhIlJ1ILGjjJQcUwicfVIokN9X79fxllkkqQ&session_state=7515cdc9-6fd7-4e93-b526-f0a3c23401fd#" + ) with allure.step("Signing in"): # sign_in(driver, config) self.logger.info("Signed in successful") @@ -81,14 +86,22 @@ class Test_interface_2: @allure.parent_suite("Tests for web interface") @allure.suite("Tests for Interface 2") @allure.sub_suite("Tests for Show Result button without Contract") - @allure.title("Verify the 'show result' button functionality without selecting contracts") - @allure.feature('Interface 2') - @allure.story("Verify the 'show result' button functionality without selecting contracts") + @allure.title( + "Verify the 'show result' button functionality without selecting contracts" + ) + @allure.feature("Interface 2") + @allure.story( + "Verify the 'show result' button functionality without selecting contracts" + ) @allure.severity(allure.severity_level.CRITICAL) def test_verify_show_results_button_without_contract(self, driver, config): - self.logger.info("Verify the 'show result' button functionality without selecting contracts") - driver.get(config[ - 'url_interface_2'] + "?code=0.AVsA57e_zrnCv0-sKQYbnoaPz5D--u_XfqNDqwMZoL4vF1hbAN8.AgABBAIAAAApTwJmzXqdR4BN2miheQMYAgDs_wUA9P91FtoYFvR20dqHG0dEKQhDe1BftWl3mQWsfh9s8UJttZdIVkp4cYIkmKKhTtdGruEzTbeZjcCH6mii-K8QH5QDmQ9chVv-P5lGlTOPig_9AZ_6ZZEty4XWPFPJzseHk4JsvCk36kgebUZvfT2X2MfMEsRi-TnOC2CWhhlurm0a8mCR7WzDWHLzPHpFAwcd-GS7V3N8NUri6LDskw8E1HiOtMSs0SwwZCGw2A7GZV59MrYvNxi8YyousqNRLhdOg56PrzlhLk5l-hgm9ZX9NrKYmlNmb3LYFrxd8zgWQ5fdulyColwquQHTGZ4kRwM49xCdTEsCDprDYqopmMoPXcwvjFz69f_a9S-qT8jNhFGN_WmdoUlq02J0BSLnD-IfKNAQe5w2D6SKCYH4LLZnSSKpQrJUFosqYmxOsk34Z80enm-f4kiJhzqYt9Qwa6n70J3bF-6vXCjZdu3l5iI5CLc7QcaicPmfiOQBGu_EyFploqsW7fmqfNhJ92_W0jL_yQvgZaBtMeZH23BmVpYQGdNeBhRauBj4qaHpTkPHSUPBw26atNtq5yWOs0j6gIL_hhHljLqqvMmI3vw0B3deVxqy2MvgU40wkYYfETwaBaIVTH5OtO60EBULwN7PBCYDVOT2ZskBZnQCp29IZqam1GwZHDBHEVbw9djRXl9ZwKOUHa9mv_dKHJAB_cMtpAzOF0H6fnyzLr1kZEd-hfHUY-ZF1mrEPgUBKC5swYTfgYpBpgWjnlxMKva9tC1iHAwKSyRcsRc7BnSzBJv9ZU8UhzrhIlJ1ILGjjJQcUwicfVIokN9X79fxllkkqQ&session_state=7515cdc9-6fd7-4e93-b526-f0a3c23401fd#") + self.logger.info( + "Verify the 'show result' button functionality without selecting contracts" + ) + driver.get( + config["url_interface_2"] + + "?code=0.AVsA57e_zrnCv0-sKQYbnoaPz5D--u_XfqNDqwMZoL4vF1hbAN8.AgABBAIAAAApTwJmzXqdR4BN2miheQMYAgDs_wUA9P91FtoYFvR20dqHG0dEKQhDe1BftWl3mQWsfh9s8UJttZdIVkp4cYIkmKKhTtdGruEzTbeZjcCH6mii-K8QH5QDmQ9chVv-P5lGlTOPig_9AZ_6ZZEty4XWPFPJzseHk4JsvCk36kgebUZvfT2X2MfMEsRi-TnOC2CWhhlurm0a8mCR7WzDWHLzPHpFAwcd-GS7V3N8NUri6LDskw8E1HiOtMSs0SwwZCGw2A7GZV59MrYvNxi8YyousqNRLhdOg56PrzlhLk5l-hgm9ZX9NrKYmlNmb3LYFrxd8zgWQ5fdulyColwquQHTGZ4kRwM49xCdTEsCDprDYqopmMoPXcwvjFz69f_a9S-qT8jNhFGN_WmdoUlq02J0BSLnD-IfKNAQe5w2D6SKCYH4LLZnSSKpQrJUFosqYmxOsk34Z80enm-f4kiJhzqYt9Qwa6n70J3bF-6vXCjZdu3l5iI5CLc7QcaicPmfiOQBGu_EyFploqsW7fmqfNhJ92_W0jL_yQvgZaBtMeZH23BmVpYQGdNeBhRauBj4qaHpTkPHSUPBw26atNtq5yWOs0j6gIL_hhHljLqqvMmI3vw0B3deVxqy2MvgU40wkYYfETwaBaIVTH5OtO60EBULwN7PBCYDVOT2ZskBZnQCp29IZqam1GwZHDBHEVbw9djRXl9ZwKOUHa9mv_dKHJAB_cMtpAzOF0H6fnyzLr1kZEd-hfHUY-ZF1mrEPgUBKC5swYTfgYpBpgWjnlxMKva9tC1iHAwKSyRcsRc7BnSzBJv9ZU8UhzrhIlJ1ILGjjJQcUwicfVIokN9X79fxllkkqQ&session_state=7515cdc9-6fd7-4e93-b526-f0a3c23401fd#" + ) with allure.step("Signing in"): # sign_in(driver, config) self.logger.info("Signed in successful") @@ -107,14 +120,22 @@ class Test_interface_2: @allure.parent_suite("Tests for web interface") @allure.suite("Tests for Interface 2") @allure.sub_suite("Tests for Show Result button without Field Groups") - @allure.title("Verify the 'show result' button functionality without selecting field groups") - @allure.feature('Interface 2') - @allure.story("Verify the 'show result' button functionality without selecting field groups") + @allure.title( + "Verify the 'show result' button functionality without selecting field groups" + ) + @allure.feature("Interface 2") + @allure.story( + "Verify the 'show result' button functionality without selecting field groups" + ) @allure.severity(allure.severity_level.CRITICAL) def test_verify_show_results_button_without_field_group(self, driver, config): - self.logger.info("Verify the 'show result' button functionality without selecting field groups") - driver.get(config[ - 'url_interface_2'] + "?code=0.AVsA57e_zrnCv0-sKQYbnoaPz5D--u_XfqNDqwMZoL4vF1hbAN8.AgABBAIAAAApTwJmzXqdR4BN2miheQMYAgDs_wUA9P91FtoYFvR20dqHG0dEKQhDe1BftWl3mQWsfh9s8UJttZdIVkp4cYIkmKKhTtdGruEzTbeZjcCH6mii-K8QH5QDmQ9chVv-P5lGlTOPig_9AZ_6ZZEty4XWPFPJzseHk4JsvCk36kgebUZvfT2X2MfMEsRi-TnOC2CWhhlurm0a8mCR7WzDWHLzPHpFAwcd-GS7V3N8NUri6LDskw8E1HiOtMSs0SwwZCGw2A7GZV59MrYvNxi8YyousqNRLhdOg56PrzlhLk5l-hgm9ZX9NrKYmlNmb3LYFrxd8zgWQ5fdulyColwquQHTGZ4kRwM49xCdTEsCDprDYqopmMoPXcwvjFz69f_a9S-qT8jNhFGN_WmdoUlq02J0BSLnD-IfKNAQe5w2D6SKCYH4LLZnSSKpQrJUFosqYmxOsk34Z80enm-f4kiJhzqYt9Qwa6n70J3bF-6vXCjZdu3l5iI5CLc7QcaicPmfiOQBGu_EyFploqsW7fmqfNhJ92_W0jL_yQvgZaBtMeZH23BmVpYQGdNeBhRauBj4qaHpTkPHSUPBw26atNtq5yWOs0j6gIL_hhHljLqqvMmI3vw0B3deVxqy2MvgU40wkYYfETwaBaIVTH5OtO60EBULwN7PBCYDVOT2ZskBZnQCp29IZqam1GwZHDBHEVbw9djRXl9ZwKOUHa9mv_dKHJAB_cMtpAzOF0H6fnyzLr1kZEd-hfHUY-ZF1mrEPgUBKC5swYTfgYpBpgWjnlxMKva9tC1iHAwKSyRcsRc7BnSzBJv9ZU8UhzrhIlJ1ILGjjJQcUwicfVIokN9X79fxllkkqQ&session_state=7515cdc9-6fd7-4e93-b526-f0a3c23401fd#") + self.logger.info( + "Verify the 'show result' button functionality without selecting field groups" + ) + driver.get( + config["url_interface_2"] + + "?code=0.AVsA57e_zrnCv0-sKQYbnoaPz5D--u_XfqNDqwMZoL4vF1hbAN8.AgABBAIAAAApTwJmzXqdR4BN2miheQMYAgDs_wUA9P91FtoYFvR20dqHG0dEKQhDe1BftWl3mQWsfh9s8UJttZdIVkp4cYIkmKKhTtdGruEzTbeZjcCH6mii-K8QH5QDmQ9chVv-P5lGlTOPig_9AZ_6ZZEty4XWPFPJzseHk4JsvCk36kgebUZvfT2X2MfMEsRi-TnOC2CWhhlurm0a8mCR7WzDWHLzPHpFAwcd-GS7V3N8NUri6LDskw8E1HiOtMSs0SwwZCGw2A7GZV59MrYvNxi8YyousqNRLhdOg56PrzlhLk5l-hgm9ZX9NrKYmlNmb3LYFrxd8zgWQ5fdulyColwquQHTGZ4kRwM49xCdTEsCDprDYqopmMoPXcwvjFz69f_a9S-qT8jNhFGN_WmdoUlq02J0BSLnD-IfKNAQe5w2D6SKCYH4LLZnSSKpQrJUFosqYmxOsk34Z80enm-f4kiJhzqYt9Qwa6n70J3bF-6vXCjZdu3l5iI5CLc7QcaicPmfiOQBGu_EyFploqsW7fmqfNhJ92_W0jL_yQvgZaBtMeZH23BmVpYQGdNeBhRauBj4qaHpTkPHSUPBw26atNtq5yWOs0j6gIL_hhHljLqqvMmI3vw0B3deVxqy2MvgU40wkYYfETwaBaIVTH5OtO60EBULwN7PBCYDVOT2ZskBZnQCp29IZqam1GwZHDBHEVbw9djRXl9ZwKOUHa9mv_dKHJAB_cMtpAzOF0H6fnyzLr1kZEd-hfHUY-ZF1mrEPgUBKC5swYTfgYpBpgWjnlxMKva9tC1iHAwKSyRcsRc7BnSzBJv9ZU8UhzrhIlJ1ILGjjJQcUwicfVIokN9X79fxllkkqQ&session_state=7515cdc9-6fd7-4e93-b526-f0a3c23401fd#" + ) with allure.step("Signing in"): # sign_in(driver, config) self.logger.info("Signed in successful") @@ -134,13 +155,15 @@ class Test_interface_2: @allure.suite("Tests for Interface 2") @allure.sub_suite("Tests for Kickoff Database Integration") @allure.title("Verifying Kickoff Database Integration Functionality") - @allure.feature('Interface 2') - @allure.story('Verifying Kickoff Database Integration Functionality') + @allure.feature("Interface 2") + @allure.story("Verifying Kickoff Database Integration Functionality") @allure.severity(allure.severity_level.CRITICAL) def test_verify_Kickoff_Database_Integration(self, driver, config): self.logger.info("Verifying Kickoff Database Integration Functionality") - driver.get(config[ - 'url_interface_2'] + "?code=0.AVsA57e_zrnCv0-sKQYbnoaPz5D--u_XfqNDqwMZoL4vF1hbAN8.AgABBAIAAAApTwJmzXqdR4BN2miheQMYAgDs_wUA9P91FtoYFvR20dqHG0dEKQhDe1BftWl3mQWsfh9s8UJttZdIVkp4cYIkmKKhTtdGruEzTbeZjcCH6mii-K8QH5QDmQ9chVv-P5lGlTOPig_9AZ_6ZZEty4XWPFPJzseHk4JsvCk36kgebUZvfT2X2MfMEsRi-TnOC2CWhhlurm0a8mCR7WzDWHLzPHpFAwcd-GS7V3N8NUri6LDskw8E1HiOtMSs0SwwZCGw2A7GZV59MrYvNxi8YyousqNRLhdOg56PrzlhLk5l-hgm9ZX9NrKYmlNmb3LYFrxd8zgWQ5fdulyColwquQHTGZ4kRwM49xCdTEsCDprDYqopmMoPXcwvjFz69f_a9S-qT8jNhFGN_WmdoUlq02J0BSLnD-IfKNAQe5w2D6SKCYH4LLZnSSKpQrJUFosqYmxOsk34Z80enm-f4kiJhzqYt9Qwa6n70J3bF-6vXCjZdu3l5iI5CLc7QcaicPmfiOQBGu_EyFploqsW7fmqfNhJ92_W0jL_yQvgZaBtMeZH23BmVpYQGdNeBhRauBj4qaHpTkPHSUPBw26atNtq5yWOs0j6gIL_hhHljLqqvMmI3vw0B3deVxqy2MvgU40wkYYfETwaBaIVTH5OtO60EBULwN7PBCYDVOT2ZskBZnQCp29IZqam1GwZHDBHEVbw9djRXl9ZwKOUHa9mv_dKHJAB_cMtpAzOF0H6fnyzLr1kZEd-hfHUY-ZF1mrEPgUBKC5swYTfgYpBpgWjnlxMKva9tC1iHAwKSyRcsRc7BnSzBJv9ZU8UhzrhIlJ1ILGjjJQcUwicfVIokN9X79fxllkkqQ&session_state=7515cdc9-6fd7-4e93-b526-f0a3c23401fd#") + driver.get( + config["url_interface_2"] + + "?code=0.AVsA57e_zrnCv0-sKQYbnoaPz5D--u_XfqNDqwMZoL4vF1hbAN8.AgABBAIAAAApTwJmzXqdR4BN2miheQMYAgDs_wUA9P91FtoYFvR20dqHG0dEKQhDe1BftWl3mQWsfh9s8UJttZdIVkp4cYIkmKKhTtdGruEzTbeZjcCH6mii-K8QH5QDmQ9chVv-P5lGlTOPig_9AZ_6ZZEty4XWPFPJzseHk4JsvCk36kgebUZvfT2X2MfMEsRi-TnOC2CWhhlurm0a8mCR7WzDWHLzPHpFAwcd-GS7V3N8NUri6LDskw8E1HiOtMSs0SwwZCGw2A7GZV59MrYvNxi8YyousqNRLhdOg56PrzlhLk5l-hgm9ZX9NrKYmlNmb3LYFrxd8zgWQ5fdulyColwquQHTGZ4kRwM49xCdTEsCDprDYqopmMoPXcwvjFz69f_a9S-qT8jNhFGN_WmdoUlq02J0BSLnD-IfKNAQe5w2D6SKCYH4LLZnSSKpQrJUFosqYmxOsk34Z80enm-f4kiJhzqYt9Qwa6n70J3bF-6vXCjZdu3l5iI5CLc7QcaicPmfiOQBGu_EyFploqsW7fmqfNhJ92_W0jL_yQvgZaBtMeZH23BmVpYQGdNeBhRauBj4qaHpTkPHSUPBw26atNtq5yWOs0j6gIL_hhHljLqqvMmI3vw0B3deVxqy2MvgU40wkYYfETwaBaIVTH5OtO60EBULwN7PBCYDVOT2ZskBZnQCp29IZqam1GwZHDBHEVbw9djRXl9ZwKOUHa9mv_dKHJAB_cMtpAzOF0H6fnyzLr1kZEd-hfHUY-ZF1mrEPgUBKC5swYTfgYpBpgWjnlxMKva9tC1iHAwKSyRcsRc7BnSzBJv9ZU8UhzrhIlJ1ILGjjJQcUwicfVIokN9X79fxllkkqQ&session_state=7515cdc9-6fd7-4e93-b526-f0a3c23401fd#" + ) with allure.step("Selecting Client"): # sign_in(driver, config) self.logger.info("Signed in successful") diff --git a/Doczy.AI_Automation/utils/CustomException.py b/Doczy.AI_Automation/utils/CustomException.py index 04d55ae..80e690b 100644 --- a/Doczy.AI_Automation/utils/CustomException.py +++ b/Doczy.AI_Automation/utils/CustomException.py @@ -3,6 +3,7 @@ from selenium.common import TimeoutException class CustomError(Exception): """Base class for custom exceptions""" + def __init__(self, message): super().__init__(message) self.message = message @@ -10,6 +11,7 @@ class CustomError(Exception): class InvalidInputError(CustomError): """Exception raised for invalid inputs""" + def __init__(self, value): super().__init__(f"Invalid input: {value}") self.value = value @@ -17,6 +19,7 @@ class InvalidInputError(CustomError): class DatabaseConnectionError(CustomError): """Exception raised for database connection errors""" + def __init__(self, db_url): super().__init__(f"Database connection failed: {db_url}") self.db_url = db_url @@ -24,19 +27,26 @@ class DatabaseConnectionError(CustomError): class NoYamlFile(FileNotFoundError): def __init__(self, yaml_file_path, yaml_file_name): - super().__init__(f"Please create YAML file at location={yaml_file_path} with filename={yaml_file_name}") + super().__init__( + f"Please create YAML file at location={yaml_file_path} with filename={yaml_file_name}" + ) self.yaml_file_path = yaml_file_path self.yaml_file_name = yaml_file_name class NoElementFound(Exception): def __init__(self, xpath, timeout): - super().__init__(f"Element with \"{xpath}\" XPATH is not located in {timeout} second. Try re-run the testcase") + super().__init__( + f'Element with "{xpath}" XPATH is not located in {timeout} second. Try re-run the testcase' + ) self.xpath = xpath self.timeout = timeout + class TitleNotFound(Exception): def __init__(self, title, timeout): - super().__init__(f"Web Page with \"{title}\" title is not loaded in {timeout} second.") + super().__init__( + f'Web Page with "{title}" title is not loaded in {timeout} second.' + ) self.title = title - self.timeout = timeout \ No newline at end of file + self.timeout = timeout diff --git a/Doczy.AI_Automation/utils/element_related_methods.py b/Doczy.AI_Automation/utils/element_related_methods.py index e1bce3e..d34c4b4 100644 --- a/Doczy.AI_Automation/utils/element_related_methods.py +++ b/Doczy.AI_Automation/utils/element_related_methods.py @@ -9,43 +9,40 @@ from utils.CustomException import NoElementFound, TitleNotFound def get_clickable_element(driver, config, By, element_id): try: - element = WebDriverWait(driver, timeout=config['timeout']).until( + element = WebDriverWait(driver, timeout=config["timeout"]).until( ec.element_to_be_clickable((By, element_id)) ) except TimeoutException: # print("NoElement") - raise NoElementFound(element_id, config['timeout']) + raise NoElementFound(element_id, config["timeout"]) return element def get_element(driver, config, By, element_id): try: - return WebDriverWait(driver, timeout=config['timeout']).until( + return WebDriverWait(driver, timeout=config["timeout"]).until( ec.visibility_of_element_located((By, element_id)) ) except TimeoutException: # print("NoElement") - raise NoElementFound(element_id, config['timeout']) + raise NoElementFound(element_id, config["timeout"]) def get_invisible_element(driver, config, By, element_id): try: - return WebDriverWait(driver, timeout=config['timeout']).until( + return WebDriverWait(driver, timeout=config["timeout"]).until( ec.presence_of_element_located((By, element_id)) ) except TimeoutException: # print("NoElement") - raise NoElementFound(element_id, config['timeout']) + raise NoElementFound(element_id, config["timeout"]) def get_title(driver, config, title): try: - return WebDriverWait(driver, timeout=config['timeout']).until( + return WebDriverWait(driver, timeout=config["timeout"]).until( ec.title_contains(title) ) except TimeoutException: # print("NoElement") - raise TitleNotFound(title, config['timeout']) - - - + raise TitleNotFound(title, config["timeout"]) diff --git a/Doczy.AI_Automation/utils/reader.py b/Doczy.AI_Automation/utils/reader.py index 31be075..ff0eb27 100644 --- a/Doczy.AI_Automation/utils/reader.py +++ b/Doczy.AI_Automation/utils/reader.py @@ -7,7 +7,7 @@ from utils.CustomException import NoYamlFile def read_yaml(file_path, file_name): try: - with open(file_path+file_name, 'r') as file: + with open(file_path + file_name, "r") as file: return yaml.safe_load(file) except Exception as e: raise NoYamlFile(file_path, file_name) diff --git a/airflow/dags/Qa_Dag.py b/airflow/dags/Qa_Dag.py index 2896f1b..848534a 100644 --- a/airflow/dags/Qa_Dag.py +++ b/airflow/dags/Qa_Dag.py @@ -5,47 +5,44 @@ from airflow.operators.python import PythonOperator from datetime import datetime from airflow import DAG from airflow.providers.snowflake.hooks.snowflake import SnowflakeHook + # from openpyxl.workbook import Workbook import pandas as pd import boto3 import io -SNOWFLAKE_CONN_ID="doczy_dev_snowflake" -TAGS=["QC","dev","etl","QC-Summery"] -DAG_ID="qc_report_dag" +SNOWFLAKE_CONN_ID = "doczy_dev_snowflake" +TAGS = ["QC", "dev", "etl", "QC-Summery"] +DAG_ID = "qc_report_dag" bucket = "doczy-dev-infra-mwaa-resources" object_key = "outputs/" -DATABASE="DOCZY_DEV" +DATABASE = "DOCZY_DEV" SCHEMA = "STG" TABLE_NAME = "TRAINING_DATA_RAW" # Trigger rules -ALL_SUCCESS = 'all_success' -ALL_FAILED = 'all_failed' -ALL_DONE = 'all_done' -ONE_SUCCESS = 'one_success' -ONE_FAILED = 'one_failed' +ALL_SUCCESS = "all_success" +ALL_FAILED = "all_failed" +ALL_DONE = "all_done" +ONE_SUCCESS = "one_success" +ONE_FAILED = "one_failed" +args = {"owner": "Airflow", "start_date": datetime(2022, 1, 1), "retries": 0} +dag = DAG(dag_id=DAG_ID, default_args=args, schedule=None, tags=TAGS) -args = {"owner": "Airflow", "start_date": datetime(2022, 1, 1), "retries":0 } -dag = DAG( - dag_id=DAG_ID, default_args=args, schedule=None, - tags=TAGS -) - def getData(): # Setup connection to Snowflake dwh_hook = SnowflakeHook(snowflake_conn_id=SNOWFLAKE_CONN_ID) conn = dwh_hook.get_conn() # Get the raw connection - + # Your query and the database details qc_query = f"SELECT * FROM {DATABASE}.{SCHEMA}.{TABLE_NAME}" - + # Fetch data into a Pandas DataFrame df = pd.read_sql(qc_query, conn) @@ -55,10 +52,12 @@ def getData(): # Create a buffer to hold the data with io.StringIO() as csv_buffer: df.to_csv(csv_buffer, index=False) - + # Save the data to S3 response = s3.put_object( - Bucket=bucket, Key=object_key+'snowflake_table_results.csv' , Body=csv_buffer.getvalue() + Bucket=bucket, + Key=object_key + "snowflake_table_results.csv", + Body=csv_buffer.getvalue(), ) status = response.get("ResponseMetadata", {}).get("HTTPStatusCode") @@ -66,29 +65,34 @@ def getData(): if status == 200: print(f"Successful S3 put_object response. Status - {status}") else: - raise AirflowFailException(f"Unsuccessful S3 put_object response. Status - {status}") - + raise AirflowFailException( + f"Unsuccessful S3 put_object response. Status - {status}" + ) + oldData = df newData = df with io.BytesIO() as output: - with pd.ExcelWriter(output, engine='xlsxwriter') as writer: + with pd.ExcelWriter(output, engine="xlsxwriter") as writer: oldData.to_excel(writer, sheet_name="Old") newData.to_excel(writer, sheet_name="new") dat = oldData.compare(newData, align_axis=0, keep_shape=True) dat.to_excel(writer, sheet_name="qc") response = s3.put_object( - Bucket=bucket, Key=object_key+'QC_Report.xlsx' , Body=output.getvalue() + Bucket=bucket, Key=object_key + "QC_Report.xlsx", Body=output.getvalue() ) print("Dataframe is written to S3 successfully.") + with dag: - begin_job = EmptyOperator(task_id='Begin') + begin_job = EmptyOperator(task_id="Begin") - get_data_from_snowflake = PythonOperator(task_id="get_data_from_snowflake", python_callable=getData) + get_data_from_snowflake = PythonOperator( + task_id="get_data_from_snowflake", python_callable=getData + ) - end_job = EmptyOperator(task_id='End') + end_job = EmptyOperator(task_id="End") -begin_job >> get_data_from_snowflake >> end_job \ No newline at end of file +begin_job >> get_data_from_snowflake >> end_job diff --git a/airflow/dags/cicd_test_dag.py b/airflow/dags/cicd_test_dag.py index 20e7351..990f8b6 100644 --- a/airflow/dags/cicd_test_dag.py +++ b/airflow/dags/cicd_test_dag.py @@ -42,13 +42,13 @@ logger = logging.getLogger(__name__) # Recommended to use this in the beginning of the script -SNOWFLAKE_CONN_ID = "doczy_dev_snowflake" # Specific for every database +SNOWFLAKE_CONN_ID = "doczy_dev_snowflake" # Specific for every database DAG_ID = "cicd_testing_dag" # Must be unique for every dag -DATABASE="XXX" +DATABASE = "XXX" -default_params = {"Database": "", "Schema":""} +default_params = {"Database": "", "Schema": ""} -''' +""" Tags will be used to filter DAGs on the Airflow UI. Guidelines for TAGs: @@ -58,12 +58,11 @@ Guidelines for TAGs: "medical" / "pharmacy" - based on the scenario "etl" / "adhoc" - based on the nature of the dataload -''' -TAGS=["adhoc","dev"] # MANDATORY +""" +TAGS = ["adhoc", "dev"] # MANDATORY - -''' +""" TRIGGER RULES for a task: all_success: (default) all parents have succeeded @@ -74,19 +73,17 @@ one_success: fires as soon as at least one parent succeeds, it does not wait for none_failed: all parents have not failed (failed or upstream_failed) i.e. all parents have succeeded or been skipped none_skipped: no parent is in a skipped state, i.e. all parents are in a success, failed, or upstream_failed state dummy: dependencies are just for show, trigger at will -''' -ALL_SUCCESS = 'all_success' -ALL_FAILED = 'all_failed' -ALL_DONE = 'all_done' -ONE_SUCCESS = 'one_success' -ONE_FAILED = 'one_failed' - +""" +ALL_SUCCESS = "all_success" +ALL_FAILED = "all_failed" +ALL_DONE = "all_done" +ONE_SUCCESS = "one_success" +ONE_FAILED = "one_failed" dag = DAG( # These args will get passed on to each operator # You can override them on a per-task basis during operator initialization - # 'queue': 'bash_queue', # 'pool': 'backfill', # 'priority_weight': 10, @@ -98,30 +95,28 @@ dag = DAG( # 'on_success_callback': some_other_function, # 'on_retry_callback': another_function, # 'sla_miss_callback': yet_another_function, - - DAG_ID, # Mandatory for every dag - start_date=datetime(2022, 1, 1), # Must be in the past + DAG_ID, # Mandatory for every dag + start_date=datetime(2022, 1, 1), # Must be in the past # Can pass snowflake conn id here instead of passing it to every task # It is needed when using Snowflake operator - default_args={"snowflake_conn_id": SNOWFLAKE_CONN_ID, 'retries': 0}, + default_args={"snowflake_conn_id": SNOWFLAKE_CONN_ID, "retries": 0}, tags=TAGS, catchup=False, # True will run the dag on the specified frequency for backdated DAGs. Not applicate in out workflow - schedule=None, # If there is no fixed schedule, then always pass None explicitly - params = default_params - + schedule=None, # If there is no fixed schedule, then always pass None explicitly + params=default_params, ) - -def python_op_eg(table_name,params): + +def python_op_eg(table_name, params): # Snowflake hook is used to fetch connection and cursor dwh_hook = SnowflakeHook(snowflake_conn_id=SNOWFLAKE_CONN_ID) conn = dwh_hook.get_conn() curr = conn.cursor() - DATABASE = params['Database'] - schema_name = params['Schema'] + DATABASE = params["Database"] + schema_name = params["Schema"] query_output = curr.execute(f"SELECT * from {DATABASE}.{schema_name}.{table_name}") - + # Alternative way to fetch query result # result = dwh_hook.get_first(f"select max(audit_sid) from {DATABASE}.stg.CLAIM_MED_STAGING") @@ -131,26 +126,28 @@ def python_op_eg(table_name,params): conn.close -# Best practice to have an empty start at the beginning and end -begin_job = EmptyOperator(task_id='Begin') +# Best practice to have an empty start at the beginning and end +begin_job = EmptyOperator(task_id="Begin") -python_task = PythonOperator(task_id='get_info_using_python_op', # Task ID has to be unique only inside a dags +python_task = PythonOperator( + task_id="get_info_using_python_op", # Task ID has to be unique only inside a dags python_callable=python_op_eg, - dag=dag, - op_kwargs={'table_name':'DIM_AUDIT',} # Variables can be passed to the python function using op_kwargs - ) + dag=dag, + op_kwargs={ + "table_name": "DIM_AUDIT", + }, # Variables can be passed to the python function using op_kwargs +) # Snowflake operator does not offer the option to display query results -# We can use this for executing stored procedures. +# We can use this for executing stored procedures. sf_task = SnowflakeOperator( - task_id="get_info_using_sf_op", - sql=f"select * from {DATABASE}.STG.DIM_AUDIT", - dag=dag + task_id="get_info_using_sf_op", + sql=f"select * from {DATABASE}.STG.DIM_AUDIT", + dag=dag, ) - -end_job = EmptyOperator(task_id='End') +end_job = EmptyOperator(task_id="End") # EXECUTING TASKS -begin_job >> python_task >> sf_task >> end_job \ No newline at end of file +begin_job >> python_task >> sf_task >> end_job diff --git a/airflow/dags/client_name_dag.py b/airflow/dags/client_name_dag.py index 138d64a..597002c 100644 --- a/airflow/dags/client_name_dag.py +++ b/airflow/dags/client_name_dag.py @@ -15,34 +15,36 @@ logger = logging.getLogger(__name__) SNOWFLAKE_CONN_ID = "doczy_dev_snowflake" DAG_ID = "load_client_config" -DATABASE="DOCZY_DEV" +DATABASE = "DOCZY_DEV" # bucket = "airflow-data-ingestion" -TAGS=["dev","config_interface","dataload"] - +TAGS = ["dev", "config_interface", "dataload"] def process_csv_in_s3(**kwargs): s3_hook = S3Hook(aws_conn_id="aws_default") bucket_name = "doczy-dev-infra-raw-data-ingestion" - key_prefix = 'client_names_openair/' + key_prefix = "client_names_openair/" # List objects within the specified bucket and key prefix objects = s3_hook.list_keys(bucket_name=bucket_name, prefix=key_prefix) - + # Sort the objects by their last modified date and select the latest if objects: - latest_file_key = sorted(objects)[-1] # Assuming file names include a timestamp or incrementing number + latest_file_key = sorted(objects)[ + -1 + ] # Assuming file names include a timestamp or incrementing number # Get the latest file from S3 file_content = s3_hook.read_key(latest_file_key, bucket_name) - + # Convert string to DataFrame from io import StringIO + df = pd.read_csv(StringIO(file_content)) # Apply the generate_s3_path function - df['s3_path'] = df['customer_name'].apply(generate_s3_path) + df["s3_path"] = df["customer_name"].apply(generate_s3_path) # Convert DataFrame to CSV string csv_buffer = StringIO() @@ -54,49 +56,55 @@ def process_csv_in_s3(**kwargs): string_data=csv_content, key=latest_file_key, bucket_name=bucket_name, - replace=True + replace=True, ) # Push the latest file name to XCom - if 'latest_file_key' in locals(): - kwargs['ti'].xcom_push(key='file_name', value=latest_file_key.split('/')[-1]) + if "latest_file_key" in locals(): + kwargs["ti"].xcom_push(key="file_name", value=latest_file_key.split("/")[-1]) else: print("No files found in the specified path.") + def generate_s3_path(client_name): # Your function as provided - client_name = client_name.rstrip('.') - pattern = r'[^0-9a-zA-Z!_.()*\'-]' - multi_underscore_pattern = r'_{2,}' - intermediate_name = re.sub(pattern, '_', client_name) - final_name = re.sub(multi_underscore_pattern, '_', intermediate_name) + client_name = client_name.rstrip(".") + pattern = r"[^0-9a-zA-Z!_.()*\'-]" + multi_underscore_pattern = r"_{2,}" + intermediate_name = re.sub(pattern, "_", client_name) + final_name = re.sub(multi_underscore_pattern, "_", intermediate_name) final_name = final_name.lower() - base_s3_path = 's3://' + base_s3_path = "s3://" return base_s3_path + final_name + "/" + def call_stored_proc(proc_name, **kwargs): - # Pull the file name from XCom - file_name = kwargs['ti'].xcom_pull(task_ids='process_csv', key='file_name') + # Pull the file name from XCom + file_name = kwargs["ti"].xcom_pull(task_ids="process_csv", key="file_name") logger.info(f"File name: {file_name}") - dwh_hook = SnowflakeHook(snowflake_conn_id=SNOWFLAKE_CONN_ID) + dwh_hook = SnowflakeHook(snowflake_conn_id=SNOWFLAKE_CONN_ID) with dwh_hook.get_conn() as conn: # dwh_hook.set_autocommit(conn,autocommit=False) cur = conn.cursor() cur.execute(f"CALL {DATABASE}.STG.{proc_name}('{file_name}');") result = cur.fetchone() - if result[0] == 'Setup, Load, and Audit Complete': - logger.info('PROCEDURE EXECUTED SUCCESSFULLY') + if result[0] == "Setup, Load, and Audit Complete": + logger.info("PROCEDURE EXECUTED SUCCESSFULLY") else: - raise AirflowFailException("Check the DAG logs for more information. ERROR FROM SNOWFLAKE: ", result) + raise AirflowFailException( + "Check the DAG logs for more information. ERROR FROM SNOWFLAKE: ", + result, + ) logger.info(f"QUERY EXECUTION RESULT: {str(result)}") + default_args = { - 'owner': 'airflow', - 'depends_on_past': False, - 'start_date': datetime(2024, 3, 28), - 'retries': 1, - 'retry_delay': timedelta(minutes=5) + "owner": "airflow", + "depends_on_past": False, + "start_date": datetime(2024, 3, 28), + "retries": 1, + "retry_delay": timedelta(minutes=5), } dag = DAG( @@ -104,17 +112,17 @@ dag = DAG( default_args=default_args, start_date=datetime(2024, 3, 28), catchup=False, - description='Process a CSV in S3 and replace it', - schedule_interval='@daily' + description="Process a CSV in S3 and replace it", + schedule_interval="@daily", ) -begin_job = EmptyOperator(task_id='Begin') +begin_job = EmptyOperator(task_id="Begin") process_csv_task = PythonOperator( - task_id='process_csv', + task_id="process_csv", python_callable=process_csv_in_s3, provide_context=True, - dag=dag + dag=dag, ) load_client_config = PythonOperator( @@ -122,12 +130,10 @@ load_client_config = PythonOperator( python_callable=call_stored_proc, dag=dag, provide_context=True, - op_kwargs={'proc_name':'LOAD_CLIENT_CONFIG'} + op_kwargs={"proc_name": "LOAD_CLIENT_CONFIG"}, ) +end_job = EmptyOperator(task_id="End") - -end_job = EmptyOperator(task_id='End') - -begin_job >> process_csv_task >> load_client_config >> end_job +begin_job >> process_csv_task >> load_client_config >> end_job diff --git a/airflow/dags/config_interface_dag.py b/airflow/dags/config_interface_dag.py index 9f83064..3bc011f 100644 --- a/airflow/dags/config_interface_dag.py +++ b/airflow/dags/config_interface_dag.py @@ -1,5 +1,3 @@ - - from __future__ import annotations import os @@ -20,81 +18,85 @@ logging.basicConfig(level=logging.INFO) logger = logging.getLogger(__name__) - SNOWFLAKE_CONN_ID = "doczy_dev_snowflake" DAG_ID = "load_request_and_contract_submissions" -DATABASE="DOCZY_DEV" +DATABASE = "DOCZY_DEV" # bucket = "airflow-data-ingestion" -TAGS=["dev","contract_config_interface","dataload"] +TAGS = ["dev", "contract_config_interface", "dataload"] # Trigger rules -ALL_SUCCESS = 'all_success' -ALL_FAILED = 'all_failed' -ALL_DONE = 'all_done' -ONE_SUCCESS = 'one_success' -ONE_FAILED = 'one_failed' +ALL_SUCCESS = "all_success" +ALL_FAILED = "all_failed" +ALL_DONE = "all_done" +ONE_SUCCESS = "one_success" +ONE_FAILED = "one_failed" # Passing empty params for now, this will be overridden by the payload from the trigger # These params can also be set from the Airflow UI while manually triggering the DAG -default_params = {"request_submission_file_name": "", "contract_config_file_name":""} +default_params = {"request_submission_file_name": "", "contract_config_file_name": ""} # This will be replaced with the payload from the event after API connection is setup # training_results_file_name = "training_results_sample.csv" # attempt_logs_file_name = "attempt_logs_sample.csv" -def call_stored_proc(proc_name,file_type, params): - dwh_hook = SnowflakeHook(snowflake_conn_id=SNOWFLAKE_CONN_ID) +def call_stored_proc(proc_name, file_type, params): + dwh_hook = SnowflakeHook(snowflake_conn_id=SNOWFLAKE_CONN_ID) with dwh_hook.get_conn() as conn: # dwh_hook.set_autocommit(conn,autocommit=False) cur = conn.cursor() # Added new parameter file_type to determine the file name to be passed to the stored procedure # The bucket name is set by default to "doczy-dev-infra-raw-data-ingestion" and the files should be ALWAYS save under training_interface/ path for now - if file_type == 'request_submission': - file_name = params['request_submission_file_name'] - elif file_type == 'contract_config': - file_name = params['contract_config_file_name'] + if file_type == "request_submission": + file_name = params["request_submission_file_name"] + elif file_type == "contract_config": + file_name = params["contract_config_file_name"] cur.execute(f"CALL {DATABASE}.STG.{proc_name}('{file_name}');") result = cur.fetchone() - if result[0] == 'Setup, Load, and Audit Complete': - logger.info('PROCEDURE EXECUTED SUCCESSFULLY') + if result[0] == "Setup, Load, and Audit Complete": + logger.info("PROCEDURE EXECUTED SUCCESSFULLY") else: - raise AirflowFailException("Check the DAG logs for more information. ERROR FROM SNOWFLAKE: ", result) + raise AirflowFailException( + "Check the DAG logs for more information. ERROR FROM SNOWFLAKE: ", + result, + ) logger.info(f"QUERY EXECUTION RESULT: {str(result)}") + dag = DAG( DAG_ID, start_date=datetime(2024, 1, 1), - default_args={"snowflake_conn_id": SNOWFLAKE_CONN_ID, "retries":0}, + default_args={"snowflake_conn_id": SNOWFLAKE_CONN_ID, "retries": 0}, tags=TAGS, catchup=False, schedule=None, - params = default_params + params=default_params, ) - -begin_job = EmptyOperator(task_id='Begin') + +begin_job = EmptyOperator(task_id="Begin") load_request_submission = PythonOperator( task_id="load_request_submissions", python_callable=call_stored_proc, dag=dag, - op_kwargs={'proc_name':'LOAD_REQUEST_SUBMISSION', 'file_type':'request_submission'} + op_kwargs={ + "proc_name": "LOAD_REQUEST_SUBMISSION", + "file_type": "request_submission", + }, ) load_contract_config = PythonOperator( task_id="load_contract_config", python_callable=call_stored_proc, dag=dag, - op_kwargs={'proc_name':'LOAD_CONTRACT_CONFIG', 'file_type':'contract_config'} + op_kwargs={"proc_name": "LOAD_CONTRACT_CONFIG", "file_type": "contract_config"}, ) +end_job = EmptyOperator(task_id="End") -end_job = EmptyOperator(task_id='End') - - -begin_job >> load_request_submission >> load_contract_config >> end_job \ No newline at end of file +begin_job >> load_request_submission >> load_contract_config >> end_job diff --git a/airflow/dags/pipeline_processed_outputs_dag.py b/airflow/dags/pipeline_processed_outputs_dag.py index cc41556..1c687f7 100644 --- a/airflow/dags/pipeline_processed_outputs_dag.py +++ b/airflow/dags/pipeline_processed_outputs_dag.py @@ -1,5 +1,3 @@ - - from __future__ import annotations import os @@ -20,64 +18,71 @@ logging.basicConfig(level=logging.INFO) logger = logging.getLogger(__name__) - SNOWFLAKE_CONN_ID = "doczy_dev_snowflake" DAG_ID = "load_doczy_processed_outputs" -DATABASE="DOCZY_DEV" +DATABASE = "DOCZY_DEV" -TAGS=["dev","config_interface","dataload"] +TAGS = ["dev", "config_interface", "dataload"] # Trigger rules -ALL_SUCCESS = 'all_success' -ALL_FAILED = 'all_failed' -ALL_DONE = 'all_done' -ONE_SUCCESS = 'one_success' -ONE_FAILED = 'one_failed' +ALL_SUCCESS = "all_success" +ALL_FAILED = "all_failed" +ALL_DONE = "all_done" +ONE_SUCCESS = "one_success" +ONE_FAILED = "one_failed" # Passing empty params for now, this will be overridden by the payload from the trigger # These params can also be set from the Airflow UI while manually triggering the DAG # This will be replaced with the payload from the event after API connection is setup default_params = {"pipeline_processed_output": ""} -def call_stored_proc(proc_name,file_type, params): - dwh_hook = SnowflakeHook(snowflake_conn_id=SNOWFLAKE_CONN_ID) + +def call_stored_proc(proc_name, file_type, params): + dwh_hook = SnowflakeHook(snowflake_conn_id=SNOWFLAKE_CONN_ID) with dwh_hook.get_conn() as conn: # dwh_hook.set_autocommit(conn,autocommit=False) cur = conn.cursor() # Added new parameter file_type to determine the file name to be passed to the stored procedure - if file_type == 'pipeline_processed_output': - file_name = params['pipeline_processed_output'] + if file_type == "pipeline_processed_output": + file_name = params["pipeline_processed_output"] cur.execute(f"CALL {DATABASE}.STG.{proc_name}('{file_name}');") result = cur.fetchone() - if result[0] == 'Setup, Load, and Audit Complete': - logger.info('PROCEDURE EXECUTED SUCCESSFULLY') + if result[0] == "Setup, Load, and Audit Complete": + logger.info("PROCEDURE EXECUTED SUCCESSFULLY") else: - raise AirflowFailException("Check the DAG logs for more information. ERROR FROM SNOWFLAKE: ", result) + raise AirflowFailException( + "Check the DAG logs for more information. ERROR FROM SNOWFLAKE: ", + result, + ) logger.info(f"QUERY EXECUTION RESULT: {str(result)}") + dag = DAG( DAG_ID, start_date=datetime(2024, 1, 1), - default_args={"snowflake_conn_id": SNOWFLAKE_CONN_ID, "retries":0}, + default_args={"snowflake_conn_id": SNOWFLAKE_CONN_ID, "retries": 0}, tags=TAGS, catchup=False, schedule=None, - params = default_params + params=default_params, ) - -begin_job = EmptyOperator(task_id='Begin') + +begin_job = EmptyOperator(task_id="Begin") load_pipeline_processed_output = PythonOperator( task_id="load_pipeline_processed_output", python_callable=call_stored_proc, dag=dag, - op_kwargs={'proc_name':'LOAD_DOCZY_PIPELINE_PROCESSED_OUTPUT', 'file_type':'pipeline_processed_output'} + op_kwargs={ + "proc_name": "LOAD_DOCZY_PIPELINE_PROCESSED_OUTPUT", + "file_type": "pipeline_processed_output", + }, ) -end_job = EmptyOperator(task_id='End') +end_job = EmptyOperator(task_id="End") begin_job >> load_pipeline_processed_output >> end_job diff --git a/airflow/dags/raw_training_data_dag.py b/airflow/dags/raw_training_data_dag.py index f8156d4..6bbc79b 100644 --- a/airflow/dags/raw_training_data_dag.py +++ b/airflow/dags/raw_training_data_dag.py @@ -1,5 +1,3 @@ - - from __future__ import annotations import os @@ -20,89 +18,101 @@ logging.basicConfig(level=logging.INFO) logger = logging.getLogger(__name__) - SNOWFLAKE_CONN_ID = "doczy_dev_snowflake" DAG_ID = "load_raw_training_data" -DATABASE="DOCZY_DEV" +DATABASE = "DOCZY_DEV" # bucket = "airflow-data-ingestion" -TAGS=["dev","training_data","dataload"] +TAGS = ["dev", "training_data", "dataload"] # Trigger rules -ALL_SUCCESS = 'all_success' -ALL_FAILED = 'all_failed' -ALL_DONE = 'all_done' -ONE_SUCCESS = 'one_success' -ONE_FAILED = 'one_failed' +ALL_SUCCESS = "all_success" +ALL_FAILED = "all_failed" +ALL_DONE = "all_done" +ONE_SUCCESS = "one_success" +ONE_FAILED = "one_failed" # Passing empty params for now, this will be overridden by the payload from the trigger # These params can also be set from the Airflow UI while manually triggering the DAG -default_params = {"column_config_file_name": "", "raw_training_data_file_name":"", "business_config_file_name":""} +default_params = { + "column_config_file_name": "", + "raw_training_data_file_name": "", + "business_config_file_name": "", +} # This will be replaced with the payload from the event after API connection is setup # training_results_file_name = "training_results_sample.csv" # attempt_logs_file_name = "attempt_logs_sample.csv" -def call_stored_proc(proc_name,file_type, params): - dwh_hook = SnowflakeHook(snowflake_conn_id=SNOWFLAKE_CONN_ID) +def call_stored_proc(proc_name, file_type, params): + dwh_hook = SnowflakeHook(snowflake_conn_id=SNOWFLAKE_CONN_ID) with dwh_hook.get_conn() as conn: # dwh_hook.set_autocommit(conn,autocommit=False) cur = conn.cursor() # Added new parameter file_type to determine the file name to be passed to the stored procedure # The bucket name is set by default to "doczy-dev-infra-raw-data-ingestion" and the files should be ALWAYS save under training_interface/ path for now - if file_type == 'column_config': - file_name = params['column_config_file_name'] - elif file_type == 'raw_training_data': - file_name = params['raw_training_data_file_name'] - elif file_type == 'business_config': - file_name = params['business_config_file_name'] + if file_type == "column_config": + file_name = params["column_config_file_name"] + elif file_type == "raw_training_data": + file_name = params["raw_training_data_file_name"] + elif file_type == "business_config": + file_name = params["business_config_file_name"] cur.execute(f"CALL {DATABASE}.STG.{proc_name}('{file_name}');") result = cur.fetchone() - if result[0] == 'Setup, Load, and Audit Complete': - logger.info('PROCEDURE EXECUTED SUCCESSFULLY') + if result[0] == "Setup, Load, and Audit Complete": + logger.info("PROCEDURE EXECUTED SUCCESSFULLY") else: - raise AirflowFailException("Check the DAG logs for more information. ERROR FROM SNOWFLAKE: ", result) + raise AirflowFailException( + "Check the DAG logs for more information. ERROR FROM SNOWFLAKE: ", + result, + ) logger.info(f"QUERY EXECUTION RESULT: {str(result)}") + dag = DAG( DAG_ID, start_date=datetime(2024, 1, 1), - default_args={"snowflake_conn_id": SNOWFLAKE_CONN_ID, "retries":0}, + default_args={"snowflake_conn_id": SNOWFLAKE_CONN_ID, "retries": 0}, tags=TAGS, catchup=False, schedule=None, - params = default_params + params=default_params, ) - -begin_job = EmptyOperator(task_id='Begin') + +begin_job = EmptyOperator(task_id="Begin") load_training_results = PythonOperator( task_id="load_column_config", python_callable=call_stored_proc, dag=dag, - op_kwargs={'proc_name':'LOAD_COLUMN_CONFIG', 'file_type':'column_config'} + op_kwargs={"proc_name": "LOAD_COLUMN_CONFIG", "file_type": "column_config"}, ) load_attempt_logs_sp = PythonOperator( task_id="load_raw_training_data", python_callable=call_stored_proc, dag=dag, - op_kwargs={'proc_name':'LOAD_TRAINING_DATA_RAW', 'file_type':'raw_training_data'} + op_kwargs={"proc_name": "LOAD_TRAINING_DATA_RAW", "file_type": "raw_training_data"}, ) load_business_config = PythonOperator( task_id="load_business_config", python_callable=call_stored_proc, dag=dag, - op_kwargs={'proc_name':'LOAD_BUSINESS_CONFIG', 'file_type':'business_config'} + op_kwargs={"proc_name": "LOAD_BUSINESS_CONFIG", "file_type": "business_config"}, ) - -end_job = EmptyOperator(task_id='End') +end_job = EmptyOperator(task_id="End") -begin_job >> load_training_results >> load_attempt_logs_sp >> load_business_config >> end_job \ No newline at end of file +( + begin_job + >> load_training_results + >> load_attempt_logs_sp + >> load_business_config + >> end_job +) diff --git a/airflow/dags/training_results_dag.py b/airflow/dags/training_results_dag.py index 7c01caa..36a4768 100644 --- a/airflow/dags/training_results_dag.py +++ b/airflow/dags/training_results_dag.py @@ -1,5 +1,3 @@ - - from __future__ import annotations import os @@ -20,81 +18,82 @@ logging.basicConfig(level=logging.INFO) logger = logging.getLogger(__name__) - SNOWFLAKE_CONN_ID = "doczy_dev_snowflake" DAG_ID = "load_training_results" -DATABASE="DOCZY_DEV" +DATABASE = "DOCZY_DEV" # bucket = "airflow-data-ingestion" -TAGS=["dev","training_interface","dataload"] +TAGS = ["dev", "training_interface", "dataload"] # Trigger rules -ALL_SUCCESS = 'all_success' -ALL_FAILED = 'all_failed' -ALL_DONE = 'all_done' -ONE_SUCCESS = 'one_success' -ONE_FAILED = 'one_failed' +ALL_SUCCESS = "all_success" +ALL_FAILED = "all_failed" +ALL_DONE = "all_done" +ONE_SUCCESS = "one_success" +ONE_FAILED = "one_failed" # Passing empty params for now, this will be overridden by the payload from the trigger # These params can also be set from the Airflow UI while manually triggering the DAG -default_params = {"training_results_file_name": "", "attempt_logs_file_name":""} +default_params = {"training_results_file_name": "", "attempt_logs_file_name": ""} # This will be replaced with the payload from the event after API connection is setup # training_results_file_name = "training_results_sample.csv" # attempt_logs_file_name = "attempt_logs_sample.csv" -def call_stored_proc(proc_name,file_type, params): - dwh_hook = SnowflakeHook(snowflake_conn_id=SNOWFLAKE_CONN_ID) +def call_stored_proc(proc_name, file_type, params): + dwh_hook = SnowflakeHook(snowflake_conn_id=SNOWFLAKE_CONN_ID) with dwh_hook.get_conn() as conn: # dwh_hook.set_autocommit(conn,autocommit=False) cur = conn.cursor() # Added new parameter file_type to determine the file name to be passed to the stored procedure # The bucket name is set by default to "doczy-dev-infra-raw-data-ingestion" and the files should be ALWAYS save under training_interface/ path for now - if file_type == 'training_results': - file_name = params['training_results_file_name'] - elif file_type == 'attempt_logs': - file_name = params['attempt_logs_file_name'] + if file_type == "training_results": + file_name = params["training_results_file_name"] + elif file_type == "attempt_logs": + file_name = params["attempt_logs_file_name"] cur.execute(f"CALL {DATABASE}.STG.{proc_name}('{file_name}');") result = cur.fetchone() - if result[0] == 'Setup, Load, and Audit Complete': - logger.info('PROCEDURE EXECUTED SUCCESSFULLY') + if result[0] == "Setup, Load, and Audit Complete": + logger.info("PROCEDURE EXECUTED SUCCESSFULLY") else: - raise AirflowFailException("Check the DAG logs for more information. ERROR FROM SNOWFLAKE: ", result) + raise AirflowFailException( + "Check the DAG logs for more information. ERROR FROM SNOWFLAKE: ", + result, + ) logger.info(f"QUERY EXECUTION RESULT: {str(result)}") + dag = DAG( DAG_ID, start_date=datetime(2024, 1, 1), - default_args={"snowflake_conn_id": SNOWFLAKE_CONN_ID, "retries":0}, + default_args={"snowflake_conn_id": SNOWFLAKE_CONN_ID, "retries": 0}, tags=TAGS, catchup=False, schedule=None, - params = default_params + params=default_params, ) - -begin_job = EmptyOperator(task_id='Begin') + +begin_job = EmptyOperator(task_id="Begin") load_training_results = PythonOperator( task_id="load_training_results", python_callable=call_stored_proc, dag=dag, - op_kwargs={'proc_name':'LOAD_TRAINING_RESULTS', 'file_type':'training_results'} + op_kwargs={"proc_name": "LOAD_TRAINING_RESULTS", "file_type": "training_results"}, ) load_attempt_logs_sp = PythonOperator( task_id="load_attempt_logs", python_callable=call_stored_proc, dag=dag, - op_kwargs={'proc_name':'LOAD_ATTEMPT_LOGS', 'file_type':'attempt_logs'} + op_kwargs={"proc_name": "LOAD_ATTEMPT_LOGS", "file_type": "attempt_logs"}, ) +end_job = EmptyOperator(task_id="End") -end_job = EmptyOperator(task_id='End') - - -begin_job >> load_training_results >> load_attempt_logs_sp >> end_job \ No newline at end of file +begin_job >> load_training_results >> load_attempt_logs_sp >> end_job diff --git a/airflow/dags/uat_Qa_Dag.py b/airflow/dags/uat_Qa_Dag.py index bfda702..3f56c9f 100644 --- a/airflow/dags/uat_Qa_Dag.py +++ b/airflow/dags/uat_Qa_Dag.py @@ -5,47 +5,44 @@ from airflow.operators.python import PythonOperator from datetime import datetime from airflow import DAG from airflow.providers.snowflake.hooks.snowflake import SnowflakeHook + # from openpyxl.workbook import Workbook import pandas as pd import boto3 import io -SNOWFLAKE_CONN_ID="doczy_uat_snowflake" -TAGS=["QC","uat","etl","QC-Summery"] -DAG_ID="qc_report_dag" +SNOWFLAKE_CONN_ID = "doczy_uat_snowflake" +TAGS = ["QC", "uat", "etl", "QC-Summery"] +DAG_ID = "qc_report_dag" bucket = "doczy-dev-infra-mwaa-resources" object_key = "outputs/" -DATABASE="DOCZY_UAT" +DATABASE = "DOCZY_UAT" SCHEMA = "STG" TABLE_NAME = "TRAINING_DATA_RAW" # Trigger rules -ALL_SUCCESS = 'all_success' -ALL_FAILED = 'all_failed' -ALL_DONE = 'all_done' -ONE_SUCCESS = 'one_success' -ONE_FAILED = 'one_failed' +ALL_SUCCESS = "all_success" +ALL_FAILED = "all_failed" +ALL_DONE = "all_done" +ONE_SUCCESS = "one_success" +ONE_FAILED = "one_failed" +args = {"owner": "Airflow", "start_date": datetime(2022, 1, 1), "retries": 0} +dag = DAG(dag_id=DAG_ID, default_args=args, schedule=None, tags=TAGS) -args = {"owner": "Airflow", "start_date": datetime(2022, 1, 1), "retries":0 } -dag = DAG( - dag_id=DAG_ID, default_args=args, schedule=None, - tags=TAGS -) - def getData(): # Setup connection to Snowflake dwh_hook = SnowflakeHook(snowflake_conn_id=SNOWFLAKE_CONN_ID) conn = dwh_hook.get_conn() # Get the raw connection - + # Your query and the database details qc_query = f"SELECT * FROM {DATABASE}.{SCHEMA}.{TABLE_NAME}" - + # Fetch data into a Pandas DataFrame df = pd.read_sql(qc_query, conn) @@ -55,10 +52,12 @@ def getData(): # Create a buffer to hold the data with io.StringIO() as csv_buffer: df.to_csv(csv_buffer, index=False) - + # Save the data to S3 response = s3.put_object( - Bucket=bucket, Key=object_key+'snowflake_table_results.csv' , Body=csv_buffer.getvalue() + Bucket=bucket, + Key=object_key + "snowflake_table_results.csv", + Body=csv_buffer.getvalue(), ) status = response.get("ResponseMetadata", {}).get("HTTPStatusCode") @@ -66,29 +65,34 @@ def getData(): if status == 200: print(f"Successful S3 put_object response. Status - {status}") else: - raise AirflowFailException(f"Unsuccessful S3 put_object response. Status - {status}") - + raise AirflowFailException( + f"Unsuccessful S3 put_object response. Status - {status}" + ) + oldData = df newData = df with io.BytesIO() as output: - with pd.ExcelWriter(output, engine='xlsxwriter') as writer: + with pd.ExcelWriter(output, engine="xlsxwriter") as writer: oldData.to_excel(writer, sheet_name="Old") newData.to_excel(writer, sheet_name="new") dat = oldData.compare(newData, align_axis=0, keep_shape=True) dat.to_excel(writer, sheet_name="qc") response = s3.put_object( - Bucket=bucket, Key=object_key+'QC_Report.xlsx' , Body=output.getvalue() + Bucket=bucket, Key=object_key + "QC_Report.xlsx", Body=output.getvalue() ) print("Dataframe is written to S3 successfully.") + with dag: - begin_job = EmptyOperator(task_id='Begin') + begin_job = EmptyOperator(task_id="Begin") - get_data_from_snowflake = PythonOperator(task_id="get_data_from_snowflake", python_callable=getData) + get_data_from_snowflake = PythonOperator( + task_id="get_data_from_snowflake", python_callable=getData + ) - end_job = EmptyOperator(task_id='End') + end_job = EmptyOperator(task_id="End") -begin_job >> get_data_from_snowflake >> end_job \ No newline at end of file +begin_job >> get_data_from_snowflake >> end_job diff --git a/airflow/dags/uat_config_interface_dag .py b/airflow/dags/uat_config_interface_dag .py index bcca34f..e36d851 100644 --- a/airflow/dags/uat_config_interface_dag .py +++ b/airflow/dags/uat_config_interface_dag .py @@ -1,5 +1,3 @@ - - from __future__ import annotations import os @@ -20,81 +18,85 @@ logging.basicConfig(level=logging.INFO) logger = logging.getLogger(__name__) - SNOWFLAKE_CONN_ID = "doczy_uat_snowflake" DAG_ID = "load_request_and_contract_submissions" -DATABASE="DOCZY_UAT" +DATABASE = "DOCZY_UAT" # bucket = "airflow-data-ingestion" -TAGS=["uat","contract_config_interface","dataload"] +TAGS = ["uat", "contract_config_interface", "dataload"] # Trigger rules -ALL_SUCCESS = 'all_success' -ALL_FAILED = 'all_failed' -ALL_DONE = 'all_done' -ONE_SUCCESS = 'one_success' -ONE_FAILED = 'one_failed' +ALL_SUCCESS = "all_success" +ALL_FAILED = "all_failed" +ALL_DONE = "all_done" +ONE_SUCCESS = "one_success" +ONE_FAILED = "one_failed" # Passing empty params for now, this will be overridden by the payload from the trigger # These params can also be set from the Airflow UI while manually triggering the DAG -default_params = {"request_submission_file_name": "", "contract_config_file_name":""} +default_params = {"request_submission_file_name": "", "contract_config_file_name": ""} # This will be replaced with the payload from the event after API connection is setup # training_results_file_name = "training_results_sample.csv" # attempt_logs_file_name = "attempt_logs_sample.csv" -def call_stored_proc(proc_name,file_type, params): - dwh_hook = SnowflakeHook(snowflake_conn_id=SNOWFLAKE_CONN_ID) +def call_stored_proc(proc_name, file_type, params): + dwh_hook = SnowflakeHook(snowflake_conn_id=SNOWFLAKE_CONN_ID) with dwh_hook.get_conn() as conn: # dwh_hook.set_autocommit(conn,autocommit=False) cur = conn.cursor() # Added new parameter file_type to determine the file name to be passed to the stored procedure # The bucket name is set by default to "doczy-dev-infra-raw-data-ingestion" and the files should be ALWAYS save under training_interface/ path for now - if file_type == 'request_submission': - file_name = params['request_submission_file_name'] - elif file_type == 'contract_config': - file_name = params['contract_config_file_name'] + if file_type == "request_submission": + file_name = params["request_submission_file_name"] + elif file_type == "contract_config": + file_name = params["contract_config_file_name"] cur.execute(f"CALL {DATABASE}.STG.{proc_name}('{file_name}');") result = cur.fetchone() - if result[0] == 'Setup, Load, and Audit Complete': - logger.info('PROCEDURE EXECUTED SUCCESSFULLY') + if result[0] == "Setup, Load, and Audit Complete": + logger.info("PROCEDURE EXECUTED SUCCESSFULLY") else: - raise AirflowFailException("Check the DAG logs for more information. ERROR FROM SNOWFLAKE: ", result) + raise AirflowFailException( + "Check the DAG logs for more information. ERROR FROM SNOWFLAKE: ", + result, + ) logger.info(f"QUERY EXECUTION RESULT: {str(result)}") + dag = DAG( DAG_ID, start_date=datetime(2024, 1, 1), - default_args={"snowflake_conn_id": SNOWFLAKE_CONN_ID, "retries":0}, + default_args={"snowflake_conn_id": SNOWFLAKE_CONN_ID, "retries": 0}, tags=TAGS, catchup=False, schedule=None, - params = default_params + params=default_params, ) - -begin_job = EmptyOperator(task_id='Begin') + +begin_job = EmptyOperator(task_id="Begin") load_request_submission = PythonOperator( task_id="load_request_submissions", python_callable=call_stored_proc, dag=dag, - op_kwargs={'proc_name':'LOAD_REQUEST_SUBMISSION', 'file_type':'request_submission'} + op_kwargs={ + "proc_name": "LOAD_REQUEST_SUBMISSION", + "file_type": "request_submission", + }, ) load_contract_config = PythonOperator( task_id="load_contract_config", python_callable=call_stored_proc, dag=dag, - op_kwargs={'proc_name':'LOAD_CONTRACT_CONFIG', 'file_type':'contract_config'} + op_kwargs={"proc_name": "LOAD_CONTRACT_CONFIG", "file_type": "contract_config"}, ) +end_job = EmptyOperator(task_id="End") -end_job = EmptyOperator(task_id='End') - - -begin_job >> load_request_submission >> load_contract_config >> end_job \ No newline at end of file +begin_job >> load_request_submission >> load_contract_config >> end_job diff --git a/airflow/dags/uat_raw_training_data_dag.py b/airflow/dags/uat_raw_training_data_dag.py index c78fda5..f22aede 100644 --- a/airflow/dags/uat_raw_training_data_dag.py +++ b/airflow/dags/uat_raw_training_data_dag.py @@ -1,5 +1,3 @@ - - from __future__ import annotations import os @@ -20,89 +18,101 @@ logging.basicConfig(level=logging.INFO) logger = logging.getLogger(__name__) - SNOWFLAKE_CONN_ID = "doczy_uat_snowflake" DAG_ID = "load_raw_training_data" -DATABASE="DOCZY_UAT" +DATABASE = "DOCZY_UAT" # bucket = "airflow-data-ingestion" -TAGS=["uat","training_data","dataload"] +TAGS = ["uat", "training_data", "dataload"] # Trigger rules -ALL_SUCCESS = 'all_success' -ALL_FAILED = 'all_failed' -ALL_DONE = 'all_done' -ONE_SUCCESS = 'one_success' -ONE_FAILED = 'one_failed' +ALL_SUCCESS = "all_success" +ALL_FAILED = "all_failed" +ALL_DONE = "all_done" +ONE_SUCCESS = "one_success" +ONE_FAILED = "one_failed" # Passing empty params for now, this will be overridden by the payload from the trigger # These params can also be set from the Airflow UI while manually triggering the DAG -default_params = {"column_config_file_name": "", "raw_training_data_file_name":"", "business_config_file_name":""} +default_params = { + "column_config_file_name": "", + "raw_training_data_file_name": "", + "business_config_file_name": "", +} # This will be replaced with the payload from the event after API connection is setup # training_results_file_name = "training_results_sample.csv" # attempt_logs_file_name = "attempt_logs_sample.csv" -def call_stored_proc(proc_name,file_type, params): - dwh_hook = SnowflakeHook(snowflake_conn_id=SNOWFLAKE_CONN_ID) +def call_stored_proc(proc_name, file_type, params): + dwh_hook = SnowflakeHook(snowflake_conn_id=SNOWFLAKE_CONN_ID) with dwh_hook.get_conn() as conn: # dwh_hook.set_autocommit(conn,autocommit=False) cur = conn.cursor() # Added new parameter file_type to determine the file name to be passed to the stored procedure # The bucket name is set by default to "doczy-dev-infra-raw-data-ingestion" and the files should be ALWAYS save under training_interface/ path for now - if file_type == 'column_config': - file_name = params['column_config_file_name'] - elif file_type == 'raw_training_data': - file_name = params['raw_training_data_file_name'] - elif file_type == 'business_config': - file_name = params['business_config_file_name'] + if file_type == "column_config": + file_name = params["column_config_file_name"] + elif file_type == "raw_training_data": + file_name = params["raw_training_data_file_name"] + elif file_type == "business_config": + file_name = params["business_config_file_name"] cur.execute(f"CALL {DATABASE}.STG.{proc_name}('{file_name}');") result = cur.fetchone() - if result[0] == 'Setup, Load, and Audit Complete': - logger.info('PROCEDURE EXECUTED SUCCESSFULLY') + if result[0] == "Setup, Load, and Audit Complete": + logger.info("PROCEDURE EXECUTED SUCCESSFULLY") else: - raise AirflowFailException("Check the DAG logs for more information. ERROR FROM SNOWFLAKE: ", result) + raise AirflowFailException( + "Check the DAG logs for more information. ERROR FROM SNOWFLAKE: ", + result, + ) logger.info(f"QUERY EXECUTION RESULT: {str(result)}") + dag = DAG( DAG_ID, start_date=datetime(2024, 1, 1), - default_args={"snowflake_conn_id": SNOWFLAKE_CONN_ID, "retries":0}, + default_args={"snowflake_conn_id": SNOWFLAKE_CONN_ID, "retries": 0}, tags=TAGS, catchup=False, schedule=None, - params = default_params + params=default_params, ) - -begin_job = EmptyOperator(task_id='Begin') + +begin_job = EmptyOperator(task_id="Begin") load_training_results = PythonOperator( task_id="load_column_config", python_callable=call_stored_proc, dag=dag, - op_kwargs={'proc_name':'LOAD_COLUMN_CONFIG', 'file_type':'column_config'} + op_kwargs={"proc_name": "LOAD_COLUMN_CONFIG", "file_type": "column_config"}, ) load_attempt_logs_sp = PythonOperator( task_id="load_raw_training_data", python_callable=call_stored_proc, dag=dag, - op_kwargs={'proc_name':'LOAD_TRAINING_DATA_RAW', 'file_type':'raw_training_data'} + op_kwargs={"proc_name": "LOAD_TRAINING_DATA_RAW", "file_type": "raw_training_data"}, ) load_business_config = PythonOperator( task_id="load_business_config", python_callable=call_stored_proc, dag=dag, - op_kwargs={'proc_name':'LOAD_BUSINESS_CONFIG', 'file_type':'business_config'} + op_kwargs={"proc_name": "LOAD_BUSINESS_CONFIG", "file_type": "business_config"}, ) - -end_job = EmptyOperator(task_id='End') +end_job = EmptyOperator(task_id="End") -begin_job >> load_training_results >> load_attempt_logs_sp >> load_business_config >> end_job \ No newline at end of file +( + begin_job + >> load_training_results + >> load_attempt_logs_sp + >> load_business_config + >> end_job +) diff --git a/fieldExtraction/src/bottom_up_funcs.py b/fieldExtraction/src/bottom_up_funcs.py index a1253ae..ae728ea 100644 --- a/fieldExtraction/src/bottom_up_funcs.py +++ b/fieldExtraction/src/bottom_up_funcs.py @@ -1,4 +1,3 @@ - import dict_operations import postprocessing_funcs import prompts @@ -8,12 +7,13 @@ import utils import difflib + def run_bottom_up(filename, text_dict): """ Processes the text of a document using a two-tiered Bottom Up approach to extract key financial and operational information. - This function first runs BOTTOM_UP_PRIMARY and BOTTOM_UP_SECONDARY prompts, then processes these initial results to further - refine and structure them into a dictionary form. + This function first runs BOTTOM_UP_PRIMARY and BOTTOM_UP_SECONDARY prompts, then processes these initial results to further + refine and structure them into a dictionary form. Parameters: filename (str): The name of the file being processed, used to tag output data. @@ -28,13 +28,13 @@ def run_bottom_up(filename, text_dict): answer_strings = run_bottom_up_primary(text_dict, 8000) answer_dicts = dict_operations.primary_string_to_dict(answer_strings, filename) answer_dicts_filtered = postprocessing_funcs.filter_service_column(answer_dicts) - + # Bottom Up Secondary results_dicts = run_bottom_up_secondary(answer_dicts_filtered, text_dict, 8000) for d in results_dicts: - d['Filename'] = filename - return results_dicts # List of dictionaries + d["Filename"] = filename + return results_dicts # List of dictionaries def run_bottom_up_primary(text_dict, tokens): @@ -51,15 +51,17 @@ def run_bottom_up_primary(text_dict, tokens): Returns: dict: A dictionary where each key is a page number and the value is the response from the language model. """ - #chunk_dict = preprocess.chunk_text(text_dict) + # chunk_dict = preprocess.chunk_text(text_dict) answer_dict = {} - #for page_number in chunk_dict.keys(): - #if '%' in chunk_dict[page_number] or '$' in chunk_dict[page_number]: - #prompt = prompts.BOTTOM_UP_PRIMARY(chunk_dict[page_number], config.CLIENT_NAME) + # for page_number in chunk_dict.keys(): + # if '%' in chunk_dict[page_number] or '$' in chunk_dict[page_number]: + # prompt = prompts.BOTTOM_UP_PRIMARY(chunk_dict[page_number], config.CLIENT_NAME) for page_number in text_dict.keys(): if utils.contains_reimbursement(text_dict, page_number): # Run Primary - prompt = prompts.BOTTOM_UP_PRIMARY(text_dict[page_number], config.CLIENT_NAME) + prompt = prompts.BOTTOM_UP_PRIMARY( + text_dict[page_number], config.CLIENT_NAME + ) answer = claude_funcs.invoke_claude_3(prompt, max_tokens=tokens) answer_dict[page_number] = answer return answer_dict @@ -69,10 +71,10 @@ def run_bottom_up_secondary(answer_dicts, text_dict, tokens): """ Executes the secondary Bottom Up processing phase on the results obtained from the primary Bottom Up analysis. - This function enhances the primary results with additional analyses based on configured conditions. Invokes + This function enhances the primary results with additional analyses based on configured conditions. Invokes Claude 3 with tailored prompts to generate structured information that complements the initial results. - Each piece of data processed possibly undergoes several rounds of checks and transformations, ensuring detailed and + Each piece of data processed possibly undergoes several rounds of checks and transformations, ensuring detailed and comprehensive output. Parameters: @@ -81,43 +83,50 @@ def run_bottom_up_secondary(answer_dicts, text_dict, tokens): tokens (int): Token limit for language model invocations. Returns: - list of dict: A list of dictionaries containing enriched and finalized structured data from both the primary and + list of dict: A list of dictionaries containing enriched and finalized structured data from both the primary and secondary analyses. """ # New rows for lesser of temp_dicts = [] - for d in answer_dicts: + for d in answer_dicts: # Bottom Up Lesser if config.RUN_LESSER: lesser_object = run_bottom_up_lesser(d.copy(), text_dict.copy(), tokens) - temp_dicts.append(lesser_object[1]) # Add original object from d - if lesser_object[0]: temp_dicts.append(lesser_object[0]) # If lesser, add lesser object + temp_dicts.append(lesser_object[1]) # Add original object from d + if lesser_object[0]: + temp_dicts.append(lesser_object[0]) # If lesser, add lesser object else: temp_dicts.append(d) - + # Add additional fields to each row final_dicts = [] for d in temp_dicts: if d is not None: - page_num = d['page_num'] - + page_num = d["page_num"] + # Bottom Up Methodology if config.RUN_METHODOLOGY: prompt = prompts.BOTTOM_UP_METHODOLOGY(d) - methodology_answer = claude_funcs.invoke_claude_3(prompt, max_tokens=100) - d['REIMBURSEMENT_METHODOLOGY'] = methodology_answer - + methodology_answer = claude_funcs.invoke_claude_3( + prompt, max_tokens=100 + ) + d["REIMBURSEMENT_METHODOLOGY"] = methodology_answer + # # Bottom Up FS if config.RUN_FS: prompt = prompts.BOTTOM_UP_FS(d) - fs_answer = claude_funcs.invoke_claude_3(prompt, model_id=config.MODEL_ID_CLAUDE3_HAIKU, max_tokens=tokens) + fs_answer = claude_funcs.invoke_claude_3( + prompt, model_id=config.MODEL_ID_CLAUDE3_HAIKU, max_tokens=tokens + ) fs_dict = dict_operations.secondary_string_to_dict(fs_answer) d.update(fs_dict) # # Bottom Up Exception/Escalator if config.RUN_EXCEPTION: prompt = prompts.BOTTOM_UP_EXCEPT_ESC(d, text_dict[page_num]) - exc_answer = claude_funcs.invoke_claude_3(prompt, model_id=config.MODEL_ID_CLAUDE3_HAIKU, max_tokens=tokens) + exc_answer = claude_funcs.invoke_claude_3( + prompt, model_id=config.MODEL_ID_CLAUDE3_HAIKU, max_tokens=tokens + ) exc_dict = dict_operations.secondary_string_to_dict(exc_answer) d.update(exc_dict) @@ -129,7 +138,7 @@ def run_bottom_up_secondary(answer_dicts, text_dict, tokens): d.update(codes_dict) final_dicts.append(d) - + return final_dicts @@ -153,62 +162,74 @@ def run_bottom_up_lesser(d, text_dict, tokens=4000): containing the original or updated data depending on the presence of such language. """ - prompt = prompts.BOTTOM_UP_LESSER(d, text_dict[d['page_num']]) + prompt = prompts.BOTTOM_UP_LESSER(d, text_dict[d["page_num"]]) lesser_of_answer = claude_funcs.invoke_claude_3(prompt, max_tokens=tokens) - lesser_of_dict_list = dict_operations.primary_string_to_dict({d['page_num'] : lesser_of_answer}, d['Filename']) + lesser_of_dict_list = dict_operations.primary_string_to_dict( + {d["page_num"]: lesser_of_answer}, d["Filename"] + ) # print(lesser_of_dict_list) # If lesser of is Y if len(lesser_of_dict_list) > 1: lesser_of_dict = get_lesser_of_dict(lesser_of_dict_list, d) else: - d['LESSER_OF_LANGUAGE_IND'] = 'N' - d['GREATER_OF_LANGUAGE_IND'] = 'N' - return ({}, d) - - lesser_of_dict['page_num'] = d['page_num'] - if lesser_of_dict['LESSER_OF_LANGUAGE_IND'] == 'Y' or lesser_of_dict['GREATER_OF_LANGUAGE_IND'] == 'Y': - lesser_of_dict['SERVICE'] = d['SERVICE'] - return (lesser_of_dict, d) - else: - d['LESSER_OF_LANGUAGE_IND'] = 'N' - d['GREATER_OF_LANGUAGE_IND'] = 'N' + d["LESSER_OF_LANGUAGE_IND"] = "N" + d["GREATER_OF_LANGUAGE_IND"] = "N" return ({}, d) + lesser_of_dict["page_num"] = d["page_num"] + if ( + lesser_of_dict["LESSER_OF_LANGUAGE_IND"] == "Y" + or lesser_of_dict["GREATER_OF_LANGUAGE_IND"] == "Y" + ): + lesser_of_dict["SERVICE"] = d["SERVICE"] + return (lesser_of_dict, d) + else: + d["LESSER_OF_LANGUAGE_IND"] = "N" + d["GREATER_OF_LANGUAGE_IND"] = "N" + return ({}, d) def get_least_similar(dict_list, original_dict, field): # print(f"Running similarity for {field}") - - min_similarity = float('inf') + + min_similarity = float("inf") least_similar_dict = None for dictionary in dict_list: # Extract the methodology text field_text = dictionary[field] # print(f"Field text: {field_text}") - + # Compute similarity using difflib - similarity = difflib.SequenceMatcher(None, str(field_text), str(original_dict[field])).ratio() + similarity = difflib.SequenceMatcher( + None, str(field_text), str(original_dict[field]) + ).ratio() # print(f"Simlarity: {similarity}") - + # If the similarity is less than the current minimum, update the minimum and the corresponding dictionary if similarity < min_similarity: min_similarity = similarity least_similar_dict = dictionary - + return least_similar_dict + def get_lesser_of_dict(dict_list, original_dict): - unique_rates = list({d['REIMBURSEMENT_RATE'] for d in dict_list}) - unique_fees = list({d['REIMBURSEMENT_FLAT_FEE'] for d in dict_list}) + unique_rates = list({d["REIMBURSEMENT_RATE"] for d in dict_list}) + unique_fees = list({d["REIMBURSEMENT_FLAT_FEE"] for d in dict_list}) - # If the rates are different, return the lesser of dict with a different rate than original + # If the rates are different, return the lesser of dict with a different rate than original if len(unique_rates) > 1: - least_similar_dict = get_least_similar(dict_list, original_dict, 'REIMBURSEMENT_RATE') + least_similar_dict = get_least_similar( + dict_list, original_dict, "REIMBURSEMENT_RATE" + ) elif len(unique_fees) > 1: - least_similar_dict = get_least_similar(dict_list, original_dict, 'REIMBURSEMENT_FLAT_FEE') + least_similar_dict = get_least_similar( + dict_list, original_dict, "REIMBURSEMENT_FLAT_FEE" + ) else: - least_similar_dict = get_least_similar(dict_list, original_dict, 'FULL_METHODOLOGY') + least_similar_dict = get_least_similar( + dict_list, original_dict, "FULL_METHODOLOGY" + ) return least_similar_dict - diff --git a/fieldExtraction/src/carveouts.py b/fieldExtraction/src/carveouts.py index 13843ad..675cb88 100644 --- a/fieldExtraction/src/carveouts.py +++ b/fieldExtraction/src/carveouts.py @@ -3,11 +3,12 @@ from difflib import SequenceMatcher import re import time -import prompts +import prompts import claude_funcs + def get_closest_substring_match(val, valid_values): - """ Returns the first match from valid_values using a case-insensitive substring match. """ + """Returns the first match from valid_values using a case-insensitive substring match.""" if pd.isna(val): return None val = val.strip().upper() @@ -16,6 +17,7 @@ def get_closest_substring_match(val, valid_values): return valid_value return None + def get_best_carveout_from_claude(carveout_list, service): print("using claude to get best carveout...") prompt = f"Given the service '{service}', please choose the best matching carveout from the following list: {', '.join(carveout_list)}. ONLY SELECT ONE FROM THE LIST. DO NOT RETURN A SENTENCE" @@ -27,6 +29,7 @@ def get_best_carveout_from_claude(carveout_list, service): print(f"Error occurred: {e}. Waiting for 60 seconds before retrying...") time.sleep(60) + def check_prov_type_similarity(prov_type, carveout): print("checking similarity between prov_type and carveout with claude...") prompt = f"Do the provider type '{prov_type}' and the carveout '{carveout}' mean the same thing or are they very similar? ONLY ANSWER WITH 'True' OR 'False'" @@ -38,6 +41,7 @@ def check_prov_type_similarity(prov_type, carveout): print(f"Error occurred: {e}. Waiting for 60 seconds before retrying...") time.sleep(60) + def label_services(filepath, carveout_list, output_filepath): # Read the CSV file df = pd.read_csv(filepath) @@ -46,13 +50,24 @@ def label_services(filepath, carveout_list, output_filepath): df = df.head(500) # Define primary service terms - primary_terms = ['Covered Services', 'Inpatient Services', 'Outpatient Services', 'Physician Services', - 'Inpatient', 'Outpatient', 'Physician', 'Medical', 'Surgical', 'Diagnostic', 'Therapeutic'] + primary_terms = [ + "Covered Services", + "Inpatient Services", + "Outpatient Services", + "Physician Services", + "Inpatient", + "Outpatient", + "Physician", + "Medical", + "Surgical", + "Diagnostic", + "Therapeutic", + ] # Initialize columns for the results - df['carveout_matched'] = '' - df['label'] = '' - df['IS_CARVEOUT'] = '' + df["carveout_matched"] = "" + df["label"] = "" + df["IS_CARVEOUT"] = "" # Initialize a nested dictionary to track occurrences for each Filename and TD_LOB occurrence_tracker = {} @@ -61,10 +76,10 @@ def label_services(filepath, carveout_list, output_filepath): # Iterate through the rows in the DataFrame for index, row in df.iterrows(): - service = str(row['SERVICE']) - prov_type = str(row['PROV_TYPE']) if pd.notna(row['PROV_TYPE']) else '' - td_lob = row['TD_LOB'] - filename = row['Filename'] + service = str(row["SERVICE"]) + prov_type = str(row["PROV_TYPE"]) if pd.notna(row["PROV_TYPE"]) else "" + td_lob = row["TD_LOB"] + filename = row["Filename"] # Initialize the nested dictionary for the filename and TD_LOB if not present if filename not in occurrence_tracker: @@ -82,49 +97,57 @@ def label_services(filepath, carveout_list, output_filepath): # Determine the label based on the count of this primary term for the given Filename and TD_LOB count = occurrence_tracker[filename][td_lob][primary] if count == 0: - df.at[index, 'label'] = 'primary' + df.at[index, "label"] = "primary" elif count == 1: - df.at[index, 'label'] = 'secondary' + df.at[index, "label"] = "secondary" elif count == 2: - df.at[index, 'label'] = 'tertiary' + df.at[index, "label"] = "tertiary" else: - df.at[index, 'label'] = 'additional' + df.at[index, "label"] = "additional" occurrence_tracker[filename][td_lob][primary] += 1 - df.at[index, 'IS_CARVEOUT'] = 'N' + df.at[index, "IS_CARVEOUT"] = "N" else: # If PROV_TYPE is blank, continue with carveout determination - best_carveout_response = get_best_carveout_from_claude(carveout_list, service) + best_carveout_response = get_best_carveout_from_claude( + carveout_list, service + ) matched_carveout = best_carveout_response # Use the model response directly - df.at[index, 'carveout_matched'] = matched_carveout + df.at[index, "carveout_matched"] = matched_carveout if matched_carveout: # Check similarity between PROV_TYPE and carveout - prov_type_similarity_score = SequenceMatcher(None, prov_type, matched_carveout).ratio() + prov_type_similarity_score = SequenceMatcher( + None, prov_type, matched_carveout + ).ratio() if prov_type_similarity_score < 0.8: - prov_type_similarity = check_prov_type_similarity(prov_type, matched_carveout) + prov_type_similarity = check_prov_type_similarity( + prov_type, matched_carveout + ) else: prov_type_similarity = "True" - if re.search(r'\b(True|Yes)\b', prov_type_similarity, re.IGNORECASE): + if re.search(r"\b(True|Yes)\b", prov_type_similarity, re.IGNORECASE): if matched_carveout not in occurrence_tracker[filename][td_lob]: - print('claude returned a match between provider type and service') + print( + "claude returned a match between provider type and service" + ) occurrence_tracker[filename][td_lob][matched_carveout] = 0 count = occurrence_tracker[filename][td_lob][matched_carveout] if count == 0: - df.at[index, 'label'] = 'primary' + df.at[index, "label"] = "primary" elif count == 1: - df.at[index, 'label'] = 'secondary' + df.at[index, "label"] = "secondary" elif count == 2: - df.at[index, 'label'] = 'tertiary' + df.at[index, "label"] = "tertiary" else: - df.at[index, 'label'] = 'additional' + df.at[index, "label"] = "additional" occurrence_tracker[filename][td_lob][matched_carveout] += 1 - df.at[index, 'IS_CARVEOUT'] = 'N' + df.at[index, "IS_CARVEOUT"] = "N" else: - df.at[index, 'IS_CARVEOUT'] = 'Y' + df.at[index, "IS_CARVEOUT"] = "Y" else: - df.at[index, 'IS_CARVEOUT'] = 'Y' + df.at[index, "IS_CARVEOUT"] = "Y" carveout_count += 1 @@ -136,4 +159,5 @@ def label_services(filepath, carveout_list, output_filepath): # Final save of the updated DataFrame to a new CSV file df.to_csv(output_filepath, index=False) -label_services('service.csv', prompts.get_carveout_list(), 'carveouts50.csv') + +label_services("service.csv", prompts.get_carveout_list(), "carveouts50.csv") diff --git a/fieldExtraction/src/claude_funcs.py b/fieldExtraction/src/claude_funcs.py index 9d1e9f8..75abfac 100644 --- a/fieldExtraction/src/claude_funcs.py +++ b/fieldExtraction/src/claude_funcs.py @@ -1,9 +1,9 @@ - import json import anthropic import config + # Claude calls def invoke_claude_2(prompt, max_tokens): """ @@ -21,28 +21,30 @@ def invoke_claude_2(prompt, max_tokens): str: The text generated by the Claude 2 model in response to the input prompt. """ body = json.dumps( - {"prompt": anthropic.HUMAN_PROMPT + prompt + anthropic.AI_PROMPT, - "max_tokens_to_sample": max_tokens, - "temperature":0.0, - "top_p":1, - "top_k":250, - "stop_sequences":[anthropic.HUMAN_PROMPT] + { + "prompt": anthropic.HUMAN_PROMPT + prompt + anthropic.AI_PROMPT, + "max_tokens_to_sample": max_tokens, + "temperature": 0.0, + "top_p": 1, + "top_k": 250, + "stop_sequences": [anthropic.HUMAN_PROMPT], } - ) + ) response = config.BEDROCK_RUNTIME.invoke_model( - body=body, - modelId=config.MODEL_ID_CLAUDE2, - accept="application/json", - contentType="application/json" + body=body, + modelId=config.MODEL_ID_CLAUDE2, + accept="application/json", + contentType="application/json", ) response_body = json.loads(response.get("body").read()) - response_text = response_body['completion'] + response_text = response_body["completion"] return response_text -def invoke_claude_3(prompt, model_id=config.MODEL_ID_CLAUDE3_SONNET, max_tokens = 0): + +def invoke_claude_3(prompt, model_id=config.MODEL_ID_CLAUDE3_SONNET, max_tokens=0): """ Invokes the Claude 3 language model with specified parameters to generate a response based on the input prompt. @@ -57,32 +59,25 @@ def invoke_claude_3(prompt, model_id=config.MODEL_ID_CLAUDE3_SONNET, max_tokens Returns: str: The text generated by the Claude 2 model in response to the input prompt. """ - prompt = prompt - body = json.dumps({ - "anthropic_version": "bedrock-2023-05-31", - "max_tokens": max_tokens, - "temperature": 0.0, - "messages": [ - { - "role": "user", - "content": [ - { - "type": "text", - "text":prompt - } - ] - } - ] - } - ) + prompt = prompt + body = json.dumps( + { + "anthropic_version": "bedrock-2023-05-31", + "max_tokens": max_tokens, + "temperature": 0.0, + "messages": [ + {"role": "user", "content": [{"type": "text", "text": prompt}]} + ], + } + ) response = config.BEDROCK_RUNTIME.invoke_model( - body=body, - modelId=model_id, - accept="application/json", - contentType="application/json" + body=body, + modelId=model_id, + accept="application/json", + contentType="application/json", ) response_body = json.loads(response.get("body").read()) - return (response_body['content'][0]['text']) \ No newline at end of file + return response_body["content"][0]["text"] diff --git a/fieldExtraction/src/config.py b/fieldExtraction/src/config.py index 0a1ca49..94cfbdb 100644 --- a/fieldExtraction/src/config.py +++ b/fieldExtraction/src/config.py @@ -1,40 +1,59 @@ - - from datetime import datetime import boto3 -#from llama_index.llms.bedrock import Bedrock -#pip install llama-index-llms-bedrock + +# from llama_index.llms.bedrock import Bedrock +# pip install llama-index-llms-bedrock # General Settings -TEST = False # True to run test prompt - just for testing model connection +TEST = False # True to run test prompt - just for testing model connection VERBOSE = True -CLIENT_NAME = '' +CLIENT_NAME = "" TODAY = datetime.now().strftime("%Y%m%d") # Valid values -VALID_LOBS = ['MEDICARE', 'MEDICARE ADVANTAGE', 'MEDICAID', 'MARKETPLACE', 'COMMERCIAL', 'GROUP', 'MEDICARE-MEDICAID'] -VALID_PROGRAMS = ['CHIP', 'CHIP-P', 'CHIP-PERINATE', 'STAR', 'STAR+PLUS', 'MA', 'DUAL SPECIAL NEEDS PLAN', 'DSNP', 'DUAL'] -VALID_NETWORKS = ['HMO', 'PPO', 'EPO', 'POS', 'FFS'] - +VALID_LOBS = [ + "MEDICARE", + "MEDICARE ADVANTAGE", + "MEDICAID", + "MARKETPLACE", + "COMMERCIAL", + "GROUP", + "MEDICARE-MEDICAID", +] +VALID_PROGRAMS = [ + "CHIP", + "CHIP-P", + "CHIP-PERINATE", + "STAR", + "STAR+PLUS", + "MA", + "DUAL SPECIAL NEEDS PLAN", + "DSNP", + "DUAL", +] +VALID_NETWORKS = ["HMO", "PPO", "EPO", "POS", "FFS"] + # Input Settings -READ_MODE = '_LOCAL_' # OR '_S3_' +READ_MODE = "_LOCAL_" # OR '_S3_' FILTER_ALREADY_PROCESSED = True # Output Settings -WRITE_OUTPUT = True # True writes csvs, False prints result in console but no output written -OUTPUT_DIRECTORY = 'output' -CONSOLIDATED_OUTPUT_DIRECTORY = 'output_consolidated' -OUTPUT_CSV_PATH = f'consolidated_output_{TODAY}.csv' -TD_RESULTS_NAME = 'td_results.csv' -BU_RESULTS_NAME = 'bu_results.csv' -UNPROCESSED_RESULTS_NAME = 'combined_results_unprocessed.csv' -PROCESSED_RESULTS_NAME = 'combined_results_post_processed.csv' +WRITE_OUTPUT = ( + True # True writes csvs, False prints result in console but no output written +) +OUTPUT_DIRECTORY = "output" +CONSOLIDATED_OUTPUT_DIRECTORY = "output_consolidated" +OUTPUT_CSV_PATH = f"consolidated_output_{TODAY}.csv" +TD_RESULTS_NAME = "td_results.csv" +BU_RESULTS_NAME = "bu_results.csv" +UNPROCESSED_RESULTS_NAME = "combined_results_unprocessed.csv" +PROCESSED_RESULTS_NAME = "combined_results_post_processed.csv" # Multithread Settings MAX_WORKERS = 10 # Prompt Debugging -RUN_PRIMARY = True # Always True +RUN_PRIMARY = True # Always True RUN_LOB = True RUN_LESSER = True RUN_METHODOLOGY = True @@ -44,53 +63,79 @@ RUN_CODES = True # Postprocessing Settings FUZZY_MATCH_THRESHOLD = 0.8 -VALID_COLUMNS = ['Filename', 'page_num', 'SERVICE', 'REIMBURSEMENT_FLAT_FEE', 'REIMBURSEMENT_RATE', 'FULL_METHODOLOGY', -'REIMBURSEMENT_METHODOLOGY', 'REIMBURSEMENT_FEE_SCHEDULE', 'REIMBURSEMENT_FEE_SCHEDULE_VERSION', -'LESSER_OF_LANGUAGE_IND', 'GREATER_OF_LANGUAGE_IND', -'CONTRACT_LOB', 'CONTRACT_NETWORK', 'PRODUCT', 'CONTRACT_PROGRAM', 'CONTRACT_MARKETPLACE_METAL_LEVEL', 'PROV_TYPE', -'REIMBURSEMENT_PROC_CODES', 'REIMBURSEMENT_PROC_CODE_MODIFIERS', 'REIMBURSEMENT_REVENUE_CODES', 'REIMBURSEMENT_STATUS_INDICATOR_CODES', -'REIMBURSEMENT_DIAG_CODES', 'REIMBURSEMENT_GROUPER_CODES', 'REIMBURSEMENT_GROUPER', 'REIMBURSEMENT_PLACEOFSERVICE_CODES', 'REIMBURSEMENT_ADMITTYPE_CODES', -'REIMBURSEMENT_EXCEPTION_IND', 'REIMBURSEMENT_DESCRIBE_EXCEPTION', -'RATE_ESCALATOR_IND', 'RATE_ESCALATOR_BASIS', 'RATE_ESCALATOR_YEARLY_PERCENT_INCREASE', -'Corrected_LOB', 'Corrected_PROGRAM', 'Corrected_NETWORK'] +VALID_COLUMNS = [ + "Filename", + "page_num", + "SERVICE", + "REIMBURSEMENT_FLAT_FEE", + "REIMBURSEMENT_RATE", + "FULL_METHODOLOGY", + "REIMBURSEMENT_METHODOLOGY", + "REIMBURSEMENT_FEE_SCHEDULE", + "REIMBURSEMENT_FEE_SCHEDULE_VERSION", + "LESSER_OF_LANGUAGE_IND", + "GREATER_OF_LANGUAGE_IND", + "CONTRACT_LOB", + "CONTRACT_NETWORK", + "PRODUCT", + "CONTRACT_PROGRAM", + "CONTRACT_MARKETPLACE_METAL_LEVEL", + "PROV_TYPE", + "REIMBURSEMENT_PROC_CODES", + "REIMBURSEMENT_PROC_CODE_MODIFIERS", + "REIMBURSEMENT_REVENUE_CODES", + "REIMBURSEMENT_STATUS_INDICATOR_CODES", + "REIMBURSEMENT_DIAG_CODES", + "REIMBURSEMENT_GROUPER_CODES", + "REIMBURSEMENT_GROUPER", + "REIMBURSEMENT_PLACEOFSERVICE_CODES", + "REIMBURSEMENT_ADMITTYPE_CODES", + "REIMBURSEMENT_EXCEPTION_IND", + "REIMBURSEMENT_DESCRIBE_EXCEPTION", + "RATE_ESCALATOR_IND", + "RATE_ESCALATOR_BASIS", + "RATE_ESCALATOR_YEARLY_PERCENT_INCREASE", + "Corrected_LOB", + "Corrected_PROGRAM", + "Corrected_NETWORK", +] # AWS Keys -AWS_ACCESS_KEY_ID="ASIAZTMXAXNXHQSWW2R5" -AWS_SECRET_ACCESS_KEY="W9zH4nDGaae8SuJqcPDXyTHFpQ4xV43eRjbDfPGo" -AWS_SESSION_TOKEN="IQoJb3JpZ2luX2VjEIX//////////wEaCXVzLWVhc3QtMiJIMEYCIQCjQVKSzQSWuO8bDdeGjDUJNUc1NwOA795yLFiBGDSslwIhAOccVQAzlBZyOVtU5mZztHVRbqEbrFigeg2aPzlczG1QKoMDCG8QABoMNjYwMTMxMDY4NzgyIgyP2pSGtJAhII6aZIMq4AI3XnnJgTtwbj6u15RcBlTCYgyn+SqIASHvEaJzyWHP/b7Vf5ZNyWKBsccgIZ590De/biM4Gb7tLsnWYCkFuyj3XswHKgyInBDqDg6b4olgyMGoDxrDmNUI28tQxB5XpwKxsas/azzwNPAgRAyXWMK4b1DcslPGpTRBBi+eBF+hTLiGk1jQnip0qR+GuJnCwXe66gcIGmgscjVxGSg2Q5n0O7VD2UOrr5T7BhvnGLiNtlrXqd3f+BFVk6/SzmuQJgsK33rtsmQUI7R5LXlyZEfH3cEkCXrdch6bFO4Wu53oBmiDtNEkle5n95TXCnXC1EchWzuDA6iP061kLRF3jdihMls8St/RoNACCR41YtOuU5oH1n5qXGvH1HQzOvBuSQBCKIqVGgB7hU63mbJlQCdX+0qd0jgteXFHMyjrwr604zPCxdcfbyehAwpxJR8xLVMimgxJqvyzvAoUjzW1sZPPMKG4rLUGOqUBOobUcGB1mI6GHaV6p6+t+qU/BG9FCt9BZbC/IJsHQkWWemCWpvV6lGVM5JPBmmgGQc3LiwH8g5JSHnoUDjTWlLvlHjXZ2ODLJP+aDFcuS4b1WZmXedL6RTsMvg8qhiELEB7rR+6jEmBmXvh56c+R4mkCG5AiWRBZ/ah910JnvKy5wQTOGJGzNxflyiJmwoO2JOZAKBLT0jg78iG1T3DMJ2qc+Qx7" +AWS_ACCESS_KEY_ID = "ASIAZTMXAXNXHQSWW2R5" +AWS_SECRET_ACCESS_KEY = "W9zH4nDGaae8SuJqcPDXyTHFpQ4xV43eRjbDfPGo" +AWS_SESSION_TOKEN = "IQoJb3JpZ2luX2VjEIX//////////wEaCXVzLWVhc3QtMiJIMEYCIQCjQVKSzQSWuO8bDdeGjDUJNUc1NwOA795yLFiBGDSslwIhAOccVQAzlBZyOVtU5mZztHVRbqEbrFigeg2aPzlczG1QKoMDCG8QABoMNjYwMTMxMDY4NzgyIgyP2pSGtJAhII6aZIMq4AI3XnnJgTtwbj6u15RcBlTCYgyn+SqIASHvEaJzyWHP/b7Vf5ZNyWKBsccgIZ590De/biM4Gb7tLsnWYCkFuyj3XswHKgyInBDqDg6b4olgyMGoDxrDmNUI28tQxB5XpwKxsas/azzwNPAgRAyXWMK4b1DcslPGpTRBBi+eBF+hTLiGk1jQnip0qR+GuJnCwXe66gcIGmgscjVxGSg2Q5n0O7VD2UOrr5T7BhvnGLiNtlrXqd3f+BFVk6/SzmuQJgsK33rtsmQUI7R5LXlyZEfH3cEkCXrdch6bFO4Wu53oBmiDtNEkle5n95TXCnXC1EchWzuDA6iP061kLRF3jdihMls8St/RoNACCR41YtOuU5oH1n5qXGvH1HQzOvBuSQBCKIqVGgB7hU63mbJlQCdX+0qd0jgteXFHMyjrwr604zPCxdcfbyehAwpxJR8xLVMimgxJqvyzvAoUjzW1sZPPMKG4rLUGOqUBOobUcGB1mI6GHaV6p6+t+qU/BG9FCt9BZbC/IJsHQkWWemCWpvV6lGVM5JPBmmgGQc3LiwH8g5JSHnoUDjTWlLvlHjXZ2ODLJP+aDFcuS4b1WZmXedL6RTsMvg8qhiELEB7rR+6jEmBmXvh56c+R4mkCG5AiWRBZ/ah910JnvKy5wQTOGJGzNxflyiJmwoO2JOZAKBLT0jg78iG1T3DMJ2qc+Qx7" # File Paths -LOCAL_PATH = 'data/texas_childrens' # Replace with local +LOCAL_PATH = "data/texas_childrens" # Replace with local # S3 Settings -S3_CLIENT = boto3.client('s3', - region_name="us-east-2", - aws_access_key_id=AWS_ACCESS_KEY_ID, - aws_secret_access_key=AWS_SECRET_ACCESS_KEY, - aws_session_token=AWS_SESSION_TOKEN - ) +S3_CLIENT = boto3.client( + "s3", + region_name="us-east-2", + aws_access_key_id=AWS_ACCESS_KEY_ID, + aws_secret_access_key=AWS_SECRET_ACCESS_KEY, + aws_session_token=AWS_SESSION_TOKEN, +) BUCKET = "doczy-dev-infra-textract" -PREFIX = "batches/batch_1/contract-text-file/" # replace with s3 path +PREFIX = "batches/batch_1/contract-text-file/" # replace with s3 path # Bedrock Settings -BEDROCK_RUNTIME = boto3.client(service_name="bedrock-runtime", - region_name="us-east-1", - aws_access_key_id=AWS_ACCESS_KEY_ID, - aws_secret_access_key=AWS_SECRET_ACCESS_KEY, - aws_session_token=AWS_SESSION_TOKEN +BEDROCK_RUNTIME = boto3.client( + service_name="bedrock-runtime", + region_name="us-east-1", + aws_access_key_id=AWS_ACCESS_KEY_ID, + aws_secret_access_key=AWS_SECRET_ACCESS_KEY, + aws_session_token=AWS_SESSION_TOKEN, ) -MODEL_ID_CLAUDE3_HAIKU = 'anthropic.claude-3-haiku-20240307-v1:0' -MODEL_ID_CLAUDE3_SONNET = 'anthropic.claude-3-sonnet-20240229-v1:0' -MODEL_ID_CLAUDE2 = 'anthropic.claude-instant-v1' +MODEL_ID_CLAUDE3_HAIKU = "anthropic.claude-3-haiku-20240307-v1:0" +MODEL_ID_CLAUDE3_SONNET = "anthropic.claude-3-sonnet-20240229-v1:0" +MODEL_ID_CLAUDE2 = "anthropic.claude-instant-v1" # LLama Model # LLAMA_MODEL = Bedrock( - # model=MODEL_ID_CLAUDE3_HAIKU, - # aws_access_key_id=AWS_ACCESS_KEY_ID, - # aws_secret_access_key=AWS_SECRET_ACCESS_KEY, - # aws_session_token=AWS_SESSION_TOKEN, - # region_name="us-east-1" - # ) - - - +# model=MODEL_ID_CLAUDE3_HAIKU, +# aws_access_key_id=AWS_ACCESS_KEY_ID, +# aws_secret_access_key=AWS_SECRET_ACCESS_KEY, +# aws_session_token=AWS_SESSION_TOKEN, +# region_name="us-east-1" +# ) diff --git a/fieldExtraction/src/consolidate_output.py b/fieldExtraction/src/consolidate_output.py index a7c6f76..ec95383 100644 --- a/fieldExtraction/src/consolidate_output.py +++ b/fieldExtraction/src/consolidate_output.py @@ -1,4 +1,3 @@ - import os import pandas as pd @@ -12,7 +11,7 @@ import config # all_dfs.append(df) # except: # continue - + # final_df = pd.concat(all_dfs, ignore_index=True) # print(final_df.shape) # print(len(final_df.Filename.unique())) @@ -20,9 +19,9 @@ import config # final_df.to_csv(os.path.join(config.CONSOLIDATED_OUTPUT_DIRECTORY, config.OUTPUT_CSV_PATH)) -ac_path = 'T:\\AArete Client Work\\Texas Children’s Health Plan\\Restricted\\Doczy_Output\\TX_Childrens_1-1.csv' -b_path = 'T:\\AArete Client Work\\Texas Children’s Health Plan\\Restricted\\Doczy_Output\\TX_Childrens_1-N.csv' -output_path = f'T:\\AArete Client Work\\Texas Children’s Health Plan\\Restricted\\Doczy_Output\\{config.TODAY}_TX_Childrens_460_Doczy_AI_Output.csv' +ac_path = "T:\\AArete Client Work\\Texas Children’s Health Plan\\Restricted\\Doczy_Output\\TX_Childrens_1-1.csv" +b_path = "T:\\AArete Client Work\\Texas Children’s Health Plan\\Restricted\\Doczy_Output\\TX_Childrens_1-N.csv" +output_path = f"T:\\AArete Client Work\\Texas Children’s Health Plan\\Restricted\\Doczy_Output\\{config.TODAY}_TX_Childrens_460_Doczy_AI_Output.csv" # Read the files into pandas DataFrames ac = pd.read_csv(ac_path) @@ -31,8 +30,7 @@ b = pd.read_csv(b_path) print(b.shape) # Merge the DataFrames on the 'Filename' column -merged_df = pd.merge(ac, b, on='Filename', how='right') +merged_df = pd.merge(ac, b, on="Filename", how="right") merged_df.to_csv(output_path) print(merged_df.shape) - diff --git a/fieldExtraction/src/dict_operations.py b/fieldExtraction/src/dict_operations.py index 78ee6fa..561ca17 100644 --- a/fieldExtraction/src/dict_operations.py +++ b/fieldExtraction/src/dict_operations.py @@ -1,8 +1,8 @@ - import re import json from collections import defaultdict + def secondary_string_to_dict(dict_string): """ Converts a string representation of a dictionary into an actual dictionary object, handling potential formatting issues. @@ -17,20 +17,21 @@ def secondary_string_to_dict(dict_string): Returns: dict: The dictionary obtained from parsing the cleaned and corrected string. """ - dict_string = dict_string.replace('<<<', '') - dict_string = dict_string.replace('>>>', '') - dict_string = dict_string.replace('\n', '') # Remove new line + dict_string = dict_string.replace("<<<", "") + dict_string = dict_string.replace(">>>", "") + dict_string = dict_string.replace("\n", "") # Remove new line - start_index = dict_string.find('{') - end_index = dict_string.rfind('}') + 1 + start_index = dict_string.find("{") + end_index = dict_string.rfind("}") + 1 dict_substring = dict_string[start_index:end_index] try: result_dict = json.loads(dict_substring) except: dict_substring = dict_substring.replace("'", '"') - result_dict = json.loads(dict_substring) + result_dict = json.loads(dict_substring) return result_dict + def primary_string_to_dict(string_dict, filename): """ Converts a dictionary of strings, where each string represents multiple dictionary entries, into a list of dictionaries, @@ -49,24 +50,25 @@ def primary_string_to_dict(string_dict, filename): their page number and filename. """ data = [] - pattern = r'\{.*?\}' + pattern = r"\{.*?\}" for page_num in string_dict.keys(): primary_list = string_dict[page_num] - dicts = primary_list.split('[')[1] # Strip front - dicts = dicts.split(']')[0] # Strip back - dicts = dicts.replace('\n', '') # Remove new lines - dicts = dicts.replace('<<<', '') - dicts = dicts.replace('>>>', '') + dicts = primary_list.split("[")[1] # Strip front + dicts = dicts.split("]")[0] # Strip back + dicts = dicts.replace("\n", "") # Remove new lines + dicts = dicts.replace("<<<", "") + dicts = dicts.replace(">>>", "") dicts = re.sub(r"(? 1:\n", - " page_num = bu['page_num']\n", + " page_num = bu[\"page_num\"]\n", " page_text = text_dict.get(page_num, \"\")\n", - " prompt = GENERATE_PROMPT(bu, td_on_page, page_text, field_name=key, values=merged[key])\n", + " prompt = GENERATE_PROMPT(\n", + " bu, td_on_page, page_text, field_name=key, values=merged[key]\n", + " )\n", " response = claude_funcs.invoke_claude_3(prompt, max_tokens=4000)\n", " try:\n", " selected_value = response.strip()\n", " merged[key] = selected_value\n", " except Exception as e:\n", " print(f\"Error processing LLM response for key {key}: {e}\")\n", - " merged[key] = ', '.join(merged[key])\n", - " \n", + " merged[key] = \", \".join(merged[key])\n", + "\n", " merged_results.append(merged)\n", - " \n", + "\n", " return merged_results\n", "\n", + "\n", "def process_file(filename, input_dict):\n", " contract_text = input_dict[filename]\n", "\n", " # Preprocess\n", " contract_text = preprocess.clean_newlines(contract_text)\n", " text_dict = preprocess.split_text(contract_text)\n", - " \n", + "\n", " # Run Top Down\n", - " td_results = prompt_funcs.run_top_down(filename, text_dict) # Returns list of dictionaries for each page\n", + " td_results = prompt_funcs.run_top_down(\n", + " filename, text_dict\n", + " ) # Returns list of dictionaries for each page\n", " print(\"TD Results:\", td_results)\n", - " \n", + "\n", " # Run Bottom Up\n", - " bu_results = prompt_funcs.run_bottom_up(filename, text_dict) # Returns list of dictionaries\n", + " bu_results = prompt_funcs.run_bottom_up(\n", + " filename, text_dict\n", + " ) # Returns list of dictionaries\n", " print(\"BU Results:\", bu_results)\n", - " \n", + "\n", " # Combine\n", " combined_results = merge_results(clean_td(td_results), bu_results, text_dict)\n", " print(\"Combined Results:\", combined_results)\n", - " \n", + "\n", " # Create directories\n", " base_filename = os.path.splitext(filename)[0]\n", - " output_dir = os.path.join('output', base_filename)\n", + " output_dir = os.path.join(\"output\", base_filename)\n", " os.makedirs(output_dir, exist_ok=True)\n", - " \n", + "\n", " # Save results\n", - " pd.DataFrame(td_results).to_csv(os.path.join(output_dir, 'td_results.csv'), index=False)\n", - " pd.DataFrame(bu_results).to_csv(os.path.join(output_dir, 'bu_results.csv'), index=False)\n", - " pd.DataFrame(combined_results).to_csv(os.path.join(output_dir, 'combined_results.csv'), index=False)\n", - " \n", + " pd.DataFrame(td_results).to_csv(\n", + " os.path.join(output_dir, \"td_results.csv\"), index=False\n", + " )\n", + " pd.DataFrame(bu_results).to_csv(\n", + " os.path.join(output_dir, \"bu_results.csv\"), index=False\n", + " )\n", + " pd.DataFrame(combined_results).to_csv(\n", + " os.path.join(output_dir, \"combined_results.csv\"), index=False\n", + " )\n", + "\n", + "\n", "def process_all_files(input_dict):\n", " for filename in input_dict.keys():\n", " process_file(filename, input_dict)\n", " print(f\"Processed {filename}\")\n", "\n", + "\n", "# Example usage\n", - "input_dict = utils.read_input(path='subset')\n", - "process_all_files(input_dict)\n" + "input_dict = utils.read_input(path=\"subset\")\n", + "process_all_files(input_dict)" ] }, { @@ -362,17 +390,18 @@ "import postprocess\n", "import utils\n", "\n", + "\n", "def clean_td(td):\n", " td_clean = []\n", " for d in td:\n", " new_d = {}\n", " for k, v in d.items():\n", - " if 'DATE' in k:\n", + " if \"DATE\" in k:\n", " new_d[k] = v if isinstance(v, list) else [v]\n", - " elif k not in ['page_num', 'Filename']:\n", - " if isinstance(v, str) and ',' in v:\n", - " new_d[k] = [item.strip() for item in v.split(',')]\n", - " elif v == 'N/A':\n", + " elif k not in [\"page_num\", \"Filename\"]:\n", + " if isinstance(v, str) and \",\" in v:\n", + " new_d[k] = [item.strip() for item in v.split(\",\")]\n", + " elif v == \"N/A\":\n", " new_d[k] = []\n", " else:\n", " new_d[k] = [v] if isinstance(v, str) else v\n", @@ -381,24 +410,29 @@ " td_clean.append(new_d)\n", " return td_clean\n", "\n", + "\n", "def get_unique_keys(dicts):\n", " keys = set()\n", " for d in dicts:\n", " keys.update(d.keys())\n", " return keys\n", "\n", + "\n", "def get_unique_date_fields(td_results):\n", " unique_date_fields = {}\n", " for td in td_results:\n", " for key, value in td.items():\n", - " if 'DATE' in key:\n", + " if \"DATE\" in key:\n", " if key not in unique_date_fields:\n", " unique_date_fields[key] = set()\n", - " unique_date_fields[key].update(value if isinstance(value, list) else [value])\n", + " unique_date_fields[key].update(\n", + " value if isinstance(value, list) else [value]\n", + " )\n", " for key in unique_date_fields:\n", " unique_date_fields[key] = list(unique_date_fields[key])\n", " return unique_date_fields\n", "\n", + "\n", "def GENERATE_PROMPT(bu_dict, td_dicts, page, field_name=None, values=None):\n", " if field_name and values:\n", " return f\"\"\"### PAGE START ### {page} ### PAGE END\n", @@ -429,96 +463,115 @@ "Write ONLY the number of the dictionary that is most associated with the Service and Methodology of interest. Do not return any additional text beyond the digit itself. \n", "\"\"\"\n", "\n", + "\n", "def merge_results(td_results, bu_results, text_dict):\n", " all_keys = get_unique_keys(td_results) | get_unique_keys(bu_results)\n", " print(\"All Keys:\", all_keys)\n", - " \n", + "\n", " date_fields = get_unique_date_fields(td_results)\n", " print(\"Date Fields:\", date_fields)\n", - " \n", + "\n", " merged_results = []\n", " for bu in bu_results:\n", " merged = {key: bu.get(key, \"\") for key in all_keys}\n", - " td_on_page = [td for td in td_results if td['page_num'] == bu['page_num']]\n", - " \n", + " td_on_page = [td for td in td_results if td[\"page_num\"] == bu[\"page_num\"]]\n", + "\n", " for td in td_on_page:\n", " for key, value in td.items():\n", " if not merged[key]:\n", " merged[key] = value\n", - " elif isinstance(value, list) and value and not isinstance(merged[key], list):\n", + " elif (\n", + " isinstance(value, list)\n", + " and value\n", + " and not isinstance(merged[key], list)\n", + " ):\n", " merged[key] = value\n", " elif isinstance(value, list) and value:\n", " merged[key].extend(value)\n", "\n", " for date_key, date_values in date_fields.items():\n", - " if date_key not in merged or not merged[date_key]:\n", - " merged[date_key] = date_values\n", - " else:\n", - " merged[date_key].extend([val for val in date_values if val not in merged[date_key]])\n", + " if date_key not in merged or not merged[date_key]:\n", + " merged[date_key] = date_values\n", + " else:\n", + " merged[date_key].extend(\n", + " [val for val in date_values if val not in merged[date_key]]\n", + " )\n", "\n", " for key in all_keys:\n", " if isinstance(merged[key], list) and len(merged[key]) > 1:\n", - " page_num = bu['page_num']\n", + " page_num = bu[\"page_num\"]\n", " page_text = text_dict.get(page_num, \"\")\n", - " prompt = GENERATE_PROMPT(bu, td_on_page, page_text, field_name=key, values=merged[key])\n", + " prompt = GENERATE_PROMPT(\n", + " bu, td_on_page, page_text, field_name=key, values=merged[key]\n", + " )\n", " response = claude_funcs.invoke_claude_3(prompt, max_tokens=4000)\n", " try:\n", " selected_value = response.strip()\n", " merged[key] = selected_value\n", " except Exception as e:\n", " print(f\"Error processing LLM response for key {key}: {e}\")\n", - " merged[key] = ', '.join(merged[key])\n", - " \n", + " merged[key] = \", \".join(merged[key])\n", + "\n", " merged_results.append(merged)\n", - " \n", + "\n", " return merged_results\n", "\n", + "\n", "def process_file(filename, input_dict):\n", " contract_text = input_dict[filename]\n", "\n", " # Preprocess\n", " contract_text = preprocess.clean_newlines(contract_text)\n", " text_dict = preprocess.split_text(contract_text)\n", - " \n", + "\n", " # Run Top Down\n", - " td_results = prompt_funcs.run_top_down(filename, text_dict) # Returns list of dictionaries for each page\n", + " td_results = prompt_funcs.run_top_down(\n", + " filename, text_dict\n", + " ) # Returns list of dictionaries for each page\n", " print(\"TD Results:\", td_results)\n", - " \n", + "\n", " # Run Bottom Up\n", - " bu_results = prompt_funcs.run_bottom_up(filename, text_dict) # Returns list of dictionaries\n", + " bu_results = prompt_funcs.run_bottom_up(\n", + " filename, text_dict\n", + " ) # Returns list of dictionaries\n", " print(\"BU Results:\", bu_results)\n", - " \n", + "\n", " # Combine\n", " combined_results = merge_results(clean_td(td_results), bu_results, text_dict)\n", " print(\"Combined Results:\", combined_results)\n", - " \n", + "\n", " # Convert to DataFrame\n", " combined_df = pd.DataFrame(combined_results)\n", - " \n", + "\n", " # Post-process combined results\n", " # post_processed_combined_df = postprocess.postprocess_results(combined_df)\n", " # print(\"Post-processed Results:\", post_processed_combined_df)\n", - " \n", + "\n", " # Create directories\n", " base_filename = os.path.splitext(filename)[0]\n", - " output_dir = os.path.join('output', base_filename)\n", + " output_dir = os.path.join(\"output\", base_filename)\n", " os.makedirs(output_dir, exist_ok=True)\n", - " \n", + "\n", " # Save results\n", - " pd.DataFrame(td_results).to_csv(os.path.join(output_dir, 'td_results.csv'), index=False)\n", - " pd.DataFrame(bu_results).to_csv(os.path.join(output_dir, 'bu_results.csv'), index=False)\n", - " combined_df.to_csv(os.path.join(output_dir, 'combined_results.csv'), index=False)\n", + " pd.DataFrame(td_results).to_csv(\n", + " os.path.join(output_dir, \"td_results.csv\"), index=False\n", + " )\n", + " pd.DataFrame(bu_results).to_csv(\n", + " os.path.join(output_dir, \"bu_results.csv\"), index=False\n", + " )\n", + " combined_df.to_csv(os.path.join(output_dir, \"combined_results.csv\"), index=False)\n", " # post_processed_combined_df.to_csv(os.path.join(output_dir, 'combined_results.csv'), index=False)\n", "\n", + "\n", "def process_all_files(input_dict):\n", " for filename in input_dict.keys():\n", " process_file(filename, input_dict)\n", " print(f\"Processed {filename}\")\n", "\n", + "\n", "# Example usage\n", - "input_dict = utils.read_input(path='../data/')\n", - "process_all_files(input_dict)\n", - "\n" + "input_dict = utils.read_input(path=\"../data/\")\n", + "process_all_files(input_dict)" ] }, { @@ -528,7 +581,7 @@ "outputs": [], "source": [ "# Usage example:\n", - "#save_combined_to_csv(combined, 'combined_output.csv')" + "# save_combined_to_csv(combined, 'combined_output.csv')" ] }, { @@ -537,7 +590,7 @@ "metadata": {}, "outputs": [], "source": [ - "utils.consolidate_individual(input_folder='../results/', output_folder='../output/')" + "utils.consolidate_individual(input_folder=\"../results/\", output_folder=\"../output/\")" ] } ], diff --git a/fieldExtraction/src/main.py b/fieldExtraction/src/main.py index 56f534f..f824ad0 100644 --- a/fieldExtraction/src/main.py +++ b/fieldExtraction/src/main.py @@ -1,4 +1,3 @@ - """ main.py Doczy AI - Pricing Before Carveouts: Main Execution Module @@ -21,30 +20,37 @@ import config import claude_funcs import file_processing + def main(): if config.TEST: - print(claude_funcs.invoke_claude_3("Write 'test', nothing more.", max_tokens = 10)) + print( + claude_funcs.invoke_claude_3("Write 'test', nothing more.", max_tokens=10) + ) else: # Read contract txt input_dict = utils.read_input() - + if config.FILTER_ALREADY_PROCESSED: input_dict = utils.filter_already_processed(input_dict) - + def process_item(item): key, value = item - if utils.contains_reimbursement(item[1],0): + if utils.contains_reimbursement(item[1], 0): try: file_processing.process_file(item) - #prompt_funcs.bottom_up(item) + # prompt_funcs.bottom_up(item) except Exception as e: # Print the error message and traceback print(f"Error processing item {key}: {e}") traceback.print_exc() - with concurrent.futures.ThreadPoolExecutor(max_workers=config.MAX_WORKERS) as executor: + with concurrent.futures.ThreadPoolExecutor( + max_workers=config.MAX_WORKERS + ) as executor: # Process each item individually - futures = [executor.submit(process_item, item) for item in input_dict.items()] + futures = [ + executor.submit(process_item, item) for item in input_dict.items() + ] # Wait for all futures to complete for future in concurrent.futures.as_completed(futures): @@ -58,5 +64,6 @@ def main(): # if config.WRITE_OUTPUT: # utils.consolidate_csvs(config.CONSOLIDATED_OUTPUT_DIRECTORY, config.OUTPUT_CSV_PATH) + if __name__ == "__main__": main() diff --git a/fieldExtraction/src/merge_funcs.py b/fieldExtraction/src/merge_funcs.py index 0b5bd75..cf6e2f3 100644 --- a/fieldExtraction/src/merge_funcs.py +++ b/fieldExtraction/src/merge_funcs.py @@ -1,4 +1,3 @@ - import prompts import claude_funcs @@ -11,42 +10,46 @@ def prompt_select_multiple(key, value, bu_dict, page): def merge_results(td_results, bu_results, text_dict): def find_td_dict(page_num): - """ Helper function to find dictionary in td_results with matching page_num """ + """Helper function to find dictionary in td_results with matching page_num""" for td in td_results: - if td['page_num'] == page_num: + if td["page_num"] == page_num: return td return None - + def get_value(td_dict, key, page_num, text_dict): - """ Helper function to handle value selection based on the rules specified. """ + """Helper function to handle value selection based on the rules specified.""" if td_dict is None: # If there's no matching page, use recursion to look one page earlier if int(page_num) > 1: - return get_value(find_td_dict(int(page_num) - 1), key, int(page_num) - 1, text_dict) + return get_value( + find_td_dict(int(page_num) - 1), key, int(page_num) - 1, text_dict + ) else: return "N/A" # Base case if no previous pages exist else: value = td_dict.get(key, "N/A") if value == "N/A" or value == "": # If value is 'N/A' or empty, use recursion to look one page earlier - return get_value(find_td_dict(int(page_num) - 1), key, int(page_num) - 1, text_dict) - elif ',' in value and 'DATE' not in key: + return get_value( + find_td_dict(int(page_num) - 1), key, int(page_num) - 1, text_dict + ) + elif "," in value and "DATE" not in key: # Randomly select one of the comma-separated values return prompt_select_multiple(key, value, bu_dict, text_dict[page_num]) else: return value - + merged_results = [] for bu_dict in bu_results: - page_num = bu_dict['page_num'] + page_num = bu_dict["page_num"] td_dict = find_td_dict(page_num) - + new_dict = bu_dict.copy() # Start with the bu_dict's data # Add or overwrite keys from td_dict if td_dict: for key, value in td_dict.items(): - if key not in ['Filename', 'page_num']: + if key not in ["Filename", "page_num"]: new_dict[key] = get_value(td_dict, key, page_num, text_dict) - + merged_results.append(new_dict) - return merged_results \ No newline at end of file + return merged_results diff --git a/fieldExtraction/src/postprocess.py b/fieldExtraction/src/postprocess.py index bde4f62..d0c2414 100644 --- a/fieldExtraction/src/postprocess.py +++ b/fieldExtraction/src/postprocess.py @@ -4,40 +4,52 @@ import re import postprocessing_funcs import config - + def postprocess_results(combined_df): # Define valid values dictionary valid_values_dict = { - 'CONTRACT_LOB': config.VALID_LOBS, - 'CONTRACT_PROGRAM': config.VALID_PROGRAMS, - 'CONTRACT_NETWORK': config.VALID_NETWORKS + "CONTRACT_LOB": config.VALID_LOBS, + "CONTRACT_PROGRAM": config.VALID_PROGRAMS, + "CONTRACT_NETWORK": config.VALID_NETWORKS, } - + # Sanitize the combined data try: df = postprocessing_funcs.sanitize_combined(combined_df) except Exception as e: print(f"Postprocessing Error - santize_combined : {e}") - + # Clean specific columns for exact matches - for field_list in [['CONTRACT_LOB', 'Corrected_LOB', config.VALID_LOBS], ['CONTRACT_PROGRAM', 'Corrected_PROGRAM', config.VALID_PROGRAMS], ['CONTRACT_NETWORK', 'Corrected_NETWORK', config.VALID_NETWORKS]]: + for field_list in [ + ["CONTRACT_LOB", "Corrected_LOB", config.VALID_LOBS], + ["CONTRACT_PROGRAM", "Corrected_PROGRAM", config.VALID_PROGRAMS], + ["CONTRACT_NETWORK", "Corrected_NETWORK", config.VALID_NETWORKS], + ]: try: - postprocessing_funcs.clean_columns_combined(df, field_list[0], field_list[2], field_list[1]) + postprocessing_funcs.clean_columns_combined( + df, field_list[0], field_list[2], field_list[1] + ) except Exception as e: print(f"Postprocessing Error - clean_columns_combined : {e}") - + try: - postprocessing_funcs.clean_columns_combined_fuzzy(df, field_list[0], field_list[2], config.FUZZY_MATCH_THRESHOLD) + postprocessing_funcs.clean_columns_combined_fuzzy( + df, field_list[0], field_list[2], config.FUZZY_MATCH_THRESHOLD + ) except Exception as e: print(f"Postprocessing Error - clean_columns_combined_fuzzy : {e}") - + # Correct misplaced values across columns try: - df = postprocessing_funcs.correct_misplaced_values(df, ['CONTRACT_LOB', 'CONTRACT_PROGRAM', 'CONTRACT_NETWORK'], valid_values_dict) + df = postprocessing_funcs.correct_misplaced_values( + df, + ["CONTRACT_LOB", "CONTRACT_PROGRAM", "CONTRACT_NETWORK"], + valid_values_dict, + ) except Exception as e: print(f"Postprocessing Error - correct_misplaced_values : {e}") - + # Move percentages and large numbers to correct columns try: df = postprocessing_funcs.move_percentage_to_rate(df) @@ -56,15 +68,17 @@ def postprocess_results(combined_df): except Exception as e: print(f"Postprocessing Error - adjust_reimbursement_rate : {e}") - # -- Deprecated -- + # -- Deprecated -- # Clean dates so that Termination and Effective aren't the same - # try: + # try: # df = postprocessing_funcs.clean_dates(df) # except Exception as e: # print(f"Postprocessing Error - clean_dates : {e}") # Re order columns - column_order = [col for col in config.VALID_COLUMNS if col in df.columns] + [col for col in df.columns if col not in config.VALID_COLUMNS] + column_order = [col for col in config.VALID_COLUMNS if col in df.columns] + [ + col for col in df.columns if col not in config.VALID_COLUMNS + ] df = df[column_order] return df diff --git a/fieldExtraction/src/postprocessing_funcs.py b/fieldExtraction/src/postprocessing_funcs.py index 0f45fe5..b36be29 100644 --- a/fieldExtraction/src/postprocessing_funcs.py +++ b/fieldExtraction/src/postprocessing_funcs.py @@ -1,10 +1,10 @@ - import pandas as pd import re import difflib import config + def sanitize_value(value): """ Sanitizes a given value by handling nulls, lists, and strings to ensure consistency in formatting. @@ -22,15 +22,16 @@ def sanitize_value(value): """ try: if isinstance(value, list): - return ', '.join(str(v) for v in value) + return ", ".join(str(v) for v in value) elif isinstance(value, str): - value = value.strip('[]') - return ', '.join([item.strip(" '") for item in value.split(',')]) + value = value.strip("[]") + return ", ".join([item.strip(" '") for item in value.split(",")]) elif pd.isna(value): return "N/A" except: return value + def exact_match(val, valid_values): """ Compares a given string value against a list of valid values to determine if there is an exact match, case-insensitively. @@ -52,6 +53,7 @@ def exact_match(val, valid_values): return valid_val return None + def clean_columns_combined(df, column_name, valid_values, new_column_name): """ Cleans and standardizes entries in a specified column of a DataFrame based on a list of valid values, creating a new column with standardized values. @@ -75,13 +77,13 @@ def clean_columns_combined(df, column_name, valid_values, new_column_name): def update_column(entry): if pd.notna(entry): - terms = entry.split(',') + terms = entry.split(",") for term in terms: match = exact_match(term, valid_values) if match: return match return None - + df[new_column_name] = df[column_name].apply(update_column) cleaned_values = df[new_column_name].unique() return original_values, cleaned_values, changes @@ -92,7 +94,7 @@ def get_closest_match(val, valid_values, similarity_threshold=0.7): Finds the closest match for a given string from a list of valid values, using a similarity threshold. This function cleans the input value and compares it to each cleaned value in the valid values list using a fuzzy - matching technique. It returns the closest match that meets or exceeds the specified similarity threshold. If no + matching technique. It returns the closest match that meets or exceeds the specified similarity threshold. If no matches meet the threshold, the function returns None. Parameters: @@ -106,7 +108,9 @@ def get_closest_match(val, valid_values, similarity_threshold=0.7): if pd.isna(val): return None val = val.strip().upper() - matches = difflib.get_close_matches(val, [v.upper() for v in valid_values], n=1, cutoff=similarity_threshold) + matches = difflib.get_close_matches( + val, [v.upper() for v in valid_values], n=1, cutoff=similarity_threshold + ) return matches[0] if matches else None @@ -131,17 +135,19 @@ def clean_columns_combined_fuzzy(df, column_name, valid_values, threshold): changes = {} df[column_name] = df[column_name].apply(sanitize_value) original_values = df[column_name].unique() - + def log_and_clean(entry): if pd.notna(entry): - words = entry.split(',') + words = entry.split(",") cleaned_words = [] for word in words: - cleaned_word = get_closest_match(word.strip(), valid_values, similarity_threshold=threshold) + cleaned_word = get_closest_match( + word.strip(), valid_values, similarity_threshold=threshold + ) if cleaned_word and word.strip().upper() != cleaned_word: changes[word.strip()] = cleaned_word cleaned_words.append(cleaned_word if cleaned_word else word.strip()) - return ', '.join(cleaned_words) + return ", ".join(cleaned_words) return None df[column_name] = df[column_name].apply(log_and_clean) @@ -149,7 +155,6 @@ def clean_columns_combined_fuzzy(df, column_name, valid_values, threshold): return original_values, cleaned_values, changes - def correct_misplaced_values(df, columns, valid_values_dict): """ Corrects misplaced values within specified columns of a DataFrame based on a dictionary of valid values for each column. @@ -170,20 +175,32 @@ def correct_misplaced_values(df, columns, valid_values_dict): for index, row in df.iterrows(): for col in columns: if pd.notna(row[col]): - terms = row[col].split(',') + terms = row[col].split(",") for term in terms: term = term.strip() for target_col, valid_values in valid_values_dict.items(): if target_col != col: match = exact_match(term, valid_values) if match: - if pd.isna(row[target_col]) or not row[target_col].strip(): + if ( + pd.isna(row[target_col]) + or not row[target_col].strip() + ): df.at[index, target_col] = match df.at[index, col] = None else: current_value = row[target_col].strip() - if get_closest_match(match, [current_value], config.FUZZY_MATCH_THRESHOLD) is None: - df.at[index, 'Corrected_' + target_col] = f"Found {term} in {col} cell" + if ( + get_closest_match( + match, + [current_value], + config.FUZZY_MATCH_THRESHOLD, + ) + is None + ): + df.at[index, "Corrected_" + target_col] = ( + f"Found {term} in {col} cell" + ) df.at[index, col] = None return df @@ -204,13 +221,28 @@ def filter_service_column(d): list: A new list of dictionaries with items containing specified keywords in the 'SERVICE' key removed. """ keywords = [ - 'LIABILITY', 'RISK', 'LOBBYING', 'DAMAGES', 'CONFIDENTIALITY', 'ARBITRATION', 'FALSE CLAIMS ACT', 'UNSPECIFIED', - 'AUDIT', 'INTEREST', 'N/A', 'BUSINESS', 'COMPLIANCE', 'Medical Assistance Program', 'MATERIAL SUBCONTRACT', 'GIFTS', 'GRATUITIES', + "LIABILITY", + "RISK", + "LOBBYING", + "DAMAGES", + "CONFIDENTIALITY", + "ARBITRATION", + "FALSE CLAIMS ACT", + "UNSPECIFIED", + "AUDIT", + "INTEREST", + "N/A", + "BUSINESS", + "COMPLIANCE", + "Medical Assistance Program", + "MATERIAL SUBCONTRACT", + "GIFTS", + "GRATUITIES", ] - pattern = '|'.join(keywords) + pattern = "|".join(keywords) regex = re.compile(pattern, re.IGNORECASE) - filtered_list = [item for item in d if not regex.search(item.get('SERVICE', ''))] + filtered_list = [item for item in d if not regex.search(item.get("SERVICE", ""))] return filtered_list @@ -229,13 +261,28 @@ def move_percentage_to_rate(df): Returns: pandas.DataFrame: The DataFrame with percentage values moved from the 'REIMBURSEMENT_FLAT_FEE' to the 'REIMBURSEMENT_RATE' column. """ + def move_percentage(value): - if isinstance(value, str) and '%' in value: + if isinstance(value, str) and "%" in value: return True return False - - df['REIMBURSEMENT_RATE'] = df.apply(lambda row: row['REIMBURSEMENT_FLAT_FEE'] if move_percentage(row['REIMBURSEMENT_FLAT_FEE']) else row['REIMBURSEMENT_RATE'], axis=1) - df['REIMBURSEMENT_FLAT_FEE'] = df.apply(lambda row: None if move_percentage(row['REIMBURSEMENT_FLAT_FEE']) else row['REIMBURSEMENT_FLAT_FEE'], axis=1) + + df["REIMBURSEMENT_RATE"] = df.apply( + lambda row: ( + row["REIMBURSEMENT_FLAT_FEE"] + if move_percentage(row["REIMBURSEMENT_FLAT_FEE"]) + else row["REIMBURSEMENT_RATE"] + ), + axis=1, + ) + df["REIMBURSEMENT_FLAT_FEE"] = df.apply( + lambda row: ( + None + if move_percentage(row["REIMBURSEMENT_FLAT_FEE"]) + else row["REIMBURSEMENT_FLAT_FEE"] + ), + axis=1, + ) return df @@ -255,12 +302,15 @@ def set_rate_to_zero_if_not_covered(df): Returns: pandas.DataFrame: The updated DataFrame with adjusted 'REIMURSEMENT_RATE' values where applicable. """ + def check_and_set_rate(row): - if pd.notna(row['FULL_METHODOLOGY']) and re.search(r'not covered', row['FULL_METHODOLOGY'], re.IGNORECASE): + if pd.notna(row["FULL_METHODOLOGY"]) and re.search( + r"not covered", row["FULL_METHODOLOGY"], re.IGNORECASE + ): return 0 - return row['REIMBURSEMENT_RATE'] - - df['REIMBURSEMENT_RATE'] = df.apply(check_and_set_rate, axis=1) + return row["REIMBURSEMENT_RATE"] + + df["REIMBURSEMENT_RATE"] = df.apply(check_and_set_rate, axis=1) return df @@ -279,13 +329,18 @@ def adjust_reimbursement_rate(df): Returns: pandas.DataFrame: The updated DataFrame with recalibrated 'REIMBURSEMENT_RATE' values based on the specified methodology condition. """ + def adjust_rate(row): - if pd.notna(row['FULL_METHODOLOGY']) and re.search(r'case insensitive', row['FULL_METHODOLOGY'], re.IGNORECASE): - if pd.notna(row['REIMBURSEMENT_RATE']) and isinstance(row['REIMBURSEMENT_RATE'], (int, float)): - return 100 - row['REIMBURSEMENT_RATE'] - return row['REIMBURSEMENT_RATE'] - - df['REIMBURSEMENT_RATE'] = df.apply(adjust_rate, axis=1) + if pd.notna(row["FULL_METHODOLOGY"]) and re.search( + r"case insensitive", row["FULL_METHODOLOGY"], re.IGNORECASE + ): + if pd.notna(row["REIMBURSEMENT_RATE"]) and isinstance( + row["REIMBURSEMENT_RATE"], (int, float) + ): + return 100 - row["REIMBURSEMENT_RATE"] + return row["REIMBURSEMENT_RATE"] + + df["REIMBURSEMENT_RATE"] = df.apply(adjust_rate, axis=1) return df @@ -310,12 +365,12 @@ def clean_td(td): for d in td: new_d = {} for k, v in d.items(): - if 'DATE' in k: + if "DATE" in k: new_d[k] = v if isinstance(v, list) else [v] - elif k not in ['page_num', 'Filename']: - if isinstance(v, str) and ',' in v: - new_d[k] = [item.strip() for item in v.split(',')] - elif v == 'N/A': + elif k not in ["page_num", "Filename"]: + if isinstance(v, str) and "," in v: + new_d[k] = [item.strip() for item in v.split(",")] + elif v == "N/A": new_d[k] = [] else: new_d[k] = [v] if isinstance(v, str) else v @@ -325,7 +380,6 @@ def clean_td(td): return td_clean - def sanitize_combined(df): """ Sanitizes all columns in a DataFrame by applying a predefined sanitization function to each value. @@ -345,6 +399,7 @@ def sanitize_combined(df): df[column] = df[column].apply(sanitize_value) return df + def extract_codes_CPT(value): """ Extracts and formats CPT codes from a given input value. @@ -364,16 +419,17 @@ def extract_codes_CPT(value): """ if pd.isna(value): return value - - value = re.sub(r'\bthrough\b', '-', str(value), flags=re.IGNORECASE) - - pattern = r'\b([A-Z]\d{4}|\d{5})(-[A-Z]{2})?\b' + + value = re.sub(r"\bthrough\b", "-", str(value), flags=re.IGNORECASE) + + pattern = r"\b([A-Z]\d{4}|\d{5})(-[A-Z]{2})?\b" matches = re.findall(pattern, value) if matches: - return ', '.join([''.join(match) for match in matches]) + return ", ".join(["".join(match) for match in matches]) else: return "" + def extract_codes_Diagnosis(value): """ Extracts and formats diagnosis codes from a given input value, typically adhering to ICD (International Classification of Diseases) formats. @@ -392,14 +448,15 @@ def extract_codes_Diagnosis(value): """ if pd.isna(value): return value - - pattern = r'\b[A-Z][0-9][A-Z0-9]{1,4}(\.[A-Z0-9]{1,4})?\b' + + pattern = r"\b[A-Z][0-9][A-Z0-9]{1,4}(\.[A-Z0-9]{1,4})?\b" matches = re.findall(pattern, value) if matches: - return ', '.join(matches) + return ", ".join(matches) else: return "" + def extract_codes_Revenue(value): """ Extracts revenue codes from a given input value. Revenue codes are typically numerical codes of three to four digits. @@ -417,14 +474,15 @@ def extract_codes_Revenue(value): """ if pd.isna(value): return value - - pattern = r'\b\d{3,4}\b' + + pattern = r"\b\d{3,4}\b" matches = re.findall(pattern, value) if matches: - return ', '.join(matches) + return ", ".join(matches) else: return "" + def clean_code_column(df, column_name, extract_function): """ Applies a specified function to clean and extract codes from a specific column in a DataFrame. @@ -443,13 +501,15 @@ def clean_code_column(df, column_name, extract_function): pandas.DataFrame: The DataFrame with the specified column updated with cleaned and processed codes. """ df[column_name] = df[column_name].apply(extract_function) - return df + return df def clean_dates(df): for idx, row in df.iterrows(): - if row['LOB_PRICING_TERMS_EFFECTIVE_DATE'] == row['LOB_PRICING_TERMS_TERMINATION_DATE']: - df.loc[idx, 'LOB_PRICING_TERMS_TERMINATION_DATE'] = 'N/A' + if ( + row["LOB_PRICING_TERMS_EFFECTIVE_DATE"] + == row["LOB_PRICING_TERMS_TERMINATION_DATE"] + ): + df.loc[idx, "LOB_PRICING_TERMS_TERMINATION_DATE"] = "N/A" return df - diff --git a/fieldExtraction/src/preprocess.py b/fieldExtraction/src/preprocess.py index d2bb5e1..f4cab0c 100644 --- a/fieldExtraction/src/preprocess.py +++ b/fieldExtraction/src/preprocess.py @@ -1,7 +1,7 @@ - from itertools import groupby, count import re + def clean_newlines(contract_text): """ Cleans up isolated newlines in a contract text, converting them into spaces to ensure text continuity. @@ -17,7 +17,7 @@ def clean_newlines(contract_text): Returns: str: The cleaned contract text with isolated newlines replaced by spaces. """ - cleaned_text = re.sub(r'(?>>{word}<<<' + if "%" in word or "$" in word: + word = f">>>{word}<<<" highlighted_words.append(word) - text_dict[page] = ' '.join(highlighted_words) + text_dict[page] = " ".join(highlighted_words) return text_dict @@ -69,8 +69,8 @@ def chunk_text(text_dict): """ Creates text chunks from a dictionary of page texts, focusing on pages with special characters (percentages and dollar amounts). - This function identifies pages that contain '%' or '$' signs and includes those pages along with their immediate neighbors - (previous and next pages) to form chunks. The chunks are then grouped and concatenated into single text blocks for easier + This function identifies pages that contain '%' or '$' signs and includes those pages along with their immediate neighbors + (previous and next pages) to form chunks. The chunks are then grouped and concatenated into single text blocks for easier processing. Each chunk is stored in a new dictionary where the keys represent the range of pages included in the chunk. Parameters: @@ -79,19 +79,28 @@ def chunk_text(text_dict): Returns: dict: A new dictionary where each key is a string representing the range of pages in a chunk, and each value is the concatenated text of those pages. """ - special_pages = {int(page): text for page, text in text_dict.items() if (('%' in text) or ('$' in text) or ('percent' in text))} + special_pages = { + int(page): text + for page, text in text_dict.items() + if (("%" in text) or ("$" in text) or ("percent" in text)) + } page_numbers = sorted(special_pages.keys()) chunk_page_numbers = [] for page in page_numbers: - chunk_page_numbers.append(page-1) + chunk_page_numbers.append(page - 1) chunk_page_numbers.append(page) - chunk_page_numbers.append(page+1) + chunk_page_numbers.append(page + 1) chunk_page_numbers = list(set(chunk_page_numbers)) - chunk_page_numbers = [page for page in chunk_page_numbers if str(page) in text_dict.keys()] - chunks = [list(group) for key, group in groupby(chunk_page_numbers, lambda x, c=count(): x - next(c))] + chunk_page_numbers = [ + page for page in chunk_page_numbers if str(page) in text_dict.keys() + ] + chunks = [ + list(group) + for key, group in groupby(chunk_page_numbers, lambda x, c=count(): x - next(c)) + ] chunk_dict = {} for item in chunks: - dict_key = f'{min(item)}-{max(item)}' + dict_key = f"{min(item)}-{max(item)}" text = "".join([text_dict[str(page)] for page in item]) chunk_dict[dict_key] = text return chunk_dict @@ -111,7 +120,7 @@ def clean_billed_charges(contract_text): Returns: str: The cleaned text with appropriate replacements made for 'billed charges'. """ - + substrings = [ "Physician's Billed Charges", "Provider's Billed Charges", @@ -136,17 +145,21 @@ def clean_billed_charges(contract_text): indices.append(index) index = lower_contract_text.find(lower_s, index + 1) return indices - + for s in substrings: indices = find_substring_indices(contract_text, s) index_adder = 0 if indices: for i in indices: - index = i+index_adder + index = i + index_adder end_index = index + max_substring_length match_part = contract_text[index:end_index] - previous = contract_text[max(0, index-30):index] - if '%' not in previous: - contract_text = contract_text[0:index] + f" 100% of {match_part}" + contract_text[end_index:] + previous = contract_text[max(0, index - 30) : index] + if "%" not in previous: + contract_text = ( + contract_text[0:index] + + f" 100% of {match_part}" + + contract_text[end_index:] + ) index_adder += 9 - return contract_text.replace(' ', ' ') + return contract_text.replace(" ", " ") diff --git a/fieldExtraction/src/prompts.py b/fieldExtraction/src/prompts.py index 92d138e..44fb69d 100644 --- a/fieldExtraction/src/prompts.py +++ b/fieldExtraction/src/prompts.py @@ -1,5 +1,3 @@ - - def GENERATE_PROMPT(bu_dict, td_dicts, page, field_name=None, values=None): if field_name and values: return f"""### PAGE START ### {page} ### PAGE END @@ -105,6 +103,7 @@ If both LESSER_OF_LANGUAGE_IND and GREATER_OF_LANGUAGE_IND are 'N', return a lis Only return the list of dictionaries, with no other commentary or explanation. Ensure you abide by proper JSON formatting. """ + def BOTTOM_UP_EXCEPT_ESC(d, page): return f"""### PAGE START ### {page} ### PAGE END @@ -124,6 +123,7 @@ RATE_ESCALATOR_YEARLY_PERCENT_INCREASE: What is the percent value of increase or Only return the dictionary, with no other commentary or explanation. Ensure you abide by proper JSON formatting. """ + def BOTTOM_UP_FS(d): return f"""Analyze the reimbursement terms listed below: Service: {d['SERVICE']} @@ -138,6 +138,7 @@ REIMBURSEMENT_FEE_SCHEDULE_VERSION : If the Methodology is based on a Fee Schedu Only return the dictionary, with no other commentary or explanation. Ensure you abide by proper JSON formatting. """ + def BOTTOM_UP_METHODOLOGY(d): return f"""Analyze the full methodology text listed below: {d['FULL_METHODOLOGY']} @@ -252,37 +253,212 @@ ONLY return the answer, with no other commentary or explanation. def get_carveout_list(): - return ['Emergency Department/Emergency Room', 'Emergency Department', 'Emergency Room', 'Observation', 'Surgery', 'General Surgery', -'Ambulatory Surgery', 'Intensive Care', 'Trauma', 'HIV', 'Human Immunodeficiency Virus', 'Major Joint Replacement', 'Transplant', 'OBGYN', -'Obstetrician/Gynecologist', 'Obstetrician', 'Gynecologist', 'Opthalmology & Vision', 'Opthalmology', 'Vision', 'Never Events', -'Medically Unnecessary', 'Medically Unnecessary Procedures', 'Not Medically Necessary', 'Not Medically Necessary Procedures', -'Experimental', 'Investigational', 'Unlisted Codes', 'Durable Medical Equipment', 'DME', 'Prosthetics & Orthotics', 'Prosthetics', -'Orthotics', 'Implants', 'Hearing Aids', 'Hearing', 'Anesthesia', 'Anesthesiology', 'Medical Pharmacy', 'Physician Administered Drugs', -'Global', 'Bundled or Unbundled Codes', 'Bundled Codes', 'Unbundled Codes', 'Multiple Procedure Reductions', 'Second Surgery', 'Subsequent Surgeries', -'Non-Behavioral Health Mid-Level Professionals', 'Physician Assistant', 'PA', 'Nurse Practicioner', 'NP', 'Non-Physician', -'Non-Physician Health Professionals', 'Technical Component', 'Professional Component', 'Laboratory', 'Pathology', 'Lab', 'Path', 'Lab/Path', -'Laboratory/Pathology', 'Radiology', 'Imaging', 'Radiology/Imaging', 'Mammography', 'Diagnostic', 'Pre-Admission Procedures', -'Post-Discharge Procedures', 'Readmission', 'Status Indicators', 'Stop Loss', 'Emergency Medical', 'EMS', 'NICU', 'Neonatal Intensive Care Unit', -'Vaccine for Children', 'VFC', 'Surgical Assistant', 'Physicians/Clinical Psychologists', 'Doctor of Nursing Practice', 'Osteopathic Medicine', -'Clinical Psychology', 'Audiologist', 'Chiropractors', 'Registered Dietician', 'AUD', 'DC', 'RD', 'Board Certified Behavioral Analysis', 'Behavioral Analysis', -'BCBA', 'Independent Licensures', 'Licensed Professional Counselor', 'Marriage and Family Therapist', 'Substance Abuse Counselor', 'Clinical Social Worker', -'Behavioral Health Outpatient Clinic', 'Physical Therapist', 'Occupational Therapist', 'Speech Therapist', 'PT', 'OT', 'ST', 'Transportation', -'Primary Care', 'Primary Care Behavioral Health', 'Behavioral Physician', 'Clinical Psychologist', 'Mid-Level Practicioner', 'Dentist', 'Dental', -'Supplies and Devices', 'Supplies', 'Devices', 'Immunizations', 'Obstetrical Epidural', 'Pediatric Subspecialties', 'Orthopedic Surgery', 'All Other Specialists', -'Specialty Care Physician', 'Pediatric Primary Care', 'Nurse Anesthetist', 'Early Periodic Screening, Diagnostic, and Treatment', 'EPSDT', -'Ancillary', 'Oncology', 'Cancer', 'Inpatient Physical Rehabilitation', 'Inpatient Rehabilitation', 'Outpatient Rehabilitation', 'Rehabilitation', 'Non-Behavioral Health Rehabilitation', -'Extracorporeal Shock Wave Lithotripsy', 'Cardiac', 'Special Care Unit', 'Skilled Nursing', 'Infusion', 'Specialty Care', -'Specialty Care Physician', 'Organ Acquisition', 'Blood Products', 'Blood Products Outpatient', 'Blood Products Inpatient', 'High Cost Drugs', 'Sleep Studies', -'NICU', 'Newborn Intensive Care Unit', 'Extracorporeal Membrane Oxygenation', 'Burns', 'Kyphoplasty', 'Cryosurgical Ablation of the Prostate', -'Transurethral Thermal Ablation', 'TUMT', 'Transurethral Needle Ablation', 'TUNA', 'Hyperbaric Treatment', 'Clinic Visit', 'Boarder Baby', -'Pediatric Intensive Care Unit', 'PICU', 'Psychiatric', 'Mental Health', 'Behavioral Health', 'Substance Abuse', -'Behavioral Health and Substance Abuse', 'Sub-Acute Facility Care', 'Unrouped Inpatient', 'All Other Acute', -'Neurology', 'Neurology Subspecialties', 'Automatic Implantable Cardioverter Defibrillator', 'Percutaneous Transluminal Coronary Angioplasty', -'Non-Coronary Angioplasty', 'Cardiac Catheters', 'Cardiovascular Surgery', 'Cardiac Surgery', 'Cesarean Birth', 'Cesarean Section', 'C-Section', -'Gamma-Knife Radio-Surgery Outpatient', 'DaVinci Robotic Assisted Surgery', 'Outpatient Electrophysiology with Ablation', 'Outpatient Electrophysiology', -'Magnetic Resonance Image', 'MRI', 'Computed Tomography Scan', 'CT Scan', 'Radiation Therapy', 'Dialysis', 'Gastric Bypass', 'Lap Band', -'Obesity', 'Laparoscopic Cholecystectomy', 'Lap Chole', 'Laparoscopic Hysterectomy', 'Laparoscopic Prostatectomy', 'Laparoscopic Hysteroscopy', -'Treatment Room', 'Wound Care', 'Cardiac Computed Tomograpy & Angiography', 'Positron Emission Tomography Scan', 'PET Scan', 'Hematology'] + return [ + "Emergency Department/Emergency Room", + "Emergency Department", + "Emergency Room", + "Observation", + "Surgery", + "General Surgery", + "Ambulatory Surgery", + "Intensive Care", + "Trauma", + "HIV", + "Human Immunodeficiency Virus", + "Major Joint Replacement", + "Transplant", + "OBGYN", + "Obstetrician/Gynecologist", + "Obstetrician", + "Gynecologist", + "Opthalmology & Vision", + "Opthalmology", + "Vision", + "Never Events", + "Medically Unnecessary", + "Medically Unnecessary Procedures", + "Not Medically Necessary", + "Not Medically Necessary Procedures", + "Experimental", + "Investigational", + "Unlisted Codes", + "Durable Medical Equipment", + "DME", + "Prosthetics & Orthotics", + "Prosthetics", + "Orthotics", + "Implants", + "Hearing Aids", + "Hearing", + "Anesthesia", + "Anesthesiology", + "Medical Pharmacy", + "Physician Administered Drugs", + "Global", + "Bundled or Unbundled Codes", + "Bundled Codes", + "Unbundled Codes", + "Multiple Procedure Reductions", + "Second Surgery", + "Subsequent Surgeries", + "Non-Behavioral Health Mid-Level Professionals", + "Physician Assistant", + "PA", + "Nurse Practicioner", + "NP", + "Non-Physician", + "Non-Physician Health Professionals", + "Technical Component", + "Professional Component", + "Laboratory", + "Pathology", + "Lab", + "Path", + "Lab/Path", + "Laboratory/Pathology", + "Radiology", + "Imaging", + "Radiology/Imaging", + "Mammography", + "Diagnostic", + "Pre-Admission Procedures", + "Post-Discharge Procedures", + "Readmission", + "Status Indicators", + "Stop Loss", + "Emergency Medical", + "EMS", + "NICU", + "Neonatal Intensive Care Unit", + "Vaccine for Children", + "VFC", + "Surgical Assistant", + "Physicians/Clinical Psychologists", + "Doctor of Nursing Practice", + "Osteopathic Medicine", + "Clinical Psychology", + "Audiologist", + "Chiropractors", + "Registered Dietician", + "AUD", + "DC", + "RD", + "Board Certified Behavioral Analysis", + "Behavioral Analysis", + "BCBA", + "Independent Licensures", + "Licensed Professional Counselor", + "Marriage and Family Therapist", + "Substance Abuse Counselor", + "Clinical Social Worker", + "Behavioral Health Outpatient Clinic", + "Physical Therapist", + "Occupational Therapist", + "Speech Therapist", + "PT", + "OT", + "ST", + "Transportation", + "Primary Care", + "Primary Care Behavioral Health", + "Behavioral Physician", + "Clinical Psychologist", + "Mid-Level Practicioner", + "Dentist", + "Dental", + "Supplies and Devices", + "Supplies", + "Devices", + "Immunizations", + "Obstetrical Epidural", + "Pediatric Subspecialties", + "Orthopedic Surgery", + "All Other Specialists", + "Specialty Care Physician", + "Pediatric Primary Care", + "Nurse Anesthetist", + "Early Periodic Screening, Diagnostic, and Treatment", + "EPSDT", + "Ancillary", + "Oncology", + "Cancer", + "Inpatient Physical Rehabilitation", + "Inpatient Rehabilitation", + "Outpatient Rehabilitation", + "Rehabilitation", + "Non-Behavioral Health Rehabilitation", + "Extracorporeal Shock Wave Lithotripsy", + "Cardiac", + "Special Care Unit", + "Skilled Nursing", + "Infusion", + "Specialty Care", + "Specialty Care Physician", + "Organ Acquisition", + "Blood Products", + "Blood Products Outpatient", + "Blood Products Inpatient", + "High Cost Drugs", + "Sleep Studies", + "NICU", + "Newborn Intensive Care Unit", + "Extracorporeal Membrane Oxygenation", + "Burns", + "Kyphoplasty", + "Cryosurgical Ablation of the Prostate", + "Transurethral Thermal Ablation", + "TUMT", + "Transurethral Needle Ablation", + "TUNA", + "Hyperbaric Treatment", + "Clinic Visit", + "Boarder Baby", + "Pediatric Intensive Care Unit", + "PICU", + "Psychiatric", + "Mental Health", + "Behavioral Health", + "Substance Abuse", + "Behavioral Health and Substance Abuse", + "Sub-Acute Facility Care", + "Unrouped Inpatient", + "All Other Acute", + "Neurology", + "Neurology Subspecialties", + "Automatic Implantable Cardioverter Defibrillator", + "Percutaneous Transluminal Coronary Angioplasty", + "Non-Coronary Angioplasty", + "Cardiac Catheters", + "Cardiovascular Surgery", + "Cardiac Surgery", + "Cesarean Birth", + "Cesarean Section", + "C-Section", + "Gamma-Knife Radio-Surgery Outpatient", + "DaVinci Robotic Assisted Surgery", + "Outpatient Electrophysiology with Ablation", + "Outpatient Electrophysiology", + "Magnetic Resonance Image", + "MRI", + "Computed Tomography Scan", + "CT Scan", + "Radiation Therapy", + "Dialysis", + "Gastric Bypass", + "Lap Band", + "Obesity", + "Laparoscopic Cholecystectomy", + "Lap Chole", + "Laparoscopic Hysterectomy", + "Laparoscopic Prostatectomy", + "Laparoscopic Hysteroscopy", + "Treatment Room", + "Wound Care", + "Cardiac Computed Tomograpy & Angiography", + "Positron Emission Tomography Scan", + "PET Scan", + "Hematology", + ] def prompt_contract_lob(page): @@ -319,8 +495,9 @@ def prompt_product(page): Extract and list all identified plan or product names. If multiple names are found, return them separated by commas. If no specific plan name is found, return 'N/A'. Only return the extracted information and no other text or sentences or explanations. Only the final values. """ + def prompt_network_name(page): - #MODIFY TO DICTS FOR MAPPING + # MODIFY TO DICTS FOR MAPPING return f"""### PAGE START ### {page} ### PAGE END Extract all network names mentioned on the page. A 'network name' in the context of health insurance refers to the type of managed care organization involved in the delivery of healthcare services. These names often indicate the structure or model of care delivery, which may define the relationships between insurers, healthcare providers, and insured individuals. @@ -340,6 +517,7 @@ def prompt_service_area(page): Extract and return all found attributes separated by commas in a single string. If no specific service areas are mentioned in the text, return 'N/A'. only return the extracted information and no other text or sentences or explanations. Only the final values. """ + def prompt_effective_date(page): return f"""### PAGE START ### {page} ### PAGE END @@ -349,6 +527,7 @@ def prompt_effective_date(page): DO NOT REPEAT THE SAME DATE MULTIPLE TIMES. """ + def prompt_termination_date(page): return f"""### PAGE START ### {page} ### PAGE END @@ -360,51 +539,65 @@ def prompt_termination_date(page): DO NOT REPEAT THE SAME DATE MULTIPLE TIMES. """ + def prompt_metal_level(page): return f"""### PAGE START ### {page} ### PAGE END If the Line of Business is identified as Marketplace, extract all instances of metal levels such as Platinum, Gold, Silver, or Bronze. List each metal level. If no metal levels are found or if the LOB is not Marketplace, return 'N/A'. only return the extracted information and no other text or sentences or explanations. Only the final values seperated by commas.""" + def prompt_claim_discount_ind(lob): return f"For this line of business: {lob}, return with a 'Y' or 'NO' only whether the claim for this LOB is discounted. Only return as 'Y' or 'NO'." + def prompt_claim_discount_percent_rate(lob): return f"For this LOB only: {lob}, extract the claim discount percent rate." + def prompt_claim_discount_start_date(lob): return f"For this LOB only: {lob}, extract the claim discount start date." + def prompt_claim_discount_termination_date(lob): return f"For this LOB only: {lob}, extract the claim discount termination date." + def prompt_premium_ind(lob): return f"""For this Line of Business{lob}, indicate with a 'Y' or 'N' only whether the premium for this LOB is subject to any adjustments. Return only 'Y' or 'N'.""" + def prompt_premium_percent(lob): return f"""For this Line of Business{lob}, extract the percentage rate of premium adjustment if applicable. Provide the percentage as a numerical value only.""" + def prompt_premium_start_date(lob): return f"""For this Line of Business{lob}, identify the start date of the premium adjustments. Return the date in MM/DD/YYYY format only.""" + def prompt_premium_termination_date(): return """For this Line of Business, determine the termination date of the premium adjustments. Provide the date in MM/DD/YYYY format only.""" + def prompt_sequestration_ind(): return """For this Line of Business, indicate with a 'Y' or 'N' whether there is any sequestration applied. Return only 'Y' or 'N'.""" + def prompt_sequestration_rate(): return """For this Line of Business, extract the rate of sequestration as a percentage. Return the rate as a numerical value only.""" + def prompt_sequestration_start_date(): return """For this Line of Business, determine the start date for sequestration. Return the date in MM/DD/YYYY format only.""" + def prompt_sequestration_termination_date(): return """For this Line of Business, identify the termination date for sequestration. Provide the date in MM/DD/YYYY format only.""" + def prompt_penalties_ind(): return """For this Line of Business, indicate with a 'Y' or 'N' only whether there are any penalties applied. Return only 'Y' or 'N'.""" + def prompt_penalties_rate(): return """For this Line of Business, extract the rate of any penalties applied as a percentage. Provide the rate as a numerical value only.""" - diff --git a/fieldExtraction/src/table_funcs.py b/fieldExtraction/src/table_funcs.py index 411873f..af762f0 100644 --- a/fieldExtraction/src/table_funcs.py +++ b/fieldExtraction/src/table_funcs.py @@ -1,20 +1,22 @@ - import re import ast + def clean_tables(text): def replace_colon(text): # Define a function to use in re.sub to check the context of the match def replacer(match): # Extract the character after ':' to check if it's '[' following_text = match.group(1) - if following_text.strip().startswith('['): - return match.group(0) # Return the original match (':' and whatever follows) + if following_text.strip().startswith("["): + return match.group( + 0 + ) # Return the original match (':' and whatever follows) else: - return '' + match.group(1) # Replace ':' with ';' and return the rest + return "" + match.group(1) # Replace ':' with ';' and return the rest # Use a regular expression to find ':' and the text that follows - pattern = r':(\s*.)' + pattern = r":(\s*.)" replaced_text = re.sub(pattern, replacer, text) return replaced_text @@ -25,20 +27,22 @@ def clean_tables(text): if "'" not in content and (len(content) > 0): return content # Return just the content without brackets else: - return match.group(0) # Return the original match if it contains a quote - + return match.group( + 0 + ) # Return the original match if it contains a quote # Use a regular expression to find brackets and the text within - pattern = r'\[([^\[\]]*?)\]' + pattern = r"\[([^\[\]]*?)\]" cleaned_text = re.sub(pattern, replacer, text) return cleaned_text - + # Clean tables - text = text.strip('{}') # Remove brackets + text = text.strip("{}") # Remove brackets text = replace_colon(text) text = remove_unquoted_brackets(text) return text + def convert_to_dict(table_text): """ Converts a formatted table text into a dictionary where each key maps to a list of values. @@ -54,22 +58,22 @@ def convert_to_dict(table_text): """ table_text = clean_tables(table_text) # print(table_text, '\n') - table_text_list = table_text.split(']') + table_text_list = table_text.split("]") table_text_list = [item for item in table_text_list if len(item) > 0] final_dict = {} num_elements = 0 for key_value_text in table_text_list: - key = key_value_text.split(':')[0].strip(', \'') + key = key_value_text.split(":")[0].strip(", '") try: - value = ':'.join(key_value_text.split(':')[1:]).strip()+']' + value = ":".join(key_value_text.split(":")[1:]).strip() + "]" # print(value, '\n') if len(value) > 1: value_list = ast.literal_eval(value) final_dict[key] = value_list num_elements = len(value_list) except: - value = ':'.join(key_value_text.split(':')[1:]).strip()+"']" + value = ":".join(key_value_text.split(":")[1:]).strip() + "']" # print(value, '\n') if len(value) > 1: value_list = ast.literal_eval(value) @@ -96,13 +100,13 @@ def format_table(table_json, table_size): """ table_text = "" keys = [key for key in table_json.keys()] - table_text += ' : '.join(keys) + "\n" + table_text += " : ".join(keys) + "\n" for i in range(table_size): for key in keys: if table_json[key][i]: - #table_text += key + ': ' + table_json[key][i] + ', ' - table_text += table_json[key][i] + ' : ' - table_text += r'\n' + # table_text += key + ': ' + table_json[key][i] + ', ' + table_text += table_json[key][i] + " : " + table_text += r"\n" return table_text @@ -123,17 +127,23 @@ def align_and_format_tables(text_dict): aligned_text_dict = {} for key, text in text_dict.items(): aligned_text = text - if 'Table Start' in text: - table_texts = re.findall(r'-------Table Start--------(.*?)-------Table End--------', text, re.DOTALL) + if "Table Start" in text: + table_texts = re.findall( + r"-------Table Start--------(.*?)-------Table End--------", + text, + re.DOTALL, + ) for table_text in table_texts: # print(table_text, '\n') - + try: # Extract pretable text - pretable = table_text.split('{')[0].strip() + pretable = table_text.split("{")[0].strip() # Extract and format table text - #table_only = "{" + table_text.split('{', 1)[1].rsplit('}', 1)[0].replace("'", '"') + "}" - table_only = "{" + table_text.split('{', 1)[1].rsplit('}', 1)[0] + "}" + # table_only = "{" + table_text.split('{', 1)[1].rsplit('}', 1)[0].replace("'", '"') + "}" + table_only = ( + "{" + table_text.split("{", 1)[1].rsplit("}", 1)[0] + "}" + ) table, table_size = convert_to_dict(table_only) table_formatted = format_table(table, table_size) @@ -141,16 +151,22 @@ def align_and_format_tables(text_dict): table_formatted = table_text # Align - if text.count(pretable) == 2: # One table + if text.count(pretable) == 2: # One table # Remove original table aligned_text = text.replace(table_text, "") # Align new table - aligned_text = aligned_text.replace(pretable, pretable + table_formatted) - + aligned_text = aligned_text.replace( + pretable, pretable + table_formatted + ) + else: - aligned_text = aligned_text.replace(table_text, pretable + ' ' + table_formatted) - aligned_text_dict[key] = aligned_text.replace("-------Table Start--------", "").replace("-------Table End--------", "") + aligned_text = aligned_text.replace( + table_text, pretable + " " + table_formatted + ) + aligned_text_dict[key] = aligned_text.replace( + "-------Table Start--------", "" + ).replace("-------Table End--------", "") else: aligned_text_dict[key] = text - - return aligned_text_dict \ No newline at end of file + + return aligned_text_dict diff --git a/fieldExtraction/src/test.py b/fieldExtraction/src/test.py index e92ec15..58128c6 100644 --- a/fieldExtraction/src/test.py +++ b/fieldExtraction/src/test.py @@ -1,4 +1,3 @@ - import utils import config @@ -7,46 +6,43 @@ import pandas as pd import boto3 +s3 = boto3.Session( + aws_access_key_id="ASIA6GBMBVWOJOOHVXUE", + aws_secret_access_key="YsXnoYhY599uOifnQ+hZnqQJpWTyjhZ8XppJr43V", + aws_session_token="IQoJb3JpZ2luX2VjEHcaCXVzLWVhc3QtMiJHMEUCIBdFMYRAL20l87z9xZJzlvOW8H7+jOIoI2f6zOkNtMdQAiEAomTpoulfRpLcPhzbNFQVR5atGD8xbM49/gKX0rNDZYgqgQMIYRAAGgw5NzUwNDk5NjA4NjAiDGoxXZfA4RDhWkpflCreAmgKA3i7fZYV6z/e5WejrcnSfffciNhUKpen98pi1gJT6xgB9Cy4k19DAkF8ab8dQcyTO80K20hd2QjzwaTXDDiIpXYJq1TRHiJVN+g0QKOx10l4KOo7g8hGPzD/QPmA80aTxC96PqRklnJkUGNKe9zfZ7nCujts/i/rs0oAeuRCCnEKsh3lOtj4YEhoQGsNlKUsx4Pfh2cTn/JTZ4hma0zO8HnVfPh2f4i4hpa8Ula/0arXZJrkJyHdbQV85w+lmjypILQwAy6kQUd57lURDXDF1uPtziKZ2WqkozpEdeblmUUAT24rXKBhv3m86oqN59pb1IiFoDE7IOmJgRqNbh58OhOwifijMRWYouvSFM8pcBDTD725UmTzNu5GKgGJE00Q4PyBu5K2YXX1rHFmOhr29wq11mPMBPVRTF/zMLjG8ZloHYQWtt8VZwso8WE7ezV7RyScz0SBwBvhOX0mMPSwqbUGOqYBhMdcVuF9yw0Norg5U7G8SDl2WmiYCh/Anfea87h/1KzPs6ZNphtuSLcaH9+C7hVx5DVAJUW7gI0xx5jhPgqHbcptDPkNWCL129URUMPFOmuRiPxyTl0Xl5jSeh0Mj+RJz81OCWaEQyTQYOTtNUir77f1KAm8y+ClFrkYFrk6uz6HyiYknmcxAokjdNMIOS/83O7GTUv1Vo6/fIXo2MKX7gxhuNUoNw==", +).client("s3") -s3 = boto3.Session(aws_access_key_id='ASIA6GBMBVWOJOOHVXUE', -aws_secret_access_key='YsXnoYhY599uOifnQ+hZnqQJpWTyjhZ8XppJr43V', -aws_session_token='IQoJb3JpZ2luX2VjEHcaCXVzLWVhc3QtMiJHMEUCIBdFMYRAL20l87z9xZJzlvOW8H7+jOIoI2f6zOkNtMdQAiEAomTpoulfRpLcPhzbNFQVR5atGD8xbM49/gKX0rNDZYgqgQMIYRAAGgw5NzUwNDk5NjA4NjAiDGoxXZfA4RDhWkpflCreAmgKA3i7fZYV6z/e5WejrcnSfffciNhUKpen98pi1gJT6xgB9Cy4k19DAkF8ab8dQcyTO80K20hd2QjzwaTXDDiIpXYJq1TRHiJVN+g0QKOx10l4KOo7g8hGPzD/QPmA80aTxC96PqRklnJkUGNKe9zfZ7nCujts/i/rs0oAeuRCCnEKsh3lOtj4YEhoQGsNlKUsx4Pfh2cTn/JTZ4hma0zO8HnVfPh2f4i4hpa8Ula/0arXZJrkJyHdbQV85w+lmjypILQwAy6kQUd57lURDXDF1uPtziKZ2WqkozpEdeblmUUAT24rXKBhv3m86oqN59pb1IiFoDE7IOmJgRqNbh58OhOwifijMRWYouvSFM8pcBDTD725UmTzNu5GKgGJE00Q4PyBu5K2YXX1rHFmOhr29wq11mPMBPVRTF/zMLjG8ZloHYQWtt8VZwso8WE7ezV7RyScz0SBwBvhOX0mMPSwqbUGOqYBhMdcVuF9yw0Norg5U7G8SDl2WmiYCh/Anfea87h/1KzPs6ZNphtuSLcaH9+C7hVx5DVAJUW7gI0xx5jhPgqHbcptDPkNWCL129URUMPFOmuRiPxyTl0Xl5jSeh0Mj+RJz81OCWaEQyTQYOTtNUir77f1KAm8y+ClFrkYFrk6uz6HyiYknmcxAokjdNMIOS/83O7GTUv1Vo6/fIXo2MKX7gxhuNUoNw==' -).client('s3') - def get_s3_files(bucket_name, folder): - paginator = s3.get_paginator('list_objects_v2') + paginator = s3.get_paginator("list_objects_v2") result = paginator.paginate(Bucket=bucket_name, Prefix=folder) files = set() for page in result: - if 'Contents' in page: - for obj in page['Contents']: - file_name = obj['Key'] + if "Contents" in page: + for obj in page["Contents"]: + file_name = obj["Key"] if file_name: # Ensure it's not a folder files.add(file_name) return files -bucket_name = 'texas-children-files' -folder = '510_text_files/' + +bucket_name = "texas-children-files" +folder = "510_text_files/" files_in_folder = get_s3_files(bucket_name, folder) -source_bucket='texas-children-files' +source_bucket = "texas-children-files" for source_key in files_in_folder: - filename = source_key.split('510_text_files/')[1] + filename = source_key.split("510_text_files/")[1] print(filename) - if filename not in os.listdir('data/texas_childrens'): + if filename not in os.listdir("data/texas_childrens"): try: - response = s3.get_object(Bucket=source_bucket, Key=source_key)['Body'] - file_content = response.read().decode('utf-8') - - with open(f'data/texas_childrens/{filename}', 'w') as file: + response = s3.get_object(Bucket=source_bucket, Key=source_key)["Body"] + file_content = response.read().decode("utf-8") + + with open(f"data/texas_childrens/{filename}", "w") as file: file.write(file_content) except: pass - - - - diff --git a/fieldExtraction/src/textract_template.py b/fieldExtraction/src/textract_template.py index 510d8e1..9a4ef6a 100644 --- a/fieldExtraction/src/textract_template.py +++ b/fieldExtraction/src/textract_template.py @@ -1,4 +1,3 @@ - """ This script was written to extract .txt files from pdf for adhoc client runs. Ensure you have permissions for API gateway, S3, Lambda function and SQS to execute this (DEVELOPER & above roles in DEV & UAT for Doczy should suffice) @@ -13,86 +12,75 @@ The output text file can be found in the client_bucket/contract_text_file/batch_ And the final LLM parsed outputs can be found in client_bucket/final_output/batch_123456 """ - - - import os import boto3 import requests import json import boto3 from datetime import datetime -import os +import os from botocore.auth import SigV4Auth from botocore.awsrequest import AWSRequest from botocore.credentials import get_credentials from botocore.session import Session import time + # Function to upload files to S3 def upload_files_to_s3(directory, bucket, batch_id): contract_list = [] for filename in os.listdir(directory): - if filename.endswith('.pdf'): + if filename.endswith(".pdf"): file_path = os.path.join(directory, filename) - s3_key = f'contracts-landing-zone/{batch_id}/{filename}' - + s3_key = f"contracts-landing-zone/{batch_id}/{filename}" + # Upload file to S3 s3_client.upload_file(file_path, bucket, s3_key) - print(f'Uploaded {filename} to S3 bucket\n') + print(f"Uploaded {filename} to S3 bucket\n") # Add tags to the uploaded file s3_client.put_object_tagging( Bucket=bucket, Key=s3_key, - Tagging={ - 'TagSet': [ - { - 'Key': 'BatchId', - 'Value': batch_id - } - ] + Tagging={"TagSet": [{"Key": "BatchId", "Value": batch_id}]}, + ) + print(f"Added tags to {filename}\n") + # Add file details to contract list + # For now we can leave it as is since only A, C have been operationalized + contract_list.append( + { + "contract_name": filename, + "groups": [ + "A", # This needs to be dynamic in UI 1 based on what group has been selected + "C", + ], + "contract_source_path": s3_key, } ) - print(f'Added tags to {filename}\n') - # Add file details to contract list - # For now we can leave it as is since only A, C have been operationalized - contract_list.append({ - "contract_name": filename, - "groups": [ - "A", # This needs to be dynamic in UI 1 based on what group has been selected - "C" - ], - "contract_source_path": s3_key - }) return contract_list - - def create_batch(client_bucket, create_batch_url): - myobj = { "client-bucket-name": client_bucket } + myobj = {"client-bucket-name": client_bucket} # Call create batch API endpoint - response = requests.post(create_batch_url, json = myobj) + response = requests.post(create_batch_url, json=myobj) if response.status_code >= 200 and response.status_code < 300: try: - new_batch_id = json.loads(json.loads(response.text)['body'])['batch_id'] - landing_zone = json.loads(json.loads(response.text)['body'])['landing_zone'] + new_batch_id = json.loads(json.loads(response.text)["body"])["batch_id"] + landing_zone = json.loads(json.loads(response.text)["body"])["landing_zone"] except: print(myobj) print(response.text) - new_batch_id = 'failed_cases' - landing_zone = 'contracts_landing_zone' + new_batch_id = "failed_cases" + landing_zone = "contracts_landing_zone" else: print(response.text) - new_batch_id = 'failed_cases' - landing_zone = 'contracts_landing_zone' + new_batch_id = "failed_cases" + landing_zone = "contracts_landing_zone" return new_batch_id, landing_zone - - def list_filtered_files(bucket_name, prefix, start_date, end_date, profile_name): """ List all files in an S3 bucket filtered by date range. @@ -108,30 +96,29 @@ def list_filtered_files(bucket_name, prefix, start_date, end_date, profile_name) - list of str: the keys of the filtered files """ # Parse the dates - start_date = datetime.fromisoformat(start_date.replace('Z', '+00:00')) - end_date = datetime.fromisoformat(end_date.replace('Z', '+00:00')) + start_date = datetime.fromisoformat(start_date.replace("Z", "+00:00")) + end_date = datetime.fromisoformat(end_date.replace("Z", "+00:00")) # Initialize a session using the specified profile session = boto3.Session(profile_name=profile_name) - s3_client = session.client('s3') + s3_client = session.client("s3") # List objects in the bucket with the specified prefix - paginator = s3_client.get_paginator('list_objects_v2') + paginator = s3_client.get_paginator("list_objects_v2") page_iterator = paginator.paginate(Bucket=bucket_name, Prefix=prefix) # Filter files by date filtered_files = [] for page in page_iterator: - if 'Contents' in page: - for obj in page['Contents']: - last_modified = obj['LastModified'] + if "Contents" in page: + for obj in page["Contents"]: + last_modified = obj["LastModified"] if start_date <= last_modified <= end_date: - filtered_files.append(obj['Key']) + filtered_files.append(obj["Key"]) return filtered_files - def download_files(bucket_name, file_keys, profile_name, local_directory): """ Download files from an S3 bucket. @@ -144,7 +131,7 @@ def download_files(bucket_name, file_keys, profile_name, local_directory): """ # Initialize a session using the specified profile session = boto3.Session(profile_name=profile_name) - s3_client = session.client('s3') + s3_client = session.client("s3") # Ensure the local directory exists if not os.path.exists(local_directory): @@ -160,37 +147,37 @@ def download_files(bucket_name, file_keys, profile_name, local_directory): print("Download completed.") - if __name__ == "__main__": # Define your variables - s3_bucket = 'doczyai-use2-u-cn1-s3-textract-processing-001' + s3_bucket = "doczyai-use2-u-cn1-s3-textract-processing-001" # batch_id = 'batch_100524101551' - client_name = 'Priority Health' - username = 'ADHOC USER' + client_name = "Priority Health" + username = "ADHOC USER" # These endpoints are in UAT, please change them to DEV if there are access issues with UAT - api_endpoint = 'https://29gm8cek03.execute-api.us-east-2.amazonaws.com/dev/trigger-pipeline' - create_batch_url = "https://29gm8cek03.execute-api.us-east-2.amazonaws.com/dev/create-batch" + api_endpoint = ( + "https://29gm8cek03.execute-api.us-east-2.amazonaws.com/dev/trigger-pipeline" + ) + create_batch_url = ( + "https://29gm8cek03.execute-api.us-east-2.amazonaws.com/dev/create-batch" + ) - - - session = boto3.Session(profile_name='temp_cred') # Change the profile name to the one you have in your .aws/credentials file - s3_client = session.client('s3') - - + session = boto3.Session( + profile_name="temp_cred" + ) # Change the profile name to the one you have in your .aws/credentials file + s3_client = session.client("s3") # Use current directory as PDF directory pdf_directory = "C:\\Doczy\\Priority Health\\For_Textract\\For_Textract\\UAT- test East Paris Surgical Center copy" - # Create batch batch_id, landing_zone = create_batch(s3_bucket, create_batch_url) print(f"Batch ID: {batch_id}") print(f"Landing Zone: {landing_zone}") # batch_id = 'batch_110624213433' - # landing_zone = 'contracts_landing_zone/batch_110624213433/' + # landing_zone = 'contracts_landing_zone/batch_110624213433/' - if batch_id == 'failed_cases': + if batch_id == "failed_cases": print("Batch creation failed. Exiting...") exit() # Upload files and get contract list @@ -202,16 +189,12 @@ if __name__ == "__main__": "batch_id": batch_id, "client_name": client_name, "username": username, - "contract_list": contract_list + "contract_list": contract_list, } # Make POST request to API response = requests.post(api_endpoint, json=data) - - - - # Print response print(response.status_code) print(response.json()) @@ -219,18 +202,22 @@ if __name__ == "__main__": ################## # Get text files - prefix = f'contract-text-file/{batch_id}/' # if you have a specific prefix (folder) in your bucket + prefix = f"contract-text-file/{batch_id}/" # if you have a specific prefix (folder) in your bucket # These dates are to filter the contracts in case there are older contracts in the same batch - start_date = '2024-06-04T00:00:00Z' # ISO 8601 format - end_date = '2024-06-14T23:59:59Z' # ISO 8601 format - profile_name = 'temp_cred' - local_directory = 'C:\\Doczy\\Priority Health\\For_Textract\\For_Textract\\text otuput' # local directory to save files + start_date = "2024-06-04T00:00:00Z" # ISO 8601 format + end_date = "2024-06-14T23:59:59Z" # ISO 8601 format + profile_name = "temp_cred" + local_directory = "C:\\Doczy\\Priority Health\\For_Textract\\For_Textract\\text otuput" # local directory to save files # List filtered files - time.sleep(30) # Waiting for the text files to be generated, this may take longer and files may not be available after 30 seconds sometimes - filtered_files = list_filtered_files(s3_bucket, prefix, start_date, end_date, profile_name) + time.sleep( + 30 + ) # Waiting for the text files to be generated, this may take longer and files may not be available after 30 seconds sometimes + filtered_files = list_filtered_files( + s3_bucket, prefix, start_date, end_date, profile_name + ) print(f"Filtered files: {len(filtered_files)}") # Download files - download_files(s3_bucket, filtered_files, profile_name, local_directory) \ No newline at end of file + download_files(s3_bucket, filtered_files, profile_name, local_directory) diff --git a/fieldExtraction/src/top_down_funcs.py b/fieldExtraction/src/top_down_funcs.py index e35c218..505e1e2 100644 --- a/fieldExtraction/src/top_down_funcs.py +++ b/fieldExtraction/src/top_down_funcs.py @@ -1,10 +1,10 @@ - import prompts import claude_funcs import utils import json + def run_top_down(filename, text_dict): """ Executes the Top Down processing strategy on a provided dictionary of text pages, extracting structured data based on specified prompts. @@ -31,17 +31,17 @@ def run_top_down(filename, text_dict): prompt = prompts.TOP_DOWN_PRIMARY(page_text) answer = claude_funcs.invoke_claude_3(prompt, max_tokens=4000) answer_dict = json.loads(answer) - answer_dict.update({'page_num' : page_num, 'Filename' : filename}) - #answer_dicts = dict_operations.primary_string_to_dict({page_num : answer}, filename) # Convert to dictionaries + answer_dict.update({"page_num": page_num, "Filename": filename}) + # answer_dicts = dict_operations.primary_string_to_dict({page_num : answer}, filename) # Convert to dictionaries all_results.append(answer_dict) - + # td_primary = [i for d in all_results for i in d] # Consolidate to one list # print(td_primary) # Secondaries td_final = top_down_secondary(all_results, text_dict) - return td_final # List of list of dictionaries + return td_final # List of list of dictionaries def top_down_secondary(td_results, text_dict): @@ -63,11 +63,13 @@ def top_down_secondary(td_results, text_dict): updated_dicts = [] for d in td_results: # Run Metal Level - d['CONTRACT_MARKETPLACE_METAL_LEVEL'] = run_top_down_metal_level(d, text_dict[d['page_num']]) - + d["CONTRACT_MARKETPLACE_METAL_LEVEL"] = run_top_down_metal_level( + d, text_dict[d["page_num"]] + ) + # Dates # d['LOB_PRICING_TERMS_EFFECTIVE_DATE'] = run_top_down_date('EFFECTIVE' , d, text_dict[d['page_num']]) - + # # if auto_renewal == N: else: 'N/A' # d['LOB_PRICING_TERMS_TERMINATION_DATE'] = run_top_down_date('TERMINATION', d, text_dict[d['page_num']]) @@ -75,7 +77,6 @@ def top_down_secondary(td_results, text_dict): return updated_dicts - def run_top_down_metal_level(d, page): """ Identifies and extracts the metal level of a contract within a specific line of business (LOB) from the provided page text. @@ -91,18 +92,20 @@ def run_top_down_metal_level(d, page): Returns: str: The metal level of the contract as determined by the analysis, or 'N/A' if the contract LOB is not applicable. """ - - if 'MARKETPLACE' in str(d['CONTRACT_LOB']).upper() or 'COMMERCIAL' in str(d['CONTRACT_LOB']).upper(): - prompt = prompts.TOP_DOWN_METAL_LEVEL(d['CONTRACT_LOB'], page) + + if ( + "MARKETPLACE" in str(d["CONTRACT_LOB"]).upper() + or "COMMERCIAL" in str(d["CONTRACT_LOB"]).upper() + ): + prompt = prompts.TOP_DOWN_METAL_LEVEL(d["CONTRACT_LOB"], page) answer = claude_funcs.invoke_claude_3(prompt, max_tokens=4000) else: - answer = 'N/A' + answer = "N/A" return answer - def run_top_down_date(type_, d, page): - """ DEPRECATED + """DEPRECATED Extracts specific date-related information from a page of text using a Top Down processing approach. This function tailors the extraction to focus on either 'EFFECTIVE' or 'TERMINATION' dates by constructing @@ -118,7 +121,7 @@ def run_top_down_date(type_, d, page): Returns: str: The extracted date as a string, based on the model's interpretation of the input prompt and text context. """ - formatted_d = utils.format_td_check([d], ['Filename', 'page_num']) - prompt = prompts.TOP_DOWN_DATE('EFFECTIVE', formatted_d, page) + formatted_d = utils.format_td_check([d], ["Filename", "page_num"]) + prompt = prompts.TOP_DOWN_DATE("EFFECTIVE", formatted_d, page) answer = claude_funcs.invoke_claude_3(prompt, max_tokens=4000) return answer diff --git a/fieldExtraction/src/utils.py b/fieldExtraction/src/utils.py index 83e2b2e..451868d 100644 --- a/fieldExtraction/src/utils.py +++ b/fieldExtraction/src/utils.py @@ -1,4 +1,3 @@ - import os import re import pandas as pd @@ -6,85 +5,97 @@ import shutil import config + def read_local(file_path): # Check if the file is a text file - if os.path.isfile(file_path) and file_path.endswith('.txt'): + if os.path.isfile(file_path) and file_path.endswith(".txt"): try: # First attempt to open the file with UTF-8 encoding - with open(file_path, 'r', encoding='utf-8') as file: + with open(file_path, "r", encoding="utf-8") as file: file_contents = file.read() return file_contents except UnicodeDecodeError: # If UTF-8 fails, try reading the file with ANSI encoding try: - with open(file_path, 'r', encoding='cp1252') as file: + with open(file_path, "r", encoding="cp1252") as file: file_contents = file.read() return file_contents except UnicodeDecodeError: # If ANSI also fails, log an error message or handle it accordingly print(f"Failed to decode {file_path} with UTF-8 and cp1252 encodings.") + def read_s3(): s3_client = config.S3_CLIENT objects = s3_client.list_objects_v2(Bucket=config.BUCKET, Prefix=config.PREFIX) file_list = [] - for obj in objects['Contents']: - if not obj['Key'].endswith('/'): - file_list.append(obj['Key']) + for obj in objects["Contents"]: + if not obj["Key"].endswith("/"): + file_list.append(obj["Key"]) contract_list = sorted(file_list) files = {} for contract in contract_list: data = s3_client.get_object(Bucket=config.BUCKET, Key=contract) - contents = data['Body'].read() - context = contents.decode('utf-8') + contents = data["Body"].read() + context = contents.decode("utf-8") path, filename = os.path.split(contract) files[filename] = context return files + def read_input(path=config.LOCAL_PATH, mode=config.READ_MODE): - if mode == '_LOCAL_': + if mode == "_LOCAL_": files = {} for file in os.listdir(path): full_path = os.path.join(path, file) file_text = read_local(full_path) files[file] = file_text return files - elif mode == '_S3_': + elif mode == "_S3_": return read_s3() -def consolidate_individual(input_folder='results', output_folder='output'): + +def consolidate_individual(input_folder="results", output_folder="output"): dfs = [] for filename in os.listdir(input_folder): - if filename.endswith('.csv'): + if filename.endswith(".csv"): filepath = os.path.join(input_folder, filename) dfs.append(pd.read_csv(filepath)) - + # Remove temp folder - if input_folder=='temp' and os.path.exists(input_folder): + if input_folder == "temp" and os.path.exists(input_folder): shutil.rmtree(input_folder) - + consolidated_df = pd.concat(dfs, ignore_index=True) # Write output version = 1 - existing_files = [filename for filename in os.listdir(output_folder) if filename.startswith(f'consolidated_results_{config.TODAY}')] + existing_files = [ + filename + for filename in os.listdir(output_folder) + if filename.startswith(f"consolidated_results_{config.TODAY}") + ] if existing_files: - versions = [int(file.split('_v')[1].split('.')[0]) for file in existing_files if '_v' in file] + versions = [ + int(file.split("_v")[1].split(".")[0]) + for file in existing_files + if "_v" in file + ] if versions: version = max(versions) + 1 - filename = f'consolidated_results_{config.TODAY}_v{version}.csv' + filename = f"consolidated_results_{config.TODAY}_v{version}.csv" consolidated_df.to_csv(os.path.join(output_folder, filename)) def preprocess_text_file(file_path): # Attempt to open the file with UTF-8 encoding first try: - with open(file_path, 'r', encoding='utf-8') as file: + with open(file_path, "r", encoding="utf-8") as file: text = file.read() except UnicodeDecodeError: # If UTF-8 fails, try reading the file with ANSI (cp1252) encoding try: - with open(file_path, 'r', encoding='cp1252') as file: + with open(file_path, "r", encoding="cp1252") as file: text = file.read() except UnicodeDecodeError: # If ANSI also fails, log an error message or handle it accordingly @@ -92,7 +103,7 @@ def preprocess_text_file(file_path): return [] # Split the text into pages based on a specific marker - pages = re.split(r'Start of Page No\. = \d+', text) + pages = re.split(r"Start of Page No\. = \d+", text) return pages @@ -100,16 +111,18 @@ def format_td_check(td_dicts, dont_include_list): final_str = "" dict_count = 1 for td_dict in td_dicts: - final_str += str(dict_count) + '. ' + final_str += str(dict_count) + ". " for k in td_dict.keys(): if k not in dont_include_list: - final_str += k + ': ' + td_dict[k] + ', ' - final_str += '\n' + final_str += k + ": " + td_dict[k] + ", " + final_str += "\n" dict_count += 1 return final_str -def consolidate_csvs(output_dir=config.CONSOLIDATED_OUTPUT_DIRECTORY, output_file=config.OUTPUT_CSV_PATH): +def consolidate_csvs( + output_dir=config.CONSOLIDATED_OUTPUT_DIRECTORY, output_file=config.OUTPUT_CSV_PATH +): os.makedirs(output_dir, exist_ok=True) df_list = [] @@ -130,23 +143,28 @@ def consolidate_csvs(output_dir=config.CONSOLIDATED_OUTPUT_DIRECTORY, output_fil print(f"All CSV files have been consolidated into {output_file}") - def contains_reimbursement(text, page): if isinstance(text, dict): - return page.isdigit() and ('%' in text[page] or '$' in text[page] or 'percent' in text[page].lower()) + return page.isdigit() and ( + "%" in text[page] or "$" in text[page] or "percent" in text[page].lower() + ) elif isinstance(text, str): - return ('%' in text or '$' in text or 'percent' in text.lower()) + return "%" in text or "$" in text or "percent" in text.lower() else: print("contains_reimbursement - Invalid data type") - def filter_already_processed(input_dict): already_processed = [] for folder_name in os.listdir(config.OUTPUT_DIRECTORY): - if config.PROCESSED_RESULTS_NAME in os.listdir(os.path.join(config.OUTPUT_DIRECTORY, folder_name)): - already_processed.append(folder_name+'.txt') + if config.PROCESSED_RESULTS_NAME in os.listdir( + os.path.join(config.OUTPUT_DIRECTORY, folder_name) + ): + already_processed.append(folder_name + ".txt") - input_dict = {key : input_dict[key] for key in input_dict.keys() if key not in already_processed} + input_dict = { + key: input_dict[key] + for key in input_dict.keys() + if key not in already_processed + } return input_dict - diff --git a/git_diff_changes.py b/git_diff_changes.py index 948c28a..a529aa4 100644 --- a/git_diff_changes.py +++ b/git_diff_changes.py @@ -19,10 +19,19 @@ def check_bitbucket_auth_header(): def get_last_successful_commit(module): if os.getenv("BITBUCKET_CI") != "true": - rc, result = run_command("git rev-parse HEAD", log_output=True, log_cmd=True, log_prefix=f"[{module}]") + rc, result = run_command( + "git rev-parse HEAD", + log_output=True, + log_cmd=True, + log_prefix=f"[{module}]", + ) return result - for commit in subprocess.check_output(["git", "log", "--format=%h", "-n", "30"]).decode().split(): + for commit in ( + subprocess.check_output(["git", "log", "--format=%h", "-n", "30"]) + .decode() + .split() + ): url = f"https://api.bitbucket.org/2.0/repositories/{BITBUCKET_REPO_OWNER}/{BITBUCKET_REPO_SLUG}/commit/{commit}/statuses/" LOGGER.info(f"commit={commit}, url={url}") basic_url_headers = { @@ -31,16 +40,17 @@ def get_last_successful_commit(module): response = requests.get(url, headers=basic_url_headers).json() LOGGER.debug(f"response={response}") - for item in response.get('values', []): - name = item['name'] + for item in response.get("values", []): + name = item["name"] if "security/snyk" in name or "iacbot" in name: continue - commit_state = item['state'] - ref_name = item['refname'] - updated_on = item['updated_on'] + commit_state = item["state"] + ref_name = item["refname"] + updated_on = item["updated_on"] LOGGER.info( - f"COMMIT={commit}, COMMIT_STATE={commit_state}, REF_NAME={ref_name}, NAME={name}, UPDATED_ON={updated_on}") + f"COMMIT={commit}, COMMIT_STATE={commit_state}, REF_NAME={ref_name}, NAME={name}, UPDATED_ON={updated_on}" + ) if commit_state == "SUCCESSFUL" and ref_name.startswith(BB_PIPELINE_BRANCH): return commit @@ -50,12 +60,20 @@ def get_last_successful_commit(module): def get_changes_output(module, last_successful_commit): if os.getenv("BITBUCKET_CI") == "true": - _, result = run_command(f"git diff --dirstat=files,0 {last_successful_commit} | sed -E 's/^[ 0-9.]+% //g'", - log_output=True, log_cmd=True, log_prefix=f"[{module}]") + _, result = run_command( + f"git diff --dirstat=files,0 {last_successful_commit} | sed -E 's/^[ 0-9.]+% //g'", + log_output=True, + log_cmd=True, + log_prefix=f"[{module}]", + ) return result else: - _, result = run_command("git diff --dirstat=files,0 HEAD~1 | sed -E 's/^[ 0-9.]+% //g'", - log_output=True, log_cmd=True, log_prefix=f"[{module}]") + _, result = run_command( + "git diff --dirstat=files,0 HEAD~1 | sed -E 's/^[ 0-9.]+% //g'", + log_output=True, + log_cmd=True, + log_prefix=f"[{module}]", + ) return result @@ -81,7 +99,11 @@ def is_deploy_module(module, dir_prefix): for mod, pattern in special_checks: pattern_found = any(pattern in change for change in changes_output) if module == mod and pattern_found: - LOGGER.info(decorate_warn(f"Marking {module} as changed based on the special pattern={pattern}")) + LOGGER.info( + decorate_warn( + f"Marking {module} as changed based on the special pattern={pattern}" + ) + ) return True if f"{dir_prefix}{module}/" in changes_output: diff --git a/streamlit/constants.py b/streamlit/constants.py index 8ccab24..f64dced 100644 --- a/streamlit/constants.py +++ b/streamlit/constants.py @@ -4,9 +4,17 @@ import os from chromadb.config import Settings # https://python.langchain.com/en/latest/modules/indexes/document_loaders/examples/excel.html?highlight=xlsx#microsoft-excel -from langchain_community.document_loaders import CSVLoader, PDFMinerLoader, TextLoader, UnstructuredExcelLoader, Docx2txtLoader -from langchain_community.document_loaders import UnstructuredFileLoader, UnstructuredMarkdownLoader - +from langchain_community.document_loaders import ( + CSVLoader, + PDFMinerLoader, + TextLoader, + UnstructuredExcelLoader, + Docx2txtLoader, +) +from langchain_community.document_loaders import ( + UnstructuredFileLoader, + UnstructuredMarkdownLoader, +) # load_dotenv() @@ -17,7 +25,7 @@ ROOT_DIRECTORY = "\\\\amznfsxuofkyi1z.aarete.local\\SharedFiles\\AArete Client W SOURCE_DIRECTORY = "SOURCE_DOCUMENTS" OUTPUT_DIRECTORY = f"{ROOT_DIRECTORY}\\Output" -PERSIST_DIRECTORY = 'DB' +PERSIST_DIRECTORY = "DB" MODELS_PATH = "C:\\Users\\Public\\models" @@ -190,48 +198,78 @@ MODEL_BASENAME = "llama-2-7b-chat.Q4_K_M.gguf" # MODEL_BASENAME = "model.safetensors.awq" - ########################################################################################################################################## ## CONSTANTS FOR INFRATRUCTURE -# SSO User list -USER_LIST = ['maamseek@aarete.com', 'smahdavian@aarete.com', 'ahinge@aarete.com', 'akadam@aarete.com', 'pkatariya@aarete.com' -, 'piragavarapu@aarete.com', 'umistry@aarete.com', 'ahutchison@aarete.com', 'bgrunst@aarete.com', 'ddimeglio@aarete.com' -, 'vnair@aarete.com', 'kminhas@aarete.com', 'fmohiuddin@aarete.com', 'slitewka@aarete.com', 'qdoest@aarete.com', 'bkoryga@aarete.com', 'bcielecki@aarete.com', 'mszymanski@aarete.com','hupreti@aarete.com', -'sshingare@aarete.com', 'vsrinivasan@aarete.com' ] +# SSO User list +USER_LIST = [ + "maamseek@aarete.com", + "smahdavian@aarete.com", + "ahinge@aarete.com", + "akadam@aarete.com", + "pkatariya@aarete.com", + "piragavarapu@aarete.com", + "umistry@aarete.com", + "ahutchison@aarete.com", + "bgrunst@aarete.com", + "ddimeglio@aarete.com", + "vnair@aarete.com", + "kminhas@aarete.com", + "fmohiuddin@aarete.com", + "slitewka@aarete.com", + "qdoest@aarete.com", + "bkoryga@aarete.com", + "bcielecki@aarete.com", + "mszymanski@aarete.com", + "hupreti@aarete.com", + "sshingare@aarete.com", + "vsrinivasan@aarete.com", +] # DOCZY DEV -DOCZY_PIPELINE_URL_DEV = 'https://4lzhid1s0h.execute-api.us-east-2.amazonaws.com/dev/trigger-pipeline' -DOCZY_REDIRECT_URL_DEV = 'https://doczydev.aarete.com:850' -DOCZY_CREATE_BATCH_URL_DEV = 'https://4lzhid1s0h.execute-api.us-east-2.amazonaws.com/dev/create-batch' +DOCZY_PIPELINE_URL_DEV = ( + "https://4lzhid1s0h.execute-api.us-east-2.amazonaws.com/dev/trigger-pipeline" +) +DOCZY_REDIRECT_URL_DEV = "https://doczydev.aarete.com:850" +DOCZY_CREATE_BATCH_URL_DEV = ( + "https://4lzhid1s0h.execute-api.us-east-2.amazonaws.com/dev/create-batch" +) # DOCZY UAT -DOCZY_PIPELINE_URL_UAT = 'https://29gm8cek03.execute-api.us-east-2.amazonaws.com/dev/trigger-pipeline' -DOCZY_REDIRECT_URL_UAT = 'https://doczyuat.aarete.com:850' -DOCZY_CREATE_BATCH_URL_UAT = 'https://29gm8cek03.execute-api.us-east-2.amazonaws.com/dev/create-batch' +DOCZY_PIPELINE_URL_UAT = ( + "https://29gm8cek03.execute-api.us-east-2.amazonaws.com/dev/trigger-pipeline" +) +DOCZY_REDIRECT_URL_UAT = "https://doczyuat.aarete.com:850" +DOCZY_CREATE_BATCH_URL_UAT = ( + "https://29gm8cek03.execute-api.us-east-2.amazonaws.com/dev/create-batch" +) # DOCZY PROD -DOCZY_PIPELINE_URL_PROD = 'https://d612isd3ja.execute-api.us-east-2.amazonaws.com/dev/trigger-pipeline' -DOCZY_REDIRECT_URL_PROD = 'https://doczy.aarete.com:850' -DOCZY_CREATE_BATCH_URL_PROD = 'https://d612isd3ja.execute-api.us-east-2.amazonaws.com/dev/create-batch' +DOCZY_PIPELINE_URL_PROD = ( + "https://d612isd3ja.execute-api.us-east-2.amazonaws.com/dev/trigger-pipeline" +) +DOCZY_REDIRECT_URL_PROD = "https://doczy.aarete.com:850" +DOCZY_CREATE_BATCH_URL_PROD = ( + "https://d612isd3ja.execute-api.us-east-2.amazonaws.com/dev/create-batch" +) -# SNOWFLAKE DEV DATABASE -SNOWFLAKE_ACCOUNT_LOCATOR="aarete-doczyai", -DEV_DB_ROLE = "DEVADMIN", -DEV_WH="DEV_XS", -DEV_DB="DOCZY_DEV", -DEV_STAGING_SCHEMA="STG" +# SNOWFLAKE DEV DATABASE +SNOWFLAKE_ACCOUNT_LOCATOR = ("aarete-doczyai",) +DEV_DB_ROLE = ("DEVADMIN",) +DEV_WH = ("DEV_XS",) +DEV_DB = ("DOCZY_DEV",) +DEV_STAGING_SCHEMA = "STG" # SNOWFLAKE UAT DATABASE -UAT_DB_ROLE = "UATADMIN", -UAT_WH="DEV_XS", -UAT_DB="DOCZY_UAT", -UAT_STAGING_SCHEMA="STG" +UAT_DB_ROLE = ("UATADMIN",) +UAT_WH = ("DEV_XS",) +UAT_DB = ("DOCZY_UAT",) +UAT_STAGING_SCHEMA = "STG" # SNOWFLAKE PROD DATABASE -PROD_DB_ROLE = "PRODADMIN", -PROD_WH="DEV_XS", -PROD_DB="DOCZY_PROD", -PROD_STAGING_SCHEMA="STG" \ No newline at end of file +PROD_DB_ROLE = ("PRODADMIN",) +PROD_WH = ("DEV_XS",) +PROD_DB = ("DOCZY_PROD",) +PROD_STAGING_SCHEMA = "STG" diff --git a/streamlit/ingest.py b/streamlit/ingest.py index 6046c77..5396a36 100644 --- a/streamlit/ingest.py +++ b/streamlit/ingest.py @@ -22,27 +22,30 @@ from constants import ( import boto3 from langchain.embeddings.bedrock import BedrockEmbeddings + def file_log(logentry): - file1 = open("file_ingest.log","a") - file1.write(logentry + "\n") - file1.close() - print(logentry + "\n") + file1 = open("file_ingest.log", "a") + file1.write(logentry + "\n") + file1.close() + print(logentry + "\n") + def load_single_document(file_path: str) -> Document: # Loads a single document from a file path try: - file_extension = os.path.splitext(file_path)[1] - loader_class = DOCUMENT_MAP.get(file_extension) - if loader_class: - file_log(file_path + ' loaded.') - loader = loader_class(file_path) - else: - file_log(file_path + ' document type is undefined.') - raise ValueError("Document type is undefined") - return loader.load()[0] + file_extension = os.path.splitext(file_path)[1] + loader_class = DOCUMENT_MAP.get(file_extension) + if loader_class: + file_log(file_path + " loaded.") + loader = loader_class(file_path) + else: + file_log(file_path + " document type is undefined.") + raise ValueError("Document type is undefined") + return loader.load()[0] except Exception as ex: - file_log('%s loading error: \n%s' % (file_path, ex)) - return None + file_log("%s loading error: \n%s" % (file_path, ex)) + return None + def load_document_batch(filepaths): logging.info("Loading document batch") @@ -52,12 +55,12 @@ def load_document_batch(filepaths): futures = [exe.submit(load_single_document, name) for name in filepaths] # collect data if futures is None: - file_log(name + ' failed to submit') - return None + file_log(name + " failed to submit") + return None else: - data_list = [future.result() for future in futures] - # return data and file paths - return (data_list, filepaths) + data_list = [future.result() for future in futures] + # return data and file paths + return (data_list, filepaths) def load_documents(source_dir: str) -> list[Document]: @@ -65,7 +68,7 @@ def load_documents(source_dir: str) -> list[Document]: paths = [] for root, _, files in os.walk(source_dir): for file_name in files: - print('Importing: ' + file_name) + print("Importing: " + file_name) file_extension = os.path.splitext(file_name)[1] source_file_path = os.path.join(root, file_name) if file_extension in DOCUMENT_MAP.keys(): @@ -83,12 +86,12 @@ def load_documents(source_dir: str) -> list[Document]: filepaths = paths[i : (i + chunksize)] # submit the task try: - future = executor.submit(load_document_batch, filepaths) + future = executor.submit(load_document_batch, filepaths) except Exception as ex: - file_log('executor task failed: %s' % (ex)) - future = None + file_log("executor task failed: %s" % (ex)) + future = None if future is not None: - futures.append(future) + futures.append(future) # process all results for future in as_completed(futures): # open the file and load the data @@ -96,8 +99,8 @@ def load_documents(source_dir: str) -> list[Document]: contents, _ = future.result() docs.extend(contents) except Exception as ex: - file_log('Exception: %s' % (ex)) - + file_log("Exception: %s" % (ex)) + return docs @@ -106,16 +109,18 @@ def split_documents(documents: list[Document]) -> tuple[list[Document], list[Doc text_docs, python_docs = [], [] for doc in documents: if doc is not None: - file_extension = os.path.splitext(doc.metadata["source"])[1] - if file_extension == ".py": - python_docs.append(doc) - else: - text_docs.append(doc) + file_extension = os.path.splitext(doc.metadata["source"])[1] + if file_extension == ".py": + python_docs.append(doc) + else: + text_docs.append(doc) return text_docs, python_docs + def process_in_batches(texts, batch_size): for i in range(0, len(texts), batch_size): - yield texts[i:i+batch_size] + yield texts[i : i + batch_size] + @click.command() @click.option( @@ -146,30 +151,35 @@ def process_in_batches(texts, batch_size): ), help="Device to run on. (Default is cuda)", ) - def main(device_type): # Load documents and split in chunks logging.info(f"Loading documents from {SOURCE_DIRECTORY}") - for filename in os.listdir(SOURCE_DIRECTORY): - with open(os.path.join(SOURCE_DIRECTORY, filename), 'r') as infile: + for filename in os.listdir(SOURCE_DIRECTORY): + with open(os.path.join(SOURCE_DIRECTORY, filename), "r") as infile: text = infile.read() - text_splitted = [i for i in text.split('Start of Page No. = ')] + text_splitted = [i for i in text.split("Start of Page No. = ")] for i, txt in enumerate(text_splitted): if len(txt) > 2: - page_path = os.path.join(SOURCE_DIRECTORY, f'{filename[:-4]}_page{i}.txt') - with open(page_path, 'w') as f: + page_path = os.path.join( + SOURCE_DIRECTORY, f"{filename[:-4]}_page{i}.txt" + ) + with open(page_path, "w") as f: f.write(txt) documents = [load_single_document(page_path)] os.remove(page_path) text_documents, python_documents = split_documents(documents) - text_splitter = RecursiveCharacterTextSplitter(chunk_size=1000, chunk_overlap=200) + text_splitter = RecursiveCharacterTextSplitter( + chunk_size=1000, chunk_overlap=200 + ) python_splitter = RecursiveCharacterTextSplitter.from_language( language=Language.PYTHON, chunk_size=880, chunk_overlap=200 ) texts = text_splitter.split_documents(text_documents) texts.extend(python_splitter.split_documents(python_documents)) - logging.info(f"Loaded {len(documents)} documents from {SOURCE_DIRECTORY}") + logging.info( + f"Loaded {len(documents)} documents from {SOURCE_DIRECTORY}" + ) logging.info(f"Split into {len(texts)} chunks of text") # Create embeddings @@ -206,11 +216,11 @@ def main(device_type): collection_metadata={"hnsw:space": "cosine"}, # ids = [str(filename)] ) - if __name__ == "__main__": logging.basicConfig( - format="%(asctime)s - %(levelname)s - %(filename)s:%(lineno)s - %(message)s", level=logging.INFO + format="%(asctime)s - %(levelname)s - %(filename)s:%(lineno)s - %(message)s", + level=logging.INFO, ) main() diff --git a/streamlit/interface_0.py b/streamlit/interface_0.py index 0e37a1f..0373837 100644 --- a/streamlit/interface_0.py +++ b/streamlit/interface_0.py @@ -20,20 +20,21 @@ from util import logger (REDIRECT_URI, create_batch_url, doczy_pipeline) = util.load_page_details(0) user_list = USER_LIST -if 'uploading' not in st.session_state: +if "uploading" not in st.session_state: st.session_state.uploading = False -if 'upload_key' not in st.session_state: - st.session_state.upload_key = 0 -if 'file_list' not in st.session_state: +if "upload_key" not in st.session_state: + st.session_state.upload_key = 0 +if "file_list" not in st.session_state: st.session_state.file_list = [] -if 'show_batchID' not in st.session_state: +if "show_batchID" not in st.session_state: st.session_state.show_batchID = False -if 'landing_zone' not in st.session_state: - st.session_state.landing_zone = 'contracts-landing-zone' -if 'batch_id' not in st.session_state: - st.session_state.batch_id = 'failed_cases' -if 'client_bucket' not in st.session_state: - st.session_state.client_bucket = 'default_bucket' +if "landing_zone" not in st.session_state: + st.session_state.landing_zone = "contracts-landing-zone" +if "batch_id" not in st.session_state: + st.session_state.batch_id = "failed_cases" +if "client_bucket" not in st.session_state: + st.session_state.client_bucket = "default_bucket" + def insert_upload_logs(batch_id, client_name, file_name, upload_datetime, upload_user): """ @@ -41,12 +42,12 @@ def insert_upload_logs(batch_id, client_name, file_name, upload_datetime, upload Output: status of the insert query """ try: - return 'Log inserted successfully' + return "Log inserted successfully" except Exception as e: return e -st.set_page_config(layout = "wide") +st.set_page_config(layout="wide") # # # Sidebar contents with st.sidebar: st.title("Doczy.AI ™") @@ -61,42 +62,60 @@ with st.sidebar: # st.write("Doczy") # AARETE LOGO -x,y,z = st.columns([15,2,15]) +x, y, z = st.columns([15, 2, 15]) with y: - st.image('aaretelogo.png') + st.image("aaretelogo.png") -hide_img_fs = ''' +hide_img_fs = """ -''' +""" st.markdown(hide_img_fs, unsafe_allow_html=True) -_,c1= st.columns([5,1]) +_, c1 = st.columns([5, 1]) try: util.setup_page(REDIRECT_URI) except Exception as e: st.write(f"SSO Failed = {e}") - st.session_state['user_info'] = {'mail': 'maamseek@aarete.com', 'displayName': 'Mayank Aamseek'} # RECOMMENDATION - Remove Line after development. + st.session_state["user_info"] = { + "mail": "maamseek@aarete.com", + "displayName": "Mayank Aamseek", + } # RECOMMENDATION - Remove Line after development. try: c1.write(f"User: **{st.session_state.user_info['displayName']}**") - user_mail = st.session_state.user_info['mail'] -except KeyError as e: # RECOMMENDATION - Not Working as Expected. Unreachable Code because the above Exception handles everything + user_mail = st.session_state.user_info["mail"] +except ( + KeyError +) as e: # RECOMMENDATION - Not Working as Expected. Unreachable Code because the above Exception handles everything st.write("Session Expired.") auth_url = security.get_auth_url(REDIRECT_URI) - st.markdown(f"Sign In", unsafe_allow_html=True) + st.markdown( + f"Sign In", unsafe_allow_html=True + ) st.stop() -s3_client = boto3.client('s3', # RECOMMENDATION - We can make use of AWS Access Key for enabling this in local testing. Will help a lot. +s3_client = boto3.client( + "s3", # RECOMMENDATION - We can make use of AWS Access Key for enabling this in local testing. Will help a lot. region_name="us-east-2", ) -client_list = ['doczy-ai-client-1', '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'] +client_list = [ + "doczy-ai-client-1", + "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", +] -# This is the list of client fetched from Snowflake +# This is the list of client fetched from Snowflake # TODO: Need to update the streamlit code to use the client names from this list # And use the s3 paths to save the objects for the respective client client_list, s3_paths = get_client_names() @@ -108,7 +127,9 @@ 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',(client_list), label_visibility = "collapsed", index = None) + client = st.selectbox( + "Client Name", (client_list), label_visibility="collapsed", index=None + ) client_bucket = client_s3_paths.get(client) @@ -119,20 +140,30 @@ file_row = st.columns([0.1, 0.8]) with file_row[0]: st.write("**Upload Files**") with file_row[1]: - file_list = st.file_uploader("Upload", type=['docx','tiff','pdf'], accept_multiple_files=True, label_visibility = "collapsed", help="Only PDF, TIFF and DOCX file formats are supported.", disabled=st.session_state.uploading, key = st.session_state.upload_key) + file_list = st.file_uploader( + "Upload", + type=["docx", "tiff", "pdf"], + accept_multiple_files=True, + label_visibility="collapsed", + help="Only PDF, TIFF and DOCX file formats are supported.", + disabled=st.session_state.uploading, + key=st.session_state.upload_key, + ) add_vertical_space(2) -df = pd.DataFrame(columns=['Contract Name']) -df['Contract Name'] = file_list +df = pd.DataFrame(columns=["Contract Name"]) +df["Contract Name"] = file_list file_names = [] buttons = st.columns([0.4, 0.4, 0.2]) + def set_uploading_state(): if not client == None and not len(file_list) == 0: st.session_state.file_list = file_list st.session_state.upload_key += 1 st.session_state.uploading = True + with buttons[1]: if st.button("Create Batch", on_click=set_uploading_state): file_list = st.session_state.file_list @@ -141,41 +172,57 @@ with buttons[1]: elif len(file_list) == 0: st.error("No Files Selected.") else: - with st.spinner('Running...'): - myobj = { "client-bucket-name": client_bucket } - response = requests.post(create_batch_url, json = myobj) + with st.spinner("Running..."): + myobj = {"client-bucket-name": client_bucket} + response = requests.post(create_batch_url, json=myobj) if response.status_code >= 200 and response.status_code < 300: try: - batch_id = json.loads(json.loads(response.text)['body'])['batch_id'] - landing_zone = json.loads(json.loads(response.text)['body'])['landing_zone'] - landing_zone = 'contracts-landing-zone/' + batch_id = json.loads(json.loads(response.text)["body"])[ + "batch_id" + ] + landing_zone = json.loads(json.loads(response.text)["body"])[ + "landing_zone" + ] + landing_zone = "contracts-landing-zone/" except: # st.write(myobj) # st.write(response.text) st.write("Internal Error. Reach out to Doczy.AI Team.") - batch_id = 'failed_cases' - landing_zone = 'contracts-landing-zone/' + batch_id = "failed_cases" + landing_zone = "contracts-landing-zone/" else: st.error("Failed") for uploaded_file in file_list: stringio = BytesIO(uploaded_file.getvalue()) stringio.seek(0) - s3_client.put_object(Bucket=client_bucket, Body=stringio.getvalue(), Key= - landing_zone+batch_id+'/'+str(uploaded_file.name)) - s3_client.put_object_tagging(Bucket=client_bucket, Key= - landing_zone+batch_id+'/'+str(uploaded_file.name), Tagging = {'TagSet': [ { 'Key': 'BatchId', 'Value': batch_id }]}) + s3_client.put_object( + Bucket=client_bucket, + Body=stringio.getvalue(), + Key=landing_zone + batch_id + "/" + str(uploaded_file.name), + ) + s3_client.put_object_tagging( + Bucket=client_bucket, + Key=landing_zone + batch_id + "/" + str(uploaded_file.name), + Tagging={"TagSet": [{"Key": "BatchId", "Value": batch_id}]}, + ) - # TODO: Test this insert function with snowflake - upload_log = insert_upload_logs(batch_id, client, str(uploaded_file.name), datetime.now().strftime("%Y-%m-%d %H:%M:%S"), user_mail) + # TODO: Test this insert function with snowflake + upload_log = insert_upload_logs( + batch_id, + client, + str(uploaded_file.name), + datetime.now().strftime("%Y-%m-%d %H:%M:%S"), + user_mail, + ) st.write(upload_log) - + file_names.append(str(uploaded_file.name)) st.session_state.uploading = False st.session_state.show_batchID = True st.session_state.client_bucket = client_bucket st.session_state.batch_id = batch_id st.session_state.landing_zone = landing_zone - st.rerun() + st.rerun() if st.session_state.show_batchID: batch_id = st.session_state.batch_id client_bucket = st.session_state.client_bucket @@ -185,13 +232,12 @@ with buttons[1]: st.write(f"Files uploaded to s3://{client_bucket}/{landing_zone}{batch_id}") st.session_state.show_batchID = False st.session_state.file_list = [] - st.session_state.batch_id = 'failed_cases' - st.session_state.client_bucket = 'default_bucket' - st.session_state.landing_zone = 'contracts-landing-zone' - + st.session_state.batch_id = "failed_cases" + st.session_state.client_bucket = "default_bucket" + st.session_state.landing_zone = "contracts-landing-zone" # RECOMMENDATION - Rethink about how to best manage cache - # @st.cache_data + # @st.cache_data # def convert_df(df): # return df.to_csv(index=False).encode('utf-8') diff --git a/streamlit/interface_1.py b/streamlit/interface_1.py index 2b80d9a..7662beb 100644 --- a/streamlit/interface_1.py +++ b/streamlit/interface_1.py @@ -17,7 +17,7 @@ from util import logger (REDIRECT_URI, create_batch_url, doczy_pipeline) = util.load_page_details(1) user_list = USER_LIST -st.set_page_config(layout = "wide") +st.set_page_config(layout="wide") # Sidebar contents with st.sidebar: st.title("Doczy.AI ™") @@ -32,38 +32,43 @@ with st.sidebar: # st.write("Doczy") # AARETE LOGO -x,y,z = st.columns([15,2,15]) +x, y, z = st.columns([15, 2, 15]) with y: - st.image('aaretelogo.png') + st.image("aaretelogo.png") -hide_img_fs = ''' +hide_img_fs = """ -''' +""" st.markdown(hide_img_fs, unsafe_allow_html=True) -_,c1= st.columns([5,1]) +_, c1 = st.columns([5, 1]) try: util.setup_page(REDIRECT_URI) except: st.write("SSO Failed") - st.session_state['user_info'] = {'mail': 'maamseek@aarete.com', 'displayName': 'Mayank Aamseek'} # RECOMMENDATION - Remove after dev phase + st.session_state["user_info"] = { + "mail": "maamseek@aarete.com", + "displayName": "Mayank Aamseek", + } # RECOMMENDATION - Remove after dev phase try: c1.write(f"User: **{st.session_state.user_info['displayName']}**") - user_mail = st.session_state.user_info['mail'] + user_mail = st.session_state.user_info["mail"] except KeyError as e: - # Do we add a link to get to the login page here? + # Do we add a link to get to the login page here? st.write("Session Expired.") # st.write("Please sign-in to use this app.") auth_url = security.get_auth_url(REDIRECT_URI) - st.markdown(f"Sign In", unsafe_allow_html=True) + st.markdown( + f"Sign In", unsafe_allow_html=True + ) st.stop() - -s3_client = boto3.client('s3', +s3_client = boto3.client( + "s3", region_name="us-east-2", ) @@ -79,7 +84,9 @@ 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',(client_list), label_visibility = "collapsed", index = None) + client = st.selectbox( + "Client Name", (client_list), label_visibility="collapsed", index=None + ) if client: client_bucket = client_s3_paths.get(client) @@ -87,23 +94,27 @@ if client: # to be deleted when buckets for different clients are ready; below line is added only for testing the corresponding DAG # client_bucket = 'doczyai-use2-d-cn1-s3-textract-processing-001' - batch_objects = s3_client.list_objects_v2(Bucket=client_bucket - , Prefix="contracts-landing-zone/", Delimiter='/') + batch_objects = s3_client.list_objects_v2( + Bucket=client_bucket, Prefix="contracts-landing-zone/", Delimiter="/" + ) batch_list = [] - for prefix in batch_objects['CommonPrefixes']: - batch_name = prefix['Prefix'][:-1].split('/')[-1] - batch_objects2 = s3_client.list_objects_v2(Bucket=client_bucket, Prefix="contracts-landing-zone/"+batch_name+"/", Delimiter='/') - if 'Contents' in batch_objects2 and len(batch_objects2['Contents']) > 0: - batch_list.append(prefix['Prefix'][:-1].split('/')[-1]) + for prefix in batch_objects["CommonPrefixes"]: + batch_name = prefix["Prefix"][:-1].split("/")[-1] + batch_objects2 = s3_client.list_objects_v2( + Bucket=client_bucket, + Prefix="contracts-landing-zone/" + batch_name + "/", + Delimiter="/", + ) + if "Contents" in batch_objects2 and len(batch_objects2["Contents"]) > 0: + batch_list.append(prefix["Prefix"][:-1].split("/")[-1]) # Hardcoded batch_list for testing purposes - # batch_list = ['batch_020524103737', 'batch_090524131433', 'batch_090524131607', 'batch_100524123000', 'batch_130524064322', + # batch_list = ['batch_020524103737', 'batch_090524131433', 'batch_090524131607', 'batch_100524123000', 'batch_130524064322', # 'batch_160524071331', 'batch_200524213550', 'batch_250424112237', 'batch_280524120530', 'batch_280524121721', 'batch_280524144222', # 'batch_290524123926', 'batch_290524164044', 'batch_310524102029', 'batch_310524124050', 'batch_310524162346', 'batch_310524162631'] - - #Select Box for Applying Different Sort for the Batch List + # Select Box for Applying Different Sort for the Batch List # if 'sorted_list' not in st.session_state: # st.session_state.sorted_list = batch_list @@ -118,14 +129,13 @@ if client: # col1, col2, col3, col4 = st.columns([0.5, 0.5, 0.5, 0.5]) - - # with col1: + # with col1: # sort_by = st.radio("**Sort Batch_IDs**", ('Alphabetical', 'Create Date')) # with col2: # order = st.radio('', ('Ascending','Descending')) - # with col3: + # with col3: # add_vertical_space(2) # if st.button('Apply'): # st.session_state.sorted_list = sort_list(batch_list, sort_by, order) @@ -134,7 +144,12 @@ if client: with path_row[0]: st.write("**Batch ID**") with path_row[1]: - batch_id = st.selectbox('**Batch ID**', reversed(batch_list) , label_visibility = "collapsed", index = None) + batch_id = st.selectbox( + "**Batch ID**", + reversed(batch_list), + label_visibility="collapsed", + index=None, + ) if not batch_id: batch_id = "None" @@ -144,71 +159,93 @@ if client: st.write("**Group No.**") with checks[1]: - a = st.checkbox('Unique Key', key = str(1), args="Unique") + a = st.checkbox("Unique Key", key=str(1), args="Unique") with checks[2]: - b = st.checkbox('Pricing Before Carveouts', key = str(2)) + b = st.checkbox("Pricing Before Carveouts", key=str(2)) with checks[3]: - c = st.checkbox('Contract Related', key = str(3)) + c = st.checkbox("Contract Related", key=str(3)) with checks[4]: - d = st.checkbox('Provider', key = str(4)) + d = st.checkbox("Provider", key=str(4)) with checks[5]: - e = st.checkbox('Timeline', key = str(5)) + e = st.checkbox("Timeline", key=str(5)) with checks[6]: - f = st.checkbox('Carveout Indicator', key = str(6)) + f = st.checkbox("Carveout Indicator", key=str(6)) with checks[7]: - g = st.checkbox('Carveout Methodology', key = str(7)) + g = st.checkbox("Carveout Methodology", key=str(7)) add_vertical_space(1) - df = pd.DataFrame(columns=['Contract Name', 'Unique Key','Pricing Before Carveouts' - , 'Contract Related', 'Provider', 'Timeline', 'Carveout Indicator', 'Carveout Methodology']) + df = pd.DataFrame( + columns=[ + "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=client_bucket - , Prefix="contracts-landing-zone/"+batch_id+"/", Delimiter='/') + file_objects = s3_client.list_objects_v2( + Bucket=client_bucket, + Prefix="contracts-landing-zone/" + batch_id + "/", + Delimiter="/", + ) # Hardcoded file_list for testing purposes - # file_list = ['Boilerplate_TX Amendment Mission Health Network effective_040114 MU.pdf', 'Custom_TX - MP AMENDMENT - MISSION HEALTH NETWORK - MU.pdf', + # file_list = ['Boilerplate_TX Amendment Mission Health Network effective_040114 MU.pdf', 'Custom_TX - MP AMENDMENT - MISSION HEALTH NETWORK - MU.pdf', # 'Delaware First Health_First State Homecare Agency_212260_7 MU.pdf', 'Molina Healthcare of Texas, Inc. Amendment 4 - HIX ACA__EFF 01012016_MU.pdf'] 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 + 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 + 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) + print(f"DEBUGGING: PWD= {dir_path}") + df.to_csv("temp1.csv", index=False) add_vertical_space(1) - df2 = pd.read_csv('temp1.csv') + df2 = pd.read_csv("temp1.csv") edited_df = st.data_editor(df2) - edited_df['REQUEST_USER'] = user_mail - edited_df['LATEST_FLAG BOOLEAN'] = True - edited_df['PIPELINE_KICKOFF_DATETIME'] = datetime.now().strftime("%Y-%m-%d %H:%M:%S") - edited_df['REQUEST_DATETIME'] = datetime.now().strftime("%Y-%m-%d %H:%M:%S") + edited_df["REQUEST_USER"] = user_mail + edited_df["LATEST_FLAG BOOLEAN"] = True + edited_df["PIPELINE_KICKOFF_DATETIME"] = datetime.now().strftime( + "%Y-%m-%d %H:%M:%S" + ) + edited_df["REQUEST_DATETIME"] = datetime.now().strftime("%Y-%m-%d %H:%M:%S") - @st.cache_data + @st.cache_data def convert_df(df): - return df.to_csv(index=False).encode('utf-8') + return df.to_csv(index=False).encode("utf-8") csv = convert_df(edited_df) # edited_df = edited_df.reset_index() # make sure indexes pair with number of rows # additional_info = pd.DataFrame(columns=['REQUEST_ID','T_DRIVE_PATH','CLIENT_NAME' # , 'GROUP_NAME', 'REQUEST_USERNAME', 'REQUEST_DATETIME']) - additional_info = pd.DataFrame(columns=['CLIENT_NAME', 'BATCH_ID', 'REQUEST_USERNAME', 'REQUEST_DATETIME']) - additional_info.loc[0] = [client, batch_id, st.session_state.user_info['mail'], datetime.now().strftime("%Y-%m-%d %H:%M:%S")] + additional_info = pd.DataFrame( + columns=["CLIENT_NAME", "BATCH_ID", "REQUEST_USERNAME", "REQUEST_DATETIME"] + ) + additional_info.loc[0] = [ + client, + batch_id, + st.session_state.user_info["mail"], + datetime.now().strftime("%Y-%m-%d %H:%M:%S"), + ] st.write(additional_info) st.session_state.contract_count = 0 @@ -216,32 +253,36 @@ if client: for index, row in edited_df.iterrows(): allow_run_for_contract = False group_list = [] - if row['Unique Key']: - group_list.append('A') + if row["Unique Key"]: + group_list.append("A") allow_run_for_contract = True - if row['Pricing Before Carveouts']: - group_list.append('B') + if row["Pricing Before Carveouts"]: + group_list.append("B") allow_run_for_contract = True - if row['Contract Related']: - group_list.append('C') + if row["Contract Related"]: + group_list.append("C") allow_run_for_contract = True - if row['Provider']: - group_list.append('D') + if row["Provider"]: + group_list.append("D") allow_run_for_contract = True - if row['Timeline']: - group_list.append('E') + if row["Timeline"]: + group_list.append("E") allow_run_for_contract = True - if row['Carveout Indicator']: - group_list.append('F') + if row["Carveout Indicator"]: + group_list.append("F") allow_run_for_contract = True - if row['Carveout Methodology']: - group_list.append('G') + if row["Carveout Methodology"]: + group_list.append("G") allow_run_for_contract = True - if allow_run_for_contract: st.session_state.contract_count += 1 + if allow_run_for_contract: + st.session_state.contract_count += 1 entry_dict = { - "contract_name": row['Contract Name'], + "contract_name": row["Contract Name"], "groups": group_list, - "contract_source_path": "contracts-landing-zone/"+batch_id+"/"+row['Contract Name'] + "contract_source_path": "contracts-landing-zone/" + + batch_id + + "/" + + row["Contract Name"], } contract_list.append(entry_dict) @@ -250,18 +291,20 @@ if client: "batch_id": batch_id, "client_name": client, "username": user_mail, - "contract_list": contract_list + "contract_list": contract_list, } buttons = st.columns([0.8, 0.2]) with buttons[0]: - st.download_button("Download Table", csv, "file.csv", "text/csv", key='download-csv') + st.download_button( + "Download Table", csv, "file.csv", "text/csv", key="download-csv" + ) with buttons[1]: if st.button("Run Doczy.AI Pipeline"): if not st.session_state.contract_count == len(edited_df): st.error("Select at least one Group No. for every Contract") else: - with st.spinner('Running...'): + with st.spinner("Running..."): # csv_buf = StringIO() # additional_info.to_csv(csv_buf, header=True, index=False) # csv_buf.seek(0) @@ -269,12 +312,12 @@ if client: # csv_buf = StringIO() # edited_df.to_csv(csv_buf, header=True, index=False) # csv_buf.seek(0) - # s3_client.put_object(Bucket='doczy-dev-infra-raw-data-ingestion', Body=csv_buf.getvalue(), Key='training_interface/contract_config.csv') + # s3_client.put_object(Bucket='doczy-dev-infra-raw-data-ingestion', Body=csv_buf.getvalue(), Key='training_interface/contract_config.csv') # try: # save_to_sf('load_request_and_contract_submissions', request_submission_file_name = "request_submission.csv", contract_config_file_name = "contract_config.csv") # except Exception as e: # st.write(e) - response = requests.post(doczy_pipeline, json = myobj) + response = requests.post(doczy_pipeline, json=myobj) if response.status_code >= 200 and response.status_code < 300: success_message = """
@@ -291,4 +334,3 @@ if client: """ st.markdown(failure_message, unsafe_allow_html=True) # st.write(response.text) - diff --git a/streamlit/interface_2.py b/streamlit/interface_2.py index 860b78d..24cdb5a 100644 --- a/streamlit/interface_2.py +++ b/streamlit/interface_2.py @@ -5,12 +5,20 @@ from langchain.prompts import PromptTemplate from langchain.embeddings.bedrock import BedrockEmbeddings from langchain.llms.bedrock import Bedrock from langchain_community.vectorstores import Chroma -from constants import CHROMA_SETTINGS, EMBEDDING_MODEL_NAME, PERSIST_DIRECTORY, MODEL_ID, MODEL_BASENAME, SOURCE_DIRECTORY, USER_LIST +from constants import ( + CHROMA_SETTINGS, + EMBEDDING_MODEL_NAME, + PERSIST_DIRECTORY, + MODEL_ID, + MODEL_BASENAME, + SOURCE_DIRECTORY, + USER_LIST, +) from langchain.chains import RetrievalQA import streamlit as st from streamlit_extras.add_vertical_space import add_vertical_space -from streamlit_pdf_viewer import pdf_viewer # needs to be installed on the server +from streamlit_pdf_viewer import pdf_viewer # needs to be installed on the server import os import pandas as pd import numpy as np @@ -29,7 +37,7 @@ from util import logger user_list = USER_LIST -st.set_page_config(layout = "wide") +st.set_page_config(layout="wide") # Sidebar contents with st.sidebar: st.title("Doczy.AI ™") @@ -44,60 +52,69 @@ with st.sidebar: # st.write("Doczy") # AARETE LOGO -x,y,z = st.columns([15,2,15]) +x, y, z = st.columns([15, 2, 15]) with y: - st.image('aaretelogo.png') + st.image("aaretelogo.png") -hide_img_fs = ''' +hide_img_fs = """ -''' +""" st.markdown(hide_img_fs, unsafe_allow_html=True) -_,c1= st.columns([5,1]) +_, c1 = st.columns([5, 1]) try: - util.setup_page(REDIRECT_URI) # RECOMMENDATION - Not secure enough + util.setup_page(REDIRECT_URI) # RECOMMENDATION - Not secure enough except: st.write("SSO Failed") - st.session_state['user_info'] = {'mail': 'maamseek@aarete.com', 'displayName': 'Mayank Aamseek'} # RECOMMENDATION - Remove after dev phase + st.session_state["user_info"] = { + "mail": "maamseek@aarete.com", + "displayName": "Mayank Aamseek", + } # RECOMMENDATION - Remove after dev phase try: c1.write(f"User: **{st.session_state.user_info['displayName']}**") - user_mail = st.session_state.user_info['mail'] + user_mail = st.session_state.user_info["mail"] except KeyError as e: st.write("Session Expired.") - #st.write("Please sign-in to use this app.") + # 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.markdown( + f"Sign In", unsafe_allow_html=True + ) st.stop() # remove below try except statement if comparison with actual vales is not required try: - conn = get_snowflake_conn('STG') + conn = get_snowflake_conn("STG") cur = conn.cursor() except: st.write("Conn failed, unable to fetch data from training data table in Snowflake") try: - query = 'select * from "PROMPT_CONFIG"' # RECOMMENDATION - Move query somewhere else + query = ( + 'select * from "PROMPT_CONFIG"' # RECOMMENDATION - Move query somewhere else + ) cur.execute(query) - fields = pd.DataFrame.from_records(iter(cur), columns=[x[0] for x in cur.description]) + fields = pd.DataFrame.from_records( + iter(cur), columns=[x[0] for x in cur.description] + ) - fields.rename(columns={'FIELD_DESC': 'Field Name'}, inplace = True) # RECOMMENDATION - Package into less hardcoded function or handle in DB - fields.rename(columns={'PROMPT': 'Interrogation Question?'}, inplace = True) - fields.rename(columns={'GROUP_ID': 'PRIORITY'}, inplace = True) - fields.rename(columns={'FIELD_NAME': 'SF_DB_COL_NAME'}, inplace = True) - fields.rename(columns={'FM_MODEL_ID': 'llm_selected'}, inplace = True) + fields.rename( + columns={"FIELD_DESC": "Field Name"}, inplace=True + ) # RECOMMENDATION - Package into less hardcoded function or handle in DB + fields.rename(columns={"PROMPT": "Interrogation Question?"}, inplace=True) + fields.rename(columns={"GROUP_ID": "PRIORITY"}, inplace=True) + fields.rename(columns={"FIELD_NAME": "SF_DB_COL_NAME"}, inplace=True) + fields.rename(columns={"FM_MODEL_ID": "llm_selected"}, inplace=True) except Exception as e: - st.write("Unable to fetch data from Snowflake: ",e) + st.write("Unable to fetch data from Snowflake: ", e) # fields = pd.read_csv('contract_fields.csv', encoding='unicode_escape', skipinitialspace=True) # fields = fields[~fields['SF_COL_NAME'].str.endswith('_PG', na=None)] # change the code below if contract list is fetched from snowflake -s3_client = boto3.client('s3', - region_name="us-east-2" -) +s3_client = boto3.client("s3", region_name="us-east-2") client_list, s3_paths = get_client_names() client_s3_paths = dict(zip(client_list, s3_paths)) @@ -106,35 +123,43 @@ client_row = st.columns([0.2, 0.7, 0.1]) with client_row[0]: st.write("**Client Name**") with client_row[1]: - client = st.selectbox('Client Name',(client_list), label_visibility = "collapsed", index= None) + client = st.selectbox( + "Client Name", (client_list), label_visibility="collapsed", index=None + ) if client: client_bucket = client_s3_paths.get(client) # client_bucket = 'doczyai-use2-d-cn1-s3-textract-processing-001' - batch_objects = s3_client.list_objects_v2(Bucket=client_bucket - , Prefix="textract-receiver-processed-pdfs/", Delimiter='/') + batch_objects = s3_client.list_objects_v2( + Bucket=client_bucket, Prefix="textract-receiver-processed-pdfs/", Delimiter="/" + ) batch_list = [] - for prefix in batch_objects['CommonPrefixes']: - batch_list.append(prefix['Prefix'][:-1].split('/')[-1]) + for prefix in batch_objects["CommonPrefixes"]: + batch_list.append(prefix["Prefix"][:-1].split("/")[-1]) path_row = st.columns([0.2, 0.7, 0.1]) with path_row[0]: st.write("**Batch ID**") with path_row[1]: - batch_id = st.selectbox('**Batch ID**', batch_list, label_visibility = "collapsed", index = None) + batch_id = st.selectbox( + "**Batch ID**", batch_list, label_visibility="collapsed", index=None + ) if batch_id: - objects = s3_client.list_objects_v2(Bucket=client_bucket, Prefix="textract-receiver-processed-pdfs/"+batch_id+"/") + objects = s3_client.list_objects_v2( + Bucket=client_bucket, + Prefix="textract-receiver-processed-pdfs/" + batch_id + "/", + ) file_list = [] - if 'Contents' in objects: - for obj in objects['Contents']: - if not obj['Key'].endswith('/'): - file_list.append(obj['Key']) + if "Contents" in objects: + for obj in objects["Contents"]: + if not obj["Key"].endswith("/"): + file_list.append(obj["Key"]) else: - st.error('This batch_id is empty.') + st.error("This batch_id is empty.") contract_list = sorted(file_list) @@ -142,59 +167,80 @@ if client: with file_row[0]: st.write("**Contract Name**") with file_row[1]: - file_name = st.selectbox('Select a file', ['All'] + contract_list, label_visibility = "collapsed", index= None) + file_name = st.selectbox( + "Select a file", + ["All"] + contract_list, + label_visibility="collapsed", + index=None, + ) field_row = st.columns([0.2, 0.7, 0.1]) with field_row[0]: st.write("**Field Group**") with field_row[1]: - field_group = st.selectbox('Field Group',('Unique Key', 'Contract Related', 'Pricing Before Carveouts - I' # RECOMMENDATION - Get from a separate list or snowflake - , 'Pricing Before Carveouts - II', 'Carveout Indicator, Code Type and Code #s - I' - , 'Carveout Indicator, Code Type and Code #s - II', 'Carveout Indicator, Code Type and Code #s - III' - , 'Optimize Carving Indic.', 'Carveout Method - I', 'Carveout Method - II', 'Provider' - , 'Timeline'), label_visibility = "collapsed") + field_group = st.selectbox( + "Field Group", + ( + "Unique Key", + "Contract Related", + "Pricing Before Carveouts - I", # RECOMMENDATION - Get from a separate list or snowflake + "Pricing Before Carveouts - II", + "Carveout Indicator, Code Type and Code #s - I", + "Carveout Indicator, Code Type and Code #s - II", + "Carveout Indicator, Code Type and Code #s - III", + "Optimize Carving Indic.", + "Carveout Method - I", + "Carveout Method - II", + "Provider", + "Timeline", + ), + label_visibility="collapsed", + ) - if field_group == 'Unique Key': # RECOMMENDATION - Use a Dictionary for this mapping - fields = fields[fields['PRIORITY'] == 'A'] - elif field_group == 'Contract Related': - fields = fields[fields['PRIORITY'] == 'C'] - elif field_group == 'Pricing Before Carveouts - I': - fields = fields[fields['PRIORITY'] == 'B'] + if ( + field_group == "Unique Key" + ): # RECOMMENDATION - Use a Dictionary for this mapping + fields = fields[fields["PRIORITY"] == "A"] + elif field_group == "Contract Related": + fields = fields[fields["PRIORITY"] == "C"] + elif field_group == "Pricing Before Carveouts - I": + fields = fields[fields["PRIORITY"] == "B"] fields = np.array_split(fields, 2)[0] - elif field_group == 'Pricing Before Carveouts - II': - fields = fields[fields['PRIORITY'] == 'B'] + elif field_group == "Pricing Before Carveouts - II": + fields = fields[fields["PRIORITY"] == "B"] fields = np.array_split(fields, 2)[1] - elif field_group == 'Carveout Indicator, Code Type and Code #s - I': - fields = fields[fields['PRIORITY'] == 'F'] + elif field_group == "Carveout Indicator, Code Type and Code #s - I": + fields = fields[fields["PRIORITY"] == "F"] fields = np.array_split(fields, 3)[0] - elif field_group == 'Carveout Indicator, Code Type and Code #s - II': - fields = fields[fields['PRIORITY'] == 'F'] + elif field_group == "Carveout Indicator, Code Type and Code #s - II": + fields = fields[fields["PRIORITY"] == "F"] fields = np.array_split(fields, 3)[1] - elif field_group == 'Carveout Indicator, Code Type and Code #s - III': - fields = fields[fields['PRIORITY'] == 'F'] + elif field_group == "Carveout Indicator, Code Type and Code #s - III": + fields = fields[fields["PRIORITY"] == "F"] fields = np.array_split(fields, 3)[2] - elif field_group == 'Carveout Methodology - I': - fields = fields[fields['PRIORITY'] == 'G'] + elif field_group == "Carveout Methodology - I": + fields = fields[fields["PRIORITY"] == "G"] fields = np.array_split(fields, 4)[0] - elif field_group == 'Carveout Methodology - II': - fields = fields[fields['PRIORITY'] == 'G'] + elif field_group == "Carveout Methodology - II": + fields = fields[fields["PRIORITY"] == "G"] fields = np.array_split(fields, 4)[1] - elif field_group == 'Carveout Method - III': - fields = fields[fields['PRIORITY'] == 'G'] + elif field_group == "Carveout Method - III": + fields = fields[fields["PRIORITY"] == "G"] fields = np.array_split(fields, 4)[2] - elif field_group == 'Carveout Method - IV': - fields = fields[fields['PRIORITY'] == 'G'] + elif field_group == "Carveout Method - IV": + fields = fields[fields["PRIORITY"] == "G"] fields = np.array_split(fields, 4)[3] - elif field_group == 'Provider': - fields = fields[fields['PRIORITY'] == 'D'] - elif field_group == 'Timeline': - fields = fields[fields['PRIORITY'] == 'E'] + elif field_group == "Provider": + fields = fields[fields["PRIORITY"] == "D"] + elif field_group == "Timeline": + fields = fields[fields["PRIORITY"] == "E"] if st.button("Show Results"): - query = 'select * from "DOCZY_PIPELINE_RAW_OUTPUT"' # RECOMMENDATION - Move query somewhere else + query = 'select * from "DOCZY_PIPELINE_RAW_OUTPUT"' # RECOMMENDATION - Move query somewhere else cur.execute(query) - df2 = pd.DataFrame.from_records(iter(cur), columns=[x[0] for x in cur.description]) - + df2 = pd.DataFrame.from_records( + iter(cur), columns=[x[0] for x in cur.description] + ) # Changed To Dictionary Mapping for Faster Runtime (Replace with above code when using more fields) @@ -223,8 +269,6 @@ if client: # if split is not None and index is not None: # fields = np.array_split(fields, split)[index] - - # Changed To Dictionary Mapping for Faster Runtime (Replace with above code when using more fields) # Define the mapping @@ -252,8 +296,7 @@ if client: # if split is not None and index is not None: # fields = np.array_split(fields, split)[index] - - button_cols = st.columns([1,2,6]) + button_cols = st.columns([1, 2, 6]) with button_cols[0]: if st.button("Show PDF"): if file_name == None or file_name == "All": @@ -270,28 +313,55 @@ if client: """, unsafe_allow_html=True, ) - s3_obj = s3_client.get_object(Bucket = client_bucket, Key = file_name) - data=s3_obj['Body'].read() + s3_obj = s3_client.get_object( + Bucket=client_bucket, Key=file_name + ) + data = s3_obj["Body"].read() pdf_viewer(data, width=1500) with button_cols[1]: if st.button("Show Results"): try: if field_group: - if field_group == 'Unique Key' or field_group == 'Contract Related': - query = f'select * from "DOCZY_PIPELINE_RAW_OUTPUT_AC" where batch_id = \'{batch_id}\'' - elif field_group == 'Pricing Before Carveouts': - query = f'select * from "DOCZY_PIPELINE_RAW_OUTPUT_B" where batch_id = \'{batch_id}\'' + if ( + field_group == "Unique Key" + or field_group == "Contract Related" + ): + query = f"select * from \"DOCZY_PIPELINE_RAW_OUTPUT_AC\" where batch_id = '{batch_id}'" + elif field_group == "Pricing Before Carveouts": + query = f"select * from \"DOCZY_PIPELINE_RAW_OUTPUT_B\" where batch_id = '{batch_id}'" cur.execute(query) - df2 = pd.DataFrame.from_records(iter(cur), columns=[x[0] for x in cur.description]) - else: - st.error('Please select a Field Group') - df2 = pd.DataFrame(columns=['Contract Name','Field Name', 'SF_DB_COL_NAME', 'Snippet','Page Number' - , 'Field Extracted Value', 'Actual Value','Imputed Value']) + df2 = pd.DataFrame.from_records( + iter(cur), columns=[x[0] for x in cur.description] + ) + else: + st.error("Please select a Field Group") + df2 = pd.DataFrame( + columns=[ + "Contract Name", + "Field Name", + "SF_DB_COL_NAME", + "Snippet", + "Page Number", + "Field Extracted Value", + "Actual Value", + "Imputed Value", + ] + ) except: - df2 = pd.DataFrame(columns=['Contract Name','Field Name', 'SF_DB_COL_NAME', 'Snippet','Page Number' - , 'Field Extracted Value', 'Actual Value','Imputed Value']) - st.error('There was an error in fetching the output.') - df2.to_csv('temp2.csv', index=False) + df2 = pd.DataFrame( + columns=[ + "Contract Name", + "Field Name", + "SF_DB_COL_NAME", + "Snippet", + "Page Number", + "Field Extracted Value", + "Actual Value", + "Imputed Value", + ] + ) + st.error("There was an error in fetching the output.") + df2.to_csv("temp2.csv", index=False) # if st.button("Show PDF"): # if file_name == None or file_name == "All": @@ -315,23 +385,25 @@ if client: # ) # st.markdown(pdf_display, unsafe_allow_html=True) - df2 = pd.read_csv('temp2.csv') - df2['Imputed Value'] = '' + df2 = pd.read_csv("temp2.csv") + df2["Imputed Value"] = "" edited_df = st.data_editor(df2) - @st.cache_data + @st.cache_data def convert_df(df): - return df.to_csv(index=False).encode('utf-8') + return df.to_csv(index=False).encode("utf-8") csv = convert_df(edited_df) buttons = st.columns(3) with buttons[0]: - # st.button("Save All Imputations") - st.download_button("Download Table", csv, "file.csv", "text/csv", key='download-csv') + # st.button("Save All Imputations") + st.download_button( + "Download Table", csv, "file.csv", "text/csv", key="download-csv" + ) with buttons[1]: # st.download_button("Download Table", csv, "file.csv", "text/csv", key='download-csv') st.write("") with buttons[2]: if st.button("Kickoff Database Integration"): - st.write("Stored in DB") \ No newline at end of file + st.write("Stored in DB") diff --git a/streamlit/interface_2_rag.py b/streamlit/interface_2_rag.py index e519c66..109aeb6 100644 --- a/streamlit/interface_2_rag.py +++ b/streamlit/interface_2_rag.py @@ -5,7 +5,15 @@ from langchain.prompts import PromptTemplate from langchain.embeddings.bedrock import BedrockEmbeddings from langchain.llms.bedrock import Bedrock from langchain_community.vectorstores import Chroma -from constants import CHROMA_SETTINGS, EMBEDDING_MODEL_NAME, PERSIST_DIRECTORY, MODEL_ID, MODEL_BASENAME, SOURCE_DIRECTORY, USER_LIST +from constants import ( + CHROMA_SETTINGS, + EMBEDDING_MODEL_NAME, + PERSIST_DIRECTORY, + MODEL_ID, + MODEL_BASENAME, + SOURCE_DIRECTORY, + USER_LIST, +) from langchain.chains import RetrievalQA import streamlit as st @@ -14,10 +22,10 @@ import os import pandas as pd import util -REDIRECT_URI = 'https://doczydev.aarete.com:8502' +REDIRECT_URI = "https://doczydev.aarete.com:8502" user_list = USER_LIST -st.set_page_config(layout = "wide") +st.set_page_config(layout="wide") # Sidebar contents with st.sidebar: st.title("Doczy.AI ™") @@ -32,20 +40,28 @@ with st.sidebar: # st.write("Doczy") util.setup_page(REDIRECT_URI) -_,c1= st.columns([5,1]) +_, c1 = st.columns([5, 1]) c1.write(f"User: **{st.session_state.user_info['displayName']}**") -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?'])) +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=SOURCE_DIRECTORY): filenames = os.listdir(folder_path) - selected_filename = st.selectbox('Select a file', filenames, label_visibility = "collapsed") + selected_filename = st.selectbox( + "Select a file", filenames, label_visibility="collapsed" + ) # return os.path.join(folder_path, selected_filename) return selected_filename @@ -66,16 +82,25 @@ if st.session_state.user_info['mail'] in user_list: 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") + llm_selected = st.selectbox( + "Langauge Model", + ( + "Claude 2", + "Claude Instant", + "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), 'r') as infile: + 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_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') @@ -84,8 +109,7 @@ if st.session_state.user_info['mail'] in user_list: # Setup bedrock bedrock_runtime = boto3.client( - service_name="bedrock-runtime", - region_name="us-east-1" + service_name="bedrock-runtime", region_name="us-east-1" ) embeddings = BedrockEmbeddings( @@ -100,7 +124,7 @@ if st.session_state.user_info['mail'] in user_list: RETRIEVER = DB.as_retriever(search_kwargs={"filter": {"$or": page_list}, "k": 4}) # if "LLM" not in st.session_state: - if llm_selected == 'Titan Text Express': + if llm_selected == "Titan Text Express": LLM = Bedrock( model_id="amazon.titan-text-express-v1", client=bedrock_runtime, @@ -109,9 +133,9 @@ if st.session_state.user_info['mail'] in user_list: "stopSequences": [], "temperature": 0, "topP": 1, - } + }, ) - elif llm_selected == 'Llama 2 Chat 70B': + elif llm_selected == "Llama 2 Chat 70B": LLM = Bedrock( model_id="meta.llama2-70b-chat-v1", client=bedrock_runtime, @@ -119,9 +143,9 @@ if st.session_state.user_info['mail'] in user_list: "max_gen_len": 512, "temperature": 0, # "topP": 0.9, - } + }, ) - elif llm_selected == 'Llama 2 Chat 13B': + elif llm_selected == "Llama 2 Chat 13B": LLM = Bedrock( model_id="meta.llama2-13b-chat-v1", client=bedrock_runtime, @@ -129,9 +153,9 @@ if st.session_state.user_info['mail'] in user_list: "max_gen_len": 512, "temperature": 0, # "topP": 0.9, - } + }, ) - elif llm_selected == 'Claude Instant': + elif llm_selected == "Claude Instant": LLM = Bedrock( model_id="anthropic.claude-instant-v1", client=bedrock_runtime, @@ -139,9 +163,9 @@ if st.session_state.user_info['mail'] in user_list: # "max_tokens_to_sample": 512, "temperature": 0, # "topP": 0.9, - } + }, ) - elif llm_selected == 'Claude 2': + elif llm_selected == "Claude 2": LLM = Bedrock( model_id="anthropic.claude-v2:1", client=bedrock_runtime, @@ -149,7 +173,7 @@ if st.session_state.user_info['mail'] in user_list: # "max_tokens_to_sample": 512, "temperature": 0, # "topP": 0.9, - } + }, ) st.session_state["LLM"] = LLM @@ -180,51 +204,68 @@ if st.session_state.user_info['mail'] in user_list: # 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']) + 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] + 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] + 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] + 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) + 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'] = '' + df2 = pd.read_csv("temp2.csv") + df2["Imputed Value"] = "" edited_df = st.data_editor(df2) - @st.cache_data + @st.cache_data def convert_df(df): - return df.to_csv(index=False).encode('utf-8') + return df.to_csv(index=False).encode("utf-8") csv = convert_df(edited_df) buttons = st.columns(3) with buttons[0]: - st.button("Save All Imputations") + st.button("Save All Imputations") with buttons[1]: - st.download_button("Download Table", csv, "file.csv", "text/csv", key='download-csv') + st.download_button( + "Download Table", csv, "file.csv", "text/csv", key="download-csv" + ) with buttons[2]: - st.button("Kickoff Database Integration") + st.button("Kickoff Database Integration") else: st.write("Access Denied") - - - diff --git a/streamlit/interface_3.py b/streamlit/interface_3.py index 788357c..544e9c9 100644 --- a/streamlit/interface_3.py +++ b/streamlit/interface_3.py @@ -4,8 +4,15 @@ from langchain.prompts import PromptTemplate from langchain.embeddings.bedrock import BedrockEmbeddings from langchain.llms.bedrock import Bedrock from langchain_community.vectorstores import Chroma -from constants import CHROMA_SETTINGS, EMBEDDING_MODEL_NAME, PERSIST_DIRECTORY, MODEL_ID, MODEL_BASENAME, \ - SOURCE_DIRECTORY, USER_LIST +from constants import ( + CHROMA_SETTINGS, + EMBEDDING_MODEL_NAME, + PERSIST_DIRECTORY, + MODEL_ID, + MODEL_BASENAME, + SOURCE_DIRECTORY, + USER_LIST, +) from langchain.chains import RetrievalQA import streamlit as st @@ -48,10 +55,13 @@ try: util.setup_page(redirect_uri) except: st.write("SSO Failed") - st.session_state['user_info'] = {'mail': 'maamseek@aarete.com', 'displayName': 'Mayank Aamseek'} + st.session_state["user_info"] = { + "mail": "maamseek@aarete.com", + "displayName": "Mayank Aamseek", + } try: c1.write(f"User: **{st.session_state.user_info['displayName']}**") - user_mail = st.session_state.user_info['mail'] + user_mail = st.session_state.user_info["mail"] except KeyError as e: st.write("Session Expired.") st.stop() @@ -59,165 +69,209 @@ except KeyError as e: try: sf_secrets = json.loads(get_secret()) conn = snowflake.connector.connect( - user=sf_secrets.get('user'), - password=sf_secrets.get('password'), + user=sf_secrets.get("user"), + password=sf_secrets.get("password"), account="aarete-doczyai", role="DEVADMIN", warehouse="DEV_XS", database="DOCZY_DEV", - schema="STG" + schema="STG", ) cur = conn.cursor() query = 'select * from "TRAINING_DATA_RAW"' cur.execute(query) - field_values = pd.DataFrame.from_records(iter(cur), columns=[x[0] for x in cur.description]) + field_values = pd.DataFrame.from_records( + iter(cur), columns=[x[0] for x in cur.description] + ) # st.write(field_values) - field_values['Document_Name'] = field_values['DOCUMENT_NAME'] + field_values["Document_Name"] = field_values["DOCUMENT_NAME"] # field_values['Contract ID'] = field_values['CONTRACT_TITLE'] # error('table values are incorrect') except: - field_values = pd.read_csv('contract_field_values.csv', encoding='utf-8-sig', skipinitialspace=True) + field_values = pd.read_csv( + "contract_field_values.csv", encoding="utf-8-sig", skipinitialspace=True + ) # field_values.rename(columns={'(internal) Document Name': 'Document_Name'}, inplace = True) ## field_values.rename(columns={'(Internal) Carveout ID': 'Contract ID'}, inplace = True) - field_values['Document_Name'] = field_values['DOCUMENT_NAME'] - field_values = field_values.loc[:, ~field_values.columns.str.contains('Unnamed:')] + field_values["Document_Name"] = field_values["DOCUMENT_NAME"] + field_values = field_values.loc[:, ~field_values.columns.str.contains("Unnamed:")] st.write("Local copy of TRAINING_DATA_RAW table loaded") try: # error('table is not updated') query = 'select * from "BUSINESS_CONFIG"' cur.execute(query) - fields = pd.DataFrame.from_records(iter(cur), columns=[x[0] for x in cur.description]) - fields.rename(columns={'FIELD_NAME': 'Field Name'}, inplace=True) - fields.rename(columns={'QUESTION': 'Interrogation Question?'}, inplace=True) - fields.rename(columns={'SF_COL_NAME': 'SF_DB_COL_NAME'}, inplace=True) - fields = fields[~fields['SF_DB_COL_NAME'].str.endswith('_PG', na=None)] - fields['Field Name'] = fields['SF_DB_COL_NAME'] + fields = pd.DataFrame.from_records( + iter(cur), columns=[x[0] for x in cur.description] + ) + fields.rename(columns={"FIELD_NAME": "Field Name"}, inplace=True) + fields.rename(columns={"QUESTION": "Interrogation Question?"}, inplace=True) + fields.rename(columns={"SF_COL_NAME": "SF_DB_COL_NAME"}, inplace=True) + fields = fields[~fields["SF_DB_COL_NAME"].str.endswith("_PG", na=None)] + fields["Field Name"] = fields["SF_DB_COL_NAME"] # fields.rename(columns={'FM_MODEL_ID': 'llm_selected'}, inplace = True) # error('table is empty') except: - fields = pd.read_csv('contract_fields.csv', encoding='utf-8-sig', skipinitialspace=True) + fields = pd.read_csv( + "contract_fields.csv", encoding="utf-8-sig", skipinitialspace=True + ) # fields = fields.drop_duplicates(subset='Field Name', keep="first").sort_values('Field Name') # fields['Field Name'] = fields['SF_DB_COL_NAME'] # fields = fields[~fields['Field Name'].isnull()] - fields.rename(columns={'FIELD_NAME': 'Field Name'}, inplace=True) - fields.rename(columns={'QUESTION': 'Interrogation Question?'}, inplace=True) - fields.rename(columns={'SF_COL_NAME': 'SF_DB_COL_NAME'}, inplace=True) - fields = fields[~fields['SF_DB_COL_NAME'].str.endswith('_PG', na=None)] - fields['Field Name'] = fields['SF_DB_COL_NAME'] + fields.rename(columns={"FIELD_NAME": "Field Name"}, inplace=True) + fields.rename(columns={"QUESTION": "Interrogation Question?"}, inplace=True) + fields.rename(columns={"SF_COL_NAME": "SF_DB_COL_NAME"}, inplace=True) + fields = fields[~fields["SF_DB_COL_NAME"].str.endswith("_PG", na=None)] + fields["Field Name"] = fields["SF_DB_COL_NAME"] st.write("Local copy of BUSINESS_CONFIG table loaded") try: query = 'select * from "TRAINING_ATTEMPT_LOGS"' cur.execute(query) - history = pd.DataFrame.from_records(iter(cur), columns=[x[0] for x in cur.description]) - history.rename(columns={'FIELD_NAME': 'Field Name'}, inplace=True) - history.rename(columns={'CONTRACTS_TESTED': '# Contracts Tested'}, inplace=True) - history.rename(columns={'USERNAME': 'Username'}, inplace=True) - history.rename(columns={'DATE_TIME': 'Date/Time'}, inplace=True) - history.rename(columns={'ACCURACY': 'Accuracy'}, inplace=True) - history.rename(columns={'ATTEMPT_NUM': 'Attempt #'}, inplace=True) - history = history[['Field Name', '# Contracts Tested', 'Username', 'Date/Time', 'Accuracy', 'Attempt #']] + history = pd.DataFrame.from_records( + iter(cur), columns=[x[0] for x in cur.description] + ) + history.rename(columns={"FIELD_NAME": "Field Name"}, inplace=True) + history.rename(columns={"CONTRACTS_TESTED": "# Contracts Tested"}, inplace=True) + history.rename(columns={"USERNAME": "Username"}, inplace=True) + history.rename(columns={"DATE_TIME": "Date/Time"}, inplace=True) + history.rename(columns={"ACCURACY": "Accuracy"}, inplace=True) + history.rename(columns={"ATTEMPT_NUM": "Attempt #"}, inplace=True) + history = history[ + [ + "Field Name", + "# Contracts Tested", + "Username", + "Date/Time", + "Accuracy", + "Attempt #", + ] + ] except: - history = pd.read_csv('history.csv') + history = pd.read_csv("history.csv") st.write("Local copy of TRAINING_ATTEMPT_LOGS table loaded") field_row = st.columns([0.15, 0.45, 0.4]) with field_row[0]: st.write("**Field Group**") with field_row[1]: - field_group = st.selectbox('Field Group', ('Unique and Contract Related', 'Pricing Before Carveouts - All' - , 'Pricing Before Carveouts - I', 'Pricing Before Carveouts - II', - 'Carveout Indicator, Code Type and Code #s - I' - , 'Carveout Indicator, Code Type and Code #s - II', - 'Carveout Indicator, Code Type and Code #s - III' - , 'Optimize Carving Indic.', 'Carveout Method - I', - 'Carveout Method - II', 'Provider' - , 'Timeline'), label_visibility="collapsed") + field_group = st.selectbox( + "Field Group", + ( + "Unique and Contract Related", + "Pricing Before Carveouts - All", + "Pricing Before Carveouts - I", + "Pricing Before Carveouts - II", + "Carveout Indicator, Code Type and Code #s - I", + "Carveout Indicator, Code Type and Code #s - II", + "Carveout Indicator, Code Type and Code #s - III", + "Optimize Carving Indic.", + "Carveout Method - I", + "Carveout Method - II", + "Provider", + "Timeline", + ), + label_visibility="collapsed", + ) # priorty column will be relaced by group_id in snowflake db -if field_group == 'Unique and Contract Related': - fields = fields[fields['PRIORITY'].isin(['A', 'C'])] -elif field_group == 'Pricing Before Carveouts - All': - fields = fields[fields['PRIORITY'] == 'B'] -elif field_group == 'Pricing Before Carveouts - I': - fields = fields[fields['PRIORITY'] == 'B'] +if field_group == "Unique and Contract Related": + fields = fields[fields["PRIORITY"].isin(["A", "C"])] +elif field_group == "Pricing Before Carveouts - All": + fields = fields[fields["PRIORITY"] == "B"] +elif field_group == "Pricing Before Carveouts - I": + fields = fields[fields["PRIORITY"] == "B"] fields = np.array_split(fields, 2)[0] -elif field_group == 'Pricing Before Carveouts - II': - fields = fields[fields['PRIORITY'] == 'B'] +elif field_group == "Pricing Before Carveouts - II": + fields = fields[fields["PRIORITY"] == "B"] fields = np.array_split(fields, 2)[1] -elif field_group == 'Carveout Indicator, Code Type and Code #s - I': - fields = fields[fields['PRIORITY'] == 'F'] +elif field_group == "Carveout Indicator, Code Type and Code #s - I": + fields = fields[fields["PRIORITY"] == "F"] fields = np.array_split(fields, 3)[0] -elif field_group == 'Carveout Indicator, Code Type and Code #s - II': - fields = fields[fields['PRIORITY'] == 'F'] +elif field_group == "Carveout Indicator, Code Type and Code #s - II": + fields = fields[fields["PRIORITY"] == "F"] fields = np.array_split(fields, 3)[1] -elif field_group == 'Carveout Indicator, Code Type and Code #s - III': - fields = fields[fields['PRIORITY'] == 'F'] +elif field_group == "Carveout Indicator, Code Type and Code #s - III": + fields = fields[fields["PRIORITY"] == "F"] fields = np.array_split(fields, 3)[2] -elif field_group == 'Carveout Methodology - I': - fields = fields[fields['PRIORITY'] == 'G'] +elif field_group == "Carveout Methodology - I": + fields = fields[fields["PRIORITY"] == "G"] fields = np.array_split(fields, 4)[0] -elif field_group == 'Carveout Methodology - II': - fields = fields[fields['PRIORITY'] == 'G'] +elif field_group == "Carveout Methodology - II": + fields = fields[fields["PRIORITY"] == "G"] fields = np.array_split(fields, 4)[1] -elif field_group == 'Carveout Method - III': - fields = fields[fields['PRIORITY'] == 'G'] +elif field_group == "Carveout Method - III": + fields = fields[fields["PRIORITY"] == "G"] fields = np.array_split(fields, 4)[2] -elif field_group == 'Carveout Method - IV': - fields = fields[fields['PRIORITY'] == 'G'] +elif field_group == "Carveout Method - IV": + fields = fields[fields["PRIORITY"] == "G"] fields = np.array_split(fields, 4)[3] -elif field_group == 'Provider': - fields = fields[fields['PRIORITY'] == 'D'] -elif field_group == 'Timeline': - fields = fields[fields['PRIORITY'] == 'E'] +elif field_group == "Provider": + fields = fields[fields["PRIORITY"] == "D"] +elif field_group == "Timeline": + fields = fields[fields["PRIORITY"] == "E"] -fields['Interrogation Question?'] = fields['Interrogation Question?'].fillna(' ') -field_prompt_mapping = dict(zip(fields['Field Name'], fields['Interrogation Question?'])) +fields["Interrogation Question?"] = fields["Interrogation Question?"].fillna(" ") +field_prompt_mapping = dict( + zip(fields["Field Name"], fields["Interrogation Question?"]) +) mode_row = st.columns([0.15, 0.45, 0.4]) with mode_row[0]: st.write("**Mode**") with mode_row[1]: - mode = st.selectbox('Mode', ('Single field - Non Empty values', 'Multiple fields', 'One-to-many fields'), index=0, - label_visibility="collapsed") + mode = st.selectbox( + "Mode", + ("Single field - Non Empty values", "Multiple fields", "One-to-many fields"), + 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]: - if mode == 'Single field - Non Empty values': - field = st.selectbox('Field Name', sorted(set(field_prompt_mapping.keys())), index=0, - label_visibility="collapsed") + if mode == "Single field - Non Empty values": + field = st.selectbox( + "Field Name", + sorted(set(field_prompt_mapping.keys())), + index=0, + label_visibility="collapsed", + ) else: - field = st.multiselect('Field Name', sorted(set(field_prompt_mapping.keys())), sorted( - set(field_prompt_mapping.keys())), label_visibility="collapsed") + field = st.multiselect( + "Field Name", + sorted(set(field_prompt_mapping.keys())), + sorted(set(field_prompt_mapping.keys())), + label_visibility="collapsed", + ) field_prompt_mapping = {key: field_prompt_mapping[key] for key in field} 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', '100', '200', 'Part-1', 'Part-2', 'All'), index=1, - label_visibility="collapsed") + contract_count = st.selectbox( + "Contract count", + ("1", "10", "20", "30", "50", "100", "200", "Part-1", "Part-2", "All"), + index=1, + label_visibility="collapsed", + ) -s3_client = boto3.client('s3', - region_name="us-east-2" - ) -bucket = 'doczy-dev-infra-textract' +s3_client = boto3.client("s3", region_name="us-east-2") +bucket = "doczy-dev-infra-textract" objects = s3_client.list_objects_v2(Bucket=bucket, Prefix="training-data/") file_list = [] -for obj in objects['Contents']: - if not obj['Key'].endswith('/'): - file_list.append(obj['Key']) +for obj in objects["Contents"]: + if not obj["Key"].endswith("/"): + file_list.append(obj["Key"]) # print(os.path.join('RAW_DOCUMENTS', obj['Key'].rsplit('/',1)[1])) # s3_client.download_file('doczy-dev-infra-textract', obj['Key'], os.path.join('RAW_DOCUMENTS', obj['Key'].rsplit('/',1)[1])) contract_list = sorted(file_list) # contract_list = sorted(os.listdir(SOURCE_DIRECTORY)) -if mode == 'Single field - Non Empty values': +if mode == "Single field - Non Empty values": # column_name = fields.loc[fields['Field Name'] == field, 'SF_DB_COL_NAME'].iloc[0] # field_values = field_values[~field_values[column_name].isnull()] field_values = field_values[~field_values[field].isnull()] @@ -226,58 +280,77 @@ if mode == 'Single field - Non Empty values': # st.write(df) # st.write(len(contract_list)) # contract_list = [contract for contract in contract_list if contract.rsplit('/',1)[1].replace(' MU','').replace('_MU','').replace('.txt','') in list(field_values['Document_Name'])] -contract_list = [str(contract)[:-4] + '.pdf' for contract in contract_list] -contract_list = [contract for contract in contract_list if - contract.rsplit('/', 1)[1] in list(field_values['Document_Name'])] +contract_list = [str(contract)[:-4] + ".pdf" for contract in contract_list] +contract_list = [ + contract + for contract in contract_list + if contract.rsplit("/", 1)[1] in list(field_values["Document_Name"]) +] # st.write(len(contract_list)) -if contract_count == 'All': +if contract_count == "All": contract_count = len(contract_list) -if contract_count == 'Part-1': +if contract_count == "Part-1": contract_list = np.array_split(contract_list, 2)[0] contract_count = len(contract_list) -if contract_count == 'Part-2': +if contract_count == "Part-2": contract_list = np.array_split(contract_list, 2)[1] contract_count = len(contract_list) seed_row = st.columns([0.15, 0.45, 0.4]) with seed_row[0]: - if contract_count in ['10', '20', '30', '50', '100', '200']: + if contract_count in ["10", "20", "30", "50", "100", "200"]: st.write("**Seed Value**") - elif contract_count == '1': + elif contract_count == "1": st.write("**Contract Name**") with seed_row[1]: - if contract_count in ['10', '20', '30', '50', '100', '200']: - seed_value = st.text_input("**Seed Value**", value=20, label_visibility="collapsed") + if contract_count in ["10", "20", "30", "50", "100", "200"]: + 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))) contract_list = sorted(random.choices(contract_list, k=int(contract_count))) - elif contract_count == '1': - contract_name = st.selectbox('Contract Name', (contract_list), label_visibility="collapsed") + elif contract_count == "1": + contract_name = st.selectbox( + "Contract Name", (contract_list), label_visibility="collapsed" + ) contract_list = [contract_name] 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 3 - Haiku', 'Claude 3 - Sonnet', 'Claude Instant' - , 'Llama 2 Chat 70B', 'Titan Text Express'), index=3, - label_visibility="collapsed") + llm_selected = st.selectbox( + "Langauge Model", + ( + "Claude 2", + "Claude 3 - Haiku", + "Claude 3 - Sonnet", + "Claude Instant", + "Llama 2 Chat 70B", + "Titan Text Express", + ), + index=3, + label_visibility="collapsed", + ) st.write("**Prompt**") -if mode != 'Single field - Non Empty values': +if mode != "Single field - Non Empty values": sequence_input = json.dumps(field_prompt_mapping) else: 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 = '' + 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=100, label_visibility="collapsed") + prompt = st.text_area( + "**Prompt**", sequence_input, height=100, label_visibility="collapsed" + ) # column_name = fields.loc[fields['Field Name'] == field, 'SF_DB_COL_NAME'].iloc[0] # column_list = ['Document_Name', column_name] @@ -287,11 +360,16 @@ with prompt_row[0]: # # field_values = field_values[column_list] # # field_values.rename(columns={'Document_Name': 'Contract Name', column_name: 'Actual Value Stored' # # , column_name+'_PG': 'Original Page Number'}, inplace=True) -field_values.rename(columns={'Document_Name': 'Contract Name'}, inplace=True) -if mode != 'One-to-many fields': - field_values = field_values.drop_duplicates(subset='Contract Name', keep="first").sort_values('Contract Name') +field_values.rename(columns={"Document_Name": "Contract Name"}, inplace=True) +if mode != "One-to-many fields": + field_values = field_values.drop_duplicates( + subset="Contract Name", keep="first" + ).sort_values("Contract Name") field_values = field_values[ - field_values['Contract Name'].isin([contract.rsplit('/', 1)[1] for contract in contract_list])] + field_values["Contract Name"].isin( + [contract.rsplit("/", 1)[1] for contract in contract_list] + ) +] # Setup bedrock bedrock_runtime = boto3.client( @@ -300,14 +378,22 @@ bedrock_runtime = boto3.client( ) # question = prompt -if mode != 'Single field - Non Empty values': +if mode != "Single field - Non Empty values": prompt_dict = json.loads(prompt) - prompt_dict_pg = prompt_dict | {str(k) + '_PG': "On which page can I find answer to the question - " + str( - v) for k, v in prompt_dict.items()} + prompt_dict_pg = prompt_dict | { + str(k) + "_PG": "On which page can I find answer to the question - " + str(v) + for k, v in prompt_dict.items() + } question = json.dumps(dict(sorted(prompt_dict_pg.items()))) else: question = json.dumps( - {field: prompt, str(field) + '_PG': "On which page can I find answer to the question - " + str(prompt)}) + { + field: prompt, + str(field) + + "_PG": "On which page can I find answer to the question - " + + str(prompt), + } + ) # st.write(question) question_with_schema = question @@ -317,8 +403,17 @@ attempt = 0 def run_llm(attempt, bucket, contract_list, llm_selected, field, field_values, history): # 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', 'New Extracted value', 'Confidence Level', 'Snippet', 'New Page Number' - , 'Revised Prompt', 'Result']) + df = pd.DataFrame( + columns=[ + "Contract Name", + "New Extracted value", + "Confidence Level", + "Snippet", + "New Page Number", + "Revised Prompt", + "Result", + ] + ) field_list = [] answer_list = [] @@ -330,8 +425,8 @@ def run_llm(attempt, bucket, contract_list, llm_selected, field, field_values, h for contract in contract_list: # with open(os.path.join(SOURCE_DIRECTORY, contract[:-4]+'.txt'), 'r') as infile: # context = infile.read() - data = s3_client.get_object(Bucket=bucket, Key=str(contract)[:-4] + '.txt') - contents = data['Body'].read() + data = s3_client.get_object(Bucket=bucket, Key=str(contract)[:-4] + ".txt") + contents = data["Body"].read() context = contents.decode("utf-8") # st.write(question_with_schema) @@ -350,13 +445,15 @@ def run_llm(attempt, bucket, contract_list, llm_selected, field, field_values, h "maxTokenCount": 2048, "stopSequences": [], "temperature": 0, - "topP": 0.9 + "topP": 0.9, } - body = json.dumps({"inputText": prompt_data, "textGenerationConfig": parameters}) + body = json.dumps( + {"inputText": prompt_data, "textGenerationConfig": parameters} + ) model_id = "amazon.titan-text-express-v1" # change this to use a different version from the model provider - elif llm_selected == 'Llama 2 Chat 70B': + elif llm_selected == "Llama 2 Chat 70B": context = context[:6000] prompt_data = f"""Answer the question based only on the information provided between ## and give step by step guide. You must answer in correct JSON format. ## @@ -369,13 +466,18 @@ def run_llm(attempt, bucket, contract_list, llm_selected, field, field_values, h "prompt": "[INST]" + prompt_data + "[/INST]", "max_gen_len": 2048, "temperature": 0.0, - "top_p": 0.9 + "top_p": 0.9, } body = json.dumps(payload) model_id = "meta.llama2-70b-chat-v1" - elif llm_selected in ['Claude Instant', 'Claude 2', 'Claude 3 - Haiku', 'Claude 3 - Sonnet']: - if llm_selected == 'Claude Instant': + elif llm_selected in [ + "Claude Instant", + "Claude 2", + "Claude 3 - Haiku", + "Claude 3 - Sonnet", + ]: + if llm_selected == "Claude Instant": context = context[:175000] prompt_data = f""" @@ -391,63 +493,74 @@ def run_llm(attempt, bucket, contract_list, llm_selected, field, field_values, h if llm_selected == "Claude 2": model_id = "anthropic.claude-v2:1" body = json.dumps( - {"prompt": anthropic.HUMAN_PROMPT + prompt_data + anthropic.AI_PROMPT, - "max_tokens_to_sample": 4096, - "temperature": 0.0, - "top_p": 1, - "top_k": 250, - "stop_sequences": [anthropic.HUMAN_PROMPT] - }) + { + "prompt": anthropic.HUMAN_PROMPT + + prompt_data + + anthropic.AI_PROMPT, + "max_tokens_to_sample": 4096, + "temperature": 0.0, + "top_p": 1, + "top_k": 250, + "stop_sequences": [anthropic.HUMAN_PROMPT], + } + ) elif llm_selected in ["Claude 3 - Haiku", "Claude 3 - Sonnet"]: if llm_selected == "Claude 3 - Haiku": - model_id = 'anthropic.claude-3-haiku-20240307-v1:0' + model_id = "anthropic.claude-3-haiku-20240307-v1:0" else: - model_id = 'anthropic.claude-3-sonnet-20240229-v1:0' - body = json.dumps({ - "anthropic_version": "bedrock-2023-05-31", - "max_tokens": 4096, - "messages": [ - { - "role": "user", - "content": [ - { - "type": "text", - "text": anthropic.HUMAN_PROMPT + prompt_data + anthropic.AI_PROMPT - } - ] - } - ], - "temperature": 0.0 - } + model_id = "anthropic.claude-3-sonnet-20240229-v1:0" + body = json.dumps( + { + "anthropic_version": "bedrock-2023-05-31", + "max_tokens": 4096, + "messages": [ + { + "role": "user", + "content": [ + { + "type": "text", + "text": anthropic.HUMAN_PROMPT + + prompt_data + + anthropic.AI_PROMPT, + } + ], + } + ], + "temperature": 0.0, + } ) else: model_id = "anthropic.claude-instant-v1" body = json.dumps( - {"prompt": anthropic.HUMAN_PROMPT + prompt_data + anthropic.AI_PROMPT, - "max_tokens_to_sample": 2048, - "temperature": 0.0, - "top_p": 1, - "top_k": 250, - "stop_sequences": [anthropic.HUMAN_PROMPT] - }) + { + "prompt": anthropic.HUMAN_PROMPT + + prompt_data + + anthropic.AI_PROMPT, + "max_tokens_to_sample": 2048, + "temperature": 0.0, + "top_p": 1, + "top_k": 250, + "stop_sequences": [anthropic.HUMAN_PROMPT], + } + ) try: response = bedrock_runtime.invoke_model( body=body, modelId=model_id, accept="application/json", - contentType="application/json" + contentType="application/json", ) response_body = json.loads(response.get("body").read()) if llm_selected == "Titan Text Express": response_text = response_body.get("results")[0].get("outputText") - elif llm_selected == 'Llama 2 Chat 70B': - response_text = response_body['generation'] - elif llm_selected in ['Claude Instant', 'Claude 2']: - response_text = response_body['completion'] - elif llm_selected in ['Claude 3 - Haiku', 'Claude 3 - Sonnet']: - response_text = response_body['content'][0]['text'] + elif llm_selected == "Llama 2 Chat 70B": + response_text = response_body["generation"] + elif llm_selected in ["Claude Instant", "Claude 2"]: + response_text = response_body["completion"] + elif llm_selected in ["Claude 3 - Haiku", "Claude 3 - Sonnet"]: + response_text = response_body["content"][0]["text"] except: response_text = "failed" @@ -473,7 +586,7 @@ def run_llm(attempt, bucket, contract_list, llm_selected, field, field_values, h try: response_dict = json.loads(response_text) except: - if mode != 'Single field - Non Empty values': + if mode != "Single field - Non Empty values": response_dict = {"Test field": "Failed to extract"} else: response_dict = {field: response_text.strip("{").strip("}")} @@ -491,11 +604,13 @@ def run_llm(attempt, bucket, contract_list, llm_selected, field, field_values, h # answer_l = list(response_dict)[:1] # except: # answer_l = [response_dict] - if mode != 'Single field - Non Empty values': + if mode != "Single field - Non Empty values": field_l = list(response_dict.keys()) answer_l = list(response_dict.values()) - field_dict = {k: v for k, v in response_dict.items() if not k.endswith('_PG')} - page_dict = {k: v for k, v in response_dict.items() if k.endswith('_PG')} + field_dict = { + k: v for k, v in response_dict.items() if not k.endswith("_PG") + } + page_dict = {k: v for k, v in response_dict.items() if k.endswith("_PG")} page_dict = {k[:-3]: v for k, v in response_dict.items()} field_l = list(field_dict.keys()) answer_l = list(field_dict.values()) @@ -506,7 +621,7 @@ def run_llm(attempt, bucket, contract_list, llm_selected, field, field_values, h try: page_no_l = [list(response_dict.values())[1]] except: - page_no_l = [''] + page_no_l = [""] field_list.extend(field_l) answer_list.extend(answer_l) @@ -518,14 +633,29 @@ def run_llm(attempt, bucket, contract_list, llm_selected, field, field_values, h # page_no = re.search(r'\d+', page_no).group() if page_no != " " and re.search(r'\d+', page_no) is not None else "" try: - location_l = [context.find(a, context.find("Start of Page No. = " + str(p))) if isinstance( - a, str) and a != "" else -1 for a, p in zip(answer_l, page_no_l)] + location_l = [ + ( + context.find(a, context.find("Start of Page No. = " + str(p))) + if isinstance(a, str) and a != "" + else -1 + ) + for a, p in zip(answer_l, page_no_l) + ] except: - location_l = [context.find(answer) if isinstance(answer, str) and answer != "" else -1 for answer in - answer_l] - snippet_l = [' '.join(context[:location].split('.')[-4:]) + ' ' + ' '.join(context[location:].split('. ')[:5] - ) if location != -1 else ' ' for - location in location_l] + location_l = [ + context.find(answer) if isinstance(answer, str) and answer != "" else -1 + for answer in answer_l + ] + snippet_l = [ + ( + " ".join(context[:location].split(".")[-4:]) + + " " + + " ".join(context[location:].split(". ")[:5]) + if location != -1 + else " " + ) + for location in location_l + ] # page_no_l = [" " if location == -1 else context[:location].rsplit("Start of Page No. = ", 1)[1] if len(context[:location].rsplit( # "Start of Page No. = ", 1)) > 1 else context[:location].rsplit("Start of Page No. = ", 1)[0] for location in location_l] # # st.write(location_list) @@ -533,8 +663,8 @@ def run_llm(attempt, bucket, contract_list, llm_selected, field, field_values, h snippet_list.extend(snippet_l) page_no_list.extend(page_no_l) - df['Field Name'] = field_list - df['Raw value'] = answer_list + df["Field Name"] = field_list + df["Raw value"] = answer_list # post-processing # if 'Date' in field: # date_list = [] @@ -547,32 +677,107 @@ def run_llm(attempt, bucket, contract_list, llm_selected, field, field_values, h # answer_list = date_list try: answer_list = [ - str(answer).strip("\n").strip().strip("[").strip("]").strip("{").strip("}").strip('"').rstrip('"').strip( - ' ') if answer is not None else None for answer in answer_list] - if llm_selected in ['Llama 2 Chat 13B', 'Llama 2 Chat 70B']: + ( + str(answer) + .strip("\n") + .strip() + .strip("[") + .strip("]") + .strip("{") + .strip("}") + .strip('"') + .rstrip('"') + .strip(" ") + if answer is not None + else None + ) + for answer in answer_list + ] + if 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 3 - Haiku', 'Claude 3 - Sonnet', 'Claude Instant']: - answer_list = [answer if "do not have" not in str(answer) else " " for answer in answer_list] - answer_list = [answer if "do not see" not in str(answer) else " " for answer in answer_list] - answer_list = [answer if "does not specify" not in str(answer) else " " for answer in answer_list] - answer_list = [answer if "Does not specify" not in str(answer) else " " for answer in answer_list] - answer_list = [answer if "does not explicitly" 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 "don't know" not in str(answer) else " " for answer in answer_list] - answer_list = [answer if "do not see" not in str(answer) else " " for answer in answer_list] - answer_list = [answer if "Not specified" not in str(answer) else " " for answer in answer_list] - answer_list = [answer if "Don't know" not in str(answer) else " " for answer in answer_list] - answer_list = [answer if "don't see" not in str(answer) else " " for answer in answer_list] - answer_list = [answer if "don't have" not in str(answer) else " " for answer in answer_list] - answer_list = [answer if "Does not apply" not in str(answer) else " " for answer in answer_list] - answer_list = [answer if "Nothing found" not in str(answer) else " " 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 3 - Haiku", + "Claude 3 - Sonnet", + "Claude Instant", + ]: + answer_list = [ + answer if "do not have" not in str(answer) else " " + for answer in answer_list + ] + answer_list = [ + answer if "do not see" not in str(answer) else " " + for answer in answer_list + ] + answer_list = [ + answer if "does not specify" not in str(answer) else " " + for answer in answer_list + ] + answer_list = [ + answer if "Does not specify" not in str(answer) else " " + for answer in answer_list + ] + answer_list = [ + answer if "does not explicitly" 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 "don't know" not in str(answer) else " " + for answer in answer_list + ] + answer_list = [ + answer if "do not see" not in str(answer) else " " + for answer in answer_list + ] + answer_list = [ + answer if "Not specified" not in str(answer) else " " + for answer in answer_list + ] + answer_list = [ + answer if "Don't know" not in str(answer) else " " + for answer in answer_list + ] + answer_list = [ + answer if "don't see" not in str(answer) else " " + for answer in answer_list + ] + answer_list = [ + answer if "don't have" not in str(answer) else " " + for answer in answer_list + ] + answer_list = [ + answer if "Does not apply" not in str(answer) else " " + for answer in answer_list + ] + answer_list = [ + answer if "Nothing found" 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] @@ -582,179 +787,373 @@ def run_llm(attempt, bucket, contract_list, llm_selected, field, field_values, h # to be deleted later # contract_list_f = [contract.rsplit('/',1)[1].replace(' MU','').replace('_MU','').replace('.txt','') for contract in contract_list_f] - contract_list_f = [contract.rsplit('/', 1)[1] for contract in contract_list_f] + contract_list_f = [contract.rsplit("/", 1)[1] for contract in contract_list_f] - df['Contract Name'] = contract_list_f - df['New Extracted value'] = answer_list - df['Confidence Level'] = ' ' - df['Snippet'] = snippet_list - df['New Page Number'] = page_no_list - df['New Page Number'] = df['New Page Number'].apply(lambda x: re.search(r'\d+', x).group( - ) if isinstance(x, str) and re.search(r'\d+', x) is not None else " ") - df['Revised Prompt'] = [prompt] * len(contract_list_f) + df["Contract Name"] = contract_list_f + df["New Extracted value"] = answer_list + df["Confidence Level"] = " " + df["Snippet"] = snippet_list + df["New Page Number"] = page_no_list + df["New Page Number"] = df["New Page Number"].apply( + lambda x: ( + re.search(r"\d+", x).group() + if isinstance(x, str) and re.search(r"\d+", x) is not None + else " " + ) + ) + df["Revised Prompt"] = [prompt] * len(contract_list_f) - df = pd.merge(df, fields[['Field Name', 'SF_DB_COL_NAME']], how='left', on='Field Name') + df = pd.merge( + df, fields[["Field Name", "SF_DB_COL_NAME"]], how="left", on="Field Name" + ) field_values_2 = pd.DataFrame( - columns=['Contract Name', 'LOB_Product_Network_Metal_Area_Program_Type_Speciality', 'SF_DB_COL_NAME' - , 'Actual Value Stored', 'Original Page Number']) + columns=[ + "Contract Name", + "LOB_Product_Network_Metal_Area_Program_Type_Speciality", + "SF_DB_COL_NAME", + "Actual Value Stored", + "Original Page Number", + ] + ) for file_name in contract_list: # document_name = [x for x in list(field_values['Contract Name']) if not pd.isna(x) and file_name.rsplit('/',1)[1].replace(' MU','').replace( # '_MU','').replace('.txt','') in x][0] - document_name = [x for x in list(field_values['Contract Name']) if not pd.isna(x) and file_name.rsplit('/',1)[1] in x][0] + document_name = [ + x + for x in list(field_values["Contract Name"]) + if not pd.isna(x) and file_name.rsplit("/", 1)[1] in x + ][0] - temp_df = field_values[field_values['Contract Name'] == document_name].fillna('NA') - unique_identifier = str(temp_df.at[temp_df.index[0], 'CONTRACT_LOB']) + '__' + str( - temp_df.at[temp_df.index[0], 'CONTRACT_PRODUCT']) + '__' + str( - temp_df.at[temp_df.index[0], 'CONTRACT_NETWORK']) + '__' + str( - temp_df.at[temp_df.index[0], 'CONTRACT_MARKETPLACE_METAL_LEVEL']) + '__' + str( - temp_df.at[temp_df.index[0], 'CONTRACT_SERVICE_AREA']) + '__' + str( - temp_df.at[temp_df.index[0], 'CONTRACT_PROGRAM']) + '__' + str( - temp_df.at[temp_df.index[0], 'PROV_TYPE']) + '__' + str(temp_df.at[temp_df.index[0], 'PROV_SPECIALTY']) - field_values_1 = field_values[field_values['Contract Name'] == document_name].head(1).transpose().reset_index() - field_values_1.columns = ['SF_DB_COL_NAME', 'Actual Value Stored'] - field_values_p1 = field_values_1[~field_values_1['SF_DB_COL_NAME'].str.endswith('_PG')] - field_values_p2 = field_values_1[field_values_1['SF_DB_COL_NAME'].str.endswith('_PG')] - field_values_p2.columns = ['SF_DB_COL_NAME', 'Original Page Number'] - field_values_p2["SF_DB_COL_NAME"] = field_values_p2["SF_DB_COL_NAME"].str.replace("_PG", "") + temp_df = field_values[field_values["Contract Name"] == document_name].fillna( + "NA" + ) + unique_identifier = ( + str(temp_df.at[temp_df.index[0], "CONTRACT_LOB"]) + + "__" + + str(temp_df.at[temp_df.index[0], "CONTRACT_PRODUCT"]) + + "__" + + str(temp_df.at[temp_df.index[0], "CONTRACT_NETWORK"]) + + "__" + + str(temp_df.at[temp_df.index[0], "CONTRACT_MARKETPLACE_METAL_LEVEL"]) + + "__" + + str(temp_df.at[temp_df.index[0], "CONTRACT_SERVICE_AREA"]) + + "__" + + str(temp_df.at[temp_df.index[0], "CONTRACT_PROGRAM"]) + + "__" + + str(temp_df.at[temp_df.index[0], "PROV_TYPE"]) + + "__" + + str(temp_df.at[temp_df.index[0], "PROV_SPECIALTY"]) + ) + field_values_1 = ( + field_values[field_values["Contract Name"] == document_name] + .head(1) + .transpose() + .reset_index() + ) + field_values_1.columns = ["SF_DB_COL_NAME", "Actual Value Stored"] + field_values_p1 = field_values_1[ + ~field_values_1["SF_DB_COL_NAME"].str.endswith("_PG") + ] + field_values_p2 = field_values_1[ + field_values_1["SF_DB_COL_NAME"].str.endswith("_PG") + ] + field_values_p2.columns = ["SF_DB_COL_NAME", "Original Page Number"] + field_values_p2["SF_DB_COL_NAME"] = field_values_p2[ + "SF_DB_COL_NAME" + ].str.replace("_PG", "") - field_values_1 = pd.merge(field_values_p1, field_values_p2, how='left', on=['SF_DB_COL_NAME']) - field_values_1['Contract Name'] = document_name - field_values_1['LOB_Product_Network_Metal_Area_Program_Type_Speciality'] = unique_identifier + field_values_1 = pd.merge( + field_values_p1, field_values_p2, how="left", on=["SF_DB_COL_NAME"] + ) + field_values_1["Contract Name"] = document_name + field_values_1["LOB_Product_Network_Metal_Area_Program_Type_Speciality"] = ( + unique_identifier + ) field_values_2 = pd.concat([field_values_2, field_values_1], ignore_index=True) - if mode == 'One-to-many fields': - for i in range(1, field_values[field_values['Contract Name'] == document_name].shape[0]): - unique_identifier = str(temp_df.at[temp_df.index[i], 'CONTRACT_LOB']) + '__' + str( - temp_df.at[temp_df.index[i], 'CONTRACT_PRODUCT']) + '__' + str( - temp_df.at[temp_df.index[i], 'CONTRACT_NETWORK']) + '__' + str( - temp_df.at[temp_df.index[i], 'CONTRACT_MARKETPLACE_METAL_LEVEL']) + '__' + str( - temp_df.at[temp_df.index[i], 'CONTRACT_SERVICE_AREA']) + '__' + str( - temp_df.at[temp_df.index[i], 'CONTRACT_PROGRAM']) + '__' + str( - temp_df.at[temp_df.index[i], 'PROV_TYPE']) + '__' + str( - temp_df.at[temp_df.index[i], 'PROV_SPECIALTY']) - field_values_1 = field_values[field_values['Contract Name'] == document_name].iloc[ - [i]].transpose().reset_index() - field_values_1.columns = ['SF_DB_COL_NAME', 'Actual Value Stored'] - field_values_p1 = field_values_1[~field_values_1['SF_DB_COL_NAME'].str.endswith('_PG')] - field_values_p2 = field_values_1[field_values_1['SF_DB_COL_NAME'].str.endswith('_PG')] - field_values_p2.columns = ['SF_DB_COL_NAME', 'Original Page Number'] - field_values_p2["SF_DB_COL_NAME"] = field_values_p2["SF_DB_COL_NAME"].str.replace("_PG", "") + if mode == "One-to-many fields": + for i in range( + 1, field_values[field_values["Contract Name"] == document_name].shape[0] + ): + unique_identifier = ( + str(temp_df.at[temp_df.index[i], "CONTRACT_LOB"]) + + "__" + + str(temp_df.at[temp_df.index[i], "CONTRACT_PRODUCT"]) + + "__" + + str(temp_df.at[temp_df.index[i], "CONTRACT_NETWORK"]) + + "__" + + str( + temp_df.at[temp_df.index[i], "CONTRACT_MARKETPLACE_METAL_LEVEL"] + ) + + "__" + + str(temp_df.at[temp_df.index[i], "CONTRACT_SERVICE_AREA"]) + + "__" + + str(temp_df.at[temp_df.index[i], "CONTRACT_PROGRAM"]) + + "__" + + str(temp_df.at[temp_df.index[i], "PROV_TYPE"]) + + "__" + + str(temp_df.at[temp_df.index[i], "PROV_SPECIALTY"]) + ) + field_values_1 = ( + field_values[field_values["Contract Name"] == document_name] + .iloc[[i]] + .transpose() + .reset_index() + ) + field_values_1.columns = ["SF_DB_COL_NAME", "Actual Value Stored"] + field_values_p1 = field_values_1[ + ~field_values_1["SF_DB_COL_NAME"].str.endswith("_PG") + ] + field_values_p2 = field_values_1[ + field_values_1["SF_DB_COL_NAME"].str.endswith("_PG") + ] + field_values_p2.columns = ["SF_DB_COL_NAME", "Original Page Number"] + field_values_p2["SF_DB_COL_NAME"] = field_values_p2[ + "SF_DB_COL_NAME" + ].str.replace("_PG", "") - field_values_1 = pd.merge(field_values_p1, field_values_p2, how='left', on=['SF_DB_COL_NAME']) - field_values_1['Contract Name'] = document_name - field_values_1['LOB_Product_Network_Metal_Area_Program_Type_Speciality'] = unique_identifier - field_values_2 = pd.concat([field_values_2, field_values_1], ignore_index=True) + field_values_1 = pd.merge( + field_values_p1, field_values_p2, how="left", on=["SF_DB_COL_NAME"] + ) + field_values_1["Contract Name"] = document_name + field_values_1[ + "LOB_Product_Network_Metal_Area_Program_Type_Speciality" + ] = unique_identifier + field_values_2 = pd.concat( + [field_values_2, field_values_1], ignore_index=True + ) # st.write(df) # st.write(field_values_2) - df = pd.merge(df, field_values_2, how='right', on=['Contract Name', 'SF_DB_COL_NAME']) - if mode != 'Single field - Non Empty values': - df = df[df['SF_DB_COL_NAME'].isin(list(field_prompt_mapping.keys()))] + df = pd.merge( + df, field_values_2, how="right", on=["Contract Name", "SF_DB_COL_NAME"] + ) + if mode != "Single field - Non Empty values": + df = df[df["SF_DB_COL_NAME"].isin(list(field_prompt_mapping.keys()))] else: - df = df[df['SF_DB_COL_NAME'].isin([field])] - if mode != 'One-to-many fields': - df = df.drop_duplicates(subset=['SF_DB_COL_NAME', 'Contract Name'], keep="first") - df['Original Page Number'] = df['Original Page Number'].apply(lambda x: re.search(r'\d+', x).group( - ) if isinstance(x, str) and re.search(r'\d+', x) is not None else " ") - df['Raw value 2'] = df['New Extracted value'] - df_date = df[df['SF_DB_COL_NAME'].str.contains('_DT', na=False)] - df_others = df[~df['SF_DB_COL_NAME'].str.contains('_DT', na=False)] + df = df[df["SF_DB_COL_NAME"].isin([field])] + if mode != "One-to-many fields": + df = df.drop_duplicates( + subset=["SF_DB_COL_NAME", "Contract Name"], keep="first" + ) + df["Original Page Number"] = df["Original Page Number"].apply( + lambda x: ( + re.search(r"\d+", x).group() + if isinstance(x, str) and re.search(r"\d+", x) is not None + else " " + ) + ) + df["Raw value 2"] = df["New Extracted value"] + df_date = df[df["SF_DB_COL_NAME"].str.contains("_DT", na=False)] + df_others = df[~df["SF_DB_COL_NAME"].str.contains("_DT", na=False)] - df_date['Actual Value Stored'] = pd.to_datetime(df_date['Actual Value Stored'], errors='coerce').dt.strftime( - '%Y-%m-%d').fillna(" ") - df_date['New Extracted value'] = pd.to_datetime(df_date['New Extracted value'], errors='coerce').dt.strftime( - '%Y-%m-%d').fillna(" ") + df_date["Actual Value Stored"] = ( + pd.to_datetime(df_date["Actual Value Stored"], errors="coerce") + .dt.strftime("%Y-%m-%d") + .fillna(" ") + ) + df_date["New Extracted value"] = ( + pd.to_datetime(df_date["New Extracted value"], errors="coerce") + .dt.strftime("%Y-%m-%d") + .fillna(" ") + ) df = pd.concat([df_date, df_others], ignore_index=True) - df.sort_values(['SF_DB_COL_NAME', 'Contract Name'], inplace=True) + df.sort_values(["SF_DB_COL_NAME", "Contract Name"], inplace=True) df.fillna(" ", inplace=True) - df['Raw value 3'] = df['New Extracted value'] - df['Actual Value Stored'] = df['Actual Value Stored'].apply(lambda x: x.strip() if isinstance(x, str) else '') - df['New Extracted value'] = df['New Extracted value'].apply(lambda x: x.strip() if isinstance(x, str) else '') - actual_value_list = list(df['Actual Value Stored']) - actual_value_list = [answer if str(answer) != "12 months" else "1 year" for answer in actual_value_list] - actual_value_list = [answer if str(answer) != "Fifth" else "5" for answer in actual_value_list] - actual_value_list = [answer if str(answer) != "Seventh" else "7" for answer in actual_value_list] + df["Raw value 3"] = df["New Extracted value"] + df["Actual Value Stored"] = df["Actual Value Stored"].apply( + lambda x: x.strip() if isinstance(x, str) else "" + ) + df["New Extracted value"] = df["New Extracted value"].apply( + lambda x: x.strip() if isinstance(x, str) else "" + ) + actual_value_list = list(df["Actual Value Stored"]) + actual_value_list = [ + answer if str(answer) != "12 months" else "1 year" + for answer in actual_value_list + ] + actual_value_list = [ + answer if str(answer) != "Fifth" else "5" for answer in actual_value_list + ] + actual_value_list = [ + answer if str(answer) != "Seventh" else "7" for answer in actual_value_list + ] - answer_list = list(df['New Extracted value']) - answer_list = [answer if str(answer) != "one-year" else "1 year" for answer in answer_list] - answer_list = [answer if str(answer) != "one year" else "1 year" for answer in answer_list] - answer_list = [answer if str(answer) != "one" else "1 year" for answer in answer_list] - answer_list = [answer if str(answer) != "one (1) year" else "1 year" for answer in answer_list] - answer_list = [answer if str(answer) != "twelve" else "1 year" for answer in answer_list] + answer_list = list(df["New Extracted value"]) + answer_list = [ + answer if str(answer) != "one-year" else "1 year" for answer in answer_list + ] + answer_list = [ + answer if str(answer) != "one year" else "1 year" for answer in answer_list + ] + answer_list = [ + answer if str(answer) != "one" else "1 year" for answer in answer_list + ] + answer_list = [ + answer if str(answer) != "one (1) year" else "1 year" for answer in answer_list + ] + answer_list = [ + answer if str(answer) != "twelve" else "1 year" for answer in answer_list + ] answer_list = [answer if str(answer) != "XI" else "11" for answer in answer_list] answer_list = [answer if str(answer) != "Third" else "3" for answer in answer_list] answer_list = [answer if str(answer) != "Six" else "6" for answer in answer_list] - actual_value_list = [s.replace('-', '').replace(' ', '').replace('[', '').replace(']', '').lower() for s in - actual_value_list] - answer_list = [s.replace('-', '').replace(' ', '').replace('[', '').replace(']', '').lower() for s in answer_list] - result_list = [(i in j) or (j in i) if isinstance(i, str) and isinstance( - j, str) and ((i != '') == (j != '')) else False for i, j in zip(actual_value_list, answer_list)] - df['Result'] = [str(x) for x in result_list] + actual_value_list = [ + s.replace("-", "").replace(" ", "").replace("[", "").replace("]", "").lower() + for s in actual_value_list + ] + answer_list = [ + s.replace("-", "").replace(" ", "").replace("[", "").replace("]", "").lower() + for s in answer_list + ] + result_list = [ + ( + (i in j) or (j in i) + if isinstance(i, str) and isinstance(j, str) and ((i != "") == (j != "")) + else False + ) + 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()] # df['Contract ID'] = contract_list_f - df['Contract ID'] = range(len(actual_value_list)) + df["Contract ID"] = range(len(actual_value_list)) - df = df[['Contract Name', 'Contract ID', 'LOB_Product_Network_Metal_Area_Program_Type_Speciality', 'SF_DB_COL_NAME', - 'Actual Value Stored', 'Raw value', 'Raw value 2', 'Raw value 3' - , 'New Extracted value', 'Confidence Level', 'Snippet', 'Original Page Number', 'New Page Number', - 'Revised Prompt', 'Result']] + df = df[ + [ + "Contract Name", + "Contract ID", + "LOB_Product_Network_Metal_Area_Program_Type_Speciality", + "SF_DB_COL_NAME", + "Actual Value Stored", + "Raw value", + "Raw value 2", + "Raw value 3", + "New Extracted value", + "Confidence Level", + "Snippet", + "Original Page Number", + "New Page Number", + "Revised Prompt", + "Result", + ] + ] try: - accuracy = round(sum(bool(x) for x in result_list) * 100 / len(list(df['Result'])), 2) + accuracy = round( + sum(bool(x) for x in result_list) * 100 / len(list(df["Result"])), 2 + ) except: - accuracy = 'NA' + accuracy = "NA" - if mode != 'Single field - Non Empty values': + if mode != "Single field - Non Empty values": field = field_group - history.loc[len(history.index)] = [field, str(contract_count), st.session_state.user_info['mail'], - datetime.now().strftime("%Y-%m-%d %H:%M:%S"), accuracy, attempt] + history.loc[len(history.index)] = [ + field, + str(contract_count), + st.session_state.user_info["mail"], + datetime.now().strftime("%Y-%m-%d %H:%M:%S"), + accuracy, + attempt, + ] # df.to_csv("RESULTS\\"+field.replace("?","").replace("/","_")+'-'+llm_selected+'.csv', index=False) return df, history, attempt, raw_response_text -raw_response_text = '' +raw_response_text = "" df = pd.DataFrame( - columns=['Contract Name', 'Contract ID', 'Actual Value Stored', 'New Extracted value', 'Confidence Level' - , 'Snippet', 'New Page Number', 'Revised Prompt', 'Result']) + columns=[ + "Contract Name", + "Contract ID", + "Actual Value Stored", + "New Extracted value", + "Confidence Level", + "Snippet", + "New Page Number", + "Revised Prompt", + "Result", + ] +) if st.button("Test Configuration"): - df, history, attempt, raw_response_text = run_llm(attempt, bucket, contract_list, llm_selected, field, field_values, - history) - df_1 = df[['Contract Name', 'Contract ID', 'Actual Value Stored', 'New Extracted value', 'Confidence Level' - , 'Snippet', 'Original Page Number', 'New Page Number', 'Revised Prompt', 'Result']] - df.to_csv('results.csv', index=False) - history.to_csv('history.csv', index=False) + df, history, attempt, raw_response_text = run_llm( + attempt, bucket, contract_list, llm_selected, field, field_values, history + ) + df_1 = df[ + [ + "Contract Name", + "Contract ID", + "Actual Value Stored", + "New Extracted value", + "Confidence Level", + "Snippet", + "Original Page Number", + "New Page Number", + "Revised Prompt", + "Result", + ] + ] + df.to_csv("results.csv", index=False) + history.to_csv("history.csv", index=False) # s3_client.upload_file('results.csv', bucket, key) csv_buf = StringIO() df_1.to_csv(csv_buf, header=True, index=False) csv_buf.seek(0) - s3_client.put_object(Bucket='doczy-dev-infra-raw-data-ingestion', Body=csv_buf.getvalue(), - Key='training_interface/results.csv') + s3_client.put_object( + Bucket="doczy-dev-infra-raw-data-ingestion", + Body=csv_buf.getvalue(), + Key="training_interface/results.csv", + ) csv_buf = StringIO() history.tail(1).to_csv(csv_buf, header=True, index=False) csv_buf.seek(0) - s3_client.put_object(Bucket='doczy-dev-infra-raw-data-ingestion', Body=csv_buf.getvalue(), - Key='training_interface/history.csv') - df = df[['Contract Name', 'Contract ID', 'LOB_Product_Network_Metal_Area_Program_Type_Speciality', 'SF_DB_COL_NAME', - 'Actual Value Stored', 'New Extracted value', 'Confidence Level' - , 'Snippet', 'Original Page Number', 'New Page Number', 'Revised Prompt', 'Result']] + s3_client.put_object( + Bucket="doczy-dev-infra-raw-data-ingestion", + Body=csv_buf.getvalue(), + Key="training_interface/history.csv", + ) + df = df[ + [ + "Contract Name", + "Contract ID", + "LOB_Product_Network_Metal_Area_Program_Type_Speciality", + "SF_DB_COL_NAME", + "Actual Value Stored", + "New Extracted value", + "Confidence Level", + "Snippet", + "Original Page Number", + "New Page Number", + "Revised Prompt", + "Result", + ] + ] try: - df = pd.read_csv('results.csv') - df['Result'] = df['Result'].astype('str') + df = pd.read_csv("results.csv") + df["Result"] = df["Result"].astype("str") except: - df = pd.DataFrame(columns=['Contract Name', 'New Extracted value', 'Confidence Level', 'Snippet', 'New Page Number' - , 'Revised Prompt', 'Result']) -history = pd.read_csv('history.csv') + df = pd.DataFrame( + columns=[ + "Contract Name", + "New Extracted value", + "Confidence Level", + "Snippet", + "New Page Number", + "Revised Prompt", + "Result", + ] + ) +history = pd.read_csv("history.csv") st.dataframe(df) st.dataframe(history) -# @st.cache_data +# @st.cache_data # def convert_df(df): # return df.to_csv(index=False).encode('utf-8') @@ -762,11 +1161,11 @@ st.dataframe(history) # buttons = st.columns(3) # with buttons[0]: -# st.button("Save All Imputations") +# 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.button("Kickoff Database Integration") add_vertical_space(20) st.write(field) @@ -777,6 +1176,10 @@ st.write(len(contract_list)) st.write(raw_response_text) try: - save_to_sf('load_training_results', training_results_file_name="results.csv", attempt_logs_file_name="history.csv") + save_to_sf( + "load_training_results", + training_results_file_name="results.csv", + attempt_logs_file_name="history.csv", + ) except: st.write("running locally") diff --git a/streamlit/interface_3_rag.py b/streamlit/interface_3_rag.py index eda97c5..1a6a0ee 100644 --- a/streamlit/interface_3_rag.py +++ b/streamlit/interface_3_rag.py @@ -4,7 +4,15 @@ from langchain.prompts import PromptTemplate from langchain.embeddings.bedrock import BedrockEmbeddings from langchain.llms.bedrock import Bedrock from langchain_community.vectorstores import Chroma -from constants import CHROMA_SETTINGS, EMBEDDING_MODEL_NAME, PERSIST_DIRECTORY, MODEL_ID, MODEL_BASENAME, SOURCE_DIRECTORY, USER_LIST +from constants import ( + CHROMA_SETTINGS, + EMBEDDING_MODEL_NAME, + PERSIST_DIRECTORY, + MODEL_ID, + MODEL_BASENAME, + SOURCE_DIRECTORY, + USER_LIST, +) from langchain.chains import RetrievalQA import streamlit as st @@ -17,10 +25,10 @@ import os import dateutil import util -REDIRECT_URI = 'http://172.29.20.126:8503' +REDIRECT_URI = "http://172.29.20.126:8503" user_list = USER_LIST -st.set_page_config(layout = "wide") +st.set_page_config(layout="wide") # Sidebar contents with st.sidebar: st.title("Doczy.AI ™") @@ -35,91 +43,140 @@ with st.sidebar: # st.write("Doczy") # util.setup_page(REDIRECT_URI) -# if st.session_state.user_info['mail'] in user_list: -if 'maamseek@aarete.com' in user_list: +# if st.session_state.user_info['mail'] in user_list: +if "maamseek@aarete.com" in user_list: - 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?'])) + 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?"]) + ) 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") + field = st.selectbox( + "Field Name", + sorted(set(field_prompt_mapping.keys())), + index=0, + 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") + contract_count = st.selectbox( + "Contract count", + ("1", "10", "20", "30", "50", "All"), + index=1, + label_visibility="collapsed", + ) seed_row = st.columns([0.15, 0.45, 0.4]) - 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["(internal) Document Name"]) + ] with seed_row[0]: - if contract_count in ['10', '20', '30', '50']: + if contract_count in ["10", "20", "30", "50"]: st.write("**Seed Value**") - elif contract_count == '1': + 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") + 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 = 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] - 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") + 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 = '' + 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") + 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: + 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_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') + 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') + 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( @@ -150,7 +207,7 @@ if 'maamseek@aarete.com' in user_list: # st.session_state.RETRIEVER = RETRIEVER # if "LLM" not in st.session_state: - if llm_selected == 'Titan Text Express': + if llm_selected == "Titan Text Express": LLM = Bedrock( model_id="amazon.titan-text-express-v1", client=bedrock_runtime, @@ -159,9 +216,9 @@ if 'maamseek@aarete.com' in user_list: "stopSequences": [], "temperature": 0, "topP": 1, - } + }, ) - elif llm_selected == 'Llama 2 Chat 70B': + elif llm_selected == "Llama 2 Chat 70B": LLM = Bedrock( model_id="meta.llama2-70b-chat-v1", client=bedrock_runtime, @@ -169,9 +226,9 @@ if 'maamseek@aarete.com' in user_list: "max_gen_len": 512, "temperature": 0, # "topP": 0.9, - } + }, ) - elif llm_selected == 'Llama 2 Chat 13B': + elif llm_selected == "Llama 2 Chat 13B": LLM = Bedrock( model_id="meta.llama2-13b-chat-v1", client=bedrock_runtime, @@ -179,9 +236,9 @@ if 'maamseek@aarete.com' in user_list: "max_gen_len": 512, "temperature": 0, # "topP": 0.9, - } + }, ) - elif llm_selected == 'Claude Instant': + elif llm_selected == "Claude Instant": LLM = Bedrock( model_id="anthropic.claude-instant-v1", client=bedrock_runtime, @@ -189,9 +246,9 @@ if 'maamseek@aarete.com' in user_list: # "max_tokens_to_sample": 512, "temperature": 0, # "topP": 0.9, - } + }, ) - elif llm_selected == 'Claude 2': + elif llm_selected == "Claude 2": LLM = Bedrock( model_id="anthropic.claude-v2:1", client=bedrock_runtime, @@ -199,7 +256,7 @@ if 'maamseek@aarete.com' in user_list: # "max_tokens_to_sample": 512, "temperature": 0, # "topP": 0.9, - } + }, ) st.session_state["LLM"] = LLM @@ -217,21 +274,47 @@ if 'maamseek@aarete.com' in user_list: # 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']) + 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: - history = pd.read_csv('history.csv') + history = pd.read_csv("history.csv") except: - history = pd.DataFrame(columns=['Field Name','# Contracts Tested', 'Username', 'Date/Time', 'Accuracy', 'Attempt #']) + 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']: + if llm_selected in ["Llama 2 Chat 13B", "Llama 2 Chat 70B"]: k_value = 10 else: k_value = 20 - if st.button("Test Configuration"): answer_list = [] @@ -241,7 +324,9 @@ if 'maamseek@aarete.com' in user_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}) + 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", @@ -249,7 +334,9 @@ if 'maamseek@aarete.com' in user_list: 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 = 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"] @@ -257,95 +344,178 @@ if 'maamseek@aarete.com' in user_list: doc_list.append(docs) response_list.append(response) - df['Raw value'] = answer_list + df["Raw value"] = answer_list # post-processing - if 'Date' in field: + if "Date" in field: date_list = [] for answer in answer_list: try: - extracted_date = dateutil.parser.parse(str(answer).replace('"',''), fuzzy=True).date() + 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']: + 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 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 + 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["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] + 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): + 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["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 = 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']] + 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']] + 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) + accuracy = round( + sum(bool(x) for x in result_list) * 100 / len(list(df["Result"])), 2 + ) except: - accuracy = 'NA' + accuracy = "NA" - history.loc[len(history.index)] = [field, str(contract_count), None, datetime.now().strftime("%Y-%m-%d %H:%M:%S"), accuracy, attempt] + 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) + 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.cache_data + # @st.cache_data # def convert_df(df): # return df.to_csv(index=False).encode('utf-8') @@ -353,17 +523,14 @@ if 'maamseek@aarete.com' in user_list: # buttons = st.columns(3) # with buttons[0]: - # st.button("Save All Imputations") + # 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.button("Kickoff Database Integration") st.write(column_name) st.write(len(contract_list)) else: st.write("Access Denied") - - diff --git a/streamlit/local/local_interface_0.py b/streamlit/local/local_interface_0.py index 1c6a006..dc13547 100644 --- a/streamlit/local/local_interface_0.py +++ b/streamlit/local/local_interface_0.py @@ -7,35 +7,39 @@ import pandas as pd from io import StringIO from datetime import datetime import boto3 + # import util import requests + # from sf_conn import get_secret, save_to_sf from io import StringIO, BytesIO import time + # from sf_conn import get_client_names, insert_upload_logs # from constants import USER_LIST -create_batch_url = 'https://lfksus2t62.execute-api.us-east-2.amazonaws.com/dev/create-batch' +create_batch_url = ( + "https://lfksus2t62.execute-api.us-east-2.amazonaws.com/dev/create-batch" +) # REDIRECT_URI = 'https://doczydev.aarete.com:8500' # user_list = USER_LIST -if 'uploading' not in st.session_state: +if "uploading" not in st.session_state: st.session_state.uploading = False - def insert_upload_logs(batch_id, client_name, file_name, upload_datetime, upload_user): """ Input: batch_id, client_name, file_name, upload_datetime, upload_user Output: status of the insert query """ try: - return 'Log inserted successfully' + return "Log inserted successfully" except Exception as e: return e -st.set_page_config(layout = "wide") +st.set_page_config(layout="wide") # # # Sidebar contents # # with st.sidebar: # # st.title("Doczy.AI ™") @@ -43,38 +47,52 @@ st.set_page_config(layout = "wide") # # """ # # ## About # # This app extracts data from contracts - + # # """ # # ) # # add_vertical_space(15) # # # st.write("Doczy") # # -_,c1= st.columns([5,1]) +_, c1 = st.columns([5, 1]) # try: # util.setup_page(REDIRECT_URI) # except Exception as e: # st.write(f"SSO Failed = {e}") # st.session_state['user_info'] = {'mail': 'maamseek@aarete.com', 'displayName': 'Mayank Aamseek'} -st.session_state['user_info'] = {'mail': 'maamseek@aarete.com', 'displayName': 'Mayank Aamseek'} +st.session_state["user_info"] = { + "mail": "maamseek@aarete.com", + "displayName": "Mayank Aamseek", +} try: c1.write(f"User: **{st.session_state.user_info['displayName']}**") - user_mail = st.session_state.user_info['mail'] + user_mail = st.session_state.user_info["mail"] except KeyError as e: st.write("Session Expired.") st.stop() print(st.session_state) -s3_client = boto3.client('s3', +s3_client = boto3.client( + "s3", region_name="us-east-2", ) -client_list = ['doczy-ai-client-1', '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'] +client_list = [ + "doczy-ai-client-1", + "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", +] -# This is the list of client fetched from Snowflake +# This is the list of client fetched from Snowflake # TODO: Need to update the streamlit code to use the client names from this list # And use the s3 paths to save the objects for the respective client # client_list, s3_paths = get_client_names() @@ -84,31 +102,42 @@ 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',(client_list), label_visibility = "collapsed", index = None) # MODIFIED for Ticket DOC-344 + client = st.selectbox( + "Client Name", (client_list), label_visibility="collapsed", index=None + ) # MODIFIED for Ticket DOC-344 # client_bucket = client_s3_paths.get(client) # to be deleted when client buckets are created -client_bucket = 'doczy-ai-client-1' +client_bucket = "doczy-ai-client-1" file_row = st.columns([0.1, 0.8]) with file_row[0]: st.write("**Upload Files**") with file_row[1]: - file_list = st.file_uploader("Upload", type=['docx','tiff','pdf'], accept_multiple_files=True, label_visibility = "collapsed", help="Only PDF, TIFF and DOCX file formats are supported.", disabled=st.session_state.uploading) + file_list = st.file_uploader( + "Upload", + type=["docx", "tiff", "pdf"], + accept_multiple_files=True, + label_visibility="collapsed", + help="Only PDF, TIFF and DOCX file formats are supported.", + disabled=st.session_state.uploading, + ) add_vertical_space(2) -df = pd.DataFrame(columns=['Contract Name']) -df['Contract Name'] = file_list +df = pd.DataFrame(columns=["Contract Name"]) +df["Contract Name"] = file_list file_names = [] buttons = st.columns([0.4, 0.4, 0.2]) + def set_uploading_state(): if not client == None and not len(file_list) == 0: st.session_state.uploading = True + with buttons[1]: - if st.button("Create Batch", on_click = set_uploading_state): + if st.button("Create Batch", on_click=set_uploading_state): if client == None: st.error("No Client Name Selected.") elif len(file_list) == 0: @@ -133,20 +162,17 @@ with buttons[1]: # stringio.seek(0) # s3_client.put_object(Bucket=client_bucket, Body=stringio.getvalue(), Key= # landing_zone+batch_id+'/'+str(uploaded_file.name)) - - # # TODO: Test this insert function with snowflake + + # # TODO: Test this insert function with snowflake # upload_log = insert_upload_logs(batch_id, client, str(uploaded_file.name), datetime.now().strftime("%Y-%m-%d %H:%M:%S"), user_mail) # st.write(upload_log) - + # file_names.append(str(uploaded_file.name)) # st.write(f"{batch_id} created") st.session_state.uploading = False # st.write(f"Files uploaded to s3://{client_bucket}/{landing_zone}{batch_id}") - - - - # @st.cache_data + # @st.cache_data # def convert_df(df): # return df.to_csv(index=False).encode('utf-8') diff --git a/streamlit/local/local_interface_1.py b/streamlit/local/local_interface_1.py index fd03888..cb5accb 100644 --- a/streamlit/local/local_interface_1.py +++ b/streamlit/local/local_interface_1.py @@ -6,9 +6,11 @@ import pandas as pd from io import StringIO from datetime import datetime import boto3 + # import util import requests import time + # from sf_conn import get_client_names, get_secret, save_to_sf # from constants import USER_LIST, DOCZY_PIPELINE_URL_DEV @@ -16,7 +18,7 @@ import time # REDIRECT_URI = 'https://doczydev.aarete.com:8501' # user_list = USER_LIST -st.set_page_config(layout = "wide") +st.set_page_config(layout="wide") # # # Sidebar contents # # with st.sidebar: # # st.title("Doczy.AI ™") @@ -24,37 +26,50 @@ st.set_page_config(layout = "wide") # # """ # # ## About # # This app extracts data from contracts - + # # """ # # ) # # add_vertical_space(15) # # # st.write("Doczy") -_,c1= st.columns([5,1]) +_, c1 = st.columns([5, 1]) # try: # util.setup_page(REDIRECT_URI) # except: # st.write("SSO Failed") # st.session_state['user_info'] = {'mail': 'maamseek@aarete.com', 'displayName': 'Mayank Aamseek'} -st.session_state['user_info'] = {'mail': 'maamseek@aarete.com', 'displayName': 'Mayank Aamseek'} +st.session_state["user_info"] = { + "mail": "maamseek@aarete.com", + "displayName": "Mayank Aamseek", +} try: c1.write(f"User: **{st.session_state.user_info['displayName']}**") - user_mail = st.session_state.user_info['mail'] + user_mail = st.session_state.user_info["mail"] except KeyError as e: - # Do we add a link to get to the login page here? + # Do we add a link to get to the login page here? st.write("Session Expired.") st.stop() - -s3_client = boto3.client('s3', +s3_client = boto3.client( + "s3", region_name="us-east-2", ) # # to be replaced with snowflake data -client_list = ['doczy-ai-client-1', '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'] +client_list = [ + "doczy-ai-client-1", + "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", +] # client_list, s3_paths = get_client_names() # client_s3_paths = dict(zip(client_list, s3_paths)) @@ -63,54 +78,79 @@ 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',(client_list), label_visibility = "collapsed", index = None) # MODIFIED for Ticket DOC-344 + client = st.selectbox( + "Client Name", (client_list), label_visibility="collapsed", index=None + ) # MODIFIED for Ticket DOC-344 # client_bucket = client_s3_paths.get(client) # # to be deleted when buckets for different clients are ready; below line is added only for testing the corresponding DAG -client_bucket = 'doczy-ai-client-1' +client_bucket = "doczy-ai-client-1" # batch_objects = s3_client.list_objects_v2(Bucket=client_bucket -# , Prefix="contracts_landing_zone/", Delimiter='/') +# , Prefix="contracts_landing_zone/", Delimiter='/') # batch_list = [] # for prefix in batch_objects['CommonPrefixes']: # batch_list.append(prefix['Prefix'][:-1].split('/')[-1]) -batch_list = ['batch_020524103737', 'batch_090524131433', 'batch_090524131607', 'batch_100524123000', 'batch_130524064322', - 'batch_160524071331', 'batch_200524213550', 'batch_250424112237', 'batch_280524120530', 'batch_280524121721', 'batch_280524144222', - 'batch_290524123926', 'batch_290524164044', 'batch_310524102029', 'batch_310524124050', 'batch_310524162346', 'batch_310524162631'] +batch_list = [ + "batch_020524103737", + "batch_090524131433", + "batch_090524131607", + "batch_100524123000", + "batch_130524064322", + "batch_160524071331", + "batch_200524213550", + "batch_250424112237", + "batch_280524120530", + "batch_280524121721", + "batch_280524144222", + "batch_290524123926", + "batch_290524164044", + "batch_310524102029", + "batch_310524124050", + "batch_310524162346", + "batch_310524162631", +] -if 'sorted_list' not in st.session_state: +if "sorted_list" not in st.session_state: st.session_state.sorted_list = batch_list + def sort_list(ex_list, sort_by, order): - if sort_by == 'Alphabetical': - ex_list = sorted(ex_list, reverse=(order == 'Descending')) - elif sort_by == 'Create Date': - ex_list = ex_list if order == 'Ascending' else list(reversed(ex_list)) + if sort_by == "Alphabetical": + ex_list = sorted(ex_list, reverse=(order == "Descending")) + elif sort_by == "Create Date": + ex_list = ex_list if order == "Ascending" else list(reversed(ex_list)) return ex_list + col1, col2, col3, col4 = st.columns([0.5, 0.5, 0.5, 0.5]) -with col1: - sort_by = st.radio("**Sort Batch_IDs**", ('Alphabetical', 'Create Date')) +with col1: + sort_by = st.radio("**Sort Batch_IDs**", ("Alphabetical", "Create Date")) with col2: - order = st.radio('', ('Ascending','Descending')) + order = st.radio("", ("Ascending", "Descending")) -with col3: +with col3: add_vertical_space(2) - if st.button('Apply'): + if st.button("Apply"): st.session_state.sorted_list = sort_list(batch_list, sort_by, order) path_row = st.columns([0.1, 0.8]) with path_row[0]: st.write("**Batch ID**") with path_row[1]: - batch_id = st.selectbox('**Batch ID**', st.session_state.sorted_list, label_visibility = "collapsed", index = None) # MODIFIED for Ticket DOC-344 + batch_id = st.selectbox( + "**Batch ID**", + st.session_state.sorted_list, + label_visibility="collapsed", + index=None, + ) # MODIFIED for Ticket DOC-344 if not batch_id: batch_id = "None" @@ -120,68 +160,91 @@ with checks[0]: st.write("**Group No.**") with checks[1]: - a = st.checkbox('Unique Key', key = str(1), args="Unique") + a = st.checkbox("Unique Key", key=str(1), args="Unique") with checks[2]: - b = st.checkbox('Pricing Before Carveouts', key = str(2)) + b = st.checkbox("Pricing Before Carveouts", key=str(2)) with checks[3]: - c = st.checkbox('Contract Related', key = str(3)) + c = st.checkbox("Contract Related", key=str(3)) with checks[4]: - d = st.checkbox('Provider', key = str(4)) + d = st.checkbox("Provider", key=str(4)) with checks[5]: - e = st.checkbox('Timeline', key = str(5)) + e = st.checkbox("Timeline", key=str(5)) with checks[6]: - f = st.checkbox('Carveout Indicator', key = str(6)) + f = st.checkbox("Carveout Indicator", key=str(6)) with checks[7]: - g = st.checkbox('Carveout Methodology', key = str(7)) + g = st.checkbox("Carveout Methodology", key=str(7)) add_vertical_space(1) -df = pd.DataFrame(columns=['Contract Name', 'Unique Key','Pricing Before Carveouts' - , 'Contract Related', 'Provider', 'Timeline', 'Carveout Indicator', 'Carveout Methodology']) +df = pd.DataFrame( + columns=[ + "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=client_bucket # , Prefix="contracts_landing_zone/"+batch_id+"/", Delimiter='/') -file_list = ['Boilerplate_TX Amendment Mission Health Network effective_040114 MU.pdf', 'Custom_TX - MP AMENDMENT - MISSION HEALTH NETWORK - MU.pdf', - 'Delaware First Health_First State Homecare Agency_212260_7 MU.pdf', 'Molina Healthcare of Texas, Inc. Amendment 4 - HIX ACA__EFF 01012016_MU.pdf'] +file_list = [ + "Boilerplate_TX Amendment Mission Health Network effective_040114 MU.pdf", + "Custom_TX - MP AMENDMENT - MISSION HEALTH NETWORK - MU.pdf", + "Delaware First Health_First State Homecare Agency_212260_7 MU.pdf", + "Molina Healthcare of Texas, Inc. Amendment 4 - HIX ACA__EFF 01012016_MU.pdf", +] 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["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 + 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) + print(f"DEBUGGING: PWD= {dir_path}") + df.to_csv("temp1.csv", index=False) add_vertical_space(1) -df2 = pd.read_csv('temp1.csv') +df2 = pd.read_csv("temp1.csv") edited_df = st.data_editor(df2) -edited_df['REQUEST_USER'] = user_mail -edited_df['LATEST_FLAG BOOLEAN'] = True -edited_df['PIPELINE_KICKOFF_DATETIME'] = datetime.now().strftime("%Y-%m-%d %H:%M:%S") -edited_df['REQUEST_DATETIME'] = datetime.now().strftime("%Y-%m-%d %H:%M:%S") +edited_df["REQUEST_USER"] = user_mail +edited_df["LATEST_FLAG BOOLEAN"] = True +edited_df["PIPELINE_KICKOFF_DATETIME"] = datetime.now().strftime("%Y-%m-%d %H:%M:%S") +edited_df["REQUEST_DATETIME"] = datetime.now().strftime("%Y-%m-%d %H:%M:%S") -@st.cache_data + +@st.cache_data def convert_df(df): - return df.to_csv(index=False).encode('utf-8') + return df.to_csv(index=False).encode("utf-8") + csv = convert_df(edited_df) # edited_df = edited_df.reset_index() # make sure indexes pair with number of rows # additional_info = pd.DataFrame(columns=['REQUEST_ID','T_DRIVE_PATH','CLIENT_NAME' # , 'GROUP_NAME', 'REQUEST_USERNAME', 'REQUEST_DATETIME']) -additional_info = pd.DataFrame(columns=['CLIENT_NAME', 'BATCH_ID', 'REQUEST_USERNAME', 'REQUEST_DATETIME']) -additional_info.loc[0] = [client, batch_id, st.session_state.user_info['mail'], datetime.now().strftime("%Y-%m-%d %H:%M:%S")] +additional_info = pd.DataFrame( + columns=["CLIENT_NAME", "BATCH_ID", "REQUEST_USERNAME", "REQUEST_DATETIME"] +) +additional_info.loc[0] = [ + client, + batch_id, + st.session_state.user_info["mail"], + datetime.now().strftime("%Y-%m-%d %H:%M:%S"), +] st.write(additional_info) st.session_state.contract_count = 0 @@ -189,32 +252,36 @@ contract_list = [] for index, row in edited_df.iterrows(): allow_run_for_contract = False group_list = [] - if row['Unique Key']: - group_list.append('Unique Key') + if row["Unique Key"]: + group_list.append("Unique Key") allow_run_for_contract = True - if row['Pricing Before Carveouts']: - group_list.append('Pricing Before Carveouts') + if row["Pricing Before Carveouts"]: + group_list.append("Pricing Before Carveouts") allow_run_for_contract = True - if row['Contract Related']: - group_list.append('Contract Related') + if row["Contract Related"]: + group_list.append("Contract Related") allow_run_for_contract = True - if row['Provider']: - group_list.append('Provider') + if row["Provider"]: + group_list.append("Provider") allow_run_for_contract = True - if row['Timeline']: - group_list.append('Timeline') + if row["Timeline"]: + group_list.append("Timeline") allow_run_for_contract = True - if row['Carveout Indicator']: - group_list.append('Carveout Indicator') + if row["Carveout Indicator"]: + group_list.append("Carveout Indicator") allow_run_for_contract = True - if row['Carveout Methodology']: - group_list.append('Carveout Methodology') + if row["Carveout Methodology"]: + group_list.append("Carveout Methodology") allow_run_for_contract = True - if allow_run_for_contract: st.session_state.contract_count += 1 + if allow_run_for_contract: + st.session_state.contract_count += 1 entry_dict = { - "contract_name": row['Contract Name'], + "contract_name": row["Contract Name"], "groups": group_list, - "contract_source_path": "contracts_landing_zone/"+batch_id+"/"+row['Contract Name'] + "contract_source_path": "contracts_landing_zone/" + + batch_id + + "/" + + row["Contract Name"], } contract_list.append(entry_dict) @@ -223,18 +290,22 @@ myobj = { "batch_id": batch_id, "client_name": client, "username": user_mail, - "contract_list": contract_list + "contract_list": contract_list, } buttons = st.columns([0.8, 0.2]) with buttons[0]: - st.download_button("Download Table", csv, "file.csv", "text/csv", key='download-csv') + st.download_button( + "Download Table", csv, "file.csv", "text/csv", key="download-csv" + ) with buttons[1]: if st.button("Run Doczy.AI Pipeline"): if not st.session_state.contract_count == len(edited_df): st.error("Select at least one Group No. for every Contract") else: - with st.spinner('Running...'): # Feedback to User while API endpoint sends response for DOC-342 + with st.spinner( + "Running..." + ): # Feedback to User while API endpoint sends response for DOC-342 # csv_buf = StringIO() # additional_info.to_csv(csv_buf, header=True, index=False) # csv_buf.seek(0) @@ -242,7 +313,7 @@ with buttons[1]: # csv_buf = StringIO() # edited_df.to_csv(csv_buf, header=True, index=False) # csv_buf.seek(0) - # s3_client.put_object(Bucket='doczy-dev-infra-raw-data-ingestion', Body=csv_buf.getvalue(), Key='training_interface/contract_config.csv') + # s3_client.put_object(Bucket='doczy-dev-infra-raw-data-ingestion', Body=csv_buf.getvalue(), Key='training_interface/contract_config.csv') # try: # save_to_sf('load_request_and_contract_submissions', request_submission_file_name = "request_submission.csv", contract_config_file_name = "contract_config.csv") # except Exception as e: @@ -256,4 +327,3 @@ with buttons[1]: # st.write(response.text) time.sleep(5) st.write("Success!") - diff --git a/streamlit/local/local_interface_2.py b/streamlit/local/local_interface_2.py index ad8192f..9b774e4 100644 --- a/streamlit/local/local_interface_2.py +++ b/streamlit/local/local_interface_2.py @@ -5,27 +5,30 @@ from langchain.prompts import PromptTemplate from langchain.embeddings.bedrock import BedrockEmbeddings from langchain.llms.bedrock import Bedrock from langchain_community.vectorstores import Chroma + # from constants import CHROMA_SETTINGS, EMBEDDING_MODEL_NAME, PERSIST_DIRECTORY, MODEL_ID, MODEL_BASENAME, SOURCE_DIRECTORY, USER_LIST from langchain.chains import RetrievalQA import streamlit as st from streamlit_extras.add_vertical_space import add_vertical_space -from streamlit_pdf_viewer import pdf_viewer # needs to be installed on the server +from streamlit_pdf_viewer import pdf_viewer # needs to be installed on the server import os import pandas as pd import numpy as np + # import util import anthropic from pydantic import BaseModel from typing import List import re import base64 + # from sf_conn import get_snowflake_conn -REDIRECT_URI = 'https://doczydev.aarete.com:8502' +REDIRECT_URI = "https://doczydev.aarete.com:8502" # user_list = USER_LIST -st.set_page_config(layout = "wide") +st.set_page_config(layout="wide") # Sidebar contents # with st.sidebar: # st.title("Doczy.AI ™") @@ -33,22 +36,25 @@ st.set_page_config(layout = "wide") # """ # ## About # This app extracts data from contracts - + # """ # ) # add_vertical_space(15) # # st.write("Doczy") -_,c1= st.columns([5,1]) +_, c1 = st.columns([5, 1]) # try: # util.setup_page(REDIRECT_URI) # except: # st.write("SSO Failed") # st.session_state['user_info'] = {'mail': 'maamseek@aarete.com', 'displayName': 'Mayank Aamseek'} -st.session_state['user_info'] = {'mail': 'maamseek@aarete.com', 'displayName': 'Mayank Aamseek'} +st.session_state["user_info"] = { + "mail": "maamseek@aarete.com", + "displayName": "Mayank Aamseek", +} try: c1.write(f"User: **{st.session_state.user_info['displayName']}**") - user_mail = st.session_state.user_info['mail'] + user_mail = st.session_state.user_info["mail"] except KeyError as e: st.write("Session Expired.") st.stop() @@ -66,8 +72,10 @@ except KeyError as e: # # field_values = field_values.loc[:, ~field_values.columns.str.contains('Unnamed:')] # st.write("Conn failed, unable to fetch data from training data table in Snowflake") -field_values = pd.read_csv('contract_field_values.csv', encoding='unicode_escape', skipinitialspace=True) -field_values = field_values.loc[:, ~field_values.columns.str.contains('Unnamed:')] +field_values = pd.read_csv( + "contract_field_values.csv", encoding="unicode_escape", skipinitialspace=True +) +field_values = field_values.loc[:, ~field_values.columns.str.contains("Unnamed:")] # try: # query = 'select * from "PROMPT_CONFIG"' @@ -78,20 +86,20 @@ field_values = field_values.loc[:, ~field_values.columns.str.contains('Unnamed:' # fields.rename(columns={'PROMPT': 'Interrogation Question?'}, inplace = True) # fields.rename(columns={'GROUP_ID': 'PRIORITY'}, inplace = True) # fields.rename(columns={'FIELD_NAME': 'SF_DB_COL_NAME'}, inplace = True) -# fields.rename(columns={'FM_MODEL_ID': 'llm_selected'}, inplace = True) +# fields.rename(columns={'FM_MODEL_ID': 'llm_selected'}, inplace = True) # except Exception as e: # st.write("Unable to fetch data from Snowflake: ",e) # # fields = pd.read_csv('contract_fields.csv', encoding='unicode_escape', skipinitialspace=True) # # fields = fields[~fields['SF_COL_NAME'].str.endswith('_PG', na=None)] -fields = pd.read_csv('contract_fields.csv', encoding='unicode_escape', skipinitialspace=True) +fields = pd.read_csv( + "contract_fields.csv", encoding="unicode_escape", skipinitialspace=True +) # fields = fields[fields['SF_DB_COL_NAME'].str.endswith('_PG', na=None)] # change the code below if contract list is fetched from snowflake -s3_client = boto3.client('s3', - region_name="us-east-2" -) -bucket = 'doczy-dev-infra-textract' +s3_client = boto3.client("s3", region_name="us-east-2") +bucket = "doczy-dev-infra-textract" # objects = s3_client.list_objects_v2(Bucket=bucket, Prefix="training-data/") # file_list = [] # for obj in objects['Contents']: @@ -99,11 +107,21 @@ bucket = 'doczy-dev-infra-textract' # file_list.append(obj['Key']) # # to be replaced with snowflake data -client_list = ['doczy-ai-client-1', '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'] +client_list = [ + "doczy-ai-client-1", + "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", +] -#Replace client_list with this to get client names from s3 buckets +# Replace client_list with this to get client names from s3 buckets # client_list, s3_paths = get_client_names() # client_s3_paths = dict(zip(client_list, s3_paths)) @@ -111,15 +129,17 @@ client_row = st.columns([0.2, 0.7, 0.1]) with client_row[0]: st.write("**Client Name**") with client_row[1]: - client = st.selectbox('Client Name',(client_list), label_visibility = "collapsed", index= None) + client = st.selectbox( + "Client Name", (client_list), label_visibility="collapsed", index=None + ) # batch_objects = s3_client.list_objects_v2(Bucket=client_bucket -# , Prefix="contracts_landing_zone/", Delimiter='/') +# , Prefix="contracts_landing_zone/", Delimiter='/') + +batch_list = ["d_12132", "d_13345", "i_23423", "i_72223", "b_12345", "b_33452"] -batch_list = ['d_12132', 'd_13345','i_23423', 'i_72223', 'b_12345', 'b_33452'] - # for prefix in batch_objects['CommonPrefixes']: # batch_list.append(prefix['Prefix'][:-1].split('/')[-1]) @@ -127,10 +147,12 @@ path_row = st.columns([0.2, 0.7, 0.1]) with path_row[0]: st.write("**Batch ID**") with path_row[1]: - batch_id = st.selectbox('**Batch ID**', batch_list, label_visibility = "collapsed", index=None) + batch_id = st.selectbox( + "**Batch ID**", batch_list, label_visibility="collapsed", index=None + ) -#Need to populate file_list with the list of files from the selected client from client_list +# Need to populate file_list with the list of files from the selected client from client_list # client_bucket = client_s3_paths.get(client) # bucket = client_bucket # objects = s3_client.list_objects_v2(Bucket=bucket, Prefix="training-data/") @@ -139,60 +161,82 @@ with path_row[1]: # if not obj['Key'].endswith('/'): # file_list.append(obj['Key']) -file_list = ['Contract_Training_Exercise_Pricing.pdf', 'Contract_Training_Exercise_SLA.pdf'] +file_list = [ + "Contract_Training_Exercise_Pricing.pdf", + "Contract_Training_Exercise_SLA.pdf", +] contract_list = sorted(file_list) file_row = st.columns([0.2, 0.7, 0.1]) with file_row[0]: st.write("**Contract Name**") with file_row[1]: - file_name = st.selectbox('Select a file', ['All'] + contract_list, label_visibility = "collapsed", index= None) # MODIFIED - Append 'All' in the front instead of at the end + file_name = st.selectbox( + "Select a file", + ["All"] + contract_list, + label_visibility="collapsed", + index=None, + ) # MODIFIED - Append 'All' in the front instead of at the end field_row = st.columns([0.2, 0.7, 0.1]) with field_row[0]: st.write("**Field Group**") with field_row[1]: - field_group = st.selectbox('Field Group',('Unique Key', 'Contract Related', 'Pricing Before Carveouts - I' - , 'Pricing Before Carveouts - II', 'Carveout Indicator, Code Type and Code #s - I' - , 'Carveout Indicator, Code Type and Code #s - II', 'Carveout Indicator, Code Type and Code #s - III' - , 'Optimize Carving Indic.', 'Carveout Method - I', 'Carveout Method - II', 'Provider' - , 'Timeline'), label_visibility = "collapsed", index = None) + field_group = st.selectbox( + "Field Group", + ( + "Unique Key", + "Contract Related", + "Pricing Before Carveouts - I", + "Pricing Before Carveouts - II", + "Carveout Indicator, Code Type and Code #s - I", + "Carveout Indicator, Code Type and Code #s - II", + "Carveout Indicator, Code Type and Code #s - III", + "Optimize Carving Indic.", + "Carveout Method - I", + "Carveout Method - II", + "Provider", + "Timeline", + ), + label_visibility="collapsed", + index=None, + ) -if field_group == 'Unique Key': - fields = fields[fields['PRIORITY'] == 'A'] -elif field_group == 'Contract Related': - fields = fields[fields['PRIORITY'] == 'C'] -elif field_group == 'Pricing Before Carveouts - I': - fields = fields[fields['PRIORITY'] == 'B'] +if field_group == "Unique Key": + fields = fields[fields["PRIORITY"] == "A"] +elif field_group == "Contract Related": + fields = fields[fields["PRIORITY"] == "C"] +elif field_group == "Pricing Before Carveouts - I": + fields = fields[fields["PRIORITY"] == "B"] fields = np.array_split(fields, 2)[0] -elif field_group == 'Pricing Before Carveouts - II': - fields = fields[fields['PRIORITY'] == 'B'] +elif field_group == "Pricing Before Carveouts - II": + fields = fields[fields["PRIORITY"] == "B"] fields = np.array_split(fields, 2)[1] -elif field_group == 'Carveout Indicator, Code Type and Code #s - I': - fields = fields[fields['PRIORITY'] == 'F'] +elif field_group == "Carveout Indicator, Code Type and Code #s - I": + fields = fields[fields["PRIORITY"] == "F"] fields = np.array_split(fields, 3)[0] -elif field_group == 'Carveout Indicator, Code Type and Code #s - II': - fields = fields[fields['PRIORITY'] == 'F'] +elif field_group == "Carveout Indicator, Code Type and Code #s - II": + fields = fields[fields["PRIORITY"] == "F"] fields = np.array_split(fields, 3)[1] -elif field_group == 'Carveout Indicator, Code Type and Code #s - III': - fields = fields[fields['PRIORITY'] == 'F'] +elif field_group == "Carveout Indicator, Code Type and Code #s - III": + fields = fields[fields["PRIORITY"] == "F"] fields = np.array_split(fields, 3)[2] -elif field_group == 'Carveout Methodology - I': - fields = fields[fields['PRIORITY'] == 'G'] +elif field_group == "Carveout Methodology - I": + fields = fields[fields["PRIORITY"] == "G"] fields = np.array_split(fields, 4)[0] -elif field_group == 'Carveout Methodology - II': - fields = fields[fields['PRIORITY'] == 'G'] +elif field_group == "Carveout Methodology - II": + fields = fields[fields["PRIORITY"] == "G"] fields = np.array_split(fields, 4)[1] -elif field_group == 'Carveout Method - III': - fields = fields[fields['PRIORITY'] == 'G'] +elif field_group == "Carveout Method - III": + fields = fields[fields["PRIORITY"] == "G"] fields = np.array_split(fields, 4)[2] -elif field_group == 'Carveout Method - IV': - fields = fields[fields['PRIORITY'] == 'G'] +elif field_group == "Carveout Method - IV": + fields = fields[fields["PRIORITY"] == "G"] fields = np.array_split(fields, 4)[3] -elif field_group == 'Provider': - fields = fields[fields['PRIORITY'] == 'D'] -elif field_group == 'Timeline': - fields = fields[fields['PRIORITY'] == 'E'] +elif field_group == "Provider": + fields = fields[fields["PRIORITY"] == "D"] +elif field_group == "Timeline": + fields = fields[fields["PRIORITY"] == "E"] if st.button("Show Results"): # query = 'select * from "DOCZY_PIPELINE_RAW_OUTPUT"' @@ -200,9 +244,19 @@ if st.button("Show Results"): # df2 = pd.DataFrame.from_records(iter(cur), columns=[x[0] for x in cur.description]) # get this dataframe from snowflake table - df2 = pd.DataFrame(columns=['Contract Name','Field Name', 'SF_DB_COL_NAME', 'Snippet','Page Number' - , 'Field Extracted Value', 'Actual Value','Imputed Value']) - df2.to_csv('temp2.csv', index=False) + df2 = pd.DataFrame( + columns=[ + "Contract Name", + "Field Name", + "SF_DB_COL_NAME", + "Snippet", + "Page Number", + "Field Extracted Value", + "Actual Value", + "Imputed Value", + ] + ) + df2.to_csv("temp2.csv", index=False) st.subheader("Showing " + client + " : " + file_name) if st.button("Show PDF"): @@ -244,27 +298,27 @@ if st.button("Show PDF"): # ) # st.markdown(pdf_display, unsafe_allow_html=True) -df2 = pd.read_csv('temp2.csv') -df2['Imputed Value'] = '' +df2 = pd.read_csv("temp2.csv") +df2["Imputed Value"] = "" edited_df = st.data_editor(df2) -@st.cache_data + +@st.cache_data def convert_df(df): - return df.to_csv(index=False).encode('utf-8') + return df.to_csv(index=False).encode("utf-8") + csv = convert_df(edited_df) buttons = st.columns(3) with buttons[0]: - # st.button("Save All Imputations") - st.download_button("Download Table", csv, "file.csv", "text/csv", key='download-csv') + # st.button("Save All Imputations") + st.download_button( + "Download Table", csv, "file.csv", "text/csv", key="download-csv" + ) with buttons[1]: # st.download_button("Download Table", csv, "file.csv", "text/csv", key='download-csv') st.write("") with buttons[2]: if st.button("Kickoff Database Integration"): st.write("Stored in DB") - - - - diff --git a/streamlit/multipage/Interface_0.py b/streamlit/multipage/Interface_0.py index a2bf8ea..f4ccba0 100644 --- a/streamlit/multipage/Interface_0.py +++ b/streamlit/multipage/Interface_0.py @@ -20,20 +20,21 @@ from util import logger (REDIRECT_URI, create_batch_url, doczy_pipeline) = util.load_page_details(5) user_list = USER_LIST -if 'uploading' not in st.session_state: +if "uploading" not in st.session_state: st.session_state.uploading = False -if 'upload_key' not in st.session_state: - st.session_state.upload_key = 0 -if 'file_list' not in st.session_state: +if "upload_key" not in st.session_state: + st.session_state.upload_key = 0 +if "file_list" not in st.session_state: st.session_state.file_list = [] -if 'show_batchID' not in st.session_state: +if "show_batchID" not in st.session_state: st.session_state.show_batchID = False -if 'landing_zone' not in st.session_state: - st.session_state.landing_zone = 'contracts-landing-zone' -if 'batch_id' not in st.session_state: - st.session_state.batch_id = 'failed_cases' -if 'client_bucket' not in st.session_state: - st.session_state.client_bucket = 'default_bucket' +if "landing_zone" not in st.session_state: + st.session_state.landing_zone = "contracts-landing-zone" +if "batch_id" not in st.session_state: + st.session_state.batch_id = "failed_cases" +if "client_bucket" not in st.session_state: + st.session_state.client_bucket = "default_bucket" + def insert_upload_logs(batch_id, client_name, file_name, upload_datetime, upload_user): """ @@ -41,12 +42,12 @@ def insert_upload_logs(batch_id, client_name, file_name, upload_datetime, upload Output: status of the insert query """ try: - return 'Log inserted successfully' + return "Log inserted successfully" except Exception as e: return e -st.set_page_config(layout = "wide") +st.set_page_config(layout="wide") # # # Sidebar contents with st.sidebar: st.title("Doczy.AI ™") @@ -61,42 +62,58 @@ with st.sidebar: # st.write("Doczy") # AARETE LOGO -x,y,z = st.columns([15,2,15]) +x, y, z = st.columns([15, 2, 15]) with y: - st.image('aaretelogo.png') + st.image("aaretelogo.png") -hide_img_fs = ''' +hide_img_fs = """ -''' +""" st.markdown(hide_img_fs, unsafe_allow_html=True) -_,c1= st.columns([5,1]) +_, c1 = st.columns([5, 1]) try: util.setup_page(REDIRECT_URI) except Exception as e: st.write(f"SSO Failed = {e}") - st.session_state['user_info'] = {'mail': 'maamseek@aarete.com', 'displayName': 'Mayank Aamseek'} + st.session_state["user_info"] = { + "mail": "maamseek@aarete.com", + "displayName": "Mayank Aamseek", + } try: c1.write(f"User: **{st.session_state.user_info['displayName']}**") - user_mail = st.session_state.user_info['mail'] + user_mail = st.session_state.user_info["mail"] except KeyError as e: st.write("Session Expired.") auth_url = security.get_auth_url(REDIRECT_URI) - st.markdown(f"Sign In", unsafe_allow_html=True) + st.markdown( + f"Sign In", unsafe_allow_html=True + ) st.stop() -s3_client = boto3.client('s3', +s3_client = boto3.client( + "s3", region_name="us-east-2", ) -client_list = ['doczy-ai-client-1', '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'] +client_list = [ + "doczy-ai-client-1", + "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", +] -# This is the list of client fetched from Snowflake +# This is the list of client fetched from Snowflake # TODO: Need to update the streamlit code to use the client names from this list # And use the s3 paths to save the objects for the respective client client_list, s3_paths = get_client_names() @@ -108,7 +125,9 @@ 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',(client_list), label_visibility = "collapsed", index = None) + client = st.selectbox( + "Client Name", (client_list), label_visibility="collapsed", index=None + ) client_bucket = client_s3_paths.get(client) @@ -119,20 +138,30 @@ file_row = st.columns([0.1, 0.8]) with file_row[0]: st.write("**Upload Files**") with file_row[1]: - file_list = st.file_uploader("Upload", type=['docx','tiff','pdf'], accept_multiple_files=True, label_visibility = "collapsed", help="Only PDF, TIFF and DOCX file formats are supported.", disabled=st.session_state.uploading, key = st.session_state.upload_key) + file_list = st.file_uploader( + "Upload", + type=["docx", "tiff", "pdf"], + accept_multiple_files=True, + label_visibility="collapsed", + help="Only PDF, TIFF and DOCX file formats are supported.", + disabled=st.session_state.uploading, + key=st.session_state.upload_key, + ) add_vertical_space(2) -df = pd.DataFrame(columns=['Contract Name']) -df['Contract Name'] = file_list +df = pd.DataFrame(columns=["Contract Name"]) +df["Contract Name"] = file_list file_names = [] buttons = st.columns([0.4, 0.4, 0.2]) + def set_uploading_state(): if not client == None and not len(file_list) == 0: st.session_state.file_list = file_list st.session_state.upload_key += 1 st.session_state.uploading = True + with buttons[1]: if st.button("Create Batch", on_click=set_uploading_state): file_list = st.session_state.file_list @@ -141,41 +170,57 @@ with buttons[1]: elif len(file_list) == 0: st.error("No Files Selected.") else: - with st.spinner('Running...'): - myobj = { "client-bucket-name": client_bucket } - response = requests.post(create_batch_url, json = myobj) + with st.spinner("Running..."): + myobj = {"client-bucket-name": client_bucket} + response = requests.post(create_batch_url, json=myobj) if response.status_code >= 200 and response.status_code < 300: try: - batch_id = json.loads(json.loads(response.text)['body'])['batch_id'] - landing_zone = json.loads(json.loads(response.text)['body'])['landing_zone'] - landing_zone = 'contracts-landing-zone/' + batch_id = json.loads(json.loads(response.text)["body"])[ + "batch_id" + ] + landing_zone = json.loads(json.loads(response.text)["body"])[ + "landing_zone" + ] + landing_zone = "contracts-landing-zone/" except: # st.write(myobj) # st.write(response.text) st.write("Internal Error. Reach out to Doczy.AI Team.") - batch_id = 'failed_cases' - landing_zone = 'contracts-landing-zone/' + batch_id = "failed_cases" + landing_zone = "contracts-landing-zone/" else: st.error("Failed") for uploaded_file in file_list: stringio = BytesIO(uploaded_file.getvalue()) stringio.seek(0) - s3_client.put_object(Bucket=client_bucket, Body=stringio.getvalue(), Key= - landing_zone+batch_id+'/'+str(uploaded_file.name)) - s3_client.put_object_tagging(Bucket=client_bucket, Key= - landing_zone+batch_id+'/'+str(uploaded_file.name), Tagging = {'TagSet': [ { 'Key': 'BatchId', 'Value': batch_id }]}) + s3_client.put_object( + Bucket=client_bucket, + Body=stringio.getvalue(), + Key=landing_zone + batch_id + "/" + str(uploaded_file.name), + ) + s3_client.put_object_tagging( + Bucket=client_bucket, + Key=landing_zone + batch_id + "/" + str(uploaded_file.name), + Tagging={"TagSet": [{"Key": "BatchId", "Value": batch_id}]}, + ) - # TODO: Test this insert function with snowflake - upload_log = insert_upload_logs(batch_id, client, str(uploaded_file.name), datetime.now().strftime("%Y-%m-%d %H:%M:%S"), user_mail) + # TODO: Test this insert function with snowflake + upload_log = insert_upload_logs( + batch_id, + client, + str(uploaded_file.name), + datetime.now().strftime("%Y-%m-%d %H:%M:%S"), + user_mail, + ) st.write(upload_log) - + file_names.append(str(uploaded_file.name)) st.session_state.uploading = False st.session_state.show_batchID = True st.session_state.client_bucket = client_bucket st.session_state.batch_id = batch_id st.session_state.landing_zone = landing_zone - st.rerun() + st.rerun() if st.session_state.show_batchID: batch_id = st.session_state.batch_id client_bucket = st.session_state.client_bucket @@ -185,13 +230,11 @@ with buttons[1]: st.write(f"Files uploaded to s3://{client_bucket}/{landing_zone}{batch_id}") st.session_state.show_batchID = False st.session_state.file_list = [] - st.session_state.batch_id = 'failed_cases' - st.session_state.client_bucket = 'default_bucket' - st.session_state.landing_zone = 'contracts-landing-zone' + st.session_state.batch_id = "failed_cases" + st.session_state.client_bucket = "default_bucket" + st.session_state.landing_zone = "contracts-landing-zone" - - - # @st.cache_data + # @st.cache_data # def convert_df(df): # return df.to_csv(index=False).encode('utf-8') diff --git a/streamlit/multipage/constants.py b/streamlit/multipage/constants.py index 8ccab24..f64dced 100644 --- a/streamlit/multipage/constants.py +++ b/streamlit/multipage/constants.py @@ -4,9 +4,17 @@ import os from chromadb.config import Settings # https://python.langchain.com/en/latest/modules/indexes/document_loaders/examples/excel.html?highlight=xlsx#microsoft-excel -from langchain_community.document_loaders import CSVLoader, PDFMinerLoader, TextLoader, UnstructuredExcelLoader, Docx2txtLoader -from langchain_community.document_loaders import UnstructuredFileLoader, UnstructuredMarkdownLoader - +from langchain_community.document_loaders import ( + CSVLoader, + PDFMinerLoader, + TextLoader, + UnstructuredExcelLoader, + Docx2txtLoader, +) +from langchain_community.document_loaders import ( + UnstructuredFileLoader, + UnstructuredMarkdownLoader, +) # load_dotenv() @@ -17,7 +25,7 @@ ROOT_DIRECTORY = "\\\\amznfsxuofkyi1z.aarete.local\\SharedFiles\\AArete Client W SOURCE_DIRECTORY = "SOURCE_DOCUMENTS" OUTPUT_DIRECTORY = f"{ROOT_DIRECTORY}\\Output" -PERSIST_DIRECTORY = 'DB' +PERSIST_DIRECTORY = "DB" MODELS_PATH = "C:\\Users\\Public\\models" @@ -190,48 +198,78 @@ MODEL_BASENAME = "llama-2-7b-chat.Q4_K_M.gguf" # MODEL_BASENAME = "model.safetensors.awq" - ########################################################################################################################################## ## CONSTANTS FOR INFRATRUCTURE -# SSO User list -USER_LIST = ['maamseek@aarete.com', 'smahdavian@aarete.com', 'ahinge@aarete.com', 'akadam@aarete.com', 'pkatariya@aarete.com' -, 'piragavarapu@aarete.com', 'umistry@aarete.com', 'ahutchison@aarete.com', 'bgrunst@aarete.com', 'ddimeglio@aarete.com' -, 'vnair@aarete.com', 'kminhas@aarete.com', 'fmohiuddin@aarete.com', 'slitewka@aarete.com', 'qdoest@aarete.com', 'bkoryga@aarete.com', 'bcielecki@aarete.com', 'mszymanski@aarete.com','hupreti@aarete.com', -'sshingare@aarete.com', 'vsrinivasan@aarete.com' ] +# SSO User list +USER_LIST = [ + "maamseek@aarete.com", + "smahdavian@aarete.com", + "ahinge@aarete.com", + "akadam@aarete.com", + "pkatariya@aarete.com", + "piragavarapu@aarete.com", + "umistry@aarete.com", + "ahutchison@aarete.com", + "bgrunst@aarete.com", + "ddimeglio@aarete.com", + "vnair@aarete.com", + "kminhas@aarete.com", + "fmohiuddin@aarete.com", + "slitewka@aarete.com", + "qdoest@aarete.com", + "bkoryga@aarete.com", + "bcielecki@aarete.com", + "mszymanski@aarete.com", + "hupreti@aarete.com", + "sshingare@aarete.com", + "vsrinivasan@aarete.com", +] # DOCZY DEV -DOCZY_PIPELINE_URL_DEV = 'https://4lzhid1s0h.execute-api.us-east-2.amazonaws.com/dev/trigger-pipeline' -DOCZY_REDIRECT_URL_DEV = 'https://doczydev.aarete.com:850' -DOCZY_CREATE_BATCH_URL_DEV = 'https://4lzhid1s0h.execute-api.us-east-2.amazonaws.com/dev/create-batch' +DOCZY_PIPELINE_URL_DEV = ( + "https://4lzhid1s0h.execute-api.us-east-2.amazonaws.com/dev/trigger-pipeline" +) +DOCZY_REDIRECT_URL_DEV = "https://doczydev.aarete.com:850" +DOCZY_CREATE_BATCH_URL_DEV = ( + "https://4lzhid1s0h.execute-api.us-east-2.amazonaws.com/dev/create-batch" +) # DOCZY UAT -DOCZY_PIPELINE_URL_UAT = 'https://29gm8cek03.execute-api.us-east-2.amazonaws.com/dev/trigger-pipeline' -DOCZY_REDIRECT_URL_UAT = 'https://doczyuat.aarete.com:850' -DOCZY_CREATE_BATCH_URL_UAT = 'https://29gm8cek03.execute-api.us-east-2.amazonaws.com/dev/create-batch' +DOCZY_PIPELINE_URL_UAT = ( + "https://29gm8cek03.execute-api.us-east-2.amazonaws.com/dev/trigger-pipeline" +) +DOCZY_REDIRECT_URL_UAT = "https://doczyuat.aarete.com:850" +DOCZY_CREATE_BATCH_URL_UAT = ( + "https://29gm8cek03.execute-api.us-east-2.amazonaws.com/dev/create-batch" +) # DOCZY PROD -DOCZY_PIPELINE_URL_PROD = 'https://d612isd3ja.execute-api.us-east-2.amazonaws.com/dev/trigger-pipeline' -DOCZY_REDIRECT_URL_PROD = 'https://doczy.aarete.com:850' -DOCZY_CREATE_BATCH_URL_PROD = 'https://d612isd3ja.execute-api.us-east-2.amazonaws.com/dev/create-batch' +DOCZY_PIPELINE_URL_PROD = ( + "https://d612isd3ja.execute-api.us-east-2.amazonaws.com/dev/trigger-pipeline" +) +DOCZY_REDIRECT_URL_PROD = "https://doczy.aarete.com:850" +DOCZY_CREATE_BATCH_URL_PROD = ( + "https://d612isd3ja.execute-api.us-east-2.amazonaws.com/dev/create-batch" +) -# SNOWFLAKE DEV DATABASE -SNOWFLAKE_ACCOUNT_LOCATOR="aarete-doczyai", -DEV_DB_ROLE = "DEVADMIN", -DEV_WH="DEV_XS", -DEV_DB="DOCZY_DEV", -DEV_STAGING_SCHEMA="STG" +# SNOWFLAKE DEV DATABASE +SNOWFLAKE_ACCOUNT_LOCATOR = ("aarete-doczyai",) +DEV_DB_ROLE = ("DEVADMIN",) +DEV_WH = ("DEV_XS",) +DEV_DB = ("DOCZY_DEV",) +DEV_STAGING_SCHEMA = "STG" # SNOWFLAKE UAT DATABASE -UAT_DB_ROLE = "UATADMIN", -UAT_WH="DEV_XS", -UAT_DB="DOCZY_UAT", -UAT_STAGING_SCHEMA="STG" +UAT_DB_ROLE = ("UATADMIN",) +UAT_WH = ("DEV_XS",) +UAT_DB = ("DOCZY_UAT",) +UAT_STAGING_SCHEMA = "STG" # SNOWFLAKE PROD DATABASE -PROD_DB_ROLE = "PRODADMIN", -PROD_WH="DEV_XS", -PROD_DB="DOCZY_PROD", -PROD_STAGING_SCHEMA="STG" \ No newline at end of file +PROD_DB_ROLE = ("PRODADMIN",) +PROD_WH = ("DEV_XS",) +PROD_DB = ("DOCZY_PROD",) +PROD_STAGING_SCHEMA = "STG" diff --git a/streamlit/multipage/pages/1_Interface_1.py b/streamlit/multipage/pages/1_Interface_1.py index 221f0d9..a6a2d4d 100644 --- a/streamlit/multipage/pages/1_Interface_1.py +++ b/streamlit/multipage/pages/1_Interface_1.py @@ -17,7 +17,7 @@ from util import logger (REDIRECT_URI, create_batch_url, doczy_pipeline) = util.load_page_details(5) user_list = USER_LIST -st.set_page_config(layout = "wide") +st.set_page_config(layout="wide") # Sidebar contents with st.sidebar: st.title("Doczy.AI ™") @@ -32,38 +32,43 @@ with st.sidebar: # st.write("Doczy") # AARETE LOGO -x,y,z = st.columns([15,2,15]) +x, y, z = st.columns([15, 2, 15]) with y: - st.image('aaretelogo.png') + st.image("aaretelogo.png") -hide_img_fs = ''' +hide_img_fs = """ -''' +""" st.markdown(hide_img_fs, unsafe_allow_html=True) -_,c1= st.columns([5,1]) +_, c1 = st.columns([5, 1]) try: util.setup_page(REDIRECT_URI) except: st.write("SSO Failed") - st.session_state['user_info'] = {'mail': 'maamseek@aarete.com', 'displayName': 'Mayank Aamseek'} + st.session_state["user_info"] = { + "mail": "maamseek@aarete.com", + "displayName": "Mayank Aamseek", + } try: c1.write(f"User: **{st.session_state.user_info['displayName']}**") - user_mail = st.session_state.user_info['mail'] + user_mail = st.session_state.user_info["mail"] except KeyError as e: - # Do we add a link to get to the login page here? + # Do we add a link to get to the login page here? st.write("Session Expired.") # st.write("Please sign-in to use this app.") auth_url = security.get_auth_url(REDIRECT_URI) - st.markdown(f"Sign In", unsafe_allow_html=True) + st.markdown( + f"Sign In", unsafe_allow_html=True + ) st.stop() - -s3_client = boto3.client('s3', +s3_client = boto3.client( + "s3", region_name="us-east-2", ) @@ -79,7 +84,9 @@ 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',(client_list), label_visibility = "collapsed", index = None) + client = st.selectbox( + "Client Name", (client_list), label_visibility="collapsed", index=None + ) if client: client_bucket = client_s3_paths.get(client) @@ -87,51 +94,60 @@ if client: # to be deleted when buckets for different clients are ready; below line is added only for testing the corresponding DAG # client_bucket = 'doczyai-use2-d-cn1-s3-textract-processing-001' - batch_objects = s3_client.list_objects_v2(Bucket=client_bucket - , Prefix="contracts-landing-zone/", Delimiter='/') + batch_objects = s3_client.list_objects_v2( + Bucket=client_bucket, Prefix="contracts-landing-zone/", Delimiter="/" + ) batch_list = [] - for prefix in batch_objects['CommonPrefixes']: - batch_name = prefix['Prefix'][:-1].split('/')[-1] - batch_objects2 = s3_client.list_objects_v2(Bucket=client_bucket, Prefix="contracts-landing-zone/"+batch_name+"/", Delimiter='/') - if 'KeyCount' in batch_objects2 and batch_objects2['KeyCount'] > 1: - batch_list.append(prefix['Prefix'][:-1].split('/')[-1]) + for prefix in batch_objects["CommonPrefixes"]: + batch_name = prefix["Prefix"][:-1].split("/")[-1] + batch_objects2 = s3_client.list_objects_v2( + Bucket=client_bucket, + Prefix="contracts-landing-zone/" + batch_name + "/", + Delimiter="/", + ) + if "KeyCount" in batch_objects2 and batch_objects2["KeyCount"] > 1: + batch_list.append(prefix["Prefix"][:-1].split("/")[-1]) # Hardcoded batch_list for testing purposes - # batch_list = ['batch_020524103737', 'batch_090524131433', 'batch_090524131607', 'batch_100524123000', 'batch_130524064322', + # batch_list = ['batch_020524103737', 'batch_090524131433', 'batch_090524131607', 'batch_100524123000', 'batch_130524064322', # 'batch_160524071331', 'batch_200524213550', 'batch_250424112237', 'batch_280524120530', 'batch_280524121721', 'batch_280524144222', # 'batch_290524123926', 'batch_290524164044', 'batch_310524102029', 'batch_310524124050', 'batch_310524162346', 'batch_310524162631'] - if 'sorted_list' not in st.session_state: + if "sorted_list" not in st.session_state: st.session_state.sorted_list = batch_list def sort_list(ex_list, sort_by, order): - if sort_by == 'Alphabetical': - ex_list = sorted(ex_list, reverse=(order == 'Descending')) - elif sort_by == 'Create Date': - ex_list = ex_list if order == 'Ascending' else list(reversed(ex_list)) + if sort_by == "Alphabetical": + ex_list = sorted(ex_list, reverse=(order == "Descending")) + elif sort_by == "Create Date": + ex_list = ex_list if order == "Ascending" else list(reversed(ex_list)) return ex_list col1, col2, col3, col4 = st.columns([0.5, 0.5, 0.5, 0.5]) - - with col1: - sort_by = st.radio("**Sort Batch_IDs**", ('Alphabetical', 'Create Date')) + with col1: + sort_by = st.radio("**Sort Batch_IDs**", ("Alphabetical", "Create Date")) with col2: - order = st.radio('', ('Ascending','Descending')) + order = st.radio("", ("Ascending", "Descending")) - with col3: + with col3: add_vertical_space(2) - if st.button('Apply'): + if st.button("Apply"): st.session_state.sorted_list = sort_list(batch_list, sort_by, order) path_row = st.columns([0.1, 0.8]) with path_row[0]: st.write("**Batch ID**") with path_row[1]: - batch_id = st.selectbox('**Batch ID**', st.session_state.sorted_list, label_visibility = "collapsed", index = None) + batch_id = st.selectbox( + "**Batch ID**", + st.session_state.sorted_list, + label_visibility="collapsed", + index=None, + ) if not batch_id: batch_id = "None" @@ -141,71 +157,93 @@ if client: st.write("**Group No.**") with checks[1]: - a = st.checkbox('Unique Key', key = str(1), args="Unique") + a = st.checkbox("Unique Key", key=str(1), args="Unique") with checks[2]: - b = st.checkbox('Pricing Before Carveouts', key = str(2)) + b = st.checkbox("Pricing Before Carveouts", key=str(2)) with checks[3]: - c = st.checkbox('Contract Related', key = str(3)) + c = st.checkbox("Contract Related", key=str(3)) with checks[4]: - d = st.checkbox('Provider', key = str(4)) + d = st.checkbox("Provider", key=str(4)) with checks[5]: - e = st.checkbox('Timeline', key = str(5)) + e = st.checkbox("Timeline", key=str(5)) with checks[6]: - f = st.checkbox('Carveout Indicator', key = str(6)) + f = st.checkbox("Carveout Indicator", key=str(6)) with checks[7]: - g = st.checkbox('Carveout Methodology', key = str(7)) + g = st.checkbox("Carveout Methodology", key=str(7)) add_vertical_space(1) - df = pd.DataFrame(columns=['Contract Name', 'Unique Key','Pricing Before Carveouts' - , 'Contract Related', 'Provider', 'Timeline', 'Carveout Indicator', 'Carveout Methodology']) + df = pd.DataFrame( + columns=[ + "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=client_bucket - , Prefix="contracts-landing-zone/"+batch_id+"/", Delimiter='/') + file_objects = s3_client.list_objects_v2( + Bucket=client_bucket, + Prefix="contracts-landing-zone/" + batch_id + "/", + Delimiter="/", + ) # Hardcoded file_list for testing purposes - # file_list = ['Boilerplate_TX Amendment Mission Health Network effective_040114 MU.pdf', 'Custom_TX - MP AMENDMENT - MISSION HEALTH NETWORK - MU.pdf', + # file_list = ['Boilerplate_TX Amendment Mission Health Network effective_040114 MU.pdf', 'Custom_TX - MP AMENDMENT - MISSION HEALTH NETWORK - MU.pdf', # 'Delaware First Health_First State Homecare Agency_212260_7 MU.pdf', 'Molina Healthcare of Texas, Inc. Amendment 4 - HIX ACA__EFF 01012016_MU.pdf'] 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 + 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 + 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) + print(f"DEBUGGING: PWD= {dir_path}") + df.to_csv("temp1.csv", index=False) add_vertical_space(1) - df2 = pd.read_csv('temp1.csv') + df2 = pd.read_csv("temp1.csv") edited_df = st.data_editor(df2) - edited_df['REQUEST_USER'] = user_mail - edited_df['LATEST_FLAG BOOLEAN'] = True - edited_df['PIPELINE_KICKOFF_DATETIME'] = datetime.now().strftime("%Y-%m-%d %H:%M:%S") - edited_df['REQUEST_DATETIME'] = datetime.now().strftime("%Y-%m-%d %H:%M:%S") + edited_df["REQUEST_USER"] = user_mail + edited_df["LATEST_FLAG BOOLEAN"] = True + edited_df["PIPELINE_KICKOFF_DATETIME"] = datetime.now().strftime( + "%Y-%m-%d %H:%M:%S" + ) + edited_df["REQUEST_DATETIME"] = datetime.now().strftime("%Y-%m-%d %H:%M:%S") - @st.cache_data + @st.cache_data def convert_df(df): - return df.to_csv(index=False).encode('utf-8') + return df.to_csv(index=False).encode("utf-8") csv = convert_df(edited_df) # edited_df = edited_df.reset_index() # make sure indexes pair with number of rows # additional_info = pd.DataFrame(columns=['REQUEST_ID','T_DRIVE_PATH','CLIENT_NAME' # , 'GROUP_NAME', 'REQUEST_USERNAME', 'REQUEST_DATETIME']) - additional_info = pd.DataFrame(columns=['CLIENT_NAME', 'BATCH_ID', 'REQUEST_USERNAME', 'REQUEST_DATETIME']) - additional_info.loc[0] = [client, batch_id, st.session_state.user_info['mail'], datetime.now().strftime("%Y-%m-%d %H:%M:%S")] + additional_info = pd.DataFrame( + columns=["CLIENT_NAME", "BATCH_ID", "REQUEST_USERNAME", "REQUEST_DATETIME"] + ) + additional_info.loc[0] = [ + client, + batch_id, + st.session_state.user_info["mail"], + datetime.now().strftime("%Y-%m-%d %H:%M:%S"), + ] st.write(additional_info) st.session_state.contract_count = 0 @@ -213,32 +251,36 @@ if client: for index, row in edited_df.iterrows(): allow_run_for_contract = False group_list = [] - if row['Unique Key']: - group_list.append('A') + if row["Unique Key"]: + group_list.append("A") allow_run_for_contract = True - if row['Pricing Before Carveouts']: - group_list.append('B') + if row["Pricing Before Carveouts"]: + group_list.append("B") allow_run_for_contract = True - if row['Contract Related']: - group_list.append('C') + if row["Contract Related"]: + group_list.append("C") allow_run_for_contract = True - if row['Provider']: - group_list.append('D') + if row["Provider"]: + group_list.append("D") allow_run_for_contract = True - if row['Timeline']: - group_list.append('E') + if row["Timeline"]: + group_list.append("E") allow_run_for_contract = True - if row['Carveout Indicator']: - group_list.append('F') + if row["Carveout Indicator"]: + group_list.append("F") allow_run_for_contract = True - if row['Carveout Methodology']: - group_list.append('G') + if row["Carveout Methodology"]: + group_list.append("G") allow_run_for_contract = True - if allow_run_for_contract: st.session_state.contract_count += 1 + if allow_run_for_contract: + st.session_state.contract_count += 1 entry_dict = { - "contract_name": row['Contract Name'], + "contract_name": row["Contract Name"], "groups": group_list, - "contract_source_path": "contracts-landing-zone/"+batch_id+"/"+row['Contract Name'] + "contract_source_path": "contracts-landing-zone/" + + batch_id + + "/" + + row["Contract Name"], } contract_list.append(entry_dict) @@ -247,18 +289,20 @@ if client: "batch_id": batch_id, "client_name": client, "username": user_mail, - "contract_list": contract_list + "contract_list": contract_list, } buttons = st.columns([0.8, 0.2]) with buttons[0]: - st.download_button("Download Table", csv, "file.csv", "text/csv", key='download-csv') + st.download_button( + "Download Table", csv, "file.csv", "text/csv", key="download-csv" + ) with buttons[1]: if st.button("Run Doczy.AI Pipeline"): if not st.session_state.contract_count == len(edited_df): st.error("Select at least one Group No. for every Contract") else: - with st.spinner('Running...'): + with st.spinner("Running..."): # csv_buf = StringIO() # additional_info.to_csv(csv_buf, header=True, index=False) # csv_buf.seek(0) @@ -266,12 +310,12 @@ if client: # csv_buf = StringIO() # edited_df.to_csv(csv_buf, header=True, index=False) # csv_buf.seek(0) - # s3_client.put_object(Bucket='doczy-dev-infra-raw-data-ingestion', Body=csv_buf.getvalue(), Key='training_interface/contract_config.csv') + # s3_client.put_object(Bucket='doczy-dev-infra-raw-data-ingestion', Body=csv_buf.getvalue(), Key='training_interface/contract_config.csv') # try: # save_to_sf('load_request_and_contract_submissions', request_submission_file_name = "request_submission.csv", contract_config_file_name = "contract_config.csv") # except Exception as e: # st.write(e) - response = requests.post(doczy_pipeline, json = myobj) + response = requests.post(doczy_pipeline, json=myobj) if response.status_code >= 200 and response.status_code < 300: st.write("Success") st.write("Current processing time for A & C: 1 min") @@ -280,4 +324,3 @@ if client: else: st.write("Failed") # st.write(response.text) - diff --git a/streamlit/multipage/pages/2_Interface_2.py b/streamlit/multipage/pages/2_Interface_2.py index c25c454..71d7ea4 100644 --- a/streamlit/multipage/pages/2_Interface_2.py +++ b/streamlit/multipage/pages/2_Interface_2.py @@ -5,12 +5,20 @@ from langchain.prompts import PromptTemplate from langchain.embeddings.bedrock import BedrockEmbeddings from langchain.llms.bedrock import Bedrock from langchain_community.vectorstores import Chroma -from constants import CHROMA_SETTINGS, EMBEDDING_MODEL_NAME, PERSIST_DIRECTORY, MODEL_ID, MODEL_BASENAME, SOURCE_DIRECTORY, USER_LIST +from constants import ( + CHROMA_SETTINGS, + EMBEDDING_MODEL_NAME, + PERSIST_DIRECTORY, + MODEL_ID, + MODEL_BASENAME, + SOURCE_DIRECTORY, + USER_LIST, +) from langchain.chains import RetrievalQA import streamlit as st from streamlit_extras.add_vertical_space import add_vertical_space -from streamlit_pdf_viewer import pdf_viewer # needs to be installed on the server +from streamlit_pdf_viewer import pdf_viewer # needs to be installed on the server import os import pandas as pd import numpy as np @@ -28,7 +36,7 @@ from util import logger (REDIRECT_URI, create_batch_url, doczy_pipeline) = util.load_page_details(5) user_list = USER_LIST -st.set_page_config(layout = "wide") +st.set_page_config(layout="wide") # Sidebar contents with st.sidebar: st.title("Doczy.AI ™") @@ -43,37 +51,42 @@ with st.sidebar: # st.write("Doczy") # AARETE LOGO -x,y,z = st.columns([15,2,15]) +x, y, z = st.columns([15, 2, 15]) with y: - st.image('aaretelogo.png') + st.image("aaretelogo.png") -hide_img_fs = ''' +hide_img_fs = """ -''' +""" st.markdown(hide_img_fs, unsafe_allow_html=True) -_,c1= st.columns([5,1]) +_, c1 = st.columns([5, 1]) try: util.setup_page(REDIRECT_URI) except: st.write("SSO Failed") - st.session_state['user_info'] = {'mail': 'maamseek@aarete.com', 'displayName': 'Mayank Aamseek'} + st.session_state["user_info"] = { + "mail": "maamseek@aarete.com", + "displayName": "Mayank Aamseek", + } try: c1.write(f"User: **{st.session_state.user_info['displayName']}**") - user_mail = st.session_state.user_info['mail'] + user_mail = st.session_state.user_info["mail"] except KeyError as e: st.write("Session Expired.") - #st.write("Please sign-in to use this app.") + # 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.markdown( + f"Sign In", unsafe_allow_html=True + ) st.stop() # remove below try except statement if comparison with actual vales is not required try: - conn = get_snowflake_conn('STG') + conn = get_snowflake_conn("STG") cur = conn.cursor() except: st.write("Conn failed, unable to fetch data from training data table in Snowflake") @@ -81,22 +94,22 @@ except: try: query = 'select * from "PROMPT_CONFIG"' cur.execute(query) - fields = pd.DataFrame.from_records(iter(cur), columns=[x[0] for x in cur.description]) + fields = pd.DataFrame.from_records( + iter(cur), columns=[x[0] for x in cur.description] + ) - fields.rename(columns={'FIELD_DESC': 'Field Name'}, inplace = True) - fields.rename(columns={'PROMPT': 'Interrogation Question?'}, inplace = True) - fields.rename(columns={'GROUP_ID': 'PRIORITY'}, inplace = True) - fields.rename(columns={'FIELD_NAME': 'SF_DB_COL_NAME'}, inplace = True) - fields.rename(columns={'FM_MODEL_ID': 'llm_selected'}, inplace = True) + fields.rename(columns={"FIELD_DESC": "Field Name"}, inplace=True) + fields.rename(columns={"PROMPT": "Interrogation Question?"}, inplace=True) + fields.rename(columns={"GROUP_ID": "PRIORITY"}, inplace=True) + fields.rename(columns={"FIELD_NAME": "SF_DB_COL_NAME"}, inplace=True) + fields.rename(columns={"FM_MODEL_ID": "llm_selected"}, inplace=True) except Exception as e: - st.write("Unable to fetch data from Snowflake: ",e) + st.write("Unable to fetch data from Snowflake: ", e) # fields = pd.read_csv('contract_fields.csv', encoding='unicode_escape', skipinitialspace=True) # fields = fields[~fields['SF_COL_NAME'].str.endswith('_PG', na=None)] # change the code below if contract list is fetched from snowflake -s3_client = boto3.client('s3', - region_name="us-east-2" -) +s3_client = boto3.client("s3", region_name="us-east-2") client_list, s3_paths = get_client_names() client_s3_paths = dict(zip(client_list, s3_paths)) @@ -105,35 +118,43 @@ client_row = st.columns([0.2, 0.7, 0.1]) with client_row[0]: st.write("**Client Name**") with client_row[1]: - client = st.selectbox('Client Name',(client_list), label_visibility = "collapsed", index= None) + client = st.selectbox( + "Client Name", (client_list), label_visibility="collapsed", index=None + ) if client: client_bucket = client_s3_paths.get(client) # client_bucket = 'doczyai-use2-d-cn1-s3-textract-processing-001' - batch_objects = s3_client.list_objects_v2(Bucket=client_bucket - , Prefix="textract-receiver-processed-pdfs/", Delimiter='/') + batch_objects = s3_client.list_objects_v2( + Bucket=client_bucket, Prefix="textract-receiver-processed-pdfs/", Delimiter="/" + ) batch_list = [] - for prefix in batch_objects['CommonPrefixes']: - batch_list.append(prefix['Prefix'][:-1].split('/')[-1]) + for prefix in batch_objects["CommonPrefixes"]: + batch_list.append(prefix["Prefix"][:-1].split("/")[-1]) path_row = st.columns([0.2, 0.7, 0.1]) with path_row[0]: st.write("**Batch ID**") with path_row[1]: - batch_id = st.selectbox('**Batch ID**', batch_list, label_visibility = "collapsed", index = None) + batch_id = st.selectbox( + "**Batch ID**", batch_list, label_visibility="collapsed", index=None + ) if batch_id: - objects = s3_client.list_objects_v2(Bucket=client_bucket, Prefix="textract-receiver-processed-pdfs/"+batch_id+"/") + objects = s3_client.list_objects_v2( + Bucket=client_bucket, + Prefix="textract-receiver-processed-pdfs/" + batch_id + "/", + ) file_list = [] - if 'Contents' in objects: - for obj in objects['Contents']: - if not obj['Key'].endswith('/'): - file_list.append(obj['Key']) + if "Contents" in objects: + for obj in objects["Contents"]: + if not obj["Key"].endswith("/"): + file_list.append(obj["Key"]) else: - st.error('This batch_id is empty.') + st.error("This batch_id is empty.") contract_list = sorted(file_list) @@ -141,14 +162,23 @@ if client: with file_row[0]: st.write("**Contract Name**") with file_row[1]: - file_name = st.selectbox('Select a file', ['All'] + contract_list, label_visibility = "collapsed", index= None) + file_name = st.selectbox( + "Select a file", + ["All"] + contract_list, + label_visibility="collapsed", + index=None, + ) field_row = st.columns([0.2, 0.7, 0.1]) with field_row[0]: st.write("**Field Group**") with field_row[1]: - field_group = st.selectbox('Field Group',('Unique Key', 'Pricing Before Carveouts', 'Contract Related'), - label_visibility = "collapsed", index = None) + field_group = st.selectbox( + "Field Group", + ("Unique Key", "Pricing Before Carveouts", "Contract Related"), + label_visibility="collapsed", + index=None, + ) # field_group = st.selectbox('Field Group',('Unique Key', 'Contract Related', 'Pricing Before Carveouts - I' # , 'Pricing Before Carveouts - II', 'Carveout Indicator, Code Type and Code #s - I' # , 'Carveout Indicator, Code Type and Code #s - II', 'Carveout Indicator, Code Type and Code #s - III' @@ -191,7 +221,7 @@ if client: # elif field_group == 'Timeline': # fields = fields[fields['PRIORITY'] == 'E'] - button_cols = st.columns([1,1,8]) + button_cols = st.columns([1, 1, 8]) with button_cols[0]: if st.button("Show PDF"): if file_name == None or file_name == "All": @@ -208,25 +238,43 @@ if client: """, unsafe_allow_html=True, ) - s3_obj = s3_client.get_object(Bucket = client_bucket, Key = file_name) - data=s3_obj['Body'].read() + s3_obj = s3_client.get_object( + Bucket=client_bucket, Key=file_name + ) + data = s3_obj["Body"].read() pdf_viewer(data, width=1500) with button_cols[1]: if st.button("Show Results"): try: if field_group: - if field_group == 'Unique Key' or field_group == 'Contract Related': - query = f'select * from "DOCZY_PIPELINE_RAW_OUTPUT_AC" where batch_id = \'{batch_id}\'' - elif field_group == 'Pricing Before Carveouts': - query = f'select * from "DOCZY_PIPELINE_RAW_OUTPUT_B" where batch_id = \'{batch_id}\'' + if ( + field_group == "Unique Key" + or field_group == "Contract Related" + ): + query = f"select * from \"DOCZY_PIPELINE_RAW_OUTPUT_AC\" where batch_id = '{batch_id}'" + elif field_group == "Pricing Before Carveouts": + query = f"select * from \"DOCZY_PIPELINE_RAW_OUTPUT_B\" where batch_id = '{batch_id}'" cur.execute(query) - df2 = pd.DataFrame.from_records(iter(cur), columns=[x[0] for x in cur.description]) - else: st.error('Please select a Field Group') + df2 = pd.DataFrame.from_records( + iter(cur), columns=[x[0] for x in cur.description] + ) + else: + st.error("Please select a Field Group") except: - df2 = pd.DataFrame(columns=['Contract Name','Field Name', 'SF_DB_COL_NAME', 'Snippet','Page Number' - , 'Field Extracted Value', 'Actual Value','Imputed Value']) - st.error('There was an error in fetching the output.') - df2.to_csv('temp2.csv', index=False) + df2 = pd.DataFrame( + columns=[ + "Contract Name", + "Field Name", + "SF_DB_COL_NAME", + "Snippet", + "Page Number", + "Field Extracted Value", + "Actual Value", + "Imputed Value", + ] + ) + st.error("There was an error in fetching the output.") + df2.to_csv("temp2.csv", index=False) # if st.button("Show PDF"): # if file_name == None or file_name == "All": @@ -250,23 +298,25 @@ if client: # ) # st.markdown(pdf_display, unsafe_allow_html=True) - df2 = pd.read_csv('temp2.csv') - df2['Imputed Value'] = '' + df2 = pd.read_csv("temp2.csv") + df2["Imputed Value"] = "" edited_df = st.data_editor(df2) - @st.cache_data + @st.cache_data def convert_df(df): - return df.to_csv(index=False).encode('utf-8') + return df.to_csv(index=False).encode("utf-8") csv = convert_df(edited_df) buttons = st.columns(3) with buttons[0]: - # st.button("Save All Imputations") - st.download_button("Download Table", csv, "file.csv", "text/csv", key='download-csv') + # st.button("Save All Imputations") + st.download_button( + "Download Table", csv, "file.csv", "text/csv", key="download-csv" + ) with buttons[1]: # st.download_button("Download Table", csv, "file.csv", "text/csv", key='download-csv') st.write("") with buttons[2]: if st.button("Kickoff Database Integration"): - st.write("Stored in DB") \ No newline at end of file + st.write("Stored in DB") diff --git a/streamlit/multipage/pages/3_Data_Dictionary.py b/streamlit/multipage/pages/3_Data_Dictionary.py index 46b4475..7b82a32 100644 --- a/streamlit/multipage/pages/3_Data_Dictionary.py +++ b/streamlit/multipage/pages/3_Data_Dictionary.py @@ -4,7 +4,15 @@ from langchain.prompts import PromptTemplate from langchain.embeddings.bedrock import BedrockEmbeddings from langchain.llms.bedrock import Bedrock from langchain_community.vectorstores import Chroma -from constants import CHROMA_SETTINGS, EMBEDDING_MODEL_NAME, PERSIST_DIRECTORY, MODEL_ID, MODEL_BASENAME, SOURCE_DIRECTORY, USER_LIST +from constants import ( + CHROMA_SETTINGS, + EMBEDDING_MODEL_NAME, + PERSIST_DIRECTORY, + MODEL_ID, + MODEL_BASENAME, + SOURCE_DIRECTORY, + USER_LIST, +) from langchain.chains import RetrievalQA import streamlit as st @@ -28,7 +36,7 @@ import pandas as pd (REDIRECT_URI, create_batch_url, doczy_pipeline) = util.load_page_details(5) user_list = USER_LIST -st.set_page_config(layout = "wide") +st.set_page_config(layout="wide") # Sidebar contents with st.sidebar: st.title("Doczy.AI ™") @@ -42,19 +50,22 @@ with st.sidebar: add_vertical_space(15) # st.write("Doczy") -_,c1= st.columns([5,1]) +_, c1 = st.columns([5, 1]) # util.setup_page(REDIRECT_URI) try: util.setup_page(REDIRECT_URI) except: st.write("SSO Failed") - st.session_state['user_info'] = {'mail': 'maamseek@aarete.com', 'displayName': 'Mayank Aamseek'} + st.session_state["user_info"] = { + "mail": "maamseek@aarete.com", + "displayName": "Mayank Aamseek", + } try: c1.write(f"User: **{st.session_state.user_info['displayName']}**") - user_mail = st.session_state.user_info['mail'] + user_mail = st.session_state.user_info["mail"] except KeyError as e: st.write("Session Expired.") - st. stop() + st.stop() -data_dict = pd.read_csv('contract_fields.csv') +data_dict = pd.read_csv("contract_fields.csv") st.data_editor(data_dict, height=700) diff --git a/streamlit/multipage/security.py b/streamlit/multipage/security.py index b3bcee3..de10d32 100644 --- a/streamlit/multipage/security.py +++ b/streamlit/multipage/security.py @@ -6,19 +6,17 @@ from botocore.exceptions import ClientError import json # Replace with your own values -CLIENT_ID = 'effafe90-7ed7-43a3-ab03-19a0be2f1758' -# CLIENT_SECRET = 'bjQ8Q~lpR2uBcGI34VDu16t73doz8Crj0YY_~dgD' +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 = 'https://172.29.20.102:8501' +AUTHORITY = "https://login.microsoftonline.com/organizations/" +SCOPE = ["User.Read"] +REDIRECT_URI = "https://172.29.20.102:8501" # Initialize boto3 client to interact with AWS Secrets Manager - - def get_secret(): secret_name = "doczy-sso-azure-app-key" @@ -30,48 +28,51 @@ def get_secret(): # service_name='secretsmanager', # region_name=region_name # ) - client = boto3.client('secretsmanager', region_name=region_name) + client = boto3.client("secretsmanager", region_name=region_name) try: - get_secret_value_response = client.get_secret_value( - SecretId=secret_name - ) + get_secret_value_response = client.get_secret_value(SecretId=secret_name) except ClientError as e: raise e - secret = get_secret_value_response['SecretString'] - secret = json.loads(secret)['CLIENT_SECRET'] + secret = get_secret_value_response["SecretString"] + secret = json.loads(secret)["CLIENT_SECRET"] return secret CLIENT_SECRET = get_secret() -app = msal.ConfidentialClientApplication(CLIENT_ID, authority=AUTHORITY, client_credential=CLIENT_SECRET) +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 + @st.cache_data 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'] + 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) + 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 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["access_token"] = access_token st.session_state - - \ No newline at end of file diff --git a/streamlit/multipage/sf_conn.py b/streamlit/multipage/sf_conn.py index ac2898f..c765e92 100644 --- a/streamlit/multipage/sf_conn.py +++ b/streamlit/multipage/sf_conn.py @@ -19,43 +19,38 @@ def get_secret(): # Create a Secrets Manager client session = boto3.session.Session() - client = session.client( - service_name='secretsmanager', - region_name=region_name - ) + client = session.client(service_name="secretsmanager", region_name=region_name) try: - get_secret_value_response = client.get_secret_value( - SecretId=secret_name - ) + get_secret_value_response = client.get_secret_value(SecretId=secret_name) except ClientError as e: # For a list of exceptions thrown, see # https://docs.aws.amazon.com/secretsmanager/latest/apireference/API_GetSecretValue.html raise e - secret = get_secret_value_response['SecretString'] + secret = get_secret_value_response["SecretString"] return secret + # get_secret() -# TODO: This function needs to be changed to accept Kwargs + +# TODO: This function needs to be changed to accept Kwargs # The function name should be more generic and cofnigurable def save_to_sf(dag_name, **kwargs): - mwaa_env_name = 'doczy-dev-infra-mwaa' + mwaa_env_name = "doczy-dev-infra-mwaa" dag_name = dag_name - mwaa_cli_command = 'dags trigger' + mwaa_cli_command = "dags trigger" # Create the client with the specified profile session = boto3.Session() - client = session.client('mwaa', region_name='us-east-2') + client = session.client("mwaa", region_name="us-east-2") # get web token - mwaa_cli_token = client.create_cli_token( - Name=mwaa_env_name - ) - - conn = http.client.HTTPSConnection(mwaa_cli_token['WebServerHostname']) + mwaa_cli_token = client.create_cli_token(Name=mwaa_env_name) + + conn = http.client.HTTPSConnection(mwaa_cli_token["WebServerHostname"]) # This section passes the payload to the MWAA CLI # The file parameters should be added dynamically in streamlit, once the file names are passed while triggering the dag, the data will be ingested @@ -67,8 +62,8 @@ def save_to_sf(dag_name, **kwargs): payload = mwaa_cli_command + " " + dag_name + " --conf '{}'".format(conf) headers = { - 'Authorization': 'Bearer ' + mwaa_cli_token['CliToken'], - 'Content-Type': 'text/plain' + "Authorization": "Bearer " + mwaa_cli_token["CliToken"], + "Content-Type": "text/plain", } conn.request("POST", "/aws_mwaa/cli", payload, headers) res = conn.getresponse() @@ -86,16 +81,17 @@ def get_snowflake_conn(schema): secret = get_secret() secret_dict = eval(secret) conn = snowflake.connector.connect( - user=secret_dict['user'], - password=secret_dict['password'], - account=secret_dict['account_alias'], - warehouse=secret_dict['warehouse'], - database=secret_dict['database'], - role=secret_dict['ROLE'], - schema=schema + user=secret_dict["user"], + password=secret_dict["password"], + account=secret_dict["account_alias"], + warehouse=secret_dict["warehouse"], + database=secret_dict["database"], + role=secret_dict["ROLE"], + schema=schema, ) return conn + def get_client_names(): """ Input: None @@ -103,7 +99,7 @@ def get_client_names(): """ try: # Get conn from snowflake_conn function for STG schema - conn = get_snowflake_conn('STG') + conn = get_snowflake_conn("STG") cursor = conn.cursor() # Query to get client names and their s3_paths @@ -112,10 +108,10 @@ def get_client_names(): # Create 2 lists from the query results client_names = [] s3_paths = [] - + for row in cursor: client_names.append(row[0]) - s3_paths.append(row[1]) # Extracting the bucket name from the s3 path + s3_paths.append(row[1]) # Extracting the bucket name from the s3 path cursor.close() conn.close() return client_names, s3_paths @@ -130,12 +126,12 @@ def insert_upload_logs(batch_id, client_name, file_name, upload_datetime, upload Output: status of the insert query """ try: - conn = get_snowflake_conn('STG') + conn = get_snowflake_conn("STG") cursor = conn.cursor() query = f"INSERT INTO STG.CONTRACT_UPLOAD_LOGS (BATCH_ID, CLIENT_NAME, FILE_NAME, UPLOAD_DATETIME, UPLOAD_USER) VALUES ('{batch_id}', '{client_name}', '{file_name}', '{upload_datetime}', '{upload_user}')" cursor.execute(query) cursor.close() conn.close() - return 'Log inserted successfully' + return "Log inserted successfully" except Exception as e: return e diff --git a/streamlit/multipage/util.py b/streamlit/multipage/util.py index 6497716..fa676ed 100644 --- a/streamlit/multipage/util.py +++ b/streamlit/multipage/util.py @@ -6,19 +6,19 @@ import logging from logging.handlers import RotatingFileHandler # Ensure the log directory exists -log_dir = '/home/ubuntu/doczy.ai/streamlit' +log_dir = "/home/ubuntu/doczy.ai/streamlit" if not os.path.exists(log_dir): os.makedirs(log_dir, exist_ok=True) # Configure logging with RotatingFileHandler -log_file = f'{log_dir}/interface.log' +log_file = f"{log_dir}/interface.log" rotating_handler = RotatingFileHandler( log_file, - maxBytes=10*1024*1024, # 10 MB - backupCount=5 # Keep up to 5 backup files + maxBytes=10 * 1024 * 1024, # 10 MB + backupCount=5, # Keep up to 5 backup files ) rotating_handler.setLevel(logging.INFO) -formatter = logging.Formatter('%(asctime)s %(levelname)s [%(filename)s] %(message)s') +formatter = logging.Formatter("%(asctime)s %(levelname)s [%(filename)s] %(message)s") rotating_handler.setFormatter(formatter) logger = logging.getLogger() @@ -32,38 +32,52 @@ def setup_page(redirect_uril): # page_icon="👋", # ) - if st.query_params.get('code'): + if st.query_params.get("code"): security.handle_redirect(redirect_uril) - access_token = st.session_state.get('access_token') + 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 + 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_uril) - st.markdown(f"Sign In", unsafe_allow_html=True) + st.markdown( + f"Sign In", unsafe_allow_html=True + ) st.stop() def load_page_details(interface): - env_var = os.environ.get('ENVIRONMENT', 'DEV') + env_var = os.environ.get("ENVIRONMENT", "DEV") logger.info(f"env_var={env_var}") - if env_var == 'UAT': + if env_var == "UAT": logger.info(constants.DOCZY_REDIRECT_URL_UAT + str(interface)) - return (constants.DOCZY_REDIRECT_URL_UAT + str(interface), constants.DOCZY_CREATE_BATCH_URL_UAT, - constants.DOCZY_PIPELINE_URL_UAT) - elif env_var == 'DEV': + return ( + constants.DOCZY_REDIRECT_URL_UAT + str(interface), + constants.DOCZY_CREATE_BATCH_URL_UAT, + constants.DOCZY_PIPELINE_URL_UAT, + ) + elif env_var == "DEV": logger.info(constants.DOCZY_REDIRECT_URL_DEV + str(interface)) - return (constants.DOCZY_REDIRECT_URL_DEV + str(interface), constants.DOCZY_CREATE_BATCH_URL_DEV, - constants.DOCZY_PIPELINE_URL_DEV) - elif env_var == 'PROD': + return ( + constants.DOCZY_REDIRECT_URL_DEV + str(interface), + constants.DOCZY_CREATE_BATCH_URL_DEV, + constants.DOCZY_PIPELINE_URL_DEV, + ) + elif env_var == "PROD": logger.info(constants.DOCZY_REDIRECT_URL_PROD + str(interface)) - return (constants.DOCZY_REDIRECT_URL_PROD + str(interface), constants.DOCZY_CREATE_BATCH_URL_PROD, - constants.DOCZY_PIPELINE_URL_PROD) + return ( + constants.DOCZY_REDIRECT_URL_PROD + str(interface), + constants.DOCZY_CREATE_BATCH_URL_PROD, + constants.DOCZY_PIPELINE_URL_PROD, + ) else: logger.info(constants.DOCZY_REDIRECT_URL_DEV + str(interface)) - return (constants.DOCZY_REDIRECT_URL_DEV + str(interface), constants.DOCZY_CREATE_BATCH_URL_DEV, - constants.DOCZY_PIPELINE_URL_DEV) + return ( + constants.DOCZY_REDIRECT_URL_DEV + str(interface), + constants.DOCZY_CREATE_BATCH_URL_DEV, + constants.DOCZY_PIPELINE_URL_DEV, + ) diff --git a/streamlit/security.py b/streamlit/security.py index d7465d2..2333230 100644 --- a/streamlit/security.py +++ b/streamlit/security.py @@ -6,19 +6,18 @@ from botocore.exceptions import ClientError import json # Replace with your own values -CLIENT_ID = 'effafe90-7ed7-43a3-ab03-19a0be2f1758' -# CLIENT_SECRET = 'bjQ8Q~lpR2uBcGI34VDu16t73doz8Crj0YY_~dgD' +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 = 'https://172.29.20.102:8501' +AUTHORITY = "https://login.microsoftonline.com/organizations/" +SCOPE = ["User.Read"] +REDIRECT_URI = "https://172.29.20.102:8501" # Initialize boto3 client to interact with AWS Secrets Manager - -#URL Masking +# URL Masking def clear_url(): js_code = """ window.history.replaceState({}, document.title, window.location.pathname); @@ -37,54 +36,56 @@ def get_secret(): # service_name='secretsmanager', # region_name=region_name # ) - client = boto3.client('secretsmanager', region_name=region_name) + client = boto3.client("secretsmanager", region_name=region_name) try: - get_secret_value_response = client.get_secret_value( - SecretId=secret_name - ) + get_secret_value_response = client.get_secret_value(SecretId=secret_name) except ClientError as e: raise e - secret = get_secret_value_response['SecretString'] - secret = json.loads(secret)['CLIENT_SECRET'] + secret = get_secret_value_response["SecretString"] + secret = json.loads(secret)["CLIENT_SECRET"] return secret CLIENT_SECRET = get_secret() -app = msal.ConfidentialClientApplication(CLIENT_ID, authority=AUTHORITY, client_credential=CLIENT_SECRET) +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 ) + auth_url = app.get_authorization_request_url(SCOPE, redirect_uri=REDIRECT_URI) return auth_url + @st.cache_data 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'] + 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) + 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 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["access_token"] = access_token st.session_state - - -#URL Masking Pt.2 -if 'access_token' in st.session_state: +# URL Masking Pt.2 +if "access_token" in st.session_state: clear_url() - \ No newline at end of file diff --git a/streamlit/sf_conn.py b/streamlit/sf_conn.py index ac2898f..c765e92 100644 --- a/streamlit/sf_conn.py +++ b/streamlit/sf_conn.py @@ -19,43 +19,38 @@ def get_secret(): # Create a Secrets Manager client session = boto3.session.Session() - client = session.client( - service_name='secretsmanager', - region_name=region_name - ) + client = session.client(service_name="secretsmanager", region_name=region_name) try: - get_secret_value_response = client.get_secret_value( - SecretId=secret_name - ) + get_secret_value_response = client.get_secret_value(SecretId=secret_name) except ClientError as e: # For a list of exceptions thrown, see # https://docs.aws.amazon.com/secretsmanager/latest/apireference/API_GetSecretValue.html raise e - secret = get_secret_value_response['SecretString'] + secret = get_secret_value_response["SecretString"] return secret + # get_secret() -# TODO: This function needs to be changed to accept Kwargs + +# TODO: This function needs to be changed to accept Kwargs # The function name should be more generic and cofnigurable def save_to_sf(dag_name, **kwargs): - mwaa_env_name = 'doczy-dev-infra-mwaa' + mwaa_env_name = "doczy-dev-infra-mwaa" dag_name = dag_name - mwaa_cli_command = 'dags trigger' + mwaa_cli_command = "dags trigger" # Create the client with the specified profile session = boto3.Session() - client = session.client('mwaa', region_name='us-east-2') + client = session.client("mwaa", region_name="us-east-2") # get web token - mwaa_cli_token = client.create_cli_token( - Name=mwaa_env_name - ) - - conn = http.client.HTTPSConnection(mwaa_cli_token['WebServerHostname']) + mwaa_cli_token = client.create_cli_token(Name=mwaa_env_name) + + conn = http.client.HTTPSConnection(mwaa_cli_token["WebServerHostname"]) # This section passes the payload to the MWAA CLI # The file parameters should be added dynamically in streamlit, once the file names are passed while triggering the dag, the data will be ingested @@ -67,8 +62,8 @@ def save_to_sf(dag_name, **kwargs): payload = mwaa_cli_command + " " + dag_name + " --conf '{}'".format(conf) headers = { - 'Authorization': 'Bearer ' + mwaa_cli_token['CliToken'], - 'Content-Type': 'text/plain' + "Authorization": "Bearer " + mwaa_cli_token["CliToken"], + "Content-Type": "text/plain", } conn.request("POST", "/aws_mwaa/cli", payload, headers) res = conn.getresponse() @@ -86,16 +81,17 @@ def get_snowflake_conn(schema): secret = get_secret() secret_dict = eval(secret) conn = snowflake.connector.connect( - user=secret_dict['user'], - password=secret_dict['password'], - account=secret_dict['account_alias'], - warehouse=secret_dict['warehouse'], - database=secret_dict['database'], - role=secret_dict['ROLE'], - schema=schema + user=secret_dict["user"], + password=secret_dict["password"], + account=secret_dict["account_alias"], + warehouse=secret_dict["warehouse"], + database=secret_dict["database"], + role=secret_dict["ROLE"], + schema=schema, ) return conn + def get_client_names(): """ Input: None @@ -103,7 +99,7 @@ def get_client_names(): """ try: # Get conn from snowflake_conn function for STG schema - conn = get_snowflake_conn('STG') + conn = get_snowflake_conn("STG") cursor = conn.cursor() # Query to get client names and their s3_paths @@ -112,10 +108,10 @@ def get_client_names(): # Create 2 lists from the query results client_names = [] s3_paths = [] - + for row in cursor: client_names.append(row[0]) - s3_paths.append(row[1]) # Extracting the bucket name from the s3 path + s3_paths.append(row[1]) # Extracting the bucket name from the s3 path cursor.close() conn.close() return client_names, s3_paths @@ -130,12 +126,12 @@ def insert_upload_logs(batch_id, client_name, file_name, upload_datetime, upload Output: status of the insert query """ try: - conn = get_snowflake_conn('STG') + conn = get_snowflake_conn("STG") cursor = conn.cursor() query = f"INSERT INTO STG.CONTRACT_UPLOAD_LOGS (BATCH_ID, CLIENT_NAME, FILE_NAME, UPLOAD_DATETIME, UPLOAD_USER) VALUES ('{batch_id}', '{client_name}', '{file_name}', '{upload_datetime}', '{upload_user}')" cursor.execute(query) cursor.close() conn.close() - return 'Log inserted successfully' + return "Log inserted successfully" except Exception as e: return e diff --git a/streamlit/util.py b/streamlit/util.py index 6497716..fa676ed 100644 --- a/streamlit/util.py +++ b/streamlit/util.py @@ -6,19 +6,19 @@ import logging from logging.handlers import RotatingFileHandler # Ensure the log directory exists -log_dir = '/home/ubuntu/doczy.ai/streamlit' +log_dir = "/home/ubuntu/doczy.ai/streamlit" if not os.path.exists(log_dir): os.makedirs(log_dir, exist_ok=True) # Configure logging with RotatingFileHandler -log_file = f'{log_dir}/interface.log' +log_file = f"{log_dir}/interface.log" rotating_handler = RotatingFileHandler( log_file, - maxBytes=10*1024*1024, # 10 MB - backupCount=5 # Keep up to 5 backup files + maxBytes=10 * 1024 * 1024, # 10 MB + backupCount=5, # Keep up to 5 backup files ) rotating_handler.setLevel(logging.INFO) -formatter = logging.Formatter('%(asctime)s %(levelname)s [%(filename)s] %(message)s') +formatter = logging.Formatter("%(asctime)s %(levelname)s [%(filename)s] %(message)s") rotating_handler.setFormatter(formatter) logger = logging.getLogger() @@ -32,38 +32,52 @@ def setup_page(redirect_uril): # page_icon="👋", # ) - if st.query_params.get('code'): + if st.query_params.get("code"): security.handle_redirect(redirect_uril) - access_token = st.session_state.get('access_token') + 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 + 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_uril) - st.markdown(f"Sign In", unsafe_allow_html=True) + st.markdown( + f"Sign In", unsafe_allow_html=True + ) st.stop() def load_page_details(interface): - env_var = os.environ.get('ENVIRONMENT', 'DEV') + env_var = os.environ.get("ENVIRONMENT", "DEV") logger.info(f"env_var={env_var}") - if env_var == 'UAT': + if env_var == "UAT": logger.info(constants.DOCZY_REDIRECT_URL_UAT + str(interface)) - return (constants.DOCZY_REDIRECT_URL_UAT + str(interface), constants.DOCZY_CREATE_BATCH_URL_UAT, - constants.DOCZY_PIPELINE_URL_UAT) - elif env_var == 'DEV': + return ( + constants.DOCZY_REDIRECT_URL_UAT + str(interface), + constants.DOCZY_CREATE_BATCH_URL_UAT, + constants.DOCZY_PIPELINE_URL_UAT, + ) + elif env_var == "DEV": logger.info(constants.DOCZY_REDIRECT_URL_DEV + str(interface)) - return (constants.DOCZY_REDIRECT_URL_DEV + str(interface), constants.DOCZY_CREATE_BATCH_URL_DEV, - constants.DOCZY_PIPELINE_URL_DEV) - elif env_var == 'PROD': + return ( + constants.DOCZY_REDIRECT_URL_DEV + str(interface), + constants.DOCZY_CREATE_BATCH_URL_DEV, + constants.DOCZY_PIPELINE_URL_DEV, + ) + elif env_var == "PROD": logger.info(constants.DOCZY_REDIRECT_URL_PROD + str(interface)) - return (constants.DOCZY_REDIRECT_URL_PROD + str(interface), constants.DOCZY_CREATE_BATCH_URL_PROD, - constants.DOCZY_PIPELINE_URL_PROD) + return ( + constants.DOCZY_REDIRECT_URL_PROD + str(interface), + constants.DOCZY_CREATE_BATCH_URL_PROD, + constants.DOCZY_PIPELINE_URL_PROD, + ) else: logger.info(constants.DOCZY_REDIRECT_URL_DEV + str(interface)) - return (constants.DOCZY_REDIRECT_URL_DEV + str(interface), constants.DOCZY_CREATE_BATCH_URL_DEV, - constants.DOCZY_PIPELINE_URL_DEV) + return ( + constants.DOCZY_REDIRECT_URL_DEV + str(interface), + constants.DOCZY_CREATE_BATCH_URL_DEV, + constants.DOCZY_PIPELINE_URL_DEV, + ) diff --git a/streamlit_multipage/Interface_1.py b/streamlit_multipage/Interface_1.py index e6182a1..d2dd4a6 100644 --- a/streamlit_multipage/Interface_1.py +++ b/streamlit_multipage/Interface_1.py @@ -6,18 +6,18 @@ from datetime import datetime import time import constants -if 'uploading' not in st.session_state: +if "uploading" not in st.session_state: st.session_state.uploading = False -if 'upload_key' not in st.session_state: - st.session_state.upload_key = 0 -if 'file_list' not in st.session_state: +if "upload_key" not in st.session_state: + st.session_state.upload_key = 0 +if "file_list" not in st.session_state: st.session_state.file_list = [] -if 'show_batchID' not in st.session_state: +if "show_batchID" not in st.session_state: st.session_state.show_batchID = False -if 'numFiles' not in st.session_state: +if "numFiles" not in st.session_state: st.session_state.numFiles = 0 -st.set_page_config(layout = "wide") +st.set_page_config(layout="wide") # # # Sidebar contents with st.sidebar: st.title("Doczy.AI ™") @@ -32,23 +32,26 @@ with st.sidebar: # st.write("Doczy") # AARETE LOGO -x,y,z = st.columns([15,2,15]) +x, y, z = st.columns([15, 2, 15]) with y: - st.image('aaretelogo.png') + st.image("aaretelogo.png") -hide_img_fs = ''' +hide_img_fs = """ -''' +""" st.markdown(hide_img_fs, unsafe_allow_html=True) -_,c1= st.columns([5,1]) -st.session_state['user_info'] = {'mail': 'piragavarapu@aarete.com', 'displayName': 'Priya Iragavarapu'} +_, c1 = st.columns([5, 1]) +st.session_state["user_info"] = { + "mail": "piragavarapu@aarete.com", + "displayName": "Priya Iragavarapu", +} try: c1.write(f"User: **{st.session_state.user_info['displayName']}**") - user_mail = st.session_state.user_info['mail'] + user_mail = st.session_state.user_info["mail"] except KeyError as e: st.write("Session Expired.") st.stop() @@ -57,38 +60,51 @@ 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',(constants.client_list), label_visibility = "collapsed", index = None) + client = st.selectbox( + "Client Name", (constants.client_list), label_visibility="collapsed", index=None + ) file_row = st.columns([0.1, 0.8]) with file_row[0]: st.write("**Upload Files**") with file_row[1]: - file_list = st.file_uploader("Upload", type=['docx','tiff','pdf'], accept_multiple_files=True, label_visibility = "collapsed", help="Only PDF, TIFF and DOCX file formats are supported.", disabled=st.session_state.uploading, key = st.session_state.upload_key) + file_list = st.file_uploader( + "Upload", + type=["docx", "tiff", "pdf"], + accept_multiple_files=True, + label_visibility="collapsed", + help="Only PDF, TIFF and DOCX file formats are supported.", + disabled=st.session_state.uploading, + key=st.session_state.upload_key, + ) add_vertical_space(2) -df = pd.DataFrame(columns=['Contract Name']) -df['Contract Name'] = file_list +df = pd.DataFrame(columns=["Contract Name"]) +df["Contract Name"] = file_list file_names = [] buttons = st.columns([0.4, 0.4, 0.2]) + def set_uploading_state(): if not client == None and not len(file_list) == 0: st.session_state.file_list = file_list st.session_state.upload_key += 1 st.session_state.uploading = True + def save_uploaded_files(uploaded_files, save_path): try: if not os.path.exists(save_path): os.makedirs(save_path) for uploaded_file in uploaded_files: - with open(os.path.join(save_path, uploaded_file.name), 'wb') as f: + with open(os.path.join(save_path, uploaded_file.name), "wb") as f: f.write(uploaded_file.getbuffer()) return True except Exception as e: st.error(f"Error saving files: {e}") return False - + + with buttons[1]: if st.button("Create Batch", on_click=set_uploading_state): file_list = st.session_state.file_list @@ -97,12 +113,12 @@ with buttons[1]: elif len(file_list) == 0: st.error("No Files Selected.") else: - with st.spinner('Running...'): + with st.spinner("Running..."): time.sleep(2) st.session_state.numFiles = len(file_list) st.session_state.uploading = False st.session_state.show_batchID = True - st.rerun() + st.rerun() if st.session_state.show_batchID: st.write(f"{st.session_state.numFiles} files uploaded successfully!") @@ -110,43 +126,52 @@ with buttons[1]: add_vertical_space(1) -df = pd.DataFrame(columns=['Contract Name', 'Run Flag']) +df = pd.DataFrame(columns=["Contract Name", "Run Flag"]) -df['Contract Name'] = [file.name for file in st.session_state.file_list] +df["Contract Name"] = [file.name for file in st.session_state.file_list] # df['Contract Name'] = st.session_state.file_list -df['Run Flag'] = True +df["Run Flag"] = True dir_path = os.path.dirname(os.path.realpath(__file__)) -print(f'DEBUGGING: PWD= {dir_path}') -df.to_csv('temp1.csv', index=False) +print(f"DEBUGGING: PWD= {dir_path}") +df.to_csv("temp1.csv", index=False) add_vertical_space(1) -df2 = pd.read_csv('temp1.csv') +df2 = pd.read_csv("temp1.csv") edited_df = st.data_editor(df2) -edited_df['REQUEST_USER'] = user_mail -edited_df['LATEST_FLAG BOOLEAN'] = True -edited_df['PIPELINE_KICKOFF_DATETIME'] = datetime.now().strftime("%Y-%m-%d %H:%M:%S") -edited_df['REQUEST_DATETIME'] = datetime.now().strftime("%Y-%m-%d %H:%M:%S") +edited_df["REQUEST_USER"] = user_mail +edited_df["LATEST_FLAG BOOLEAN"] = True +edited_df["PIPELINE_KICKOFF_DATETIME"] = datetime.now().strftime("%Y-%m-%d %H:%M:%S") +edited_df["REQUEST_DATETIME"] = datetime.now().strftime("%Y-%m-%d %H:%M:%S") -@st.cache_data + +@st.cache_data def convert_df(df): - return df.to_csv(index=False).encode('utf-8') + return df.to_csv(index=False).encode("utf-8") + csv = convert_df(edited_df) -additional_info = pd.DataFrame(columns=['CLIENT_NAME', 'REQUEST_USERNAME', 'REQUEST_DATETIME']) -additional_info.loc[0] = [client, st.session_state.user_info['mail'], datetime.now().strftime("%Y-%m-%d %H:%M:%S")] +additional_info = pd.DataFrame( + columns=["CLIENT_NAME", "REQUEST_USERNAME", "REQUEST_DATETIME"] +) +additional_info.loc[0] = [ + client, + st.session_state.user_info["mail"], + datetime.now().strftime("%Y-%m-%d %H:%M:%S"), +] st.write(additional_info) buttons = st.columns([0.8, 0.2]) with buttons[0]: - st.download_button("Download Table", csv, "file.csv", "text/csv", key='download-csv') + st.download_button( + "Download Table", csv, "file.csv", "text/csv", key="download-csv" + ) with buttons[1]: if st.button("Run Doczy.AI Pipeline"): - with st.spinner('Running...'): + with st.spinner("Running..."): time.sleep(2) st.write("Success!") st.write("Current processing time for Contract Terms: 1 min") st.write("Current processing time for Reimbursement: 10 mins") - diff --git a/streamlit_multipage/constants.py b/streamlit_multipage/constants.py index 537db8f..8a49dcb 100644 --- a/streamlit_multipage/constants.py +++ b/streamlit_multipage/constants.py @@ -1,3 +1,3 @@ -current_user = {'mail':'piragavarapu@aarete.com', 'name':'Priya Iragavarapu'} +current_user = {"mail": "piragavarapu@aarete.com", "name": "Priya Iragavarapu"} # current_user = {'mail':'agupta@aarete.com', 'name':'Aryan Gupta'} -client_list = ['Prime Care', 'Health Guard', 'Vita Shield'] \ No newline at end of file +client_list = ["Prime Care", "Health Guard", "Vita Shield"] diff --git a/streamlit_multipage/pages/1_Interface_2.py b/streamlit_multipage/pages/1_Interface_2.py index a8cb1c3..71214ec 100644 --- a/streamlit_multipage/pages/1_Interface_2.py +++ b/streamlit_multipage/pages/1_Interface_2.py @@ -7,7 +7,7 @@ from typing import List import os import re -st.set_page_config(layout = "wide") +st.set_page_config(layout="wide") # Sidebar contents with st.sidebar: st.title("Doczy.AI ™") @@ -20,23 +20,26 @@ with st.sidebar: ) # AARETE LOGO -x,y,z = st.columns([15,2,15]) +x, y, z = st.columns([15, 2, 15]) with y: - st.image('aaretelogo.png') + st.image("aaretelogo.png") -hide_img_fs = ''' +hide_img_fs = """ -''' +""" st.markdown(hide_img_fs, unsafe_allow_html=True) -_,c1= st.columns([5,1]) -st.session_state['user_info'] = {'mail': constants.current_user['mail'], 'displayName': constants.current_user['name']} +_, c1 = st.columns([5, 1]) +st.session_state["user_info"] = { + "mail": constants.current_user["mail"], + "displayName": constants.current_user["name"], +} try: c1.write(f"User: **{st.session_state.user_info['displayName']}**") - user_mail = st.session_state.user_info['mail'] + user_mail = st.session_state.user_info["mail"] except KeyError as e: st.write("Session Expired.") st.stop() @@ -50,7 +53,7 @@ except KeyError as e: # fields.rename(columns={'PROMPT': 'Interrogation Question?'}, inplace = True) # fields.rename(columns={'GROUP_ID': 'PRIORITY'}, inplace = True) # fields.rename(columns={'FIELD_NAME': 'SF_DB_COL_NAME'}, inplace = True) -# fields.rename(columns={'FM_MODEL_ID': 'llm_selected'}, inplace = True) +# fields.rename(columns={'FM_MODEL_ID': 'llm_selected'}, inplace = True) # except Exception as e: # st.write("Unable to fetch data from Snowflake: ",e) # # fields = pd.read_csv('contract_fields.csv', encoding='unicode_escape', skipinitialspace=True) @@ -60,15 +63,17 @@ client_row = st.columns([0.2, 0.7, 0.1]) with client_row[0]: st.write("**Client Name**") with client_row[1]: - client = st.selectbox('Client Name',(constants.client_list), label_visibility = "collapsed", index= None) + client = st.selectbox( + "Client Name", (constants.client_list), label_visibility="collapsed", index=None + ) if client: try: - folder_path = os.path.join(os.getcwd(), f'temp/{client}') + folder_path = os.path.join(os.getcwd(), f"temp/{client}") files = os.listdir(folder_path) for f in files: - print(f.lower().endswith('.pdf')) - file_list = [f for f in files if f.lower().endswith('.pdf')] + print(f.lower().endswith(".pdf")) + file_list = [f for f in files if f.lower().endswith(".pdf")] except Exception as e: print(f"An error occurred: {e}") file_list = [] @@ -78,16 +83,25 @@ if client: with file_row[0]: st.write("**Contract Name**") with file_row[1]: - file_name = st.selectbox('Select a file', ['All'] + contract_list, label_visibility = "collapsed", index= None) + file_name = st.selectbox( + "Select a file", + ["All"] + contract_list, + label_visibility="collapsed", + index=None, + ) field_row = st.columns([0.2, 0.7, 0.1]) with field_row[0]: st.write("**Field Group**") with field_row[1]: - field_group = st.selectbox('Field Group',('Unique Key', 'Pricing Before Carveouts', 'Contract Related'), - label_visibility = "collapsed", index = None) - - button_cols = st.columns([1,1,8]) + field_group = st.selectbox( + "Field Group", + ("Unique Key", "Pricing Before Carveouts", "Contract Related"), + label_visibility="collapsed", + index=None, + ) + + button_cols = st.columns([1, 1, 8]) with button_cols[0]: if st.button("Show PDF"): if file_name == None or file_name == "All": @@ -106,31 +120,95 @@ if client: ) pdf_viewer(os.path.join(folder_path, file_name), width=1500) with button_cols[1]: - df2 = pd.read_csv('temp2.csv') + df2 = pd.read_csv("temp2.csv") if st.button("Show Results"): try: - csv_name = re.sub(r'\.pdf', '.csv', os.path.join(folder_path, file_name), flags=re.IGNORECASE) + csv_name = re.sub( + r"\.pdf", + ".csv", + os.path.join(folder_path, file_name), + flags=re.IGNORECASE, + ) df2 = pd.read_csv(csv_name) except: - df2 = pd.DataFrame(columns=['Filename','Agreement_Name (Contract Title)','PAYER NAME','Health Plan State','Affiliate (Y/N)','Credentialing Application Indicator','Term Clause','Evergreen, Fixed or Hard Term','Termination Date','Termination Upon Notice - Days','Termination With Cause - Days','Amend Contract Upon notice Flag (Y/N)','Timeframe  to Object - Days','Assignments Clause  (Y/N)','Contract Effective Date','IRS #','IRS_Name' - ,'NPI (10-digits)','NPI_NAME','PROV_GROUP_TIN_SIGNATORY','PROV_TIN_OTHER','PROV_NPI_OTHER','Notice to Provider Name','Notice to Provider Address','Sequestration Language','Sequestration Reductions, included [Medicare only] (Y/N)','PROV_TIN_OTHER.1','PROV_NPI_OTHER.1','Parent Agreement Code','Pages','page_num', - 'Attachment/Exhibit','Line of Business','Provider Type','Provider Type - Level 2','Service Type','Plan Type','Lesser of Logic Language, included (Y/N)','Lesser of Rate','Reimb. Methodology','Reimb. Methodology_short','If rate is % of Payer or MCR [STANDARD]','If rate is % of Payer or MCR [STANDARD]_Short','If rate is Flat Fee [STANDARD]','Default Term','Default Rate','Inclusion of essential RBRVS "Fee Source" Language (Y/N)','CDM Neutralization Language, included (Y/N)','Chargemaster Protection Language','Exclusions','Not to Exceed','Escalator or COLA (Y/N)','Escalator I, Eff. Date','IP/OP','IP - DSH/IME/UC, included (Y/N)','IP - Stoploss Catastrophic Threshold']) - - df2['Imputed Value'] = '' - df2.to_csv('temp2.csv', index=False) + df2 = pd.DataFrame( + columns=[ + "Filename", + "Agreement_Name (Contract Title)", + "PAYER NAME", + "Health Plan State", + "Affiliate (Y/N)", + "Credentialing Application Indicator", + "Term Clause", + "Evergreen, Fixed or Hard Term", + "Termination Date", + "Termination Upon Notice - Days", + "Termination With Cause - Days", + "Amend Contract Upon notice Flag (Y/N)", + "Timeframe  to Object - Days", + "Assignments Clause  (Y/N)", + "Contract Effective Date", + "IRS #", + "IRS_Name", + "NPI (10-digits)", + "NPI_NAME", + "PROV_GROUP_TIN_SIGNATORY", + "PROV_TIN_OTHER", + "PROV_NPI_OTHER", + "Notice to Provider Name", + "Notice to Provider Address", + "Sequestration Language", + "Sequestration Reductions, included [Medicare only] (Y/N)", + "PROV_TIN_OTHER.1", + "PROV_NPI_OTHER.1", + "Parent Agreement Code", + "Pages", + "page_num", + "Attachment/Exhibit", + "Line of Business", + "Provider Type", + "Provider Type - Level 2", + "Service Type", + "Plan Type", + "Lesser of Logic Language, included (Y/N)", + "Lesser of Rate", + "Reimb. Methodology", + "Reimb. Methodology_short", + "If rate is % of Payer or MCR [STANDARD]", + "If rate is % of Payer or MCR [STANDARD]_Short", + "If rate is Flat Fee [STANDARD]", + "Default Term", + "Default Rate", + 'Inclusion of essential RBRVS "Fee Source" Language (Y/N)', + "CDM Neutralization Language, included (Y/N)", + "Chargemaster Protection Language", + "Exclusions", + "Not to Exceed", + "Escalator or COLA (Y/N)", + "Escalator I, Eff. Date", + "IP/OP", + "IP - DSH/IME/UC, included (Y/N)", + "IP - Stoploss Catastrophic Threshold", + ] + ) + + df2["Imputed Value"] = "" + df2.to_csv("temp2.csv", index=False) edited_df = st.data_editor(df2) - @st.cache_data + @st.cache_data def convert_df(df): - return df.to_csv(index=False).encode('utf-8') + return df.to_csv(index=False).encode("utf-8") csv = convert_df(edited_df) buttons = st.columns(3) with buttons[0]: - # st.button("Save All Imputations") - st.download_button("Download Table", csv, "file.csv", "text/csv", key='download-csv') + # st.button("Save All Imputations") + st.download_button( + "Download Table", csv, "file.csv", "text/csv", key="download-csv" + ) with buttons[1]: # st.download_button("Download Table", csv, "file.csv", "text/csv", key='download-csv') - st.write("") \ No newline at end of file + st.write("") diff --git a/terminal.py b/terminal.py index 0b0f224..be1a382 100644 --- a/terminal.py +++ b/terminal.py @@ -9,7 +9,9 @@ from typing import Callable, Optional, Dict, Tuple, List LOGGER = logging.getLogger(__name__) LOGGER.setLevel(logging.INFO) handler = logging.StreamHandler() -handler.setFormatter(logging.Formatter('%(asctime)s - %(name)s - %(levelname)s - %(message)s')) +handler.setFormatter( + logging.Formatter("%(asctime)s - %(name)s - %(levelname)s - %(message)s") +) LOGGER.addHandler(handler) @@ -18,11 +20,11 @@ class SensitiveFormatter(logging.Formatter): @staticmethod def mask_password(log: str) -> str: - return re.sub(r'password=([^\s]+)', r'password=*****', log) + return re.sub(r"password=([^\s]+)", r"password=*****", log) @staticmethod def mask_api_key(log: str) -> str: - return re.sub(r'api_key=([^\s]+)', r'api_key=*****', log) + return re.sub(r"api_key=([^\s]+)", r"api_key=*****", log) @staticmethod def _mask(s: str) -> str: @@ -41,7 +43,7 @@ def prepare_logger() -> logging.Logger: return LOGGER LOGGER = logging.getLogger(__name__) LOGGER.setLevel(logging.INFO) - log_format = '%(asctime)s %(filename)s:%(lineno)-4s [%(levelname)s] %(message)s' + log_format = "%(asctime)s %(filename)s:%(lineno)-4s [%(levelname)s] %(message)s" handler = logging.StreamHandler() handler.setFormatter(SensitiveFormatter(log_format)) LOGGER.addHandler(handler) @@ -61,43 +63,67 @@ def signal_handler(sig, frame, process): process.terminate() -def run_command(command: str, log_output: bool = False, decorate_logs: bool = True, log_cmd: bool = False, - log_prefix: str = "", envs: Optional[str] = None, secrets: Optional[Dict[str, str]] = None, - failure_callback: Optional[Callable[[str], None]] = None, - cwd: Optional[str] = None) -> Tuple[int, List[str]]: +def run_command( + command: str, + log_output: bool = False, + decorate_logs: bool = True, + log_cmd: bool = False, + log_prefix: str = "", + envs: Optional[str] = None, + secrets: Optional[Dict[str, str]] = None, + failure_callback: Optional[Callable[[str], None]] = None, + cwd: Optional[str] = None, +) -> Tuple[int, List[str]]: new_env = os.environ.copy() if secrets: for key, value in secrets.items(): new_env[key] = os.path.expandvars(value) full_command = f"{envs} {command}" if envs else command - process = subprocess.Popen(full_command, shell=True, cwd=cwd, env=new_env, - stdout=subprocess.PIPE, stderr=subprocess.STDOUT, text=True) + process = subprocess.Popen( + full_command, + shell=True, + cwd=cwd, + env=new_env, + stdout=subprocess.PIPE, + stderr=subprocess.STDOUT, + text=True, + ) # Setup signal handler original_sigint_handler = signal.getsignal(signal.SIGINT) signal.signal(signal.SIGINT, lambda sig, frame: signal_handler(sig, frame, process)) log_prefix_env = os.getenv("LOG_PREFIX", "") - log_prefix_full = f"[{log_prefix_env}] {log_prefix}" if log_prefix_env else log_prefix + log_prefix_full = ( + f"[{log_prefix_env}] {log_prefix}" if log_prefix_env else log_prefix + ) if log_cmd: LOGGER.info( - f"{get_blue_shade(log_prefix_full)}{log_prefix_full}\033[0m: Running command: \033[1;34m{command}\033[0m") + f"{get_blue_shade(log_prefix_full)}{log_prefix_full}\033[0m: Running command: \033[1;34m{command}\033[0m" + ) if envs: LOGGER.info(f"Using additional envs: \033[1;34m{envs}\033[0m") if secrets: - LOGGER.info(f"Using additional environment variables with secrets: \033[1;34m{list(secrets.keys())}\033[0m") + LOGGER.info( + f"Using additional environment variables with secrets: \033[1;34m{list(secrets.keys())}\033[0m" + ) cmd_output = [] while True: - output = process.stdout.readline() if process.stdout else '' + output = process.stdout.readline() if process.stdout else "" if output: output_line = output.strip() if log_output: log_statement = ( - f"{get_blue_shade(log_prefix_full)}{log_prefix_full}\033[0m: " - f"{output_line}") if log_prefix else output_line + ( + f"{get_blue_shade(log_prefix_full)}{log_prefix_full}\033[0m: " + f"{output_line}" + ) + if log_prefix + else output_line + ) LOGGER.info(log_statement) if decorate_logs else print(log_statement) cmd_output.append(output_line) elif process.poll() is not None: diff --git a/terraform.py b/terraform.py index 3ab19f9..aaf2679 100644 --- a/terraform.py +++ b/terraform.py @@ -14,15 +14,37 @@ LOGGER = prepare_logger() LOGGER.info(sys.argv) -parser = argparse.ArgumentParser(description='Run Terraform init with dynamic backend configuration.') -parser.add_argument('terraform_command', type=str, choices=['plan', 'apply'], - help='The Terraform command to execute (plan, apply, init).') -parser.add_argument('--environment', type=str, required=True, help='The environment name i.e. dev, prod.') -parser.add_argument('--module', type=str, required=False, help='The module name.') -parser.add_argument('--init-extra-args', type=str, required=False, help='Terraform init extra arguments') -parser.add_argument('--plan-extra-args', type=str, required=False, help='Terraform plan extra arguments') -parser.add_argument('--apply-extra-args', type=str, required=False, help='Terraform plan extra arguments') -parser.add_argument('--threads', type=int, default=4, help='Number of threads to use (default: 4)') +parser = argparse.ArgumentParser( + description="Run Terraform init with dynamic backend configuration." +) +parser.add_argument( + "terraform_command", + type=str, + choices=["plan", "apply"], + help="The Terraform command to execute (plan, apply, init).", +) +parser.add_argument( + "--environment", + type=str, + required=True, + help="The environment name i.e. dev, prod.", +) +parser.add_argument("--module", type=str, required=False, help="The module name.") +parser.add_argument( + "--init-extra-args", type=str, required=False, help="Terraform init extra arguments" +) +parser.add_argument( + "--plan-extra-args", type=str, required=False, help="Terraform plan extra arguments" +) +parser.add_argument( + "--apply-extra-args", + type=str, + required=False, + help="Terraform plan extra arguments", +) +parser.add_argument( + "--threads", type=int, default=4, help="Number of threads to use (default: 4)" +) args = parser.parse_args() @@ -34,12 +56,14 @@ apply_extra_args = args.apply_extra_args terraform_threads = args.threads module_name = args.module -bitbucket_ci = os.getenv('BITBUCKET_CI') +bitbucket_ci = os.getenv("BITBUCKET_CI") def change_dir(path): if path: - _, git_root = run_command('git rev-parse --show-toplevel', log_cmd=True, log_output=True) + _, git_root = run_command( + "git rev-parse --show-toplevel", log_cmd=True, log_output=True + ) LOGGER.info(f"git_root={git_root}") absolute_path = os.path.join(git_root[0], path) return absolute_path @@ -50,7 +74,7 @@ def change_dir(path): def select_vars_file(): - tf_vars_file = os.getenv('TFVARS_FILE_PATH', f'vars-{environment}.tfvars') + tf_vars_file = os.getenv("TFVARS_FILE_PATH", f"vars-{environment}.tfvars") if tf_vars_file: LOGGER.info(f"Using {tf_vars_file} as vars-file...") else: @@ -59,34 +83,43 @@ def select_vars_file(): return tf_vars_file -def prepare_tf_init_command(stack_path, init_extra_args, state_environment, tf_state_bucket_name): - command = (f'terraform init --reconfigure ' - f'-backend-config="dynamodb_table=doczyai-use2-{get_environment_value(environment)}-infra-dyd-terraform-lock" ' - f'-backend-config="bucket=doczyai-use2-{get_environment_value(environment)}-infra-s3-terraform-state" ' - f'-backend-config="key=terraform/{stack_path}/terraform.tfstate" -backend-config="region=us-east-2"') +def prepare_tf_init_command( + stack_path, init_extra_args, state_environment, tf_state_bucket_name +): + command = ( + f"terraform init --reconfigure " + f'-backend-config="dynamodb_table=doczyai-use2-{get_environment_value(environment)}-infra-dyd-terraform-lock" ' + f'-backend-config="bucket=doczyai-use2-{get_environment_value(environment)}-infra-s3-terraform-state" ' + f'-backend-config="key=terraform/{stack_path}/terraform.tfstate" -backend-config="region=us-east-2"' + ) if init_extra_args: command = f"{command} {init_extra_args}" return command def get_environment_value(environment): - environment_map = { - 'prod': 'p', - 'uat': 'u', - 'qa': 'q', - 'dev': 'd' - } + environment_map = {"prod": "p", "uat": "u", "qa": "q", "dev": "d"} - return environment_map.get(environment, None) # Returns None if the environment is not found + return environment_map.get( + environment, None + ) # Returns None if the environment is not found -def prepare_tf_main_command(environment, stack_path, state_environment, tf_state_bucket_name, vars_file, - plan_extra_args, - apply_extra_args): +def prepare_tf_main_command( + environment, + stack_path, + state_environment, + tf_state_bucket_name, + vars_file, + plan_extra_args, + apply_extra_args, +): common_command = f'-var-file=vars-{environment}.tfvars -var "aws_region=us-east-2" -var "environment={environment}"' if raw_terraform_command == "plan": - path_part = stack_path.split('/')[1] if '/' in stack_path else stack_path - command = f"terraform plan -out .{path_part}.{environment}.tfplan {common_command}" + path_part = stack_path.split("/")[1] if "/" in stack_path else stack_path + command = ( + f"terraform plan -out .{path_part}.{environment}.tfplan {common_command}" + ) if plan_extra_args: command = f"{command} {plan_extra_args}" elif raw_terraform_command == "apply": @@ -96,15 +129,23 @@ def prepare_tf_main_command(environment, stack_path, state_environment, tf_state return command -tf_root_module_path = os.getenv('TF_ROOT_MODULE_PATH') if not module_name else module_name +tf_root_module_path = ( + os.getenv("TF_ROOT_MODULE_PATH") if not module_name else module_name +) def do_terraform(stack_path): LOGGER.info(f"[{stack_path}]: stack_path={stack_path}") - is_deploy_module_path = is_deploy_module(stack_path, '') - LOGGER.info(f"[{stack_path}]: is_deploy_module_path={decorate_white_bold(is_deploy_module_path)}") + is_deploy_module_path = is_deploy_module(stack_path, "") + LOGGER.info( + f"[{stack_path}]: is_deploy_module_path={decorate_white_bold(is_deploy_module_path)}" + ) if not is_deploy_module_path: - LOGGER.info(decorate_warn(f"[{stack_path}]: Module is not found in git changes. Skipping deploy.")) + LOGGER.info( + decorate_warn( + f"[{stack_path}]: Module is not found in git changes. Skipping deploy." + ) + ) return True, "OK" absolute_path = change_dir(stack_path) @@ -112,23 +153,42 @@ def do_terraform(stack_path): state_environment = "dev-new" if environment == "dev" else environment tf_state_bucket_name = f"pi2new-{state_environment}-terraform-state-file" - terraform_init_command = prepare_tf_init_command(stack_path, init_extra_args, state_environment, - tf_state_bucket_name) - rc_init, init_result = run_command(terraform_init_command, log_output=True, log_cmd=True, log_prefix=stack_path, - cwd=absolute_path) - + terraform_init_command = prepare_tf_init_command( + stack_path, init_extra_args, state_environment, tf_state_bucket_name + ) + rc_init, init_result = run_command( + terraform_init_command, + log_output=True, + log_cmd=True, + log_prefix=stack_path, + cwd=absolute_path, + ) if rc_init > 0: raise Exception(f"{stack_path}: Failed to run terraform init") vars_file = select_vars_file() - terraform_main_command = prepare_tf_main_command(environment, stack_path, state_environment, tf_state_bucket_name, - vars_file, plan_extra_args, apply_extra_args) + terraform_main_command = prepare_tf_main_command( + environment, + stack_path, + state_environment, + tf_state_bucket_name, + vars_file, + plan_extra_args, + apply_extra_args, + ) - rc_main, command_result = run_command(terraform_main_command, log_output=True, log_cmd=True, - log_prefix=stack_path, cwd=absolute_path) + rc_main, command_result = run_command( + terraform_main_command, + log_output=True, + log_cmd=True, + log_prefix=stack_path, + cwd=absolute_path, + ) if rc_init or rc_main > 0: - raise RuntimeError(f"Unexpected Terraform error for command: {terraform_main_command}") + raise RuntimeError( + f"Unexpected Terraform error for command: {terraform_main_command}" + ) else: return True, "OK" @@ -137,17 +197,23 @@ def do_terraform(stack_path): def find_paths_with_file(file_name): try: - completed_process = subprocess.run(['git', 'ls-files', f'*{file_name}'], check=True, text=True, - stdout=subprocess.PIPE) + completed_process = subprocess.run( + ["git", "ls-files", f"*{file_name}"], + check=True, + text=True, + stdout=subprocess.PIPE, + ) paths = set() for file_path in completed_process.stdout.splitlines(): - dir_path = '/'.join(file_path.split('/')[:-1]) + dir_path = "/".join(file_path.split("/")[:-1]) if dir_path.startswith("modules/"): continue if dir_path.startswith("infra/"): - dir_path = dir_path[len("infra/"):] - if dir_path.startswith("kinesis") or dir_path.startswith("aws-transfer-sftp"): + dir_path = dir_path[len("infra/") :] + if dir_path.startswith("kinesis") or dir_path.startswith( + "aws-transfer-sftp" + ): paths.add(dir_path) return paths @@ -174,10 +240,14 @@ def wait_for_dependencies(dependencies, completed, failed, path): pending_dependencies = [dep for dep in dependencies if dep not in completed] if any(dep in failed for dep in pending_dependencies): failed_deps = [dep for dep in pending_dependencies if dep in failed] - raise Exception(f"{path}: Cannot proceed. Dependencies failed: {failed_deps}") + raise Exception( + f"{path}: Cannot proceed. Dependencies failed: {failed_deps}" + ) if not pending_dependencies: return - LOGGER.warning(f"{path}: Waiting for dependencies to complete: {pending_dependencies}") + LOGGER.warning( + f"{path}: Waiting for dependencies to complete: {pending_dependencies}" + ) time.sleep(5) @@ -193,9 +263,15 @@ def execute_with_dependencies(paths): dependencies = read_dependencies(path) if dependencies: # Submit a future for waiting dependencies - dep_future = executor.submit(wait_for_dependencies, dependencies, completed, failed, path) + dep_future = executor.submit( + wait_for_dependencies, dependencies, completed, failed, path + ) # Submit the main task to be executed after dependencies are resolved - future = executor.submit(lambda p=dep_future, x=path: do_terraform(x) if p.result() is None else None) + future = executor.submit( + lambda p=dep_future, x=path: ( + do_terraform(x) if p.result() is None else None + ) + ) else: future = executor.submit(do_terraform, path) future_to_path[future] = path @@ -204,7 +280,9 @@ def execute_with_dependencies(paths): path = future_to_path[future] try: result = future.result() - if result is None or not result[0]: # Assuming do_terraform returns (success: bool, message: str) + if ( + result is None or not result[0] + ): # Assuming do_terraform returns (success: bool, message: str) errored_modules.add(path) failed[path] = "Failed due to internal error" # Mark as failed else: @@ -228,7 +306,7 @@ if __name__ == "__main__": if tf_root_module_path: do_terraform(tf_root_module_path) else: - paths = find_paths_with_file('provider.tf') + paths = find_paths_with_file("provider.tf") for path in paths: LOGGER.warning(f"Discovered: {path}") # TODO finish later if we want multi modulue deploy in one go diff --git a/textract-pipeline/src/lambda/batch-creation/index.py b/textract-pipeline/src/lambda/batch-creation/index.py index a957056..b00bea9 100644 --- a/textract-pipeline/src/lambda/batch-creation/index.py +++ b/textract-pipeline/src/lambda/batch-creation/index.py @@ -11,9 +11,10 @@ logger = logging.getLogger() logger.setLevel(logging.INFO) # Initialize S3 client -s3_client = boto3.client('s3') +s3_client = boto3.client("s3") DATABASE_LOGGING_LAMBDA_FUNCTION_NAME = "" + # Function to load configuration from S3 def load_config_from_s3(bucket_name, file_key): """ @@ -24,12 +25,14 @@ def load_config_from_s3(bucket_name, file_key): Returns: dict: Configuration dictionary parsed from the file. """ - logger.info(f"Loading configuration from S3. Bucket: {bucket_name}, File Key: {file_key}") - + logger.info( + f"Loading configuration from S3. Bucket: {bucket_name}, File Key: {file_key}" + ) + try: # Download the config file from S3 response = s3_client.get_object(Bucket=bucket_name, Key=file_key) - config_content = response['Body'].read().decode('utf-8') + config_content = response["Body"].read().decode("utf-8") # Parse the config file config_parser = ConfigParser() @@ -38,7 +41,9 @@ def load_config_from_s3(bucket_name, file_key): # Convert the configuration to a dictionary config_dict = {} for section in config_parser.sections(): - config_dict[section] = {key.upper(): value for key, value in config_parser.items(section)} + config_dict[section] = { + key.upper(): value for key, value in config_parser.items(section) + } logger.info("Configuration loaded successfully.") return config_dict @@ -56,17 +61,19 @@ def create_folder_in_s3(client_bucket_name, CONTRACTS_LANDNING_ZONE): Returns: str: Batch ID for the created folder. """ - logger.info(f"Creating folder in S3 bucket: {client_bucket_name}, Landing Zone: {CONTRACTS_LANDNING_ZONE}") + logger.info( + f"Creating folder in S3 bucket: {client_bucket_name}, Landing Zone: {CONTRACTS_LANDNING_ZONE}" + ) # Generate timestamp ID - timestamp_id = datetime.now().strftime('%d%m%y%H%M%S') + timestamp_id = datetime.now().strftime("%d%m%y%H%M%S") # Generate Batch ID - batch_id = "batch_"+timestamp_id + batch_id = "batch_" + timestamp_id # Folder key (name) in S3 bucket folder_key = f"{CONTRACTS_LANDNING_ZONE}".format(batch_id) - + # Create the folder in S3 bucket try: s3_client.put_object(Bucket=client_bucket_name, Key=folder_key) @@ -83,68 +90,82 @@ def lambda_handler(event, context): """ logger.info(f"Received event: {event}") - if 'client-bucket-name' in event: - client_bucket_name = event['client-bucket-name'] + if "client-bucket-name" in event: + client_bucket_name = event["client-bucket-name"] + + client_id = event.get("client_id", "default_client") - client_id = event.get('client_id', 'default_client') - # Read environment variables - property_file_path = os.environ.get('PROPERTY_FILE_S3_PATH', '') + property_file_path = os.environ.get("PROPERTY_FILE_S3_PATH", "") file_path_array = property_file_path.split("/") - + global DATABASE_LOGGING_LAMBDA_FUNCTION_NAME - DATABASE_LOGGING_LAMBDA_FUNCTION_NAME = os.environ.get('DATABASE_LOGGING_LAMBDA_FUNCTION_NAME', '') - + DATABASE_LOGGING_LAMBDA_FUNCTION_NAME = os.environ.get( + "DATABASE_LOGGING_LAMBDA_FUNCTION_NAME", "" + ) + # Extract config_file_path CONFIG_FILE_PATH = "/".join(file_path_array[1:]) # Load config file try: config_dict = load_config_from_s3(client_bucket_name, CONFIG_FILE_PATH) - CONTRACTS_LANDNING_ZONE = config_dict['FOLDER_LOCATIONS']['CONTRACTS_LANDNING_ZONE'] + CONTRACTS_LANDNING_ZONE = config_dict["FOLDER_LOCATIONS"][ + "CONTRACTS_LANDNING_ZONE" + ] batch_id = create_folder_in_s3(client_bucket_name, CONTRACTS_LANDNING_ZONE) - generate_batch_logs_input(batch_id,100,'',client_id) + generate_batch_logs_input(batch_id, 100, "", client_id) return { - 'statusCode': 200, - 'body': json.dumps({'batch_id': batch_id, 'landing_zone': CONTRACTS_LANDNING_ZONE.format(batch_id)}) + "statusCode": 200, + "body": json.dumps( + { + "batch_id": batch_id, + "landing_zone": CONTRACTS_LANDNING_ZONE.format(batch_id), + } + ), } except Exception as e: logger.error(f"Error processing event: {str(e)}") return { - 'statusCode': 500, - 'body': json.dumps({'error': 'Internal server error'}) + "statusCode": 500, + "body": json.dumps({"error": "Internal server error"}), } else: logger.error("'client-bucket-name' is missing in the input.") return { - 'statusCode': 400, - 'body': json.dumps({'error': "'client-bucket-name' is missing in the input."}) + "statusCode": 400, + "body": json.dumps( + {"error": "'client-bucket-name' is missing in the input."} + ), } -def generate_batch_logs_input(batch_id,no_of_doc,user_name,client_id): - current_time = datetime.now().isoformat() - - logger.info('DATABASE_LOGGING_LAMBDA_FUNCTION_NAME: ' + str(DATABASE_LOGGING_LAMBDA_FUNCTION_NAME)) - - data = { +def generate_batch_logs_input(batch_id, no_of_doc, user_name, client_id): + current_time = datetime.now().isoformat() + + logger.info( + "DATABASE_LOGGING_LAMBDA_FUNCTION_NAME: " + + str(DATABASE_LOGGING_LAMBDA_FUNCTION_NAME) + ) + + data = { "operation": "insert", "table": "BATCH_LOGS", "data": { "BATCH_ID": batch_id, - "EXECUTION_START_TIME" : current_time, - "NO_OF_DOCUMENTS" : no_of_doc, + "EXECUTION_START_TIME": current_time, + "NO_OF_DOCUMENTS": no_of_doc, "USER_NAME": user_name, - "CLIENT_ID" : client_id - } - } - logger.info(f"Request: {data}") - - lambda_client = boto3.client('lambda') - response = lambda_client.invoke( - FunctionName=DATABASE_LOGGING_LAMBDA_FUNCTION_NAME, - InvocationType='Event', # Asynchronous invocation - Payload=json.dumps(data).encode('utf-8') - ) - - logger.info("Generated document logs input successfully") - return response \ No newline at end of file + "CLIENT_ID": client_id, + }, + } + logger.info(f"Request: {data}") + + lambda_client = boto3.client("lambda") + response = lambda_client.invoke( + FunctionName=DATABASE_LOGGING_LAMBDA_FUNCTION_NAME, + InvocationType="Event", # Asynchronous invocation + Payload=json.dumps(data).encode("utf-8"), + ) + + logger.info("Generated document logs input successfully") + return response diff --git a/textract-pipeline/src/lambda/call-bedrock/index.py b/textract-pipeline/src/lambda/call-bedrock/index.py index e3463a8..f5c4eae 100644 --- a/textract-pipeline/src/lambda/call-bedrock/index.py +++ b/textract-pipeline/src/lambda/call-bedrock/index.py @@ -8,6 +8,7 @@ from botocore.exceptions import ClientError logger = logging.getLogger() logger.setLevel(logging.INFO) + # Lambda function handler def lambda_handler(event, context): logger.info("Lambda function invoked.") @@ -20,20 +21,23 @@ def lambda_handler(event, context): logger.error("Too many requests. Please wait before trying again.") raise + # Function to invoke the bedrock_llm with necessary parameters def invoke(wrapper, event): logger.info(f"Event details: {event}") # Extract parameters from the event - prompt = event['prompt'] - model_id = event['model_id'] - max_gen_len = event['max_gen_len'] - temperature = event['temperature'] - top_p = event['top_p'] + prompt = event["prompt"] + model_id = event["model_id"] + max_gen_len = event["max_gen_len"] + temperature = event["temperature"] + top_p = event["top_p"] try: # Invoke the bedrock_llm using the provided wrapper - completion = wrapper.invoke_llm(prompt, model_id, max_gen_len, temperature, top_p) + completion = wrapper.invoke_llm( + prompt, model_id, max_gen_len, temperature, top_p + ) logger.info("Bedrock LLM invoked successfully.") return completion @@ -42,28 +46,30 @@ def invoke(wrapper, event): logger.exception(f"Couldn't invoke model {model_id}. Error: {str(e)}") raise + # Function to call the Bedrock LLM def call_bedrock_llm(event): logger.info("Calling Bedrock LLM.") # Initialize Bedrock Runtime client client = boto3.client(service_name="bedrock-runtime", region_name="us-east-1") - + # Create an instance of BedrockRuntimeWrapper wrapper = BedrockRuntimeWrapper(client) - + # Invoke the wrapper and return the result try: answer = invoke(wrapper, event) logger.info("Bedrock LLM call completed.") return answer except ClientError as e: - if e.response['Error']['Code'] == 'ThrottlingException': + if e.response["Error"]["Code"] == "ThrottlingException": raise ThrottlingException else: logger.error(f"Error calling Bedrock LLM: {str(e)}") raise + # Class to wrap Bedrock Runtime functionality class BedrockRuntimeWrapper: # Constructor to initialize the wrapper with a Bedrock Runtime client @@ -76,9 +82,13 @@ class BedrockRuntimeWrapper: try: # Prepare the request body - if 'claude' in model_id.lower(): + if "claude" in model_id.lower(): # If yes, change the parameter name to 'max_tokens_to_sample' for claude - prompt='Human:'+prompt+'\n\nAssistant:You read and understand the USA healthcrae contract and able to answer questions based on given contract.' + prompt = ( + "Human:" + + prompt + + "\n\nAssistant:You read and understand the USA healthcrae contract and able to answer questions based on given contract." + ) body = { "prompt": prompt, "temperature": temperature, @@ -98,7 +108,7 @@ class BedrockRuntimeWrapper: response = self.bedrock_runtime_client.invoke_model( modelId=model_id, body=json.dumps(body) ) - + # Parse the response body from JSON response_body = json.loads(response["body"].read()) logger.info(f"Response Body from Bedrock LLM: {response_body}") @@ -111,5 +121,6 @@ class BedrockRuntimeWrapper: logger.error(f"Problem in invoking model. Error: {str(e)}") raise + class ThrottlingException(Exception): pass diff --git a/textract-pipeline/src/lambda/database-interface-get/index.py b/textract-pipeline/src/lambda/database-interface-get/index.py index dab457d..8155638 100644 --- a/textract-pipeline/src/lambda/database-interface-get/index.py +++ b/textract-pipeline/src/lambda/database-interface-get/index.py @@ -1,14 +1,14 @@ import json import snowflake.connector -import boto3 +import boto3 import snowflake.connector import logging from datetime import datetime import os -logging.basicConfig(level=logging.INFO, format='%(levelname)s: %(message)s', force=True) -logging.getLogger('snowflake.connector').setLevel(logging.WARNING) +logging.basicConfig(level=logging.INFO, format="%(levelname)s: %(message)s", force=True) +logging.getLogger("snowflake.connector").setLevel(logging.WARNING) logging.getLogger("botocore").setLevel(logging.WARNING) logger = logging.getLogger(__name__) logger.setLevel(logging.INFO) @@ -43,21 +43,23 @@ json.loads(return_object['body']) query column values: json_event[0]['AUDIT_SID'] """ + # Function to convert non-serializable types def default_converter(o): if isinstance(o, datetime): return o.isoformat() raise TypeError("Type not serializable") + def get_secret(secrets_name: str): """Get credentials from Secret Manager as dict""" - secrets_manager = boto3.client('secretsmanager') + secrets_manager = boto3.client("secretsmanager") get_secret_value_response = secrets_manager.get_secret_value(SecretId=secrets_name) - if 'SecretString' in get_secret_value_response: - secret_json = get_secret_value_response['SecretString'] + if "SecretString" in get_secret_value_response: + secret_json = get_secret_value_response["SecretString"] else: - secret_json = base64.b64decode(get_secret_value_response['SecretBinary']) + secret_json = base64.b64decode(get_secret_value_response["SecretBinary"]) return json.loads(secret_json) @@ -66,16 +68,25 @@ def get_snowflake_db_connection(secrets_name: str): """Create connection to Snowflake db.""" try: con_params = get_secret(secrets_name) - account = con_params['account_locator'] - user = con_params['user'] - password = con_params['password'] - database = con_params['database'].upper() - warehouse = con_params['warehouse'] - role= con_params['role'] - logger.info(f'Using credentials: account={account}, user={user}, password=***, database={database}, ' - f'warehouse={warehouse}') - snowflake_connection = snowflake.connector.connect(account=account, user=user, password=password, database=database, - warehouse=warehouse, autocommit=True, schema='STG') + account = con_params["account_locator"] + user = con_params["user"] + password = con_params["password"] + database = con_params["database"].upper() + warehouse = con_params["warehouse"] + role = con_params["role"] + logger.info( + f"Using credentials: account={account}, user={user}, password=***, database={database}, " + f"warehouse={warehouse}" + ) + snowflake_connection = snowflake.connector.connect( + account=account, + user=user, + password=password, + database=database, + warehouse=warehouse, + autocommit=True, + schema="STG", + ) logger.info(snowflake_connection) cursor = snowflake_connection.cursor() cursor.execute("SELECT * FROM STG.DIM_AUDIT LIMIT 10;") @@ -84,20 +95,19 @@ def get_snowflake_db_connection(secrets_name: str): except Exception as e: return e -SECRET_MANAGER_NAME = os.environ.get('SECRET_MANAGER_NAME', '') + +SECRET_MANAGER_NAME = os.environ.get("SECRET_MANAGER_NAME", "") conn = get_snowflake_db_connection(SECRET_MANAGER_NAME) - def lambda_handler(event, context): - - + cursor = conn.cursor() try: # Execute the SQL query provided in the event payload - cursor.execute(event['query']) - + cursor.execute(event["query"]) + # Extract column names from the cursor description column_names = [col[0] for col in cursor.description] logger.info(f"column_names:: {column_names}") @@ -106,18 +116,14 @@ def lambda_handler(event, context): result = [dict(zip(column_names, row)) for row in rows] return { - 'statusCode': 200, - 'body': json.dumps(result, default=default_converter) + "statusCode": 200, + "body": json.dumps(result, default=default_converter), } except Exception as e: logger.error(f"An error occurred: {str(e)}") - return { - 'statusCode': 400, - 'body': json.dumps(str(e)) - } + return {"statusCode": 400, "body": json.dumps(str(e))} finally: pass # Ensure that resources are cleaned up # cursor.close() # conn.close() - diff --git a/textract-pipeline/src/lambda/database-interface/index.py b/textract-pipeline/src/lambda/database-interface/index.py index 30594cd..999aa13 100644 --- a/textract-pipeline/src/lambda/database-interface/index.py +++ b/textract-pipeline/src/lambda/database-interface/index.py @@ -1,6 +1,6 @@ import json import snowflake.connector -import boto3 +import boto3 import snowflake.connector import logging import os @@ -95,8 +95,8 @@ Data is a dictionary with the columns and values to be inserted or updated """ -logging.basicConfig(level=logging.INFO, format='%(levelname)s: %(message)s', force=True) -logging.getLogger('snowflake.connector').setLevel(logging.WARNING) +logging.basicConfig(level=logging.INFO, format="%(levelname)s: %(message)s", force=True) +logging.getLogger("snowflake.connector").setLevel(logging.WARNING) logging.getLogger("botocore").setLevel(logging.WARNING) logger = logging.getLogger(__name__) logger.setLevel(logging.INFO) @@ -104,13 +104,13 @@ logger.setLevel(logging.INFO) def get_secret(secrets_name: str): """Get credentials from Secret Manager as dict""" - secrets_manager = boto3.client('secretsmanager') + secrets_manager = boto3.client("secretsmanager") get_secret_value_response = secrets_manager.get_secret_value(SecretId=secrets_name) - if 'SecretString' in get_secret_value_response: - secret_json = get_secret_value_response['SecretString'] + if "SecretString" in get_secret_value_response: + secret_json = get_secret_value_response["SecretString"] else: - secret_json = base64.b64decode(get_secret_value_response['SecretBinary']) + secret_json = base64.b64decode(get_secret_value_response["SecretBinary"]) return json.loads(secret_json) @@ -119,16 +119,26 @@ def get_snowflake_db_connection(secrets_name: str): """Create connection to Snowflake db.""" try: con_params = get_secret(secrets_name) - account = con_params['account_locator'] - user = con_params['user'] - password = con_params['password'] - database = con_params['database'].upper() - warehouse = con_params['warehouse'] - role= con_params['role'] - logger.info(f'Using credentials: account={account}, user={user}, password=***, database={database}, ' - f'warehouse={warehouse}') - snowflake_connection = snowflake.connector.connect(account=account, user=user, password=password, database=database, - warehouse=warehouse,role=role, autocommit=True, schema='STG') + account = con_params["account_locator"] + user = con_params["user"] + password = con_params["password"] + database = con_params["database"].upper() + warehouse = con_params["warehouse"] + role = con_params["role"] + logger.info( + f"Using credentials: account={account}, user={user}, password=***, database={database}, " + f"warehouse={warehouse}" + ) + snowflake_connection = snowflake.connector.connect( + account=account, + user=user, + password=password, + database=database, + warehouse=warehouse, + role=role, + autocommit=True, + schema="STG", + ) logger.info(snowflake_connection) cursor = snowflake_connection.cursor() cursor.execute("SELECT * FROM STG.DIM_AUDIT LIMIT 10;") @@ -136,61 +146,101 @@ def get_snowflake_db_connection(secrets_name: str): return snowflake_connection except Exception as e: return e - + # Leaving this statement outside the lambda_handler function to reuse the connection - # Secret has been setup to use the logging service account -SECRET_MANAGER_NAME = os.environ.get('SECRET_MANAGER_NAME', '') +# Secret has been setup to use the logging service account +SECRET_MANAGER_NAME = os.environ.get("SECRET_MANAGER_NAME", "") conn = get_snowflake_db_connection(SECRET_MANAGER_NAME) # cur = conn.cursor() # logger.info(f"CURSOR OBJECT: {cur}") + def construct_doc_insert_sql(data): """ Constructs the SQL for an insert operation Sample return value: - INSERT INTO STG.DOCUMENT_LOGS (DOCUMENT_ID,BATCH_ID, JOB_ID, STAGE, TEXTRACT_STATUS, BUCKET_NAME, FILE_NAME, FILE_PATH, DOCUMENT_TYPE, PAYER_SIGNED, PROVIDER_SIGNED, GROUP_ID, CREATED_TIME, MODIFIED_TIME, CREATED_BY, MODIFIED_BY, ORIGINAL_FILE_EXTENSION, NO_OF_PAGES, FILE_SIZE) + INSERT INTO STG.DOCUMENT_LOGS (DOCUMENT_ID,BATCH_ID, JOB_ID, STAGE, TEXTRACT_STATUS, BUCKET_NAME, FILE_NAME, FILE_PATH, DOCUMENT_TYPE, PAYER_SIGNED, PROVIDER_SIGNED, GROUP_ID, CREATED_TIME, MODIFIED_TIME, CREATED_BY, MODIFIED_BY, ORIGINAL_FILE_EXTENSION, NO_OF_PAGES, FILE_SIZE) VALUES ('doc_1212','batch_101', 'J123456', 'Processing', 'Success', 'doc-bucket', 'file1.pdf', '/documents/2023/', 'Report', True, False, 'G100', '2024-03-21 10:00:00', '2024-03-21 10:00:00', 'admin', 'admin', 'pdf', 10, 1048576); """ - columns = ', '.join(data.keys()) - values = ', '.join(["'" + str(value).replace("'", "''") + "'" if isinstance(value, str) else str(value) for value in data.values()]) + columns = ", ".join(data.keys()) + values = ", ".join( + [ + ( + "'" + str(value).replace("'", "''") + "'" + if isinstance(value, str) + else str(value) + ) + for value in data.values() + ] + ) sql = f"INSERT INTO STG.DOCUMENT_LOGS ({columns}) VALUES ({values});" return sql + def construct_doczy_pipeline_insert_sql(data): """ Constructs the SQL for an insert operation Sample return value: - INSERT INTO STG.DOCUMENT_LOGS (DOCUMENT_ID,BATCH_ID, JOB_ID, STAGE, TEXTRACT_STATUS, BUCKET_NAME, FILE_NAME, FILE_PATH, DOCUMENT_TYPE, PAYER_SIGNED, PROVIDER_SIGNED, GROUP_ID, CREATED_TIME, MODIFIED_TIME, CREATED_BY, MODIFIED_BY, ORIGINAL_FILE_EXTENSION, NO_OF_PAGES, FILE_SIZE) + INSERT INTO STG.DOCUMENT_LOGS (DOCUMENT_ID,BATCH_ID, JOB_ID, STAGE, TEXTRACT_STATUS, BUCKET_NAME, FILE_NAME, FILE_PATH, DOCUMENT_TYPE, PAYER_SIGNED, PROVIDER_SIGNED, GROUP_ID, CREATED_TIME, MODIFIED_TIME, CREATED_BY, MODIFIED_BY, ORIGINAL_FILE_EXTENSION, NO_OF_PAGES, FILE_SIZE) VALUES ('doc_1212','batch_101', 'J123456', 'Processing', 'Success', 'doc-bucket', 'file1.pdf', '/documents/2023/', 'Report', True, False, 'G100', '2024-03-21 10:00:00', '2024-03-21 10:00:00', 'admin', 'admin', 'pdf', 10, 1048576); """ - columns = ', '.join(data.keys()) - values = ', '.join(["'" + str(value).replace("'", "''") + "'" if isinstance(value, str) else str(value) for value in data.values()]) + columns = ", ".join(data.keys()) + values = ", ".join( + [ + ( + "'" + str(value).replace("'", "''") + "'" + if isinstance(value, str) + else str(value) + ) + for value in data.values() + ] + ) sql = f"INSERT INTO STG.DOCZY_PIPELINE_RAW_OUTPUT ({columns}) VALUES ({values});" return sql + def construct_doczy_pipeline_update_sql(data, document_id): """ Constructs the SQL for an update operation - Sample return value: + Sample return value: UPDATE STG.DOCUMENT_LOGS SET TEXTRACT_STATUS = 'Failed', MODIFIED_TIME = '2024-03-22 15:00:00', MODIFIED_BY = 'admin' WHERE DOCUMENT_ID = 1001; """ - set_clauses = ', '.join([f"{key} = '" + str(value).replace("'", "''") + "'" if isinstance(value, str) else f"{key} = {value}" for key, value in data.items()]) + set_clauses = ", ".join( + [ + ( + f"{key} = '" + str(value).replace("'", "''") + "'" + if isinstance(value, str) + else f"{key} = {value}" + ) + for key, value in data.items() + ] + ) sql = f"UPDATE STG.DOCZY_PIPELINE_RAW_OUTPUT SET {set_clauses} WHERE DOCUMENT_ID = '{document_id}';" logger.info(f"Executing the following logging SQL statement: {sql}") return sql + def construct_doc_update_sql(data, document_id): """ Constructs the SQL for an update operation - Sample return value: + Sample return value: UPDATE STG.DOCUMENT_LOGS SET TEXTRACT_STATUS = 'Failed', MODIFIED_TIME = '2024-03-22 15:00:00', MODIFIED_BY = 'admin' WHERE DOCUMENT_ID = 1001; """ - set_clauses = ', '.join([f"{key} = '" + str(value).replace("'", "''") + "'" if isinstance(value, str) else f"{key} = {value}" for key, value in data.items()]) + set_clauses = ", ".join( + [ + ( + f"{key} = '" + str(value).replace("'", "''") + "'" + if isinstance(value, str) + else f"{key} = {value}" + ) + for key, value in data.items() + ] + ) sql = f"UPDATE STG.DOCUMENT_LOGS SET {set_clauses} WHERE DOCUMENT_ID = '{document_id}';" logger.info(f"Executing the following logging SQL statement: {sql}") return sql @@ -203,11 +253,21 @@ def construct_batch_insert_sql(data): Sample return value: INSERT INTO STG.BATCH_LOGS (BATCH_ID,CLIENT_ID, EXECUTION_START_TIME, NO_OF_DOCUMENTS, USER_NAME) VALUES ('C200', '2024-03-21 09:00:00', 150, 'batch_processor'); """ - columns = ', '.join(data.keys()) - values = ', '.join(["'" + str(value).replace("'", "''") + "'" if isinstance(value, str) else str(value) for value in data.values()]) + columns = ", ".join(data.keys()) + values = ", ".join( + [ + ( + "'" + str(value).replace("'", "''") + "'" + if isinstance(value, str) + else str(value) + ) + for value in data.values() + ] + ) sql = f"INSERT INTO STG.BATCH_LOGS ({columns}) VALUES ({values});" return sql + def construct_client_insert_sql(data): """ Constructs the SQL for an insert operation @@ -215,8 +275,17 @@ def construct_client_insert_sql(data): Sample return value: INSERT INTO STG.CLIENT_LOGS (CLIENT_ID, CLIENT_NAME, BUCKET_NAME) VALUES ('CL300', 'Acme Corporation', 'acme-docs'); """ - columns = ', '.join(data.keys()) - values = ', '.join(["'" + str(value).replace("'", "''") + "'" if isinstance(value, str) else str(value) for value in data.values()]) + columns = ", ".join(data.keys()) + values = ", ".join( + [ + ( + "'" + str(value).replace("'", "''") + "'" + if isinstance(value, str) + else str(value) + ) + for value in data.values() + ] + ) sql = f"INSERT INTO STG.CLIENT_LOGS ({columns}) VALUES ({values});" return sql @@ -228,10 +297,20 @@ def construct_client_update_sql(data, client_id): Sample return value: UPDATE STG.CLIENT_LOGS SET BUCKET_NAME = 'new-acme-docs' WHERE CLIENT_ID = 'CL300'; """ - set_clauses = ', '.join([f"{key} = '" + str(value).replace("'", "''") + "'" if isinstance(value, str) else f"{key} = {value}" for key, value in data.items()]) + set_clauses = ", ".join( + [ + ( + f"{key} = '" + str(value).replace("'", "''") + "'" + if isinstance(value, str) + else f"{key} = {value}" + ) + for key, value in data.items() + ] + ) sql = f"UPDATE STG.CLIENT_LOGS SET {set_clauses} WHERE CLIENT_ID = '{client_id}';" return sql + def construct_batch_update_sql(data, batch_id): """ Constructs the SQL for an update operation @@ -239,7 +318,16 @@ def construct_batch_update_sql(data, batch_id): Sample return value: UPDATE STG.BATCH_LOGS SET NO_OF_DOCUMENTS = 155, USER_NAME = 'updated_processor' WHERE BATCH_ID = 1; """ - set_clauses = ', '.join([f"{key} = '" + str(value).replace("'", "''") + "'" if isinstance(value, str) else f"{key} = {value}" for key, value in data.items()]) + set_clauses = ", ".join( + [ + ( + f"{key} = '" + str(value).replace("'", "''") + "'" + if isinstance(value, str) + else f"{key} = {value}" + ) + for key, value in data.items() + ] + ) sql = f"UPDATE STG.BATCH_LOGS SET {set_clauses} WHERE BATCH_ID = '{batch_id}';" return sql @@ -249,48 +337,48 @@ def lambda_handler(event, context): # Extract operation type and payload from event # res = cur.execute("SELECT CURRENT_DATABASE();").fetchone() # logger.info(f"current database: {res}") - operation = event['operation'] # 'insert' or 'update' - data = event['data'] - table = event['table'] - + operation = event["operation"] # 'insert' or 'update' + data = event["data"] + table = event["table"] + with conn.cursor() as cur: try: - if table == 'DOCUMENT_LOGS': - if operation == 'insert': + if table == "DOCUMENT_LOGS": + if operation == "insert": sql = construct_doc_insert_sql(data) - elif operation == 'update': - document_id = data.pop('DOCUMENT_ID', None) + elif operation == "update": + document_id = data.pop("DOCUMENT_ID", None) sql = construct_doc_update_sql(data, document_id) - - elif table == 'BATCH_LOGS': - if operation == 'insert': + + elif table == "BATCH_LOGS": + if operation == "insert": sql = construct_batch_insert_sql(data) - elif operation == 'update': - batch_id = data.pop('BATCH_ID', None) + elif operation == "update": + batch_id = data.pop("BATCH_ID", None) sql = construct_batch_update_sql(data, batch_id) - elif table == 'DOCZY_PIPELINE_RAW_OUTPUT': - if operation == 'insert': + elif table == "DOCZY_PIPELINE_RAW_OUTPUT": + if operation == "insert": sql = construct_doczy_pipeline_insert_sql(data) - #elif operation == 'update': - #batch_id = data.pop('BATCH_ID', None) - #sql = construct_doczy_pipeline_update_sql(data, batch_id) - - elif table == 'CLIENT_LOGS': - if operation == 'insert': + # elif operation == 'update': + # batch_id = data.pop('BATCH_ID', None) + # sql = construct_doczy_pipeline_update_sql(data, batch_id) + + elif table == "CLIENT_LOGS": + if operation == "insert": sql = construct_client_insert_sql(data) - elif operation == 'update': - client_id = data.pop('CLIENT_ID', None) + elif operation == "update": + client_id = data.pop("CLIENT_ID", None) sql = construct_client_update_sql(data, client_id) else: raise ValueError("Unsupported table.") - + response = cur.execute(sql) logger.info(f"Executing the following logging SQL statement: {sql}") - return {'statusCode': 200, 'body': json.dumps('Operation successful')} - + return {"statusCode": 200, "body": json.dumps("Operation successful")} + except Exception as e: - return {'statusCode': 400, 'body': json.dumps(str(e))} + return {"statusCode": 400, "body": json.dumps(str(e))} finally: # conn.close() # Need to close the connection to avoid reaching the limit of open connections diff --git a/textract-pipeline/src/lambda/docx-to-pdf/index.py b/textract-pipeline/src/lambda/docx-to-pdf/index.py index 254e0f7..3130871 100644 --- a/textract-pipeline/src/lambda/docx-to-pdf/index.py +++ b/textract-pipeline/src/lambda/docx-to-pdf/index.py @@ -17,24 +17,22 @@ logger = logging.getLogger(__name__) logger.setLevel(logging.INFO) # Initialize S3 & Textract clients -s3_client = boto3.client('s3') +s3_client = boto3.client("s3") DATABASE_LOGGING_LAMBDA_FUNCTION_NAME = "" -LIBRE_OFFICE_INSTALL_DIR = '/tmp/instdir' +LIBRE_OFFICE_INSTALL_DIR = "/tmp/instdir" + # Function to get s3 object tags def get_s3_object_tags(bucket_name, object_key): try: # Get object tags - response = s3_client.get_object_tagging( - Bucket=bucket_name, - Key=object_key - ) + response = s3_client.get_object_tagging(Bucket=bucket_name, Key=object_key) # Extract tags from the response and convert to dictionary - tags_list = response['TagSet'] - tags_dict = {tag['Key']: tag['Value'] for tag in tags_list} + tags_list = response["TagSet"] + tags_dict = {tag["Key"]: tag["Value"] for tag in tags_list} return tags_dict @@ -43,12 +41,13 @@ def get_s3_object_tags(bucket_name, object_key): print(f"Error: {e}") return None + # Function to retrieve configuration values from S3 def load_config_from_s3(bucket_name, file_key): - + # Download the config file from S3 response = s3_client.get_object(Bucket=bucket_name, Key=file_key) - config_content = response['Body'].read().decode('utf-8') + config_content = response["Body"].read().decode("utf-8") # Parse the config file config_parser = ConfigParser() @@ -57,15 +56,22 @@ def load_config_from_s3(bucket_name, file_key): # Convert the configuration to a dictionary config_dict = {} for section in config_parser.sections(): - config_dict[section] = {key.upper(): value for key, value in config_parser.items(section)} + config_dict[section] = { + key.upper(): value for key, value in config_parser.items(section) + } return config_dict + # Function to move file within S3 def move_file_within_s3(source_bucket, source_path, destination_path): try: # Copy the file to the destination folder - s3_client.copy_object(Bucket=source_bucket, CopySource={'Bucket': source_bucket, 'Key': source_path}, Key=destination_path) + s3_client.copy_object( + Bucket=source_bucket, + CopySource={"Bucket": source_bucket, "Key": source_path}, + Key=destination_path, + ) # Delete the file from the source folder s3_client.delete_object(Bucket=source_bucket, Key=source_path) @@ -74,13 +80,16 @@ def move_file_within_s3(source_bucket, source_path, destination_path): except Exception as e: logger.error(f"Error moving file: {e}") + def load_libre_office(): - if os.path.exists(LIBRE_OFFICE_INSTALL_DIR) and os.path.isdir(LIBRE_OFFICE_INSTALL_DIR): - print('We have a cached copy of LibreOffice, skipping extraction') + if os.path.exists(LIBRE_OFFICE_INSTALL_DIR) and os.path.isdir( + LIBRE_OFFICE_INSTALL_DIR + ): + print("We have a cached copy of LibreOffice, skipping extraction") else: - print('No cached copy of LibreOffice, extracting tar stream from Brotli file.') + print("No cached copy of LibreOffice, extracting tar stream from Brotli file.") buffer = BytesIO() - with open('/opt/lo.tar.br', 'rb') as brotli_file: + with open("/opt/lo.tar.br", "rb") as brotli_file: d = brotli.Decompressor() while True: chunk = brotli_file.read(1024) @@ -88,28 +97,35 @@ def load_libre_office(): if len(chunk) < 1024: break buffer.seek(0) - print('Extracting tar stream to /tmp for caching.') + print("Extracting tar stream to /tmp for caching.") with tarfile.open(fileobj=buffer) as tar: - tar.extractall('/tmp') - print('Done caching LibreOffice!') - return f'{LIBRE_OFFICE_INSTALL_DIR}/program/soffice.bin' - + tar.extractall("/tmp") + print("Done caching LibreOffice!") + return f"{LIBRE_OFFICE_INSTALL_DIR}/program/soffice.bin" + + def download_from_s3(bucket, key, download_path): s3 = boto3.client("s3") s3.download_file(bucket, key, download_path) - + + def upload_to_s3(file_path, bucket, key, tags): s3 = boto3.client("s3") response = s3.upload_file(file_path, bucket, key, ExtraArgs={"Tagging": tags}) print(response) - + + def convert_word_to_pdf(soffice_path, word_file_path, output_dir): print(word_file_path) conv_cmd = f"{soffice_path} --headless --norestore --invisible --nodefault --nofirststartwizard --nolockcheck --nologo --convert-to pdf:writer_pdf_Export --outdir {output_dir} {word_file_path}" print(conv_cmd) - response = subprocess.run(conv_cmd.split(), stdout=subprocess.PIPE, stderr=subprocess.PIPE) + response = subprocess.run( + conv_cmd.split(), stdout=subprocess.PIPE, stderr=subprocess.PIPE + ) if response.returncode != 0: - response = subprocess.run(conv_cmd.split(), stdout=subprocess.PIPE, stderr=subprocess.PIPE) + response = subprocess.run( + conv_cmd.split(), stdout=subprocess.PIPE, stderr=subprocess.PIPE + ) print(response.returncode, response.stdout, response.stderr) if len(response.stderr) > 0: @@ -120,20 +136,26 @@ def convert_word_to_pdf(soffice_path, word_file_path, output_dir): if response.returncode != 0: return False - + return True + def lambda_handler(event, context): try: # Read environment variables - property_file_path = os.environ.get('PROPERTY_FILE_S3_PATH', '') + property_file_path = os.environ.get("PROPERTY_FILE_S3_PATH", "") global DATABASE_LOGGING_LAMBDA_FUNCTION_NAME - DATABASE_LOGGING_LAMBDA_FUNCTION_NAME = os.environ.get('DATABASE_LOGGING_LAMBDA_FUNCTION_NAME', '') + DATABASE_LOGGING_LAMBDA_FUNCTION_NAME = os.environ.get( + "DATABASE_LOGGING_LAMBDA_FUNCTION_NAME", "" + ) batch_id = "" - - logger.info('DATABASE_LOGGING_LAMBDA_FUNCTION_NAME: ' + str(DATABASE_LOGGING_LAMBDA_FUNCTION_NAME)) + + logger.info( + "DATABASE_LOGGING_LAMBDA_FUNCTION_NAME: " + + str(DATABASE_LOGGING_LAMBDA_FUNCTION_NAME) + ) # Read config.properties file_path_array = property_file_path.split("/") @@ -145,13 +167,13 @@ def lambda_handler(event, context): S3_BUCKET_NAME = file_path_array[0] CONFIG_FILE_PATH = "/".join(file_path_array[1:]) - logger.info(f'S3_BUCKET_NAME: {S3_BUCKET_NAME}') - logger.info(f'CONFIG_FILE_PATH: {CONFIG_FILE_PATH}') + logger.info(f"S3_BUCKET_NAME: {S3_BUCKET_NAME}") + logger.info(f"CONFIG_FILE_PATH: {CONFIG_FILE_PATH}") # Load config file config_dict = load_config_from_s3(S3_BUCKET_NAME, CONFIG_FILE_PATH) - #logger.info('## CONFIG DICTIONARY\r' + str(config_dict)) + # logger.info('## CONFIG DICTIONARY\r' + str(config_dict)) print(event) @@ -160,44 +182,78 @@ def lambda_handler(event, context): files_failure = 0 # Process each message from the SQS event - for record in event['Records']: + for record in event["Records"]: # Extract the message body from the record - record_body = json.loads(record['body']) + record_body = json.loads(record["body"]) - for sqs_record in record_body['Records']: + for sqs_record in record_body["Records"]: # decode source path - source_path = unquote_plus(sqs_record['s3']['object']['key']) + source_path = unquote_plus(sqs_record["s3"]["object"]["key"]) logging.info("SOURCE_PATH: {source_path} ") - + # Get file tags tags_dict = get_s3_object_tags(S3_BUCKET_NAME, source_path) - + if "BatchId" in tags_dict: - batch_id = tags_dict['BatchId'] + batch_id = tags_dict["BatchId"] else: logger.error(f"BatchId not found in file tag") return - + # Extract configuration values - ALL_PDF_LOCATION = config_dict['FOLDER_LOCATIONS']['ALL_PDF_LOCATION'].format(batch_id) # SOURCE_LOCATION - SOURCE_DOCX_LOCATION = config_dict['FOLDER_LOCATIONS']['SOURCE_DOCX_LOCATION'].format(batch_id) # SOURCE_DOCX_LOCATION - SOURCE_DOCX_PROCESSED_LOCATION = config_dict['FOLDER_LOCATIONS']['SOURCE_DOCX_PROCESSED_LOCATION'].format(batch_id) # SOURCE_DOCX_LOCATION - SOURCE_DOCX_UNPROCESSED_LOCATION = config_dict['FOLDER_LOCATIONS']['SOURCE_DOCX_UNPROCESSED_LOCATION'].format(batch_id) # SOURCE_DOCX_LOCATION - logger.info('ALL_PDF_LOCATION: ' + ALL_PDF_LOCATION) - logger.info('SOURCE_DOCX_LOCATION: ' + SOURCE_DOCX_LOCATION) - logger.info('SOURCE_DOCX_PROCESSED_LOCATION: ' + SOURCE_DOCX_PROCESSED_LOCATION) - logger.info('SOURCE_DOCX_UNPROCESSED_LOCATION: ' + SOURCE_DOCX_UNPROCESSED_LOCATION) + ALL_PDF_LOCATION = config_dict["FOLDER_LOCATIONS"][ + "ALL_PDF_LOCATION" + ].format( + batch_id + ) # SOURCE_LOCATION + SOURCE_DOCX_LOCATION = config_dict["FOLDER_LOCATIONS"][ + "SOURCE_DOCX_LOCATION" + ].format( + batch_id + ) # SOURCE_DOCX_LOCATION + SOURCE_DOCX_PROCESSED_LOCATION = config_dict["FOLDER_LOCATIONS"][ + "SOURCE_DOCX_PROCESSED_LOCATION" + ].format( + batch_id + ) # SOURCE_DOCX_LOCATION + SOURCE_DOCX_UNPROCESSED_LOCATION = config_dict["FOLDER_LOCATIONS"][ + "SOURCE_DOCX_UNPROCESSED_LOCATION" + ].format( + batch_id + ) # SOURCE_DOCX_LOCATION + logger.info("ALL_PDF_LOCATION: " + ALL_PDF_LOCATION) + logger.info("SOURCE_DOCX_LOCATION: " + SOURCE_DOCX_LOCATION) + logger.info( + "SOURCE_DOCX_PROCESSED_LOCATION: " + + SOURCE_DOCX_PROCESSED_LOCATION + ) + logger.info( + "SOURCE_DOCX_UNPROCESSED_LOCATION: " + + SOURCE_DOCX_UNPROCESSED_LOCATION + ) # Verify source path have valid file extension - if os.path.splitext(source_path)[1].lower() not in ['.doc','.docx']: - print('File type not supported: ', os.path.splitext(source_path)[1]) - files_failure+=1 + if os.path.splitext(source_path)[1].lower() not in [ + ".doc", + ".docx", + ]: + print( + "File type not supported: ", + os.path.splitext(source_path)[1], + ) + files_failure += 1 continue - processed_destination_path = SOURCE_DOCX_PROCESSED_LOCATION + source_path.replace(SOURCE_DOCX_LOCATION,"") - unprocessed_destination_path = SOURCE_DOCX_UNPROCESSED_LOCATION + source_path.replace(SOURCE_DOCX_LOCATION,"") + processed_destination_path = ( + SOURCE_DOCX_PROCESSED_LOCATION + + source_path.replace(SOURCE_DOCX_LOCATION, "") + ) + unprocessed_destination_path = ( + SOURCE_DOCX_UNPROCESSED_LOCATION + + source_path.replace(SOURCE_DOCX_LOCATION, "") + ) temporary_filename = "conversion_file" key_prefix, base_name = os.path.split(source_path) @@ -206,122 +262,129 @@ def lambda_handler(event, context): output_dir = "/tmp" # Logic to consider subfolder at source location - temp_path = source_path.replace(SOURCE_DOCX_LOCATION,"") + temp_path = source_path.replace(SOURCE_DOCX_LOCATION, "") key_prefix, base_name = os.path.split(temp_path) - + print("key_prefix len ", len(key_prefix)) - - if len(key_prefix) > 0 and key_prefix != None or key_prefix != "": + + if len(key_prefix) > 0 and key_prefix != None or key_prefix != "": destination_path = ALL_PDF_LOCATION + key_prefix + "/" else: destination_path = ALL_PDF_LOCATION - + print("destination_path : ", destination_path) - + # Uncomment below line while testing this lambda # destination_path = "pdf_converted_files/" - + files[source_path] = False # Load Libreoffice library libreoffice_exec_path = "" - if os.path.isfile('/opt/lo.tar.br'): - logging.info('compressed Libreoffice found!') + if os.path.isfile("/opt/lo.tar.br"): + logging.info("compressed Libreoffice found!") libreoffice_exec_path = load_libre_office() else: - print('libreoffice Layer not found!') - return {'body': 'libreoffice Layer not found!'} + print("libreoffice Layer not found!") + return {"body": "libreoffice Layer not found!"} logging.info("DOWNLOAD_PATH: {download_path}") download_from_s3(S3_BUCKET_NAME, source_path, download_path) - logging.info('Downloading Finished!') - - print("Files list after download: ",os.listdir('/tmp')) + logging.info("Downloading Finished!") - logger.info('Starting Conversion') + print("Files list after download: ", os.listdir("/tmp")) - is_converted = convert_word_to_pdf(libreoffice_exec_path, download_path, output_dir) + logger.info("Starting Conversion") + + is_converted = convert_word_to_pdf( + libreoffice_exec_path, download_path, output_dir + ) output_filepath = f"{output_dir}/{temporary_filename}.pdf" - - if is_converted and os.path.isfile(output_filepath): - logger.info('Conversion Success!') - file_name, _ = os.path.splitext(base_name) - files_success+=1 - files[source_path] = True - logger.info(f'Saving converted file in Bucket {S3_BUCKET_NAME} with Key : {destination_path}{file_name}.pdf') - + if is_converted and os.path.isfile(output_filepath): + logger.info("Conversion Success!") + file_name, _ = os.path.splitext(base_name) + files_success += 1 + files[source_path] = True + + logger.info( + f"Saving converted file in Bucket {S3_BUCKET_NAME} with Key : {destination_path}{file_name}.pdf" + ) + tags = urlencode(tags_dict) # Upload file to S3 location - upload_to_s3(output_filepath, S3_BUCKET_NAME, f"{destination_path}{file_name}.pdf", tags) + upload_to_s3( + output_filepath, + S3_BUCKET_NAME, + f"{destination_path}{file_name}.pdf", + tags, + ) # Move file to processed folder - move_file_within_s3(S3_BUCKET_NAME, source_path, processed_destination_path) - + move_file_within_s3( + S3_BUCKET_NAME, source_path, processed_destination_path + ) + logger.info("File converted: " + source_path) - generate_document_logs_input(file_name,"PDF_CONVERSION_SUCCESS") + generate_document_logs_input( + file_name, "PDF_CONVERSION_SUCCESS" + ) else: # Move file to unprocessed folder - move_file_within_s3(S3_BUCKET_NAME, source_path, unprocessed_destination_path) - generate_document_logs_input(file_name,"PDF_CONVERSION_FAILED") - files_failure+=1 + move_file_within_s3( + S3_BUCKET_NAME, source_path, unprocessed_destination_path + ) + generate_document_logs_input(file_name, "PDF_CONVERSION_FAILED") + files_failure += 1 logger.error("File not converted: " + source_path) - + files_converted = { - 'statusCode': 200, - 'body': { - 'Files_Processed': len(files), - 'Success' : files_success, - 'Failure': files_failure - } - } + "statusCode": 200, + "body": { + "Files_Processed": len(files), + "Success": files_success, + "Failure": files_failure, + }, + } logger.info(files_converted) - + return files_converted except ClientError as e: # Handle specific Textract client errors error_message = f"Error in pdf conversion operation: {e}" logger.error(error_message) - return { - 'statusCode': 500, - 'body': error_message - } + return {"statusCode": 500, "body": error_message} except Exception as e: # Handle other exceptions error_message = f"Unexpected error: {e}" logger.error(error_message) - return { - 'statusCode': 500, - 'body': error_message - } - + return {"statusCode": 500, "body": error_message} -def generate_document_logs_input(document_id,stage): - current_time = datetime.datetime.now().isoformat() - - data = { + +def generate_document_logs_input(document_id, stage): + current_time = datetime.datetime.now().isoformat() + + data = { "operation": "update", "table": "DOCUMENT_LOGS", "data": { "DOCUMENT_ID": document_id, - "STAGE" : stage, + "STAGE": stage, "MODIFIED_TIME": current_time, - "MODIFIED_BY": "DOCX_TO_PDF" - } - } - logger.info(f"Request: {data}") - - lambda_client = boto3.client('lambda') - response = lambda_client.invoke( - FunctionName=DATABASE_LOGGING_LAMBDA_FUNCTION_NAME, - InvocationType='Event', # Asynchronous invocation - Payload=json.dumps(data).encode('utf-8') - ) - - logger.info("Generated document logs input successfully") - return response + "MODIFIED_BY": "DOCX_TO_PDF", + }, + } + logger.info(f"Request: {data}") - \ No newline at end of file + lambda_client = boto3.client("lambda") + response = lambda_client.invoke( + FunctionName=DATABASE_LOGGING_LAMBDA_FUNCTION_NAME, + InvocationType="Event", # Asynchronous invocation + Payload=json.dumps(data).encode("utf-8"), + ) + + logger.info("Generated document logs input successfully") + return response diff --git a/textract-pipeline/src/lambda/insert-record/index.py b/textract-pipeline/src/lambda/insert-record/index.py index 73c2263..e44ac23 100644 --- a/textract-pipeline/src/lambda/insert-record/index.py +++ b/textract-pipeline/src/lambda/insert-record/index.py @@ -11,15 +11,16 @@ logger = logging.getLogger() logger.setLevel(logging.INFO) # Initialize S3 clients -s3_client = boto3.client('s3') +s3_client = boto3.client("s3") DATABASE_LOGGING_LAMBDA_FUNCTION_NAME = "" + # Function to retrieve configuration values from S3 def load_config_from_s3(bucket_name, file_key): - + # Download the config file from S3 response = s3_client.get_object(Bucket=bucket_name, Key=file_key) - config_content = response['Body'].read().decode('utf-8') + config_content = response["Body"].read().decode("utf-8") # Parse the config file config_parser = ConfigParser() @@ -28,156 +29,204 @@ def load_config_from_s3(bucket_name, file_key): # Convert the configuration to a dictionary config_dict = {} for section in config_parser.sections(): - config_dict[section] = {key.upper(): value for key, value in config_parser.items(section)} + config_dict[section] = { + key.upper(): value for key, value in config_parser.items(section) + } return config_dict + def lambda_handler(event, context): # Extracting message body from the event - message_body = json.loads(event['Records'][0]['body']) - + message_body = json.loads(event["Records"][0]["body"]) + # Extracting relevant information from the message - s3_bucket = message_body.get('s3_bucket') - batch_id = message_body.get('batch_id') - username = message_body.get('username') - client_name = message_body.get('client_name') - contract_list = message_body.get('contract_list', []) - + s3_bucket = message_body.get("s3_bucket") + batch_id = message_body.get("batch_id") + username = message_body.get("username") + client_name = message_body.get("client_name") + contract_list = message_body.get("contract_list", []) + # Load config file # Read environment variables - property_file_path = os.environ.get('PROPERTY_FILE_S3_PATH', '') + property_file_path = os.environ.get("PROPERTY_FILE_S3_PATH", "") global DATABASE_LOGGING_LAMBDA_FUNCTION_NAME - DATABASE_LOGGING_LAMBDA_FUNCTION_NAME = os.environ.get('DATABASE_LOGGING_LAMBDA_FUNCTION_NAME', '') - - logger.info('DATABASE_LOGGING_LAMBDA_FUNCTION_NAME: ' + str(DATABASE_LOGGING_LAMBDA_FUNCTION_NAME)) - + DATABASE_LOGGING_LAMBDA_FUNCTION_NAME = os.environ.get( + "DATABASE_LOGGING_LAMBDA_FUNCTION_NAME", "" + ) + + logger.info( + "DATABASE_LOGGING_LAMBDA_FUNCTION_NAME: " + + str(DATABASE_LOGGING_LAMBDA_FUNCTION_NAME) + ) + # Read config.properties file_path_array = property_file_path.split("/") - + # Valid if file_path_array has more than 2 elements if len(file_path_array) > 1: - + # Extract BUCKET_NAME and config_file_path S3_BUCKET_NAME = file_path_array[0] CONFIG_FILE_PATH = "/".join(file_path_array[1:]) # Load config file config_dict = load_config_from_s3(S3_BUCKET_NAME, CONFIG_FILE_PATH) - ALL_PDF_LOCATION = config_dict['FOLDER_LOCATIONS']['ALL_PDF_LOCATION'].format(batch_id) # ALL_PDF_LOCATION - SOURCE_DOCX_LOCATION = config_dict['FOLDER_LOCATIONS']['SOURCE_DOCX_LOCATION'].format(batch_id) # SOURCE_DOCX_LOCATION - CONTRACTS_LANDNING_ZONE = config_dict['FOLDER_LOCATIONS']['CONTRACTS_LANDNING_ZONE'].format(batch_id) - tags_dict = {'BatchId': batch_id, "ClientName": client_name} + ALL_PDF_LOCATION = config_dict["FOLDER_LOCATIONS"]["ALL_PDF_LOCATION"].format( + batch_id + ) # ALL_PDF_LOCATION + SOURCE_DOCX_LOCATION = config_dict["FOLDER_LOCATIONS"][ + "SOURCE_DOCX_LOCATION" + ].format( + batch_id + ) # SOURCE_DOCX_LOCATION + CONTRACTS_LANDNING_ZONE = config_dict["FOLDER_LOCATIONS"][ + "CONTRACTS_LANDNING_ZONE" + ].format(batch_id) + tags_dict = {"BatchId": batch_id, "ClientName": client_name} additional_tags = urlencode(tags_dict) logger.info("Tagging: " + additional_tags) # Processing each contract in the contract list for contract in contract_list: - contract_name = contract.get('contract_name') - groups = contract.get('groups', []) - groups_str = ', '.join(groups) - contract_source_path = contract.get('contract_source_path') - + contract_name = contract.get("contract_name") + groups = contract.get("groups", []) + groups_str = ", ".join(groups) + contract_source_path = contract.get("contract_source_path") + # Perform your processing logic here - + # Example processing: logging information logger.info(f"Processing contract: {contract_name}") logger.info(f"Groups: {groups_str}") logger.info(f"Contract source path: {contract_source_path}") - file_name_without_extension, file_extension = os.path.splitext(contract_name) - file_move_status = file_distribution(s3_bucket,contract_source_path,ALL_PDF_LOCATION,SOURCE_DOCX_LOCATION,CONTRACTS_LANDNING_ZONE,file_extension,additional_tags) + file_name_without_extension, file_extension = os.path.splitext( + contract_name + ) + file_move_status = file_distribution( + s3_bucket, + contract_source_path, + ALL_PDF_LOCATION, + SOURCE_DOCX_LOCATION, + CONTRACTS_LANDNING_ZONE, + file_extension, + additional_tags, + ) stage = "NEW" if not file_move_status: stage = "FILE_NOT_FOUND" - - response = generate_document_logs_input(file_name_without_extension,batch_id, stage, s3_bucket, - contract_name, contract_source_path, groups_str, username, file_extension) + + response = generate_document_logs_input( + file_name_without_extension, + batch_id, + stage, + s3_bucket, + contract_name, + contract_source_path, + groups_str, + username, + file_extension, + ) logger.info(f"Response: {response}") - - - # Assuming processing is successful, returning a success message - return { - 'statusCode': 200, - 'body': json.dumps('Processing completed successfully') - } + return {"statusCode": 200, "body": json.dumps("Processing completed successfully")} -def file_distribution(s3_bucket,contract_source_path,ALL_PDF_LOCATION,SOURCE_DOCX_LOCATION,CONTRACTS_LANDNING_ZONE,file_extension,additional_tags): - if file_extension.lower() in [".pdf",".filepart"]: + +def file_distribution( + s3_bucket, + contract_source_path, + ALL_PDF_LOCATION, + SOURCE_DOCX_LOCATION, + CONTRACTS_LANDNING_ZONE, + file_extension, + additional_tags, +): + if file_extension.lower() in [".pdf", ".filepart"]: # Copy to ALL_PDF_LOCATION source_key = contract_source_path - destination_key = ALL_PDF_LOCATION + contract_source_path.replace(CONTRACTS_LANDNING_ZONE,"") - - return move_object_within_bucket(s3_bucket, source_key, destination_key, additional_tags) - + destination_key = ALL_PDF_LOCATION + contract_source_path.replace( + CONTRACTS_LANDNING_ZONE, "" + ) + + return move_object_within_bucket( + s3_bucket, source_key, destination_key, additional_tags + ) # If doc/docx copy to doc folder - elif file_extension.lower() in [".docx",".doc"]: + elif file_extension.lower() in [".docx", ".doc"]: # Copy to SOURCE_DOCX_LOCATION source_key = contract_source_path - destination_key = SOURCE_DOCX_LOCATION + contract_source_path.replace(CONTRACTS_LANDNING_ZONE,"") - - return move_object_within_bucket(s3_bucket, source_key, destination_key, additional_tags) - - + destination_key = SOURCE_DOCX_LOCATION + contract_source_path.replace( + CONTRACTS_LANDNING_ZONE, "" + ) -def generate_document_logs_input(document_id,batch_id, stage, bucket_name, - file_name, file_path, group_id, - created_by, original_file_extension): + return move_object_within_bucket( + s3_bucket, source_key, destination_key, additional_tags + ) + + +def generate_document_logs_input( + document_id, + batch_id, + stage, + bucket_name, + file_name, + file_path, + group_id, + created_by, + original_file_extension, +): current_time = datetime.datetime.now().isoformat() data = { "operation": "insert", "table": "DOCUMENT_LOGS", "data": { - "DOCUMENT_ID" : document_id, + "DOCUMENT_ID": document_id, "BATCH_ID": batch_id, "STAGE": stage, "BUCKET_NAME": bucket_name, "FILE_NAME": file_name, "FILE_PATH": file_path, "GROUP_ID": group_id, - "CREATED_TIME" : current_time, + "CREATED_TIME": current_time, "CREATED_BY": created_by, - "ORIGINAL_FILE_EXTENSION": original_file_extension - } + "ORIGINAL_FILE_EXTENSION": original_file_extension, + }, } logger.info(f"Request: {data}") - lambda_client = boto3.client('lambda') + lambda_client = boto3.client("lambda") response = lambda_client.invoke( FunctionName=DATABASE_LOGGING_LAMBDA_FUNCTION_NAME, - InvocationType='Event', # Asynchronous invocation - Payload=json.dumps(data).encode('utf-8') + InvocationType="Event", # Asynchronous invocation + Payload=json.dumps(data).encode("utf-8"), ) logger.info("Generated document logs input successfully") return response - def move_object_within_bucket(bucket_name, source_key, destination_key, tags=None): """ Move an object within the same S3 bucket and add additional tags if provided. """ - copy_source = { - 'Bucket': bucket_name, - 'Key': source_key - } + copy_source = {"Bucket": bucket_name, "Key": source_key} copy_object_args = { - 'Bucket': bucket_name, - 'Key': destination_key, - 'CopySource': copy_source, - 'TaggingDirective': 'REPLACE' + "Bucket": bucket_name, + "Key": destination_key, + "CopySource": copy_source, + "TaggingDirective": "REPLACE", } logger.info(f"Destination Folder Details: {copy_object_args}") if tags: - copy_object_args['Tagging'] = tags + copy_object_args["Tagging"] = tags try: # Copy the object to the new location @@ -186,10 +235,10 @@ def move_object_within_bucket(bucket_name, source_key, destination_key, tags=Non # Delete the original object s3_client.delete_object(Bucket=bucket_name, Key=source_key) - msg = f"Object moved successfully. New object Key: {destination_key}" + msg = f"Object moved successfully. New object Key: {destination_key}" logger.info(msg) return True except Exception as e: - msg = f"Error moving object: {e}" + msg = f"Error moving object: {e}" logger.error(msg) return False diff --git a/textract-pipeline/src/lambda/pdf-validation/index.py b/textract-pipeline/src/lambda/pdf-validation/index.py index 4604760..9e4f7a7 100644 --- a/textract-pipeline/src/lambda/pdf-validation/index.py +++ b/textract-pipeline/src/lambda/pdf-validation/index.py @@ -15,7 +15,7 @@ logger = logging.getLogger() logger.setLevel(logging.INFO) # Initialize S3 & Textract clients -s3_client = boto3.client('s3') +s3_client = boto3.client("s3") DATABASE_LOGGING_LAMBDA_FUNCTION_NAME = "" @@ -23,14 +23,11 @@ DATABASE_LOGGING_LAMBDA_FUNCTION_NAME = "" def get_s3_object_tags(bucket_name, object_key): try: # Get object tags - response = s3_client.get_object_tagging( - Bucket=bucket_name, - Key=object_key - ) + response = s3_client.get_object_tagging(Bucket=bucket_name, Key=object_key) # Extract tags from the response and convert to dictionary - tags_list = response['TagSet'] - tags_dict = {tag['Key']: tag['Value'] for tag in tags_list} + tags_list = response["TagSet"] + tags_dict = {tag["Key"]: tag["Value"] for tag in tags_list} print(f"get_s3_object_tags: response={response}") return tags_dict @@ -44,7 +41,7 @@ def get_s3_object_tags(bucket_name, object_key): def load_config_from_s3(bucket_name, file_key): # Download the config file from S3 response = s3_client.get_object(Bucket=bucket_name, Key=file_key) - config_content = response['Body'].read().decode('utf-8') + config_content = response["Body"].read().decode("utf-8") # Parse the config file config_parser = ConfigParser() @@ -53,16 +50,20 @@ def load_config_from_s3(bucket_name, file_key): # Convert the configuration to a dictionary config_dict = {} for section in config_parser.sections(): - config_dict[section] = {key.upper(): value for key, value in config_parser.items(section)} + config_dict[section] = { + key.upper(): value for key, value in config_parser.items(section) + } return config_dict def lambda_handler(event, context): # Read environment variables - property_file_path = os.environ.get('PROPERTY_FILE_S3_PATH', '') + property_file_path = os.environ.get("PROPERTY_FILE_S3_PATH", "") global DATABASE_LOGGING_LAMBDA_FUNCTION_NAME - DATABASE_LOGGING_LAMBDA_FUNCTION_NAME = os.environ.get('DATABASE_LOGGING_LAMBDA_FUNCTION_NAME', '') + DATABASE_LOGGING_LAMBDA_FUNCTION_NAME = os.environ.get( + "DATABASE_LOGGING_LAMBDA_FUNCTION_NAME", "" + ) batch_id = "" @@ -76,8 +77,8 @@ def lambda_handler(event, context): S3_BUCKET_NAME = file_path_array[0] CONFIG_FILE_PATH = "/".join(file_path_array[1:]) - logger.info(f'S3_BUCKET_NAME: {S3_BUCKET_NAME}') - logger.info(f'CONFIG_FILE_PATH: {CONFIG_FILE_PATH}') + logger.info(f"S3_BUCKET_NAME: {S3_BUCKET_NAME}") + logger.info(f"CONFIG_FILE_PATH: {CONFIG_FILE_PATH}") # Load config file config_dict = load_config_from_s3(S3_BUCKET_NAME, CONFIG_FILE_PATH) @@ -87,15 +88,15 @@ def lambda_handler(event, context): print(event) # Process each message from the SQS event - for record in event['Records']: + for record in event["Records"]: # Extract the message body from the record - record_body = json.loads(record['body']) + record_body = json.loads(record["body"]) - for sqs_record in record_body['Records']: + for sqs_record in record_body["Records"]: # Retrieve the S3 bucket and key from the event - key = unquote_plus(sqs_record['s3']['object']['key']) + key = unquote_plus(sqs_record["s3"]["object"]["key"]) logging.info("SOURCE_PATH: {key} ") logger.info(f"Processing file: s3://{S3_BUCKET_NAME}/{key}") @@ -104,7 +105,7 @@ def lambda_handler(event, context): tags_dict = get_s3_object_tags(S3_BUCKET_NAME, key) if "BatchId" in tags_dict: - batch_id = tags_dict['BatchId'] + batch_id = tags_dict["BatchId"] else: logger.error("BatchId not found in file tag") return @@ -112,10 +113,15 @@ def lambda_handler(event, context): logger.info(f"Batch ID: {batch_id}") # Extract configuration values - INVALID_PDF_FILE_LOCATION = config_dict['FOLDER_LOCATIONS']['INVALID_PDF_FILE_LOCATION'].format( - batch_id) - SOURCE_LOCATION = config_dict['FOLDER_LOCATIONS']['SOURCE_LOCATION'].format(batch_id) - ALL_PDF_LOCATION = config_dict['FOLDER_LOCATIONS']['ALL_PDF_LOCATION'].format(batch_id) + INVALID_PDF_FILE_LOCATION = config_dict["FOLDER_LOCATIONS"][ + "INVALID_PDF_FILE_LOCATION" + ].format(batch_id) + SOURCE_LOCATION = config_dict["FOLDER_LOCATIONS"][ + "SOURCE_LOCATION" + ].format(batch_id) + ALL_PDF_LOCATION = config_dict["FOLDER_LOCATIONS"][ + "ALL_PDF_LOCATION" + ].format(batch_id) logger.info(f"INVALID_PDF_FILE_LOCATION : {INVALID_PDF_FILE_LOCATION}") logger.info(f"SOURCE_LOCATION : {SOURCE_LOCATION}") @@ -132,20 +138,25 @@ def lambda_handler(event, context): document_id = get_filename_from_path(key) document_id = os.path.splitext(document_id)[0] - logger.info(f"Processing S3 file - Bucket: {S3_BUCKET_NAME}, Key: {key}") + logger.info( + f"Processing S3 file - Bucket: {S3_BUCKET_NAME}, Key: {key}" + ) # Check if the file has a '.filepart' extension - if key.lower().endswith('.filepart'): - new_key = key[:-9] + '.pdf' # Rename to have a '.pdf' extension - s3_client.copy_object(Bucket=S3_BUCKET_NAME, CopySource={'Bucket': S3_BUCKET_NAME, 'Key': key}, - Key=new_key) + if key.lower().endswith(".filepart"): + new_key = key[:-9] + ".pdf" # Rename to have a '.pdf' extension + s3_client.copy_object( + Bucket=S3_BUCKET_NAME, + CopySource={"Bucket": S3_BUCKET_NAME, "Key": key}, + Key=new_key, + ) s3_client.delete_object(Bucket=S3_BUCKET_NAME, Key=key) key = new_key # Update key to the new filename # Validate number of pages try: s3_file = s3_client.get_object(Bucket=S3_BUCKET_NAME, Key=key) - pdf_reader = PdfReader(BytesIO(s3_file['Body'].read())) + pdf_reader = PdfReader(BytesIO(s3_file["Body"].read())) if pdf_reader.is_encrypted: logger.info("File is encrypted. Trying to decrypt.") @@ -154,66 +165,148 @@ def lambda_handler(event, context): logger.info("File decrypted successfully.") except Exception as e: logger.error(f"Failed to decrypt file: {e}") - move_to_unprocessed(S3_BUCKET_NAME, key, INVALID_PDF_FILE_LOCATION, - sub_folder_with_filename) - generate_document_logs_input(document_id, num_pages, file_size, - "CONTRACT_VALIDATION - ENCRYPTED FILE") + move_to_unprocessed( + S3_BUCKET_NAME, + key, + INVALID_PDF_FILE_LOCATION, + sub_folder_with_filename, + ) + generate_document_logs_input( + document_id, + num_pages, + file_size, + "CONTRACT_VALIDATION - ENCRYPTED FILE", + ) return # Validate file size - file_size = s3_client.head_object(Bucket=S3_BUCKET_NAME, Key=key)['ContentLength'] + file_size = s3_client.head_object(Bucket=S3_BUCKET_NAME, Key=key)[ + "ContentLength" + ] logger.info(f"File size: {file_size} bytes") if file_size > 500 * 1024 * 1024: # 500MB - logger.info("File size exceeds 500MB. Moving to 'unprocessed' folder.") - move_to_unprocessed(S3_BUCKET_NAME, key, INVALID_PDF_FILE_LOCATION, sub_folder_with_filename) - generate_document_logs_input(document_id, num_pages, file_size, - "CONTRACT_VALIDATION - INVALID FILE SIZE") + logger.info( + "File size exceeds 500MB. Moving to 'unprocessed' folder." + ) + move_to_unprocessed( + S3_BUCKET_NAME, + key, + INVALID_PDF_FILE_LOCATION, + sub_folder_with_filename, + ) + generate_document_logs_input( + document_id, + num_pages, + file_size, + "CONTRACT_VALIDATION - INVALID FILE SIZE", + ) return num_pages = len(pdf_reader.pages) logger.info(f"Number of pages: {num_pages}") if num_pages > 3000: - logger.info("Number of pages exceeds 3000. Moving to 'unprocessed' folder.") - move_to_unprocessed(S3_BUCKET_NAME, key, INVALID_PDF_FILE_LOCATION, sub_folder_with_filename) - generate_document_logs_input(document_id, num_pages, file_size, - "CONTRACT_VALIDATION - INVALID NUMBER OF PAGES") + logger.info( + "Number of pages exceeds 3000. Moving to 'unprocessed' folder." + ) + move_to_unprocessed( + S3_BUCKET_NAME, + key, + INVALID_PDF_FILE_LOCATION, + sub_folder_with_filename, + ) + generate_document_logs_input( + document_id, + num_pages, + file_size, + "CONTRACT_VALIDATION - INVALID NUMBER OF PAGES", + ) return except botocore.exceptions.ClientError as e: - if e.response['Error']['Code'] == '404': + if e.response["Error"]["Code"] == "404": logger.error(f"File not found: s3://{S3_BUCKET_NAME}/{key}") # Handle the case where the file doesn't exist return else: - logger.error(f"Error checking file size: {str(e)}. Moving to 'unprocessed' folder.") - move_to_unprocessed(S3_BUCKET_NAME, key, INVALID_PDF_FILE_LOCATION, sub_folder_with_filename) - generate_document_logs_input(document_id, num_pages, file_size, - "CONTRACT_VALIDATION - ERROR FILE SIZE") + logger.error( + f"Error checking file size: {str(e)}. Moving to 'unprocessed' folder." + ) + move_to_unprocessed( + S3_BUCKET_NAME, + key, + INVALID_PDF_FILE_LOCATION, + sub_folder_with_filename, + ) + generate_document_logs_input( + document_id, + num_pages, + file_size, + "CONTRACT_VALIDATION - ERROR FILE SIZE", + ) return except PdfReadError as e: - logger.exception(f"PdfReadError: {str(e)}. Moving to 'unprocessed' folder.") - move_to_unprocessed(S3_BUCKET_NAME, key, INVALID_PDF_FILE_LOCATION, sub_folder_with_filename) - generate_document_logs_input(document_id, num_pages, file_size, - "CONTRACT_VALIDATION - ERROR FILE SIZE") + logger.exception( + f"PdfReadError: {str(e)}. Moving to 'unprocessed' folder." + ) + move_to_unprocessed( + S3_BUCKET_NAME, + key, + INVALID_PDF_FILE_LOCATION, + sub_folder_with_filename, + ) + generate_document_logs_input( + document_id, + num_pages, + file_size, + "CONTRACT_VALIDATION - ERROR FILE SIZE", + ) return except Exception as e: - logger.error(f"Error checking file size: {str(e)}. Moving to 'unprocessed' folder.") - move_to_unprocessed(S3_BUCKET_NAME, key, INVALID_PDF_FILE_LOCATION, sub_folder_with_filename) - generate_document_logs_input(document_id, num_pages, file_size, - "CONTRACT_VALIDATION - ERROR FILE SIZE") + logger.error( + f"Error checking file size: {str(e)}. Moving to 'unprocessed' folder." + ) + move_to_unprocessed( + S3_BUCKET_NAME, + key, + INVALID_PDF_FILE_LOCATION, + sub_folder_with_filename, + ) + generate_document_logs_input( + document_id, + num_pages, + file_size, + "CONTRACT_VALIDATION - ERROR FILE SIZE", + ) return # Validate password protection and resolution if not is_resolution_valid(pdf_reader): - logger.info("File has invalid resolution. Moving to 'unprocessed' folder.") - move_to_unprocessed(S3_BUCKET_NAME, key, INVALID_PDF_FILE_LOCATION, sub_folder_with_filename) - generate_document_logs_input(document_id, num_pages, file_size, - "CONTRACT_VALIDATION - INVALID RESOLUTION") + logger.info( + "File has invalid resolution. Moving to 'unprocessed' folder." + ) + move_to_unprocessed( + S3_BUCKET_NAME, + key, + INVALID_PDF_FILE_LOCATION, + sub_folder_with_filename, + ) + generate_document_logs_input( + document_id, + num_pages, + file_size, + "CONTRACT_VALIDATION - INVALID RESOLUTION", + ) return else: - logger.info("File passed all conditions. Moving to 'processed' folder.") - move_to_processed(S3_BUCKET_NAME, key, SOURCE_LOCATION, sub_folder_with_filename) + logger.info( + "File passed all conditions. Moving to 'processed' folder." + ) + move_to_processed( + S3_BUCKET_NAME, key, SOURCE_LOCATION, sub_folder_with_filename + ) - generate_document_logs_input(document_id, num_pages, file_size, "CONTRACT_VALIDATION") + generate_document_logs_input( + document_id, num_pages, file_size, "CONTRACT_VALIDATION" + ) def is_password_protected(pdf_reader): @@ -231,17 +324,27 @@ def is_resolution_valid(pdf_reader, max_resolution=3000 * 4000): def move_to_processed(bucket, key, valid_file_location, sub_folder_with_filename): # Replace processed with valid_file_location like in unprocessed method - s3_client.copy_object(Bucket=bucket, CopySource={'Bucket': bucket, 'Key': key}, - Key=f'{valid_file_location}{sub_folder_with_filename}') + s3_client.copy_object( + Bucket=bucket, + CopySource={"Bucket": bucket, "Key": key}, + Key=f"{valid_file_location}{sub_folder_with_filename}", + ) s3_client.delete_object(Bucket=bucket, Key=key) - logger.info(f"File moved to 'processed' folder: s3://{bucket}/{valid_file_location}{sub_folder_with_filename}") + logger.info( + f"File moved to 'processed' folder: s3://{bucket}/{valid_file_location}{sub_folder_with_filename}" + ) def move_to_unprocessed(bucket, key, invalid_file_location, sub_folder_with_filename): - s3_client.copy_object(Bucket=bucket, CopySource={'Bucket': bucket, 'Key': key}, - Key=f'{invalid_file_location}{sub_folder_with_filename}') + s3_client.copy_object( + Bucket=bucket, + CopySource={"Bucket": bucket, "Key": key}, + Key=f"{invalid_file_location}{sub_folder_with_filename}", + ) s3_client.delete_object(Bucket=bucket, Key=key) - logger.info(f"File NOT ;] moved to 'unprocessed' folder: s3://{bucket}/{invalid_file_location}{sub_folder_with_filename}") + logger.info( + f"File NOT ;] moved to 'unprocessed' folder: s3://{bucket}/{invalid_file_location}{sub_folder_with_filename}" + ) def get_filename_from_path(full_path): @@ -251,7 +354,10 @@ def get_filename_from_path(full_path): def generate_document_logs_input(document_id, num_pages, file_size, stage): current_time = datetime.datetime.now().isoformat() - logger.info('DATABASE_LOGGING_LAMBDA_FUNCTION_NAME: ' + str(DATABASE_LOGGING_LAMBDA_FUNCTION_NAME)) + logger.info( + "DATABASE_LOGGING_LAMBDA_FUNCTION_NAME: " + + str(DATABASE_LOGGING_LAMBDA_FUNCTION_NAME) + ) data = { "operation": "update", @@ -262,16 +368,16 @@ def generate_document_logs_input(document_id, num_pages, file_size, stage): "NO_OF_PAGES": num_pages, "FILE_SIZE": file_size, "MODIFIED_TIME": current_time, - "MODIFIED_BY": "CONTRACT_VALIDATION" - } + "MODIFIED_BY": "CONTRACT_VALIDATION", + }, } logger.info(f"Request: {data}") - lambda_client = boto3.client('lambda') + lambda_client = boto3.client("lambda") response = lambda_client.invoke( FunctionName=DATABASE_LOGGING_LAMBDA_FUNCTION_NAME, - InvocationType='Event', # Asynchronous invocation - Payload=json.dumps(data).encode('utf-8') + InvocationType="Event", # Asynchronous invocation + Payload=json.dumps(data).encode("utf-8"), ) logger.info(f"Generated document logs input successfully : {response}") diff --git a/textract-pipeline/src/lambda/prompt-orchestrator/bottom_up_funcs.py b/textract-pipeline/src/lambda/prompt-orchestrator/bottom_up_funcs.py index c85ee5a..c16a6a2 100644 --- a/textract-pipeline/src/lambda/prompt-orchestrator/bottom_up_funcs.py +++ b/textract-pipeline/src/lambda/prompt-orchestrator/bottom_up_funcs.py @@ -1,4 +1,3 @@ - import dict_operations import postprocess_funcs import prompts @@ -7,12 +6,13 @@ import claude_funcs import difflib + def run_bottom_up(filename, text_dict): """ Processes the text of a document using a two-tiered Bottom Up approach to extract key financial and operational information. - This function first runs BOTTOM_UP_PRIMARY and BOTTOM_UP_SECONDARY prompts, then processes these initial results to further - refine and structure them into a dictionary form. + This function first runs BOTTOM_UP_PRIMARY and BOTTOM_UP_SECONDARY prompts, then processes these initial results to further + refine and structure them into a dictionary form. Parameters: filename (str): The name of the file being processed, used to tag output data. @@ -27,7 +27,7 @@ def run_bottom_up(filename, text_dict): answer_strings = run_bottom_up_primary(text_dict, 8000) answer_dicts = dict_operations.primary_string_to_dict(answer_strings, filename) answer_dicts_filtered = postprocess_funcs.filter_service_column(answer_dicts) - + # Bottom Up Secondary results_dicts = run_bottom_up_secondary(answer_dicts_filtered, text_dict, 8000) @@ -36,9 +36,9 @@ def run_bottom_up(filename, text_dict): # Add Filename for d in results_dicts: - d['Filename'] = filename - - return results_dicts # List of dictionaries + d["Filename"] = filename + + return results_dicts # List of dictionaries def run_bottom_up_primary(text_dict, tokens): @@ -55,15 +55,21 @@ def run_bottom_up_primary(text_dict, tokens): Returns: dict: A dictionary where each key is a page number and the value is the response from the language model. """ - #chunk_dict = preprocess.chunk_text(text_dict) + # chunk_dict = preprocess.chunk_text(text_dict) answer_dict = {} - #for page_number in chunk_dict.keys(): - #if '%' in chunk_dict[page_number] or '$' in chunk_dict[page_number]: - #prompt = prompts.BOTTOM_UP_PRIMARY(chunk_dict[page_number], config.CLIENT_NAME) + # for page_number in chunk_dict.keys(): + # if '%' in chunk_dict[page_number] or '$' in chunk_dict[page_number]: + # prompt = prompts.BOTTOM_UP_PRIMARY(chunk_dict[page_number], config.CLIENT_NAME) for page_number in text_dict.keys(): - if page_number.isdigit() and ('%' in text_dict[page_number] or '$' in text_dict[page_number] or 'percent' in text_dict[page_number].lower()): + if page_number.isdigit() and ( + "%" in text_dict[page_number] + or "$" in text_dict[page_number] + or "percent" in text_dict[page_number].lower() + ): # Run Primary - prompt = prompts.BOTTOM_UP_PRIMARY(text_dict[page_number], config.CLIENT_NAME) + prompt = prompts.BOTTOM_UP_PRIMARY( + text_dict[page_number], config.CLIENT_NAME + ) answer = claude_funcs.invoke_claude_3(prompt, max_tokens=tokens) answer_dict[page_number] = answer return answer_dict @@ -73,10 +79,10 @@ def run_bottom_up_secondary(answer_dicts, text_dict, tokens): """ Executes the secondary Bottom Up processing phase on the results obtained from the primary Bottom Up analysis. - This function enhances the primary results with additional analyses based on configured conditions. Invokes + This function enhances the primary results with additional analyses based on configured conditions. Invokes Claude 3 with tailored prompts to generate structured information that complements the initial results. - Each piece of data processed possibly undergoes several rounds of checks and transformations, ensuring detailed and + Each piece of data processed possibly undergoes several rounds of checks and transformations, ensuring detailed and comprehensive output. Parameters: @@ -85,43 +91,50 @@ def run_bottom_up_secondary(answer_dicts, text_dict, tokens): tokens (int): Token limit for language model invocations. Returns: - list of dict: A list of dictionaries containing enriched and finalized structured data from both the primary and + list of dict: A list of dictionaries containing enriched and finalized structured data from both the primary and secondary analyses. """ # New rows for lesser of temp_dicts = [] - for d in answer_dicts: + for d in answer_dicts: # Bottom Up Lesser if config.RUN_LESSER: lesser_object = run_bottom_up_lesser(d.copy(), text_dict.copy(), tokens) - temp_dicts.append(lesser_object[1]) # Add original object from d - if lesser_object[0]: temp_dicts.append(lesser_object[0]) # If lesser, add lesser object + temp_dicts.append(lesser_object[1]) # Add original object from d + if lesser_object[0]: + temp_dicts.append(lesser_object[0]) # If lesser, add lesser object else: temp_dicts.append(d) - + # Add additional fields to each row final_dicts = [] for d in temp_dicts: if d is not None: - page_num = d['page_num'] - + page_num = d["page_num"] + # Bottom Up Methodology if config.RUN_METHODOLOGY: prompt = prompts.BOTTOM_UP_METHODOLOGY(d) - methodology_answer = claude_funcs.invoke_claude_3(prompt, max_tokens=100) - d['REIMBURSEMENT_METHODOLOGY'] = methodology_answer - + methodology_answer = claude_funcs.invoke_claude_3( + prompt, max_tokens=100 + ) + d["REIMBURSEMENT_METHODOLOGY"] = methodology_answer + # # Bottom Up FS if config.RUN_FS: prompt = prompts.BOTTOM_UP_FS(d) - fs_answer = claude_funcs.invoke_claude_3(prompt, model_id=config.MODEL_ID_CLAUDE3_HAIKU, max_tokens=tokens) + fs_answer = claude_funcs.invoke_claude_3( + prompt, model_id=config.MODEL_ID_CLAUDE3_HAIKU, max_tokens=tokens + ) fs_dict = dict_operations.secondary_string_to_dict(fs_answer) d.update(fs_dict) # # Bottom Up Exception/Escalator if config.RUN_EXCEPTION: prompt = prompts.BOTTOM_UP_EXCEPT_ESC(d, text_dict[page_num]) - exc_answer = claude_funcs.invoke_claude_3(prompt, model_id=config.MODEL_ID_CLAUDE3_HAIKU, max_tokens=tokens) + exc_answer = claude_funcs.invoke_claude_3( + prompt, model_id=config.MODEL_ID_CLAUDE3_HAIKU, max_tokens=tokens + ) exc_dict = dict_operations.secondary_string_to_dict(exc_answer) d.update(exc_dict) @@ -133,9 +146,8 @@ def run_bottom_up_secondary(answer_dicts, text_dict, tokens): d.update(codes_dict) final_dicts.append(d) - - return final_dicts + return final_dicts def run_bottom_up_lesser(d, text_dict, tokens=4000): @@ -158,50 +170,59 @@ def run_bottom_up_lesser(d, text_dict, tokens=4000): containing the original or updated data depending on the presence of such language. """ - def get_least_similar(dict_list, original_dict, field): - min_similarity = float('inf') + def get_least_similar(dict_list, original_dict, field): + min_similarity = float("inf") least_similar_dict = None for dictionary in dict_list: field_text = dictionary[field] - similarity = difflib.SequenceMatcher(None, str(field_text), str(original_dict[field])).ratio() + similarity = difflib.SequenceMatcher( + None, str(field_text), str(original_dict[field]) + ).ratio() if similarity < min_similarity: min_similarity = similarity least_similar_dict = dictionary return least_similar_dict - + def get_lesser_of_dict(dict_list, original_dict): - unique_rates = list({d['REIMBURSEMENT_RATE'] for d in dict_list}) - unique_fees = list({d['REIMBURSEMENT_FLAT_FEE'] for d in dict_list}) + unique_rates = list({d["REIMBURSEMENT_RATE"] for d in dict_list}) + unique_fees = list({d["REIMBURSEMENT_FLAT_FEE"] for d in dict_list}) if len(unique_rates) > 1: - least_similar_dict = get_least_similar(dict_list, original_dict, 'REIMBURSEMENT_RATE') + least_similar_dict = get_least_similar( + dict_list, original_dict, "REIMBURSEMENT_RATE" + ) elif len(unique_fees) > 1: - least_similar_dict = get_least_similar(dict_list, original_dict, 'REIMBURSEMENT_FLAT_FEE') + least_similar_dict = get_least_similar( + dict_list, original_dict, "REIMBURSEMENT_FLAT_FEE" + ) else: - least_similar_dict = get_least_similar(dict_list, original_dict, 'FULL_METHODOLOGY') + least_similar_dict = get_least_similar( + dict_list, original_dict, "FULL_METHODOLOGY" + ) return least_similar_dict - - prompt = prompts.BOTTOM_UP_LESSER(d, text_dict[d['page_num']]) + prompt = prompts.BOTTOM_UP_LESSER(d, text_dict[d["page_num"]]) lesser_of_answer = claude_funcs.invoke_claude_3(prompt, max_tokens=tokens) - lesser_of_dict_list = dict_operations.primary_string_to_dict({d['page_num'] : lesser_of_answer}, d['Filename']) + lesser_of_dict_list = dict_operations.primary_string_to_dict( + {d["page_num"]: lesser_of_answer}, d["Filename"] + ) # print(lesser_of_dict_list) # If lesser of is Y if len(lesser_of_dict_list) > 1: lesser_of_dict = get_lesser_of_dict(lesser_of_dict_list, d) else: - d['LESSER_OF_LANGUAGE_IND'] = 'N' - d['GREATER_OF_LANGUAGE_IND'] = 'N' + d["LESSER_OF_LANGUAGE_IND"] = "N" + d["GREATER_OF_LANGUAGE_IND"] = "N" return ({}, d) - - lesser_of_dict['page_num'] = d['page_num'] - if lesser_of_dict['LESSER_OF_LANGUAGE_IND'] == 'Y' or lesser_of_dict['GREATER_OF_LANGUAGE_IND'] == 'Y': - lesser_of_dict['SERVICE'] = d['SERVICE'] + + lesser_of_dict["page_num"] = d["page_num"] + if ( + lesser_of_dict["LESSER_OF_LANGUAGE_IND"] == "Y" + or lesser_of_dict["GREATER_OF_LANGUAGE_IND"] == "Y" + ): + lesser_of_dict["SERVICE"] = d["SERVICE"] return (lesser_of_dict, d) else: - d['LESSER_OF_LANGUAGE_IND'] = 'N' - d['GREATER_OF_LANGUAGE_IND'] = 'N' + d["LESSER_OF_LANGUAGE_IND"] = "N" + d["GREATER_OF_LANGUAGE_IND"] = "N" return ({}, d) - - - diff --git a/textract-pipeline/src/lambda/prompt-orchestrator/carveouts.py b/textract-pipeline/src/lambda/prompt-orchestrator/carveouts.py index 13843ad..675cb88 100644 --- a/textract-pipeline/src/lambda/prompt-orchestrator/carveouts.py +++ b/textract-pipeline/src/lambda/prompt-orchestrator/carveouts.py @@ -3,11 +3,12 @@ from difflib import SequenceMatcher import re import time -import prompts +import prompts import claude_funcs + def get_closest_substring_match(val, valid_values): - """ Returns the first match from valid_values using a case-insensitive substring match. """ + """Returns the first match from valid_values using a case-insensitive substring match.""" if pd.isna(val): return None val = val.strip().upper() @@ -16,6 +17,7 @@ def get_closest_substring_match(val, valid_values): return valid_value return None + def get_best_carveout_from_claude(carveout_list, service): print("using claude to get best carveout...") prompt = f"Given the service '{service}', please choose the best matching carveout from the following list: {', '.join(carveout_list)}. ONLY SELECT ONE FROM THE LIST. DO NOT RETURN A SENTENCE" @@ -27,6 +29,7 @@ def get_best_carveout_from_claude(carveout_list, service): print(f"Error occurred: {e}. Waiting for 60 seconds before retrying...") time.sleep(60) + def check_prov_type_similarity(prov_type, carveout): print("checking similarity between prov_type and carveout with claude...") prompt = f"Do the provider type '{prov_type}' and the carveout '{carveout}' mean the same thing or are they very similar? ONLY ANSWER WITH 'True' OR 'False'" @@ -38,6 +41,7 @@ def check_prov_type_similarity(prov_type, carveout): print(f"Error occurred: {e}. Waiting for 60 seconds before retrying...") time.sleep(60) + def label_services(filepath, carveout_list, output_filepath): # Read the CSV file df = pd.read_csv(filepath) @@ -46,13 +50,24 @@ def label_services(filepath, carveout_list, output_filepath): df = df.head(500) # Define primary service terms - primary_terms = ['Covered Services', 'Inpatient Services', 'Outpatient Services', 'Physician Services', - 'Inpatient', 'Outpatient', 'Physician', 'Medical', 'Surgical', 'Diagnostic', 'Therapeutic'] + primary_terms = [ + "Covered Services", + "Inpatient Services", + "Outpatient Services", + "Physician Services", + "Inpatient", + "Outpatient", + "Physician", + "Medical", + "Surgical", + "Diagnostic", + "Therapeutic", + ] # Initialize columns for the results - df['carveout_matched'] = '' - df['label'] = '' - df['IS_CARVEOUT'] = '' + df["carveout_matched"] = "" + df["label"] = "" + df["IS_CARVEOUT"] = "" # Initialize a nested dictionary to track occurrences for each Filename and TD_LOB occurrence_tracker = {} @@ -61,10 +76,10 @@ def label_services(filepath, carveout_list, output_filepath): # Iterate through the rows in the DataFrame for index, row in df.iterrows(): - service = str(row['SERVICE']) - prov_type = str(row['PROV_TYPE']) if pd.notna(row['PROV_TYPE']) else '' - td_lob = row['TD_LOB'] - filename = row['Filename'] + service = str(row["SERVICE"]) + prov_type = str(row["PROV_TYPE"]) if pd.notna(row["PROV_TYPE"]) else "" + td_lob = row["TD_LOB"] + filename = row["Filename"] # Initialize the nested dictionary for the filename and TD_LOB if not present if filename not in occurrence_tracker: @@ -82,49 +97,57 @@ def label_services(filepath, carveout_list, output_filepath): # Determine the label based on the count of this primary term for the given Filename and TD_LOB count = occurrence_tracker[filename][td_lob][primary] if count == 0: - df.at[index, 'label'] = 'primary' + df.at[index, "label"] = "primary" elif count == 1: - df.at[index, 'label'] = 'secondary' + df.at[index, "label"] = "secondary" elif count == 2: - df.at[index, 'label'] = 'tertiary' + df.at[index, "label"] = "tertiary" else: - df.at[index, 'label'] = 'additional' + df.at[index, "label"] = "additional" occurrence_tracker[filename][td_lob][primary] += 1 - df.at[index, 'IS_CARVEOUT'] = 'N' + df.at[index, "IS_CARVEOUT"] = "N" else: # If PROV_TYPE is blank, continue with carveout determination - best_carveout_response = get_best_carveout_from_claude(carveout_list, service) + best_carveout_response = get_best_carveout_from_claude( + carveout_list, service + ) matched_carveout = best_carveout_response # Use the model response directly - df.at[index, 'carveout_matched'] = matched_carveout + df.at[index, "carveout_matched"] = matched_carveout if matched_carveout: # Check similarity between PROV_TYPE and carveout - prov_type_similarity_score = SequenceMatcher(None, prov_type, matched_carveout).ratio() + prov_type_similarity_score = SequenceMatcher( + None, prov_type, matched_carveout + ).ratio() if prov_type_similarity_score < 0.8: - prov_type_similarity = check_prov_type_similarity(prov_type, matched_carveout) + prov_type_similarity = check_prov_type_similarity( + prov_type, matched_carveout + ) else: prov_type_similarity = "True" - if re.search(r'\b(True|Yes)\b', prov_type_similarity, re.IGNORECASE): + if re.search(r"\b(True|Yes)\b", prov_type_similarity, re.IGNORECASE): if matched_carveout not in occurrence_tracker[filename][td_lob]: - print('claude returned a match between provider type and service') + print( + "claude returned a match between provider type and service" + ) occurrence_tracker[filename][td_lob][matched_carveout] = 0 count = occurrence_tracker[filename][td_lob][matched_carveout] if count == 0: - df.at[index, 'label'] = 'primary' + df.at[index, "label"] = "primary" elif count == 1: - df.at[index, 'label'] = 'secondary' + df.at[index, "label"] = "secondary" elif count == 2: - df.at[index, 'label'] = 'tertiary' + df.at[index, "label"] = "tertiary" else: - df.at[index, 'label'] = 'additional' + df.at[index, "label"] = "additional" occurrence_tracker[filename][td_lob][matched_carveout] += 1 - df.at[index, 'IS_CARVEOUT'] = 'N' + df.at[index, "IS_CARVEOUT"] = "N" else: - df.at[index, 'IS_CARVEOUT'] = 'Y' + df.at[index, "IS_CARVEOUT"] = "Y" else: - df.at[index, 'IS_CARVEOUT'] = 'Y' + df.at[index, "IS_CARVEOUT"] = "Y" carveout_count += 1 @@ -136,4 +159,5 @@ def label_services(filepath, carveout_list, output_filepath): # Final save of the updated DataFrame to a new CSV file df.to_csv(output_filepath, index=False) -label_services('service.csv', prompts.get_carveout_list(), 'carveouts50.csv') + +label_services("service.csv", prompts.get_carveout_list(), "carveouts50.csv") diff --git a/textract-pipeline/src/lambda/prompt-orchestrator/claude_funcs.py b/textract-pipeline/src/lambda/prompt-orchestrator/claude_funcs.py index 9d1e9f8..75abfac 100644 --- a/textract-pipeline/src/lambda/prompt-orchestrator/claude_funcs.py +++ b/textract-pipeline/src/lambda/prompt-orchestrator/claude_funcs.py @@ -1,9 +1,9 @@ - import json import anthropic import config + # Claude calls def invoke_claude_2(prompt, max_tokens): """ @@ -21,28 +21,30 @@ def invoke_claude_2(prompt, max_tokens): str: The text generated by the Claude 2 model in response to the input prompt. """ body = json.dumps( - {"prompt": anthropic.HUMAN_PROMPT + prompt + anthropic.AI_PROMPT, - "max_tokens_to_sample": max_tokens, - "temperature":0.0, - "top_p":1, - "top_k":250, - "stop_sequences":[anthropic.HUMAN_PROMPT] + { + "prompt": anthropic.HUMAN_PROMPT + prompt + anthropic.AI_PROMPT, + "max_tokens_to_sample": max_tokens, + "temperature": 0.0, + "top_p": 1, + "top_k": 250, + "stop_sequences": [anthropic.HUMAN_PROMPT], } - ) + ) response = config.BEDROCK_RUNTIME.invoke_model( - body=body, - modelId=config.MODEL_ID_CLAUDE2, - accept="application/json", - contentType="application/json" + body=body, + modelId=config.MODEL_ID_CLAUDE2, + accept="application/json", + contentType="application/json", ) response_body = json.loads(response.get("body").read()) - response_text = response_body['completion'] + response_text = response_body["completion"] return response_text -def invoke_claude_3(prompt, model_id=config.MODEL_ID_CLAUDE3_SONNET, max_tokens = 0): + +def invoke_claude_3(prompt, model_id=config.MODEL_ID_CLAUDE3_SONNET, max_tokens=0): """ Invokes the Claude 3 language model with specified parameters to generate a response based on the input prompt. @@ -57,32 +59,25 @@ def invoke_claude_3(prompt, model_id=config.MODEL_ID_CLAUDE3_SONNET, max_tokens Returns: str: The text generated by the Claude 2 model in response to the input prompt. """ - prompt = prompt - body = json.dumps({ - "anthropic_version": "bedrock-2023-05-31", - "max_tokens": max_tokens, - "temperature": 0.0, - "messages": [ - { - "role": "user", - "content": [ - { - "type": "text", - "text":prompt - } - ] - } - ] - } - ) + prompt = prompt + body = json.dumps( + { + "anthropic_version": "bedrock-2023-05-31", + "max_tokens": max_tokens, + "temperature": 0.0, + "messages": [ + {"role": "user", "content": [{"type": "text", "text": prompt}]} + ], + } + ) response = config.BEDROCK_RUNTIME.invoke_model( - body=body, - modelId=model_id, - accept="application/json", - contentType="application/json" + body=body, + modelId=model_id, + accept="application/json", + contentType="application/json", ) response_body = json.loads(response.get("body").read()) - return (response_body['content'][0]['text']) \ No newline at end of file + return response_body["content"][0]["text"] diff --git a/textract-pipeline/src/lambda/prompt-orchestrator/config.py b/textract-pipeline/src/lambda/prompt-orchestrator/config.py index 0ee313b..afb1e88 100644 --- a/textract-pipeline/src/lambda/prompt-orchestrator/config.py +++ b/textract-pipeline/src/lambda/prompt-orchestrator/config.py @@ -1,39 +1,58 @@ - - from datetime import datetime import boto3 -#from llama_index.llms.bedrock import Bedrock -#pip install llama-index-llms-bedrock + +# from llama_index.llms.bedrock import Bedrock +# pip install llama-index-llms-bedrock # General Settings -TEST = False # True to run test prompt - just for testing model connection +TEST = False # True to run test prompt - just for testing model connection VERBOSE = True -CLIENT_NAME = '' +CLIENT_NAME = "" TODAY = datetime.now().strftime("%Y%m%d") # Valid values -VALID_LOBS = ['MEDICARE', 'MEDICARE ADVANTAGE', 'MEDICAID', 'MARKETPLACE', 'COMMERCIAL', 'GROUP', 'MEDICARE-MEDICAID'] -VALID_PROGRAMS = ['CHIP', 'CHIP-P', 'CHIP-PERINATE', 'STAR', 'STAR+PLUS', 'MA', 'DUAL SPECIAL NEEDS PLAN', 'DSNP', 'DUAL'] -VALID_NETWORKS = ['HMO', 'PPO', 'EPO', 'POS', 'FFS'] - +VALID_LOBS = [ + "MEDICARE", + "MEDICARE ADVANTAGE", + "MEDICAID", + "MARKETPLACE", + "COMMERCIAL", + "GROUP", + "MEDICARE-MEDICAID", +] +VALID_PROGRAMS = [ + "CHIP", + "CHIP-P", + "CHIP-PERINATE", + "STAR", + "STAR+PLUS", + "MA", + "DUAL SPECIAL NEEDS PLAN", + "DSNP", + "DUAL", +] +VALID_NETWORKS = ["HMO", "PPO", "EPO", "POS", "FFS"] + # Input Settings -READ_MODE = '_LOCAL_' # OR '_S3_' +READ_MODE = "_LOCAL_" # OR '_S3_' # Output Settings -WRITE_OUTPUT = True # True writes csvs, False prints result in console but no output written -OUTPUT_DIRECTORY = 'output' -CONSOLIDATED_OUTPUT_DIRECTORY = 'output_consolidated' -OUTPUT_CSV_PATH = f'consolidated_output_{TODAY}.csv' -TD_RESULTS_NAME = 'td_results.csv' -BU_RESULTS_NAME = 'bu_results.csv' -UNPROCESSED_RESULTS_NAME = 'combined_results_unprocessed.csv' -PROCESSED_RESULTS_NAME = 'combined_results_post_processed.csv' +WRITE_OUTPUT = ( + True # True writes csvs, False prints result in console but no output written +) +OUTPUT_DIRECTORY = "output" +CONSOLIDATED_OUTPUT_DIRECTORY = "output_consolidated" +OUTPUT_CSV_PATH = f"consolidated_output_{TODAY}.csv" +TD_RESULTS_NAME = "td_results.csv" +BU_RESULTS_NAME = "bu_results.csv" +UNPROCESSED_RESULTS_NAME = "combined_results_unprocessed.csv" +PROCESSED_RESULTS_NAME = "combined_results_post_processed.csv" # Multithread Settings MAX_WORKERS = 3 # Prompt Debugging -RUN_PRIMARY = True # Always True +RUN_PRIMARY = True # Always True RUN_LOB = True RUN_LESSER = True RUN_METHODOLOGY = True @@ -43,65 +62,87 @@ RUN_CODES = True # Postprocessing Settings FUZZY_MATCH_THRESHOLD = 0.8 -VALID_COLUMNS = ['Filename', - 'SERVICE', 'SERVICE_PG', - 'REIMBURSEMENT_FLAT_FEE', 'REIMBURSEMENT_FLAT_FEE_PG', - 'REIMBURSEMENT_RATE', 'REIMBURSEMENT_RATE_PG', - 'FULL_METHODOLOGY', 'FULL_METHODOLOGY_PG', - 'REIMBURSEMENT_METHODOLOGY', 'REIMBURSEMENT_METHODOLOGY_PG', - 'REIMBURSEMENT_FEE_SCHEDULE', 'REIMBURSEMENT_FEE_SCHEDULE_PG', - 'REIMBURSEMENT_FEE_SCHEDULE_VERSION', 'REIMBURSEMENT_FEE_SCHEDULE_VERSION_PG', - 'LESSER_OF_LANGUAGE_IND', 'LESSER_OF_LANGUAGE_IND_PG', - 'GREATER_OF_LANGUAGE_IND', 'GREATER_OF_LANGUAGE_IND_PG', - 'CONTRACT_LOB', 'CONTRACT_LOB_PG', - 'CONTRACT_MARKETPLACE_METAL_LEVEL', 'CONTRACT_MARKETPLACE_METAL_LEVEL_PG', - 'CONTRACT_NETWORK', 'CONTRACT_NETWORK_PG', - 'PRODUCT', 'PRODUCT_PG', - 'CONTRACT_PROGRAM', 'CONTRACT_PROGRAM_PG', - 'REIMBURSEMENT_PROC_CODES', 'REIMBURSEMENT_PROC_CODES_PG', - 'REIMBURSEMENT_PROC_CODE_MODIFIERS', 'REIMBURSEMENT_PROC_CODE_MODIFIERS_PG', - 'REIMBURSEMENT_REVENUE_CODES', 'REIMBURSEMENT_REVENUE_CODES_PG', - 'REIMBURSEMENT_STATUS_INDICATOR_CODES', 'REIMBURSEMENT_STATUS_INDICATOR_CODES_PG', - 'REIMBURSEMENT_DIAG_CODES', 'REIMBURSEMENT_DIAG_CODES_PG', - 'REIMBURSEMENT_GROUPER_CODES', 'REIMBURSEMENT_GROUPER_CODES_PG', - 'REIMBURSEMENT_GROUPER', 'REIMBURSEMENT_GROUPER_PG', - 'REIMBURSEMENT_PLACEOFSERVICE_CODES', 'REIMBURSEMENT_PLACEOFSERVICE_CODES_PG', - 'REIMBURSEMENT_ADMITTYPE_CODES', 'REIMBURSEMENT_ADMITTYPE_CODES_PG', - 'REIMBURSEMENT_EXCEPTION_IND', 'REIMBURSEMENT_EXCEPTION_IND_PG', - 'REIMBURSEMENT_DESCRIBE_EXCEPTION', 'REIMBURSEMENT_DESCRIBE_EXCEPTION_PG', - 'LOB_PRICING_TERMS_EFFECTIVE_DATE', 'LOB_PRICING_TERMS_EFFECTIVE_DATE_PG', - 'LOB_PRICING_TERMS_TERMINATION_DATE', 'LOB_PRICING_TERMS_TERMINATION_DATE_PG', - 'Corrected_LOB', - 'Corrected_PROGRAM', - 'Corrected_NETWORK'] +VALID_COLUMNS = [ + "Filename", + "SERVICE", + "SERVICE_PG", + "REIMBURSEMENT_FLAT_FEE", + "REIMBURSEMENT_FLAT_FEE_PG", + "REIMBURSEMENT_RATE", + "REIMBURSEMENT_RATE_PG", + "FULL_METHODOLOGY", + "FULL_METHODOLOGY_PG", + "REIMBURSEMENT_METHODOLOGY", + "REIMBURSEMENT_METHODOLOGY_PG", + "REIMBURSEMENT_FEE_SCHEDULE", + "REIMBURSEMENT_FEE_SCHEDULE_PG", + "REIMBURSEMENT_FEE_SCHEDULE_VERSION", + "REIMBURSEMENT_FEE_SCHEDULE_VERSION_PG", + "LESSER_OF_LANGUAGE_IND", + "LESSER_OF_LANGUAGE_IND_PG", + "GREATER_OF_LANGUAGE_IND", + "GREATER_OF_LANGUAGE_IND_PG", + "CONTRACT_LOB", + "CONTRACT_LOB_PG", + "CONTRACT_MARKETPLACE_METAL_LEVEL", + "CONTRACT_MARKETPLACE_METAL_LEVEL_PG", + "CONTRACT_NETWORK", + "CONTRACT_NETWORK_PG", + "PRODUCT", + "PRODUCT_PG", + "CONTRACT_PROGRAM", + "CONTRACT_PROGRAM_PG", + "REIMBURSEMENT_PROC_CODES", + "REIMBURSEMENT_PROC_CODES_PG", + "REIMBURSEMENT_PROC_CODE_MODIFIERS", + "REIMBURSEMENT_PROC_CODE_MODIFIERS_PG", + "REIMBURSEMENT_REVENUE_CODES", + "REIMBURSEMENT_REVENUE_CODES_PG", + "REIMBURSEMENT_STATUS_INDICATOR_CODES", + "REIMBURSEMENT_STATUS_INDICATOR_CODES_PG", + "REIMBURSEMENT_DIAG_CODES", + "REIMBURSEMENT_DIAG_CODES_PG", + "REIMBURSEMENT_GROUPER_CODES", + "REIMBURSEMENT_GROUPER_CODES_PG", + "REIMBURSEMENT_GROUPER", + "REIMBURSEMENT_GROUPER_PG", + "REIMBURSEMENT_PLACEOFSERVICE_CODES", + "REIMBURSEMENT_PLACEOFSERVICE_CODES_PG", + "REIMBURSEMENT_ADMITTYPE_CODES", + "REIMBURSEMENT_ADMITTYPE_CODES_PG", + "REIMBURSEMENT_EXCEPTION_IND", + "REIMBURSEMENT_EXCEPTION_IND_PG", + "REIMBURSEMENT_DESCRIBE_EXCEPTION", + "REIMBURSEMENT_DESCRIBE_EXCEPTION_PG", + "LOB_PRICING_TERMS_EFFECTIVE_DATE", + "LOB_PRICING_TERMS_EFFECTIVE_DATE_PG", + "LOB_PRICING_TERMS_TERMINATION_DATE", + "LOB_PRICING_TERMS_TERMINATION_DATE_PG", + "Corrected_LOB", + "Corrected_PROGRAM", + "Corrected_NETWORK", +] # File Paths -LOCAL_PATH = 'data/test/' # Replace with local +LOCAL_PATH = "data/test/" # Replace with local # S3 Settings -S3_CLIENT = boto3.client('s3', - region_name="us-east-2" - ) +S3_CLIENT = boto3.client("s3", region_name="us-east-2") BUCKET = "doczy-dev-infra-textract" -PREFIX = "batches/batch_1/contract-text-file/" # replace with s3 path +PREFIX = "batches/batch_1/contract-text-file/" # replace with s3 path # Bedrock Settings -BEDROCK_RUNTIME = boto3.client(service_name="bedrock-runtime", - region_name="us-east-1" -) -MODEL_ID_CLAUDE3_HAIKU = 'anthropic.claude-3-haiku-20240307-v1:0' -MODEL_ID_CLAUDE3_SONNET = 'anthropic.claude-3-sonnet-20240229-v1:0' -MODEL_ID_CLAUDE2 = 'anthropic.claude-instant-v1' +BEDROCK_RUNTIME = boto3.client(service_name="bedrock-runtime", region_name="us-east-1") +MODEL_ID_CLAUDE3_HAIKU = "anthropic.claude-3-haiku-20240307-v1:0" +MODEL_ID_CLAUDE3_SONNET = "anthropic.claude-3-sonnet-20240229-v1:0" +MODEL_ID_CLAUDE2 = "anthropic.claude-instant-v1" # LLama Model # LLAMA_MODEL = Bedrock( - # model=MODEL_ID_CLAUDE3_HAIKU, - # aws_access_key_id=AWS_ACCESS_KEY_ID, - # aws_secret_access_key=AWS_SECRET_ACCESS_KEY, - # aws_session_token=AWS_SESSION_TOKEN, - # region_name="us-east-1" - # ) - - - +# model=MODEL_ID_CLAUDE3_HAIKU, +# aws_access_key_id=AWS_ACCESS_KEY_ID, +# aws_secret_access_key=AWS_SECRET_ACCESS_KEY, +# aws_session_token=AWS_SESSION_TOKEN, +# region_name="us-east-1" +# ) diff --git a/textract-pipeline/src/lambda/prompt-orchestrator/consolidate_output.py b/textract-pipeline/src/lambda/prompt-orchestrator/consolidate_output.py index 3a7b3bf..0ce6f46 100644 --- a/textract-pipeline/src/lambda/prompt-orchestrator/consolidate_output.py +++ b/textract-pipeline/src/lambda/prompt-orchestrator/consolidate_output.py @@ -1,4 +1,3 @@ - import os import pandas as pd @@ -6,13 +5,19 @@ import config all_dfs = [] for file in os.listdir(config.OUTPUT_DIRECTORY): - if config.PROCESSED_RESULTS_NAME in os.listdir(os.path.join(config.OUTPUT_DIRECTORY, file)): + if config.PROCESSED_RESULTS_NAME in os.listdir( + os.path.join(config.OUTPUT_DIRECTORY, file) + ): try: - df = pd.read_csv(f'{config.OUTPUT_DIRECTORY}/{file}/{config.PROCESSED_RESULTS_NAME}') + df = pd.read_csv( + f"{config.OUTPUT_DIRECTORY}/{file}/{config.PROCESSED_RESULTS_NAME}" + ) all_dfs.append(df) except: continue - + final_df = pd.concat(all_dfs, ignore_index=True) -final_df.to_excel(os.path.join(config.CONSOLIDATED_OUTPUT_DIRECTORY, config.OUTPUT_CSV_PATH)) \ No newline at end of file +final_df.to_excel( + os.path.join(config.CONSOLIDATED_OUTPUT_DIRECTORY, config.OUTPUT_CSV_PATH) +) diff --git a/textract-pipeline/src/lambda/prompt-orchestrator/dict_operations.py b/textract-pipeline/src/lambda/prompt-orchestrator/dict_operations.py index 78ee6fa..561ca17 100644 --- a/textract-pipeline/src/lambda/prompt-orchestrator/dict_operations.py +++ b/textract-pipeline/src/lambda/prompt-orchestrator/dict_operations.py @@ -1,8 +1,8 @@ - import re import json from collections import defaultdict + def secondary_string_to_dict(dict_string): """ Converts a string representation of a dictionary into an actual dictionary object, handling potential formatting issues. @@ -17,20 +17,21 @@ def secondary_string_to_dict(dict_string): Returns: dict: The dictionary obtained from parsing the cleaned and corrected string. """ - dict_string = dict_string.replace('<<<', '') - dict_string = dict_string.replace('>>>', '') - dict_string = dict_string.replace('\n', '') # Remove new line + dict_string = dict_string.replace("<<<", "") + dict_string = dict_string.replace(">>>", "") + dict_string = dict_string.replace("\n", "") # Remove new line - start_index = dict_string.find('{') - end_index = dict_string.rfind('}') + 1 + start_index = dict_string.find("{") + end_index = dict_string.rfind("}") + 1 dict_substring = dict_string[start_index:end_index] try: result_dict = json.loads(dict_substring) except: dict_substring = dict_substring.replace("'", '"') - result_dict = json.loads(dict_substring) + result_dict = json.loads(dict_substring) return result_dict + def primary_string_to_dict(string_dict, filename): """ Converts a dictionary of strings, where each string represents multiple dictionary entries, into a list of dictionaries, @@ -49,24 +50,25 @@ def primary_string_to_dict(string_dict, filename): their page number and filename. """ data = [] - pattern = r'\{.*?\}' + pattern = r"\{.*?\}" for page_num in string_dict.keys(): primary_list = string_dict[page_num] - dicts = primary_list.split('[')[1] # Strip front - dicts = dicts.split(']')[0] # Strip back - dicts = dicts.replace('\n', '') # Remove new lines - dicts = dicts.replace('<<<', '') - dicts = dicts.replace('>>>', '') + dicts = primary_list.split("[")[1] # Strip front + dicts = dicts.split("]")[0] # Strip back + dicts = dicts.replace("\n", "") # Remove new lines + dicts = dicts.replace("<<<", "") + dicts = dicts.replace(">>>", "") dicts = re.sub(r"(? 1: if response_text.rsplit("}", 1)[0].strip()[-1] == '"': response_text = response_text.rsplit("}", 1)[0] + "}" - elif response_text.rsplit("}", 1)[0].strip()[-1] == '}': + elif response_text.rsplit("}", 1)[0].strip()[-1] == "}": response_text = response_text.rsplit("}", 1)[0] else: response_text = response_text.rstrip(",") + "}" @@ -384,144 +405,340 @@ def json_parsing(response_text, context): field_l = list(response_dict.keys()) answer_l = list(response_dict.values()) - field_dict = {k: v for k, v in response_dict.items() if not k.endswith('_PG')} - page_dict = {k: v for k, v in response_dict.items() if k.endswith('_PG')} + field_dict = {k: v for k, v in response_dict.items() if not k.endswith("_PG")} + page_dict = {k: v for k, v in response_dict.items() if k.endswith("_PG")} page_dict = {k[:-3]: v for k, v in response_dict.items()} field_l = list(field_dict.keys()) answer_l = list(field_dict.values()) page_no_l = [page_dict.get(x, "") for x in field_l] try: - location_l = [context.find(a, context.find("Start of Page No. = "+str(p))) if isinstance( - a, str) and a != "" else -1 for a, p in zip(answer_l, page_no_l)] + location_l = [ + ( + context.find(a, context.find("Start of Page No. = " + str(p))) + if isinstance(a, str) and a != "" + else -1 + ) + for a, p in zip(answer_l, page_no_l) + ] except: - location_l = [context.find(answer) if isinstance(answer, str) and answer != "" else -1 for answer in answer_l] - snippet_l = [' '.join(context[:location].split('.')[-4:]) + ' ' + ' '.join(context[location:].split('. ')[:5]) if location != -1 else ' ' for location in location_l] + location_l = [ + context.find(answer) if isinstance(answer, str) and answer != "" else -1 + for answer in answer_l + ] + snippet_l = [ + ( + " ".join(context[:location].split(".")[-4:]) + + " " + + " ".join(context[location:].split(". ")[:5]) + if location != -1 + else " " + ) + for location in location_l + ] logger.info(f"Parsed fields: {field_l}") return field_l, answer_l, page_no_l, location_l, snippet_l + def post_processing(answer_list): logger.info("Post-processing answers.") try: - answer_list = [" " if any(val in str(answer) for val in ['do not have', 'do not see', 'does not specify', 'Does not specify', 'does not explicitly', 'N/A', "don't know", "do not see", 'Not specified', "Don't know", "don't see", "don't have", "Does not apply", "Nothing found", "None"]) else answer for answer in answer_list] + answer_list = [ + ( + " " + if any( + val in str(answer) + for val in [ + "do not have", + "do not see", + "does not specify", + "Does not specify", + "does not explicitly", + "N/A", + "don't know", + "do not see", + "Not specified", + "Don't know", + "don't see", + "don't have", + "Does not apply", + "Nothing found", + "None", + ] + ) + else answer + ) + for answer in answer_list + ] answer_list = [str(answer).rstrip(".") for answer in answer_list] - answer_list = [answer if str(answer) != "one-year" else "1 year" for answer in answer_list] - answer_list = [answer if str(answer) != "one year" else "1 year" for answer in answer_list] - answer_list = [answer if str(answer) != "one" else "1 year" for answer in answer_list] - answer_list = [answer if str(answer) != "one (1) year" else "1 year" for answer in answer_list] - answer_list = [answer if str(answer) != "twelve" else "1 year" for answer in answer_list] - answer_list = [answer if str(answer) != "XI" else "11" for answer in answer_list] - answer_list = [answer if str(answer) != "Third" else "3" for answer in answer_list] - answer_list = [answer if str(answer) != "Six" else "6" for answer in answer_list] + answer_list = [ + answer if str(answer) != "one-year" else "1 year" for answer in answer_list + ] + answer_list = [ + answer if str(answer) != "one year" else "1 year" for answer in answer_list + ] + answer_list = [ + answer if str(answer) != "one" else "1 year" for answer in answer_list + ] + answer_list = [ + answer if str(answer) != "one (1) year" else "1 year" + for answer in answer_list + ] + answer_list = [ + answer if str(answer) != "twelve" else "1 year" for answer in answer_list + ] + answer_list = [ + answer if str(answer) != "XI" else "11" for answer in answer_list + ] + answer_list = [ + answer if str(answer) != "Third" else "3" for answer in answer_list + ] + answer_list = [ + answer if str(answer) != "Six" else "6" for answer in answer_list + ] logger.info(f"Post-processed answers: {answer_list}") return answer_list except: logger.error("Error in post-processing answers.") return ["post processing error"] * len(answer_list) + def compare_with_actuals(df): logger.info("Comparing extracted values with actual values.") query = 'SELECT * FROM "TRAINING_DATA_RAW"' field_values = read_from_db(query) - field_values['Contract Name'] = field_values['DOCUMENT_NAME'] - field_values_2 = pd.DataFrame(columns=['Contract Name', 'SF_DB_COL_NAME', 'Actual Value Stored', 'Original Page Number']) + field_values["Contract Name"] = field_values["DOCUMENT_NAME"] + field_values_2 = pd.DataFrame( + columns=[ + "Contract Name", + "SF_DB_COL_NAME", + "Actual Value Stored", + "Original Page Number", + ] + ) - for contract in set(field_values['Contract Name']): - field_values_1 = field_values[field_values['Contract Name'] == contract].head(1).transpose().reset_index() - field_values_1.columns = ['SF_DB_COL_NAME', 'Actual Value Stored'] - field_values_p1 = field_values_1[~field_values_1['SF_DB_COL_NAME'].str.endswith('_PG')] - field_values_p2 = field_values_1[field_values_1['SF_DB_COL_NAME'].str.endswith('_PG')] - field_values_p2.columns = ['SF_DB_COL_NAME', 'Original Page Number'] - field_values_p2["SF_DB_COL_NAME"] = field_values_p2["SF_DB_COL_NAME"].str.replace("_PG", "") + for contract in set(field_values["Contract Name"]): + field_values_1 = ( + field_values[field_values["Contract Name"] == contract] + .head(1) + .transpose() + .reset_index() + ) + field_values_1.columns = ["SF_DB_COL_NAME", "Actual Value Stored"] + field_values_p1 = field_values_1[ + ~field_values_1["SF_DB_COL_NAME"].str.endswith("_PG") + ] + field_values_p2 = field_values_1[ + field_values_1["SF_DB_COL_NAME"].str.endswith("_PG") + ] + field_values_p2.columns = ["SF_DB_COL_NAME", "Original Page Number"] + field_values_p2["SF_DB_COL_NAME"] = field_values_p2[ + "SF_DB_COL_NAME" + ].str.replace("_PG", "") - field_values_1 = pd.merge(field_values_p1, field_values_p2, how='left', on=['SF_DB_COL_NAME']) - field_values_1['Contract Name'] = contract + field_values_1 = pd.merge( + field_values_p1, field_values_p2, how="left", on=["SF_DB_COL_NAME"] + ) + field_values_1["Contract Name"] = contract field_values_2 = pd.concat([field_values_2, field_values_1], ignore_index=True) - df.rename(columns={'Field Name': 'SF_DB_COL_NAME'}, inplace=True) - df['Contract Name'] = df['Contract Name'].str[:-3] + 'pdf' - fields_selected = list(df['SF_DB_COL_NAME']) - contracted_selected = list(df['Contract Name']) - df = pd.merge(df, field_values_2, how='right', on=['Contract Name', 'SF_DB_COL_NAME']) - df = df[df['SF_DB_COL_NAME'].isin(fields_selected)] - df = df[df['Contract Name'].isin(contracted_selected)] + df.rename(columns={"Field Name": "SF_DB_COL_NAME"}, inplace=True) + df["Contract Name"] = df["Contract Name"].str[:-3] + "pdf" + fields_selected = list(df["SF_DB_COL_NAME"]) + contracted_selected = list(df["Contract Name"]) + df = pd.merge( + df, field_values_2, how="right", on=["Contract Name", "SF_DB_COL_NAME"] + ) + df = df[df["SF_DB_COL_NAME"].isin(fields_selected)] + df = df[df["Contract Name"].isin(contracted_selected)] - df['Original Page Number'] = df['Original Page Number'].apply(lambda x: re.search(r'\d+', x).group() if isinstance(x, str) and re.search(r'\d+', x) is not None else " ") - df_date = df[df['SF_DB_COL_NAME'].str.contains('_DT', na=False)] - df_others = df[~df['SF_DB_COL_NAME'].str.contains('_DT', na=False)] - df_date['Actual Value Stored'] = pd.to_datetime(df_date['Actual Value Stored'], errors='coerce').dt.strftime('%Y-%m-%d').fillna(" ") - df_date['New Extracted value'] = pd.to_datetime(df_date['New Extracted value'], errors='coerce').dt.strftime('%Y-%m-%d').fillna(" ") + df["Original Page Number"] = df["Original Page Number"].apply( + lambda x: ( + re.search(r"\d+", x).group() + if isinstance(x, str) and re.search(r"\d+", x) is not None + else " " + ) + ) + df_date = df[df["SF_DB_COL_NAME"].str.contains("_DT", na=False)] + df_others = df[~df["SF_DB_COL_NAME"].str.contains("_DT", na=False)] + df_date["Actual Value Stored"] = ( + pd.to_datetime(df_date["Actual Value Stored"], errors="coerce") + .dt.strftime("%Y-%m-%d") + .fillna(" ") + ) + df_date["New Extracted value"] = ( + pd.to_datetime(df_date["New Extracted value"], errors="coerce") + .dt.strftime("%Y-%m-%d") + .fillna(" ") + ) df = pd.concat([df_date, df_others], ignore_index=True) - df.sort_values(['SF_DB_COL_NAME', 'Contract Name'], inplace=True) + df.sort_values(["SF_DB_COL_NAME", "Contract Name"], inplace=True) df.fillna(" ", inplace=True) - df['Actual Value Stored'] = df['Actual Value Stored'].apply(lambda x: x.strip() if isinstance(x, str) else '') - df['New Extracted value'] = df['New Extracted value'].apply(lambda x: x.strip() if isinstance(x, str) else '') - actual_value_list = list(df['Actual Value Stored']) - actual_value_list = [answer if str(answer) != "12 months" else "1 year" for answer in actual_value_list] - actual_value_list = [answer if str(answer) != "Fifth" else "5" for answer in actual_value_list] - actual_value_list = [answer if str(answer) != "Seventh" else "7" for answer in actual_value_list] - answer_list = list(df['New Extracted value']) - actual_value_list = [s.replace('-', '').replace(' ', '').replace('[', '').replace(']', '').lower() for s in actual_value_list] - answer_list = [s.replace('-', '').replace(' ', '').replace('[', '').replace(']', '').lower() for s in answer_list] - result_list = [(i in j) or (j in i) if isinstance(i, str) and isinstance(j, str) and ((i != '') == (j != '')) else False for i, j in zip(actual_value_list, answer_list)] - df['Result'] = [str(x) for x in result_list] + df["Actual Value Stored"] = df["Actual Value Stored"].apply( + lambda x: x.strip() if isinstance(x, str) else "" + ) + df["New Extracted value"] = df["New Extracted value"].apply( + lambda x: x.strip() if isinstance(x, str) else "" + ) + actual_value_list = list(df["Actual Value Stored"]) + actual_value_list = [ + answer if str(answer) != "12 months" else "1 year" + for answer in actual_value_list + ] + actual_value_list = [ + answer if str(answer) != "Fifth" else "5" for answer in actual_value_list + ] + actual_value_list = [ + answer if str(answer) != "Seventh" else "7" for answer in actual_value_list + ] + answer_list = list(df["New Extracted value"]) + actual_value_list = [ + s.replace("-", "").replace(" ", "").replace("[", "").replace("]", "").lower() + for s in actual_value_list + ] + answer_list = [ + s.replace("-", "").replace(" ", "").replace("[", "").replace("]", "").lower() + for s in answer_list + ] + result_list = [ + ( + (i in j) or (j in i) + if isinstance(i, str) and isinstance(j, str) and ((i != "") == (j != "")) + else False + ) + for i, j in zip(actual_value_list, answer_list) + ] + df["Result"] = [str(x) for x in result_list] - df = df[['Contract Name', 'SF_DB_COL_NAME', 'Actual Value Stored', 'Raw value', 'New Extracted value', 'Confidence Level', 'Snippet', 'Original Page Number', 'New Page Number', 'Revised Prompt', 'Result']] - - accuracy = round(sum(bool(x) for x in result_list) * 100 / len(list(df['Result'])) if len(df['Result']) > 0 else 0, 2) + df = df[ + [ + "Contract Name", + "SF_DB_COL_NAME", + "Actual Value Stored", + "Raw value", + "New Extracted value", + "Confidence Level", + "Snippet", + "Original Page Number", + "New Page Number", + "Revised Prompt", + "Result", + ] + ] + + accuracy = round( + ( + sum(bool(x) for x in result_list) * 100 / len(list(df["Result"])) + if len(df["Result"]) > 0 + else 0 + ), + 2, + ) logger.info(f"ACCURACY: {accuracy}") history = read_from_db('SELECT * FROM "TRAINING_ATTEMPT_LOGS"') - history.rename(columns={'FIELD_NAME': 'Field Name', 'CONTRACTS_TESTED': '# Contracts Tested', 'USERNAME': 'Username', 'DATE_TIME': 'Date/Time', 'ACCURACY': 'Accuracy', 'ATTEMPT_NUM': 'Attempt #'}, inplace=True) - history = history[['Field Name', '# Contracts Tested', 'Username', 'Date/Time', 'Accuracy', 'Attempt #']] - history.loc[len(history.index)] = [df['SF_DB_COL_NAME'].iloc[0], str(df['Contract Name'].nunique()), 'pipeline', datetime.now().strftime("%Y-%m-%d %H:%M:%S"), accuracy, 0] + history.rename( + columns={ + "FIELD_NAME": "Field Name", + "CONTRACTS_TESTED": "# Contracts Tested", + "USERNAME": "Username", + "DATE_TIME": "Date/Time", + "ACCURACY": "Accuracy", + "ATTEMPT_NUM": "Attempt #", + }, + inplace=True, + ) + history = history[ + [ + "Field Name", + "# Contracts Tested", + "Username", + "Date/Time", + "Accuracy", + "Attempt #", + ] + ] + history.loc[len(history.index)] = [ + df["SF_DB_COL_NAME"].iloc[0], + str(df["Contract Name"].nunique()), + "pipeline", + datetime.now().strftime("%Y-%m-%d %H:%M:%S"), + accuracy, + 0, + ] logger.info(f"Comparison completed. Accuracy: {accuracy}") return df, history + def main(bucket, object_key, field_groups): logger.info(f"Processing contract: {object_key} from bucket: {bucket}") try: data = s3_client.get_object(Bucket=bucket, Key=object_key) - contents = data['Body'].read() + contents = data["Body"].read() context = contents.decode("utf-8") # Fetch batch_id from object tags - batch_id = get_s3_object_tags(bucket, object_key).get('BatchId') + batch_id = get_s3_object_tags(bucket, object_key).get("BatchId") logger.info(f"Batch ID: {batch_id}") - df = pd.DataFrame(columns=['Contract Name', 'Field Name', 'Raw value', 'New Extracted value', 'Confidence Level', 'Snippet', 'New Page Number', 'Revised Prompt']) - field_list, answer_list, snippet_list, page_no_list, contract_list_f = [], [], [], [], [] + df = pd.DataFrame( + columns=[ + "Contract Name", + "Field Name", + "Raw value", + "New Extracted value", + "Confidence Level", + "Snippet", + "New Page Number", + "Revised Prompt", + ] + ) + field_list, answer_list, snippet_list, page_no_list, contract_list_f = ( + [], + [], + [], + [], + [], + ) question = prepare_prompt(*field_groups) response_text = invoke_llm(context, question, llm_selected=LLM_SELECTED) # logger.info(f"Response text: {response_text}, CONTEXT: {context}") - field_l, answer_l, page_no_l, location_l, snippet_l = json_parsing(response_text, context) + field_l, answer_l, page_no_l, location_l, snippet_l = json_parsing( + response_text, context + ) field_list.extend(field_l) answer_list.extend(answer_l) - contract_list_f.extend([object_key]*len(field_l)) + contract_list_f.extend([object_key] * len(field_l)) snippet_list.extend(snippet_l) page_no_list.extend(page_no_l) - - df['Field Name'] = field_list - df['Raw value'] = answer_list - contract_list_f = [contract.rsplit('/', 1)[1] for contract in contract_list_f] - df['Contract Name'] = contract_list_f + + df["Field Name"] = field_list + df["Raw value"] = answer_list + contract_list_f = [contract.rsplit("/", 1)[1] for contract in contract_list_f] + df["Contract Name"] = contract_list_f answer_list = post_processing(answer_list) - df['New Extracted value'] = answer_list + df["New Extracted value"] = answer_list # df['Confidence Level'] = ' ' - df['Snippet'] = snippet_list - df['New Page Number'] = page_no_list - df['New Page Number'] = df['New Page Number'].apply(lambda x: re.search(r'\d+', x).group() if isinstance(x, str) and re.search(r'\d+', x) is not None else " ") + df["Snippet"] = snippet_list + df["New Page Number"] = page_no_list + df["New Page Number"] = df["New Page Number"].apply( + lambda x: ( + re.search(r"\d+", x).group() + if isinstance(x, str) and re.search(r"\d+", x) is not None + else " " + ) + ) # df['Revised Prompt'] = [question] * len(contract_list_f) # Need to add batch_id to the df - df['Batch_id'] = batch_id + df["Batch_id"] = batch_id # Remove column confidence level and revised prompt from df - df = df.drop(columns=['Confidence Level']) - df = df.drop(columns=['Revised Prompt']) - + df = df.drop(columns=["Confidence Level"]) + df = df.drop(columns=["Revised Prompt"]) + # TODO: uncomment accuracy function after testing # df, history = compare_with_actuals(df) csv_buf = StringIO() @@ -529,18 +746,25 @@ def main(bucket, object_key, field_groups): csv_buf.seek(0) timestamp = datetime.now().strftime("%Y%m%d%H%M%S") file_name = f"results_{batch_id}_AC_{timestamp}.csv" - s3_client.put_object(Bucket=bucket, Body=csv_buf.getvalue(), Key=f"final_output/{batch_id}/{file_name}") + s3_client.put_object( + Bucket=bucket, + Body=csv_buf.getvalue(), + Key=f"final_output/{batch_id}/{file_name}", + ) csv_buf = StringIO() df.to_csv(csv_buf, header=True, index=False) csv_buf.seek(0) - - # Saving data to snowflake raw data ingestion bucket & calling stored proc - s3_client.put_object(Bucket=snowflake_ingestion_bucket, Body=csv_buf.getvalue(), Key=f"doczy_pipeline_output/{file_name}") + + # Saving data to snowflake raw data ingestion bucket & calling stored proc + s3_client.put_object( + Bucket=snowflake_ingestion_bucket, + Body=csv_buf.getvalue(), + Key=f"doczy_pipeline_output/{file_name}", + ) cur = snowflake_conn.cursor() load_data_query = f"CALL LOAD_DOCZY_PIPELINE_RAW_OUTPUT('{file_name}')" cur.execute(load_data_query) - # history.tail(1).to_csv(csv_buf, header=True, index=False) # csv_buf.seek(0) # timestamp = datetime.now().strftime("%Y%m%d%H%M%S") @@ -550,6 +774,6 @@ def main(bucket, object_key, field_groups): # logger.info("Results saved to S3 and database.") # save_to_sf('load_training_results', training_results_file_name="results.csv", attempt_logs_file_name="history.csv") except Exception as e: - logger.error(f"Error processing contract: {e} \n TRACEBACK: {traceback.format_exc()}") - - + logger.error( + f"Error processing contract: {e} \n TRACEBACK: {traceback.format_exc()}" + ) diff --git a/textract-pipeline/src/lambda/prompt-orchestrator/integration_testing.ipynb b/textract-pipeline/src/lambda/prompt-orchestrator/integration_testing.ipynb index 8df98fe..a2d3e25 100644 --- a/textract-pipeline/src/lambda/prompt-orchestrator/integration_testing.ipynb +++ b/textract-pipeline/src/lambda/prompt-orchestrator/integration_testing.ipynb @@ -124,17 +124,18 @@ "import os\n", "import pandas as pd\n", "\n", + "\n", "def clean_td(td):\n", " td_clean = []\n", " for d in td:\n", " new_d = {}\n", " for k, v in d.items():\n", - " if 'DATE' in k:\n", + " if \"DATE\" in k:\n", " new_d[k] = v if isinstance(v, list) else [v]\n", - " elif k not in ['page_num', 'Filename']:\n", - " if isinstance(v, str) and ',' in v:\n", - " new_d[k] = [item.strip() for item in v.split(',')]\n", - " elif v == 'N/A':\n", + " elif k not in [\"page_num\", \"Filename\"]:\n", + " if isinstance(v, str) and \",\" in v:\n", + " new_d[k] = [item.strip() for item in v.split(\",\")]\n", + " elif v == \"N/A\":\n", " new_d[k] = []\n", " else:\n", " new_d[k] = [v] if isinstance(v, str) else v\n", @@ -143,25 +144,30 @@ " td_clean.append(new_d)\n", " return td_clean\n", "\n", + "\n", "def get_unique_keys(dicts):\n", " keys = set()\n", " for d in dicts:\n", " keys.update(d.keys())\n", " return keys\n", "\n", + "\n", "def get_unique_date_fields(td_results):\n", " unique_date_fields = {}\n", " for td in td_results:\n", " for key, value in td.items():\n", - " if 'DATE' in key:\n", + " if \"DATE\" in key:\n", " if key not in unique_date_fields:\n", " unique_date_fields[key] = set()\n", - " unique_date_fields[key].update(value if isinstance(value, list) else [value])\n", + " unique_date_fields[key].update(\n", + " value if isinstance(value, list) else [value]\n", + " )\n", " # Convert sets to lists for consistency\n", " for key in unique_date_fields:\n", " unique_date_fields[key] = list(unique_date_fields[key])\n", " return unique_date_fields\n", "\n", + "\n", "def GENERATE_PROMPT(bu_dict, td_dicts, page, field_name=None, values=None):\n", " if field_name and values:\n", " return f\"\"\"### PAGE START ### {page} ### PAGE END\n", @@ -192,23 +198,28 @@ "Write ONLY the number of the dictionary that is most associated with the Service and Methodology of interest. Do not return any additional text beyond the digit itself. \n", "\"\"\"\n", "\n", + "\n", "def merge_results(td_results, bu_results, text_dict):\n", " all_keys = get_unique_keys(td_results) | get_unique_keys(bu_results)\n", " print(\"All Keys:\", all_keys)\n", - " \n", + "\n", " date_fields = get_unique_date_fields(td_results)\n", " print(\"Date Fields:\", date_fields)\n", - " \n", + "\n", " merged_results = []\n", " for bu in bu_results:\n", " merged = {key: bu.get(key, \"\") for key in all_keys}\n", - " td_on_page = [td for td in td_results if td['page_num'] == bu['page_num']]\n", - " \n", + " td_on_page = [td for td in td_results if td[\"page_num\"] == bu[\"page_num\"]]\n", + "\n", " for td in td_on_page:\n", " for key, value in td.items():\n", " if not merged[key]:\n", " merged[key] = value\n", - " elif isinstance(value, list) and value and not isinstance(merged[key], list):\n", + " elif (\n", + " isinstance(value, list)\n", + " and value\n", + " and not isinstance(merged[key], list)\n", + " ):\n", " merged[key] = value\n", " elif isinstance(value, list) and value:\n", " merged[key].extend(value)\n", @@ -218,63 +229,80 @@ " if date_key not in merged or not merged[date_key]:\n", " merged[date_key] = date_values\n", " else:\n", - " merged[date_key].extend([val for val in date_values if val not in merged[date_key]])\n", + " merged[date_key].extend(\n", + " [val for val in date_values if val not in merged[date_key]]\n", + " )\n", "\n", " # Use LLM for disambiguation and inference if necessary\n", " for key in all_keys:\n", " if isinstance(merged[key], list) and len(merged[key]) > 1:\n", - " page_num = bu['page_num']\n", + " page_num = bu[\"page_num\"]\n", " page_text = text_dict.get(page_num, \"\")\n", - " prompt = GENERATE_PROMPT(bu, td_on_page, page_text, field_name=key, values=merged[key])\n", + " prompt = GENERATE_PROMPT(\n", + " bu, td_on_page, page_text, field_name=key, values=merged[key]\n", + " )\n", " response = claude_funcs.invoke_claude_3(prompt, max_tokens=4000)\n", " try:\n", " selected_value = response.strip()\n", " merged[key] = selected_value\n", " except Exception as e:\n", " print(f\"Error processing LLM response for key {key}: {e}\")\n", - " merged[key] = ', '.join(merged[key])\n", - " \n", + " merged[key] = \", \".join(merged[key])\n", + "\n", " merged_results.append(merged)\n", - " \n", + "\n", " return merged_results\n", "\n", + "\n", "def process_file(filename, input_dict):\n", " contract_text = input_dict[filename]\n", "\n", " # Preprocess\n", " contract_text = preprocess.clean_newlines(contract_text)\n", " text_dict = preprocess.split_text(contract_text)\n", - " \n", + "\n", " # Run Top Down\n", - " td_results = prompt_funcs.run_top_down(filename, text_dict) # Returns list of dictionaries for each page\n", + " td_results = prompt_funcs.run_top_down(\n", + " filename, text_dict\n", + " ) # Returns list of dictionaries for each page\n", " print(\"TD Results:\", td_results)\n", - " \n", + "\n", " # Run Bottom Up\n", - " bu_results = prompt_funcs.run_bottom_up(filename, text_dict) # Returns list of dictionaries\n", + " bu_results = prompt_funcs.run_bottom_up(\n", + " filename, text_dict\n", + " ) # Returns list of dictionaries\n", " print(\"BU Results:\", bu_results)\n", - " \n", + "\n", " # Combine\n", " combined_results = merge_results(clean_td(td_results), bu_results, text_dict)\n", " print(\"Combined Results:\", combined_results)\n", - " \n", + "\n", " # Create directories\n", " base_filename = os.path.splitext(filename)[0]\n", - " output_dir = os.path.join('output', base_filename)\n", + " output_dir = os.path.join(\"output\", base_filename)\n", " os.makedirs(output_dir, exist_ok=True)\n", - " \n", + "\n", " # Save results\n", - " pd.DataFrame(td_results).to_csv(os.path.join(output_dir, 'td_results.csv'), index=False)\n", - " pd.DataFrame(bu_results).to_csv(os.path.join(output_dir, 'bu_results.csv'), index=False)\n", - " pd.DataFrame(combined_results).to_csv(os.path.join(output_dir, 'combined_results.csv'), index=False)\n", - " \n", + " pd.DataFrame(td_results).to_csv(\n", + " os.path.join(output_dir, \"td_results.csv\"), index=False\n", + " )\n", + " pd.DataFrame(bu_results).to_csv(\n", + " os.path.join(output_dir, \"bu_results.csv\"), index=False\n", + " )\n", + " pd.DataFrame(combined_results).to_csv(\n", + " os.path.join(output_dir, \"combined_results.csv\"), index=False\n", + " )\n", + "\n", + "\n", "def process_all_files(input_dict):\n", " for filename in input_dict.keys():\n", " process_file(filename, input_dict)\n", " print(f\"Processed {filename}\")\n", "\n", + "\n", "# Example usage\n", - "input_dict = utils.read_input(path='subset')\n", - "process_all_files(input_dict)\n" + "input_dict = utils.read_input(path=\"subset\")\n", + "process_all_files(input_dict)" ] }, { @@ -362,17 +390,18 @@ "import postprocess\n", "import utils\n", "\n", + "\n", "def clean_td(td):\n", " td_clean = []\n", " for d in td:\n", " new_d = {}\n", " for k, v in d.items():\n", - " if 'DATE' in k:\n", + " if \"DATE\" in k:\n", " new_d[k] = v if isinstance(v, list) else [v]\n", - " elif k not in ['page_num', 'Filename']:\n", - " if isinstance(v, str) and ',' in v:\n", - " new_d[k] = [item.strip() for item in v.split(',')]\n", - " elif v == 'N/A':\n", + " elif k not in [\"page_num\", \"Filename\"]:\n", + " if isinstance(v, str) and \",\" in v:\n", + " new_d[k] = [item.strip() for item in v.split(\",\")]\n", + " elif v == \"N/A\":\n", " new_d[k] = []\n", " else:\n", " new_d[k] = [v] if isinstance(v, str) else v\n", @@ -381,24 +410,29 @@ " td_clean.append(new_d)\n", " return td_clean\n", "\n", + "\n", "def get_unique_keys(dicts):\n", " keys = set()\n", " for d in dicts:\n", " keys.update(d.keys())\n", " return keys\n", "\n", + "\n", "def get_unique_date_fields(td_results):\n", " unique_date_fields = {}\n", " for td in td_results:\n", " for key, value in td.items():\n", - " if 'DATE' in key:\n", + " if \"DATE\" in key:\n", " if key not in unique_date_fields:\n", " unique_date_fields[key] = set()\n", - " unique_date_fields[key].update(value if isinstance(value, list) else [value])\n", + " unique_date_fields[key].update(\n", + " value if isinstance(value, list) else [value]\n", + " )\n", " for key in unique_date_fields:\n", " unique_date_fields[key] = list(unique_date_fields[key])\n", " return unique_date_fields\n", "\n", + "\n", "def GENERATE_PROMPT(bu_dict, td_dicts, page, field_name=None, values=None):\n", " if field_name and values:\n", " return f\"\"\"### PAGE START ### {page} ### PAGE END\n", @@ -429,96 +463,115 @@ "Write ONLY the number of the dictionary that is most associated with the Service and Methodology of interest. Do not return any additional text beyond the digit itself. \n", "\"\"\"\n", "\n", + "\n", "def merge_results(td_results, bu_results, text_dict):\n", " all_keys = get_unique_keys(td_results) | get_unique_keys(bu_results)\n", " print(\"All Keys:\", all_keys)\n", - " \n", + "\n", " date_fields = get_unique_date_fields(td_results)\n", " print(\"Date Fields:\", date_fields)\n", - " \n", + "\n", " merged_results = []\n", " for bu in bu_results:\n", " merged = {key: bu.get(key, \"\") for key in all_keys}\n", - " td_on_page = [td for td in td_results if td['page_num'] == bu['page_num']]\n", - " \n", + " td_on_page = [td for td in td_results if td[\"page_num\"] == bu[\"page_num\"]]\n", + "\n", " for td in td_on_page:\n", " for key, value in td.items():\n", " if not merged[key]:\n", " merged[key] = value\n", - " elif isinstance(value, list) and value and not isinstance(merged[key], list):\n", + " elif (\n", + " isinstance(value, list)\n", + " and value\n", + " and not isinstance(merged[key], list)\n", + " ):\n", " merged[key] = value\n", " elif isinstance(value, list) and value:\n", " merged[key].extend(value)\n", "\n", " for date_key, date_values in date_fields.items():\n", - " if date_key not in merged or not merged[date_key]:\n", - " merged[date_key] = date_values\n", - " else:\n", - " merged[date_key].extend([val for val in date_values if val not in merged[date_key]])\n", + " if date_key not in merged or not merged[date_key]:\n", + " merged[date_key] = date_values\n", + " else:\n", + " merged[date_key].extend(\n", + " [val for val in date_values if val not in merged[date_key]]\n", + " )\n", "\n", " for key in all_keys:\n", " if isinstance(merged[key], list) and len(merged[key]) > 1:\n", - " page_num = bu['page_num']\n", + " page_num = bu[\"page_num\"]\n", " page_text = text_dict.get(page_num, \"\")\n", - " prompt = GENERATE_PROMPT(bu, td_on_page, page_text, field_name=key, values=merged[key])\n", + " prompt = GENERATE_PROMPT(\n", + " bu, td_on_page, page_text, field_name=key, values=merged[key]\n", + " )\n", " response = claude_funcs.invoke_claude_3(prompt, max_tokens=4000)\n", " try:\n", " selected_value = response.strip()\n", " merged[key] = selected_value\n", " except Exception as e:\n", " print(f\"Error processing LLM response for key {key}: {e}\")\n", - " merged[key] = ', '.join(merged[key])\n", - " \n", + " merged[key] = \", \".join(merged[key])\n", + "\n", " merged_results.append(merged)\n", - " \n", + "\n", " return merged_results\n", "\n", + "\n", "def process_file(filename, input_dict):\n", " contract_text = input_dict[filename]\n", "\n", " # Preprocess\n", " contract_text = preprocess.clean_newlines(contract_text)\n", " text_dict = preprocess.split_text(contract_text)\n", - " \n", + "\n", " # Run Top Down\n", - " td_results = prompt_funcs.run_top_down(filename, text_dict) # Returns list of dictionaries for each page\n", + " td_results = prompt_funcs.run_top_down(\n", + " filename, text_dict\n", + " ) # Returns list of dictionaries for each page\n", " print(\"TD Results:\", td_results)\n", - " \n", + "\n", " # Run Bottom Up\n", - " bu_results = prompt_funcs.run_bottom_up(filename, text_dict) # Returns list of dictionaries\n", + " bu_results = prompt_funcs.run_bottom_up(\n", + " filename, text_dict\n", + " ) # Returns list of dictionaries\n", " print(\"BU Results:\", bu_results)\n", - " \n", + "\n", " # Combine\n", " combined_results = merge_results(clean_td(td_results), bu_results, text_dict)\n", " print(\"Combined Results:\", combined_results)\n", - " \n", + "\n", " # Convert to DataFrame\n", " combined_df = pd.DataFrame(combined_results)\n", - " \n", + "\n", " # Post-process combined results\n", " # post_processed_combined_df = postprocess.postprocess_results(combined_df)\n", " # print(\"Post-processed Results:\", post_processed_combined_df)\n", - " \n", + "\n", " # Create directories\n", " base_filename = os.path.splitext(filename)[0]\n", - " output_dir = os.path.join('output', base_filename)\n", + " output_dir = os.path.join(\"output\", base_filename)\n", " os.makedirs(output_dir, exist_ok=True)\n", - " \n", + "\n", " # Save results\n", - " pd.DataFrame(td_results).to_csv(os.path.join(output_dir, 'td_results.csv'), index=False)\n", - " pd.DataFrame(bu_results).to_csv(os.path.join(output_dir, 'bu_results.csv'), index=False)\n", - " combined_df.to_csv(os.path.join(output_dir, 'combined_results.csv'), index=False)\n", + " pd.DataFrame(td_results).to_csv(\n", + " os.path.join(output_dir, \"td_results.csv\"), index=False\n", + " )\n", + " pd.DataFrame(bu_results).to_csv(\n", + " os.path.join(output_dir, \"bu_results.csv\"), index=False\n", + " )\n", + " combined_df.to_csv(os.path.join(output_dir, \"combined_results.csv\"), index=False)\n", " # post_processed_combined_df.to_csv(os.path.join(output_dir, 'combined_results.csv'), index=False)\n", "\n", + "\n", "def process_all_files(input_dict):\n", " for filename in input_dict.keys():\n", " process_file(filename, input_dict)\n", " print(f\"Processed {filename}\")\n", "\n", + "\n", "# Example usage\n", - "input_dict = utils.read_input(path='../data/')\n", - "process_all_files(input_dict)\n", - "\n" + "input_dict = utils.read_input(path=\"../data/\")\n", + "process_all_files(input_dict)" ] }, { @@ -528,7 +581,7 @@ "outputs": [], "source": [ "# Usage example:\n", - "#save_combined_to_csv(combined, 'combined_output.csv')" + "# save_combined_to_csv(combined, 'combined_output.csv')" ] }, { @@ -537,7 +590,7 @@ "metadata": {}, "outputs": [], "source": [ - "utils.consolidate_individual(input_folder='../results/', output_folder='../output/')" + "utils.consolidate_individual(input_folder=\"../results/\", output_folder=\"../output/\")" ] } ], diff --git a/textract-pipeline/src/lambda/prompt-orchestrator/main.py b/textract-pipeline/src/lambda/prompt-orchestrator/main.py index 4dc2748..8e580c3 100644 --- a/textract-pipeline/src/lambda/prompt-orchestrator/main.py +++ b/textract-pipeline/src/lambda/prompt-orchestrator/main.py @@ -1,4 +1,3 @@ - """ main.py Doczy AI - Pricing Before Carveouts: Main Execution Module @@ -20,31 +19,44 @@ import config import claude_funcs import file_processing + def main(): if config.TEST: - print(claude_funcs.invoke_claude_3("Write 'test', nothing more.", max_tokens = 10)) + print( + claude_funcs.invoke_claude_3("Write 'test', nothing more.", max_tokens=10) + ) else: # Read contract txt input_dict = utils.read_input() # already_processed = [s.split('.txt_results')[0]+'.txt' for s in os.listdir('results/')] # input_dict = {key : input_dict[key] for key in input_dict.keys() if key not in already_processed} - + all_results = [] for item in input_dict.items(): file_result = file_processing.process_file(item) all_results.append(file_result) - + # Write Consolidated Output if config.WRITE_OUTPUT: all_results_df = pd.concat(all_results).reset_index(drop=True) - all_results_df.drop('page_num', axis=1, inplace=True) - + all_results_df.drop("page_num", axis=1, inplace=True) + # Reorder columns - final_df = all_results_df[[col for col in config.VALID_COLUMNS if col in all_results_df.columns] + [col for col in all_results_df.columns if col not in config.VALID_COLUMNS]] + final_df = all_results_df[ + [col for col in config.VALID_COLUMNS if col in all_results_df.columns] + + [ + col + for col in all_results_df.columns + if col not in config.VALID_COLUMNS + ] + ] print(final_df.columns) - final_df.to_csv(os.path.join(config.CONSOLIDATED_OUTPUT_DIRECTORY, config.OUTPUT_CSV_PATH)) + final_df.to_csv( + os.path.join( + config.CONSOLIDATED_OUTPUT_DIRECTORY, config.OUTPUT_CSV_PATH + ) + ) if __name__ == "__main__": main() - diff --git a/textract-pipeline/src/lambda/prompt-orchestrator/merge_funcs.py b/textract-pipeline/src/lambda/prompt-orchestrator/merge_funcs.py index 49c7b5c..bb18a6a 100644 --- a/textract-pipeline/src/lambda/prompt-orchestrator/merge_funcs.py +++ b/textract-pipeline/src/lambda/prompt-orchestrator/merge_funcs.py @@ -1,4 +1,3 @@ - import prompts import claude_funcs @@ -8,44 +7,53 @@ def prompt_select_multiple(key, value, bu_dict, page): answer = claude_funcs.invoke_claude_3(prompt, max_tokens=4000) return answer + def merge_results(td_results, bu_results, text_dict): def find_td_dict(page_num): - """ Helper function to find dictionary in td_results with matching page_num """ + """Helper function to find dictionary in td_results with matching page_num""" for td in td_results: - if td['page_num'] == page_num: + if td["page_num"] == page_num: return td return None - + def get_value(td_dict, key, page_num, text_dict): - """ Helper function to handle value selection based on the rules specified. """ + """Helper function to handle value selection based on the rules specified.""" if td_dict is None: # If there's no matching page, use recursion to look one page earlier if int(page_num) > 1: - return get_value(find_td_dict(int(page_num) - 1), key, int(page_num) - 1, text_dict), str(int(page_num) - 1) + return get_value( + find_td_dict(int(page_num) - 1), key, int(page_num) - 1, text_dict + ), str(int(page_num) - 1) else: return "N/A", "N/A" # Base case if no previous pages exist else: value = td_dict.get(key, "N/A") if value == "N/A" or value == "": # If value is 'N/A' or empty, use recursion to look one page earlier - return get_value(find_td_dict(int(page_num) - 1), key, int(page_num) - 1, text_dict), str(int(page_num) - 1) - elif ',' in value and 'DATE' not in key: + return get_value( + find_td_dict(int(page_num) - 1), key, int(page_num) - 1, text_dict + ), str(int(page_num) - 1) + elif "," in value and "DATE" not in key: # Randomly select one of the comma-separated values - return prompt_select_multiple(key, value, bu_dict, text_dict[page_num]), str(int(page_num)) + return prompt_select_multiple( + key, value, bu_dict, text_dict[page_num] + ), str(int(page_num)) else: return value, page_num - + merged_results = [] for bu_dict in bu_results: - page_num = bu_dict['page_num'] + page_num = bu_dict["page_num"] td_dict = find_td_dict(page_num) - + new_dict = bu_dict.copy() # Start with the bu_dict's data # Add or overwrite keys from td_dict if td_dict: for key, value in td_dict.items(): - if key not in ['Filename', 'page_num']: - new_dict[key], new_dict[f"{key}_PG"] = get_value(td_dict, key, page_num, text_dict) - + if key not in ["Filename", "page_num"]: + new_dict[key], new_dict[f"{key}_PG"] = get_value( + td_dict, key, page_num, text_dict + ) + merged_results.append(new_dict) - return merged_results \ No newline at end of file + return merged_results diff --git a/textract-pipeline/src/lambda/prompt-orchestrator/postprocess.py b/textract-pipeline/src/lambda/prompt-orchestrator/postprocess.py index efa3039..d7559a4 100644 --- a/textract-pipeline/src/lambda/prompt-orchestrator/postprocess.py +++ b/textract-pipeline/src/lambda/prompt-orchestrator/postprocess.py @@ -4,40 +4,52 @@ import re import postprocess_funcs import config - + def postprocess_results(combined_df): # Define valid values dictionary valid_values_dict = { - 'CONTRACT_LOB': config.VALID_LOBS, - 'CONTRACT_PROGRAM': config.VALID_PROGRAMS, - 'CONTRACT_NETWORK': config.VALID_NETWORKS + "CONTRACT_LOB": config.VALID_LOBS, + "CONTRACT_PROGRAM": config.VALID_PROGRAMS, + "CONTRACT_NETWORK": config.VALID_NETWORKS, } - + # Sanitize the combined data try: df = postprocess_funcs.sanitize_combined(combined_df) except Exception as e: print(f"Postprocessing Error - santize_combined : {e}") - + # Clean specific columns for exact matches - for field_list in [['CONTRACT_LOB', 'Corrected_LOB', config.VALID_LOBS], ['CONTRACT_PROGRAM', 'Corrected_PROGRAM', config.VALID_PROGRAMS], ['CONTRACT_NETWORK', 'Corrected_NETWORK', config.VALID_NETWORKS]]: + for field_list in [ + ["CONTRACT_LOB", "Corrected_LOB", config.VALID_LOBS], + ["CONTRACT_PROGRAM", "Corrected_PROGRAM", config.VALID_PROGRAMS], + ["CONTRACT_NETWORK", "Corrected_NETWORK", config.VALID_NETWORKS], + ]: try: - postprocess_funcs.clean_columns_combined(df, field_list[0], field_list[2], field_list[1]) + postprocess_funcs.clean_columns_combined( + df, field_list[0], field_list[2], field_list[1] + ) except Exception as e: print(f"Postprocessing Error - clean_columns_combined : {e}") - + try: - postprocess_funcs.clean_columns_combined_fuzzy(df, field_list[0], field_list[2], config.FUZZY_MATCH_THRESHOLD) + postprocess_funcs.clean_columns_combined_fuzzy( + df, field_list[0], field_list[2], config.FUZZY_MATCH_THRESHOLD + ) except Exception as e: print(f"Postprocessing Error - clean_columns_combined_fuzzy : {e}") - + # Correct misplaced values across columns try: - df = postprocess_funcs.correct_misplaced_values(df, ['CONTRACT_LOB', 'CONTRACT_PROGRAM', 'CONTRACT_NETWORK'], valid_values_dict) + df = postprocess_funcs.correct_misplaced_values( + df, + ["CONTRACT_LOB", "CONTRACT_PROGRAM", "CONTRACT_NETWORK"], + valid_values_dict, + ) except Exception as e: print(f"Postprocessing Error - correct_misplaced_values : {e}") - + # Move percentages and large numbers to correct columns try: df = postprocess_funcs.move_percentage_to_rate(df) @@ -57,36 +69,17 @@ def postprocess_results(combined_df): print(f"Postprocessing Error - adjust_reimbursement_rate : {e}") # Clean dates so that Termination and Effective aren't the same - try: + try: df = postprocess_funcs.clean_dates(df) except Exception as e: print(f"Postprocessing Error - clean_dates : {e}") return df - - - - - - - - - - - - - - - - - - - - - - + # Order columns - cols = [col for col in config.VALID_COLUMNS if col in combined_df.columns] + [col for col in combined_df.columns if col not in config.VALID_COLUMNS] - #combined_df = combined_df.loc[:, cols] + cols = [col for col in config.VALID_COLUMNS if col in combined_df.columns] + [ + col for col in combined_df.columns if col not in config.VALID_COLUMNS + ] + # combined_df = combined_df.loc[:, cols] combined_df = combined_df[cols] return combined_df diff --git a/textract-pipeline/src/lambda/prompt-orchestrator/postprocess_funcs.py b/textract-pipeline/src/lambda/prompt-orchestrator/postprocess_funcs.py index 1fec04e..e803cf2 100644 --- a/textract-pipeline/src/lambda/prompt-orchestrator/postprocess_funcs.py +++ b/textract-pipeline/src/lambda/prompt-orchestrator/postprocess_funcs.py @@ -1,10 +1,10 @@ - import pandas as pd import re import difflib import config + def sanitize_value(value): """ Sanitizes a given value by handling nulls, lists, and strings to ensure consistency in formatting. @@ -22,15 +22,16 @@ def sanitize_value(value): """ try: if isinstance(value, list): - return ', '.join(str(v) for v in value) + return ", ".join(str(v) for v in value) elif isinstance(value, str): - value = value.strip('[]') - return ', '.join([item.strip(" '") for item in value.split(',')]) + value = value.strip("[]") + return ", ".join([item.strip(" '") for item in value.split(",")]) elif pd.isna(value): return "N/A" except: return value + def exact_match(val, valid_values): """ Compares a given string value against a list of valid values to determine if there is an exact match, case-insensitively. @@ -52,6 +53,7 @@ def exact_match(val, valid_values): return valid_val return None + def clean_columns_combined(df, column_name, valid_values, new_column_name): """ Cleans and standardizes entries in a specified column of a DataFrame based on a list of valid values, creating a new column with standardized values. @@ -75,13 +77,13 @@ def clean_columns_combined(df, column_name, valid_values, new_column_name): def update_column(entry): if pd.notna(entry): - terms = entry.split(',') + terms = entry.split(",") for term in terms: match = exact_match(term, valid_values) if match: return match return None - + df[new_column_name] = df[column_name].apply(update_column) cleaned_values = df[new_column_name].unique() return original_values, cleaned_values, changes @@ -92,7 +94,7 @@ def get_closest_match(val, valid_values, similarity_threshold=0.7): Finds the closest match for a given string from a list of valid values, using a similarity threshold. This function cleans the input value and compares it to each cleaned value in the valid values list using a fuzzy - matching technique. It returns the closest match that meets or exceeds the specified similarity threshold. If no + matching technique. It returns the closest match that meets or exceeds the specified similarity threshold. If no matches meet the threshold, the function returns None. Parameters: @@ -106,7 +108,9 @@ def get_closest_match(val, valid_values, similarity_threshold=0.7): if pd.isna(val): return None val = val.strip().upper() - matches = difflib.get_close_matches(val, [v.upper() for v in valid_values], n=1, cutoff=similarity_threshold) + matches = difflib.get_close_matches( + val, [v.upper() for v in valid_values], n=1, cutoff=similarity_threshold + ) return matches[0] if matches else None @@ -131,17 +135,19 @@ def clean_columns_combined_fuzzy(df, column_name, valid_values, threshold): changes = {} df[column_name] = df[column_name].apply(sanitize_value) original_values = df[column_name].unique() - + def log_and_clean(entry): if pd.notna(entry): - words = entry.split(',') + words = entry.split(",") cleaned_words = [] for word in words: - cleaned_word = get_closest_match(word.strip(), valid_values, similarity_threshold=threshold) + cleaned_word = get_closest_match( + word.strip(), valid_values, similarity_threshold=threshold + ) if cleaned_word and word.strip().upper() != cleaned_word: changes[word.strip()] = cleaned_word cleaned_words.append(cleaned_word if cleaned_word else word.strip()) - return ', '.join(cleaned_words) + return ", ".join(cleaned_words) return None df[column_name] = df[column_name].apply(log_and_clean) @@ -149,7 +155,6 @@ def clean_columns_combined_fuzzy(df, column_name, valid_values, threshold): return original_values, cleaned_values, changes - def correct_misplaced_values(df, columns, valid_values_dict): """ Corrects misplaced values within specified columns of a DataFrame based on a dictionary of valid values for each column. @@ -170,20 +175,32 @@ def correct_misplaced_values(df, columns, valid_values_dict): for index, row in df.iterrows(): for col in columns: if pd.notna(row[col]): - terms = row[col].split(',') + terms = row[col].split(",") for term in terms: term = term.strip() for target_col, valid_values in valid_values_dict.items(): if target_col != col: match = exact_match(term, valid_values) if match: - if pd.isna(row[target_col]) or not row[target_col].strip(): + if ( + pd.isna(row[target_col]) + or not row[target_col].strip() + ): df.at[index, target_col] = match df.at[index, col] = None else: current_value = row[target_col].strip() - if get_closest_match(match, [current_value], config.FUZZY_MATCH_THRESHOLD) is None: - df.at[index, 'Corrected_' + target_col] = f"Found {term} in {col} cell" + if ( + get_closest_match( + match, + [current_value], + config.FUZZY_MATCH_THRESHOLD, + ) + is None + ): + df.at[index, "Corrected_" + target_col] = ( + f"Found {term} in {col} cell" + ) df.at[index, col] = None return df @@ -204,13 +221,28 @@ def filter_service_column(d): list: A new list of dictionaries with items containing specified keywords in the 'SERVICE' key removed. """ keywords = [ - 'LIABILITY', 'RISK', 'LOBBYING', 'DAMAGES', 'CONFIDENTIALITY', 'ARBITRATION', 'FALSE CLAIMS ACT', 'UNSPECIFIED', - 'AUDIT', 'INTEREST', 'N/A', 'BUSINESS', 'COMPLIANCE', 'Medical Assistance Program', 'MATERIAL SUBCONTRACT', 'GIFTS', 'GRATUITIES', + "LIABILITY", + "RISK", + "LOBBYING", + "DAMAGES", + "CONFIDENTIALITY", + "ARBITRATION", + "FALSE CLAIMS ACT", + "UNSPECIFIED", + "AUDIT", + "INTEREST", + "N/A", + "BUSINESS", + "COMPLIANCE", + "Medical Assistance Program", + "MATERIAL SUBCONTRACT", + "GIFTS", + "GRATUITIES", ] - pattern = '|'.join(keywords) + pattern = "|".join(keywords) regex = re.compile(pattern, re.IGNORECASE) - filtered_list = [item for item in d if not regex.search(item.get('SERVICE', ''))] + filtered_list = [item for item in d if not regex.search(item.get("SERVICE", ""))] return filtered_list @@ -229,13 +261,28 @@ def move_percentage_to_rate(df): Returns: pandas.DataFrame: The DataFrame with percentage values moved from the 'REIMBURSEMENT_FLAT_FEE' to the 'REIMBURSEMENT_RATE' column. """ + def move_percentage(value): - if isinstance(value, str) and '%' in value: + if isinstance(value, str) and "%" in value: return True return False - - df['REIMBURSEMENT_RATE'] = df.apply(lambda row: row['REIMBURSEMENT_FLAT_FEE'] if move_percentage(row['REIMBURSEMENT_FLAT_FEE']) else row['REIMBURSEMENT_RATE'], axis=1) - df['REIMBURSEMENT_FLAT_FEE'] = df.apply(lambda row: None if move_percentage(row['REIMBURSEMENT_FLAT_FEE']) else row['REIMBURSEMENT_FLAT_FEE'], axis=1) + + df["REIMBURSEMENT_RATE"] = df.apply( + lambda row: ( + row["REIMBURSEMENT_FLAT_FEE"] + if move_percentage(row["REIMBURSEMENT_FLAT_FEE"]) + else row["REIMBURSEMENT_RATE"] + ), + axis=1, + ) + df["REIMBURSEMENT_FLAT_FEE"] = df.apply( + lambda row: ( + None + if move_percentage(row["REIMBURSEMENT_FLAT_FEE"]) + else row["REIMBURSEMENT_FLAT_FEE"] + ), + axis=1, + ) return df @@ -255,12 +302,15 @@ def set_rate_to_zero_if_not_covered(df): Returns: pandas.DataFrame: The updated DataFrame with adjusted 'REIMURSEMENT_RATE' values where applicable. """ + def check_and_set_rate(row): - if pd.notna(row['FULL_METHODOLOGY']) and re.search(r'not covered', row['FULL_METHODOLOGY'], re.IGNORECASE): + if pd.notna(row["FULL_METHODOLOGY"]) and re.search( + r"not covered", row["FULL_METHODOLOGY"], re.IGNORECASE + ): return 0 - return row['REIMBURSEMENT_RATE'] - - df['REIMBURSEMENT_RATE'] = df.apply(check_and_set_rate, axis=1) + return row["REIMBURSEMENT_RATE"] + + df["REIMBURSEMENT_RATE"] = df.apply(check_and_set_rate, axis=1) return df @@ -279,13 +329,18 @@ def adjust_reimbursement_rate(df): Returns: pandas.DataFrame: The updated DataFrame with recalibrated 'REIMBURSEMENT_RATE' values based on the specified methodology condition. """ + def adjust_rate(row): - if pd.notna(row['FULL_METHODOLOGY']) and re.search(r'case insensitive', row['FULL_METHODOLOGY'], re.IGNORECASE): - if pd.notna(row['REIMBURSEMENT_RATE']) and isinstance(row['REIMBURSEMENT_RATE'], (int, float)): - return 100 - row['REIMBURSEMENT_RATE'] - return row['REIMBURSEMENT_RATE'] - - df['REIMBURSEMENT_RATE'] = df.apply(adjust_rate, axis=1) + if pd.notna(row["FULL_METHODOLOGY"]) and re.search( + r"case insensitive", row["FULL_METHODOLOGY"], re.IGNORECASE + ): + if pd.notna(row["REIMBURSEMENT_RATE"]) and isinstance( + row["REIMBURSEMENT_RATE"], (int, float) + ): + return 100 - row["REIMBURSEMENT_RATE"] + return row["REIMBURSEMENT_RATE"] + + df["REIMBURSEMENT_RATE"] = df.apply(adjust_rate, axis=1) return df @@ -310,12 +365,12 @@ def clean_td(td): for d in td: new_d = {} for k, v in d.items(): - if 'DATE' in k: + if "DATE" in k: new_d[k] = v if isinstance(v, list) else [v] - elif k not in ['page_num', 'Filename']: - if isinstance(v, str) and ',' in v: - new_d[k] = [item.strip() for item in v.split(',')] - elif v == 'N/A': + elif k not in ["page_num", "Filename"]: + if isinstance(v, str) and "," in v: + new_d[k] = [item.strip() for item in v.split(",")] + elif v == "N/A": new_d[k] = [] else: new_d[k] = [v] if isinstance(v, str) else v @@ -325,7 +380,6 @@ def clean_td(td): return td_clean - def sanitize_combined(df): """ Sanitizes all columns in a DataFrame by applying a predefined sanitization function to each value. @@ -345,6 +399,7 @@ def sanitize_combined(df): df[column] = df[column].apply(sanitize_value) return df + def extract_codes_CPT(value): """ Extracts and formats CPT codes from a given input value. @@ -364,16 +419,17 @@ def extract_codes_CPT(value): """ if pd.isna(value): return value - - value = re.sub(r'\bthrough\b', '-', str(value), flags=re.IGNORECASE) - - pattern = r'\b([A-Z]\d{4}|\d{5})(-[A-Z]{2})?\b' + + value = re.sub(r"\bthrough\b", "-", str(value), flags=re.IGNORECASE) + + pattern = r"\b([A-Z]\d{4}|\d{5})(-[A-Z]{2})?\b" matches = re.findall(pattern, value) if matches: - return ', '.join([''.join(match) for match in matches]) + return ", ".join(["".join(match) for match in matches]) else: return "" + def extract_codes_Diagnosis(value): """ Extracts and formats diagnosis codes from a given input value, typically adhering to ICD (International Classification of Diseases) formats. @@ -392,14 +448,15 @@ def extract_codes_Diagnosis(value): """ if pd.isna(value): return value - - pattern = r'\b[A-Z][0-9][A-Z0-9]{1,4}(\.[A-Z0-9]{1,4})?\b' + + pattern = r"\b[A-Z][0-9][A-Z0-9]{1,4}(\.[A-Z0-9]{1,4})?\b" matches = re.findall(pattern, value) if matches: - return ', '.join(matches) + return ", ".join(matches) else: return "" + def extract_codes_Revenue(value): """ Extracts revenue codes from a given input value. Revenue codes are typically numerical codes of three to four digits. @@ -417,14 +474,15 @@ def extract_codes_Revenue(value): """ if pd.isna(value): return value - - pattern = r'\b\d{3,4}\b' + + pattern = r"\b\d{3,4}\b" matches = re.findall(pattern, value) if matches: - return ', '.join(matches) + return ", ".join(matches) else: return "" + def clean_code_column(df, column_name, extract_function): """ Applies a specified function to clean and extract codes from a specific column in a DataFrame. @@ -443,19 +501,23 @@ def clean_code_column(df, column_name, extract_function): pandas.DataFrame: The DataFrame with the specified column updated with cleaned and processed codes. """ df[column_name] = df[column_name].apply(extract_function) - return df + return df def clean_dates(df): for idx, row in df.iterrows(): - if row['LOB_PRICING_TERMS_EFFECTIVE_DATE'] == row['LOB_PRICING_TERMS_TERMINATION_DATE']: - df.loc[idx, 'LOB_PRICING_TERMS_TERMINATION_DATE'] = 'N/A' + if ( + row["LOB_PRICING_TERMS_EFFECTIVE_DATE"] + == row["LOB_PRICING_TERMS_TERMINATION_DATE"] + ): + df.loc[idx, "LOB_PRICING_TERMS_TERMINATION_DATE"] = "N/A" return df + def add_PG_columns(results_dicts): for key in list(results_dicts[0].keys()): - if key != 'page_num': + if key != "page_num": for d in results_dicts: - d[f"{key}_PG"] = d['page_num'] + d[f"{key}_PG"] = d["page_num"] return results_dicts diff --git a/textract-pipeline/src/lambda/prompt-orchestrator/preprocess.py b/textract-pipeline/src/lambda/prompt-orchestrator/preprocess.py index d2bb5e1..f4cab0c 100644 --- a/textract-pipeline/src/lambda/prompt-orchestrator/preprocess.py +++ b/textract-pipeline/src/lambda/prompt-orchestrator/preprocess.py @@ -1,7 +1,7 @@ - from itertools import groupby, count import re + def clean_newlines(contract_text): """ Cleans up isolated newlines in a contract text, converting them into spaces to ensure text continuity. @@ -17,7 +17,7 @@ def clean_newlines(contract_text): Returns: str: The cleaned contract text with isolated newlines replaced by spaces. """ - cleaned_text = re.sub(r'(?>>{word}<<<' + if "%" in word or "$" in word: + word = f">>>{word}<<<" highlighted_words.append(word) - text_dict[page] = ' '.join(highlighted_words) + text_dict[page] = " ".join(highlighted_words) return text_dict @@ -69,8 +69,8 @@ def chunk_text(text_dict): """ Creates text chunks from a dictionary of page texts, focusing on pages with special characters (percentages and dollar amounts). - This function identifies pages that contain '%' or '$' signs and includes those pages along with their immediate neighbors - (previous and next pages) to form chunks. The chunks are then grouped and concatenated into single text blocks for easier + This function identifies pages that contain '%' or '$' signs and includes those pages along with their immediate neighbors + (previous and next pages) to form chunks. The chunks are then grouped and concatenated into single text blocks for easier processing. Each chunk is stored in a new dictionary where the keys represent the range of pages included in the chunk. Parameters: @@ -79,19 +79,28 @@ def chunk_text(text_dict): Returns: dict: A new dictionary where each key is a string representing the range of pages in a chunk, and each value is the concatenated text of those pages. """ - special_pages = {int(page): text for page, text in text_dict.items() if (('%' in text) or ('$' in text) or ('percent' in text))} + special_pages = { + int(page): text + for page, text in text_dict.items() + if (("%" in text) or ("$" in text) or ("percent" in text)) + } page_numbers = sorted(special_pages.keys()) chunk_page_numbers = [] for page in page_numbers: - chunk_page_numbers.append(page-1) + chunk_page_numbers.append(page - 1) chunk_page_numbers.append(page) - chunk_page_numbers.append(page+1) + chunk_page_numbers.append(page + 1) chunk_page_numbers = list(set(chunk_page_numbers)) - chunk_page_numbers = [page for page in chunk_page_numbers if str(page) in text_dict.keys()] - chunks = [list(group) for key, group in groupby(chunk_page_numbers, lambda x, c=count(): x - next(c))] + chunk_page_numbers = [ + page for page in chunk_page_numbers if str(page) in text_dict.keys() + ] + chunks = [ + list(group) + for key, group in groupby(chunk_page_numbers, lambda x, c=count(): x - next(c)) + ] chunk_dict = {} for item in chunks: - dict_key = f'{min(item)}-{max(item)}' + dict_key = f"{min(item)}-{max(item)}" text = "".join([text_dict[str(page)] for page in item]) chunk_dict[dict_key] = text return chunk_dict @@ -111,7 +120,7 @@ def clean_billed_charges(contract_text): Returns: str: The cleaned text with appropriate replacements made for 'billed charges'. """ - + substrings = [ "Physician's Billed Charges", "Provider's Billed Charges", @@ -136,17 +145,21 @@ def clean_billed_charges(contract_text): indices.append(index) index = lower_contract_text.find(lower_s, index + 1) return indices - + for s in substrings: indices = find_substring_indices(contract_text, s) index_adder = 0 if indices: for i in indices: - index = i+index_adder + index = i + index_adder end_index = index + max_substring_length match_part = contract_text[index:end_index] - previous = contract_text[max(0, index-30):index] - if '%' not in previous: - contract_text = contract_text[0:index] + f" 100% of {match_part}" + contract_text[end_index:] + previous = contract_text[max(0, index - 30) : index] + if "%" not in previous: + contract_text = ( + contract_text[0:index] + + f" 100% of {match_part}" + + contract_text[end_index:] + ) index_adder += 9 - return contract_text.replace(' ', ' ') + return contract_text.replace(" ", " ") diff --git a/textract-pipeline/src/lambda/prompt-orchestrator/prompts.py b/textract-pipeline/src/lambda/prompt-orchestrator/prompts.py index f174602..619f1ff 100644 --- a/textract-pipeline/src/lambda/prompt-orchestrator/prompts.py +++ b/textract-pipeline/src/lambda/prompt-orchestrator/prompts.py @@ -1,5 +1,3 @@ - - def GENERATE_PROMPT(bu_dict, td_dicts, page, field_name=None, values=None): if field_name and values: return f"""### PAGE START ### {page} ### PAGE END @@ -103,6 +101,7 @@ If both LESSER_OF_LANGUAGE_IND and GREATER_OF_LANGUAGE_IND are 'N', return a lis Only return the list of dictionaries, with no other commentary or explanation. Ensure you abide by proper JSON formatting. """ + def BOTTOM_UP_EXCEPT_ESC(d, page): return f"""### PAGE START ### {page} ### PAGE END @@ -119,6 +118,7 @@ REIMBURSEMENT_DESCRIBE_EXCEPTION: If the listed reimbursement contains some sort Only return the dictionary, with no other commentary or explanation. Ensure you abide by proper JSON formatting. """ + def BOTTOM_UP_FS(d): return f"""Analyze the reimbursement terms listed below: Service: {d['SERVICE']} @@ -133,6 +133,7 @@ REIMBURSEMENT_FEE_SCHEDULE_VERSION : If the Methodology is based on a Fee Schedu Only return the dictionary, with no other commentary or explanation. Ensure you abide by proper JSON formatting. """ + def BOTTOM_UP_METHODOLOGY(d): return f"""Analyze the full methodology text listed below: {d['FULL_METHODOLOGY']} @@ -246,40 +247,212 @@ ONLY return the answer, with no other commentary or explanation. def get_carveout_list(): - return ['Emergency Department/Emergency Room', 'Emergency Department', 'Emergency Room', 'Observation', 'Surgery', 'General Surgery', -'Ambulatory Surgery', 'Intensive Care', 'Trauma', 'HIV', 'Human Immunodeficiency Virus', 'Major Joint Replacement', 'Transplant', 'OBGYN', -'Obstetrician/Gynecologist', 'Obstetrician', 'Gynecologist', 'Opthalmology & Vision', 'Opthalmology', 'Vision', 'Never Events', -'Medically Unnecessary', 'Medically Unnecessary Procedures', 'Not Medically Necessary', 'Not Medically Necessary Procedures', -'Experimental', 'Investigational', 'Unlisted Codes', 'Durable Medical Equipment', 'DME', 'Prosthetics & Orthotics', 'Prosthetics', -'Orthotics', 'Implants', 'Hearing Aids', 'Hearing', 'Anesthesia', 'Anesthesiology', 'Medical Pharmacy', 'Physician Administered Drugs', -'Global', 'Bundled or Unbundled Codes', 'Bundled Codes', 'Unbundled Codes', 'Multiple Procedure Reductions', 'Second Surgery', 'Subsequent Surgeries', -'Non-Behavioral Health Mid-Level Professionals', 'Physician Assistant', 'PA', 'Nurse Practicioner', 'NP', 'Non-Physician', -'Non-Physician Health Professionals', 'Technical Component', 'Professional Component', 'Laboratory', 'Pathology', 'Lab', 'Path', 'Lab/Path', -'Laboratory/Pathology', 'Radiology', 'Imaging', 'Radiology/Imaging', 'Mammography', 'Diagnostic', 'Pre-Admission Procedures', -'Post-Discharge Procedures', 'Readmission', 'Status Indicators', 'Stop Loss', 'Emergency Medical', 'EMS', 'NICU', 'Neonatal Intensive Care Unit', -'Vaccine for Children', 'VFC', 'Surgical Assistant', 'Physicians/Clinical Psychologists', 'Doctor of Nursing Practice', 'Osteopathic Medicine', -'Clinical Psychology', 'Audiologist', 'Chiropractors', 'Registered Dietician', 'AUD', 'DC', 'RD', 'Board Certified Behavioral Analysis', 'Behavioral Analysis', -'BCBA', 'Independent Licensures', 'Licensed Professional Counselor', 'Marriage and Family Therapist', 'Substance Abuse Counselor', 'Clinical Social Worker', -'Behavioral Health Outpatient Clinic', 'Physical Therapist', 'Occupational Therapist', 'Speech Therapist', 'PT', 'OT', 'ST', 'Transportation', -'Primary Care', 'Primary Care Behavioral Health', 'Behavioral Physician', 'Clinical Psychologist', 'Mid-Level Practicioner', 'Dentist', 'Dental', -'Supplies and Devices', 'Supplies', 'Devices', 'Immunizations', 'Obstetrical Epidural', 'Pediatric Subspecialties', 'Orthopedic Surgery', 'All Other Specialists', -'Specialty Care Physician', 'Pediatric Primary Care', 'Nurse Anesthetist', 'Early Periodic Screening, Diagnostic, and Treatment', 'EPSDT', -'Ancillary', 'Oncology', 'Cancer', 'Inpatient Physical Rehabilitation', 'Inpatient Rehabilitation', 'Outpatient Rehabilitation', 'Rehabilitation', 'Non-Behavioral Health Rehabilitation', -'Extracorporeal Shock Wave Lithotripsy', 'Cardiac', 'Special Care Unit', 'Skilled Nursing', 'Infusion', 'Specialty Care', -'Specialty Care Physician', 'Organ Acquisition', 'Blood Products', 'Blood Products Outpatient', 'Blood Products Inpatient', 'High Cost Drugs', 'Sleep Studies', -'NICU', 'Newborn Intensive Care Unit', 'Extracorporeal Membrane Oxygenation', 'Burns', 'Kyphoplasty', 'Cryosurgical Ablation of the Prostate', -'Transurethral Thermal Ablation', 'TUMT', 'Transurethral Needle Ablation', 'TUNA', 'Hyperbaric Treatment', 'Clinic Visit', 'Boarder Baby', -'Pediatric Intensive Care Unit', 'PICU', 'Psychiatric', 'Mental Health', 'Behavioral Health', 'Substance Abuse', -'Behavioral Health and Substance Abuse', 'Sub-Acute Facility Care', 'Unrouped Inpatient', 'All Other Acute', -'Neurology', 'Neurology Subspecialties', 'Automatic Implantable Cardioverter Defibrillator', 'Percutaneous Transluminal Coronary Angioplasty', -'Non-Coronary Angioplasty', 'Cardiac Catheters', 'Cardiovascular Surgery', 'Cardiac Surgery', 'Cesarean Birth', 'Cesarean Section', 'C-Section', -'Gamma-Knife Radio-Surgery Outpatient', 'DaVinci Robotic Assisted Surgery', 'Outpatient Electrophysiology with Ablation', 'Outpatient Electrophysiology', -'Magnetic Resonance Image', 'MRI', 'Computed Tomography Scan', 'CT Scan', 'Radiation Therapy', 'Dialysis', 'Gastric Bypass', 'Lap Band', -'Obesity', 'Laparoscopic Cholecystectomy', 'Lap Chole', 'Laparoscopic Hysterectomy', 'Laparoscopic Prostatectomy', 'Laparoscopic Hysteroscopy', -'Treatment Room', 'Wound Care', 'Cardiac Computed Tomograpy & Angiography', 'Positron Emission Tomography Scan', 'PET Scan', 'Hematology'] - - - + return [ + "Emergency Department/Emergency Room", + "Emergency Department", + "Emergency Room", + "Observation", + "Surgery", + "General Surgery", + "Ambulatory Surgery", + "Intensive Care", + "Trauma", + "HIV", + "Human Immunodeficiency Virus", + "Major Joint Replacement", + "Transplant", + "OBGYN", + "Obstetrician/Gynecologist", + "Obstetrician", + "Gynecologist", + "Opthalmology & Vision", + "Opthalmology", + "Vision", + "Never Events", + "Medically Unnecessary", + "Medically Unnecessary Procedures", + "Not Medically Necessary", + "Not Medically Necessary Procedures", + "Experimental", + "Investigational", + "Unlisted Codes", + "Durable Medical Equipment", + "DME", + "Prosthetics & Orthotics", + "Prosthetics", + "Orthotics", + "Implants", + "Hearing Aids", + "Hearing", + "Anesthesia", + "Anesthesiology", + "Medical Pharmacy", + "Physician Administered Drugs", + "Global", + "Bundled or Unbundled Codes", + "Bundled Codes", + "Unbundled Codes", + "Multiple Procedure Reductions", + "Second Surgery", + "Subsequent Surgeries", + "Non-Behavioral Health Mid-Level Professionals", + "Physician Assistant", + "PA", + "Nurse Practicioner", + "NP", + "Non-Physician", + "Non-Physician Health Professionals", + "Technical Component", + "Professional Component", + "Laboratory", + "Pathology", + "Lab", + "Path", + "Lab/Path", + "Laboratory/Pathology", + "Radiology", + "Imaging", + "Radiology/Imaging", + "Mammography", + "Diagnostic", + "Pre-Admission Procedures", + "Post-Discharge Procedures", + "Readmission", + "Status Indicators", + "Stop Loss", + "Emergency Medical", + "EMS", + "NICU", + "Neonatal Intensive Care Unit", + "Vaccine for Children", + "VFC", + "Surgical Assistant", + "Physicians/Clinical Psychologists", + "Doctor of Nursing Practice", + "Osteopathic Medicine", + "Clinical Psychology", + "Audiologist", + "Chiropractors", + "Registered Dietician", + "AUD", + "DC", + "RD", + "Board Certified Behavioral Analysis", + "Behavioral Analysis", + "BCBA", + "Independent Licensures", + "Licensed Professional Counselor", + "Marriage and Family Therapist", + "Substance Abuse Counselor", + "Clinical Social Worker", + "Behavioral Health Outpatient Clinic", + "Physical Therapist", + "Occupational Therapist", + "Speech Therapist", + "PT", + "OT", + "ST", + "Transportation", + "Primary Care", + "Primary Care Behavioral Health", + "Behavioral Physician", + "Clinical Psychologist", + "Mid-Level Practicioner", + "Dentist", + "Dental", + "Supplies and Devices", + "Supplies", + "Devices", + "Immunizations", + "Obstetrical Epidural", + "Pediatric Subspecialties", + "Orthopedic Surgery", + "All Other Specialists", + "Specialty Care Physician", + "Pediatric Primary Care", + "Nurse Anesthetist", + "Early Periodic Screening, Diagnostic, and Treatment", + "EPSDT", + "Ancillary", + "Oncology", + "Cancer", + "Inpatient Physical Rehabilitation", + "Inpatient Rehabilitation", + "Outpatient Rehabilitation", + "Rehabilitation", + "Non-Behavioral Health Rehabilitation", + "Extracorporeal Shock Wave Lithotripsy", + "Cardiac", + "Special Care Unit", + "Skilled Nursing", + "Infusion", + "Specialty Care", + "Specialty Care Physician", + "Organ Acquisition", + "Blood Products", + "Blood Products Outpatient", + "Blood Products Inpatient", + "High Cost Drugs", + "Sleep Studies", + "NICU", + "Newborn Intensive Care Unit", + "Extracorporeal Membrane Oxygenation", + "Burns", + "Kyphoplasty", + "Cryosurgical Ablation of the Prostate", + "Transurethral Thermal Ablation", + "TUMT", + "Transurethral Needle Ablation", + "TUNA", + "Hyperbaric Treatment", + "Clinic Visit", + "Boarder Baby", + "Pediatric Intensive Care Unit", + "PICU", + "Psychiatric", + "Mental Health", + "Behavioral Health", + "Substance Abuse", + "Behavioral Health and Substance Abuse", + "Sub-Acute Facility Care", + "Unrouped Inpatient", + "All Other Acute", + "Neurology", + "Neurology Subspecialties", + "Automatic Implantable Cardioverter Defibrillator", + "Percutaneous Transluminal Coronary Angioplasty", + "Non-Coronary Angioplasty", + "Cardiac Catheters", + "Cardiovascular Surgery", + "Cardiac Surgery", + "Cesarean Birth", + "Cesarean Section", + "C-Section", + "Gamma-Knife Radio-Surgery Outpatient", + "DaVinci Robotic Assisted Surgery", + "Outpatient Electrophysiology with Ablation", + "Outpatient Electrophysiology", + "Magnetic Resonance Image", + "MRI", + "Computed Tomography Scan", + "CT Scan", + "Radiation Therapy", + "Dialysis", + "Gastric Bypass", + "Lap Band", + "Obesity", + "Laparoscopic Cholecystectomy", + "Lap Chole", + "Laparoscopic Hysterectomy", + "Laparoscopic Prostatectomy", + "Laparoscopic Hysteroscopy", + "Treatment Room", + "Wound Care", + "Cardiac Computed Tomograpy & Angiography", + "Positron Emission Tomography Scan", + "PET Scan", + "Hematology", + ] def prompt_contract_lob(page): @@ -316,8 +489,9 @@ def prompt_product(page): Extract and list all identified plan or product names. If multiple names are found, return them separated by commas. If no specific plan name is found, return 'N/A'. Only return the extracted information and no other text or sentences or explanations. Only the final values. """ + def prompt_network_name(page): - #MODIFY TO DICTS FOR MAPPING + # MODIFY TO DICTS FOR MAPPING return f"""### PAGE START ### {page} ### PAGE END Extract all network names mentioned on the page. A 'network name' in the context of health insurance refers to the type of managed care organization involved in the delivery of healthcare services. These names often indicate the structure or model of care delivery, which may define the relationships between insurers, healthcare providers, and insured individuals. @@ -337,6 +511,7 @@ def prompt_service_area(page): Extract and return all found attributes separated by commas in a single string. If no specific service areas are mentioned in the text, return 'N/A'. only return the extracted information and no other text or sentences or explanations. Only the final values. """ + def prompt_effective_date(page): return f"""### PAGE START ### {page} ### PAGE END @@ -346,6 +521,7 @@ def prompt_effective_date(page): DO NOT REPEAT THE SAME DATE MULTIPLE TIMES. """ + def prompt_termination_date(page): return f"""### PAGE START ### {page} ### PAGE END @@ -357,51 +533,65 @@ def prompt_termination_date(page): DO NOT REPEAT THE SAME DATE MULTIPLE TIMES. """ + def prompt_metal_level(page): return f"""### PAGE START ### {page} ### PAGE END If the Line of Business is identified as Marketplace, extract all instances of metal levels such as Platinum, Gold, Silver, or Bronze. List each metal level. If no metal levels are found or if the LOB is not Marketplace, return 'N/A'. only return the extracted information and no other text or sentences or explanations. Only the final values seperated by commas.""" + def prompt_claim_discount_ind(lob): return f"For this line of business: {lob}, return with a 'Y' or 'NO' only whether the claim for this LOB is discounted. Only return as 'Y' or 'NO'." + def prompt_claim_discount_percent_rate(lob): return f"For this LOB only: {lob}, extract the claim discount percent rate." + def prompt_claim_discount_start_date(lob): return f"For this LOB only: {lob}, extract the claim discount start date." + def prompt_claim_discount_termination_date(lob): return f"For this LOB only: {lob}, extract the claim discount termination date." + def prompt_premium_ind(lob): return f"""For this Line of Business{lob}, indicate with a 'Y' or 'N' only whether the premium for this LOB is subject to any adjustments. Return only 'Y' or 'N'.""" + def prompt_premium_percent(lob): return f"""For this Line of Business{lob}, extract the percentage rate of premium adjustment if applicable. Provide the percentage as a numerical value only.""" + def prompt_premium_start_date(lob): return f"""For this Line of Business{lob}, identify the start date of the premium adjustments. Return the date in MM/DD/YYYY format only.""" + def prompt_premium_termination_date(): return """For this Line of Business, determine the termination date of the premium adjustments. Provide the date in MM/DD/YYYY format only.""" + def prompt_sequestration_ind(): return """For this Line of Business, indicate with a 'Y' or 'N' whether there is any sequestration applied. Return only 'Y' or 'N'.""" + def prompt_sequestration_rate(): return """For this Line of Business, extract the rate of sequestration as a percentage. Return the rate as a numerical value only.""" + def prompt_sequestration_start_date(): return """For this Line of Business, determine the start date for sequestration. Return the date in MM/DD/YYYY format only.""" + def prompt_sequestration_termination_date(): return """For this Line of Business, identify the termination date for sequestration. Provide the date in MM/DD/YYYY format only.""" + def prompt_penalties_ind(): return """For this Line of Business, indicate with a 'Y' or 'N' only whether there are any penalties applied. Return only 'Y' or 'N'.""" + def prompt_penalties_rate(): return """For this Line of Business, extract the rate of any penalties applied as a percentage. Provide the rate as a numerical value only.""" - diff --git a/textract-pipeline/src/lambda/prompt-orchestrator/table_funcs.py b/textract-pipeline/src/lambda/prompt-orchestrator/table_funcs.py index 411873f..af762f0 100644 --- a/textract-pipeline/src/lambda/prompt-orchestrator/table_funcs.py +++ b/textract-pipeline/src/lambda/prompt-orchestrator/table_funcs.py @@ -1,20 +1,22 @@ - import re import ast + def clean_tables(text): def replace_colon(text): # Define a function to use in re.sub to check the context of the match def replacer(match): # Extract the character after ':' to check if it's '[' following_text = match.group(1) - if following_text.strip().startswith('['): - return match.group(0) # Return the original match (':' and whatever follows) + if following_text.strip().startswith("["): + return match.group( + 0 + ) # Return the original match (':' and whatever follows) else: - return '' + match.group(1) # Replace ':' with ';' and return the rest + return "" + match.group(1) # Replace ':' with ';' and return the rest # Use a regular expression to find ':' and the text that follows - pattern = r':(\s*.)' + pattern = r":(\s*.)" replaced_text = re.sub(pattern, replacer, text) return replaced_text @@ -25,20 +27,22 @@ def clean_tables(text): if "'" not in content and (len(content) > 0): return content # Return just the content without brackets else: - return match.group(0) # Return the original match if it contains a quote - + return match.group( + 0 + ) # Return the original match if it contains a quote # Use a regular expression to find brackets and the text within - pattern = r'\[([^\[\]]*?)\]' + pattern = r"\[([^\[\]]*?)\]" cleaned_text = re.sub(pattern, replacer, text) return cleaned_text - + # Clean tables - text = text.strip('{}') # Remove brackets + text = text.strip("{}") # Remove brackets text = replace_colon(text) text = remove_unquoted_brackets(text) return text + def convert_to_dict(table_text): """ Converts a formatted table text into a dictionary where each key maps to a list of values. @@ -54,22 +58,22 @@ def convert_to_dict(table_text): """ table_text = clean_tables(table_text) # print(table_text, '\n') - table_text_list = table_text.split(']') + table_text_list = table_text.split("]") table_text_list = [item for item in table_text_list if len(item) > 0] final_dict = {} num_elements = 0 for key_value_text in table_text_list: - key = key_value_text.split(':')[0].strip(', \'') + key = key_value_text.split(":")[0].strip(", '") try: - value = ':'.join(key_value_text.split(':')[1:]).strip()+']' + value = ":".join(key_value_text.split(":")[1:]).strip() + "]" # print(value, '\n') if len(value) > 1: value_list = ast.literal_eval(value) final_dict[key] = value_list num_elements = len(value_list) except: - value = ':'.join(key_value_text.split(':')[1:]).strip()+"']" + value = ":".join(key_value_text.split(":")[1:]).strip() + "']" # print(value, '\n') if len(value) > 1: value_list = ast.literal_eval(value) @@ -96,13 +100,13 @@ def format_table(table_json, table_size): """ table_text = "" keys = [key for key in table_json.keys()] - table_text += ' : '.join(keys) + "\n" + table_text += " : ".join(keys) + "\n" for i in range(table_size): for key in keys: if table_json[key][i]: - #table_text += key + ': ' + table_json[key][i] + ', ' - table_text += table_json[key][i] + ' : ' - table_text += r'\n' + # table_text += key + ': ' + table_json[key][i] + ', ' + table_text += table_json[key][i] + " : " + table_text += r"\n" return table_text @@ -123,17 +127,23 @@ def align_and_format_tables(text_dict): aligned_text_dict = {} for key, text in text_dict.items(): aligned_text = text - if 'Table Start' in text: - table_texts = re.findall(r'-------Table Start--------(.*?)-------Table End--------', text, re.DOTALL) + if "Table Start" in text: + table_texts = re.findall( + r"-------Table Start--------(.*?)-------Table End--------", + text, + re.DOTALL, + ) for table_text in table_texts: # print(table_text, '\n') - + try: # Extract pretable text - pretable = table_text.split('{')[0].strip() + pretable = table_text.split("{")[0].strip() # Extract and format table text - #table_only = "{" + table_text.split('{', 1)[1].rsplit('}', 1)[0].replace("'", '"') + "}" - table_only = "{" + table_text.split('{', 1)[1].rsplit('}', 1)[0] + "}" + # table_only = "{" + table_text.split('{', 1)[1].rsplit('}', 1)[0].replace("'", '"') + "}" + table_only = ( + "{" + table_text.split("{", 1)[1].rsplit("}", 1)[0] + "}" + ) table, table_size = convert_to_dict(table_only) table_formatted = format_table(table, table_size) @@ -141,16 +151,22 @@ def align_and_format_tables(text_dict): table_formatted = table_text # Align - if text.count(pretable) == 2: # One table + if text.count(pretable) == 2: # One table # Remove original table aligned_text = text.replace(table_text, "") # Align new table - aligned_text = aligned_text.replace(pretable, pretable + table_formatted) - + aligned_text = aligned_text.replace( + pretable, pretable + table_formatted + ) + else: - aligned_text = aligned_text.replace(table_text, pretable + ' ' + table_formatted) - aligned_text_dict[key] = aligned_text.replace("-------Table Start--------", "").replace("-------Table End--------", "") + aligned_text = aligned_text.replace( + table_text, pretable + " " + table_formatted + ) + aligned_text_dict[key] = aligned_text.replace( + "-------Table Start--------", "" + ).replace("-------Table End--------", "") else: aligned_text_dict[key] = text - - return aligned_text_dict \ No newline at end of file + + return aligned_text_dict diff --git a/textract-pipeline/src/lambda/prompt-orchestrator/test.py b/textract-pipeline/src/lambda/prompt-orchestrator/test.py index 55d4330..720ca32 100644 --- a/textract-pipeline/src/lambda/prompt-orchestrator/test.py +++ b/textract-pipeline/src/lambda/prompt-orchestrator/test.py @@ -62,12 +62,14 @@ Only return the list of dictionaries, with no other commentary or explanation. E """ -d = {'SERVICE': 'Inpatient Services', - 'REIMBURSEMENT_FLAT_FEE': 'N/A', - 'REIMBURSEMENT_RATE': '105%', - 'FULL_METHODOLOGY': "The lesser of 105% of the West Virginia Medicaid DRG (based on Hospital's current DHHR payment rate) or 100% of Hospital's allowable billed charges.", - 'page_num': '26', - 'Filename': 'App Regional Healthcare_Hospital Agreement_20160414_Dually Executed WV MCD-C13439125AA.txt'} +d = { + "SERVICE": "Inpatient Services", + "REIMBURSEMENT_FLAT_FEE": "N/A", + "REIMBURSEMENT_RATE": "105%", + "FULL_METHODOLOGY": "The lesser of 105% of the West Virginia Medicaid DRG (based on Hospital's current DHHR payment rate) or 100% of Hospital's allowable billed charges.", + "page_num": "26", + "Filename": "App Regional Healthcare_Hospital Agreement_20160414_Dually Executed WV MCD-C13439125AA.txt", +} answer = prompt_funcs.run_bottom_up_lesser(d, text_dict) print(answer) @@ -80,13 +82,8 @@ print(answer) # print(answer_dicts_filtered) # for d in answer_dicts_filtered: -# print(f"Original dictionary: {d}") +# print(f"Original dictionary: {d}") # # Bottom Up Lesser # if config.RUN_LESSER: # lesser_object = prompt_funcs.run_bottom_up_lesser(d.copy(), text_dict.copy(), tokens=4000) # print(f"Lesser dictionary: {lesser_object}") - - - - - diff --git a/textract-pipeline/src/lambda/prompt-orchestrator/textract_template.py b/textract-pipeline/src/lambda/prompt-orchestrator/textract_template.py index 510d8e1..9a4ef6a 100644 --- a/textract-pipeline/src/lambda/prompt-orchestrator/textract_template.py +++ b/textract-pipeline/src/lambda/prompt-orchestrator/textract_template.py @@ -1,4 +1,3 @@ - """ This script was written to extract .txt files from pdf for adhoc client runs. Ensure you have permissions for API gateway, S3, Lambda function and SQS to execute this (DEVELOPER & above roles in DEV & UAT for Doczy should suffice) @@ -13,86 +12,75 @@ The output text file can be found in the client_bucket/contract_text_file/batch_ And the final LLM parsed outputs can be found in client_bucket/final_output/batch_123456 """ - - - import os import boto3 import requests import json import boto3 from datetime import datetime -import os +import os from botocore.auth import SigV4Auth from botocore.awsrequest import AWSRequest from botocore.credentials import get_credentials from botocore.session import Session import time + # Function to upload files to S3 def upload_files_to_s3(directory, bucket, batch_id): contract_list = [] for filename in os.listdir(directory): - if filename.endswith('.pdf'): + if filename.endswith(".pdf"): file_path = os.path.join(directory, filename) - s3_key = f'contracts-landing-zone/{batch_id}/{filename}' - + s3_key = f"contracts-landing-zone/{batch_id}/{filename}" + # Upload file to S3 s3_client.upload_file(file_path, bucket, s3_key) - print(f'Uploaded {filename} to S3 bucket\n') + print(f"Uploaded {filename} to S3 bucket\n") # Add tags to the uploaded file s3_client.put_object_tagging( Bucket=bucket, Key=s3_key, - Tagging={ - 'TagSet': [ - { - 'Key': 'BatchId', - 'Value': batch_id - } - ] + Tagging={"TagSet": [{"Key": "BatchId", "Value": batch_id}]}, + ) + print(f"Added tags to {filename}\n") + # Add file details to contract list + # For now we can leave it as is since only A, C have been operationalized + contract_list.append( + { + "contract_name": filename, + "groups": [ + "A", # This needs to be dynamic in UI 1 based on what group has been selected + "C", + ], + "contract_source_path": s3_key, } ) - print(f'Added tags to {filename}\n') - # Add file details to contract list - # For now we can leave it as is since only A, C have been operationalized - contract_list.append({ - "contract_name": filename, - "groups": [ - "A", # This needs to be dynamic in UI 1 based on what group has been selected - "C" - ], - "contract_source_path": s3_key - }) return contract_list - - def create_batch(client_bucket, create_batch_url): - myobj = { "client-bucket-name": client_bucket } + myobj = {"client-bucket-name": client_bucket} # Call create batch API endpoint - response = requests.post(create_batch_url, json = myobj) + response = requests.post(create_batch_url, json=myobj) if response.status_code >= 200 and response.status_code < 300: try: - new_batch_id = json.loads(json.loads(response.text)['body'])['batch_id'] - landing_zone = json.loads(json.loads(response.text)['body'])['landing_zone'] + new_batch_id = json.loads(json.loads(response.text)["body"])["batch_id"] + landing_zone = json.loads(json.loads(response.text)["body"])["landing_zone"] except: print(myobj) print(response.text) - new_batch_id = 'failed_cases' - landing_zone = 'contracts_landing_zone' + new_batch_id = "failed_cases" + landing_zone = "contracts_landing_zone" else: print(response.text) - new_batch_id = 'failed_cases' - landing_zone = 'contracts_landing_zone' + new_batch_id = "failed_cases" + landing_zone = "contracts_landing_zone" return new_batch_id, landing_zone - - def list_filtered_files(bucket_name, prefix, start_date, end_date, profile_name): """ List all files in an S3 bucket filtered by date range. @@ -108,30 +96,29 @@ def list_filtered_files(bucket_name, prefix, start_date, end_date, profile_name) - list of str: the keys of the filtered files """ # Parse the dates - start_date = datetime.fromisoformat(start_date.replace('Z', '+00:00')) - end_date = datetime.fromisoformat(end_date.replace('Z', '+00:00')) + start_date = datetime.fromisoformat(start_date.replace("Z", "+00:00")) + end_date = datetime.fromisoformat(end_date.replace("Z", "+00:00")) # Initialize a session using the specified profile session = boto3.Session(profile_name=profile_name) - s3_client = session.client('s3') + s3_client = session.client("s3") # List objects in the bucket with the specified prefix - paginator = s3_client.get_paginator('list_objects_v2') + paginator = s3_client.get_paginator("list_objects_v2") page_iterator = paginator.paginate(Bucket=bucket_name, Prefix=prefix) # Filter files by date filtered_files = [] for page in page_iterator: - if 'Contents' in page: - for obj in page['Contents']: - last_modified = obj['LastModified'] + if "Contents" in page: + for obj in page["Contents"]: + last_modified = obj["LastModified"] if start_date <= last_modified <= end_date: - filtered_files.append(obj['Key']) + filtered_files.append(obj["Key"]) return filtered_files - def download_files(bucket_name, file_keys, profile_name, local_directory): """ Download files from an S3 bucket. @@ -144,7 +131,7 @@ def download_files(bucket_name, file_keys, profile_name, local_directory): """ # Initialize a session using the specified profile session = boto3.Session(profile_name=profile_name) - s3_client = session.client('s3') + s3_client = session.client("s3") # Ensure the local directory exists if not os.path.exists(local_directory): @@ -160,37 +147,37 @@ def download_files(bucket_name, file_keys, profile_name, local_directory): print("Download completed.") - if __name__ == "__main__": # Define your variables - s3_bucket = 'doczyai-use2-u-cn1-s3-textract-processing-001' + s3_bucket = "doczyai-use2-u-cn1-s3-textract-processing-001" # batch_id = 'batch_100524101551' - client_name = 'Priority Health' - username = 'ADHOC USER' + client_name = "Priority Health" + username = "ADHOC USER" # These endpoints are in UAT, please change them to DEV if there are access issues with UAT - api_endpoint = 'https://29gm8cek03.execute-api.us-east-2.amazonaws.com/dev/trigger-pipeline' - create_batch_url = "https://29gm8cek03.execute-api.us-east-2.amazonaws.com/dev/create-batch" + api_endpoint = ( + "https://29gm8cek03.execute-api.us-east-2.amazonaws.com/dev/trigger-pipeline" + ) + create_batch_url = ( + "https://29gm8cek03.execute-api.us-east-2.amazonaws.com/dev/create-batch" + ) - - - session = boto3.Session(profile_name='temp_cred') # Change the profile name to the one you have in your .aws/credentials file - s3_client = session.client('s3') - - + session = boto3.Session( + profile_name="temp_cred" + ) # Change the profile name to the one you have in your .aws/credentials file + s3_client = session.client("s3") # Use current directory as PDF directory pdf_directory = "C:\\Doczy\\Priority Health\\For_Textract\\For_Textract\\UAT- test East Paris Surgical Center copy" - # Create batch batch_id, landing_zone = create_batch(s3_bucket, create_batch_url) print(f"Batch ID: {batch_id}") print(f"Landing Zone: {landing_zone}") # batch_id = 'batch_110624213433' - # landing_zone = 'contracts_landing_zone/batch_110624213433/' + # landing_zone = 'contracts_landing_zone/batch_110624213433/' - if batch_id == 'failed_cases': + if batch_id == "failed_cases": print("Batch creation failed. Exiting...") exit() # Upload files and get contract list @@ -202,16 +189,12 @@ if __name__ == "__main__": "batch_id": batch_id, "client_name": client_name, "username": username, - "contract_list": contract_list + "contract_list": contract_list, } # Make POST request to API response = requests.post(api_endpoint, json=data) - - - - # Print response print(response.status_code) print(response.json()) @@ -219,18 +202,22 @@ if __name__ == "__main__": ################## # Get text files - prefix = f'contract-text-file/{batch_id}/' # if you have a specific prefix (folder) in your bucket + prefix = f"contract-text-file/{batch_id}/" # if you have a specific prefix (folder) in your bucket # These dates are to filter the contracts in case there are older contracts in the same batch - start_date = '2024-06-04T00:00:00Z' # ISO 8601 format - end_date = '2024-06-14T23:59:59Z' # ISO 8601 format - profile_name = 'temp_cred' - local_directory = 'C:\\Doczy\\Priority Health\\For_Textract\\For_Textract\\text otuput' # local directory to save files + start_date = "2024-06-04T00:00:00Z" # ISO 8601 format + end_date = "2024-06-14T23:59:59Z" # ISO 8601 format + profile_name = "temp_cred" + local_directory = "C:\\Doczy\\Priority Health\\For_Textract\\For_Textract\\text otuput" # local directory to save files # List filtered files - time.sleep(30) # Waiting for the text files to be generated, this may take longer and files may not be available after 30 seconds sometimes - filtered_files = list_filtered_files(s3_bucket, prefix, start_date, end_date, profile_name) + time.sleep( + 30 + ) # Waiting for the text files to be generated, this may take longer and files may not be available after 30 seconds sometimes + filtered_files = list_filtered_files( + s3_bucket, prefix, start_date, end_date, profile_name + ) print(f"Filtered files: {len(filtered_files)}") # Download files - download_files(s3_bucket, filtered_files, profile_name, local_directory) \ No newline at end of file + download_files(s3_bucket, filtered_files, profile_name, local_directory) diff --git a/textract-pipeline/src/lambda/prompt-orchestrator/top_down_funcs.py b/textract-pipeline/src/lambda/prompt-orchestrator/top_down_funcs.py index cbfd6a6..0397ea0 100644 --- a/textract-pipeline/src/lambda/prompt-orchestrator/top_down_funcs.py +++ b/textract-pipeline/src/lambda/prompt-orchestrator/top_down_funcs.py @@ -1,10 +1,10 @@ - import prompts import claude_funcs import utils import json + def run_top_down_metal_level(d, page): """ Identifies and extracts the metal level of a contract within a specific line of business (LOB) from the provided page text. @@ -20,12 +20,15 @@ def run_top_down_metal_level(d, page): Returns: str: The metal level of the contract as determined by the analysis, or 'N/A' if the contract LOB is not applicable. """ - - if 'MARKETPLACE' in str(d['CONTRACT_LOB']).upper() or 'COMMERCIAL' in str(d['CONTRACT_LOB']).upper(): - prompt = prompts.TOP_DOWN_METAL_LEVEL(d['CONTRACT_LOB'], page) + + if ( + "MARKETPLACE" in str(d["CONTRACT_LOB"]).upper() + or "COMMERCIAL" in str(d["CONTRACT_LOB"]).upper() + ): + prompt = prompts.TOP_DOWN_METAL_LEVEL(d["CONTRACT_LOB"], page) answer = claude_funcs.invoke_claude_3(prompt, max_tokens=4000) else: - answer = 'N/A' + answer = "N/A" return answer @@ -46,8 +49,8 @@ def run_top_down_date(type_, d, page): Returns: str: The extracted date as a string, based on the model's interpretation of the input prompt and text context. """ - formatted_d = utils.format_td_check([d], ['Filename', 'page_num']) - prompt = prompts.TOP_DOWN_DATE('EFFECTIVE', formatted_d, page) + formatted_d = utils.format_td_check([d], ["Filename", "page_num"]) + prompt = prompts.TOP_DOWN_DATE("EFFECTIVE", formatted_d, page) answer = claude_funcs.invoke_claude_3(prompt, max_tokens=4000) return answer @@ -71,19 +74,24 @@ def top_down_secondary(td_results, text_dict): updated_dicts = [] for d in td_results: # Run Metal Level - d['CONTRACT_MARKETPLACE_METAL_LEVEL'] = run_top_down_metal_level(d, text_dict[d['page_num']]) - + d["CONTRACT_MARKETPLACE_METAL_LEVEL"] = run_top_down_metal_level( + d, text_dict[d["page_num"]] + ) + # Dates - d['LOB_PRICING_TERMS_EFFECTIVE_DATE'] = run_top_down_date('EFFECTIVE' , d, text_dict[d['page_num']]) - + d["LOB_PRICING_TERMS_EFFECTIVE_DATE"] = run_top_down_date( + "EFFECTIVE", d, text_dict[d["page_num"]] + ) + # if auto_renewal == N: else: 'N/A' - d['LOB_PRICING_TERMS_TERMINATION_DATE'] = run_top_down_date('TERMINATION', d, text_dict[d['page_num']]) + d["LOB_PRICING_TERMS_TERMINATION_DATE"] = run_top_down_date( + "TERMINATION", d, text_dict[d["page_num"]] + ) updated_dicts.append(d) return updated_dicts - def run_top_down(filename, text_dict): """ Executes the Top Down processing strategy on a provided dictionary of text pages, extracting structured data based on specified prompts. @@ -110,14 +118,14 @@ def run_top_down(filename, text_dict): prompt = prompts.TOP_DOWN_PRIMARY(page_text) answer = claude_funcs.invoke_claude_3(prompt, max_tokens=4000) answer_dict = json.loads(answer) - answer_dict.update({'page_num' : page_num, 'Filename' : filename}) - #answer_dicts = dict_operations.primary_string_to_dict({page_num : answer}, filename) # Convert to dictionaries + answer_dict.update({"page_num": page_num, "Filename": filename}) + # answer_dicts = dict_operations.primary_string_to_dict({page_num : answer}, filename) # Convert to dictionaries all_results.append(answer_dict) - + # td_primary = [i for d in all_results for i in d] # Consolidate to one list # print(td_primary) # Secondaries td_final = top_down_secondary(all_results, text_dict) - return td_final # List of list of dictionaries \ No newline at end of file + return td_final # List of list of dictionaries diff --git a/textract-pipeline/src/lambda/prompt-orchestrator/utils.py b/textract-pipeline/src/lambda/prompt-orchestrator/utils.py index 43a16b0..2cab891 100644 --- a/textract-pipeline/src/lambda/prompt-orchestrator/utils.py +++ b/textract-pipeline/src/lambda/prompt-orchestrator/utils.py @@ -1,4 +1,3 @@ - import os import re import pandas as pd @@ -6,85 +5,97 @@ import shutil import config + def read_local(file_path): # Check if the file is a text file - if os.path.isfile(file_path) and file_path.endswith('.txt'): + if os.path.isfile(file_path) and file_path.endswith(".txt"): try: # First attempt to open the file with UTF-8 encoding - with open(file_path, 'r', encoding='utf-8') as file: + with open(file_path, "r", encoding="utf-8") as file: file_contents = file.read() return file_contents except UnicodeDecodeError: # If UTF-8 fails, try reading the file with ANSI encoding try: - with open(file_path, 'r', encoding='cp1252') as file: + with open(file_path, "r", encoding="cp1252") as file: file_contents = file.read() return file_contents except UnicodeDecodeError: # If ANSI also fails, log an error message or handle it accordingly print(f"Failed to decode {file_path} with UTF-8 and cp1252 encodings.") + def read_s3(): s3_client = config.S3_CLIENT objects = s3_client.list_objects_v2(Bucket=config.BUCKET, Prefix=config.PREFIX) file_list = [] - for obj in objects['Contents']: - if not obj['Key'].endswith('/'): - file_list.append(obj['Key']) + for obj in objects["Contents"]: + if not obj["Key"].endswith("/"): + file_list.append(obj["Key"]) contract_list = sorted(file_list) files = {} for contract in contract_list: data = s3_client.get_object(Bucket=config.BUCKET, Key=contract) - contents = data['Body'].read() - context = contents.decode('utf-8') + contents = data["Body"].read() + context = contents.decode("utf-8") path, filename = os.path.split(contract) files[filename] = context return files + def read_input(path=config.LOCAL_PATH, mode=config.READ_MODE): - if mode == '_LOCAL_': + if mode == "_LOCAL_": files = {} for file in os.listdir(path): full_path = os.path.join(path, file) file_text = read_local(full_path) files[file] = file_text return files - elif mode == '_S3_': + elif mode == "_S3_": return read_s3() -def consolidate_individual(input_folder='results', output_folder='output'): + +def consolidate_individual(input_folder="results", output_folder="output"): dfs = [] for filename in os.listdir(input_folder): - if filename.endswith('.csv'): + if filename.endswith(".csv"): filepath = os.path.join(input_folder, filename) dfs.append(pd.read_csv(filepath)) - + # Remove temp folder - if input_folder=='temp' and os.path.exists(input_folder): + if input_folder == "temp" and os.path.exists(input_folder): shutil.rmtree(input_folder) - + consolidated_df = pd.concat(dfs, ignore_index=True) # Write output version = 1 - existing_files = [filename for filename in os.listdir(output_folder) if filename.startswith(f'consolidated_results_{config.TODAY}')] + existing_files = [ + filename + for filename in os.listdir(output_folder) + if filename.startswith(f"consolidated_results_{config.TODAY}") + ] if existing_files: - versions = [int(file.split('_v')[1].split('.')[0]) for file in existing_files if '_v' in file] + versions = [ + int(file.split("_v")[1].split(".")[0]) + for file in existing_files + if "_v" in file + ] if versions: version = max(versions) + 1 - filename = f'consolidated_results_{config.TODAY}_v{version}.csv' + filename = f"consolidated_results_{config.TODAY}_v{version}.csv" consolidated_df.to_csv(os.path.join(output_folder, filename)) def preprocess_text_file(file_path): # Attempt to open the file with UTF-8 encoding first try: - with open(file_path, 'r', encoding='utf-8') as file: + with open(file_path, "r", encoding="utf-8") as file: text = file.read() except UnicodeDecodeError: # If UTF-8 fails, try reading the file with ANSI (cp1252) encoding try: - with open(file_path, 'r', encoding='cp1252') as file: + with open(file_path, "r", encoding="cp1252") as file: text = file.read() except UnicodeDecodeError: # If ANSI also fails, log an error message or handle it accordingly @@ -92,7 +103,7 @@ def preprocess_text_file(file_path): return [] # Split the text into pages based on a specific marker - pages = re.split(r'Start of Page No\. = \d+', text) + pages = re.split(r"Start of Page No\. = \d+", text) return pages @@ -100,11 +111,11 @@ def format_td_check(td_dicts, dont_include_list): final_str = "" dict_count = 1 for td_dict in td_dicts: - final_str += str(dict_count) + '. ' + final_str += str(dict_count) + ". " for k in td_dict.keys(): if k not in dont_include_list: - final_str += k + ': ' + td_dict[k] + ', ' - final_str += '\n' + final_str += k + ": " + td_dict[k] + ", " + final_str += "\n" dict_count += 1 return final_str @@ -126,12 +137,7 @@ def consolidate_csvs(output_dir, output_file): pass concatenated_df = pd.concat(df_list, ignore_index=True) - concatenated_df.to_csv(os.path.join(config.CONSOLIDATED_OUTPUT_DIRECTORY, output_file), index=False) + concatenated_df.to_csv( + os.path.join(config.CONSOLIDATED_OUTPUT_DIRECTORY, output_file), index=False + ) print(f"All CSV files have been consolidated into {output_file}") - - - - - - - diff --git a/textract-pipeline/src/lambda/prompt-processor/index.py b/textract-pipeline/src/lambda/prompt-processor/index.py index 30cce0f..a5d8121 100644 --- a/textract-pipeline/src/lambda/prompt-processor/index.py +++ b/textract-pipeline/src/lambda/prompt-processor/index.py @@ -16,16 +16,16 @@ logger = logging.getLogger() logger.setLevel(logging.INFO) # Initialize S3 & Textract clients -s3_client = boto3.client('s3') +s3_client = boto3.client("s3") DATABASE_LOGGING_LAMBDA_FUNCTION_NAME = "" # Function to load configuration from S3 def load_config_from_s3(bucket_name, file_key): - + # Download the config file from S3 response = s3_client.get_object(Bucket=bucket_name, Key=file_key) - config_content = response['Body'].read().decode('utf-8') + config_content = response["Body"].read().decode("utf-8") # Parse the config file config_parser = ConfigParser() @@ -34,23 +34,23 @@ def load_config_from_s3(bucket_name, file_key): # Convert the configuration to a dictionary config_dict = {} for section in config_parser.sections(): - config_dict[section] = {key.upper(): value for key, value in config_parser.items(section)} + config_dict[section] = { + key.upper(): value for key, value in config_parser.items(section) + } return config_dict + # Function to get s3 object tags def get_s3_object_tags(bucket_name, object_key): try: # Get object tags - response = s3_client.get_object_tagging( - Bucket=bucket_name, - Key=object_key - ) + response = s3_client.get_object_tagging(Bucket=bucket_name, Key=object_key) # Extract tags from the response and convert to dictionary - tags_list = response['TagSet'] - tags_dict = {tag['Key']: tag['Value'] for tag in tags_list} + tags_list = response["TagSet"] + tags_dict = {tag["Key"]: tag["Value"] for tag in tags_list} return tags_dict @@ -59,13 +59,16 @@ def get_s3_object_tags(bucket_name, object_key): logger.error(f"Error: {e}") return None + def lambda_handler(event, context): - + # Read environment variables global DATABASE_LOGGING_LAMBDA_FUNCTION_NAME - DATABASE_LOGGING_LAMBDA_FUNCTION_NAME = os.environ.get('DATABASE_LOGGING_LAMBDA_FUNCTION_NAME', '') - - property_file_path = os.environ.get('PROPERTY_FILE_S3_PATH', '') + DATABASE_LOGGING_LAMBDA_FUNCTION_NAME = os.environ.get( + "DATABASE_LOGGING_LAMBDA_FUNCTION_NAME", "" + ) + + property_file_path = os.environ.get("PROPERTY_FILE_S3_PATH", "") # Read config.properties file_path_array = property_file_path.split("/") # Valid if file_path_array has more than 2 elements @@ -74,18 +77,18 @@ def lambda_handler(event, context): S3_BUCKET_NAME = file_path_array[0] CONFIG_FILE_PATH = "/".join(file_path_array[1:]) - logger.info(f'S3_BUCKET_NAME: {S3_BUCKET_NAME}') - logger.info(f'CONFIG_FILE_PATH: {CONFIG_FILE_PATH}') + logger.info(f"S3_BUCKET_NAME: {S3_BUCKET_NAME}") + logger.info(f"CONFIG_FILE_PATH: {CONFIG_FILE_PATH}") # Load config file config_dict = load_config_from_s3(S3_BUCKET_NAME, CONFIG_FILE_PATH) try: # Process each message from the SQS event - for record in event['Records']: + for record in event["Records"]: # Extract the message body from the record logger.info(f"record: {record}") - record_body = json.loads(record['body']) + record_body = json.loads(record["body"]) logger.info(f"record_body: {record_body}") logger.info(f"record_body type: {type(record_body)}") @@ -95,22 +98,28 @@ def lambda_handler(event, context): logger.info(f"field_list: {field_list}") logger.info(f"file_name: {file_name_with_path}") - path_parts = file_name_with_path.split('/') + path_parts = file_name_with_path.split("/") # Extract the first string s3_bucket = path_parts[0] - file_name_with_subfolder = '/'.join(path_parts[1:]) + file_name_with_subfolder = "/".join(path_parts[1:]) fields = json.loads(field_list) field_names = [field["FIELD_NAME"] for field in fields] required_field_list = "'" + "', '".join(field_names) + "'" fm_model_id = fields[0]["FM_MODEL_ID"] if fields else None - + # Download the file content from S3 file_content = read_text_from_s3(file_name_with_path) - prompt_config_query = "SELECT LISTAGG(CONCAT('[{\"FIELD_NAME:\"',FIELD_NAME,'\"},{\"PROMPT\":\"',PROMPT,'\"}]'),',') FROM STG.PROMPT_CONFIG WHERE FIELD_NAME IN ("+required_field_list+")" + prompt_config_query = ( + "SELECT LISTAGG(CONCAT('[{\"FIELD_NAME:\"',FIELD_NAME,'\"},{\"PROMPT\":\"',PROMPT,'\"}]'),',') FROM STG.PROMPT_CONFIG WHERE FIELD_NAME IN (" + + required_field_list + + ")" + ) - field_prompt_group_response = call_snowflake_db_select_lambda(prompt_config_query) + field_prompt_group_response = call_snowflake_db_select_lambda( + prompt_config_query + ) data_dict = json.loads(field_prompt_group_response) body = data_dict["body"] body_dict = json.loads(body) @@ -118,37 +127,43 @@ def lambda_handler(event, context): body_str = json.dumps(body_dict) logger.info(f"Parsed body: {type(body_dict)}") logger.info(f"body_str: {body_str}") - - file_content += body_str+" Provide answers in VALID JSON in 'completion' tag format against the tag in the bracket along with the file page number. [{\"FIELD_NAME\":,\"ANSWER\":},\"PAGE_NO\":}]" + + file_content += ( + body_str + + ' Provide answers in VALID JSON in \'completion\' tag format against the tag in the bracket along with the file page number. [{"FIELD_NAME":,"ANSWER":},"PAGE_NO":}]' + ) logger.info(f"file_content: {file_content}") api_payload = { "prompt": file_content, - "model_id": 'anthropic.claude-instant-v1', - #"model_id": fm_model_id, + "model_id": "anthropic.claude-instant-v1", + # "model_id": fm_model_id, "max_gen_len": 40000, "temperature": 0.3, - "top_p": 0.5 + "top_p": 0.5, } api_response = make_api_call(api_payload) logger.info(f"API Response==: {api_response}") # Save the response text to a file - #LLM_RESPONSE_FILE_LOCATION = config_dict['FOLDER_LOCATIONS']['LLM_RESPONSE_FILE_LOCATION'].format(batch_id) + # LLM_RESPONSE_FILE_LOCATION = config_dict['FOLDER_LOCATIONS']['LLM_RESPONSE_FILE_LOCATION'].format(batch_id) LLM_RESPONSE_FILE_LOCATION = "/backup/contract-text-file/" - + logger.info(f"file_name_with_subfolder = {file_name_with_subfolder}") logger.info(f"s3_bucket = {s3_bucket}") - final_path = LLM_RESPONSE_FILE_LOCATION+file_name_with_subfolder + final_path = LLM_RESPONSE_FILE_LOCATION + file_name_with_subfolder logger.info(f"final_path = {final_path}") - + # Extract response from the API response response_text = extract_response(api_response) logger.info(f"Extracted 'generation' text: {response_text}") - new_file_name_with_subfolder = create_new_text_file(file_name_with_subfolder,prompt_batch_id) - upload_response_to_s3(response_text, s3_bucket, new_file_name_with_subfolder,"temp_tags") - save_output_to_db(final_path,response_text) - + new_file_name_with_subfolder = create_new_text_file( + file_name_with_subfolder, prompt_batch_id + ) + upload_response_to_s3( + response_text, s3_bucket, new_file_name_with_subfolder, "temp_tags" + ) + save_output_to_db(final_path, response_text) except Exception as e: # Log any unhandled exceptions @@ -156,59 +171,70 @@ def lambda_handler(event, context): logger.error(traceback.format_exc()) raise e -def save_output_to_db(document_id,response_text): + +def save_output_to_db(document_id, response_text): logger.info(f"Saving records to DB for document_id: {document_id}") logger.info(f"Saving records to DB for response_text: {response_text}") - #generate_doczy_pipeline_raw_output(document_id,sf_db_col_name, raw_value, original_page_number) - generate_doczy_pipeline_raw_output(document_id,"PAYER_NAME", "Client Company Inc.", "1") + # generate_doczy_pipeline_raw_output(document_id,sf_db_col_name, raw_value, original_page_number) + generate_doczy_pipeline_raw_output( + document_id, "PAYER_NAME", "Client Company Inc.", "1" + ) -def create_new_text_file(file_name_with_subfolder,prompt_batch_id): - parts = file_name_with_subfolder.split('/') + +def create_new_text_file(file_name_with_subfolder, prompt_batch_id): + parts = file_name_with_subfolder.split("/") # Modify the last part (filename) by adding '_temp' before the file extension - filename_parts = parts[-1].split('.') - new_filename = filename_parts[0] +'_' +str(prompt_batch_id) + '.' + filename_parts[1] + filename_parts = parts[-1].split(".") + new_filename = ( + filename_parts[0] + "_" + str(prompt_batch_id) + "." + filename_parts[1] + ) # Combine the parts back into a string parts[-1] = new_filename - new_string = '/'.join(parts) + new_string = "/".join(parts) return new_string + def read_text_from_s3(s3_path): # Initialize S3 client - s3 = boto3.client('s3') - + s3 = boto3.client("s3") + try: bucket_name, file_key = s3_path.split("/", 1) - logger.info(f"Reading file from S3 - bucket_name: {bucket_name}, file_key: {file_key}") + logger.info( + f"Reading file from S3 - bucket_name: {bucket_name}, file_key: {file_key}" + ) # Get object from S3 response = s3.get_object(Bucket=bucket_name, Key=file_key) - + # Read text from the response - text = response['Body'].read().decode('utf-8') - + text = response["Body"].read().decode("utf-8") + return text except Exception as e: print("Error:", e) return None - def download_file_from_s3(bucket, key, region): try: - logger.info(f"Downloading file from S3 - Bucket: {bucket}, Key: {key}, Region: {region}") - s3_client = boto3.client('s3', region_name=region) + logger.info( + f"Downloading file from S3 - Bucket: {bucket}, Key: {key}, Region: {region}" + ) + s3_client = boto3.client("s3", region_name=region) response = s3_client.get_object(Bucket=bucket, Key=key) - file_content = response['Body'].read().decode('utf-8') + file_content = response["Body"].read().decode("utf-8") return file_content except Exception as e: # Log the error logger.error(f"Error downloading file from S3: {str(e)}") raise e + def make_api_call(data): try: - lambda_function_name = 'bedrock_call' + lambda_function_name = "bedrock_call" input_json_payload = json.dumps(data) logger.info(input_json_payload) response = call_lambda_function(lambda_function_name, input_json_payload) @@ -220,8 +246,9 @@ def make_api_call(data): logger.error(f"Error making API call: {str(e)}") raise e -def upload_response_to_s3(response_text, bucket_name, object_key,tags): - s3_client = boto3.client('s3') + +def upload_response_to_s3(response_text, bucket_name, object_key, tags): + s3_client = boto3.client("s3") text_value = str(response_text) try: # Upload the JSON response to S3 @@ -229,16 +256,17 @@ def upload_response_to_s3(response_text, bucket_name, object_key,tags): Bucket=bucket_name, Key=object_key, Body=text_value, - ContentType='application/txt', - Tagging=tags + ContentType="application/txt", + Tagging=tags, ) logger.info(f"Saved analysis response to S3: {object_key}") except Exception as e: logger.error(f"Error saving response to S3 - {object_key}: {str(e)}") + def extract_response(api_response): try: - clean_json_response(api_response,["\n",""]) + clean_json_response(api_response, ["\n", ""]) generation_text = api_response.get("generation", "") if generation_text is None or not generation_text.strip(): @@ -246,7 +274,7 @@ def extract_response(api_response): if generation_text is None or not generation_text.strip(): return generation_text else: - start_index = generation_text.find('[') # Find the start of JSON + start_index = generation_text.find("[") # Find the start of JSON json_data_str = generation_text[start_index:] # Extract JSON data # Parse the JSON data into a Python list json_data = json.loads(json_data_str) @@ -256,21 +284,21 @@ def extract_response(api_response): # Log the error logger.error(f"Error extracting 'generation' from API response: {str(e)}") raise e - def clean_json_response(json_data, substrings_to_remove): """ Clean JSON data by removing unwanted substrings from all string values recursively. - + Args: - json_data (dict): JSON data (Python dictionary) to be cleaned. - substrings_to_remove (list): List of substrings to be removed from string values. - + Returns: - dict: Cleaned JSON data (Python dictionary). """ logger.info(f"clean_json_response = json_data : {json_data}") + # Function to recursively clean dictionary values def clean_values(obj): logger.info(f"clean_json_response = obj : {obj}") @@ -287,7 +315,7 @@ def clean_json_response(json_data, substrings_to_remove): return [clean_values(item) for item in obj] else: return obj - + # Clean JSON data cleaned_json_data = clean_values(json_data) return cleaned_json_data @@ -297,80 +325,81 @@ def get_filename_from_path(full_path): return os.path.basename(full_path) -def call_lambda_function(lambda_function_name, json_payload, aws_region='us-east-2'): +def call_lambda_function(lambda_function_name, json_payload, aws_region="us-east-2"): # Create a Lambda client - lambda_client = boto3.client('lambda', region_name=aws_region) + lambda_client = boto3.client("lambda", region_name=aws_region) try: # Invoke the Lambda function response = lambda_client.invoke( FunctionName=lambda_function_name, - InvocationType='RequestResponse', # Use 'Event' for asynchronous invocation + InvocationType="RequestResponse", # Use 'Event' for asynchronous invocation Payload=json_payload, ) # Parse and return the response payload - response_payload = json.loads(response['Payload'].read().decode('utf-8')) + response_payload = json.loads(response["Payload"].read().decode("utf-8")) return response_payload except Exception as e: # Handle any exceptions (e.g., Lambda function not found, permission issues) print(f"Error calling Lambda function: {e}") - return {'error': str(e)} - + return {"error": str(e)} + def generate_document_logs_input(document_id, stage): - current_time = datetime.datetime.now().isoformat() - - data = { + current_time = datetime.datetime.now().isoformat() + + data = { "operation": "update", "table": "DOCUMENT_LOGS", "data": { "DOCUMENT_ID": document_id, - "STAGE" : stage, + "STAGE": stage, "MODIFIED_TIME": current_time, - "MODIFIED_BY": "PROMPT_PROCESSOR" - } - } - logger.info(f"Request: {data}") - - lambda_client = boto3.client('lambda') - response = lambda_client.invoke( - FunctionName=DATABASE_LOGGING_LAMBDA_FUNCTION_NAME, - InvocationType='Event', # Asynchronous invocation - Payload=json.dumps(data).encode('utf-8') - ) - - logger.info("Generated document logs input successfully") - return response + "MODIFIED_BY": "PROMPT_PROCESSOR", + }, + } + logger.info(f"Request: {data}") + + lambda_client = boto3.client("lambda") + response = lambda_client.invoke( + FunctionName=DATABASE_LOGGING_LAMBDA_FUNCTION_NAME, + InvocationType="Event", # Asynchronous invocation + Payload=json.dumps(data).encode("utf-8"), + ) + + logger.info("Generated document logs input successfully") + return response + def call_snowflake_db_select_lambda(query): # Create a Lambda client - lambda_client = boto3.client('lambda') - + lambda_client = boto3.client("lambda") + # Specify your Lambda function name - lambda_function_name = 'doczy-dev-get-snowflake-logs' - + lambda_function_name = "doczy-dev-get-snowflake-logs" + # Prepare payload - payload = { - 'query': query - } + payload = {"query": query} try: # Call the Lambda function response = lambda_client.invoke( FunctionName=lambda_function_name, - InvocationType='RequestResponse', # Synchronous invocation - Payload=json.dumps(payload) + InvocationType="RequestResponse", # Synchronous invocation + Payload=json.dumps(payload), ) - + # Return response body - response_body = response['Payload'].read().decode('utf-8') + response_body = response["Payload"].read().decode("utf-8") return response_body except Exception as e: logger.exception(f"Error: {e}") return None -def generate_doczy_pipeline_raw_output(document_id,sf_db_col_name, raw_value, original_page_number): +def generate_doczy_pipeline_raw_output( + document_id, sf_db_col_name, raw_value, original_page_number +): current_time = datetime.datetime.now().isoformat() """ DOCUMENT_ID VARCHAR(16777216) NOT NULL, @@ -391,21 +420,21 @@ def generate_doczy_pipeline_raw_output(document_id,sf_db_col_name, raw_value, o "operation": "insert", "table": "DOCZY_PIPELINE_RAW_OUTPUT", "data": { - "DOCUMENT_ID" : document_id, + "DOCUMENT_ID": document_id, "SF_DB_COL_NAME": sf_db_col_name, - "RAW_VALUE" : raw_value, - "ORIGINAL_PAGE_NUMBER" : original_page_number, - "CREATED_TIME" : current_time - } + "RAW_VALUE": raw_value, + "ORIGINAL_PAGE_NUMBER": original_page_number, + "CREATED_TIME": current_time, + }, } logger.info(f"Request: {data}") - lambda_client = boto3.client('lambda') + lambda_client = boto3.client("lambda") response = lambda_client.invoke( - FunctionName='doczy-dev-snowflake-log', - InvocationType='Event', # Asynchronous invocation - Payload=json.dumps(data).encode('utf-8') + FunctionName="doczy-dev-snowflake-log", + InvocationType="Event", # Asynchronous invocation + Payload=json.dumps(data).encode("utf-8"), ) logger.info("Generated document logs input successfully") - return response \ No newline at end of file + return response diff --git a/textract-pipeline/src/lambda/s3-folder-details/index.py b/textract-pipeline/src/lambda/s3-folder-details/index.py index f155181..8a42ea3 100644 --- a/textract-pipeline/src/lambda/s3-folder-details/index.py +++ b/textract-pipeline/src/lambda/s3-folder-details/index.py @@ -2,52 +2,55 @@ import json import boto3 from urllib.parse import urlparse + def list_files_s3(s3_url, file_extension=None): # Parse the S3 URL to extract the bucket and key parsed_url = urlparse(s3_url) bucket = parsed_url.netloc - key = parsed_url.path.lstrip('/') + key = parsed_url.path.lstrip("/") # Create an S3 client - s3_client = boto3.client('s3') + s3_client = boto3.client("s3") # List objects in the specified S3 path response = s3_client.list_objects_v2(Bucket=bucket, Prefix=key) # Extract details of each object based on the specified file extension or retrieve all files files_details = [] - for obj in response.get('Contents', []): + for obj in response.get("Contents", []): # Check if a file extension is specified and filter based on it, or retrieve all files - if file_extension is None or file_extension == "*" or obj['Key'].lower().endswith(f".{file_extension.lower()}"): + if ( + file_extension is None + or file_extension == "*" + or obj["Key"].lower().endswith(f".{file_extension.lower()}") + ): # Extract specific details for each object file_details = { - 'Key': obj['Key'], - 'LastModified': obj['LastModified'].isoformat(), - 'Size': obj['Size'], - 'ETag': obj['ETag'] + "Key": obj["Key"], + "LastModified": obj["LastModified"].isoformat(), + "Size": obj["Size"], + "ETag": obj["ETag"], } # Append the details to the list files_details.append(file_details) return files_details + def lambda_handler(event, context): # Extract the S3 URL and file extension from the Lambda event input - s3_url = event.get('s3_url') - file_extension = event.get('file_extension') + s3_url = event.get("s3_url") + file_extension = event.get("file_extension") # Check if the S3 URL is provided if not s3_url: return { - 'statusCode': 400, - 'body': json.dumps('Error: S3 URL is missing in the input.') + "statusCode": 400, + "body": json.dumps("Error: S3 URL is missing in the input."), } # Retrieve details of files in the specified S3 path with optional file extension filtering files_details = list_files_s3(s3_url, file_extension) # Return the details as a JSON response - return { - 'statusCode': 200, - 'body': json.dumps(files_details) - } + return {"statusCode": 200, "body": json.dumps(files_details)} diff --git a/textract-pipeline/src/lambda/text-creation/index.py b/textract-pipeline/src/lambda/text-creation/index.py index bf4be48..3377f35 100644 --- a/textract-pipeline/src/lambda/text-creation/index.py +++ b/textract-pipeline/src/lambda/text-creation/index.py @@ -16,20 +16,21 @@ logger = logging.getLogger() logger.setLevel(logging.INFO) # Initialize S3 & Textract clients -s3_client = boto3.client('s3') -textract_client = boto3.client('textract') +s3_client = boto3.client("s3") +textract_client = boto3.client("textract") final_number_of_signature = 0 -TEXTRACT_SIGNATURE_JSON = '' -TEXTRACT_TABLE_FORM_JSON = '' +TEXTRACT_SIGNATURE_JSON = "" +TEXTRACT_TABLE_FORM_JSON = "" PROCESSED_LOCATION = "" DATABASE_LOGGING_LAMBDA_FUNCTION_NAME = "" + # Function to load configuration from S3 def load_config_from_s3(bucket_name, file_key): - + # Download the config file from S3 response = s3_client.get_object(Bucket=bucket_name, Key=file_key) - config_content = response['Body'].read().decode('utf-8') + config_content = response["Body"].read().decode("utf-8") # Parse the config file config_parser = ConfigParser() @@ -38,7 +39,9 @@ def load_config_from_s3(bucket_name, file_key): # Convert the configuration to a dictionary config_dict = {} for section in config_parser.sections(): - config_dict[section] = {key.upper(): value for key, value in config_parser.items(section)} + config_dict[section] = { + key.upper(): value for key, value in config_parser.items(section) + } return config_dict @@ -48,12 +51,17 @@ def lambda_handler(event, context): logger.info("Lambda function started processing.") # Read environment variables - property_file_path = os.environ.get('PROPERTY_FILE_S3_PATH', '') + property_file_path = os.environ.get("PROPERTY_FILE_S3_PATH", "") global DATABASE_LOGGING_LAMBDA_FUNCTION_NAME - DATABASE_LOGGING_LAMBDA_FUNCTION_NAME = os.environ.get('DATABASE_LOGGING_LAMBDA_FUNCTION_NAME', '') + DATABASE_LOGGING_LAMBDA_FUNCTION_NAME = os.environ.get( + "DATABASE_LOGGING_LAMBDA_FUNCTION_NAME", "" + ) batch_id = "" - - logger.info('DATABASE_LOGGING_LAMBDA_FUNCTION_NAME: ' + str(DATABASE_LOGGING_LAMBDA_FUNCTION_NAME)) + + logger.info( + "DATABASE_LOGGING_LAMBDA_FUNCTION_NAME: " + + str(DATABASE_LOGGING_LAMBDA_FUNCTION_NAME) + ) # Read config.properties file_path_array = property_file_path.split("/") @@ -64,105 +72,152 @@ def lambda_handler(event, context): # Extract BUCKET_NAME and config_file_path S3_BUCKET_NAME = file_path_array[0] CONFIG_FILE_PATH = "/".join(file_path_array[1:]) - - logger.info(f"Using S3 bucket: {S3_BUCKET_NAME} and config file path: {CONFIG_FILE_PATH}") - + + logger.info( + f"Using S3 bucket: {S3_BUCKET_NAME} and config file path: {CONFIG_FILE_PATH}" + ) # Load config file config_dict = load_config_from_s3(S3_BUCKET_NAME, CONFIG_FILE_PATH) logger.info("Config file loaded successfully.") # Process each message from the SQS event - for record in event['Records']: + for record in event["Records"]: # Extract the message body from the record - record_body = json.loads(record['body']) + record_body = json.loads(record["body"]) - for sqs_record in record_body['Records']: + for sqs_record in record_body["Records"]: # decode source path - s3_key = unquote_plus(sqs_record['s3']['object']['key']) + s3_key = unquote_plus(sqs_record["s3"]["object"]["key"]) logging.info("SOURCE_PATH: {source_path} ") - + # Read tags from staging file tags_dict = get_s3_object_tags(S3_BUCKET_NAME, s3_key) - + # Encode to URL Query parameter tags = urlencode(tags_dict) # check if key exist in tags_dict - + if "BatchId" in tags_dict: - batch_id = tags_dict['BatchId'] + batch_id = tags_dict["BatchId"] else: logger.error(f"BatchId not found in file tag") return - + logger.info(f"Batch ID: {batch_id}") - - logger.info(f"Processing file: {s3_key} in bucket: {S3_BUCKET_NAME}") + + logger.info( + f"Processing file: {s3_key} in bucket: {S3_BUCKET_NAME}" + ) # Read the Textract JSON file from S3 response = s3_client.get_object(Bucket=S3_BUCKET_NAME, Key=s3_key) - textract_json = json.loads(response['Body'].read().decode('utf-8')) + textract_json = json.loads(response["Body"].read().decode("utf-8")) logger.info("Textract JSON file loaded successfully.") # Extract text from the Textract response and add page numbers - extracted_text = '' + extracted_text = "" page_number = 1 - index_text = '' - output_text = 'Document Index\n' + index_text = "" + output_text = "Document Index\n" logger.info("Starting text extraction from Textract response.") # Create a map of all blocks and their IDs - block_id_map = {block['Id']: block for block in textract_json['Blocks']} + block_id_map = { + block["Id"]: block for block in textract_json["Blocks"] + } - target_block_types = ['LAYOUT_TITLE', 'LAYOUT_HEADER', 'LAYOUT_SECTION_HEADER'] + target_block_types = [ + "LAYOUT_TITLE", + "LAYOUT_HEADER", + "LAYOUT_SECTION_HEADER", + ] # Extract text from specific block types for document index - for block_id, block in enumerate(textract_json['Blocks']): - if block['BlockType'] in target_block_types: + for block_id, block in enumerate(textract_json["Blocks"]): + if block["BlockType"] in target_block_types: child_ids = get_child_ids(textract_json, block_id) - index_text = concatenate_text_from_child_ids(child_ids, block_id_map, block['BlockType']) + index_text = concatenate_text_from_child_ids( + child_ids, block_id_map, block["BlockType"] + ) # Append the concatenated text to the output - output_text += index_text + '\n' + output_text += index_text + "\n" global PROCESSED_LOCATION - PROCESSED_LOCATION = config_dict['FOLDER_LOCATIONS']['PROCESSED_LOCATION'].format(batch_id) - CONTRACT_TEXT_FILE_LOCATION = config_dict['FOLDER_LOCATIONS']['CONTRACT_TEXT_FILE_LOCATION'].format(batch_id) - OUTPUT_LOCATION = config_dict['FOLDER_LOCATIONS']['OUTPUT_LOCATION'].format(batch_id) + PROCESSED_LOCATION = config_dict["FOLDER_LOCATIONS"][ + "PROCESSED_LOCATION" + ].format(batch_id) + CONTRACT_TEXT_FILE_LOCATION = config_dict["FOLDER_LOCATIONS"][ + "CONTRACT_TEXT_FILE_LOCATION" + ].format(batch_id) + OUTPUT_LOCATION = config_dict["FOLDER_LOCATIONS"][ + "OUTPUT_LOCATION" + ].format(batch_id) global TEXTRACT_SIGNATURE_JSON - TEXTRACT_SIGNATURE_JSON = config_dict['FOLDER_LOCATIONS']['TEXTRACT_SIGNATURE_JSON'].format(batch_id) + TEXTRACT_SIGNATURE_JSON = config_dict["FOLDER_LOCATIONS"][ + "TEXTRACT_SIGNATURE_JSON" + ].format(batch_id) global TEXTRACT_TABLE_FORM_JSON - TEXTRACT_TABLE_FORM_JSON = config_dict['FOLDER_LOCATIONS']['TEXTRACT_TABLE_FORM_JSON'].format(batch_id) + TEXTRACT_TABLE_FORM_JSON = config_dict["FOLDER_LOCATIONS"][ + "TEXTRACT_TABLE_FORM_JSON" + ].format(batch_id) logger.info(f"PROCESSED_LOCATION : {PROCESSED_LOCATION}") - logger.info(f"CONTRACT_TEXT_FILE_LOCATION : {CONTRACT_TEXT_FILE_LOCATION}") + logger.info( + f"CONTRACT_TEXT_FILE_LOCATION : {CONTRACT_TEXT_FILE_LOCATION}" + ) logger.info(f"OUTPUT_LOCATION : {OUTPUT_LOCATION}") logger.info(f"TEXTRACT_SIGNATURE_JSON : {TEXTRACT_SIGNATURE_JSON}") - logger.info(f"TEXTRACT_TABLE_FORM_JSON : {TEXTRACT_TABLE_FORM_JSON}") + logger.info( + f"TEXTRACT_TABLE_FORM_JSON : {TEXTRACT_TABLE_FORM_JSON}" + ) # Extract sub folder with filename - sub_folder_with_filename = s3_key.replace(OUTPUT_LOCATION,"") + sub_folder_with_filename = s3_key.replace(OUTPUT_LOCATION, "") - logger.info(f"sub_folder_with_filename : {sub_folder_with_filename}") + logger.info( + f"sub_folder_with_filename : {sub_folder_with_filename}" + ) - pdf_file_name_with_path = PROCESSED_LOCATION + sub_folder_with_filename.replace('.json', '.pdf') + pdf_file_name_with_path = ( + PROCESSED_LOCATION + + sub_folder_with_filename.replace(".json", ".pdf") + ) # Download the PDF file from S3 - temp_file_path = '/tmp/'+get_filename_from_path(pdf_file_name_with_path) - s3_client.download_file(S3_BUCKET_NAME, pdf_file_name_with_path, temp_file_path) - logging.info(f"PDF file downloaded from S3: s3://{S3_BUCKET_NAME}/{pdf_file_name_with_path}") + temp_file_path = "/tmp/" + get_filename_from_path( + pdf_file_name_with_path + ) + s3_client.download_file( + S3_BUCKET_NAME, pdf_file_name_with_path, temp_file_path + ) + logging.info( + f"PDF file downloaded from S3: s3://{S3_BUCKET_NAME}/{pdf_file_name_with_path}" + ) # Extract text from PAGE block type and concatenate with page numbers - for block_id, block in enumerate(textract_json['Blocks']): - if block['BlockType'] == 'PAGE': - logger.info(f"Processing page {page_number}... Block ID: {block_id}, block type: {block['BlockType']}") - extracted_text += f'\n\nStart of Page No. = {page_number}\n' + for block_id, block in enumerate(textract_json["Blocks"]): + if block["BlockType"] == "PAGE": + logger.info( + f"Processing page {page_number}... Block ID: {block_id}, block type: {block['BlockType']}" + ) + extracted_text += f"\n\nStart of Page No. = {page_number}\n" child_ids = get_child_ids(textract_json, block_id) - extracted_text += concatenate_page_text_from_child_ids(child_ids, block_id_map,pdf_file_name_with_path,page_number,S3_BUCKET_NAME,final_number_of_signature,temp_file_path, tags) + extracted_text += concatenate_page_text_from_child_ids( + child_ids, + block_id_map, + pdf_file_name_with_path, + page_number, + S3_BUCKET_NAME, + final_number_of_signature, + temp_file_path, + tags, + ) page_number += 1 logger.info("Text extraction completed.") @@ -170,153 +225,198 @@ def lambda_handler(event, context): # Remove the temporary files os.remove(temp_file_path) logging.info(f"Temporary files removed.") - - output_text += '\n' - final_text_file_content = output_text + extracted_text - - logger.info(f'File path == {s3_key}') + + output_text += "\n" + final_text_file_content = output_text + extracted_text + + logger.info(f"File path == {s3_key}") # s3_key = get_filename_from_path(s3_key) # Upload the text file to S3 or perform further processing as needed - output_key = CONTRACT_TEXT_FILE_LOCATION + f"{sub_folder_with_filename.replace('.json', '.txt')}" # Example output key - - upload_text_file_to_s3(final_text_file_content, S3_BUCKET_NAME, output_key, s3_client, tags) - + output_key = ( + CONTRACT_TEXT_FILE_LOCATION + + f"{sub_folder_with_filename.replace('.json', '.txt')}" + ) # Example output key + + upload_text_file_to_s3( + final_text_file_content, + S3_BUCKET_NAME, + output_key, + s3_client, + tags, + ) # Log success message - logger.info(f'Text extracted and saved to S3://{S3_BUCKET_NAME}/{output_key}') + logger.info( + f"Text extracted and saved to S3://{S3_BUCKET_NAME}/{output_key}" + ) document_id = get_filename_from_path(output_key) document_id = os.path.splitext(document_id)[0] # Return the S3 URL or any other relevant information - generate_document_logs_input(document_id,True,False) + generate_document_logs_input(document_id, True, False) return { - 'statusCode': 200, - 'body': json.dumps(f'Text extracted and saved to S3://{S3_BUCKET_NAME}/{output_key}') + "statusCode": 200, + "body": json.dumps( + f"Text extracted and saved to S3://{S3_BUCKET_NAME}/{output_key}" + ), } except Exception as e: # Log error message - logger.error(f'Error processing Textract output: {e}') + logger.error(f"Error processing Textract output: {e}") # Return error response return { - 'statusCode': 500, - 'body': json.dumps(f'Error processing Textract output: {e}') + "statusCode": 500, + "body": json.dumps(f"Error processing Textract output: {e}"), } - + def get_child_ids(textract_json, block_id): child_ids = [] # Check if the block_id is within the valid range - block = textract_json['Blocks'][block_id] - #logger.info(f"Processing block: {block}") + block = textract_json["Blocks"][block_id] + # logger.info(f"Processing block: {block}") # Check if the block has 'Relationships' key - if 'Relationships' in block: + if "Relationships" in block: logger.info(f"Block has relationships: {block['Relationships']}") - for relationship in block['Relationships']: - if relationship['Type'] == 'CHILD': - child_ids.extend(relationship['Ids']) - + for relationship in block["Relationships"]: + if relationship["Type"] == "CHILD": + child_ids.extend(relationship["Ids"]) + return child_ids -def concatenate_text_from_child_ids(child_ids, block_id_map,block_type): - concatenated_text = '' +def concatenate_text_from_child_ids(child_ids, block_id_map, block_type): + concatenated_text = "" for child_id in child_ids: logger.debug(f"Processing child ID: {child_id}") - + # Use the block_id_map to access the block using the ID directly block = block_id_map.get(child_id) - - if block and 'Text' in block: - if block_type == 'LAYOUT_TABLE': - concatenated_text += block['Text']+'\n' - else : - concatenated_text += block['Text'] + ' ' + str(block['Page']) + + if block and "Text" in block: + if block_type == "LAYOUT_TABLE": + concatenated_text += block["Text"] + "\n" + else: + concatenated_text += block["Text"] + " " + str(block["Page"]) logger.debug(f"Concatenated text: {concatenated_text}") return concatenated_text.strip() - -def concatenate_page_text_from_child_ids(child_ids, block_id_map,pdf_file_name_with_path,page_number,s3_bucket,final_number_of_signature,temp_file_path,tags): - - concatenated_text = '' - table_text = '' + +def concatenate_page_text_from_child_ids( + child_ids, + block_id_map, + pdf_file_name_with_path, + page_number, + s3_bucket, + final_number_of_signature, + temp_file_path, + tags, +): + + concatenated_text = "" + table_text = "" table_found = False # Flag to indicate if a table block is found for child_id in child_ids: # Adding extensive logging to ensure the hierarchy of blocks is correct and if we are catching all elements of the document - logger.debug(f"Processing child ID: {child_id}, block type: {block_id_map[child_id]['BlockType']}") - + logger.debug( + f"Processing child ID: {child_id}, block type: {block_id_map[child_id]['BlockType']}" + ) + # Use the block_id_map to access the block using the ID directly block = block_id_map.get(child_id) - - if block and 'Text' in block: - concatenated_text += block['Text']+'\n' - elif block['BlockType'] == 'LAYOUT_TABLE': - concatenated_text += '\n'+'-------TABLE Start-----'+'\n' + if block and "Text" in block: + concatenated_text += block["Text"] + "\n" + + elif block["BlockType"] == "LAYOUT_TABLE": + concatenated_text += "\n" + "-------TABLE Start-----" + "\n" table_child_ids = get_child_table_ids(block) - table_text = concatenate_text_from_child_ids(table_child_ids, block_id_map,block['BlockType']) - table_text += '\n'+'-------TABLE End-----'+'\n' + table_text = concatenate_text_from_child_ids( + table_child_ids, block_id_map, block["BlockType"] + ) + table_text += "\n" + "-------TABLE End-----" + "\n" table_found = True - concatenated_text +=table_text - - + concatenated_text += table_text if table_found: - - textract_table_response = extract_page_and_send_to_textract(page_number-1, s3_bucket, pdf_file_name_with_path,temp_file_path,tags,feature_types=['TABLES']) - + + textract_table_response = extract_page_and_send_to_textract( + page_number - 1, + s3_bucket, + pdf_file_name_with_path, + temp_file_path, + tags, + feature_types=["TABLES"], + ) + if textract_table_response is not None: child_ids_list = [] - table_text = get_table_text(textract_table_response,child_ids_list) + table_text = get_table_text(textract_table_response, child_ids_list) table_child_ids_list = [] - table_child_ids_list = list_line_ids_from_table_blocks(textract_table_response) - lines_text = get_lines(textract_table_response,table_child_ids_list) - - + table_child_ids_list = list_line_ids_from_table_blocks( + textract_table_response + ) + lines_text = get_lines(textract_table_response, table_child_ids_list) + concatenated_text = f""" {lines_text} {table_text} """ - if final_number_of_signature==0: - signature_page_word_list = ["IN WITNESS WHEREOF","Signature"] - matched_count = strings_contained(signature_page_word_list,concatenated_text) - - if matched_count>1: + if final_number_of_signature == 0: + signature_page_word_list = ["IN WITNESS WHEREOF", "Signature"] + matched_count = strings_contained(signature_page_word_list, concatenated_text) + + if matched_count > 1: logging.info(f"matched_count = {matched_count}") - textract_signature_response = extract_page_and_send_to_textract(page_number-1, s3_bucket, pdf_file_name_with_path,temp_file_path,tags,feature_types=['SIGNATURES']) + textract_signature_response = extract_page_and_send_to_textract( + page_number - 1, + s3_bucket, + pdf_file_name_with_path, + temp_file_path, + tags, + feature_types=["SIGNATURES"], + ) logging.info(f"textract_signature_response = {textract_signature_response}") - final_number_of_signature = get_number_of_signatures(textract_signature_response) + final_number_of_signature = get_number_of_signatures( + textract_signature_response + ) logging.info(f"Number of signature found {final_number_of_signature}") - concatenated_text = concatenated_text + f"\nThis page has {final_number_of_signature} signature." - - + concatenated_text = ( + concatenated_text + + f"\nThis page has {final_number_of_signature} signature." + ) + return concatenated_text.strip() def get_number_of_signatures(json_data): blocks = json_data["Blocks"] signatures = map_blocks(blocks, "SIGNATURE") - + signature_count = 0 - + for index, signature in enumerate(signatures.values()): signature_count += 1 - + return signature_count def strings_contained(substrings, main_string): return sum(1 for substring in substrings if substring in main_string) + # Function to extract a specific page from a PDF file and send it to Textract -def extract_page_and_send_to_textract(page_number, s3_bucket, s3_key, temp_file_path, tags, feature_types=None): +def extract_page_and_send_to_textract( + page_number, s3_bucket, s3_key, temp_file_path, tags, feature_types=None +): try: # Read the PDF file and check if it's encrypted - with open(temp_file_path, 'rb') as file: + with open(temp_file_path, "rb") as file: pdf_reader = PdfReader(file) if pdf_reader.is_encrypted: logging.info("File is encrypted. Trying to decrypt.") @@ -336,43 +436,62 @@ def extract_page_and_send_to_textract(page_number, s3_bucket, s3_key, temp_file_ temp_page_file_path = temp_page_file.name page_writer = PdfWriter() page_writer.add_page(page) - with open(temp_page_file_path, 'wb') as temp_page_file: + with open(temp_page_file_path, "wb") as temp_page_file: page_writer.write(temp_page_file) - logging.info(f"Page {page_number} extracted and stored temporarily at {temp_page_file_path}") + logging.info( + f"Page {page_number} extracted and stored temporarily at {temp_page_file_path}" + ) # Call AWS Textract synchronous API to analyze the extracted page - with open(temp_page_file_path, 'rb') as page_file: - textract_response = textract_client.analyze_document(Document={'Bytes': page_file.read()}, - FeatureTypes=feature_types) + with open(temp_page_file_path, "rb") as page_file: + textract_response = textract_client.analyze_document( + Document={"Bytes": page_file.read()}, FeatureTypes=feature_types + ) logging.info(f"Page {page_number} analyzed by Textract.") # Extract sub folder with filename sub_folder_with_filename = s3_key.replace(PROCESSED_LOCATION, "") # Store the JSON response in S3 - page_json_key = sub_folder_with_filename.replace('.pdf', '') + f"_page_{page_number + 1}.json" - logging.info(f"TEXTRACT_SIGNATURE_JSON: {TEXTRACT_SIGNATURE_JSON} TEXTRACT_TABLE_FORM_JSON: {TEXTRACT_TABLE_FORM_JSON}") + page_json_key = ( + sub_folder_with_filename.replace(".pdf", "") + + f"_page_{page_number + 1}.json" + ) + logging.info( + f"TEXTRACT_SIGNATURE_JSON: {TEXTRACT_SIGNATURE_JSON} TEXTRACT_TABLE_FORM_JSON: {TEXTRACT_TABLE_FORM_JSON}" + ) - if any('SIGNATURES' in feature for feature in feature_types): + if any("SIGNATURES" in feature for feature in feature_types): page_json_key = TEXTRACT_SIGNATURE_JSON + page_json_key else: page_json_key = TEXTRACT_TABLE_FORM_JSON + page_json_key - s3_client.put_object(Bucket=s3_bucket, Key=page_json_key, - Body=json.dumps(textract_response), ContentType='application/json', Tagging=tags) - logging.info(f"Textract JSON response for page {page_number} stored in S3 with key: {page_json_key}") + s3_client.put_object( + Bucket=s3_bucket, + Key=page_json_key, + Body=json.dumps(textract_response), + ContentType="application/json", + Tagging=tags, + ) + logging.info( + f"Textract JSON response for page {page_number} stored in S3 with key: {page_json_key}" + ) # Remove the temporary files os.remove(temp_page_file_path) logging.info(f"Temporary files removed.") - logging.info(f"Page {page_number} extraction and Textract processing completed successfully.") + logging.info( + f"Page {page_number} extraction and Textract processing completed successfully." + ) return textract_response except PdfReadError as e: logging.error(f"PdfReadError: {str(e)}. Unable to process page {page_number}.") return None except Exception as e: - logging.error(f"Error extracting page {page_number} and sending to Textract: {e}") + logging.error( + f"Error extracting page {page_number} and sending to Textract: {e}" + ) return None @@ -380,29 +499,31 @@ def get_child_table_ids(block): child_ids = [] # Check if the block has 'Relationships' key - if 'Relationships' in block: - for relationship in block['Relationships']: - if relationship['Type'] == 'CHILD': - child_ids.extend(relationship['Ids']) - + if "Relationships" in block: + for relationship in block["Relationships"]: + if relationship["Type"] == "CHILD": + child_ids.extend(relationship["Ids"]) + return child_ids # Function to save text to S3 -def upload_text_file_to_s3(final_text_file_content, bucket_name, object_key, s3_client, tags): - +def upload_text_file_to_s3( + final_text_file_content, bucket_name, object_key, s3_client, tags +): + try: # Upload the JSON response to S3 s3_client.put_object( Bucket=bucket_name, Key=object_key, Body=final_text_file_content, - ContentType='text/plain', - Tagging=tags + ContentType="text/plain", + Tagging=tags, ) logger.info(f"Saved analysis response to S3: {object_key}") except Exception as e: - logger.error(f"Error saving response to S3 - {object_key}: {str(e)}") + logger.error(f"Error saving response to S3 - {object_key}: {str(e)}") # Function to get s3 object tags @@ -410,14 +531,11 @@ def get_s3_object_tags(bucket_name, object_key): try: # Get object tags - response = s3_client.get_object_tagging( - Bucket=bucket_name, - Key=object_key - ) + response = s3_client.get_object_tagging(Bucket=bucket_name, Key=object_key) # Extract tags from the response and convert to dictionary - tags_list = response['TagSet'] - tags_dict = {tag['Key']: tag['Value'] for tag in tags_list} + tags_list = response["TagSet"] + tags_dict = {tag["Key"]: tag["Value"] for tag in tags_list} return tags_dict @@ -425,17 +543,17 @@ def get_s3_object_tags(bucket_name, object_key): logger.exception(f"Error: {e}") logger.error(f"Error: {e}") return None - + + def get_filename_from_path(full_path): return os.path.basename(full_path) - def map_blocks(blocks, block_type): return {block["Id"]: block for block in blocks if block["BlockType"] == block_type} -def get_lines(json_data,child_ids_list): +def get_lines(json_data, child_ids_list): blocks = json_data["Blocks"] lines = map_blocks(blocks, "LINE") @@ -481,17 +599,17 @@ def get_table_text(json_data, child_ids_list): cell = cells.get(cell_id, {}) # Check if the cell has the 'Text' key - if 'Text' in cell: - extracted_text.append(cell['Text']) + if "Text" in cell: + extracted_text.append(cell["Text"]) else: # If 'Text' key is not present, look for child words - if 'Relationships' in cell: - for rel in cell['Relationships']: - if rel['Type'] == 'CHILD': - for child_id in rel['Ids']: + if "Relationships" in cell: + for rel in cell["Relationships"]: + if rel["Type"] == "CHILD": + for child_id in rel["Ids"]: if child_id in words: extracted_text.append(words[child_id]["Text"]) - return ' '.join(extracted_text) + return " ".join(extracted_text) def get_table_title_children_ids(block): logger.info(f"Finding table title: {block}") @@ -511,7 +629,10 @@ def get_table_text(json_data, child_ids_list): if table_title_id != "": for block2 in json_data.get("Blocks", []): - if block2["BlockType"] == "TABLE_TITLE" and block2["Id"] == table_title_id: + if ( + block2["BlockType"] == "TABLE_TITLE" + and block2["Id"] == table_title_id + ): if block2["Relationships"][0]["Type"] == "CHILD": child_ids_list.extend(block2["Relationships"][0]["Ids"]) @@ -530,7 +651,9 @@ def get_table_text(json_data, child_ids_list): def ensure_content_size(content, required_rows, required_cols): # Expand rows if needed while len(content) < required_rows: - content.append([None] * len(content[0])) # Add new rows with the same number of columns + content.append( + [None] * len(content[0]) + ) # Add new rows with the same number of columns # Expand columns if needed for row in content: @@ -540,11 +663,18 @@ def get_table_text(json_data, child_ids_list): dataframe_dicts = [] for index, table in enumerate(tables.values()): # Get all the cells belonging to this table - table_cells = [cells[cell_id] for cell_id in get_children_ids(table) if cell_id in cells] + table_cells = [ + cells[cell_id] + for cell_id in get_children_ids(table) + if cell_id in cells + ] # Get all the merged cells belonging to this table - table_merged_cells = [merged_cells[cell_id] for cell_id in get_children_ids(table, "MERGED_CELL") if - cell_id in merged_cells] + table_merged_cells = [ + merged_cells[cell_id] + for cell_id in get_children_ids(table, "MERGED_CELL") + if cell_id in merged_cells + ] # Determine the table's number of rows and columns if table_cells: # Check if table_cells is not empty n_rows = max(cell["RowIndex"] for cell in table_cells) @@ -574,9 +704,9 @@ def get_table_text(json_data, child_ids_list): # Dynamically adjust columns if j >= len(content[i]): content[i].extend([None] * (j - len(content[i]) + 1)) - content[i][j] = ' '.join(cell_contents) + content[i][j] = " ".join(cell_contents) else: - content = '' + content = "" # Handle merged cells for merged_cell in table_merged_cells: @@ -587,14 +717,16 @@ def get_table_text(json_data, child_ids_list): # Concatenate text from all words associated with the merged cell merged_text = [] - if 'Relationships' in merged_cell: - for rel in merged_cell['Relationships']: - if rel['Type'] == 'CHILD': - merged_text = get_text_from_cell_blocks(cells, rel['Ids']) + if "Relationships" in merged_cell: + for rel in merged_cell["Relationships"]: + if rel["Type"] == "CHILD": + merged_text = get_text_from_cell_blocks(cells, rel["Ids"]) for row_offset in range(row_span): for col_offset in range(col_span): - content[row_start + row_offset][col_start + col_offset] = merged_text + content[row_start + row_offset][ + col_start + col_offset + ] = merged_text logger.info(f"Table content = {content}") dataframe_dicts.append("-------Table Start--------") @@ -603,68 +735,95 @@ def get_table_text(json_data, child_ids_list): dataframe_dicts.append(get_table_footer_children_ids(table)) dataframe_dicts.append("-------Table End--------") - dataframe_dicts_str = "\n".join(str(dataframe_dict) for dataframe_dict in dataframe_dicts) + dataframe_dicts_str = "\n".join( + str(dataframe_dict) for dataframe_dict in dataframe_dicts + ) return dataframe_dicts_str except Exception as e: logger.error(f"Error in get_table_text: {e}") print(traceback.format_exc()) return None - + def list_line_ids_from_table_blocks(textract_response): line_ids = [] - - - blocks = textract_response.get('Blocks') # Use get method to handle potential NoneType - + + blocks = textract_response.get( + "Blocks" + ) # Use get method to handle potential NoneType + if blocks is not None: for block in blocks: - if block['BlockType'] == 'TABLE': - for relationship in block.get('Relationships', []): - if relationship['Type'] == 'CHILD': - for child_id in relationship.get('Ids', []): + if block["BlockType"] == "TABLE": + for relationship in block.get("Relationships", []): + if relationship["Type"] == "CHILD": + for child_id in relationship.get("Ids", []): for sub_block in blocks: - if sub_block['Id'] == child_id and sub_block['BlockType'] == 'CELL': - for sub_relationship in sub_block.get('Relationships', []): - if sub_relationship['Type'] == 'CHILD': - for word_id in sub_relationship.get('Ids', []): + if ( + sub_block["Id"] == child_id + and sub_block["BlockType"] == "CELL" + ): + for sub_relationship in sub_block.get( + "Relationships", [] + ): + if sub_relationship["Type"] == "CHILD": + for word_id in sub_relationship.get( + "Ids", [] + ): for line_block in blocks: - if line_block['BlockType'] == 'LINE': - for line_relationship in line_block.get('Relationships', []): - if line_relationship['Type'] == 'CHILD' and word_id in line_relationship['Ids']: - line_ids.append(line_block['Id']) - + if ( + line_block["BlockType"] + == "LINE" + ): + for ( + line_relationship + ) in line_block.get( + "Relationships", [] + ): + if ( + line_relationship[ + "Type" + ] + == "CHILD" + and word_id + in line_relationship[ + "Ids" + ] + ): + line_ids.append( + line_block["Id"] + ) + else: logger.warning("No blocks found in the Textract response.") - - + return line_ids -def generate_document_logs_input(document_id,payer_signed,provider_signed): - current_time = datetime.datetime.now().isoformat() - - data = { +def generate_document_logs_input(document_id, payer_signed, provider_signed): + current_time = datetime.datetime.now().isoformat() + + data = { "operation": "update", "table": "DOCUMENT_LOGS", "data": { "DOCUMENT_ID": document_id, - "STAGE" : "TEXT_EXTRACTION", + "STAGE": "TEXT_EXTRACTION", "PAYER_SIGNED": payer_signed, "PROVIDER_SIGNED": provider_signed, "MODIFIED_TIME": current_time, - "MODIFIED_BY": "TEXTRACT_OUTPUT_TEXT" - } - } - logger.info(f"Request: {data}") - - lambda_client = boto3.client('lambda') - response = lambda_client.invoke( - FunctionName=DATABASE_LOGGING_LAMBDA_FUNCTION_NAME, - InvocationType='Event', # Asynchronous invocation - Payload=json.dumps(data).encode('utf-8') - ) - - logger.info("Generated document logs input successfully") - return response \ No newline at end of file + "MODIFIED_BY": "TEXTRACT_OUTPUT_TEXT", + }, + } + logger.info(f"Request: {data}") + + lambda_client = boto3.client("lambda") + response = lambda_client.invoke( + FunctionName=DATABASE_LOGGING_LAMBDA_FUNCTION_NAME, + InvocationType="Event", # Asynchronous invocation + Payload=json.dumps(data).encode("utf-8"), + ) + + logger.info("Generated document logs input successfully") + return response diff --git a/textract-pipeline/src/lambda/textract-receiver/index.py b/textract-pipeline/src/lambda/textract-receiver/index.py index 54aca8e..4ac7b02 100644 --- a/textract-pipeline/src/lambda/textract-receiver/index.py +++ b/textract-pipeline/src/lambda/textract-receiver/index.py @@ -13,8 +13,8 @@ logger = logging.getLogger(__name__) logger.setLevel(logging.INFO) # Initialize S3 & Textract clients -s3_client = boto3.client('s3') -textract_client = boto3.client('textract') +s3_client = boto3.client("s3") +textract_client = boto3.client("textract") DATABASE_LOGGING_LAMBDA_FUNCTION_NAME = "" @@ -22,7 +22,7 @@ DATABASE_LOGGING_LAMBDA_FUNCTION_NAME = "" def load_config_from_s3(bucket_name, file_key): # Download the config file from S3 response = s3_client.get_object(Bucket=bucket_name, Key=file_key) - config_content = response['Body'].read().decode('utf-8') + config_content = response["Body"].read().decode("utf-8") # Parse the config file config_parser = ConfigParser() @@ -31,7 +31,9 @@ def load_config_from_s3(bucket_name, file_key): # Convert the configuration to a dictionary config_dict = {} for section in config_parser.sections(): - config_dict[section] = {key.upper(): value for key, value in config_parser.items(section)} + config_dict[section] = { + key.upper(): value for key, value in config_parser.items(section) + } return config_dict @@ -50,20 +52,28 @@ def generate_output_json_path(src_folder, dest_folder, s3_object): # Exponential backoff utility function -def exponential_backoff_retry(func, *args, retries=5, initial_delay=1, max_delay=32, **kwargs): +def exponential_backoff_retry( + func, *args, retries=5, initial_delay=1, max_delay=32, **kwargs +): delay = initial_delay for attempt in range(retries): try: return func(*args, **kwargs) except ClientError as e: - error_code = e.response['Error']['Code'] - if error_code == 'ProvisionedThroughputExceededException' and attempt < retries - 1: + error_code = e.response["Error"]["Code"] + if ( + error_code == "ProvisionedThroughputExceededException" + and attempt < retries - 1 + ): logger.warning( - f"ProvisionedThroughputExceededException on attempt {attempt + 1}/{retries}, retrying in {delay} seconds...") + f"ProvisionedThroughputExceededException on attempt {attempt + 1}/{retries}, retrying in {delay} seconds..." + ) time.sleep(delay) delay = min(delay * 2, max_delay) else: - logger.error(f"Exceeded maximum retries. Function {func.__name__} failed with error: {e}") + logger.error( + f"Exceeded maximum retries. Function {func.__name__} failed with error: {e}" + ) raise @@ -78,20 +88,25 @@ def get_textract_document_detection(job_id, textract_client): try: while True: if flag: - response = exponential_backoff_retry(textract_client.get_document_text_detection, JobId=job_id) + response = exponential_backoff_retry( + textract_client.get_document_text_detection, JobId=job_id + ) flag = False else: - response = exponential_backoff_retry(textract_client.get_document_text_detection, JobId=job_id, - NextToken=next_token) + response = exponential_backoff_retry( + textract_client.get_document_text_detection, + JobId=job_id, + NextToken=next_token, + ) job_status = response["JobStatus"] logger.info("Job %s status is %s.", job_id, job_status) # Merge the blocks from the current response - all_blocks.extend(response.get('Blocks', [])) + all_blocks.extend(response.get("Blocks", [])) # Check if there are more blocks to retrieve - next_token = response.get('NextToken') + next_token = response.get("NextToken") if not next_token: logger.info("No more Textract response to retrieve") break @@ -102,12 +117,12 @@ def get_textract_document_detection(job_id, textract_client): else: # Remove unnecessary keys from the last response last_response = response.copy() - last_response.pop('Blocks', None) - last_response.pop('ResponseMetadata', None) + last_response.pop("Blocks", None) + last_response.pop("ResponseMetadata", None) # logger.info("Removed 'ResponseMetadata' key") # Merge with {'Blocks': all_blocks} - final_response = {'Blocks': all_blocks} + final_response = {"Blocks": all_blocks} final_response.update(last_response) logger.info("Final Textract response is constructed") return final_response @@ -124,20 +139,25 @@ def get_textract_document_analysis(job_id, textract_client): try: while True: if flag: - response = exponential_backoff_retry(textract_client.get_document_analysis, JobId=job_id) + response = exponential_backoff_retry( + textract_client.get_document_analysis, JobId=job_id + ) flag = False else: - response = exponential_backoff_retry(textract_client.get_document_analysis, JobId=job_id, - NextToken=next_token) + response = exponential_backoff_retry( + textract_client.get_document_analysis, + JobId=job_id, + NextToken=next_token, + ) job_status = response["JobStatus"] logger.info("Job %s status is %s.", job_id, job_status) # Merge the blocks from the current response - all_blocks.extend(response.get('Blocks', [])) + all_blocks.extend(response.get("Blocks", [])) # Check if there are more blocks to retrieve - next_token = response.get('NextToken') + next_token = response.get("NextToken") if not next_token: logger.info("No more Textract response to retrieve") break @@ -148,12 +168,12 @@ def get_textract_document_analysis(job_id, textract_client): else: # Remove unnecessary keys from the last response last_response = response.copy() - last_response.pop('Blocks', None) - last_response.pop('ResponseMetadata', None) + last_response.pop("Blocks", None) + last_response.pop("ResponseMetadata", None) # logger.info("Removed 'ResponseMetadata' key") # Merge with {'Blocks': all_blocks} - final_response = {'Blocks': all_blocks} + final_response = {"Blocks": all_blocks} final_response.update(last_response) logger.info("Final Textract response is constructed") return final_response @@ -170,8 +190,8 @@ def upload_response_to_s3(response, bucket_name, object_key, s3_client, tags): Bucket=bucket_name, Key=object_key, Body=response_json, - ContentType='application/json', - Tagging=tags + ContentType="application/json", + Tagging=tags, ) logger.info(f"Saved analysis response to S3: {object_key}") except Exception as e: @@ -183,8 +203,11 @@ def upload_response_to_s3(response, bucket_name, object_key, s3_client, tags): def move_file_within_s3(source_bucket, source_path, destination_path): try: # Copy the file to the destination folder - s3_client.copy_object(Bucket=source_bucket, CopySource={'Bucket': source_bucket, 'Key': source_path}, - Key=destination_path) + s3_client.copy_object( + Bucket=source_bucket, + CopySource={"Bucket": source_bucket, "Key": source_path}, + Key=destination_path, + ) # Delete the file from the source folder s3_client.delete_object(Bucket=source_bucket, Key=source_path) @@ -198,14 +221,11 @@ def move_file_within_s3(source_bucket, source_path, destination_path): def get_s3_object_tags(bucket_name, object_key): try: # Get object tags - response = s3_client.get_object_tagging( - Bucket=bucket_name, - Key=object_key - ) + response = s3_client.get_object_tagging(Bucket=bucket_name, Key=object_key) # Extract tags from the response and convert to dictionary - tags_list = response['TagSet'] - tags_dict = {tag['Key']: tag['Value'] for tag in tags_list} + tags_list = response["TagSet"] + tags_dict = {tag["Key"]: tag["Value"] for tag in tags_list} return tags_dict @@ -221,12 +241,17 @@ def lambda_handler(event, context): print(event) # Read environment variables - property_file_path = os.environ.get('PROPERTY_FILE_S3_PATH', '') + property_file_path = os.environ.get("PROPERTY_FILE_S3_PATH", "") global DATABASE_LOGGING_LAMBDA_FUNCTION_NAME - DATABASE_LOGGING_LAMBDA_FUNCTION_NAME = os.environ.get('DATABASE_LOGGING_LAMBDA_FUNCTION_NAME', '') + DATABASE_LOGGING_LAMBDA_FUNCTION_NAME = os.environ.get( + "DATABASE_LOGGING_LAMBDA_FUNCTION_NAME", "" + ) batch_id = "" - logger.info('DATABASE_LOGGING_LAMBDA_FUNCTION_NAME: ' + str(DATABASE_LOGGING_LAMBDA_FUNCTION_NAME)) + logger.info( + "DATABASE_LOGGING_LAMBDA_FUNCTION_NAME: " + + str(DATABASE_LOGGING_LAMBDA_FUNCTION_NAME) + ) # Validate environment variables file_path_array = property_file_path.split("/") @@ -235,8 +260,8 @@ def lambda_handler(event, context): S3_BUCKET_NAME = file_path_array[0] CONFIG_FILE_PATH = "/".join(file_path_array[1:]) - logger.info(f'S3_BUCKET_NAME: {S3_BUCKET_NAME}') - logger.info(f'CONFIG_FILE_PATH: {CONFIG_FILE_PATH}') + logger.info(f"S3_BUCKET_NAME: {S3_BUCKET_NAME}") + logger.info(f"CONFIG_FILE_PATH: {CONFIG_FILE_PATH}") # Load config file config_dict = load_config_from_s3(S3_BUCKET_NAME, CONFIG_FILE_PATH) @@ -244,100 +269,119 @@ def lambda_handler(event, context): # logger.info('## CONFIG DICTIONARY\r' + str(config_dict)) # Process each message from the SQS event - for record in event['Records']: + for record in event["Records"]: # Extract the message body from the record - record_body = json.loads(record['body']) - message_body = json.loads(record_body['Message']) - logger.info('MESSAGE_BODY: ' + str(message_body)) + record_body = json.loads(record["body"]) + message_body = json.loads(record_body["Message"]) + logger.info("MESSAGE_BODY: " + str(message_body)) try: # Extract relevant information from the message body - job_id = message_body.get('JobId') - document_location = message_body.get('DocumentLocation') - s3_object_name = document_location.get('S3ObjectName') + job_id = message_body.get("JobId") + document_location = message_body.get("DocumentLocation") + s3_object_name = document_location.get("S3ObjectName") # Read tags from staging file tags_dict = get_s3_object_tags(S3_BUCKET_NAME, s3_object_name) # check if key exist in tags_dict if "BatchId" in tags_dict: - batch_id = tags_dict['BatchId'] + batch_id = tags_dict["BatchId"] else: logger.error(f"BatchId not found in file tag") return - STAGING_LOCATION = config_dict['FOLDER_LOCATIONS']['STAGING_LOCATION'].format(batch_id) - OUTPUT_LOCATION = config_dict['FOLDER_LOCATIONS']['OUTPUT_LOCATION'].format(batch_id) - PROCESSED_LOCATION = config_dict['FOLDER_LOCATIONS']['PROCESSED_LOCATION'].format(batch_id) - UNPROCESSED_LOCATION = config_dict['FOLDER_LOCATIONS']['UNPROCESSED_LOCATION'].format(batch_id) - PROCESS_TYPE = str(config_dict['OTHERS']['PROCESS_TYPE']).upper() + STAGING_LOCATION = config_dict["FOLDER_LOCATIONS"][ + "STAGING_LOCATION" + ].format(batch_id) + OUTPUT_LOCATION = config_dict["FOLDER_LOCATIONS"][ + "OUTPUT_LOCATION" + ].format(batch_id) + PROCESSED_LOCATION = config_dict["FOLDER_LOCATIONS"][ + "PROCESSED_LOCATION" + ].format(batch_id) + UNPROCESSED_LOCATION = config_dict["FOLDER_LOCATIONS"][ + "UNPROCESSED_LOCATION" + ].format(batch_id) + PROCESS_TYPE = str(config_dict["OTHERS"]["PROCESS_TYPE"]).upper() # encode to URL Query parameter tags = urlencode(tags_dict) # Check if the status is "SUCCEEDED" - if message_body.get('Status') == 'SUCCEEDED': + if message_body.get("Status") == "SUCCEEDED": document = {} if PROCESS_TYPE == "ANALYSIS": # Call the function to get document analysis using Textract - document = get_textract_document_analysis(job_id, textract_client) + document = get_textract_document_analysis( + job_id, textract_client + ) elif PROCESS_TYPE == "DETECTION": # Call the function to get document analysis using Textract - document = get_textract_document_detection(job_id, textract_client) + document = get_textract_document_detection( + job_id, textract_client + ) # Save the document analysis response to S3 - s3_object_key = generate_output_json_path(STAGING_LOCATION, OUTPUT_LOCATION, s3_object_name) - upload_response_to_s3(document, S3_BUCKET_NAME, s3_object_key, s3_client, tags) + s3_object_key = generate_output_json_path( + STAGING_LOCATION, OUTPUT_LOCATION, s3_object_name + ) + upload_response_to_s3( + document, S3_BUCKET_NAME, s3_object_key, s3_client, tags + ) # Construct the destination paths - destination_path = PROCESSED_LOCATION + s3_object_name.replace(STAGING_LOCATION, "") + destination_path = PROCESSED_LOCATION + s3_object_name.replace( + STAGING_LOCATION, "" + ) # Move file to processed folder - move_file_within_s3(S3_BUCKET_NAME, s3_object_name, destination_path) + move_file_within_s3( + S3_BUCKET_NAME, s3_object_name, destination_path + ) document_id = get_filename_from_path(s3_object_name) document_id = os.path.splitext(document_id)[0] - generate_document_logs_input(document_id, message_body.get('Status')) - success_message = 'Processed file ' + str(s3_object_key) + generate_document_logs_input( + document_id, message_body.get("Status") + ) + success_message = "Processed file " + str(s3_object_key) logger.info(success_message) - return { - 'statusCode': 200, - 'body': success_message - } + return {"statusCode": 200, "body": success_message} else: error_message = f"Skipping message with JobId {message_body.get('JobId')} as Status is not 'SUCCEEDED'" logger.error(error_message) # Construct the destination paths - destination_path = UNPROCESSED_LOCATION + s3_object_name.replace(STAGING_LOCATION, "") + destination_path = UNPROCESSED_LOCATION + s3_object_name.replace( + STAGING_LOCATION, "" + ) # Move file to unprocessed folder - move_file_within_s3(S3_BUCKET_NAME, s3_object_name, destination_path) + move_file_within_s3( + S3_BUCKET_NAME, s3_object_name, destination_path + ) document_id = get_filename_from_path(s3_object_name) document_id = os.path.splitext(document_id)[0] - generate_document_logs_input(document_id, message_body.get('Status')) - return { - 'statusCode': 500, - 'body': error_message - } + generate_document_logs_input( + document_id, message_body.get("Status") + ) + return {"statusCode": 500, "body": error_message} except Exception as e: error_message = f"Error processing message with JobId {message_body.get('JobId')}: {str(e)}" logger.error(error_message) document_id = get_filename_from_path(s3_object_name) document_id = os.path.splitext(document_id)[0] - generate_document_logs_input(document_id, message_body.get('Status')) - return { - 'statusCode': 500, - 'body': error_message - } + generate_document_logs_input(document_id, message_body.get("Status")) + return {"statusCode": 500, "body": error_message} else: - error_message = 'Incorrect value for ENVIRONMENT VARIABLES: PROPERTY_FILE_S3_PATH\r' + str(property_file_path) + error_message = ( + "Incorrect value for ENVIRONMENT VARIABLES: PROPERTY_FILE_S3_PATH\r" + + str(property_file_path) + ) logger.error(error_message) - return { - 'statusCode': 500, - 'body': error_message - } + return {"statusCode": 500, "body": error_message} def get_filename_from_path(full_path): @@ -354,16 +398,16 @@ def generate_document_logs_input(document_id, textract_status): "STAGE": "RECEIVED_FROM_TEXTRACT", "TEXTRACT_STATUS": textract_status, "MODIFIED_TIME": current_time, - "MODIFIED_BY": "TEXTRACT_RECEIVER" - } + "MODIFIED_BY": "TEXTRACT_RECEIVER", + }, } logger.info(f"Request: {data}") - lambda_client = boto3.client('lambda') + lambda_client = boto3.client("lambda") response = lambda_client.invoke( FunctionName=DATABASE_LOGGING_LAMBDA_FUNCTION_NAME, - InvocationType='Event', # Asynchronous invocation - Payload=json.dumps(data).encode('utf-8') + InvocationType="Event", # Asynchronous invocation + Payload=json.dumps(data).encode("utf-8"), ) logger.info("Generated document logs input successfully") diff --git a/textract-pipeline/src/lambda/textract-sender/index.py b/textract-pipeline/src/lambda/textract-sender/index.py index 763a7c8..8ffac54 100644 --- a/textract-pipeline/src/lambda/textract-sender/index.py +++ b/textract-pipeline/src/lambda/textract-sender/index.py @@ -13,21 +13,23 @@ logger = logging.getLogger(__name__) logger.setLevel(logging.INFO) # Initialize S3 & Textract clients -s3_client = boto3.client('s3') -textract_client = boto3.client('textract') +s3_client = boto3.client("s3") +textract_client = boto3.client("textract") DATABASE_LOGGING_LAMBDA_FUNCTION_NAME = "" + # Function to generate a Unix timestamp def generate_unix_timestamp(): # Get the current time in seconds since the epoch unix_timestamp = int(time.time()) return unix_timestamp + # Function to retrieve configuration values from S3 def load_config_from_s3(bucket_name, file_key): # Download the config file from S3 response = s3_client.get_object(Bucket=bucket_name, Key=file_key) - config_content = response['Body'].read().decode('utf-8') + config_content = response["Body"].read().decode("utf-8") # Parse the config file config_parser = ConfigParser() @@ -36,16 +38,24 @@ def load_config_from_s3(bucket_name, file_key): # Convert the configuration to a dictionary config_dict = {} for section in config_parser.sections(): - config_dict[section] = {key.upper(): value for key, value in config_parser.items(section)} + config_dict[section] = { + key.upper(): value for key, value in config_parser.items(section) + } return config_dict + # Function to move a file from source to destination in S3 def move_file_within_s3(source_bucket, source_key, destination_key): try: tags = "env=dev" # Copy the file to the destination folder - s3_client.copy_object(Bucket=source_bucket, CopySource={'Bucket': source_bucket, 'Key': source_key}, Key=destination_key, Tagging=f'{tags}') + s3_client.copy_object( + Bucket=source_bucket, + CopySource={"Bucket": source_bucket, "Key": source_key}, + Key=destination_key, + Tagging=f"{tags}", + ) # Delete the file from the source folder s3_client.delete_object(Bucket=source_bucket, Key=source_key) @@ -56,65 +66,86 @@ def move_file_within_s3(source_bucket, source_key, destination_key): except Exception as e: logger.error(f"Error moving file: {e}") + # Function to get a list of PDF files in a given S3 folder def get_pdf_files_list_from_s3(source_bucket, source_folder): file_list = [] # List S3 Object & iterate (as per max files allowed) - s3_list_response = s3_client.list_objects_v2(Bucket=source_bucket, Prefix=source_folder) + s3_list_response = s3_client.list_objects_v2( + Bucket=source_bucket, Prefix=source_folder + ) - if s3_list_response and s3_list_response['ResponseMetadata']['HTTPStatusCode'] == 200 and s3_list_response['KeyCount'] != 0: - objects = s3_list_response['Contents'] + if ( + s3_list_response + and s3_list_response["ResponseMetadata"]["HTTPStatusCode"] == 200 + and s3_list_response["KeyCount"] != 0 + ): + objects = s3_list_response["Contents"] for s3_object in objects: # Skip non-PDF files - if not s3_object['Key'].lower().endswith('.pdf'): + if not s3_object["Key"].lower().endswith(".pdf"): continue - file_list.append(s3_object['Key']) + file_list.append(s3_object["Key"]) return file_list + # Exponential backoff utility function -def exponential_backoff_retry(func, *args, retries=5, initial_delay=1, max_delay=32, **kwargs): +def exponential_backoff_retry( + func, *args, retries=5, initial_delay=1, max_delay=32, **kwargs +): delay = initial_delay for attempt in range(retries): try: return func(*args, **kwargs) except ClientError as e: - error_code = e.response['Error']['Code'] - if error_code in ['ProvisionedThroughputExceededException', 'LimitExceededException'] and attempt < retries - 1: - logger.warning(f"{error_code} on attempt {attempt + 1}/{retries}, retrying in {delay} seconds...") + error_code = e.response["Error"]["Code"] + if ( + error_code + in ["ProvisionedThroughputExceededException", "LimitExceededException"] + and attempt < retries - 1 + ): + logger.warning( + f"{error_code} on attempt {attempt + 1}/{retries}, retrying in {delay} seconds..." + ) time.sleep(delay) delay = min(delay * 2, max_delay) else: - logger.error(f"Exceeded maximum retries. Function {func.__name__} failed with error: {e}") + logger.error( + f"Exceeded maximum retries. Function {func.__name__} failed with error: {e}" + ) raise + # Function to start a Textract text detection job -def start_textract_detection_job(bucket_name, document_file_name, sns_topic_arn, sns_role_arn, job_tag): +def start_textract_detection_job( + bucket_name, document_file_name, sns_topic_arn, sns_role_arn, job_tag +): # Define the parameters for the start_document_text_detection API start_document_detection_params = { - 'DocumentLocation': { - 'S3Object': { - 'Bucket': bucket_name, - 'Name': document_file_name - } + "DocumentLocation": { + "S3Object": {"Bucket": bucket_name, "Name": document_file_name} + }, + "ClientRequestToken": "unique-token-" + + str(generate_unix_timestamp()), # Use a unique token for each request + "JobTag": job_tag, # Use a tag to identify your job + "NotificationChannel": { + "SNSTopicArn": sns_topic_arn, + "RoleArn": sns_role_arn, # Role to allow Textract service to notify SNS topic when response is ready }, - 'ClientRequestToken': 'unique-token-' + str(generate_unix_timestamp()), # Use a unique token for each request - 'JobTag': job_tag, # Use a tag to identify your job - 'NotificationChannel': { - 'SNSTopicArn': sns_topic_arn, - 'RoleArn': sns_role_arn # Role to allow Textract service to notify SNS topic when response is ready - } } - logger.info('start_document_detection_params ' + str(start_document_detection_params)) + logger.info( + "start_document_detection_params " + str(start_document_detection_params) + ) try: # Send the request to start document detection with exponential backoff textract_response = exponential_backoff_retry( textract_client.start_document_text_detection, - **start_document_detection_params + **start_document_detection_params, ) job_id = textract_response["JobId"] @@ -124,31 +155,35 @@ def start_textract_detection_job(bucket_name, document_file_name, sns_topic_arn, logger.exception("Couldn't detect text in %s.", document_file_name) raise + # Function to start a Textract analysis job -def start_textract_analysis_job(bucket_name, document_file_name, analysis_feature_type, sns_topic_arn, sns_role_arn, job_tag): +def start_textract_analysis_job( + bucket_name, + document_file_name, + analysis_feature_type, + sns_topic_arn, + sns_role_arn, + job_tag, +): # Define the parameters for the start_document_analysis API start_document_analysis_params = { - 'DocumentLocation': { - 'S3Object': { - 'Bucket': bucket_name, - 'Name': document_file_name - } + "DocumentLocation": { + "S3Object": {"Bucket": bucket_name, "Name": document_file_name} + }, + "FeatureTypes": analysis_feature_type, # Customize based on requirements + "JobTag": job_tag, # Use a tag to identify your job + "NotificationChannel": { + "SNSTopicArn": sns_topic_arn, + "RoleArn": sns_role_arn, # Role to allow Textract service to notify SNS topic when response is ready }, - 'FeatureTypes': analysis_feature_type, # Customize based on requirements - 'JobTag': job_tag, # Use a tag to identify your job - 'NotificationChannel': { - 'SNSTopicArn': sns_topic_arn, - 'RoleArn': sns_role_arn # Role to allow Textract service to notify SNS topic when response is ready - } } - logger.info('start_document_analysis_params ' + str(start_document_analysis_params)) + logger.info("start_document_analysis_params " + str(start_document_analysis_params)) try: # Send the request to start document analysis with exponential backoff textract_response = exponential_backoff_retry( - textract_client.start_document_analysis, - **start_document_analysis_params + textract_client.start_document_analysis, **start_document_analysis_params ) job_id = textract_response["JobId"] @@ -158,18 +193,16 @@ def start_textract_analysis_job(bucket_name, document_file_name, analysis_featur logger.exception("Couldn't analyze text in %s.", document_file_name) raise + # Function to get s3 object tags def get_s3_object_tags(bucket_name, object_key): try: # Get object tags - response = s3_client.get_object_tagging( - Bucket=bucket_name, - Key=object_key - ) + response = s3_client.get_object_tagging(Bucket=bucket_name, Key=object_key) # Extract tags from the response and convert to dictionary - tags_list = response['TagSet'] - tags_dict = {tag['Key']: tag['Value'] for tag in tags_list} + tags_list = response["TagSet"] + tags_dict = {tag["Key"]: tag["Value"] for tag in tags_list} return tags_dict except Exception as e: @@ -177,22 +210,30 @@ def get_s3_object_tags(bucket_name, object_key): print(f"Error: {e}") return None + # AWS Lambda handler function def lambda_handler(event, context): try: - logger.info('## ENVIRONMENT VARIABLES\r' + str(os.environ)) + logger.info("## ENVIRONMENT VARIABLES\r" + str(os.environ)) # Read environment variables - property_file_path = os.environ.get('PROPERTY_FILE_S3_PATH', '') - TEXTRACT_ROLE_ARN = os.environ.get('TEXTRACT_ROLE_ARN', '') # Textract IAM Role ARN to publish to SNS - SNS_TOPIC_ARN = os.environ.get('SNS_TOPIC_ARN', '') + property_file_path = os.environ.get("PROPERTY_FILE_S3_PATH", "") + TEXTRACT_ROLE_ARN = os.environ.get( + "TEXTRACT_ROLE_ARN", "" + ) # Textract IAM Role ARN to publish to SNS + SNS_TOPIC_ARN = os.environ.get("SNS_TOPIC_ARN", "") global DATABASE_LOGGING_LAMBDA_FUNCTION_NAME - DATABASE_LOGGING_LAMBDA_FUNCTION_NAME = os.environ.get('DATABASE_LOGGING_LAMBDA_FUNCTION_NAME', '') + DATABASE_LOGGING_LAMBDA_FUNCTION_NAME = os.environ.get( + "DATABASE_LOGGING_LAMBDA_FUNCTION_NAME", "" + ) batch_id = "" - logger.info('DATABASE_LOGGING_LAMBDA_FUNCTION_NAME: ' + str(DATABASE_LOGGING_LAMBDA_FUNCTION_NAME)) - logger.info('SNS_TOPIC_ARN: ' + SNS_TOPIC_ARN) - logger.info('TEXTRACT_ROLE_ARN: ' + TEXTRACT_ROLE_ARN) + logger.info( + "DATABASE_LOGGING_LAMBDA_FUNCTION_NAME: " + + str(DATABASE_LOGGING_LAMBDA_FUNCTION_NAME) + ) + logger.info("SNS_TOPIC_ARN: " + SNS_TOPIC_ARN) + logger.info("TEXTRACT_ROLE_ARN: " + TEXTRACT_ROLE_ARN) # Read config.properties file_path_array = property_file_path.split("/") @@ -204,72 +245,82 @@ def lambda_handler(event, context): S3_BUCKET_NAME = file_path_array[0] CONFIG_FILE_PATH = "/".join(file_path_array[1:]) - logger.info(f'S3_BUCKET_NAME: {S3_BUCKET_NAME}') - logger.info(f'CONFIG_FILE_PATH: {CONFIG_FILE_PATH}') + logger.info(f"S3_BUCKET_NAME: {S3_BUCKET_NAME}") + logger.info(f"CONFIG_FILE_PATH: {CONFIG_FILE_PATH}") # Load config file config_dict = load_config_from_s3(S3_BUCKET_NAME, CONFIG_FILE_PATH) - #logger.info('## CONFIG DICTIONARY\r' + str(config_dict)) + # logger.info('## CONFIG DICTIONARY\r' + str(config_dict)) # File count file_count = 0 # Process each message from the SQS event - for record in event['Records']: + for record in event["Records"]: # Extract the message body from the record - record_body = json.loads(record['body']) + record_body = json.loads(record["body"]) - #logger.info('Message Count: ', str(len(record_body['Records'])) ) + # logger.info('Message Count: ', str(len(record_body['Records'])) ) - for sqs_record in record_body['Records']: + for sqs_record in record_body["Records"]: # Construct the source paths - source_path = unquote_plus(sqs_record['s3']['object']['key']) + source_path = unquote_plus(sqs_record["s3"]["object"]["key"]) # Read tags from file tags_dict = get_s3_object_tags(S3_BUCKET_NAME, source_path) - JOB_TAG = config_dict['OTHERS']['JOB_TAG'] + JOB_TAG = config_dict["OTHERS"]["JOB_TAG"] # check if key exists in tags_dict if "BatchId" in tags_dict.keys(): - JOB_TAG = JOB_TAG + "-" + tags_dict['BatchId'] - batch_id = tags_dict['BatchId'] + JOB_TAG = JOB_TAG + "-" + tags_dict["BatchId"] + batch_id = tags_dict["BatchId"] else: logger.error(f"BatchId not found in file tag") return # Extract configuration values - SOURCE_LOCATION = config_dict['FOLDER_LOCATIONS']['SOURCE_LOCATION'].format(batch_id) # SOURCE_LOCATION - STAGING_LOCATION = config_dict['FOLDER_LOCATIONS']['STAGING_LOCATION'].format(batch_id) - ANALYSIS_FEATURE_TYPE = config_dict['OTHERS']['ANALYSIS_FEATURE_TYPE'].split(",") # Analysis FeatureType - SENDER_MAX_FILES = int(config_dict['OTHERS']['SENDER_MAX_FILES']) - PROCESS_TYPE = str(config_dict['OTHERS']['PROCESS_TYPE']).upper() - + SOURCE_LOCATION = config_dict["FOLDER_LOCATIONS"][ + "SOURCE_LOCATION" + ].format( + batch_id + ) # SOURCE_LOCATION + STAGING_LOCATION = config_dict["FOLDER_LOCATIONS"][ + "STAGING_LOCATION" + ].format(batch_id) + ANALYSIS_FEATURE_TYPE = config_dict["OTHERS"][ + "ANALYSIS_FEATURE_TYPE" + ].split( + "," + ) # Analysis FeatureType + SENDER_MAX_FILES = int(config_dict["OTHERS"]["SENDER_MAX_FILES"]) + PROCESS_TYPE = str(config_dict["OTHERS"]["PROCESS_TYPE"]).upper() # Construct the destination paths - destination_path = STAGING_LOCATION + source_path.replace(SOURCE_LOCATION, "") + destination_path = STAGING_LOCATION + source_path.replace( + SOURCE_LOCATION, "" + ) # Move file to staging move_file_within_s3(S3_BUCKET_NAME, source_path, destination_path) logger.info(f"Batch ID: {batch_id}") - logger.info('SOURCE_LOCATION: ' + SOURCE_LOCATION) - logger.info('STAGING_LOCATION: ' + STAGING_LOCATION) - logger.info('ANALYSIS_FEATURE_TYPE: ' + str(ANALYSIS_FEATURE_TYPE)) - logger.info('SENDER_MAX_FILES: ' + str(SENDER_MAX_FILES)) - logger.info('JOB_TAG: ' + str(JOB_TAG)) - logger.info('PROCESS_TYPE: ' + str(PROCESS_TYPE)) - + logger.info("SOURCE_LOCATION: " + SOURCE_LOCATION) + logger.info("STAGING_LOCATION: " + STAGING_LOCATION) + logger.info("ANALYSIS_FEATURE_TYPE: " + str(ANALYSIS_FEATURE_TYPE)) + logger.info("SENDER_MAX_FILES: " + str(SENDER_MAX_FILES)) + logger.info("JOB_TAG: " + str(JOB_TAG)) + logger.info("PROCESS_TYPE: " + str(PROCESS_TYPE)) job_id = "" if PROCESS_TYPE == "ANALYSIS": # Start Textract analysis job - job_id = start_textract_analysis_job ( + job_id = start_textract_analysis_job( S3_BUCKET_NAME, destination_path, ANALYSIS_FEATURE_TYPE, @@ -279,7 +330,7 @@ def lambda_handler(event, context): ) elif PROCESS_TYPE == "DETECTION": # Start Textract detection job - job_id = start_textract_detection_job ( + job_id = start_textract_detection_job( S3_BUCKET_NAME, destination_path, SNS_TOPIC_ARN, @@ -288,50 +339,49 @@ def lambda_handler(event, context): ) file_count = file_count + 1 - logger.info(str(file_count) + '. ' + str(source_path) + " Job Id: " + str(job_id)) + logger.info( + str(file_count) + + ". " + + str(source_path) + + " Job Id: " + + str(job_id) + ) time.sleep(1) - success_message = 'Total files sent to textract : '+ str(file_count) + success_message = "Total files sent to textract : " + str(file_count) logger.info(success_message) document_id = get_filename_from_path(source_path) document_id = os.path.splitext(document_id)[0] generate_document_logs_input(document_id, job_id, "SENT_SUCCESSFULLY") - return { - 'statusCode': 200, - 'body': success_message - } + return {"statusCode": 200, "body": success_message} else: - error_message = 'Incorrect value for ENVIRONMENT VARIABLES: PROPERTY_FILE_S3_PATH\r' + str(property_file_path) + error_message = ( + "Incorrect value for ENVIRONMENT VARIABLES: PROPERTY_FILE_S3_PATH\r" + + str(property_file_path) + ) logger.error(error_message) generate_document_logs_input(document_id, job_id, "FAILED_TO_SENT") - return { - 'statusCode': 500, - 'body': error_message - } + return {"statusCode": 500, "body": error_message} except ClientError as e: # Handle specific Textract client errors error_message = f"Error in Textract operation: {e}" logger.error(error_message) - return { - 'statusCode': 500, - 'body': error_message - } + return {"statusCode": 500, "body": error_message} except Exception as e: # Handle other exceptions error_message = f"Unexpected error: {e}" logger.error(error_message) - return { - 'statusCode': 500, - 'body': error_message - } + return {"statusCode": 500, "body": error_message} + def get_filename_from_path(full_path): return os.path.basename(full_path) + def generate_document_logs_input(document_id, job_id, textract_status): current_time = datetime.datetime.now().isoformat() @@ -344,16 +394,16 @@ def generate_document_logs_input(document_id, job_id, textract_status): "JOB_ID": job_id, "MODIFIED_TIME": current_time, "TEXTRACT_STATUS": textract_status, - "MODIFIED_BY": "TEXTRACT_SENDER" - } + "MODIFIED_BY": "TEXTRACT_SENDER", + }, } logger.info(f"Request: {data}") - lambda_client = boto3.client('lambda') + lambda_client = boto3.client("lambda") response = lambda_client.invoke( FunctionName=DATABASE_LOGGING_LAMBDA_FUNCTION_NAME, - InvocationType='Event', - Payload=json.dumps(data).encode('utf-8') + InvocationType="Event", + Payload=json.dumps(data).encode("utf-8"), ) logger.info("Generated document logs input successfully") diff --git a/textract-pipeline/src/lambda/tiff-to-pdf/index.py b/textract-pipeline/src/lambda/tiff-to-pdf/index.py index 015fbee..ac4c498 100644 --- a/textract-pipeline/src/lambda/tiff-to-pdf/index.py +++ b/textract-pipeline/src/lambda/tiff-to-pdf/index.py @@ -1,3 +1,3 @@ # AWS Lambda handler function def lambda_handler(event, context): - print("Sender Lambda") \ No newline at end of file + print("Sender Lambda") diff --git a/textract-pipeline/src/lambda/trigger-pipeline/index.py b/textract-pipeline/src/lambda/trigger-pipeline/index.py index 0db2334..110e6b5 100644 --- a/textract-pipeline/src/lambda/trigger-pipeline/index.py +++ b/textract-pipeline/src/lambda/trigger-pipeline/index.py @@ -13,15 +13,16 @@ logger = logging.getLogger(__name__) logger.setLevel(logging.INFO) # Initialize S3 clients -s3_client = boto3.client('s3') +s3_client = boto3.client("s3") DATABASE_LOGGING_LAMBDA_FUNCTION_NAME = "" + # Function to retrieve configuration values from S3 def load_config_from_s3(bucket_name, file_key): - + # Download the config file from S3 response = s3_client.get_object(Bucket=bucket_name, Key=file_key) - config_content = response['Body'].read().decode('utf-8') + config_content = response["Body"].read().decode("utf-8") # Parse the config file config_parser = ConfigParser() @@ -30,16 +31,20 @@ def load_config_from_s3(bucket_name, file_key): # Convert the configuration to a dictionary config_dict = {} for section in config_parser.sections(): - config_dict[section] = {key.upper(): value for key, value in config_parser.items(section)} + config_dict[section] = { + key.upper(): value for key, value in config_parser.items(section) + } return config_dict - + + def get_current_timestamp(): """ Get the current timestamp in a formatted string. """ return datetime.now().strftime("%Y-%m-%d_%H-%M-%S") - + + def get_api_input(): json = { "s3_bucket": "doczy-dev-infra-textract", @@ -48,63 +53,54 @@ def get_api_input(): "username": "", "contract_list": [ { - "contract_name": "file-sample_150kB.pdf", - "groups": [ - "group-name-1", - "group-name-2", - "group-name-3" - ], - "contract_source_path": "contracts_landing_zone/file-sample_150kB.pdf" + "contract_name": "file-sample_150kB.pdf", + "groups": ["group-name-1", "group-name-2", "group-name-3"], + "contract_source_path": "contracts_landing_zone/file-sample_150kB.pdf", }, { - "contract_name": "Sample_PDF.docx", - "groups": [ - "group-name-1", - "group-name-2", - "group-name-3" - ], - "contract_source_path": "contracts_landing_zone/Sample_PDF.docx" - } - ] + "contract_name": "Sample_PDF.docx", + "groups": ["group-name-1", "group-name-2", "group-name-3"], + "contract_source_path": "contracts_landing_zone/Sample_PDF.docx", + }, + ], } return json - + + def create_s3_folder(bucket_name, folder_name): """ Create a folder in the specified S3 bucket with the given name. """ # Ensure the folder name ends with a trailing slash - if not folder_name.endswith('/'): - folder_name += '/' - + if not folder_name.endswith("/"): + folder_name += "/" + try: # Create an empty object with a trailing slash to represent the folder - #s3_client.put_object(Bucket=bucket_name, Key=folder_name) + # s3_client.put_object(Bucket=bucket_name, Key=folder_name) print(f"Folder '{folder_name}' created successfully in bucket '{bucket_name}'") return folder_name except Exception as e: print(f"Error creating folder: {e}") - + + def move_object_within_bucket(bucket_name, source_key, destination_key, tags=None): """ Move an object within the same S3 bucket and add additional tags if provided. """ - copy_source = { - 'Bucket': bucket_name, - 'Key': source_key - } + copy_source = {"Bucket": bucket_name, "Key": source_key} copy_object_args = { - 'Bucket': bucket_name, - 'Key': destination_key, - 'CopySource': copy_source, - 'TaggingDirective': 'REPLACE' + "Bucket": bucket_name, + "Key": destination_key, + "CopySource": copy_source, + "TaggingDirective": "REPLACE", } if tags: - copy_object_args['Tagging'] = tags + copy_object_args["Tagging"] = tags print(copy_object_args) try: @@ -114,66 +110,58 @@ def move_object_within_bucket(bucket_name, source_key, destination_key, tags=Non # Delete the original object s3_client.delete_object(Bucket=bucket_name, Key=source_key) - msg = f"Object moved successfully. New object Key: {destination_key}" + msg = f"Object moved successfully. New object Key: {destination_key}" logger.info(msg) return msg except Exception as e: - msg = f"Error moving object: {e}" + msg = f"Error moving object: {e}" logger.error(msg) return msg + def lambda_handler(request, context): - + try: contract_list = request if "s3_bucket" not in request or request["s3_bucket"] == "": error_message = "s3_bucket key is missing or value is empty in request" logger.error(error_message) - return { - 'statusCode': 500, - 'body': error_message - } + return {"statusCode": 500, "body": error_message} elif "batch_id" not in request or request["batch_id"] == "": error_message = "batch_id key is missing or value is empty in request" logger.error(error_message) - return { - 'statusCode': 500, - 'body': error_message - } + return {"statusCode": 500, "body": error_message} elif "client_name" not in request or request["client_name"] == "": error_message = "client_name key is missing or value is empty in request" logger.error(error_message) - return { - 'statusCode': 500, - 'body': error_message - } - elif "contract_list" not in request or type(request["contract_list"]) is not list: + return {"statusCode": 500, "body": error_message} + elif ( + "contract_list" not in request or type(request["contract_list"]) is not list + ): error_message = "contract_list key is missing or value is empty in request" logger.error(error_message) - return { - 'statusCode': 500, - 'body': error_message - } - + return {"statusCode": 500, "body": error_message} batch_id = request["batch_id"] client_name = request["client_name"] contract_list = request["contract_list"] username = request["username"] common_dict = remove_key(request, "contract_list") - - logger.info(f'Request: {request}') - - logger.info('## ENVIRONMENT VARIABLES\r' + str(os.environ)) - + + logger.info(f"Request: {request}") + + logger.info("## ENVIRONMENT VARIABLES\r" + str(os.environ)) + # Read environment variables - property_file_path = os.environ.get('PROPERTY_FILE_S3_PATH', '') + property_file_path = os.environ.get("PROPERTY_FILE_S3_PATH", "") global DATABASE_LOGGING_LAMBDA_FUNCTION_NAME - DATABASE_LOGGING_LAMBDA_FUNCTION_NAME = os.environ.get('DATABASE_LOGGING_LAMBDA_FUNCTION_NAME', '') - INITIATE_DB_SQS_URL = os.environ.get('INITIATE_DB_SQS_URL', '') - - logger.info('INITIATE_DB_SQS_URL: ' + str(INITIATE_DB_SQS_URL)) - + DATABASE_LOGGING_LAMBDA_FUNCTION_NAME = os.environ.get( + "DATABASE_LOGGING_LAMBDA_FUNCTION_NAME", "" + ) + INITIATE_DB_SQS_URL = os.environ.get("INITIATE_DB_SQS_URL", "") + + logger.info("INITIATE_DB_SQS_URL: " + str(INITIATE_DB_SQS_URL)) + # Read config.properties file_path_array = property_file_path.split("/") @@ -183,17 +171,17 @@ def lambda_handler(request, context): docx_contracts_count = 0 other_contracts = 0 file_list = {} - + # Valid if file_path_array has more than 2 elements if len(file_path_array) > 1: - + # Extract BUCKET_NAME and config_file_path S3_BUCKET_NAME = file_path_array[0] CONFIG_FILE_PATH = "/".join(file_path_array[1:]) - - logger.info(f'S3_BUCKET_NAME: {S3_BUCKET_NAME}') - logger.info(f'CONFIG_FILE_PATH: {CONFIG_FILE_PATH}') - + + logger.info(f"S3_BUCKET_NAME: {S3_BUCKET_NAME}") + logger.info(f"CONFIG_FILE_PATH: {CONFIG_FILE_PATH}") + # batch folder in s3 bucket # batch_id = "batch_" + get_current_timestamp() # Add timestamp to the folder name # foldername_prefix = "batches/" + batch_id @@ -201,29 +189,41 @@ def lambda_handler(request, context): # Load config file config_dict = load_config_from_s3(S3_BUCKET_NAME, CONFIG_FILE_PATH) - - # Extract configuration values - CONTRACTS_LANDNING_ZONE = config_dict['FOLDER_LOCATIONS']['CONTRACTS_LANDNING_ZONE'].format(batch_id) # CONTRACTS_LANDNING_ZONE - # Uncomment below code while integrating to main pipeline - ALL_PDF_LOCATION = config_dict['FOLDER_LOCATIONS']['ALL_PDF_LOCATION'].format(batch_id) # ALL_PDF_LOCATION - SOURCE_DOCX_LOCATION = config_dict['FOLDER_LOCATIONS']['SOURCE_DOCX_LOCATION'].format(batch_id) # SOURCE_DOCX_LOCATION - - logger.info('CONTRACTS_LANDNING_ZONE: ' + CONTRACTS_LANDNING_ZONE) - logger.info('ALL_PDF_LOCATION: ' + ALL_PDF_LOCATION) - logger.info('SOURCE_DOCX_LOCATION: ' + SOURCE_DOCX_LOCATION) - - tags_dict = {'batch_id': batch_id, "client_name": client_name} + # Extract configuration values + CONTRACTS_LANDNING_ZONE = config_dict["FOLDER_LOCATIONS"][ + "CONTRACTS_LANDNING_ZONE" + ].format( + batch_id + ) # CONTRACTS_LANDNING_ZONE + + # Uncomment below code while integrating to main pipeline + ALL_PDF_LOCATION = config_dict["FOLDER_LOCATIONS"][ + "ALL_PDF_LOCATION" + ].format( + batch_id + ) # ALL_PDF_LOCATION + SOURCE_DOCX_LOCATION = config_dict["FOLDER_LOCATIONS"][ + "SOURCE_DOCX_LOCATION" + ].format( + batch_id + ) # SOURCE_DOCX_LOCATION + + logger.info("CONTRACTS_LANDNING_ZONE: " + CONTRACTS_LANDNING_ZONE) + logger.info("ALL_PDF_LOCATION: " + ALL_PDF_LOCATION) + logger.info("SOURCE_DOCX_LOCATION: " + SOURCE_DOCX_LOCATION) + + tags_dict = {"batch_id": batch_id, "client_name": client_name} additional_tags = urlencode(tags_dict) logger.info("Tagging: " + additional_tags) - + logger.info(contract_list) - + # Iterate over contracts for contract in contract_list: - total_contracts_count+=1 - contract_source_path = contract['contract_source_path'] - + total_contracts_count += 1 + contract_source_path = contract["contract_source_path"] + logger.info("CONTRACT_SOURCE: " + contract_source_path) file_prefix, basename = os.path.split(contract_source_path) filename, file_ext = os.path.splitext(basename) @@ -231,7 +231,7 @@ def lambda_handler(request, context): common_dict["contract_list"] = [contract] message_body = json.dumps(common_dict) - + # Send message to SQS message_id = send_message_to_sqs(INITIATE_DB_SQS_URL, message_body) @@ -241,70 +241,63 @@ def lambda_handler(request, context): logger.error("Failed to send message to SQS.") # If pdf or filepart copy to folder - if file_ext.lower() in [".pdf",".filepart"]: - pdf_contracts_count+=1 - + if file_ext.lower() in [".pdf", ".filepart"]: + pdf_contracts_count += 1 + # Copy to ALL_PDF_LOCATION source_key = contract_source_path file_list[source_key] = "Sent for processing" # If doc/docx copy to doc folder - elif file_ext.lower() in [".docx",".doc"]: - docx_contracts_count+=1 + elif file_ext.lower() in [".docx", ".doc"]: + docx_contracts_count += 1 # Copy to SOURCE_DOCX_LOCATION source_key = contract_source_path file_list[source_key] = "Sent for processing" - + # If tiff copy to tiff folder # elif file_ext.lower() in [".tiff"]: - - # Copy to SOURCE_TIFF_LOCATION - + + # Copy to SOURCE_TIFF_LOCATION + # Other files else: - other_contracts+=1 + other_contracts += 1 file_list[source_key] = "Not processed" - generate_batch_logs_input(batch_id,total_contracts_count,username) + generate_batch_logs_input(batch_id, total_contracts_count, username) except ClientError as e: # Handle specific Textract client errors error_message = f"Error in s3 client operation: {e}" logger.error(error_message) - return { - 'statusCode': 500, - 'body': error_message - } + return {"statusCode": 500, "body": error_message} except Exception as e: # Handle other exceptions error_message = f"Unexpected error: {e}" logger.error(error_message) - return { - 'statusCode': 500, - 'body': error_message - } - + return {"statusCode": 500, "body": error_message} + response = { - "total_files" : total_contracts_count, - "pdf_files" : pdf_contracts_count, - "doc_files" : docx_contracts_count, - "other_files" : other_contracts, - "file_list" : file_list - } - # TODO implement - return { - - 'statusCode': 200, - 'body': response + "total_files": total_contracts_count, + "pdf_files": pdf_contracts_count, + "doc_files": docx_contracts_count, + "other_files": other_contracts, + "file_list": file_list, } + # TODO implement + return {"statusCode": 200, "body": response} -def generate_batch_logs_input(batch_id,no_of_documents,user_name): +def generate_batch_logs_input(batch_id, no_of_documents, user_name): current_time = datetime.datetime.now().isoformat() - logger.info('DATABASE_LOGGING_LAMBDA_FUNCTION_NAME: ' + str(DATABASE_LOGGING_LAMBDA_FUNCTION_NAME)) + logger.info( + "DATABASE_LOGGING_LAMBDA_FUNCTION_NAME: " + + str(DATABASE_LOGGING_LAMBDA_FUNCTION_NAME) + ) data = { "operation": "update", @@ -313,42 +306,41 @@ def generate_batch_logs_input(batch_id,no_of_documents,user_name): "BATCH_ID": batch_id, "NO_OF_DOCUMENTS": no_of_documents, "USER_NAME": user_name, - "EXECUTION_START_TIME": current_time - } + "EXECUTION_START_TIME": current_time, + }, } logger.info(f"Request: {data}") - lambda_client = boto3.client('lambda') + lambda_client = boto3.client("lambda") response = lambda_client.invoke( FunctionName=DATABASE_LOGGING_LAMBDA_FUNCTION_NAME, - InvocationType='Event', # Asynchronous invocation - Payload=json.dumps(data).encode('utf-8') + InvocationType="Event", # Asynchronous invocation + Payload=json.dumps(data).encode("utf-8"), ) logger.info("Generated document logs input successfully") return response + def send_message_to_sqs(queue_url, message_body): # Initialize SQS client - sqs = boto3.client('sqs') - + sqs = boto3.client("sqs") + try: # Send message to SQS queue - response = sqs.send_message( - QueueUrl=queue_url, - MessageBody=message_body - ) + response = sqs.send_message(QueueUrl=queue_url, MessageBody=message_body) # Return message ID if successful - return response['MessageId'] + return response["MessageId"] except Exception as e: # Log any errors print("Error sending message to SQS:", e) return None + def remove_key(original_dict, key_to_remove): # Make a copy of the original dictionary new_dict = original_dict.copy() - + # Check if the key exists in the dictionary if key_to_remove in new_dict: # Remove the key from the dictionary @@ -356,4 +348,4 @@ def remove_key(original_dict, key_to_remove): else: print("Key not found in dictionary") - return new_dict \ No newline at end of file + return new_dict