Skip to content
This repository was archived by the owner on Jun 20, 2023. It is now read-only.
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 9 additions & 3 deletions metrics.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,21 +20,27 @@
from common import AV_STATUS_INFECTED


def send(env, bucket, key, status):
def send(env, bucket, key, status, signature, logurl):
if "DATADOG_API_KEY" in os.environ:
datadog.initialize() # by default uses DATADOG_API_KEY

result_metric_name = "unknown"

metric_tags = ["env:%s" % env, "bucket:%s" % bucket, "object:%s" % key]
metric_tags = [
"env:%s" % env,
"bucket:%s" % bucket,
"object:%s" % key,
"signature:%s" % signature,
]

if status == AV_STATUS_CLEAN:
result_metric_name = "clean"
elif status == AV_STATUS_INFECTED:
result_metric_name = "infected"
datadog.api.Event.create(
title="Infected S3 Object Found",
text="Virus found in s3://%s/%s." % (bucket, key),
text="Virus found in s3://%s/%s. The signature is '%s'. Cloudwatch Logs : %s"
% (bucket, key, signature, logurl),
tags=metric_tags,
)

Expand Down
33 changes: 31 additions & 2 deletions scan.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
import json
import os
from urllib.parse import unquote_plus
import urllib
from distutils.util import strtobool

import boto3
Expand Down Expand Up @@ -47,6 +48,15 @@ def event_object(event, event_source="s3"):
if event_source.upper() == "SNS":
event = json.loads(event["Records"][0]["Sns"]["Message"])

# S3 Batch
if "Records" not in event:
s3Key = urllib.unquote(event["tasks"][0]["s3Key"]).decode("utf8")
s3BucketArn = event["tasks"][0]["s3BucketArn"]
s3Bucket = s3BucketArn.split(":::")[-1]
# Create and return the object
s3 = boto3.resource("s3")
return s3.Object(s3Bucket, s3Key)

# Break down the record
records = event["Records"]
if len(records) == 0:
Expand All @@ -66,7 +76,7 @@ def event_object(event, event_source="s3"):
key_name = s3_obj["object"].get("key", None)

if key_name:
key_name = unquote_plus(key_name)
key_name = urllib.unquote_plus(key_name.encode("utf8"))

# Ensure both bucket and key exist
if (not bucket_name) or (not key_name):
Expand Down Expand Up @@ -207,6 +217,20 @@ def lambda_handler(event, context):
ENV = os.getenv("ENV", "")
EVENT_SOURCE = os.getenv("EVENT_SOURCE", "S3")

region = context.invoked_function_arn.split(":")[3]
log_group_name = context.log_group_name
log_stream_name = context.log_stream_name
log_url = (
"https://"
+ region
+ ".console.aws.amazon.com/cloudwatch/home?region="
+ region
+ "#logEvent:group="
+ log_group_name
+ ";stream="
+ log_stream_name
)

start_time = get_timestamp()
print("Script starting at %s\n" % (start_time))
s3_object = event_object(event, event_source=EVENT_SOURCE)
Expand Down Expand Up @@ -257,7 +281,12 @@ def lambda_handler(event, context):
)

metrics.send(
env=ENV, bucket=s3_object.bucket_name, key=s3_object.key, status=scan_result
env=ENV,
bucket=s3_object.bucket_name,
key=s3_object.key,
status=scan_result,
signature=scan_signature,
logurl=log_url,
)
# Delete downloaded file to free up room on re-usable lambda function container
try:
Expand Down