From 4b819cbce1cee37af1e0fc5953244659fbc1c08c Mon Sep 17 00:00:00 2001 From: Grzegorz Huber Date: Thu, 6 Jun 2024 17:51:31 +0200 Subject: [PATCH] terraform refactor - refactor prompt_orchestrator_lambda --- streamlit-server/main.tf | 2 +- .../src/lambda/prompt-orchestrator/index.py | 61 ++++++++++--------- .../aws-textract-pipeline-part3/main.tf | 20 +++--- .../aws-textract-pipeline-part3/variables.tf | 32 +++++----- .../modules/aws-textract-pipeline/main.tf | 1 + 5 files changed, 61 insertions(+), 55 deletions(-) diff --git a/streamlit-server/main.tf b/streamlit-server/main.tf index 112322c..75fc049 100644 --- a/streamlit-server/main.tf +++ b/streamlit-server/main.tf @@ -88,7 +88,7 @@ resource "aws_instance" "streamlit_server" { iam_instance_profile = aws_iam_instance_profile.ec2_instance_profile.name associate_public_ip_address = false key_name = var.key_name - + root_block_device { volume_size = 30 # Variable encrypted = true # Mandatory diff --git a/textract-pipeline/src/lambda/prompt-orchestrator/index.py b/textract-pipeline/src/lambda/prompt-orchestrator/index.py index cfacc59..77503e6 100644 --- a/textract-pipeline/src/lambda/prompt-orchestrator/index.py +++ b/textract-pipeline/src/lambda/prompt-orchestrator/index.py @@ -13,12 +13,12 @@ sqs = boto3.client('sqs') # Initialize S3 & Textract clients s3_client = boto3.client('s3') -def lambda_handler(event, context): +def lambda_handler(event, context): try: # Read environment variables property_file_path = os.environ.get('PROPERTY_FILE_S3_PATH', '') - + # Read config.properties file_path_array = property_file_path.split("/") @@ -36,14 +36,14 @@ def lambda_handler(event, context): for record in event['Records']: # Extract the message body from the record record_body = json.loads(record['body']) - #logger.info('Message Count: ', str(len(record_body['Records'])) ) + # logger.info('Message Count: ', str(len(record_body['Records'])) ) for sqs_record in record_body['Records']: # Construct the source paths source_path = unquote_plus(sqs_record['s3']['object']['key']) document_id = get_filename_from_path(source_path) document_id = os.path.splitext(document_id)[0] - - prompt_config_query = "SELECT DISTINCT PC.FIELD_NAME,LENGTH(PC.PROMPT) as PROMPT_LENGTH, PC.FM_MODEL_ID,PC.GROUP_ID FROM PROMPT_CONFIG PC JOIN DOCUMENT_LOGS DL ON DL.GROUP_ID LIKE '%'||PC.GROUP_ID ||'%' AND DL.DOCUMENT_ID = '"+document_id+"' ORDER BY PROMPT_LENGTH" + + prompt_config_query = "SELECT DISTINCT PC.FIELD_NAME,LENGTH(PC.PROMPT) as PROMPT_LENGTH, PC.FM_MODEL_ID,PC.GROUP_ID FROM PROMPT_CONFIG PC JOIN DOCUMENT_LOGS DL ON DL.GROUP_ID LIKE '%'||PC.GROUP_ID ||'%' AND DL.DOCUMENT_ID = '" + document_id + "' ORDER BY PROMPT_LENGTH" field_Prompt_group_response = call_snowflake_db_select_lambda(prompt_config_query) # Maximum prompt length per group @@ -56,7 +56,11 @@ def lambda_handler(event, context): for i, (key, batch_list) in enumerate(batches.items(), 1): group_id, fm_model_id = key for j, batch in enumerate(batch_list, 1): - send_batch_to_sqs(batch,"https://sqs.us-east-2.amazonaws.com/660131068782/doczy-prompt-processor-sqs-dev",S3_BUCKET_NAME+"/"+source_path,i) + AWS_REGION_PARAM = os.getenv('AWS_REGION_PARAM') + AWS_ACCOUNT_ID = os.getenv('AWS_ACCOUNT_ID') + ENVIRONMENT = os.getenv('ENVIRONMENT') + sqs_url = f"https://sqs.{AWS_REGION_PARAM}.amazonaws.com/{AWS_ACCOUNT_ID}/doczy-prompt-processor-sqs-{ENVIRONMENT}" + send_batch_to_sqs(batch, sqs_url, S3_BUCKET_NAME + "/" + source_path, i) return { 'statusCode': 200, 'body': json.dumps('Message sent to SQS queue') @@ -68,12 +72,12 @@ def lambda_handler(event, context): return { 'statusCode': 500, 'body': error_message - } - + } + + # Function to retrieve configuration values from S3 + -# Function to retrieve configuration values from S3 def load_config_from_s3(bucket_name, file_key): - # Download the config file from S3 response = s3_client.get_object(Bucket=bucket_name, Key=file_key) config_content = response['Body'].read().decode('utf-8') @@ -89,12 +93,13 @@ def load_config_from_s3(bucket_name, file_key): return config_dict + def get_filename_from_path(full_path): return os.path.basename(full_path) + # Function to get s3 object tags def get_s3_object_tags(bucket_name, object_key): - try: # Get object tags response = s3_client.get_object_tagging( @@ -110,15 +115,15 @@ def get_s3_object_tags(bucket_name, object_key): except Exception as e: logger.exception(f"Error: {e}") return None - + def call_snowflake_db_select_lambda(query): # Create a Lambda client lambda_client = boto3.client('lambda') - + # Specify your Lambda function name lambda_function_name = 'doczy-dev-get-snowflake-logs' - + # Prepare payload payload = { 'query': query @@ -130,14 +135,15 @@ def call_snowflake_db_select_lambda(query): InvocationType='RequestResponse', # Synchronous invocation Payload=json.dumps(payload) ) - + # Return response body response_body = response['Payload'].read().decode('utf-8') return response_body except Exception as e: logger.exception(f"Error: {e}") return None - + + def extract_group_id(json_string): try: logger.info(f"json_string: {json_string}") @@ -146,13 +152,14 @@ def extract_group_id(json_string): group_id = body[0]['GROUP_ID'] group_ids = group_id.split(',') group_ids_quoted = ["'" + group_id.strip() + "'" for group_id in group_ids] - + # Join the elements back together with commas return ','.join(group_ids_quoted) except (json.JSONDecodeError, KeyError, IndexError) as e: logger.error("Error parsing JSON or extracting GROUP_ID:", e) return None - + + def extract_body_tags(json_data): try: # Parse the JSON string @@ -167,14 +174,13 @@ def extract_body_tags(json_data): logger.info(f"tags_list: {tags_list}") return tags_list - + except (json.JSONDecodeError, KeyError) as e: logger.error("Error parsing JSON or extracting 'body' tag:", e) return None def divide_into_batches(data, max_prompt_length_per_group): - try: # Parse the input data data_dict = json.loads(data) @@ -189,7 +195,7 @@ def divide_into_batches(data, max_prompt_length_per_group): # Divide data into batches based on group_id and fm_model_id for item in body_dict: - + group_id = item["GROUP_ID"] fm_model_id = item["FM_MODEL_ID"] prompt_length = item["PROMPT_LENGTH"] @@ -212,8 +218,7 @@ def divide_into_batches(data, max_prompt_length_per_group): # Retrieve the current batch (last sub-batch) current_batch = batches[(group_id, fm_model_id)][-1] - - + # Append the current item to the current batch current_batch.append(item) # Update prompt length for the current group and model id @@ -227,16 +232,16 @@ def divide_into_batches(data, max_prompt_length_per_group): raise e -def send_batch_to_sqs(batch, queue_url,source_path,prompt_batch_id): +def send_batch_to_sqs(batch, queue_url, source_path, prompt_batch_id): # Initialize SQS client sqs = boto3.client('sqs') # Convert batch to JSON string - batch_json = json.dumps(batch) + batch_json = json.dumps(batch) try: result = {"field_list": batch_json, - "file_name": source_path, - "prompt_batch_id" : prompt_batch_id + "file_name": source_path, + "prompt_batch_id": prompt_batch_id } logger.info(f"Message sent to SQS queue: {json.dumps(result)}") # Send message to SQS queue @@ -250,4 +255,4 @@ def send_batch_to_sqs(batch, queue_url,source_path,prompt_batch_id): except Exception as e: logger.error('Error sending message to SQS queue:', e) - return False \ No newline at end of file + return False diff --git a/textract-pipeline/terraform/modules/aws-textract-pipeline-part3/main.tf b/textract-pipeline/terraform/modules/aws-textract-pipeline-part3/main.tf index a30f011..95fe99f 100644 --- a/textract-pipeline/terraform/modules/aws-textract-pipeline-part3/main.tf +++ b/textract-pipeline/terraform/modules/aws-textract-pipeline-part3/main.tf @@ -42,7 +42,7 @@ locals { sender_lambda_name = "${local.lambda_prefix}-${var.textract_sender_lambda.name}" receiver_lambda_name = "${local.lambda_prefix}-${var.textract_receiver_lambda.name}" text_creation_lambda_name = "${local.lambda_prefix}-${var.text_creation_lambda.name}" - prompt_orchestrator_lambda_name = "${local.lambda_prefix}-${var.prompt_orchestrator_lambda.name}" + prompt_orchestrator_lambda_name = "${local.lambda_prefix}-${local.prompt_orchestrator_lambda.name}" prompt_processor_lambda_name = "${local.lambda_prefix}-${var.prompt_processor_lambda.name}" call_bedrock_lambda_name = "${local.lambda_prefix}-${var.call_bedrock_lambda.name}" database_interface_lambda_name = "${local.lambda_prefix}-${var.database_interface_lambda.name}" @@ -548,22 +548,22 @@ module "prompt_orchestrator_lambda_function" { version = "7.2.2" function_name = local.prompt_orchestrator_lambda_name - description = var.prompt_orchestrator_lambda.description - handler = var.prompt_orchestrator_lambda.handler - runtime = var.prompt_orchestrator_lambda.runtime - timeout = var.prompt_orchestrator_lambda.timeout - memory_size = var.prompt_orchestrator_lambda.memory_size - ephemeral_storage_size = var.prompt_orchestrator_lambda.ephemeral_storage_size - architectures = var.prompt_orchestrator_lambda.architectures + description = local.prompt_orchestrator_lambda.description + handler = local.prompt_orchestrator_lambda.handler + runtime = local.prompt_orchestrator_lambda.runtime + timeout = local.prompt_orchestrator_lambda.timeout + memory_size = local.prompt_orchestrator_lambda.memory_size + ephemeral_storage_size = local.prompt_orchestrator_lambda.ephemeral_storage_size + architectures = local.prompt_orchestrator_lambda.architectures layers = [] - environment_variables = merge(local.environment_variables,{ + environment_variables = merge(local.environment_variables, local.prompt_orchestrator_lambda.environment_variables, { PROMPT_PROCESSOR_QUEUE_URL = var.prompt_processor_queue_url MAX_PROMPT_LENGTH_PER_GROUP = "5000" }) - source_path = var.prompt_orchestrator_lambda.source_path + source_path = local.prompt_orchestrator_lambda.source_path event_source_mapping = { sqs = { diff --git a/textract-pipeline/terraform/modules/aws-textract-pipeline-part3/variables.tf b/textract-pipeline/terraform/modules/aws-textract-pipeline-part3/variables.tf index 82718ad..2ec776b 100644 --- a/textract-pipeline/terraform/modules/aws-textract-pipeline-part3/variables.tf +++ b/textract-pipeline/terraform/modules/aws-textract-pipeline-part3/variables.tf @@ -18,6 +18,11 @@ variable "aws_region" { type = string } +# Required +variable "aws_account_id" { + type = string +} + # Required variable "vpc_id" { type = string @@ -481,21 +486,10 @@ variable "text_creation_lambda" { } # prompt-orchestrator lambda -variable "prompt_orchestrator_lambda" { - type = object({ - name = string - description = string - handler = string - runtime = string - timeout = number - memory_size = number - ephemeral_storage_size = number - architectures = list(string) - layers = list(string) - environment_variables = map(string) - source_path = string - }) - default = { + +# Use locals to merge default values with variable values +locals { + prompt_orchestrator_lambda_default = { name = "prompt-orchestrator" description = "This Lambda functions breaks prompts into chunks and send messages into SQS" handler = "index.lambda_handler" @@ -505,9 +499,15 @@ variable "prompt_orchestrator_lambda" { ephemeral_storage_size = 512 architectures = ["x86_64"] layers = [] - environment_variables = {} + environment_variables = { + AWS_REGION_PARAM = var.aws_region + AWS_ACCOUNT_ID = var.aws_account_id + ENVIRONMENT = var.environment + } source_path = "../src/lambda/prompt-orchestrator" } + + prompt_orchestrator_lambda = merge(local.prompt_orchestrator_lambda_default) } # prompt-processor lambda diff --git a/textract-pipeline/terraform/modules/aws-textract-pipeline/main.tf b/textract-pipeline/terraform/modules/aws-textract-pipeline/main.tf index 9a5dce0..3d3366b 100644 --- a/textract-pipeline/terraform/modules/aws-textract-pipeline/main.tf +++ b/textract-pipeline/terraform/modules/aws-textract-pipeline/main.tf @@ -52,6 +52,7 @@ module "textract_pipeline_part3" { project_name = var.project_name aws_region = var.aws_region + aws_account_id = var.aws_account_id environment = var.environment vpc_id = var.vpc_id secret_manager_name = var.secret_manager_name