556 lines
26 KiB
Python
556 lines
26 KiB
Python
import json
|
|
import boto3
|
|
import pandas as pd
|
|
import numpy as np
|
|
from datetime import datetime
|
|
import os
|
|
import anthropic
|
|
import re
|
|
import snowflake.connector
|
|
from io import StringIO
|
|
import logging
|
|
from urllib.parse import unquote_plus
|
|
from botocore.exceptions import ClientError
|
|
import traceback
|
|
import file_processing
|
|
|
|
# Global configuration variables
|
|
S3_REGION = "us-east-2"
|
|
BEDROCK_REGION = "us-east-1"
|
|
SNOWFLAKE_SCHEMA = "STG"
|
|
MAX_PROMPT_LENGTH_PER_GROUP = 20000 # Adjust as needed
|
|
LLM_SELECTED = "Claude 3 - Sonnet"
|
|
SNOWFLAKE_CONN_SECRET_NAME = "doczy-dev-db-svc-acc"
|
|
|
|
# Get variable ENVIRONMENT from env variables using os
|
|
env = os.environ.get('ENVIRONMENT')
|
|
# Get first character from env to get env name abbreviation
|
|
env = env[0]
|
|
# Get the bucket name based on the env that is linked with the storage integration
|
|
snowflake_ingestion_bucket = f'doczyai-use2-{env}-infra-s3-raw-data-ingestion'
|
|
|
|
# Configure logging
|
|
logger = logging.getLogger()
|
|
logger.setLevel(logging.INFO)
|
|
pd.options.mode.chained_assignment = None
|
|
|
|
# Initialize AWS clients
|
|
logger.info("Initializing AWS clients.")
|
|
sqs = boto3.client('sqs')
|
|
s3_client = boto3.client('s3', region_name=S3_REGION)
|
|
bedrock_runtime = boto3.client('bedrock-runtime', region_name=BEDROCK_REGION)
|
|
|
|
def get_secret():
|
|
|
|
secret_name = SNOWFLAKE_CONN_SECRET_NAME
|
|
region_name = "us-east-2"
|
|
|
|
# Create a Secrets Manager client
|
|
session = boto3.session.Session()
|
|
client = session.client(
|
|
service_name='secretsmanager',
|
|
region_name=region_name
|
|
)
|
|
|
|
try:
|
|
get_secret_value_response = client.get_secret_value(
|
|
SecretId=secret_name
|
|
)
|
|
except ClientError as e:
|
|
# For a list of exceptions thrown, see
|
|
# https://docs.aws.amazon.com/secretsmanager/latest/apireference/API_GetSecretValue.html
|
|
raise e
|
|
|
|
secret = get_secret_value_response['SecretString']
|
|
return secret
|
|
|
|
# Snowflake connection
|
|
logger.info("Establishing Snowflake connection.")
|
|
sf_secrets = json.loads(get_secret())
|
|
snowflake_conn = snowflake.connector.connect(
|
|
user=sf_secrets.get('user'),
|
|
password=sf_secrets.get('password'),
|
|
account=sf_secrets.get('account_locator'),
|
|
role=sf_secrets.get('role'),
|
|
warehouse=sf_secrets.get('warehouse'),
|
|
database=sf_secrets.get('database'),
|
|
schema=SNOWFLAKE_SCHEMA
|
|
)
|
|
|
|
|
|
|
|
|
|
def save_to_sf(dag_name, **kwargs):
|
|
|
|
mwaa_env_name = 'doczy-dev-infra-mwaa'
|
|
dag_name = dag_name
|
|
mwaa_cli_command = 'dags trigger'
|
|
|
|
# Create the client with the specified profile
|
|
session = boto3.Session()
|
|
client = session.client('mwaa', region_name='us-east-2')
|
|
|
|
# get web token
|
|
mwaa_cli_token = client.create_cli_token(
|
|
Name=mwaa_env_name
|
|
)
|
|
|
|
conn = http.client.HTTPSConnection(mwaa_cli_token['WebServerHostname'])
|
|
|
|
# This section passes the payload to the MWAA CLI
|
|
# The file parameters should be added dynamically in streamlit, once the file names are passed while triggering the dag, the data will be ingested
|
|
# training_results_file = "training_results_sample.csv"
|
|
# attempt_logs_file = "attempt_logs_sample.csv"
|
|
# conf = "{\"" + "training_results_file_name" + "\":\"" + {training_results_file} + "\", \"" + "attempt_logs_file_name" + "\":\"" + {attempt_logs_file} + "\"}".format(training_results_file=training_results_file, attempt_logs_file=attempt_logs_file)
|
|
|
|
conf = json.dumps(kwargs)
|
|
|
|
payload = mwaa_cli_command + " " + dag_name + " --conf '{}'".format(conf)
|
|
headers = {
|
|
'Authorization': 'Bearer ' + mwaa_cli_token['CliToken'],
|
|
'Content-Type': 'text/plain'
|
|
}
|
|
conn.request("POST", "/aws_mwaa/cli", payload, headers)
|
|
res = conn.getresponse()
|
|
data = res.read()
|
|
dict_str = data.decode("UTF-8")
|
|
mydata = ast.literal_eval(dict_str)
|
|
return payload
|
|
|
|
|
|
def parse_field_groups(field_groups):
|
|
# Split the string by comma and strip any surrounding whitespace
|
|
groups = [group.strip() for group in field_groups.split(',')]
|
|
# Convert the list to a tuple
|
|
return tuple(groups)
|
|
|
|
def lambda_handler(event, context):
|
|
try:
|
|
logger.info("Processing SQS event.")
|
|
records = json.loads(event['Records'][0]['body'])
|
|
logger.info(f"PROCESSING RECORDS:: {records}")
|
|
# Process each message from the SQS event
|
|
for record in records['Records']:
|
|
# Extract the message body from the record
|
|
logger.info(f"RECORD:: {record}")
|
|
|
|
sqs_record = record['s3']
|
|
bucket_name = sqs_record['bucket']['name']
|
|
object_key = unquote_plus(sqs_record['object']['key'])
|
|
document_id = get_filename_from_path(object_key)
|
|
document_id = os.path.splitext(document_id)[0]
|
|
|
|
logger.info(f"Processing document: {document_id}")
|
|
|
|
# Fetch batch_id from object tags
|
|
batch_id = get_s3_object_tags(bucket_name, object_key).get('BatchId')
|
|
logger.info(f"Batch ID: {batch_id}")
|
|
|
|
|
|
# Get file name from s3 bucket object key
|
|
try:
|
|
logger.info(f"OBJECT KEY:: {object_key}")
|
|
# Get file name from s3 bucket object key and remove file extension
|
|
file_name = object_key.split('/')[-1].split('.')[0]
|
|
cursor = snowflake_conn.cursor()
|
|
result = cursor.execute(f"select group_id from stg.document_logs where batch_id= '{batch_id}' and document_id='{file_name}' limit 1;")
|
|
logger.info(f"QUERY BEING EXECUTED: select group_id from stg.document_logs where batch_id= '{batch_id}' and file_name='{file_name}' limit 1;")
|
|
|
|
|
|
for rec in result:
|
|
logger.info(f"QUERY RESULT::: {rec[0]}")
|
|
field_groups = rec[0]
|
|
|
|
# logger.info(f"QUERY RESULT::: {cursor.fetchone()} & {cursor.fetchone()[0]}&fetching all {cursor.fetchall()}")
|
|
# field_groups = cursor.fetchone()[0]
|
|
# logger.info(f"RESULT FIELD GROUP:: {field_groups}")
|
|
# Check if field group is a tuple, if not convert it to tuple
|
|
field_groups = parse_field_groups(field_groups)
|
|
logger.info(f"RESULT FIELD GROUP after tuple conversion:: {field_groups}")
|
|
|
|
except Exception as e:
|
|
logger.error(f"Error fetching field groups: {e} stopping execution")
|
|
exit()
|
|
field_groups = ('A','C')
|
|
|
|
|
|
# field_groups = ('A','C')
|
|
# field_groups = 'B'
|
|
# logger.info(f"Field groups: {field_groups}")
|
|
if field_groups == ('B',):
|
|
data = s3_client.get_object(Bucket=bucket_name, Key=object_key)
|
|
contents = data['Body'].read()
|
|
context = contents.decode("utf-8")
|
|
# Get the file name from the s3 object prefix
|
|
filename = object_key.split('/')[-1]
|
|
process_b_fields(filename, context, bucket_name,batch_id)
|
|
# Prepare and process the prompt if field groups are A and C or Just A or just C
|
|
elif field_groups == ('A','C'):
|
|
main(bucket_name, object_key, field_groups)
|
|
# Process scenario where all field groups are processed
|
|
elif field_groups == ('A','B','C'):
|
|
main(bucket_name, object_key, ('A','C'))
|
|
data = s3_client.get_object(Bucket=bucket_name, Key=object_key)
|
|
contents = data['Body'].read()
|
|
context = contents.decode("utf-8")
|
|
# Get the file name from the s3 object prefix
|
|
filename = object_key.split('/')[-1]
|
|
process_b_fields(filename, context, bucket_name,batch_id)
|
|
elif field_groups == ('A',):
|
|
main(bucket_name, object_key, field_groups)
|
|
elif field_groups == ('C',):
|
|
main(bucket_name, object_key, field_groups)
|
|
elif field_groups == ('B','A') or field_groups == ('B','C') or field_groups == ('A','B') or field_groups == ('C','B'):
|
|
main(bucket_name, object_key, ('A','C'))
|
|
data = s3_client.get_object(Bucket=bucket_name, Key=object_key)
|
|
contents = data['Body'].read()
|
|
context = contents.decode("utf-8")
|
|
# Get the file name from the s3 object prefix
|
|
filename = object_key.split('/')[-1]
|
|
process_b_fields(filename, context, bucket_name,batch_id)
|
|
else:
|
|
logger.error(f"Unsupported field groups: {field_groups}")
|
|
return {
|
|
'statusCode': 500,
|
|
'body': f"Unsupported field groups: {field_groups}"
|
|
}
|
|
|
|
return {
|
|
'statusCode': 200,
|
|
'body': json.dumps('Message processed successfully')
|
|
}
|
|
except Exception as e:
|
|
logger.error(f"Unexpected error: {e}")
|
|
return {
|
|
'statusCode': 500,
|
|
'body': f"{str(e)} \n {traceback.format_exc()}"
|
|
}
|
|
|
|
|
|
def process_b_fields(filename, context, bucket, batch_id):
|
|
# Create tuple from filename and context
|
|
item = (filename, context)
|
|
df = file_processing.process_file(item)
|
|
# Adding batch_id to the last column of df
|
|
df['Batch_id'] = batch_id
|
|
csv_buf = StringIO()
|
|
df.to_csv(csv_buf, header=True, index=False)
|
|
csv_buf.seek(0)
|
|
timestamp = datetime.now().strftime("%Y%m%d%H%M%S")
|
|
file_name = f"results_{batch_id}_B_{timestamp}.csv"
|
|
s3_client.put_object(Bucket=bucket, Body=csv_buf.getvalue(), Key=f"final_output/{batch_id}/{file_name}")
|
|
csv_buf = StringIO()
|
|
df.to_csv(csv_buf, header=True, index=False)
|
|
csv_buf.seek(0)
|
|
|
|
# Saving data to snowflake raw data ingestion bucket & calling stored proc
|
|
s3_client.put_object(Bucket=snowflake_ingestion_bucket, Body=csv_buf.getvalue(), Key=f"doczy_pipeline_output/{file_name}")
|
|
cur = snowflake_conn.cursor()
|
|
load_data_query = f"CALL LOAD_DOCZY_PIPELINE_RAW_OUTPUT_B_FIELDS('{file_name}')"
|
|
cur.execute(load_data_query)
|
|
# log df shape as output
|
|
logger.info(f"B fields Processed file: {filename} with shape: {df.shape}")
|
|
|
|
|
|
|
|
def get_filename_from_path(full_path):
|
|
return os.path.basename(full_path)
|
|
|
|
def get_s3_object_tags(bucket_name, object_key):
|
|
try:
|
|
response = s3_client.get_object_tagging(Bucket=bucket_name, Key=object_key)
|
|
tags_list = response['TagSet']
|
|
tags_dict = {tag['Key']: tag['Value'] for tag in tags_list}
|
|
return tags_dict
|
|
except Exception as e:
|
|
logger.error(f"Error retrieving tags: {e}")
|
|
return None
|
|
|
|
def read_from_db(query):
|
|
try:
|
|
logger.info("Executing query on Snowflake.")
|
|
cur = snowflake_conn.cursor()
|
|
cur.execute(query)
|
|
df = pd.DataFrame.from_records(iter(cur), columns=[x[0] for x in cur.description])
|
|
logger.info(f"Query executed successfully. Retrieved {len(df)} records.")
|
|
return df
|
|
except Exception as e:
|
|
logger.error(f"Error reading from database: {e}")
|
|
return pd.DataFrame()
|
|
|
|
def prepare_prompt(*field_group):
|
|
""" function takes field group e.g. A, B as input
|
|
and returns prompt in json format e.g. {"EFFECTIVE_DT": "What is the effective date of contract", ...}
|
|
"""
|
|
|
|
# fetch prompts from database
|
|
try:
|
|
query = 'select * from "PROMPT_CONFIG"'
|
|
fields = read_from_db(query)
|
|
fields.rename(columns={'PROMPT': 'Interrogation Question?'}, inplace = True)
|
|
fields.rename(columns={'GROUP_ID': 'PRIORITY'}, inplace = True)
|
|
fields.rename(columns={'FIELD_NAME': 'Field Name'}, inplace = True)
|
|
fields['SF_DB_COL_NAME'] = fields['Field Name']
|
|
except:
|
|
fields = pd.read_csv('contract_fields.csv', encoding='utf-8-sig', skipinitialspace=True)
|
|
logger.info("Local copy of BUSINESS_CONFIG table loaded")
|
|
fields.rename(columns={'QUESTION': 'Interrogation Question?'}, inplace = True)
|
|
fields.rename(columns={'SF_COL_NAME': 'SF_DB_COL_NAME'}, inplace = True)
|
|
fields['Field Name'] = fields['SF_DB_COL_NAME']
|
|
fields = fields[~fields['SF_DB_COL_NAME'].str.endswith('_PG', na=None)]
|
|
fields = fields[fields['PRIORITY'].isin(field_group)]
|
|
fields['Interrogation Question?'] = fields['Interrogation Question?'].fillna(' ')
|
|
prompt_dict = dict(zip(fields['Field Name'], fields['Interrogation Question?']))
|
|
|
|
# create prompts to get page number
|
|
prompt_dict_pg = prompt_dict | {str(k)+'_PG': "On which page can I find answer to the question - "+str(
|
|
v) for k, v in prompt_dict.items()}
|
|
|
|
# print(f"PROMPT CONFIG RETURNED: {fields}")
|
|
question = json.dumps(dict(sorted(prompt_dict_pg.items())))
|
|
|
|
return question
|
|
|
|
def invoke_llm(context, question, llm_selected=LLM_SELECTED):
|
|
logger.info("Invoking LLM.")
|
|
prompt_data = f"""
|
|
|
|
Human: Use the following pieces of context to provide a concise answer to the questions 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. You must answer in correct JSON format.
|
|
|
|
{context}
|
|
|
|
Question: {question}
|
|
|
|
Assistant: Answer in JSON format: {{"""
|
|
|
|
if llm_selected == "Claude 3 - Sonnet":
|
|
model_id = 'anthropic.claude-3-sonnet-20240229-v1:0'
|
|
body = json.dumps({
|
|
"anthropic_version": "bedrock-2023-05-31",
|
|
"max_tokens": 4096,
|
|
"messages": [
|
|
{
|
|
"role": "user",
|
|
"content": [
|
|
{
|
|
"type": "text",
|
|
"text": anthropic.HUMAN_PROMPT + prompt_data + anthropic.AI_PROMPT
|
|
}
|
|
]
|
|
}
|
|
],
|
|
"temperature": 0.0
|
|
})
|
|
else:
|
|
raise ValueError("Unsupported LLM selected")
|
|
|
|
try:
|
|
response = bedrock_runtime.invoke_model(
|
|
body=body,
|
|
modelId=model_id,
|
|
accept="application/json",
|
|
contentType="application/json"
|
|
)
|
|
response_body = json.loads(response.get("body").read())
|
|
response_text = response_body['content'][0]['text']
|
|
# logger.info(f"LLM response: {response_text}")
|
|
except Exception as e:
|
|
logger.error(f"Error invoking LLM: {e}")
|
|
response_text = "failed"
|
|
return response_text.strip()
|
|
|
|
def json_parsing(response_text, context):
|
|
logger.info("Parsing JSON response.")
|
|
try:
|
|
if response_text.split("{", 1)[1].strip()[0] == '"':
|
|
response_text = "{" + response_text.split("{", 1)[1]
|
|
else:
|
|
response_text = "{" + response_text
|
|
except:
|
|
response_text = "{" + response_text
|
|
|
|
if len(response_text.split("}", 1)) > 1:
|
|
if response_text.rsplit("}", 1)[0].strip()[-1] == '"':
|
|
response_text = response_text.rsplit("}", 1)[0] + "}"
|
|
elif response_text.rsplit("}", 1)[0].strip()[-1] == '}':
|
|
response_text = response_text.rsplit("}", 1)[0]
|
|
else:
|
|
response_text = response_text.rstrip(",") + "}"
|
|
|
|
try:
|
|
response_dict = json.loads(response_text)
|
|
except:
|
|
response_dict = {"Test field": "Failed to extract"}
|
|
|
|
field_l = list(response_dict.keys())
|
|
answer_l = list(response_dict.values())
|
|
field_dict = {k: v for k, v in response_dict.items() if not k.endswith('_PG')}
|
|
page_dict = {k: v for k, v in response_dict.items() if k.endswith('_PG')}
|
|
page_dict = {k[:-3]: v for k, v in response_dict.items()}
|
|
field_l = list(field_dict.keys())
|
|
answer_l = list(field_dict.values())
|
|
page_no_l = [page_dict.get(x, "") for x in field_l]
|
|
|
|
try:
|
|
location_l = [context.find(a, context.find("Start of Page No. = "+str(p))) if isinstance(
|
|
a, str) and a != "" else -1 for a, p in zip(answer_l, page_no_l)]
|
|
except:
|
|
location_l = [context.find(answer) if isinstance(answer, str) and answer != "" else -1 for answer in answer_l]
|
|
snippet_l = [' '.join(context[:location].split('.')[-4:]) + ' ' + ' '.join(context[location:].split('. ')[:5]) if location != -1 else ' ' for location in location_l]
|
|
|
|
logger.info(f"Parsed fields: {field_l}")
|
|
return field_l, answer_l, page_no_l, location_l, snippet_l
|
|
|
|
def post_processing(answer_list):
|
|
logger.info("Post-processing answers.")
|
|
try:
|
|
answer_list = [" " if any(val in str(answer) for val in ['do not have', 'do not see', 'does not specify', 'Does not specify', 'does not explicitly', 'N/A', "don't know", "do not see", 'Not specified', "Don't know", "don't see", "don't have", "Does not apply", "Nothing found", "None"]) else answer for answer in answer_list]
|
|
answer_list = [str(answer).rstrip(".") for answer in answer_list]
|
|
answer_list = [answer if str(answer) != "one-year" else "1 year" for answer in answer_list]
|
|
answer_list = [answer if str(answer) != "one year" else "1 year" for answer in answer_list]
|
|
answer_list = [answer if str(answer) != "one" else "1 year" for answer in answer_list]
|
|
answer_list = [answer if str(answer) != "one (1) year" else "1 year" for answer in answer_list]
|
|
answer_list = [answer if str(answer) != "twelve" else "1 year" for answer in answer_list]
|
|
answer_list = [answer if str(answer) != "XI" else "11" for answer in answer_list]
|
|
answer_list = [answer if str(answer) != "Third" else "3" for answer in answer_list]
|
|
answer_list = [answer if str(answer) != "Six" else "6" for answer in answer_list]
|
|
logger.info(f"Post-processed answers: {answer_list}")
|
|
return answer_list
|
|
except:
|
|
logger.error("Error in post-processing answers.")
|
|
return ["post processing error"] * len(answer_list)
|
|
|
|
def compare_with_actuals(df):
|
|
logger.info("Comparing extracted values with actual values.")
|
|
query = 'SELECT * FROM "TRAINING_DATA_RAW"'
|
|
field_values = read_from_db(query)
|
|
field_values['Contract Name'] = field_values['DOCUMENT_NAME']
|
|
field_values_2 = pd.DataFrame(columns=['Contract Name', 'SF_DB_COL_NAME', 'Actual Value Stored', 'Original Page Number'])
|
|
|
|
for contract in set(field_values['Contract Name']):
|
|
field_values_1 = field_values[field_values['Contract Name'] == contract].head(1).transpose().reset_index()
|
|
field_values_1.columns = ['SF_DB_COL_NAME', 'Actual Value Stored']
|
|
field_values_p1 = field_values_1[~field_values_1['SF_DB_COL_NAME'].str.endswith('_PG')]
|
|
field_values_p2 = field_values_1[field_values_1['SF_DB_COL_NAME'].str.endswith('_PG')]
|
|
field_values_p2.columns = ['SF_DB_COL_NAME', 'Original Page Number']
|
|
field_values_p2["SF_DB_COL_NAME"] = field_values_p2["SF_DB_COL_NAME"].str.replace("_PG", "")
|
|
|
|
field_values_1 = pd.merge(field_values_p1, field_values_p2, how='left', on=['SF_DB_COL_NAME'])
|
|
field_values_1['Contract Name'] = contract
|
|
field_values_2 = pd.concat([field_values_2, field_values_1], ignore_index=True)
|
|
|
|
df.rename(columns={'Field Name': 'SF_DB_COL_NAME'}, inplace=True)
|
|
df['Contract Name'] = df['Contract Name'].str[:-3] + 'pdf'
|
|
fields_selected = list(df['SF_DB_COL_NAME'])
|
|
contracted_selected = list(df['Contract Name'])
|
|
df = pd.merge(df, field_values_2, how='right', on=['Contract Name', 'SF_DB_COL_NAME'])
|
|
df = df[df['SF_DB_COL_NAME'].isin(fields_selected)]
|
|
df = df[df['Contract Name'].isin(contracted_selected)]
|
|
|
|
df['Original Page Number'] = df['Original Page Number'].apply(lambda x: re.search(r'\d+', x).group() if isinstance(x, str) and re.search(r'\d+', x) is not None else " ")
|
|
df_date = df[df['SF_DB_COL_NAME'].str.contains('_DT', na=False)]
|
|
df_others = df[~df['SF_DB_COL_NAME'].str.contains('_DT', na=False)]
|
|
df_date['Actual Value Stored'] = pd.to_datetime(df_date['Actual Value Stored'], errors='coerce').dt.strftime('%Y-%m-%d').fillna(" ")
|
|
df_date['New Extracted value'] = pd.to_datetime(df_date['New Extracted value'], errors='coerce').dt.strftime('%Y-%m-%d').fillna(" ")
|
|
df = pd.concat([df_date, df_others], ignore_index=True)
|
|
|
|
df.sort_values(['SF_DB_COL_NAME', 'Contract Name'], inplace=True)
|
|
df.fillna(" ", inplace=True)
|
|
df['Actual Value Stored'] = df['Actual Value Stored'].apply(lambda x: x.strip() if isinstance(x, str) else '')
|
|
df['New Extracted value'] = df['New Extracted value'].apply(lambda x: x.strip() if isinstance(x, str) else '')
|
|
actual_value_list = list(df['Actual Value Stored'])
|
|
actual_value_list = [answer if str(answer) != "12 months" else "1 year" for answer in actual_value_list]
|
|
actual_value_list = [answer if str(answer) != "Fifth" else "5" for answer in actual_value_list]
|
|
actual_value_list = [answer if str(answer) != "Seventh" else "7" for answer in actual_value_list]
|
|
answer_list = list(df['New Extracted value'])
|
|
actual_value_list = [s.replace('-', '').replace(' ', '').replace('[', '').replace(']', '').lower() for s in actual_value_list]
|
|
answer_list = [s.replace('-', '').replace(' ', '').replace('[', '').replace(']', '').lower() for s in answer_list]
|
|
result_list = [(i in j) or (j in i) if isinstance(i, str) and isinstance(j, str) and ((i != '') == (j != '')) else False for i, j in zip(actual_value_list, answer_list)]
|
|
df['Result'] = [str(x) for x in result_list]
|
|
|
|
df = df[['Contract Name', 'SF_DB_COL_NAME', 'Actual Value Stored', 'Raw value', 'New Extracted value', 'Confidence Level', 'Snippet', 'Original Page Number', 'New Page Number', 'Revised Prompt', 'Result']]
|
|
|
|
accuracy = round(sum(bool(x) for x in result_list) * 100 / len(list(df['Result'])) if len(df['Result']) > 0 else 0, 2)
|
|
logger.info(f"ACCURACY: {accuracy}")
|
|
history = read_from_db('SELECT * FROM "TRAINING_ATTEMPT_LOGS"')
|
|
history.rename(columns={'FIELD_NAME': 'Field Name', 'CONTRACTS_TESTED': '# Contracts Tested', 'USERNAME': 'Username', 'DATE_TIME': 'Date/Time', 'ACCURACY': 'Accuracy', 'ATTEMPT_NUM': 'Attempt #'}, inplace=True)
|
|
history = history[['Field Name', '# Contracts Tested', 'Username', 'Date/Time', 'Accuracy', 'Attempt #']]
|
|
history.loc[len(history.index)] = [df['SF_DB_COL_NAME'].iloc[0], str(df['Contract Name'].nunique()), 'pipeline', datetime.now().strftime("%Y-%m-%d %H:%M:%S"), accuracy, 0]
|
|
|
|
logger.info(f"Comparison completed. Accuracy: {accuracy}")
|
|
return df, history
|
|
|
|
def main(bucket, object_key, field_groups):
|
|
logger.info(f"Processing contract: {object_key} from bucket: {bucket}")
|
|
try:
|
|
data = s3_client.get_object(Bucket=bucket, Key=object_key)
|
|
contents = data['Body'].read()
|
|
context = contents.decode("utf-8")
|
|
|
|
# Fetch batch_id from object tags
|
|
batch_id = get_s3_object_tags(bucket, object_key).get('BatchId')
|
|
logger.info(f"Batch ID: {batch_id}")
|
|
|
|
df = pd.DataFrame(columns=['Contract Name', 'Field Name', 'Raw value', 'New Extracted value', 'Confidence Level', 'Snippet', 'New Page Number', 'Revised Prompt'])
|
|
field_list, answer_list, snippet_list, page_no_list, contract_list_f = [], [], [], [], []
|
|
|
|
question = prepare_prompt(*field_groups)
|
|
response_text = invoke_llm(context, question, llm_selected=LLM_SELECTED)
|
|
# logger.info(f"Response text: {response_text}, CONTEXT: {context}")
|
|
field_l, answer_l, page_no_l, location_l, snippet_l = json_parsing(response_text, context)
|
|
field_list.extend(field_l)
|
|
answer_list.extend(answer_l)
|
|
contract_list_f.extend([object_key]*len(field_l))
|
|
snippet_list.extend(snippet_l)
|
|
page_no_list.extend(page_no_l)
|
|
|
|
df['Field Name'] = field_list
|
|
df['Raw value'] = answer_list
|
|
contract_list_f = [contract.rsplit('/', 1)[1] for contract in contract_list_f]
|
|
df['Contract Name'] = contract_list_f
|
|
answer_list = post_processing(answer_list)
|
|
df['New Extracted value'] = answer_list
|
|
# df['Confidence Level'] = ' '
|
|
df['Snippet'] = snippet_list
|
|
df['New Page Number'] = page_no_list
|
|
df['New Page Number'] = df['New Page Number'].apply(lambda x: re.search(r'\d+', x).group() if isinstance(x, str) and re.search(r'\d+', x) is not None else " ")
|
|
# df['Revised Prompt'] = [question] * len(contract_list_f)
|
|
# Need to add batch_id to the df
|
|
df['Batch_id'] = batch_id
|
|
|
|
# Remove column confidence level and revised prompt from df
|
|
df = df.drop(columns=['Confidence Level'])
|
|
df = df.drop(columns=['Revised Prompt'])
|
|
|
|
# TODO: uncomment accuracy function after testing
|
|
# df, history = compare_with_actuals(df)
|
|
csv_buf = StringIO()
|
|
df.to_csv(csv_buf, header=True, index=False)
|
|
csv_buf.seek(0)
|
|
timestamp = datetime.now().strftime("%Y%m%d%H%M%S")
|
|
file_name = f"results_{batch_id}_AC_{timestamp}.csv"
|
|
s3_client.put_object(Bucket=bucket, Body=csv_buf.getvalue(), Key=f"final_output/{batch_id}/{file_name}")
|
|
csv_buf = StringIO()
|
|
df.to_csv(csv_buf, header=True, index=False)
|
|
csv_buf.seek(0)
|
|
|
|
# Saving data to snowflake raw data ingestion bucket & calling stored proc
|
|
s3_client.put_object(Bucket=snowflake_ingestion_bucket, Body=csv_buf.getvalue(), Key=f"doczy_pipeline_output/{file_name}")
|
|
cur = snowflake_conn.cursor()
|
|
load_data_query = f"CALL LOAD_DOCZY_PIPELINE_RAW_OUTPUT('{file_name}')"
|
|
cur.execute(load_data_query)
|
|
|
|
|
|
# history.tail(1).to_csv(csv_buf, header=True, index=False)
|
|
# csv_buf.seek(0)
|
|
# timestamp = datetime.now().strftime("%Y%m%d%H%M%S")
|
|
# file_name = f"history_{timestamp}.csv"
|
|
# s3_client.put_object(Bucket=S3_RESULT_BUCKET, Body=csv_buf.getvalue(), Key=f"training_interface/{file_name}")
|
|
|
|
# logger.info("Results saved to S3 and database.")
|
|
# save_to_sf('load_training_results', training_results_file_name="results.csv", attempt_logs_file_name="history.csv")
|
|
except Exception as e:
|
|
logger.error(f"Error processing contract: {e} \n TRACEBACK: {traceback.format_exc()}")
|
|
|
|
|