import argparse
import ipaddress
import json
import queue
import time
from http.server import BaseHTTPRequestHandler, HTTPServer

from zeroconf import IPVersion, ServiceBrowser, ServiceInfo, ServiceListener, Zeroconf


SERVICE_TYPE = "_mdnsdemo._tcp.local."


class Handler(BaseHTTPRequestHandler):
    def do_GET(self):
        if self.path != "/hello":
            self.send_error(404)
            return
        body = b"Hello from DNS-SD demo!\n"
        self.send_response(200)
        self.send_header("Content-Type", "text/plain; charset=utf-8")
        self.send_header("Content-Length", str(len(body)))
        self.end_headers()
        self.wfile.write(body)


def serve(zc, ip, port, name):
    with HTTPServer((ip, port), Handler) as server:
        info = ServiceInfo(
            SERVICE_TYPE,
            f"{name}.{SERVICE_TYPE}",
            parsed_addresses=[ip],
            port=server.server_port,
            properties={"path": "/hello", "version": "1"},
            server=f"mdnsdemo-{ipaddress.IPv4Address(ip).packed.hex()}.local.",
        )
        zc.register_service(info, allow_name_change=True)
        try:
            print(f"Published {info.name} on {ip}:{info.port}", flush=True)
            server.serve_forever()
        finally:
            zc.unregister_service(info)


class Listener(ServiceListener):
    def __init__(self, names):
        self.names = names

    def add_service(self, zc, type_, name):
        self.names.put(name)

    def update_service(self, zc, type_, name):
        self.names.put(name)

    def remove_service(self, zc, type_, name):
        pass


def discover(zc, timeout):
    names = queue.Queue()
    browser = ServiceBrowser(zc, SERVICE_TYPE, listener=Listener(names))
    deadline = time.monotonic() + timeout
    try:
        while (remaining := deadline - time.monotonic()) > 0:
            try:
                name = names.get(timeout=remaining)
            except queue.Empty:
                break
            # 在主线程解析记录，避免阻塞发现回调；只输出地址，不自动访问陌生服务。
            info = zc.get_service_info(SERVICE_TYPE, name, timeout=1000)
            if info and info.parsed_addresses():
                print(json.dumps({
                    "name": info.name,
                    "server": info.server,
                    "addresses": info.parsed_addresses(),
                    "port": info.port,
                    "properties": info.decoded_properties,
                }, ensure_ascii=False), flush=True)
                return
        raise SystemExit("No service found before timeout")
    finally:
        browser.cancel()


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("mode", choices=["serve", "discover"])
    parser.add_argument("--ip", required=True, type=ipaddress.IPv4Address)
    parser.add_argument("--port", type=int, default=8000)
    parser.add_argument("--name", default="Demo")
    parser.add_argument("--timeout", type=float, default=10)
    args = parser.parse_args()
    if args.ip.is_unspecified or args.ip.is_multicast:
        parser.error("--ip must be a local unicast IPv4 address")
    if not 0 <= args.port <= 65535 or not 0 < args.timeout < float("inf"):
        parser.error("port must be 0..65535 and timeout must be finite and positive")
    with Zeroconf(interfaces=[str(args.ip)], ip_version=IPVersion.V4Only) as zc:
        try:
            if args.mode == "serve":
                serve(zc, str(args.ip), args.port, args.name)
            else:
                discover(zc, args.timeout)
        except KeyboardInterrupt:
            pass


if __name__ == "__main__":
    main()
