Files
han-app/VM2_services/codebase/services/message-safety/app/adapters.py
T

87 lines
3.2 KiB
Python

from __future__ import annotations
import asyncio
import ipaddress
from collections.abc import AsyncIterator
import boto3
import dns.asyncresolver
from botocore.config import Config
from botocore.exceptions import BotoCoreError, ClientError
from opentelemetry import trace
from app.contracts import Attachment
from app.file_pipeline import DependencyFailure, ObjectChanged
from app.url_policy import DnsError, DnsNxDomain
tracer = trace.get_tracer("message-safety.dependencies")
class TrustedDnsResolver:
def __init__(self, nameservers: list[str]) -> None:
self._resolver = dns.asyncresolver.Resolver(configure=not nameservers)
if nameservers:
self._resolver.nameservers = nameservers
async def resolve(
self, hostname: str
) -> tuple[ipaddress.IPv4Address | ipaddress.IPv6Address, ...]:
with tracer.start_as_current_span(
"message_safety.dns.resolve", attributes={"server.address.type": "domain"}
):
return await self._resolve(hostname)
async def _resolve(
self, hostname: str
) -> tuple[ipaddress.IPv4Address | ipaddress.IPv6Address, ...]:
found: list[ipaddress.IPv4Address | ipaddress.IPv6Address] = []
try:
for kind in ("A", "AAAA"):
try:
answer = await self._resolver.resolve(hostname, kind, lifetime=1.0)
found.extend(ipaddress.ip_address(item.address) for item in answer)
except dns.resolver.NoAnswer:
pass
except dns.resolver.NXDOMAIN as exc:
raise DnsNxDomain from exc
except dns.exception.DNSException as exc:
raise DnsError from exc
if not found:
raise DnsNxDomain
return tuple(found)
class S3VersionReader:
def __init__(self, endpoint_url: str, bucket: str, access_key: str, secret_key: str) -> None:
self.bucket = bucket
self.client = boto3.client(
"s3",
endpoint_url=endpoint_url,
aws_access_key_id=access_key,
aws_secret_access_key=secret_key,
config=Config(s3={"addressing_style": "virtual"}, retries={"max_attempts": 2}),
)
async def stream(self, attachment: Attachment) -> AsyncIterator[bytes]:
try:
response = await asyncio.to_thread(
self.client.get_object,
Bucket=self.bucket,
Key=attachment.quarantine_object_key,
VersionId=attachment.quarantine_version_id,
IfMatch=attachment.quarantine_etag,
)
body = response["Body"]
while True:
chunk = await asyncio.to_thread(body.read, 65_536)
if not chunk:
break
yield chunk
except ClientError as exc:
code = exc.response.get("Error", {}).get("Code")
if code in {"PreconditionFailed", "NoSuchKey", "NoSuchVersion"}:
raise ObjectChanged("version or ETag changed") from exc
raise DependencyFailure("S3 dependency unavailable") from exc
except BotoCoreError as exc:
raise DependencyFailure("S3 dependency unavailable") from exc