#!/usr/local/bin/python3.14 -BE

from psutil import OPENBSD

if OPENBSD:
    from openbsd import pledge, unveil

    pledge("stdio rpath inet dns tty prot_exec unveil")

from argparse import (
    ArgumentParser,
    ArgumentTypeError,
    FileType,
)
from ipaddress import ip_address
from itertools import chain
from json import load
from socket import (
    AF_INET,
    AF_INET6,
    gethostname,
)
from sys import exit

from google.protobuf.json_format import (
    MessageToDict,
    MessageToJson,
)
from grpc import RpcError
from psutil import net_if_addrs
from yandex.cloud.dns.v1.dns_zone_pb2 import RecordSet
from yandex.cloud.dns.v1.dns_zone_service_pb2 import (
    ListDnsZonesRequest,
    RecordSetDiff,
    UpsertRecordSetsMetadata,
    UpsertRecordSetsRequest,
)
from yandex.cloud.dns.v1.dns_zone_service_pb2_grpc import (
    DnsZoneServiceStub,
)
from yandex.cloud.resourcemanager.v1.cloud_service_pb2 import ListCloudsRequest
from yandex.cloud.resourcemanager.v1.cloud_service_pb2_grpc import (
    CloudServiceStub,
)
from yandex.cloud.resourcemanager.v1.folder_service_pb2 import (
    ListFoldersRequest,
)
from yandex.cloud.resourcemanager.v1.folder_service_pb2_grpc import (
    FolderServiceStub,
)
from yandexcloud import SDK

if OPENBSD:
    pledge("stdio rpath inet dns tty unveil")


def debug(msg):
    if args.debug:
        print(MessageToJson(msg))


def domain_name(arg):
    if (domain := arg.strip(".*@")) != arg:
        raise ArgumentTypeError(f"{arg}: Invalid record name")

    return domain


global_ips = lambda iface: filter(
    lambda a: (
        a.family in (AF_INET, AF_INET6) and ip_address(a.address).is_global
    ),
    iface,
)


def interface(arg):
    if arg not in interfaces:
        raise ArgumentTypeError(f"{arg}: No such interface")

    return global_ips(interfaces[arg])


KEY = "./authorized_key.json"

host, domain = gethostname().split(".", maxsplit=1)

interfaces = net_if_addrs()


parser = ArgumentParser(
    description="Dynamic DNS client for Yandex Cloud",
    usage="%(prog)s [-dhn] [-k key] [-r record] [-z zone] [interface ...]",
    color=False,
)
parser.add_argument(
    "-d", "--debug", help="show API messages (JSON)", action="store_true"
)
parser.add_argument(
    "-n", "--noop", help="test run, implies -d", action="store_true"
)
parser.add_argument(
    "-k",
    "--key",
    help=f"account credentials (default: {KEY})",
    type=FileType(),
    default=KEY,
)
parser.add_argument(
    "-r",
    "--record",
    help="(default: host name)",
    type=domain_name,
    default=host,
)
parser.add_argument(
    "-z",
    "--zone",
    help="(default: domain name)",
    type=domain_name,
    default=domain,
)
parser.add_argument(
    "interface",
    help="global IP address(es) (default: all interfaces)",
    type=interface,
    nargs="*",
    default=map(global_ips, interfaces.values()),
)
args = parser.parse_args()

if args.noop:
    args.debug = True

if not (addresses := tuple(chain.from_iterable(args.interface))):
    exit("No public addresses")


if OPENBSD:
    for rpath in (
        "/dev/urandom",
        "/etc/hosts",
        "/etc/resolv.conf",
        "/usr/local/lib/python3.14/site-packages/grpc",
        "/usr/share/zoneinfo",
    ):
        unveil(rpath, "r")
    pledge("stdio rpath inet dns")


with args.key as f:
    try:
        key = load(f)
    except ValueError as e:
        exit(f"{f.name}: {e}")

try:
    sdk = SDK(service_account_key=key)
    del key
except RuntimeError as e:
    exit(f"SDK: {e}")

service = sdk.client(CloudServiceStub)
try:
    response = service.List(ListCloudsRequest())
except RpcError as e:
    exit(f"ListCloudsRequest: {e}")
debug(response)


zone_id = None

service = sdk.client(FolderServiceStub)
for cloud in response.clouds:
    request = ListFoldersRequest(cloud_id=cloud.id)
    debug(request)

    try:
        response = service.List(request)
    except RpcError as e:
        exit(f"ListFoldersRequest: {e.details()}")
    debug(response)

    service = sdk.client(DnsZoneServiceStub)
    for folder in response.folders:
        request = ListDnsZonesRequest(
            folder_id=folder.id,
            filter=f'zone="{args.zone}."',
        )
        debug(request)

        try:
            response = service.List(request)
        except RpcError as e:
            exit(f"ListDnsZonesRequest: {e.details()}")
        debug(response)

        if zones := response.dns_zones:
            zone_id = zones[0].id
            break
    else:
        continue
    break

if not zone_id:
    exit(f"{args.zone}: No such zone")


make_set = lambda af: RecordSet(
    name=args.record,
    ttl=600,
    type={
        AF_INET6: "AAAA",
        AF_INET: "A",
    }[af],
    data=(a.address for a in filter(lambda a: a.family == af, addresses)),
)

request = UpsertRecordSetsRequest(
    dns_zone_id=zone_id,
    replacements=filter(
        lambda r: r.data,
        map(
            make_set,
            (AF_INET, AF_INET6),
        ),
    ),
)
debug(request)

if args.noop:
    exit(0)


try:
    operation = service.UpsertRecordSets(request)

    result = sdk.wait_operation_and_get_result(
        operation,
        response_type=RecordSetDiff,
        meta_type=UpsertRecordSetsMetadata,
    )
except RpcError as e:
    exit(f"UpsertRecordSets: {e.details()}")
debug(operation)


if args.debug or not (response := MessageToDict(result.response)):
    exit(0)

new, old = (
    set(response.get(k, ({},))[0].get("data", ()))
    for k in ("additions", "deletions")
)

if adds := new - old:
    print("added:", *adds)
if dels := old - new:
    print("deleted:", *dels)
