diff --git a/.gitignore b/.gitignore index 924c5d3..ed2bd94 100644 --- a/.gitignore +++ b/.gitignore @@ -55,3 +55,17 @@ terraform.tfstate .terraform.lock.hcl build/ +# Data Files +streamlit/history.csv +streamlit/RESULTS +streamlit/DB/ +streamlit/RAW_DOCUMENTS/ +streamlit/SOURCE_DOCUMENTS/ +streamlit/contract_field_values.csv +streamlit/contract_fields.csv +streamlit/sample.csv +streamlit/temp1.csv +streamlit/temp2.csv + +# env +streamlit/venv diff --git a/airflow/Qa_Dag.py b/airflow/dags/Qa_Dag.py similarity index 80% rename from airflow/Qa_Dag.py rename to airflow/dags/Qa_Dag.py index bceed96..2896f1b 100644 --- a/airflow/Qa_Dag.py +++ b/airflow/dags/Qa_Dag.py @@ -67,6 +67,19 @@ def getData(): print(f"Successful S3 put_object response. Status - {status}") else: 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: + 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() + ) print("Dataframe is written to S3 successfully.") diff --git a/airflow/dags/cicd_test_dag.py b/airflow/dags/cicd_test_dag.py new file mode 100644 index 0000000..20e7351 --- /dev/null +++ b/airflow/dags/cicd_test_dag.py @@ -0,0 +1,156 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. + + +# Mandatory imports +from __future__ import annotations +import os +from datetime import datetime +import logging +from airflow import DAG +from airflow.operators.empty import EmptyOperator +from airflow.utils.trigger_rule import TriggerRule +from airflow.models import Variable + +# SNOWFLAKE only imports +from airflow.providers.snowflake.operators.snowflake import SnowflakeOperator + +# PYTHON only imports +from airflow.providers.snowflake.hooks.snowflake import SnowflakeHook +from airflow.operators.python import PythonOperator + + +# Configuring basic logging for INFO / ERROR in our scripts + +logging.basicConfig(level=logging.INFO) +logger = logging.getLogger(__name__) + +# Recommended to use this in the beginning of the script + +SNOWFLAKE_CONN_ID = "doczy_dev_snowflake" # Specific for every database +DAG_ID = "cicd_testing_dag" # Must be unique for every dag +DATABASE="XXX" + +default_params = {"Database": "", "Schema":""} + +''' +Tags will be used to filter DAGs on the Airflow UI. + +Guidelines for TAGs: +"prod" - for dataload scripts that have been tested +"dev" - for dataloads for DEV databases / scripts in dev +"db_name" - same naming convention as Snowflake databses for clients / projects +"medical" / "pharmacy" - based on the scenario +"etl" / "adhoc" - based on the nature of the dataload + +''' +TAGS=["adhoc","dev"] # MANDATORY + + + +''' +TRIGGER RULES for a task: + +all_success: (default) all parents have succeeded +all_failed: all parents are in a failed or upstream_failed state +all_done: all parents are done with their execution +one_failed: fires as soon as at least one parent has failed, it does not wait for all parents to be done +one_success: fires as soon as at least one parent succeeds, it does not wait for all parents to be done +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' + + + +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, + # 'end_date': datetime(2016, 1, 1), + # 'wait_for_downstream': False, + # 'sla': timedelta(hours=2), + # 'execution_timeout': timedelta(seconds=300), + # 'on_failure_callback': some_function, + # '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 + # 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}, + 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 + +) + + +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'] + 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") + # max_audit_sid = result[0] + + logging.info(f"PYTHON OPERATOR OUTPUT = {query_output.fetchall()}") + conn.close + + +# 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_callable=python_op_eg, + 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. +sf_task = SnowflakeOperator( + task_id="get_info_using_sf_op", + sql=f"select * from {DATABASE}.STG.DIM_AUDIT", + dag=dag +) + + + +end_job = EmptyOperator(task_id='End') + +# EXECUTING TASKS +begin_job >> python_task >> sf_task >> end_job \ No newline at end of file diff --git a/airflow/dags/client_name_dag.py b/airflow/dags/client_name_dag.py new file mode 100644 index 0000000..138d64a --- /dev/null +++ b/airflow/dags/client_name_dag.py @@ -0,0 +1,133 @@ +import re +import pandas as pd +from airflow import DAG +from airflow.operators.python_operator import PythonOperator +from airflow.providers.amazon.aws.hooks.s3 import S3Hook +from datetime import datetime, timedelta +from airflow.providers.snowflake.hooks.snowflake import SnowflakeHook +from airflow.operators.empty import EmptyOperator +import logging +from airflow.exceptions import AirflowFailException + + +logging.basicConfig(level=logging.INFO) +logger = logging.getLogger(__name__) + +SNOWFLAKE_CONN_ID = "doczy_dev_snowflake" +DAG_ID = "load_client_config" +DATABASE="DOCZY_DEV" +# bucket = "airflow-data-ingestion" + +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/' + + # 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 + + # 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) + + # Convert DataFrame to CSV string + csv_buffer = StringIO() + df.to_csv(csv_buffer, index=False) + csv_content = csv_buffer.getvalue() + + # Replace the file in S3 with the processed content + s3_hook.load_string( + string_data=csv_content, + key=latest_file_key, + bucket_name=bucket_name, + 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]) + 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) + final_name = final_name.lower() + 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') + logger.info(f"File name: {file_name}") + 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') + else: + 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) +} + +dag = DAG( + DAG_ID, + default_args=default_args, + start_date=datetime(2024, 3, 28), + catchup=False, + description='Process a CSV in S3 and replace it', + schedule_interval='@daily' +) + +begin_job = EmptyOperator(task_id='Begin') + +process_csv_task = PythonOperator( + task_id='process_csv', + python_callable=process_csv_in_s3, + provide_context=True, + dag=dag +) + +load_client_config = PythonOperator( + task_id="load_client_config", + python_callable=call_stored_proc, + dag=dag, + provide_context=True, + op_kwargs={'proc_name':'LOAD_CLIENT_CONFIG'} +) + + + + +end_job = EmptyOperator(task_id='End') + +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 new file mode 100644 index 0000000..9f83064 --- /dev/null +++ b/airflow/dags/config_interface_dag.py @@ -0,0 +1,100 @@ + + +from __future__ import annotations + +import os +from datetime import datetime +import logging + +from airflow import DAG +from airflow.providers.snowflake.operators.snowflake import SnowflakeOperator +from airflow.providers.snowflake.hooks.snowflake import SnowflakeHook +from airflow.operators.python import PythonOperator +from airflow.utils.trigger_rule import TriggerRule +from airflow.operators.empty import EmptyOperator +from airflow.models import Variable +from airflow.exceptions import AirflowFailException + + +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" +# bucket = "airflow-data-ingestion" + +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' + +# 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":""} +# 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) + 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'] + + 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') + else: + 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}, + tags=TAGS, + catchup=False, + schedule=None, + params = default_params +) + +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'} +) + +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'} +) + + + + +end_job = EmptyOperator(task_id='End') + + +begin_job >> load_request_submission >> load_contract_config >> end_job \ No newline at end of file diff --git a/airflow/dags/raw_training_data_dag.py b/airflow/dags/raw_training_data_dag.py new file mode 100644 index 0000000..f8156d4 --- /dev/null +++ b/airflow/dags/raw_training_data_dag.py @@ -0,0 +1,108 @@ + + +from __future__ import annotations + +import os +from datetime import datetime +import logging + +from airflow import DAG +from airflow.providers.snowflake.operators.snowflake import SnowflakeOperator +from airflow.providers.snowflake.hooks.snowflake import SnowflakeHook +from airflow.operators.python import PythonOperator +from airflow.utils.trigger_rule import TriggerRule +from airflow.operators.empty import EmptyOperator +from airflow.models import Variable +from airflow.exceptions import AirflowFailException + + +logging.basicConfig(level=logging.INFO) +logger = logging.getLogger(__name__) + + + +SNOWFLAKE_CONN_ID = "doczy_dev_snowflake" +DAG_ID = "load_raw_training_data" +DATABASE="DOCZY_DEV" +# bucket = "airflow-data-ingestion" + +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' + +# 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":""} +# 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) + 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'] + + 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') + else: + 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}, + tags=TAGS, + catchup=False, + schedule=None, + params = default_params +) + +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'} +) + +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'} +) + +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'} +) + + + +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 diff --git a/airflow/training_results_dag.py b/airflow/dags/training_results_dag.py similarity index 66% rename from airflow/training_results_dag.py rename to airflow/dags/training_results_dag.py index 40cad67..7c01caa 100644 --- a/airflow/training_results_dag.py +++ b/airflow/dags/training_results_dag.py @@ -35,17 +35,27 @@ 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":""} # 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" +# training_results_file_name = "training_results_sample.csv" +# attempt_logs_file_name = "attempt_logs_sample.csv" -def call_stored_proc(proc_name,file_name): +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'] + cur.execute(f"CALL {DATABASE}.STG.{proc_name}('{file_name}');") result = cur.fetchone() if result[0] == 'Setup, Load, and Audit Complete': @@ -60,7 +70,8 @@ dag = DAG( default_args={"snowflake_conn_id": SNOWFLAKE_CONN_ID, "retries":0}, tags=TAGS, catchup=False, - schedule=None + schedule=None, + params = default_params ) begin_job = EmptyOperator(task_id='Begin') @@ -70,14 +81,14 @@ load_training_results = PythonOperator( task_id="load_training_results", python_callable=call_stored_proc, dag=dag, - op_kwargs={'proc_name':'LOAD_TRAINING_RESULTS', 'file_name':training_results_file_name} + 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_name':attempt_logs_file_name} + op_kwargs={'proc_name':'LOAD_ATTEMPT_LOGS', 'file_type':'attempt_logs'} ) diff --git a/bitbucket-pipelines.yml b/bitbucket-pipelines.yml index 16e644c..29c80e3 100644 --- a/bitbucket-pipelines.yml +++ b/bitbucket-pipelines.yml @@ -4,7 +4,7 @@ definitions: name: "Plan and Apply Terraform" image: zenika/terraform-aws-cli:latest script: - - terraform --version + - terraform --version # terraform steps here trigger: automatic # uncomment after testing # trigger: manual @@ -21,10 +21,46 @@ definitions: - if [ "$BITBUCKET_BRANCH" == "UAT" ]; then export SNOWFLAKE_DATABASE=$UAT_SNOWFLAKE_DATABASE; export SNOWFLAKE_ROLE=$UAT_SNOWFLAKE_ROLE; fi - if [ "$BITBUCKET_BRANCH" == "PROD" ]; then export SNOWFLAKE_DATABASE=$PROD_SNOWFLAKE_DATABASE; export SNOWFLAKE_ROLE=$PROD_SNOWFLAKE_ROLE; fi - schemachange -a $SNOWFLAKE_ACCOUNT -u $SNOWFLAKE_USER -r $SNOWFLAKE_ROLE -w $SNOWFLAKE_WAREHOUSE -d $SNOWFLAKE_DATABASE -c $SNOWFLAKE_DATABASE.SCHEMACHANGE.CHANGE_HISTORY --create-change-history-table -v - trigger: automatic # uncomment after testing + trigger: automatic # trigger: manual caches: - pip + - step: &streamlit_deploy + name: "Deploy streamlit to EC2" + oidc: true + image: python:3.8 + script: + - python -m pip install --upgrade pip + - apt-get update && apt-get install -y jq git + - pip install awscli + - export AWS_ROLE_ARN=arn:aws:iam::$AWS_ACCOUNT_NO:role/$OIDC_ROLE + - export AWS_WEB_IDENTITY_TOKEN_FILE=$(pwd)/web-identity-token + - echo $BITBUCKET_STEP_OIDC_TOKEN > $(pwd)/web-identity-token + - export STS_OUTPUT=$(aws sts assume-role-with-web-identity --role-arn $AWS_ROLE_ARN --role-session-name BitbucketPipeline --web-identity-token "$BITBUCKET_STEP_OIDC_TOKEN" --duration-seconds 3600) + - export AWS_ACCESS_KEY_ID=$(echo $STS_OUTPUT | jq -r '.Credentials.AccessKeyId') + - export AWS_SECRET_ACCESS_KEY=$(echo $STS_OUTPUT | jq -r '.Credentials.SecretAccessKey') + - export AWS_SESSION_TOKEN=$(echo $STS_OUTPUT | jq -r '.Credentials.SessionToken') + - aws sts get-caller-identity + - aws ssm send-command --document-name "AWS-RunShellScript" --instance-ids $DEV_INSTANCE_ID --region $AWS_DEFAULT_REGION --parameters commands='["cd /home/ubuntu/streamlit","(if [ -d doczy.ai ]; then cd doczy.ai && git fetch && git pull; else git clone git@bitbucket.org:aarete/doczy.ai.git; fi)"]' + trigger: automatic # uncomment after testing + - step: &airflow_dags_deploy + name: "Deploy Airflow DAGs to S3 bucket" + image: python:3.8 + oidc: true + script: + - python -m pip install --upgrade pip + - apt-get update && apt-get install -y jq git + - pip install awscli + - export AWS_ROLE_ARN=arn:aws:iam::$AWS_ACCOUNT_NO:role/$OIDC_ROLE + - export AWS_WEB_IDENTITY_TOKEN_FILE=$(pwd)/web-identity-token + - echo $BITBUCKET_STEP_OIDC_TOKEN > $(pwd)/web-identity-token + - export STS_OUTPUT=$(aws sts assume-role-with-web-identity --role-arn $AWS_ROLE_ARN --role-session-name BitbucketPipeline --web-identity-token "$BITBUCKET_STEP_OIDC_TOKEN" --duration-seconds 3600) + - export AWS_ACCESS_KEY_ID=$(echo $STS_OUTPUT | jq -r '.Credentials.AccessKeyId') + - export AWS_SECRET_ACCESS_KEY=$(echo $STS_OUTPUT | jq -r '.Credentials.SecretAccessKey') + - export AWS_SESSION_TOKEN=$(echo $STS_OUTPUT | jq -r '.Credentials.SessionToken') + - aws sts get-caller-identity + - aws s3 sync ./airflow/dags s3://doczy-dev-infra-mwaa-resources/dags + trigger: automatic # uncomment after testing pipelines: branches: @@ -41,7 +77,22 @@ pipelines: changesets: includePaths: - snowflake/**/* - uat: + - step: + <<: *streamlit_deploy + condition: + changesets: + includePaths: + - "streamlit/*" + - "streamlit/**/*" + - step: + <<: *airflow_dags_deploy + condition: + changesets: + includePaths: + - airflow/dags/* + + # UAT and PROD pipelines will be updated with Streamlit steps once it is tested and ready + UAT: - step: <<: *terraform_plan_apply condition: @@ -52,7 +103,7 @@ pipelines: condition: paths: - snowflake/**/*.sql - prod: + PROD: - step: <<: *terraform_plan_apply condition: diff --git a/lambda/textract_logs_lambda.py b/lambda/textract_logs_lambda.py new file mode 100644 index 0000000..0faa05d --- /dev/null +++ b/lambda/textract_logs_lambda.py @@ -0,0 +1,259 @@ +import json +import snowflake.connector +import boto3 +import snowflake.connector +import logging + + +""" +# Sample input events + +doc_input_event = { + "operation": "insert", + "table": "DOCUMENT_LOGS", + "data": { + "BATCH_ID": 101, + "JOB_ID": "J123456", + "STAGE": "Processing", + "TEXTRACT_STATUS": "Success", + "BUCKET_NAME": "doc-bucket", + "FILE_NAME": "file1.pdf", + "FILE_PATH": "/documents/2023/", + "DOCUMENT_TYPE": "Report", + "PAYER_SIGNED": True, + "PROVIDER_SIGNED": False, + "GROUP_ID": "G100", + "CREATED_TIME": "2024-03-21 10:00:00", + "MODIFIED_TIME": "2024-03-21 10:00:00", + "CREATED_BY": "admin", + "MODIFIED_BY": "admin", + "ORIGINAL_FILE_EXTENSION": "pdf", + "NO_OF_PAGES": 10, + "FILE_SIZE": 1048576 + } +} + + + +doc_update_event = { + "operation": "update", + "table": "DOCUMENT_LOGS", + "data": { + "DOCUMENT_ID": 1001, + "TEXTRACT_STATUS": "Failed", + "MODIFIED_TIME": "2024-03-22 15:00:00", + "MODIFIED_BY": "admin" + } +} + + + +batch_insert_event = { + "operation": "insert", + "table": "BATCH_LOGS", + "data": { + "CLIENT_ID": "C200", + "EXECUTION_START_TIME": "2024-03-21 09:00:00", + "NO_OF_DOCUMENTS": 150, + "USER_NAME": "batch_processor" + } +} + +batch_update_event = { + "operation": "update", + "table": "BATCH_LOGS", + "data": { + "BATCH_ID": 102, + "NO_OF_DOCUMENTS": 155, + "USER_NAME": "updated_processor" + } +} + +client_insert_event = { + "operation": "insert", + "table": "CLIENT_LOGS", + "data": { + "CLIENT_ID": "CL300", + "CLIENT_NAME": "Acme Corporation", + "BUCKET_NAME": "acme-docs" + } +} + +client_update_event = { + "operation": "update", + "table": "CLIENT_LOGS", + "data": { + "CLIENT_ID": "CL300", + "BUCKET_NAME": "new-acme-docs" + } +} +Operation can be: 'insert' or 'update' +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.getLogger("botocore").setLevel(logging.WARNING) +logger = logging.getLogger(__name__) + + +def get_secret(secrets_name: str): + """Get credentials from Secret Manager as dict""" + 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'] + else: + secret_json = base64.b64decode(get_secret_value_response['SecretBinary']) + + return json.loads(secret_json) + + +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) + logger.info(snowflake_connection) + 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 +conn = get_snowflake_db_connection('doczy-dev-db-svc-acc') +cur = conn.cursor() + + +def construct_doc_insert_sql(data): + """ + Constructs the SQL for an insert operation + + Sample return value: + INSERT INTO STG.DOCUMENT_LOGS (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 (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()]) + sql = f"INSERT INTO STG.DOCUMENT_LOGS ({columns}) VALUES ({values});" + return sql + + +def construct_doc_update_sql(data, document_id): + """ + Constructs the SQL for an update operation + + 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()]) + sql = f"UPDATE STG.DOCUMENT_LOGS SET {set_clauses} WHERE DOCUMENT_ID = {document_id};" + return sql + + +def construct_batch_insert_sql(data): + """ + Constructs the SQL for an insert operation + + Sample return value: + INSERT INTO STG.BATCH_LOGS (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()]) + 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 + + 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()]) + sql = f"INSERT INTO STG.CLIENT_LOGS ({columns}) VALUES ({values});" + return sql + + +def construct_client_update_sql(data, client_id): + """ + Constructs the SQL for an update operation + + 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()]) + 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 + + 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()]) + sql = f"UPDATE STG.BATCH_LOGS SET {set_clauses} WHERE BATCH_ID = {batch_id};" + return sql + + +# Main Lambda handler +def lambda_handler(event, context): + # Extract operation type and payload from event + + + operation = event['operation'] # 'insert' or 'update' + data = event['data'] + table = event['table'] + + + try: + if table == 'DOCUMENT_LOGS': + if operation == 'insert': + sql = construct_doc_insert_sql(data) + 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': + sql = construct_batch_insert_sql(data) + elif operation == 'update': + batch_id = data.pop('BATCH_ID', None) + sql = construct_batch_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) + sql = construct_client_update_sql(data, client_id) + else: + raise ValueError("Unsupported table.") + logger.info(f"Executing the following logging SQL statement: {sql}") + cur.execute(sql) + return {'statusCode': 200, 'body': json.dumps('Operation successful')} + + except Exception as 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 + # However, we need the conn to be in hot state for subsequent concurrent executions + # Need to decide the best approach to handle this + pass diff --git a/on_demand_scripts/training_data_pre_processing.py b/on_demand_scripts/training_data_pre_processing.py new file mode 100644 index 0000000..df9ac37 --- /dev/null +++ b/on_demand_scripts/training_data_pre_processing.py @@ -0,0 +1,113 @@ +import pandas as pd +from datetime import datetime + + +def export_column_config(column_names: list): + """ + This function exports the column names and datatypes to a csv file + This will then be ingested to the training data column config + """ + # Create a data frame from the 2 lists and export as csv with the current date and time as filename + + # Create a datatypes list that is all VARCHAR strings equal to the length of the column_names list + try: + column_datatypes = ["VARCHAR" for i in range(len(column_names))] + df = pd.DataFrame(list(zip(column_names, column_datatypes)), columns=["Column_Name", "Data_Type"]) + date = datetime.now() + timestamp = str(date.strftime("%m%d%Y_%H%M%S")) + df.to_csv(f"column_config_{timestamp}.csv", index=False) + return "Column config created successfully" + except Exception as e: + return str(e) + + +def process_xls(file_name: str): + """ + This function processes the master_doczy_db.xlsx file and creates a csv file with the processed data + This will then be ingested to the training data raw table + """ + try: + xl_df = pd.read_excel(file_name, sheet_name="Data Base", header=4) # Passing header as 4 to use sf_col as header + datatypes = xl_df.iloc[0].values.tolist() # grab the datatypes + xl_df2 = xl_df[26:] # Trim the df to remove the first 26 rows where the data is not useful + xl_df2 = xl_df2.reset_index(drop=True) + xl_df2.columns.values[7] = "DOCUMENT_NAME" # works + xl_df2 = xl_df2.iloc[:, 7:] # Drop columns before DOCUMENT_NAME + + start_idx = xl_df2.columns.get_loc('DOCUMENT_NAME') + 1 # +1 because we don't want to drop 'DOCUMENT_NAME' + + # Get index of 'CONTRACT_TITLE' column + end_idx = xl_df2.columns.get_loc('CONTRACT_TITLE') + + # Create a list of column names to drop, which are between 'DOCUMENT_NAME' and 'CONTRACT_TITLE' + cols_to_drop = xl_df2.columns[start_idx:end_idx] + + # Drop the columns + xl_df2.drop(columns=cols_to_drop, inplace=True) + + xl_df2.dropna(axis=1, how='all') + date = datetime.now() + timestamp = str(date.strftime("%m%d%Y_%H%M%S")) + + xl_df2 = xl_df2.loc[:, ~xl_df2.columns.str.startswith('Unnamed')] # Dropping any unnamed columns (Question cols without SF_COL_NAME) + + date = datetime.now() + timestamp = str(date.strftime("%m%d%Y_%H%M%S")) + + xl_df2.columns = xl_df2.columns.str.replace('.', '_', regex=False) # Replace '.' with '_' in column names so that snowflake can ingest + + xl_df2.to_csv(f"processed_training_data-{timestamp}.csv", index=False) + print("Processed training data created successfully") + return xl_df2 + + except Exception as e: + return str(e) + + +def create_business_config_table(file_name: str): + """ + This function creates a business config table from the Business excel file + Where we extact the sf_columns, interrogation question, priority, group_no and theme + """ + try: + xl_df = pd.read_excel(file_name, sheet_name="Data Base", header=2) # Passing header as 4 to use sf_col as header + xl_df = xl_df.iloc[:,12:] # Drop columns before DOCUMENT_NAME + + questions = xl_df.columns.tolist() # Grab the questions that are in the header row + field_name = xl_df.iloc[0].tolist() # Grab the field_name + sf_cols = xl_df.iloc[1].tolist() # Grab the sf_cols + priority = xl_df.iloc[3].tolist() # Grab the priority + group_no = xl_df.iloc[4].tolist() # Grab the group_no + theme = xl_df.iloc[5].tolist() # Grab the theme + + # Create a dataframe from the lists + df_internal = pd.DataFrame({'field_name':field_name,'Column_name': sf_cols, 'Question': questions, 'priority': priority, 'group_no': group_no, 'theme': theme}) + + # Drop rows where the question is 'Unnamed' and the column_name is NaN (Pandas automatically fills NaN with 'Unnamed' when reading excel files depending on the formatting) + df_cleaned = df_internal[~df_internal['Question'].str.contains('Unnamed', na=False) & ~df_internal['Column_name'].isna()] + + date = datetime.now() + timestamp = str(date.strftime("%m%d%Y_%H%M%S")) + + df_cleaned.to_csv(f'biz_config-{timestamp}.csv', index=False) + return "Business config created successfully" + except Exception as e: + return str(e) + + + +def main(): + """ + For this script to work, ensure all rows in the excel file are expanded between the header and actual values. + We need the sf_column name, group no, priority etc to be accessible for ingestion + """ + file_name = "master_doczy_db.xlsx" + xl_df = process_xls(file_name) + status = export_column_config(xl_df.columns.tolist()) + print(status) + status = create_business_config_table(file_name) + print(status) + + +if __name__ == "__main__": + main() diff --git a/snowflake/scripts/config_interface/R__001_PROMPT_CONFIG_TABLE.sql b/snowflake/scripts/config_interface/R__001_PROMPT_CONFIG_TABLE.sql new file mode 100644 index 0000000..b94878b --- /dev/null +++ b/snowflake/scripts/config_interface/R__001_PROMPT_CONFIG_TABLE.sql @@ -0,0 +1,13 @@ +-- Separation of single vs multi prompt using a flag in the table +CREATE TABLE IF NOT EXISTS STG.PROMPT_CONFIG ( + FIELD_NAME VARCHAR, + FIELD_DESC VARCHAR, + IS_REQUIRED BOOLEAN, + FIELD_DATA_TYPE VARCHAR, + PROMPT VARCHAR, + FM_MODEL_ID VARCHAR, + GROUP_ID NUMERIC, + FIELD_MAX_LENGTH NUMBER, + SAMPLE_VALUE VARCHAR, + CONTRACT_SECTION VARCHAR +); diff --git a/snowflake/scripts/config_interface/R__002_CONTRACT_CONFIG_TABLE.sql b/snowflake/scripts/config_interface/R__002_CONTRACT_CONFIG_TABLE.sql new file mode 100644 index 0000000..eba2a95 --- /dev/null +++ b/snowflake/scripts/config_interface/R__002_CONTRACT_CONFIG_TABLE.sql @@ -0,0 +1,16 @@ +-- CREATING THE CONFIG TABLE +CREATE TABLE IF NOT EXISTS STG.CONTRACT_CONFIG ( + CONTRACT_NAME VARCHAR, + BATCH_ID VARCHAR, + UNIQUE_KEY BOOLEAN, + PRICING_BEFORE_CARVEOUTS BOOLEAN, + CONTRACT_RELATED BOOLEAN, + PROVIDER BOOLEAN, + TIMELINE BOOLEAN, + CARVEOUT_INDICATOR BOOLEAN, + CARVEOUT_METHODOLOGY BOOLEAN + REQUEST_USER VARCHAR, + LATEST_FLAG BOOLEAN DEFAULT TRUE, + PIPELINE_KICKOFF_DATETIME DATETIME, + REQUEST_DATETIME DATETIME DEFAULT CURRENT_TIMESTAMP() +); \ No newline at end of file diff --git a/snowflake/scripts/config_interface/R__003_REQUEST_SUBMISSION_TABLE.sql b/snowflake/scripts/config_interface/R__003_REQUEST_SUBMISSION_TABLE.sql new file mode 100644 index 0000000..2107d73 --- /dev/null +++ b/snowflake/scripts/config_interface/R__003_REQUEST_SUBMISSION_TABLE.sql @@ -0,0 +1,6 @@ +CREATE TABLE IF NOT EXISTS STG.REQUEST_SUBMISSION ( + CLIENT_NAME VARCHAR, + BATCH_ID VARCHAR, + REQUEST_USERNAME VARCHAR, + REQUEST_DATETIME DATETIME +); diff --git a/snowflake/scripts/config_interface/R__004_LOAD_CONTRACT_CONFIG_SP.sql b/snowflake/scripts/config_interface/R__004_LOAD_CONTRACT_CONFIG_SP.sql new file mode 100644 index 0000000..802c843 --- /dev/null +++ b/snowflake/scripts/config_interface/R__004_LOAD_CONTRACT_CONFIG_SP.sql @@ -0,0 +1,72 @@ +CREATE OR REPLACE PROCEDURE STG.LOAD_CONTRACT_CONFIG(file_name VARCHAR) +RETURNS STRING +LANGUAGE SQL +EXECUTE AS CALLER +AS +$$ +DECLARE + procedure_name varchar; +BEGIN + + procedure_name := 'LOAD_CONTRACT_CONFIG'; + + call stg.log_audit(:procedure_name, 'Section 1', 99, 'START'); + + -- Create or replace stage with dynamic file name + EXECUTE IMMEDIATE 'CREATE OR REPLACE STAGE STG.CONTRACT_CONFIG_STAGE + STORAGE_INTEGRATION = dev_bucket_integration + URL = ''s3://doczy-dev-infra-raw-data-ingestion/config_interface/' || :file_name || ''' + FILE_FORMAT = (FORMAT_NAME = ''STG.CSV_HEADER'');'; + + call stg.log_audit(:procedure_name, 'Section 1', 99, 'END'); + + call stg.log_audit(:procedure_name, 'Section 2', 99, 'START'); + COPY INTO STG.CONTRACT_CONFIG FROM ( + SELECT + NULLIF(TRIM($1),'') AS CONTRACT_NAME, + NULLIF(TRIM($2),'') AS BATCH_ID + NULLIF(TRIM($3),'') AS UNIQUE_KEY, + NULLIF(TRIM($4),'') AS PRICING_BEFORE_CARVEOUTS, + NULLIF(TRIM($5),'') AS CONTRACT_RELATED, + NULLIF(TRIM($6),'') AS PROVIDER, + NULLIF(TRIM($7),'') AS TIMELINE, + NULLIF(TRIM($8),'') AS CARVEOUT_INDICATOR, + NULLIF(TRIM($9),'') AS CARVEOUT_METHODOLOGY, + NULLIF(TRIM($10),'') AS REQUEST_USER, + NULLIF(TRIM($11),'') AS LATEST_FLAG, + NULLIF(TRIM($12),'') AS PIPELINE_KICKOFF_DATETIME, + NULLIF(TRIM($13),'') AS REQUEST_DATETIME + FROM @STG.CONTRACT_CONFIG_STAGE + ) + FILE_FORMAT = (FORMAT_NAME = 'STG.CSV_HEADER') + ON_ERROR = ABORT_STATEMENT; + + call stg.log_audit(:procedure_name, 'Section 2', 99, 'END'); + + call stg.log_audit(:procedure_name, 'Section 3', 99, 'START'); + + -- Once the new data is loaded to the stage, we need to update the existing records in the table to set the LATEST_FLAG to FALSE + + contract_name := (SELECT NULLIF(TRIM($1),'') FROM @STG.CONTRACT_CONFIG_STAGE LIMIT 1); + batch_id := (SELECT NULLIF(TRIM($2),'') FROM @STG.CONTRACT_CONFIG_STAGE LIMIT 1); + request_datetime := (SELECT NULLIF(TRIM($13),'') FROM @STG.CONTRACT_CONFIG_STAGE LIMIT 1); + + UPDATE TABLE STG.CONTRACT_CONFIG + SET LATEST_FLAG = FALSE + WHERE BATCH_ID = :batch_id AND CONTRACT_NAME = :contract_name AND REQUEST_DATETIME < :request_datetime AND LATEST_FLAG = TRUE; + + call stg.log_audit(:procedure_name, 'Section 3', 99, 'END'); + + call stg.log_audit(:procedure_name, 'Section 4', 99, 'START'); + INSERT INTO DOCZY_DEV.STG.DIM_AUDIT (AUDIT_SID, TABLE_NAME, SOURCE_FILE_NAME, LOAD_DATE, SOURCE_COUNT) + SELECT STG.AUDIT_SID.NEXTVAL, :procedure_name,* + FROM + (SELECT DISTINCT METADATA$FILENAME, CURRENT_TIMESTAMP(), max(METADATA$FILE_ROW_NUMBER) from @STG.CONTRACT_CONFIG_STAGE group by 1,2); + + RETURN 'Setup, Load, and Audit Complete'; + + call stg.log_audit(:procedure_name, 'Section 4', 99, 'END'); + +END; +$$; + diff --git a/snowflake/scripts/config_interface/R__005_LOAD_REQUEST_SUBMISSION_SP.sql b/snowflake/scripts/config_interface/R__005_LOAD_REQUEST_SUBMISSION_SP.sql new file mode 100644 index 0000000..419305b --- /dev/null +++ b/snowflake/scripts/config_interface/R__005_LOAD_REQUEST_SUBMISSION_SP.sql @@ -0,0 +1,57 @@ +CREATE OR REPLACE PROCEDURE STG.LOAD_REQUEST_SUBMISSION(file_name VARCHAR) +RETURNS STRING +LANGUAGE SQL +EXECUTE AS CALLER +AS +$$ +DECLARE + procedure_name varchar; +BEGIN + + procedure_name := 'LOAD_REQUEST_SUBMISSION'; + + call stg.log_audit(:procedure_name, 'Section 1', 99, 'START'); + + -- Create or replace stage with dynamic file name + EXECUTE IMMEDIATE 'CREATE OR REPLACE STAGE STG.REQUEST_SUBMISSION_STAGE + STORAGE_INTEGRATION = dev_bucket_integration + URL = ''s3://doczy-dev-infra-raw-data-ingestion/config_interface/' || :file_name || ''' + FILE_FORMAT = (FORMAT_NAME = ''STG.CSV_HEADER'');'; + + call stg.log_audit(:procedure_name, 'Section 1', 99, 'END'); + + call stg.log_audit(:procedure_name, 'Section 2', 99, 'START'); + + COPY INTO STG.REQUEST_SUBMISSION FROM ( + SELECT + NULLIF(TRIM($2),'') AS CLIENT_NAME, + NULLIF(TRIM($3),'') AS BATCH_ID, + NULLIF(TRIM($4),'') AS REQUEST_USERNAME, + NULLIF(TRIM($5),'') AS REQUEST_DATETIME + + FROM @STG.REQUEST_SUBMISSION_STAGE + ) + FILE_FORMAT = (FORMAT_NAME = 'STG.CSV_HEADER') + ON_ERROR = ABORT_STATEMENT; + + call stg.log_audit(:procedure_name, 'Section 2', 99, 'END'); + + call stg.log_audit(:procedure_name, 'Section 3', 99, 'START'); + + INSERT INTO STG.DIM_AUDIT (AUDIT_SID, TABLE_NAME, SOURCE_FILE_NAME, LOAD_DATE, SOURCE_COUNT) + SELECT STG.AUDIT_SID.NEXTVAL, 'STG.REQUEST_SUBMISSION',* + FROM + (SELECT DISTINCT METADATA$FILENAME, CURRENT_TIMESTAMP(), max(METADATA$FILE_ROW_NUMBER) from @STG.REQUEST_SUBMISSION_STAGE group by 1,2); + + + + RETURN 'Setup, Load, and Audit Complete'; + + call stg.log_audit(:procedure_name, 'Section 3', 99, 'END'); + + + + +END; +$$; + diff --git a/snowflake/scripts/system_wide/R__001_SERVERLESS_LOGS_TABLE.sql b/snowflake/scripts/system_wide/R__001_SERVERLESS_LOGS_TABLE.sql new file mode 100644 index 0000000..63b03cf --- /dev/null +++ b/snowflake/scripts/system_wide/R__001_SERVERLESS_LOGS_TABLE.sql @@ -0,0 +1,44 @@ + + +-- Creating the 'document_logs' table +CREATE TABLE IF NOT EXISTS STG.DOCUMENT_LOGS +( + DOCUMENT_ID VARCHAR NOT NULL, + BATCH_ID VARCHAR NOT NULL, + JOB_ID VARCHAR, + STAGE VARCHAR, + TEXTRACT_STATUS VARCHAR, + BUCKET_NAME VARCHAR NOT NULL, + FILE_NAME VARCHAR NOT NULL, + FILE_PATH VARCHAR, + DOCUMENT_TYPE VARCHAR, + PAYER_SIGNED BOOLEAN, + PROVIDER_SIGNED BOOLEAN, + GROUP_ID VARCHAR, + CREATED_TIME TIMESTAMP, + MODIFIED_TIME TIMESTAMP, + CREATED_BY VARCHAR, + MODIFIED_BY VARCHAR, + ORIGINAL_FILE_EXTENSION VARCHAR, + NO_OF_PAGES NUMERIC, + FILE_SIZE NUMERIC, + PRIMARY KEY (DOCUMENT_ID) +); + +-- Creating the 'batch_logs' table +CREATE TABLE IF NOT EXISTS STG.BATCH_LOGS ( + BATCH_ID VARCHAR NOT NULL, + CLIENT_ID VARCHAR, + EXECUTION_START_TIME TIMESTAMP, + NO_OF_DOCUMENTS NUMBER, + USER_NAME VARCHAR, + PRIMARY KEY (BATCH_ID) +); + +-- Creating the 'client_logs' table +CREATE TABLE IF NOT EXISTS STG.CLIENT_LOGS ( + CLIENT_ID VARCHAR , + CLIENT_NAME VARCHAR, + BUCKET_NAME VARCHAR, + PRIMARY KEY (CLIENT_ID) +); diff --git a/snowflake/scripts/training_interface/R__001_INIT_OBJS.sql b/snowflake/scripts/training_interface/R__001_INIT_OBJS.sql index 7b7bb3c..992ece8 100644 --- a/snowflake/scripts/training_interface/R__001_INIT_OBJS.sql +++ b/snowflake/scripts/training_interface/R__001_INIT_OBJS.sql @@ -17,13 +17,13 @@ CREATE TABLE IF NOT EXISTS STG.TRAINING_RESULTS -- Create file format CREATE FILE FORMAT IF NOT EXISTS STG.CSV_HEADER - TYPE = 'CSV' + TYPE = 'CSV' FIELD_DELIMITER = ',' FILE_EXTENSION = '.csv' RECORD_DELIMITER = '\\n' DATE_FORMAT = AUTO TRIM_SPACE = TRUE - NULL_IF = ('NULL', '', 'N/A','?','~') + NULL_IF = ('NULL', '', 'N/A','?','~','\\N') SKIP_HEADER = 1 EMPTY_FIELD_AS_NULL = TRUE FIELD_OPTIONALLY_ENCLOSED_BY = '"' @@ -31,7 +31,6 @@ CREATE FILE FORMAT IF NOT EXISTS STG.CSV_HEADER -- Create staging table for attempt logs - CREATE TABLE IF NOT EXISTS STG.TRAINING_ATTEMPT_LOGS ( FIELD_NAME VARCHAR, @@ -44,3 +43,24 @@ CREATE TABLE IF NOT EXISTS STG.TRAINING_ATTEMPT_LOGS ); + +-- Create table for column config that will be used by SP to create raw training data table +CREATE TABLE IF NOT EXISTS STG.TRAINING_DATA_COLUMN_CONFIG( +COLUMN_NAME VARCHAR, +COLUMN_DATATYPE VARCHAR +); + +-- Creating a separate file format for the training data +CREATE FILE FORMAT IF NOT EXISTS STG.TRAINING_DATA_FILE_FORMAT + TYPE = 'CSV' + FIELD_DELIMITER = ',' + FILE_EXTENSION = '.csv' + RECORD_DELIMITER = '\\n' + DATE_FORMAT = AUTO + TRIM_SPACE = TRUE + NULL_IF = ('NULL', '', 'N/A','?','~','\\N') + SKIP_HEADER = 1 + EMPTY_FIELD_AS_NULL = TRUE + FIELD_OPTIONALLY_ENCLOSED_BY = '"' + error_on_column_count_mismatch=false + SKIP_BLANK_LINES = TRUE; \ No newline at end of file diff --git a/snowflake/scripts/training_interface/R__004_CREATE_TRAINING_DATA_TABLE_SP.sql b/snowflake/scripts/training_interface/R__004_CREATE_TRAINING_DATA_TABLE_SP.sql new file mode 100644 index 0000000..eeb3dbc --- /dev/null +++ b/snowflake/scripts/training_interface/R__004_CREATE_TRAINING_DATA_TABLE_SP.sql @@ -0,0 +1,28 @@ +CREATE OR REPLACE PROCEDURE STG.CREATE_TRAINING_DATA_TABLE() +RETURNS VARCHAR(16777216) +LANGUAGE SQL +EXECUTE AS CALLER +AS ' +DECLARE + dynamic_ddl STRING := ''CREATE OR REPLACE TABLE STG.TRAINING_DATA_RAW (''; + column_details RESULTSET; + first_column BOOLEAN := TRUE; + cur_config cursor FOR + SELECT column_name, column_datatype FROM STG.TRAINING_DATA_COLUMN_CONFIG; +BEGIN + -- Creating the raw training data table with all columns as VARCHAR due to the dynamic nature of the columns and fields + OPEN cur_config; + FOR rec IN cur_config DO + dynamic_ddl := dynamic_ddl || rec.column_name || '' '' || ''VARCHAR'' || '',''; + END FOR; + + dynamic_ddl := LEFT(dynamic_ddl, LENGTH(dynamic_ddl) - 1); + + dynamic_ddl := dynamic_ddl || '');''; + + EXECUTE IMMEDIATE dynamic_ddl; + + RETURN dynamic_ddl; +END; + +'; \ No newline at end of file diff --git a/snowflake/scripts/training_interface/R__005_LOAD_COLUMN_CONFIG_SP.sql b/snowflake/scripts/training_interface/R__005_LOAD_COLUMN_CONFIG_SP.sql new file mode 100644 index 0000000..5248a83 --- /dev/null +++ b/snowflake/scripts/training_interface/R__005_LOAD_COLUMN_CONFIG_SP.sql @@ -0,0 +1,56 @@ +CREATE OR REPLACE PROCEDURE STG.LOAD_COLUMN_CONFIG(file_name VARCHAR) +RETURNS STRING +LANGUAGE SQL +EXECUTE AS CALLER +AS +$$ +DECLARE + procedure_name varchar; +BEGIN + + procedure_name := 'LOAD_COLUMN_CONFIG'; + + call stg.log_audit(:procedure_name, 'Section 1', 99, 'START'); + + -- Create or replace stage with dynamic file name + EXECUTE IMMEDIATE 'CREATE OR REPLACE STAGE STG.COLUMN_CONFIG_STAGE + STORAGE_INTEGRATION = dev_bucket_integration + URL = ''s3://doczy-dev-infra-raw-data-ingestion/training_data_raw/' || :file_name || ''' + FILE_FORMAT = (FORMAT_NAME = ''STG.CSV_HEADER'');'; + + call stg.log_audit(:procedure_name, 'Section 1', 99, 'END'); + + call stg.log_audit(:procedure_name, 'Section 2', 99, 'START'); + + -- Truncate table as we are using the KILL & FILL approach + TRUNCATE TABLE STG.TRAINING_DATA_COLUMN_CONFIG; + + call stg.log_audit(:procedure_name, 'Section 2', 99, 'END'); + + call stg.log_audit(:procedure_name, 'Section 3', 99, 'START'); + + -- Copy command to load data + COPY INTO STG.TRAINING_DATA_COLUMN_CONFIG FROM ( + SELECT + NULLIF(TRIM($1), '') AS COLUMN_NAME, + NULLIF(TRIM($2), '') AS COLUMN_DATATYPE + FROM @STG.COLUMN_CONFIG_STAGE + ) + FILE_FORMAT = (FORMAT_NAME = 'STG.CSV_HEADER') + ON_ERROR = ABORT_STATEMENT; + + call stg.log_audit(:procedure_name, 'Section 3', 99, 'END'); + + call stg.log_audit(:procedure_name, 'Section 4', 99, 'START'); + + INSERT INTO DOCZY_DEV.STG.DIM_AUDIT (AUDIT_SID, TABLE_NAME, SOURCE_FILE_NAME, LOAD_DATE, SOURCE_COUNT) + SELECT STG.AUDIT_SID.NEXTVAL, :procedure_name,* + FROM + (SELECT DISTINCT METADATA$FILENAME, CURRENT_TIMESTAMP(), max(METADATA$FILE_ROW_NUMBER) from @STG.COLUMN_CONFIG_STAGE group by 1,2); + + call stg.log_audit(:procedure_name, 'Section 4', 99, 'END'); + + RETURN 'Setup, Load, and Audit Complete'; +END; +$$; + diff --git a/snowflake/scripts/training_interface/R__006_LOAD_TRAINING_DATA_RAW_SP.sql b/snowflake/scripts/training_interface/R__006_LOAD_TRAINING_DATA_RAW_SP.sql new file mode 100644 index 0000000..f1bbc08 --- /dev/null +++ b/snowflake/scripts/training_interface/R__006_LOAD_TRAINING_DATA_RAW_SP.sql @@ -0,0 +1,51 @@ +CREATE OR REPLACE PROCEDURE STG.LOAD_TRAINING_DATA_RAW(file_name VARCHAR) +RETURNS STRING +LANGUAGE SQL +EXECUTE AS CALLER +AS +$$ +DECLARE + procedure_name varchar; +BEGIN + + procedure_name := 'LOAD_TRAINING_DATA_RAW'; + + call stg.log_audit(:procedure_name, 'Section 1', 99, 'START'); + + -- Create or replace stage with dynamic file name + EXECUTE IMMEDIATE 'CREATE OR REPLACE STAGE STG.RAW_TRAINING_DATA_STAGE + STORAGE_INTEGRATION = dev_bucket_integration + URL = ''s3://doczy-dev-infra-raw-data-ingestion/training_data_raw/' || :file_name || ''' + FILE_FORMAT = (FORMAT_NAME = ''STG.CSV_HEADER'');'; + + call stg.log_audit(:procedure_name, 'Section 1', 99, 'END'); + + call stg.log_audit(:procedure_name, 'Section 2', 99, 'START'); + + -- Recreating the training data table by calling the SP. This will replace the existing table with new column definitions + call STG.CREATE_TRAINING_DATA_TABLE(); + + call stg.log_audit(:procedure_name, 'Section 2', 99, 'END'); + + call stg.log_audit(:procedure_name, 'Section 3', 99, 'START'); + + -- Copy command to load data + COPY INTO STG.TRAINING_DATA_RAW FROM @STG.RAW_TRAINING_DATA_STAGE + FILE_FORMAT = (FORMAT_NAME = 'STG.TRAINING_DATA_FILE_FORMAT') + ON_ERROR = ABORT_STATEMENT; + + call stg.log_audit(:procedure_name, 'Section 3', 99, 'END'); + + call stg.log_audit(:procedure_name, 'Section 4', 99, 'START'); + + INSERT INTO DOCZY_DEV.STG.DIM_AUDIT (AUDIT_SID, TABLE_NAME, SOURCE_FILE_NAME, LOAD_DATE, SOURCE_COUNT) + SELECT STG.AUDIT_SID.NEXTVAL, :procedure_name,* + FROM + (SELECT DISTINCT METADATA$FILENAME, CURRENT_TIMESTAMP(), max(METADATA$FILE_ROW_NUMBER) from @STG.RAW_TRAINING_DATA_STAGE group by 1,2); + + call stg.log_audit(:procedure_name, 'Section 4', 99, 'END'); + + RETURN 'Setup, Load, and Audit Complete'; +END; +$$; + diff --git a/snowflake/scripts/training_interface/R__007_CREATE_BUSINESS_CONFIG_objects.sql b/snowflake/scripts/training_interface/R__007_CREATE_BUSINESS_CONFIG_objects.sql new file mode 100644 index 0000000..2601408 --- /dev/null +++ b/snowflake/scripts/training_interface/R__007_CREATE_BUSINESS_CONFIG_objects.sql @@ -0,0 +1,8 @@ +CREATE TABLE IF NOT EXISTS STG.BUSINESS_CONFIG ( + FIELD_NAME VARCHAR , + SF_COL_NAME VARCHAR, + QUESTION VARCHAR, + PRIORITY VARCHAR, + GROUP_NO VARCHAR, + THEME VARCHAR +); \ No newline at end of file diff --git a/snowflake/scripts/training_interface/R__008_LOAD_BUSINESS_CONFIG_SP.sql b/snowflake/scripts/training_interface/R__008_LOAD_BUSINESS_CONFIG_SP.sql new file mode 100644 index 0000000..13a750e --- /dev/null +++ b/snowflake/scripts/training_interface/R__008_LOAD_BUSINESS_CONFIG_SP.sql @@ -0,0 +1,61 @@ + +CREATE OR REPLACE PROCEDURE STG.LOAD_BUSINESS_CONFIG(file_name VARCHAR) +RETURNS STRING +LANGUAGE SQL +EXECUTE AS CALLER +AS +$$ +DECLARE + procedure_name varchar; +BEGIN + + procedure_name := 'LOAD_BUSINESS_CONFIG'; + + call stg.log_audit(:procedure_name, 'Section 1', 99, 'START'); + + -- Create or replace stage with dynamic file name + EXECUTE IMMEDIATE 'CREATE OR REPLACE STAGE STG.BUSINESS_CONFIG_STAGE + STORAGE_INTEGRATION = dev_bucket_integration + URL = ''s3://doczy-dev-infra-raw-data-ingestion/training_data_raw/' || :file_name || ''' + FILE_FORMAT = (FORMAT_NAME = ''STG.CSV_HEADER'');'; + + call stg.log_audit(:procedure_name, 'Section 1', 99, 'END'); + + call stg.log_audit(:procedure_name, 'Section 2', 99, 'START'); + + -- Truncate table as we are using the KILL & FILL approach + TRUNCATE TABLE STG.BUSINESS_CONFIG; + + call stg.log_audit(:procedure_name, 'Section 2', 99, 'END'); + + call stg.log_audit(:procedure_name, 'Section 3', 99, 'START'); + + -- Copy command to load data + COPY INTO STG.BUSINESS_CONFIG FROM ( + SELECT + NULLIF(TRIM($1), '') AS FIELD_NAME, + NULLIF(TRIM($2), '') AS SF_COL_NAME, + NULLIF(TRIM($3), '') AS QUESTION, + NULLIF(TRIM($4), '') AS PRIORITY, + NULLIF(TRIM($5), '') AS GROUP_NO, + NULLIF(TRIM($6), '') AS THEME + FROM @STG.BUSINESS_CONFIG_STAGE + ) + FILE_FORMAT = (FORMAT_NAME = 'STG.CSV_HEADER') + ON_ERROR = ABORT_STATEMENT; + + call stg.log_audit(:procedure_name, 'Section 3', 99, 'END'); + + call stg.log_audit(:procedure_name, 'Section 4', 99, 'START'); + + INSERT INTO DOCZY_DEV.STG.DIM_AUDIT (AUDIT_SID, TABLE_NAME, SOURCE_FILE_NAME, LOAD_DATE, SOURCE_COUNT) + SELECT STG.AUDIT_SID.NEXTVAL, :procedure_name,* + FROM + (SELECT DISTINCT METADATA$FILENAME, CURRENT_TIMESTAMP(), max(METADATA$FILE_ROW_NUMBER) from @STG.BUSINESS_CONFIG_STAGE group by 1,2); + + call stg.log_audit(:procedure_name, 'Section 4', 99, 'END'); + + RETURN 'Setup, Load, and Audit Complete'; +END; +$$; + diff --git a/snowflake/scripts/upload_interface/R__001_CLIENT_CONFIG_TABLE.sql b/snowflake/scripts/upload_interface/R__001_CLIENT_CONFIG_TABLE.sql new file mode 100644 index 0000000..42b03e7 --- /dev/null +++ b/snowflake/scripts/upload_interface/R__001_CLIENT_CONFIG_TABLE.sql @@ -0,0 +1,7 @@ +-- Create the table to store the client configuration details for Upload UI & Config UI +CREATE TABLE IF NOT EXISTS STG.CLIENT_CONFIG( + OA_CLIENT_ID NUMERIC, + CLIENT_NAME VARCHAR, + ACTIVE_PROJECT_COUNT NUMERIC, + S3_BUCKET_PATH VARCHAR +); \ No newline at end of file diff --git a/snowflake/scripts/upload_interface/R__002_LOAD_CLIENT_CONFIG_SP.sql b/snowflake/scripts/upload_interface/R__002_LOAD_CLIENT_CONFIG_SP.sql new file mode 100644 index 0000000..f5b268f --- /dev/null +++ b/snowflake/scripts/upload_interface/R__002_LOAD_CLIENT_CONFIG_SP.sql @@ -0,0 +1,58 @@ +CREATE OR REPLACE PROCEDURE STG.LOAD_CLIENT_CONFIG(file_name VARCHAR) +RETURNS STRING +LANGUAGE SQL +EXECUTE AS CALLER +AS +$$ +DECLARE + procedure_name varchar; +BEGIN + + procedure_name := 'LOAD_CLIENT_CONFIG'; + + call stg.log_audit(:procedure_name, 'Section 1', 99, 'START'); + + -- Create or replace stage with dynamic file name + EXECUTE IMMEDIATE 'CREATE OR REPLACE STAGE STG.CLIENT_CONFIG_STAGE + STORAGE_INTEGRATION = dev_bucket_integration + URL = ''s3://doczy-dev-infra-raw-data-ingestion/client_names_openair/' || :file_name || ''' + FILE_FORMAT = (FORMAT_NAME = ''STG.CSV_HEADER'');'; + + call stg.log_audit(:procedure_name, 'Section 1', 99, 'END'); + + call stg.log_audit(:procedure_name, 'Section 2', 99, 'START'); + + -- Truncate table as we are using the KILL & FILL approach + TRUNCATE TABLE STG.CLIENT_CONFIG; + + call stg.log_audit(:procedure_name, 'Section 2', 99, 'END'); + + call stg.log_audit(:procedure_name, 'Section 3', 99, 'START'); + + -- Copy command to load data + COPY INTO STG.CLIENT_CONFIG FROM ( + SELECT + NULLIF(TRIM($1), '') AS OA_CLIENT_ID, + NULLIF(TRIM($2), '') AS CLIENT_NAME, + NULLIF(TRIM($3), '') AS ACTIVE_PROJECT_COUNT, + NULLIF(TRIM($4), '') AS S3_BUCKET_PATH, + FROM @STG.CLIENT_CONFIG_STAGE + ) + FILE_FORMAT = (FORMAT_NAME = 'STG.CSV_HEADER') + ON_ERROR = ABORT_STATEMENT; + + call stg.log_audit(:procedure_name, 'Section 3', 99, 'END'); + + call stg.log_audit(:procedure_name, 'Section 4', 99, 'START'); + + INSERT INTO DOCZY_DEV.STG.DIM_AUDIT (AUDIT_SID, TABLE_NAME, SOURCE_FILE_NAME, LOAD_DATE, SOURCE_COUNT) + SELECT STG.AUDIT_SID.NEXTVAL, :procedure_name,* + FROM + (SELECT DISTINCT METADATA$FILENAME, CURRENT_TIMESTAMP(), max(METADATA$FILE_ROW_NUMBER) from @STG.CLIENT_CONFIG_STAGE group by 1,2); + + call stg.log_audit(:procedure_name, 'Section 4', 99, 'END'); + + RETURN 'Setup, Load, and Audit Complete'; +END; +$$; + diff --git a/snowflake/scripts/upload_interface/R__003_CONTRACT_UPLOAD_LOGS.sql b/snowflake/scripts/upload_interface/R__003_CONTRACT_UPLOAD_LOGS.sql new file mode 100644 index 0000000..ac7126e --- /dev/null +++ b/snowflake/scripts/upload_interface/R__003_CONTRACT_UPLOAD_LOGS.sql @@ -0,0 +1,10 @@ +-- This table is used to store the logs of all the files that were uploaded by the user to S3 +-- Following fields are to be stored in the table: +-- Batch id & client name & file name, date time, user +CREATE TABLE IF NOT EXISTS STG.CONTRACT_UPLOAD_LOGS ( + BATCH_ID VARCHAR, + CLIENT_NAME VARCHAR, + FILE_NAME VARCHAR, + UPLOAD_DATETIME DATETIME DEFAULT CURRENT_TIMESTAMP(), + UPLOAD_USER VARCHAR +); \ No newline at end of file diff --git a/streamlit/constants.py b/streamlit/constants.py new file mode 100644 index 0000000..6f1285c --- /dev/null +++ b/streamlit/constants.py @@ -0,0 +1,190 @@ +import os + +# from dotenv import load_dotenv +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 + + + +# load_dotenv() +# ROOT_DIRECTORY = os.path.dirname(os.path.realpath(__file__)) +ROOT_DIRECTORY = "\\\\amznfsxuofkyi1z.aarete.local\\SharedFiles\\AArete Client Work\\Modahealth\\Restricted\\Moda Growth\\Artificial Intelligence\\DEFAXXER_20231207" + +# Define the folder for source and output +SOURCE_DIRECTORY = "SOURCE_DOCUMENTS" +OUTPUT_DIRECTORY = f"{ROOT_DIRECTORY}\\Output" + +PERSIST_DIRECTORY = 'DB' + +MODELS_PATH = "C:\\Users\\Public\\models" + +# Can be changed to a specific number +INGEST_THREADS = os.cpu_count() or 8 + +# Define the Chroma settings +CHROMA_SETTINGS = Settings( + anonymized_telemetry=False, + is_persistent=True, + allow_reset=True, +) + +# Context Window and Max New Tokens +CONTEXT_WINDOW_SIZE = 4096 +MAX_NEW_TOKENS = CONTEXT_WINDOW_SIZE # int(CONTEXT_WINDOW_SIZE/4) + +#### If you get a "not enough space in the buffer" error, you should reduce the values below, start with half of the original values and keep halving the value until the error stops appearing + +N_GPU_LAYERS = 100 # Llama-2-70B has 83 layers +N_BATCH = 512 + +### From experimenting with the Llama-2-7B-Chat-GGML model on 8GB VRAM, these values work: +# N_GPU_LAYERS = 20 +# N_BATCH = 512 + + +# https://python.langchain.com/en/latest/_modules/langchain/document_loaders/excel.html#UnstructuredExcelLoader +DOCUMENT_MAP = { + ".txt": TextLoader, + ".md": UnstructuredMarkdownLoader, + ".py": TextLoader, + # ".pdf": PDFMinerLoader, + ".pdf": UnstructuredFileLoader, + ".csv": CSVLoader, + ".xls": UnstructuredExcelLoader, + ".xlsx": UnstructuredExcelLoader, + ".docx": Docx2txtLoader, + ".doc": Docx2txtLoader, +} + +# Default Instructor Model +EMBEDDING_MODEL_NAME = "hkunlp/instructor-large" # Uses 1.5 GB of VRAM (High Accuracy with lower VRAM usage) + +#### +#### OTHER EMBEDDING MODEL OPTIONS +#### + +# EMBEDDING_MODEL_NAME = "hkunlp/instructor-xl" # Uses 5 GB of VRAM (Most Accurate of all models) +# EMBEDDING_MODEL_NAME = "intfloat/e5-large-v2" # Uses 1.5 GB of VRAM (A little less accurate than instructor-large) +# EMBEDDING_MODEL_NAME = "intfloat/e5-base-v2" # Uses 0.5 GB of VRAM (A good model for lower VRAM GPUs) +# EMBEDDING_MODEL_NAME = "all-MiniLM-L6-v2" # Uses 0.2 GB of VRAM (Less accurate but fastest - only requires 150mb of vram) + +#### +#### MULTILINGUAL EMBEDDING MODELS +#### + +# EMBEDDING_MODEL_NAME = "intfloat/multilingual-e5-large" # Uses 2.5 GB of VRAM +# EMBEDDING_MODEL_NAME = "intfloat/multilingual-e5-base" # Uses 1.2 GB of VRAM + + +#### SELECT AN OPEN SOURCE LLM (LARGE LANGUAGE MODEL) +# Select the Model ID and model_basename +# load the LLM for generating Natural Language responses + +#### GPU VRAM Memory required for LLM Models (ONLY) by Billion Parameter value (B Model) +#### Does not include VRAM used by Embedding Models - which use an additional 2GB-7GB of VRAM depending on the model. +#### +#### (B Model) (float32) (float16) (GPTQ 8bit) (GPTQ 4bit) +#### 7b 28 GB 14 GB 7 GB - 9 GB 3.5 GB - 5 GB +#### 13b 52 GB 26 GB 13 GB - 15 GB 6.5 GB - 8 GB +#### 32b 130 GB 65 GB 32.5 GB - 35 GB 16.25 GB - 19 GB +#### 65b 260.8 GB 130.4 GB 65.2 GB - 67 GB 32.6 GB - - 35 GB + +# MODEL_ID = "TheBloke/Llama-2-7B-Chat-GGML" +# MODEL_BASENAME = "llama-2-7b-chat.ggmlv3.q4_0.bin" + +#### +#### (FOR GGUF MODELS) +#### + +# MODEL_ID = "TheBloke/Llama-2-13b-Chat-GGUF" +# MODEL_BASENAME = "llama-2-13b-chat.Q4_K_M.gguf" + +MODEL_ID = "TheBloke/Llama-2-7b-Chat-GGUF" +MODEL_BASENAME = "llama-2-7b-chat.Q4_K_M.gguf" + +# MODEL_ID = "TheBloke/Mistral-7B-Instruct-v0.1-GGUF" +# MODEL_BASENAME = "mistral-7b-instruct-v0.1.Q8_0.gguf" + +# MODEL_ID = "TheBloke/Llama-2-70b-Chat-GGUF" +# MODEL_BASENAME = "llama-2-70b-chat.Q4_K_M.gguf" + +#### +#### (FOR HF MODELS) +#### + +# MODEL_ID = "NousResearch/Llama-2-7b-chat-hf" +# MODEL_BASENAME = None +# MODEL_ID = "TheBloke/vicuna-7B-1.1-HF" +# MODEL_BASENAME = None +# MODEL_ID = "TheBloke/Wizard-Vicuna-7B-Uncensored-HF" +# MODEL_ID = "TheBloke/guanaco-7B-HF" +# MODEL_ID = 'NousResearch/Nous-Hermes-13b' # Requires ~ 23GB VRAM. Using STransformers +# alongside will 100% create OOM on 24GB cards. +# llm = load_model(device_type, model_id=model_id) + +#### +#### (FOR GPTQ QUANTIZED) Select a llm model based on your GPU and VRAM GB. Does not include Embedding Models VRAM usage. +#### + +##### 48GB VRAM Graphics Cards (RTX 6000, RTX A6000 and other 48GB VRAM GPUs) ##### + +### 65b GPTQ LLM Models for 48GB GPUs (*** With best embedding model: hkunlp/instructor-xl ***) +# MODEL_ID = "TheBloke/guanaco-65B-GPTQ" +# MODEL_BASENAME = "model.safetensors" +# MODEL_ID = "TheBloke/Airoboros-65B-GPT4-2.0-GPTQ" +# MODEL_BASENAME = "model.safetensors" +# MODEL_ID = "TheBloke/gpt4-alpaca-lora_mlp-65B-GPTQ" +# MODEL_BASENAME = "model.safetensors" +# MODEL_ID = "TheBloke/Upstage-Llama1-65B-Instruct-GPTQ" +# MODEL_BASENAME = "model.safetensors" + +##### 24GB VRAM Graphics Cards (RTX 3090 - RTX 4090 (35% Faster) - RTX A5000 - RTX A5500) ##### + +### 13b GPTQ Models for 24GB GPUs (*** With best embedding model: hkunlp/instructor-xl ***) +# MODEL_ID = "TheBloke/Wizard-Vicuna-13B-Uncensored-GPTQ" +# MODEL_BASENAME = "Wizard-Vicuna-13B-Uncensored-GPTQ-4bit-128g.compat.no-act-order.safetensors" +# MODEL_ID = "TheBloke/vicuna-13B-v1.5-GPTQ" +# MODEL_BASENAME = "model.safetensors" +# MODEL_ID = "TheBloke/Nous-Hermes-13B-GPTQ" +# MODEL_BASENAME = "nous-hermes-13b-GPTQ-4bit-128g.no-act.order" +# MODEL_ID = "TheBloke/WizardLM-13B-V1.2-GPTQ" +# MODEL_BASENAME = "gptq_model-4bit-128g.safetensors + +### 30b GPTQ Models for 24GB GPUs (*** Requires using intfloat/e5-base-v2 instead of hkunlp/instructor-large as embedding model ***) +# MODEL_ID = "TheBloke/Wizard-Vicuna-30B-Uncensored-GPTQ" +# MODEL_BASENAME = "Wizard-Vicuna-30B-Uncensored-GPTQ-4bit--1g.act.order.safetensors" +# MODEL_ID = "TheBloke/WizardLM-30B-Uncensored-GPTQ" +# MODEL_BASENAME = "WizardLM-30B-Uncensored-GPTQ-4bit.act-order.safetensors" + +##### 8-10GB VRAM Graphics Cards (RTX 3080 - RTX 3080 Ti - RTX 3070 Ti - 3060 Ti - RTX 2000 Series, Quadro RTX 4000, 5000, 6000) ##### +### (*** Requires using intfloat/e5-small-v2 instead of hkunlp/instructor-large as embedding model ***) + +### 7b GPTQ Models for 8GB GPUs +# MODEL_ID = "TheBloke/Wizard-Vicuna-7B-Uncensored-GPTQ" +# MODEL_BASENAME = "Wizard-Vicuna-7B-Uncensored-GPTQ-4bit-128g.no-act.order.safetensors" +# MODEL_ID = "TheBloke/WizardLM-7B-uncensored-GPTQ" +# MODEL_BASENAME = "WizardLM-7B-uncensored-GPTQ-4bit-128g.compat.no-act-order.safetensors" +# MODEL_ID = "TheBloke/wizardLM-7B-GPTQ" +# MODEL_BASENAME = "wizardLM-7B-GPTQ-4bit.compat.no-act-order.safetensors" + +#### +#### (FOR GGML) (Quantized cpu+gpu+mps) models - check if they support llama.cpp +#### + +# MODEL_ID = "TheBloke/wizard-vicuna-13B-GGML" +# MODEL_BASENAME = "wizard-vicuna-13B.ggmlv3.q4_0.bin" +# MODEL_BASENAME = "wizard-vicuna-13B.ggmlv3.q6_K.bin" +# MODEL_BASENAME = "wizard-vicuna-13B.ggmlv3.q2_K.bin" +# MODEL_ID = "TheBloke/orca_mini_3B-GGML" +# MODEL_BASENAME = "orca-mini-3b.ggmlv3.q4_0.bin" + +#### +#### (FOR AWQ QUANTIZED) Select a llm model based on your GPU and VRAM GB. Does not include Embedding Models VRAM usage. +### (*** MODEL_BASENAME is not actually used but have to contain .awq so the correct model loading is used ***) +### (*** Compute capability 7.5 (sm75) and CUDA Toolkit 11.8+ are required ***) +#### +# MODEL_ID = "TheBloke/Llama-2-7B-Chat-AWQ" +# MODEL_BASENAME = "model.safetensors.awq" diff --git a/streamlit/ingest.py b/streamlit/ingest.py new file mode 100644 index 0000000..c5e5833 --- /dev/null +++ b/streamlit/ingest.py @@ -0,0 +1,219 @@ +import logging +import os +from concurrent.futures import ProcessPoolExecutor, ThreadPoolExecutor, as_completed + +import click +import torch +from langchain.docstore.document import Document +from langchain.embeddings import HuggingFaceInstructEmbeddings +from langchain.text_splitter import Language, RecursiveCharacterTextSplitter +from langchain.vectorstores import Chroma +import uuid + +from constants import ( + CHROMA_SETTINGS, + DOCUMENT_MAP, + EMBEDDING_MODEL_NAME, + INGEST_THREADS, + PERSIST_DIRECTORY, + SOURCE_DIRECTORY, +) + +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") + +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] + except Exception as ex: + file_log('%s loading error: \n%s' % (file_path, ex)) + return None + +def load_document_batch(filepaths): + logging.info("Loading document batch") + # create a thread pool + with ThreadPoolExecutor(len(filepaths)) as exe: + # load files + 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 + else: + 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]: + # Loads all documents from the source documents directory, including nested folders + paths = [] + for root, _, files in os.walk(source_dir): + for file_name in files: + 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(): + paths.append(source_file_path) + + # Have at least one worker and at most INGEST_THREADS workers + n_workers = min(INGEST_THREADS, max(len(paths), 1)) + chunksize = round(len(paths) / n_workers) + docs = [] + with ProcessPoolExecutor(n_workers) as executor: + futures = [] + # split the load operations into chunks + for i in range(0, len(paths), chunksize): + # select a chunk of filenames + filepaths = paths[i : (i + chunksize)] + # submit the task + try: + future = executor.submit(load_document_batch, filepaths) + except Exception as ex: + file_log('executor task failed: %s' % (ex)) + future = None + if future is not None: + futures.append(future) + # process all results + for future in as_completed(futures): + # open the file and load the data + try: + contents, _ = future.result() + docs.extend(contents) + except Exception as ex: + file_log('Exception: %s' % (ex)) + + return docs + + +def split_documents(documents: list[Document]) -> tuple[list[Document], list[Document]]: + # Splits documents for correct Text Splitter + 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) + 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] + +@click.command() +@click.option( + "--device_type", + default="cuda" if torch.cuda.is_available() else "cpu", + type=click.Choice( + [ + "cpu", + "cuda", + "ipu", + "xpu", + "mkldnn", + "opengl", + "opencl", + "ideep", + "hip", + "ve", + "fpga", + "ort", + "xla", + "lazy", + "vulkan", + "mps", + "meta", + "hpu", + "mtia", + ], + ), + 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: + text = infile.read() + 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: + 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) + 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"Split into {len(texts)} chunks of text") + + # Create embeddings + # embeddings = HuggingFaceInstructEmbeddings( + # model_name=EMBEDDING_MODEL_NAME, + # model_kwargs={"device": device_type}, + # ) + bedrock_runtime = boto3.client( + service_name="bedrock-runtime", + region_name="us-east-1", + aws_access_key_id="ASIAZTMXAXNXD3TOUOAJ", + aws_secret_access_key="pk7k69CqXZPB/bf2hdFsW+47D5WYkoWXFdxQ633X", + aws_session_token="IQoJb3JpZ2luX2VjELj//////////wEaCXVzLWVhc3QtMiJIMEYCIQCqcyLSwLeN9RX6tz+TgB4VMabHYwPcA3bTzx7xZ5b75QIhAMiOKbDlmPZ+eHkeskzelK4M9gWtS7GqlzJq5qmWbbrOKo4DCKH//////////wEQABoMNjYwMTMxMDY4NzgyIgzPuv0vReC65Mo/bQcq4gLXtII+VTaxES+JrjHHWpWywdmpsJneN6bcB3U37z7J8BFA5aaUqkATKmwJI0brmk5ZJKL8SpgBEC2TNdA/V9nzbMlf2HPunhEv6OPLzSWp5iJqaNL945MP764CbkYvfN9QWd6durUv1WgGZRNcbMzXg2UFsxKcRql795vtOmL207+R7uIouWl73So7NaCkEgaj4FdEJ9lbnfvWFeNcBlHbjwUx8e9EJjwm8D60OkTdS4w7Q3EacoEKLO94/kp2RtsaggAUV33OcvO/32VwYJzhRJYuveQUZnIzfkmybGtrXkkWLqMO9pls1bkTmIjaeMwcL8Uo7oowR9sTFCT87rY711yIYBGVOjkN9mfavPH4FSCNeeI6ta5aXqa7iVDU7rPerpFtle2i1VGfTW5bKoaWxjO13IdqI4yI9Pgibl+FuVeRB7md2tDuS6SMAuX4qpWnMVDLKCItlAxXlZdmVhuDFf8wr+7rrgY6pQH02EIrvCkYYXbJd8u7t78hOam3lSNPz+nghQHA5ppl7TnyRBL5/9Rfqp5Ib8y54HdC4hbzz/7w6lvfS0QbHE3Q6r2GC6XWk9L4FAxD8pyZkIYbtwf/WueX1g0r+0x7uLCxKxbWsYum/bigyvNxsDyFIcQcUOm+OJsXGPH3Z7z+qWKni5N9RRA4SEEpdSTRkAWipDtgVNII6kd/eWvyE74d+LYUrmA=" + ) + embeddings = BedrockEmbeddings( + client=bedrock_runtime, + model_id="amazon.titan-embed-text-v1", + ) + + """ + db = Chroma.from_documents( + texts, + embeddings, + persist_directory=PERSIST_DIRECTORY, + client_settings=CHROMA_SETTINGS, + ) + """ + + # for batch_texts in process_in_batches(texts, 20000): # https://github.com/PromtEngineer/localGPT/issues/489y + print(filename) + + db = Chroma.from_documents( + texts, + embeddings, + persist_directory=PERSIST_DIRECTORY, + client_settings=CHROMA_SETTINGS, + 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 + ) + main() diff --git a/streamlit/interface_1.py b/streamlit/interface_1.py new file mode 100644 index 0000000..a02d1c4 --- /dev/null +++ b/streamlit/interface_1.py @@ -0,0 +1,129 @@ +import streamlit as st +from streamlit_extras.add_vertical_space import add_vertical_space +import os +import streamlit as st +import pandas as pd +from io import StringIO +from datetime import datetime +import boto3 +import util + +REDIRECT_URI = 'http://172.29.20.126:8501' +user_list = ['maamseek@aarete.com', 'smahdavian@aarete.com', 'ahinge@aarete.com', 'akadam@aarete.com' +, 'piragavarapu@aarete.com', 'umistry@aarete.com'] +st.set_page_config(layout = "wide") +# # Sidebar contents +# with st.sidebar: +# st.title("Doczy.AI ™") +# st.markdown( +# """ +# ## About +# This app extracts data from contracts + +# """ +# ) +# add_vertical_space(15) +# # st.write("Doczy") + + +util.setup_page(REDIRECT_URI) +if st.session_state.user_info['mail'] in user_list: + + s3_client = boto3.client('s3', + region_name="us-east-1", + ) + objects = s3_client.list_objects_v2(Bucket='doczy-dev-infra-textract' + , Prefix="batches/batch_1/", Delimiter='/') + + folder_list = [] + for prefix in objects['CommonPrefixes']: + folder_list.append(prefix['Prefix'][:-1].split('/')[-1]) + + client_row = st.columns([0.1, 0.8]) + with client_row[0]: + st.write("**Client Name**") + with client_row[1]: + client = st.selectbox('Client Name',(folder_list), index=7, label_visibility = "collapsed") + + folder_objects = s3_client.list_objects_v2(Bucket='doczy-dev-infra-textract' + , Prefix="batches/batch_1/"+client+"/", Delimiter='/') + + folder_list_2 = [] + for prefix in folder_objects['CommonPrefixes']: + folder_list_2.append(prefix['Prefix'][:-1].split('/')[-1]) + + path_row = st.columns([0.1, 0.8]) + with path_row[0]: + st.write("**Path to folder**") + with path_row[1]: + Directory = st.selectbox('**Path to folder**', folder_list_2, label_visibility = "collapsed") + + checks = st.columns([0.1, 0.12, 0.12, 0.12, 0.12, 0.12, 0.12, 0.12]) + with checks[0]: + st.write("**Group No.**") + with checks[1]: + a = st.checkbox('Unique Key', key = str(1)) + with checks[2]: + b = st.checkbox('Pricing Before Carveouts', key = str(2)) + with checks[3]: + c = st.checkbox('Contract Related', key = str(3)) + with checks[4]: + d = st.checkbox('Provider', key = str(4)) + with checks[5]: + e = st.checkbox('Timeline', key = str(5)) + with checks[6]: + f = st.checkbox('Carveout Indicator', key = str(6)) + with checks[7]: + g = st.checkbox('Carveout Methodology', key = str(7)) + + add_vertical_space(1) + + df = pd.DataFrame(columns=['Request ID','Contract ID','Contract Name','Unique Key','Pricing Before Carveouts' + , 'Contract Related', 'Provider', 'Timeline', 'Carveout Indicator', 'Carveout Methodology']) + file_list = [] + file_objects = s3_client.list_objects_v2(Bucket='doczy-dev-infra-textract' + , Prefix="batches/batch_1/"+client+"/"+Directory+"/", Delimiter='/') + + if st.button("Read the contracts from Path"): + for obj in file_objects.get('Contents',[]): + if not obj['Key'].endswith('/'): + file_list.append(obj['Key'].split('/')[-1]) + + df['Contract Name'] = file_list + df['Request ID'] = range(len(file_list)) + df['Contract ID'] = file_list + df['Unique Key'] = a + df['Pricing Before Carveouts'] = b + df['Contract Related'] = c + df['Provider'] = d + df['Timeline'] = e + df['Carveout Indicator'] = f + df['Carveout Methodology'] = g + dir_path = os.path.dirname(os.path.realpath(__file__)) + print(f'DEBUGGING: PWD= {dir_path}') + df.to_csv('temp1.csv', index=False) + + add_vertical_space(1) + + # df_copy = df.set_index(df.columns[0]).copy() + df2 = pd.read_csv('temp1.csv') + edited_df = st.data_editor(df2) + + + @st.cache_data + def convert_df(df): + return df.to_csv(index=False).encode('utf-8') + + csv = convert_df(edited_df) + + buttons = st.columns([0.8, 0.2]) + with buttons[0]: + st.download_button("Download Table", csv, "file.csv", "text/csv", key='download-csv') + with buttons[1]: + st.button("Run Doczy.AI Pipeline") + +else: + st.write("Access Denied") + + + diff --git a/streamlit/interface_2.py b/streamlit/interface_2.py new file mode 100644 index 0000000..8ccbfa2 --- /dev/null +++ b/streamlit/interface_2.py @@ -0,0 +1,206 @@ +import json + +import boto3 +from langchain.prompts import PromptTemplate +from langchain.embeddings.bedrock import BedrockEmbeddings +from langchain.llms.bedrock import Bedrock +from langchain_community.vectorstores import Chroma +from constants import CHROMA_SETTINGS, EMBEDDING_MODEL_NAME, PERSIST_DIRECTORY, MODEL_ID, MODEL_BASENAME, SOURCE_DIRECTORY +from langchain.chains import RetrievalQA + +import streamlit as st +from streamlit_extras.add_vertical_space import add_vertical_space +import os +import pandas as pd +import util + +REDIRECT_URI = 'http://172.29.20.126:8502' +user_list = ['maamseek@aarete.com', 'smahdavian@aarete.com', 'ahinge@aarete.com', 'akadam@aarete.com' +, 'piragavarapu@aarete.com', 'umistry@aarete.com', 'ahutchison@aarete.com', 'bgrunst@aarete.com', 'ddimeglio@aarete.com' +, 'vnair@aarete.com'] + +st.set_page_config(layout = "wide") +# Sidebar contents +with st.sidebar: + st.title("Doczy.AI ™") + st.markdown( + """ + ## About + This app extracts data from contracts + + """ + ) + add_vertical_space(15) + # st.write("Doczy") + +util.setup_page(REDIRECT_URI) +if st.session_state.user_info['mail'] in user_list: + + fields = pd.read_csv('contract_fields.csv', encoding='unicode_escape', skipinitialspace=True) + fields = fields[fields['PRIORITY'] == 'A'] + fields = fields.drop_duplicates(subset='Field Name', keep="first").sort_values('Field Name') + fields['Interrogation Question?'] = fields['Interrogation Question?'].fillna(' ') + field_prompt_mapping = dict(zip(fields['Field Name'], fields['Interrogation Question?'])) + + def file_selector(folder_path=SOURCE_DIRECTORY): + filenames = os.listdir(folder_path) + selected_filename = st.selectbox('Select a file', filenames, label_visibility = "collapsed") + # return os.path.join(folder_path, selected_filename) + return selected_filename + + file_row = st.columns([0.2, 0.7, 0.1]) + with file_row[0]: + st.write("**Contract Name**") + with file_row[1]: + # file_name = st.text_input("**Contract Name**", label_visibility = "collapsed") + file_name = file_selector() + + # lob_row = st.columns([0.2, 0.7, 0.1]) + # with lob_row[0]: + # st.write("**LOB**") + # with lob_row[1]: + # lob = st.selectbox('LOB',('Medicare', 'Medicaid'), label_visibility = "collapsed") + + llm_row = st.columns([0.2, 0.7, 0.1]) + with llm_row[0]: + st.write("**Langauge Model**") + with llm_row[1]: + llm_selected = st.selectbox('Langauge Model',('Llama 2 Chat 13B', 'Llama 2 Chat 70B', 'Titan Text Express'), label_visibility = "collapsed") + + page_list = [] + with open(os.path.join(SOURCE_DIRECTORY, file_name), 'r') as infile: + text = infile.read() + page_count = text.count('Start of Page No. = ') + for page in range(page_count+1): + file_path = "SOURCE_DOCUMENTS\\" + f'{file_name[:-4]}_page{page}.txt' + dict_with_pages = { 'source': { '$eq': file_path }} + page_list.append(dict_with_pages) + + # AWS_ACCESS_KEY_ID = os.getenv('AWS_ACCESS_KEY_ID') + # AWS_SECRET_ACCESS_KEY = os.getenv('AWS_SECRET_ACCESS_KEY') + # AWS_SESSION_TOKEN=os.getenv('AWS_SESSION_TOKEN') + + # Setup bedrock + bedrock_runtime = boto3.client( + service_name="bedrock-runtime", + region_name="us-east-1" + ) + + embeddings = BedrockEmbeddings( + client=bedrock_runtime, + model_id="amazon.titan-embed-text-v1", + ) + DB = Chroma( + persist_directory=PERSIST_DIRECTORY, + embedding_function=embeddings, + client_settings=CHROMA_SETTINGS, + ) + RETRIEVER = DB.as_retriever(search_kwargs={"filter": {"$or": page_list}, "k": 4}) + + if llm_selected == 'Titan Text Express': + LLM = Bedrock( + model_id="amazon.titan-text-express-v1", + client=bedrock_runtime, + model_kwargs={ + "maxTokenCount": 4096, + "stopSequences": [], + "temperature": 0, + "topP": 1, + } + ) + elif llm_selected == 'Llama 2 Chat 70B': + LLM = Bedrock( + model_id="meta.llama2-70b-chat-v1", + client=bedrock_runtime, + model_kwargs={ + "max_gen_len": 512, + "temperature": 0, + # "topP": 0.9, + } + ) + else: + LLM = Bedrock( + model_id="meta.llama2-13b-chat-v1", + client=bedrock_runtime, + model_kwargs={ + "max_gen_len": 512, + "temperature": 0, + # "topP": 0.9, + } + ) + + template = """ + + Use the following pieces of context to answer the question at the end. If you don't know the answer,\ + just say that you don't know, don't try to make up an answer. + + {context} + + Question: {question} + Answer:""" + prompt = PromptTemplate(input_variables=["context", "question"], template=template) + + QA = RetrievalQA.from_chain_type( + llm=LLM, + chain_type="stuff", + retriever=RETRIEVER, + return_source_documents=True, + chain_type_kwargs={"prompt": prompt}, + ) + + # query = "In which state or states is the Contract applicable? Answer in one or two words. State name: " + # response = QA({"query":query}) + # st.write(query) + # st.write(response['result']) + # st.write("-----------") + # st.write(response) + + # clicked = st.button("Show Results") + df = pd.DataFrame(columns=['Contract Name','Field Name','Snippet','Page Number','Confidence Level', + 'Field Extracted Value','Imputed Value']) + field_list = list(field_prompt_mapping.keys()) + query_list = [field_prompt_mapping[x] for x in field_list] + score_list = [DB.similarity_search_with_relevance_scores(query, k=4, filter={"$or": page_list}) for query in query_list] + confidence_list = [] + for score in score_list: + confidence_list.append(max(d[1] for d in score)) + # st.write(confidence_list) + + if st.button("Show Results"): + response_list = [QA({"query":query}) for query in query_list] + answer_list = [response['result'] for response in response_list] + doc_list = [response['source_documents'] for response in response_list] + snippet_list = [str(doc[0].page_content) for doc in doc_list] + page_no_list = [int(str(doc[0].metadata["source"]).rsplit('_page')[1].replace('.txt',''))+1 for doc in doc_list] + + df['Field Name'] = field_list + df['Contract Name'] = file_name + df['Snippet'] = snippet_list + df['Page Number'] = page_no_list + df['Confidence Level'] = confidence_list + df['Field Extracted Value'] = answer_list + df.to_csv('temp2.csv', index=False) + + df2 = pd.read_csv('temp2.csv') + df2['Imputed Value'] = '' + edited_df = st.data_editor(df2) + + @st.cache_data + def convert_df(df): + return df.to_csv(index=False).encode('utf-8') + + csv = convert_df(edited_df) + + buttons = st.columns(3) + with buttons[0]: + st.button("Save All Imputations") + with buttons[1]: + st.download_button("Download Table", csv, "file.csv", "text/csv", key='download-csv') + with buttons[2]: + st.button("Kickoff Database Integration") + +else: + st.write("Access Denied") + + + diff --git a/streamlit/interface_3.py b/streamlit/interface_3.py new file mode 100644 index 0000000..aa49067 --- /dev/null +++ b/streamlit/interface_3.py @@ -0,0 +1,370 @@ +import json +import boto3 +from langchain.prompts import PromptTemplate +from langchain.embeddings.bedrock import BedrockEmbeddings +from langchain.llms.bedrock import Bedrock +from langchain_community.vectorstores import Chroma +from constants import CHROMA_SETTINGS, EMBEDDING_MODEL_NAME, PERSIST_DIRECTORY, MODEL_ID, MODEL_BASENAME, SOURCE_DIRECTORY +from langchain.chains import RetrievalQA + +import streamlit as st +from streamlit_extras.add_vertical_space import add_vertical_space + +import pandas as pd +from datetime import datetime +import random +import os +import dateutil +import util + +REDIRECT_URI = 'http://172.29.20.126:8503' +user_list = ['maamseek@aarete.com', 'smahdavian@aarete.com', 'ahinge@aarete.com', 'akadam@aarete.com' +, 'piragavarapu@aarete.com', 'umistry@aarete.com', 'ahutchison@aarete.com', 'bgrunst@aarete.com', 'ddimeglio@aarete.com' +, 'vnair@aarete.com'] + +st.set_page_config(layout = "wide") +# Sidebar contents +with st.sidebar: + st.title("Doczy.AI ™") + st.markdown( + """ + ## About + This app extracts data from contracts + + """ + ) + add_vertical_space(15) + # st.write("Doczy") + +util.setup_page(REDIRECT_URI) +if st.session_state.user_info['mail'] in user_list: + + fields = pd.read_csv('contract_fields.csv', encoding='unicode_escape', skipinitialspace=True) + fields = fields[fields['PRIORITY'].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") + + 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") + + 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'])] + + with seed_row[0]: + if contract_count in ['10', '20', '30', '50']: + st.write("**Seed Value**") + elif contract_count == '1': + st.write("**Contract Name**") + with seed_row[1]: + if contract_count in ['10', '20', '30', '50']: + seed_value = st.text_input("**Seed Value**", value = 20, label_visibility = "collapsed") + random.seed(seed_value) + contract_list = sorted(random.choices(os.listdir(SOURCE_DIRECTORY), k=int(contract_count))) + elif contract_count == '1': + contract_name = st.selectbox('Contract Name', (contract_list), label_visibility = "collapsed") + contract_list = [contract_name] + + + llm_row = st.columns([0.15, 0.45, 0.4]) + with llm_row[0]: + st.write("**Langauge Model**") + with llm_row[1]: + llm_selected = st.selectbox('Langauge Model',('Claude 2', 'Claude Instant', 'Llama 2 Chat 13B', 'Llama 2 Chat 70B' + , 'Titan Text Express'), label_visibility = "collapsed") + + st.write("**Prompt**") + sequence_input = field_prompt_mapping.get(field) + prompt_row = st.columns([0.8, 0.2]) + with prompt_row[1]: + if st.button("Clear Prompt"): + sequence_input = '' + if st.button("Back to default"): + prompt = sequence_input + st.button("Save Prompt") + with prompt_row[0]: + prompt = st.text_area("**Prompt**", sequence_input, height = 150, label_visibility = "collapsed") + + page_list_all = [] + for contract in contract_list: + page_list = [] + with open(os.path.join(SOURCE_DIRECTORY, contract[:-4]+'.txt'), 'r') as infile: + text = infile.read() + page_count = text.count('Start of Page No. = ') + for page in range(page_count+1): + file_path = "SOURCE_DOCUMENTS\\" + f'{contract[:-4]}_page{page}.txt' + dict_with_pages = { 'source': { '$eq': file_path }} + page_list.append(dict_with_pages) + page_list_all.append(page_list) + contract_txt_mapping = dict(zip(contract_list, page_list_all)) + + column_name = fields.loc[fields['Field Name'] == field, 'SF_DB_COL_NAME'].iloc[0] + column_list = ['(internal) Document Name', '(Internal) Carveout ID', column_name] + if column_name+'_PG' in list(field_values.columns): + column_list.append(column_name+'_PG') + field_values = field_values[column_list] + field_values.rename(columns={'(internal) Document Name': 'Contract Name', column_name: 'Actual Value Stored' + , '(Internal) Carveout ID': 'Contract ID', column_name+'_PG': 'Original Page Number'}, inplace=True) + field_values = field_values.drop_duplicates(subset='Contract Name', keep="first").sort_values('Contract Name') + + # Setup bedrock + bedrock_runtime = boto3.client( + service_name="bedrock-runtime", + region_name="us-east-1", + ) + + # Define the retreiver + # load the vectorstore + if "EMBEDDINGS" not in st.session_state: + EMBEDDINGS = BedrockEmbeddings( + client=bedrock_runtime, + model_id="amazon.titan-embed-text-v1", + ) + st.session_state.EMBEDDINGS = EMBEDDINGS + + if "DB" not in st.session_state: + DB = Chroma( + persist_directory=PERSIST_DIRECTORY, + embedding_function=st.session_state.EMBEDDINGS, + client_settings=CHROMA_SETTINGS, + ) + st.session_state.DB = DB + + # if "RETRIEVER" not in st.session_state: + # # { "source": { '$eq': "SOURCE_DOCUMENTS\\A.1_UH_Health_System_eff_2_1_08 (1)_page0.txt"} } + # RETRIEVER = DB.as_retriever(search_kwargs={"filter": { "source": { '$eq': "SOURCE_DOCUMENTS\\A.1_UH_Health_System_eff_2_1_08 (1)_page0.txt"} }, "k": 2}) + # st.session_state.RETRIEVER = RETRIEVER + + # if "LLM" not in st.session_state: + if llm_selected == 'Titan Text Express': + LLM = Bedrock( + model_id="amazon.titan-text-express-v1", + client=bedrock_runtime, + model_kwargs={ + "maxTokenCount": 512, + "stopSequences": [], + "temperature": 0, + "topP": 1, + } + ) + elif llm_selected == 'Llama 2 Chat 70B': + LLM = Bedrock( + model_id="meta.llama2-70b-chat-v1", + client=bedrock_runtime, + model_kwargs={ + "max_gen_len": 512, + "temperature": 0, + # "topP": 0.9, + } + ) + elif llm_selected == 'Llama 2 Chat 13B': + LLM = Bedrock( + model_id="meta.llama2-13b-chat-v1", + client=bedrock_runtime, + model_kwargs={ + "max_gen_len": 512, + "temperature": 0, + # "topP": 0.9, + } + ) + elif llm_selected == 'Claude Instant': + LLM = Bedrock( + model_id="anthropic.claude-instant-v1", + client=bedrock_runtime, + model_kwargs={ + # "max_tokens_to_sample": 512, + "temperature": 0, + # "topP": 0.9, + } + ) + elif llm_selected == 'Claude 2': + LLM = Bedrock( + model_id="anthropic.claude-v2:1", + client=bedrock_runtime, + model_kwargs={ + # "max_tokens_to_sample": 512, + "temperature": 0, + # "topP": 0.9, + } + ) + st.session_state["LLM"] = LLM + + # if "QA" not in st.session_state: + # prompt, memory = model_memory() + + # QA = RetrievalQA.from_chain_type( + # llm=LLM, + # chain_type="stuff", + # retriever=RETRIEVER, + # return_source_documents=True, + # chain_type_kwargs={"prompt": prompt, "memory": memory}, + # ) + # st.session_state["QA"] = QA + + # df = pd.DataFrame(columns=['Contract Name','Contract ID','Actual Value Stored','New Extracted value','Confidence Level','Snippet', + # 'Original Page Number', 'New Page Number', 'Revised Prompt', 'Result']) + df = pd.DataFrame(columns=['Contract Name','Raw value','New Extracted value','Confidence Level','Snippet1','Snippet2','Snippet3' + ,'Snippet4','Snippet5','Snippet6','Snippet7','Snippet8','Snippet9','Snippet10','New Page Number' + , 'Revised Prompt', 'Result']) + try: + history = pd.read_csv('history.csv') + except: + history = pd.DataFrame(columns=['Field Name','# Contracts Tested', 'Username', 'Date/Time', 'Accuracy', 'Attempt #']) + attempt = 0 + + if llm_selected in ['Llama 2 Chat 13B', 'Llama 2 Chat 70B']: + k_value = 10 + else: + k_value = 20 + + + if st.button("Test Configuration"): + + answer_list = [] + doc_list = [] + response_list = [] + score_list = [] + attempt = attempt + 1 + + for page_list in page_list_all: + RETRIEVER = st.session_state.DB.as_retriever(search_kwargs={"filter": {"$or": page_list}, "k": k_value}) + QA = RetrievalQA.from_chain_type( + llm=st.session_state["LLM"], + chain_type="stuff", + retriever=RETRIEVER, + return_source_documents=True, + # chain_type_kwargs={"prompt": prompt, "memory": None}, + ) + score = st.session_state.DB.similarity_search_with_relevance_scores(prompt, k=4, filter={"$or": page_list}) + score_list.append(max(d[1] for d in score)) + response = QA(prompt) + answer, docs = response["result"], response["source_documents"] + answer_list.append(answer) + doc_list.append(docs) + response_list.append(response) + + df['Raw value'] = answer_list + # post-processing + if 'Date' in field: + date_list = [] + for answer in answer_list: + try: + extracted_date = dateutil.parser.parse(str(answer).replace('"',''), fuzzy=True).date() + except: + extracted_date = " " + date_list.append(extracted_date) + answer_list = date_list + elif llm_selected in ['Llama 2 Chat 13B', 'Llama 2 Chat 70B']: + answer_list = [answer.rstrip(".") for answer in answer_list] + answer_list = [answer if "I don't know" not in str(answer) else " " for answer in answer_list] + answer_list = [answer if "N/A" not in str(answer) else " " for answer in answer_list] + answer_list = [answer if "does not contain" not in str(answer) else " " for answer in answer_list] + answer_list = [answer if "None" not in str(answer) else " " for answer in answer_list] + answer_list = [answer if "Not specified in the contract" not in str(answer) else " " for answer in answer_list] + answer_list = [answer if "Not applicable" not in str(answer) else " " for answer in answer_list] + elif llm_selected in ['Claude 2', 'Claude Instant']: + answer_list = [answer if "Unfortunately, I do not have enough context" not in str(answer) else " " for answer in answer_list] + answer_list = [answer.rstrip(".") for answer in answer_list] + else: + answer_list = [answer.rstrip(".") for answer in answer_list] + # answer_list = [str(x).rsplit(':',1)[0] if len(str(x).rsplit(':',1)) < 2 else str(x).rsplit(':',1)[1] for x in answer_list] + + df['Contract Name'] = contract_list + # to be deleted later + df['Contract Name'] = [contract.replace(' MU','').replace('_MU','').replace('.txt','') for contract in contract_list] + + df['New Extracted value'] = answer_list + df['Confidence Level'] = [round(score, 2) for score in score_list] + Snippet = [] + count = 0 + for i in range(int(contract_count)): + for j in range(10): + try: + content = str(doc_list[i][j].page_content) + except: + content = " " + Snippet.append(content) + # df['Snippet1'] = [str(doc[0].page_content) for doc in doc_list] + df['Snippet1'] = Snippet[:int(contract_count)] + df['Snippet2'] = Snippet[int(contract_count):2*int(contract_count)] + df['Snippet3'] = Snippet[2*int(contract_count):3*int(contract_count)] + df['Snippet4'] = Snippet[3*int(contract_count):4*int(contract_count)] + df['Snippet5'] = Snippet[4*int(contract_count):5*int(contract_count)] + df['Snippet6'] = Snippet[5*int(contract_count):6*int(contract_count)] + df['Snippet7'] = Snippet[6*int(contract_count):7*int(contract_count)] + df['Snippet8'] = Snippet[7*int(contract_count):8*int(contract_count)] + df['Snippet9'] = Snippet[8*int(contract_count):9*int(contract_count)] + df['Snippet10'] = Snippet[9*int(contract_count):] + df['New Page Number'] = [int(str(doc[0].metadata["source"]).rsplit('_page')[1].replace('.txt',''))+1 for doc in doc_list] + df['Revised Prompt'] = [prompt] * len(contract_list) + + df = pd.merge(df, field_values, how ='left', on ='Contract Name') + + answer_list = list(df['New Extracted value']) + df['Actual Value Stored'] = pd.to_datetime(df['Actual Value Stored'],errors='coerce').dt.date + df.fillna(" ", inplace=True) + actual_value_list = list(df['Actual Value Stored']) + result_list = [i==j for i, j in zip(actual_value_list, answer_list)] + df['Result'] = [str(x) for x in result_list] + df = df[~df['Contract ID'].isnull()] + if 'Original Page Number' in df.columns: + df = df[['Contract Name','Contract ID','Actual Value Stored','Raw value','New Extracted value','Confidence Level' + ,'Snippet1','Snippet2','Snippet3','Snippet4','Snippet5','Snippet6','Snippet7','Snippet8','Snippet9','Snippet10' + ,'Original Page Number', 'New Page Number', 'Revised Prompt', 'Result']] + else: + df = df[['Contract Name','Contract ID','Actual Value Stored','Raw value','New Extracted value','Confidence Level' + ,'Snippet1','Snippet2','Snippet3','Snippet4','Snippet5','Snippet6','Snippet7','Snippet8','Snippet9','Snippet10' + , 'New Page Number', 'Revised Prompt', 'Result']] + + try: + accuracy = round(sum(bool(x) for x in result_list) * 100 / len(list(df['Result'])), 2) + except: + accuracy = 'NA' + + history.loc[len(history.index)] = [field, str(contract_count), None, datetime.now().strftime("%Y-%m-%d %H:%M:%S"), accuracy, attempt] + # df.to_csv("RESULTS\\"+field.replace("?","").replace("/","_")+'-'+llm_selected+'.csv', index=False) + history.to_csv('history.csv', index=False) + + # df_copy = df.set_index(df.columns[0]).copy() + # df_2_copy = history.set_index(history.columns[0]).copy() + st.dataframe(df) + st.dataframe(history) + + + # @st.cache_data + # def convert_df(df): + # return df.to_csv(index=False).encode('utf-8') + + # csv = convert_df(edited_df) + + # buttons = st.columns(3) + # with buttons[0]: + # st.button("Save All Imputations") + # with buttons[1]: + # st.download_button("Download Table", csv, "file.csv", "text/csv", key='download-csv') + # with buttons[2]: + # st.button("Kickoff Database Integration") + + + st.write(column_name) + st.write(len(contract_list)) + +else: + st.write("Access Denied") + + diff --git a/streamlit/requirements.txt b/streamlit/requirements.txt new file mode 100644 index 0000000..94f93cf --- /dev/null +++ b/streamlit/requirements.txt @@ -0,0 +1,39 @@ +# Natural Language Processing +langchain==0.0.267 +chromadb==0.4.6 +pdfminer.six==20221105 +InstructorEmbedding +sentence-transformers==2.2.2 +faiss-cpu +huggingface_hub +transformers +autoawq +protobuf==3.20.2; sys_platform != 'darwin' +protobuf==3.20.2; sys_platform == 'darwin' and platform_machine != 'arm64' +protobuf==3.20.3; sys_platform == 'darwin' and platform_machine == 'arm64' +auto-gptq==0.2.2 +docx2txt +unstructured +unstructured[pdf] + +# Utilities +urllib3==1.26.6 +accelerate +bitsandbytes ; sys_platform != 'win32' +bitsandbytes-windows ; sys_platform == 'win32' +click +flask +requests + +# Streamlit related +streamlit +Streamlit-extras + +# Excel File Manipulation +openpyxl +numpy>=1.22.2 + +# AWS related +boto3 +awscli + diff --git a/streamlit/security.py b/streamlit/security.py new file mode 100644 index 0000000..1337f58 --- /dev/null +++ b/streamlit/security.py @@ -0,0 +1,41 @@ +import streamlit as st +import msal +import requests + +# Replace with your own values +CLIENT_ID = 'effafe90-7ed7-43a3-ab03-19a0be2f1758' +CLIENT_SECRET = 'bjQ8Q~lpR2uBcGI34VDu16t73doz8Crj0YY_~dgD' +# TENANT_ID = '' + +AUTHORITY = 'https://login.microsoftonline.com/organizations/' +SCOPE = ['User.Read'] +# REDIRECT_URI = 'http://localhost:8501' + + +app = msal.ConfidentialClientApplication(CLIENT_ID, authority=AUTHORITY, client_credential=CLIENT_SECRET) + + +def get_auth_url(REDIRECT_URI): + auth_url = app.get_authorization_request_url(SCOPE, redirect_uri=REDIRECT_URI) + return auth_url + + +def get_token_from_code(auth_code, REDIRECT_URI): + app = msal.ConfidentialClientApplication(CLIENT_ID, authority=AUTHORITY, client_credential=CLIENT_SECRET) + result = app.acquire_token_by_authorization_code(auth_code, scopes=SCOPE, redirect_uri=REDIRECT_URI) + return result['access_token'] + + +def get_user_info(access_token): + headers = {'Authorization': f'Bearer {access_token}'} + response = requests.get('https://graph.microsoft.com/v1.0/me', headers=headers) + return response.json() + + +def handle_redirect(REDIRECT_URI): + if not st.session_state.get('access_token'): + code = st.query_params.get('code') + if code: + access_token = get_token_from_code(code, REDIRECT_URI) + st.session_state['access_token'] = access_token + st.session_state \ No newline at end of file diff --git a/streamlit/util.py b/streamlit/util.py new file mode 100644 index 0000000..7e4565a --- /dev/null +++ b/streamlit/util.py @@ -0,0 +1,24 @@ + +import streamlit as st +import security + +def setup_page(REDIRECT_URI): + # st.set_page_config( + # page_title=page_title, + # page_icon="👋", + # ) + + if st.query_params.get('code'): + security.handle_redirect(REDIRECT_URI) + + access_token = st.session_state.get('access_token') + + if access_token: + user_info = security.get_user_info(access_token) + st.session_state['user_info'] = user_info + return True + else: + st.write("Please sign-in to use this app.") + auth_url = security.get_auth_url(REDIRECT_URI) + st.markdown(f"Sign In", unsafe_allow_html=True) + st.stop() \ No newline at end of file