document_redaction / tools /malware_scan.py
seanpedrickcase's picture
Sync: Fixes to image annotation component for compatibility with gradio 6.26.0
2677bc3
Raw
History Blame Contribute Delete
19.3 kB
"""
GuardDuty Malware Protection for S3 — stage Gradio uploads, poll scan tags, fail closed.
When enabled (``SCAN_UPLOADS_FOR_MALWARE`` + ``RUN_AWS_FUNCTIONS`` + bucket), each file
is uploaded to a dedicated scan bucket. GuardDuty tags objects with
``GuardDutyMalwareScanStatus``. Only ``NO_THREATS_FOUND`` is accepted. Staged S3
objects are always deleted in ``finally`` after upload (success or failure).
"""
from __future__ import annotations
import logging
import os
import time
import uuid
from typing import Any
import boto3
import gradio as gr
from botocore.exceptions import BotoCoreError, ClientError
try:
from boto3.exceptions import S3UploadFailedError
except ImportError: # pragma: no cover - defensive for older boto3
S3UploadFailedError = type("S3UploadFailedError", (Exception,), {})
from tools.aws_functions import _effective_aws_region
from tools.config import (
MALWARE_SCAN_POLL_INTERVAL_SEC,
MALWARE_SCAN_S3_BUCKET,
MALWARE_SCAN_S3_PREFIX,
MALWARE_SCAN_SHOW_CHECKING_INFO,
MALWARE_SCAN_TIMEOUT_SEC,
malware_scan_enabled,
)
from tools.example_data_paths import is_bundled_example_file
from tools.secure_path_utils import secure_join
_logger = logging.getLogger(__name__)
GUARDDUTY_MALWARE_SCAN_TAG_KEY = "GuardDutyMalwareScanStatus"
CLEAN_SCAN_STATUS = "NO_THREATS_FOUND"
USER_REJECT_MESSAGE = (
"Upload rejected: the file did not pass malware scanning. "
"Please contact your administrator if you believe this is an error."
)
USER_SERVICE_ERROR_MESSAGE = (
"Upload could not be scanned due to a configuration or permissions "
"error. Please contact your administrator."
)
USER_NOT_SCANNED_MESSAGE = (
"Upload blocked: file(s) must pass malware scanning before processing. "
"Please upload again or contact your administrator."
)
MALWARE_SCAN_CHECKING_INFO_MESSAGE = "Scanning uploaded file(s). Please wait."
MALWARE_SCAN_SUCCESS_INFO_MESSAGE = (
"Malware scan complete: no issues detected in the uploaded file(s)."
)
# (abspath, mtime) -> True after a successful scan (Pi upload + submit fallback)
_recent_clean_scans: dict[tuple[str, float], bool] = {}
class MalwareScanRejectedError(Exception):
"""Raised when a file fails malware scanning or scan did not complete cleanly."""
class MalwareScanServiceError(Exception):
"""Raised when malware scanning cannot run (S3, IAM, network, etc.)."""
def normalize_gradio_file_paths(file_input: Any) -> list[str]:
"""Normalize Gradio File component values to local path strings."""
if file_input is None:
return []
# Gradio 6 ``ListFiles`` wrapper (``file_count="multiple"`` payloads).
root = getattr(file_input, "root", None)
if root is not None and not isinstance(file_input, (dict, str, bytes)):
file_input = root
if isinstance(file_input, dict):
name = file_input.get("name") or file_input.get("path")
return [os.path.abspath(str(name))] if name else []
if isinstance(file_input, str):
return [os.path.abspath(file_input)] if file_input.strip() else []
if not isinstance(file_input, (list, tuple)):
path = getattr(file_input, "name", None) or getattr(file_input, "path", None)
return [os.path.abspath(str(path))] if path else []
paths: list[str] = []
for item in file_input:
paths.extend(normalize_gradio_file_paths(item))
return paths
def _file_scan_cache_key(local_path: str) -> tuple[str, float] | None:
try:
abspath = os.path.abspath(local_path)
return abspath, os.path.getmtime(abspath)
except OSError:
return None
def already_scanned_clean(local_path: str) -> bool:
key = _file_scan_cache_key(local_path)
return bool(key and _recent_clean_scans.get(key))
def mark_scanned_clean(local_path: str) -> None:
key = _file_scan_cache_key(local_path)
if key:
_recent_clean_scans[key] = True
def path_is_malware_clean(local_path: str) -> bool:
"""True when a path may proceed without a new GuardDuty staging scan."""
if already_scanned_clean(local_path):
return True
if is_bundled_example_file(local_path):
mark_scanned_clean(local_path)
_logger.info("Skipping malware scan for bundled example file: %s", local_path)
return True
return False
def mark_gradio_example_files_malware_clean(file_input: Any) -> None:
"""Mark Gradio Example file paths clean (direct path or temp copy of demo asset)."""
if not malware_scan_enabled():
return
for path in normalize_gradio_file_paths(file_input):
path_is_malware_clean(path)
def mark_app_generated_files_malware_clean(file_input: Any) -> None:
"""Mark app-written output paths as scan-clean (e.g. OCR CSVs after redaction)."""
if not malware_scan_enabled():
return
for path in normalize_gradio_file_paths(file_input):
if os.path.isfile(path):
mark_scanned_clean(path)
def clear_scan_cache() -> None:
"""Test helper — reset the in-process clean-scan cache."""
_recent_clean_scans.clear()
def clear_scan_cache_for_path(local_path: str) -> None:
"""Drop all clean-scan cache entries for ``local_path`` (any mtime)."""
try:
abspath = os.path.abspath(local_path)
except OSError:
return
stale_keys = [key for key in _recent_clean_scans if key[0] == abspath]
for key in stale_keys:
del _recent_clean_scans[key]
if stale_keys:
_logger.info("Cleared malware scan cache for removed upload path: %s", abspath)
def clear_scan_cache_for_paths(paths: list[str]) -> None:
"""Drop clean-scan cache entries for each path in ``paths``."""
for path in paths:
clear_scan_cache_for_path(path)
def _staging_s3_key(local_path: str) -> str:
basename = os.path.basename(local_path)
prefix = (MALWARE_SCAN_S3_PREFIX or "").strip()
if prefix and not prefix.endswith("/"):
prefix = prefix + "/"
unique_name = f"{uuid.uuid4().hex}_{basename}"
return secure_join(prefix, unique_name).replace("\\", "/")
def _read_scan_status(s3_client: Any, bucket: str, key: str) -> str | None:
try:
response = s3_client.get_object_tagging(Bucket=bucket, Key=key)
except ClientError as exc:
code = (exc.response.get("Error") or {}).get("Code", "")
if code in {"AccessDenied", "AllAccessDisabled", "UnauthorizedAccess"}:
_raise_service_error(
f"Malware scan cannot read object tags on s3://{bucket}/{key}", exc
)
_logger.warning(
"get_object_tagging failed for s3://%s/%s (%s): %s",
bucket,
key,
code or "ClientError",
exc,
)
return None
except Exception as exc:
_logger.warning(
"get_object_tagging failed for s3://%s/%s: %s", bucket, key, exc
)
return None
for tag in response.get("TagSet", []):
if tag.get("Key") == GUARDDUTY_MALWARE_SCAN_TAG_KEY:
return tag.get("Value")
return None
def _poll_scan_status(s3_client: Any, bucket: str, key: str) -> str:
deadline = time.monotonic() + MALWARE_SCAN_TIMEOUT_SEC
interval = max(0.5, float(MALWARE_SCAN_POLL_INTERVAL_SEC))
_logger.info("Waiting for GuardDuty malware scan on s3://%s/%s", bucket, key)
last_progress_log = time.monotonic()
while time.monotonic() < deadline:
status = _read_scan_status(s3_client, bucket, key)
if status is not None:
_logger.info(
"GuardDuty malware scan status for s3://%s/%s: %s",
bucket,
key,
status,
)
return status
now = time.monotonic()
if now - last_progress_log >= 10.0:
_logger.info(
"Still waiting for GuardDuty scan tag on s3://%s/%s", bucket, key
)
last_progress_log = now
time.sleep(interval)
_logger.warning(
"Malware scan timed out waiting for GuardDuty tag on s3://%s/%s",
bucket,
key,
)
raise MalwareScanRejectedError(
"Malware scan timed out before a result was available."
)
def _delete_staging_object(s3_client: Any, bucket: str, key: str) -> None:
try:
s3_client.delete_object(Bucket=bucket, Key=key)
_logger.info(
"Deleted staging object s3://%s/%s after malware scan", bucket, key
)
except Exception as exc:
_logger.warning(
"Failed to delete staging object s3://%s/%s: %s", bucket, key, exc
)
def _raise_service_error(context: str, exc: Exception) -> None:
"""Log a technical AWS/boto failure and raise a user-safe service error."""
_logger.exception("%s: %s", context, exc)
raise MalwareScanServiceError(USER_SERVICE_ERROR_MESSAGE) from exc
def _halt_gradio_upload(message: str, *, title: str = "Upload blocked") -> None:
"""Show a Gradio warning and halt the upload event chain."""
# gr.Warning(message)
raise gr.Error(message, title=title)
def clear_gradio_file_upload(_file_input: Any = None):
"""Clear a File component after a failed malware scan upload."""
return gr.update(value=None)
def _gradio_updates(count: int, **kwargs: Any) -> Any:
"""Return one ``gr.update`` or a tuple of them for multi-output event handlers."""
if count == 1:
return gr.update(**kwargs)
return tuple(gr.update(**kwargs) for _ in range(count))
def make_malware_scan_disable_outputs(button_count: int):
"""Return a handler that disables ``button_count`` Gradio buttons."""
def _disable(*_args: Any):
if not malware_scan_enabled():
return _gradio_updates(button_count)
return _gradio_updates(button_count, interactive=False)
return _disable
def make_malware_scan_upload_start(button_count: int):
"""
Return an ``.upload(...)`` handler that scans uploads then disables buttons.
Runs ``scan_gradio_file_upload`` on the upload event payload (not a chained
``.success()`` step) so each replacement upload is scanned reliably.
"""
def _start(file_input: Any):
scan_gradio_file_upload(file_input)
if not malware_scan_enabled():
return _gradio_updates(button_count)
return _gradio_updates(button_count, interactive=False)
return _start
def handle_gradio_file_deleted(delete_data: gr.DeletedFileData) -> None:
"""
Gradio ``.delete(...)`` handler — invalidate scan cache for the removed file.
Requires ``gr.DeletedFileData`` type hint so Gradio injects event data (``Any``
is not detected and the handler would never receive the deleted path).
"""
if not malware_scan_enabled():
return
paths = normalize_gradio_file_paths(delete_data.file)
clear_scan_cache_for_paths(paths)
def make_malware_scan_enable_outputs(button_count: int):
"""Return a handler that re-enables ``button_count`` Gradio buttons."""
def _enable(*_args: Any):
return _gradio_updates(button_count, interactive=True)
return _enable
def make_malware_scan_upload_failure_outputs(button_count: int):
"""Return a handler that clears the file input and re-enables buttons."""
def _clear_file_and_enable_buttons(_file_input: Any = None):
if button_count <= 0:
return gr.update(value=None)
updates: list[Any] = [gr.update(value=None)]
updates.extend(gr.update(interactive=True) for _ in range(button_count))
return tuple(updates)
return _clear_file_and_enable_buttons
def bind_malware_scan_upload(
file_input: Any,
buttons: Any,
*,
api_visibility: str = "undocumented",
) -> Any:
"""
Wire upload-time malware scan, button disable/enable, and failure file-clear.
Attaches ``.upload(scan)``, ``.success(re-enable)``, and ``.failure(clear)``.
Use for File inputs that have no extra ``.success()`` steps after the scan.
"""
if isinstance(buttons, (list, tuple)):
button_list = list(buttons)
else:
button_list = [buttons]
if not button_list:
raise ValueError("bind_malware_scan_upload requires at least one button")
n = len(button_list)
return (
file_input.upload(
fn=make_malware_scan_upload_start(n),
inputs=[file_input],
outputs=button_list,
queue=True,
api_visibility=api_visibility,
)
.success(
fn=make_malware_scan_enable_outputs(n),
inputs=None,
outputs=button_list,
queue=False,
api_visibility=api_visibility,
)
.failure(
fn=make_malware_scan_upload_failure_outputs(n),
outputs=[file_input, *button_list],
queue=False,
api_visibility=api_visibility,
)
)
def require_files_malware_scanned(file_input: Any) -> None:
"""
Fail closed when malware scanning is enabled and paths are not scan-clean.
Call at the start of prepare/redact handlers so button-click paths cannot
bypass a failed upload scan. Scans any unclean paths that the upload event
missed (e.g. stale Gradio component state on replace-after-delete).
"""
if not malware_scan_enabled():
return
paths = normalize_gradio_file_paths(file_input)
if not paths:
return
if any(not path_is_malware_clean(path) for path in paths):
scan_gradio_file_upload(file_input)
if any(not path_is_malware_clean(path) for path in paths):
_halt_gradio_upload(USER_NOT_SCANNED_MESSAGE)
def _gradio_info(message: str) -> None:
"""Show a Gradio info toast and mirror the same text to container logs."""
_logger.info(message)
print(message, flush=True)
gr.Info(message)
def _notify_malware_scan_in_progress() -> None:
"""Optional Gradio info toast + log line while a malware scan runs."""
if not MALWARE_SCAN_SHOW_CHECKING_INFO:
return
_gradio_info(MALWARE_SCAN_CHECKING_INFO_MESSAGE)
def _notify_malware_scan_complete() -> None:
"""Optional Gradio info toast + log line after a clean malware scan."""
if not MALWARE_SCAN_SHOW_CHECKING_INFO:
return
_gradio_info(MALWARE_SCAN_SUCCESS_INFO_MESSAGE)
def scan_local_file_for_malware(local_path: str) -> None:
"""
Upload ``local_path`` to the malware scan bucket, poll GuardDuty tag, fail closed.
Staged S3 object is deleted in ``finally``. Raises ``MalwareScanRejectedError``
when the upload must be blocked.
"""
if not malware_scan_enabled():
return
if not local_path or not os.path.isfile(local_path):
raise MalwareScanRejectedError("Upload file is missing or not readable.")
if path_is_malware_clean(local_path):
return
bucket = MALWARE_SCAN_S3_BUCKET.strip()
s3_key = _staging_s3_key(local_path)
region = _effective_aws_region()
s3_client = boto3.client("s3", region_name=region or None)
uploaded = False
try:
_logger.info(
"Uploading %s to s3://%s/%s for malware scan",
local_path,
bucket,
s3_key,
)
try:
s3_client.upload_file(local_path, bucket, s3_key)
except (ClientError, S3UploadFailedError, BotoCoreError) as exc:
_raise_service_error("Malware scan S3 upload failed", exc)
uploaded = True
_logger.info(
"Malware scan upload complete; polling GuardDuty tag for %s", s3_key
)
status = _poll_scan_status(s3_client, bucket, s3_key)
if status != CLEAN_SCAN_STATUS:
raise MalwareScanRejectedError(
f"Malware scan result was not clean (status={status!r})."
)
mark_scanned_clean(local_path)
_logger.info("Malware scan passed for %s", local_path)
finally:
if uploaded:
_delete_staging_object(s3_client, bucket, s3_key)
def _delete_local_file(local_path: str) -> None:
try:
if os.path.isfile(local_path):
os.remove(local_path)
_logger.info("Deleted rejected upload file: %s", local_path)
except OSError as exc:
_logger.warning("Could not delete rejected upload file %s: %s", local_path, exc)
def scan_gradio_file_upload(file_input: Any) -> None:
"""
Gradio upload handler: scan all uploaded paths; raise ``gr.Error`` on rejection.
Intended for ``.upload(..., outputs=[])`` chains — does not return component updates.
"""
tic = time.monotonic()
if not malware_scan_enabled():
return
paths = normalize_gradio_file_paths(file_input)
if not paths:
_logger.warning("Malware scan upload handler received no file paths")
return
# Force a fresh GuardDuty scan for every user upload. Gradio may reuse the
# same temp path after remove-and-replace; the (path, mtime) cache entry
# from the previous file would otherwise skip scanning the replacement.
for path in paths:
if not is_bundled_example_file(path):
clear_scan_cache_for_path(path)
cached_paths = [path for path in paths if path_is_malware_clean(path)]
paths_to_scan = [path for path in paths if path not in cached_paths]
if cached_paths:
_logger.info(
"Skipping malware scan for %d already-clean path(s): %s",
len(cached_paths),
cached_paths,
)
if not paths_to_scan:
return
_notify_malware_scan_in_progress()
rejected_paths: list[str] = []
for path in paths_to_scan:
try:
scan_local_file_for_malware(path)
except MalwareScanServiceError as exc:
_logger.warning("Malware scan service error for upload %s: %s", path, exc)
_delete_local_file(path)
_halt_gradio_upload(USER_SERVICE_ERROR_MESSAGE)
except MalwareScanRejectedError as exc:
_logger.warning("Malware scan rejected upload %s: %s", path, exc)
_delete_local_file(path)
rejected_paths.append(path)
if rejected_paths:
_halt_gradio_upload(USER_REJECT_MESSAGE, title="Upload rejected")
_logger.info(
"Malware scan finished successfully for %d file(s)",
len(paths_to_scan),
)
toc = time.monotonic()
_logger.info(f"Malware scan finished in {round(toc - tic, 1)} seconds")
print(f"Malware scan finished in {round(toc - tic, 1)} seconds", flush=True)
_notify_malware_scan_complete()
def ensure_upload_scanned_for_malware(local_path: str | None) -> None:
"""
Scan a single path if enabled and not already in the clean cache.
Used by Pi agent submit fallback when upload events did not run (e.g. Examples).
"""
if not local_path or not malware_scan_enabled():
return
if path_is_malware_clean(local_path):
return
_notify_malware_scan_in_progress()
try:
scan_local_file_for_malware(local_path)
except MalwareScanServiceError:
if local_path:
_delete_local_file(local_path)
raise
except MalwareScanRejectedError:
if local_path:
_delete_local_file(local_path)
raise
_notify_malware_scan_complete()