#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
dnsdiag - SysAK DNS 综合诊断工具

双模式设计：
  模式一（一键诊断）：零参数自动跑完 配置检查 + 延迟测试 + 解析验证
    sysak dnsdiag
    sysak dnsdiag -d example.com

  模式二（精细控制）：通过子命令 + 参数实现细粒度诊断
    sysak dnsdiag latency -d example.com -n 10 -s 8.8.8.8 --no-cache
    sysak dnsdiag failure -d example.com --full
    sysak dnsdiag trace -d 10 --top
"""

import sys
import os

# 确保模块搜索路径包含当前脚本所在目录
script_dir = os.path.dirname(os.path.abspath(__file__))
if script_dir not in sys.path:
    sys.path.insert(0, script_dir)

import argparse

# 已知子命令列表
SUBCMDS = ('latency', 'failure', 'trace')


# =============================================================================
# 一键诊断模式
# =============================================================================

def run_oneclick(args):
    """
    一键诊断模式：自动执行完整诊断流水线。

    流程：
    [1/3] 配置检查 — 验证 resolv.conf/nsswitch.conf 配置是否正确
    [2/3] 延迟测试 — 对所有 nameserver 发送查询，测量 RTT
    [3/3] 解析验证 — 发送实际 DNS 查询，检查是否能正确解析

    最终汇总输出结构化结论和建议。
    """
    from dnsdiag_common import (
        read_resolv_conf,
        read_nsswitch_conf,
        send_dns_query,
        get_default_domain,
        format_status,
    )
    from dnsdiag_latency import latency_check
    from dnsdiag_failure import failure_check

    # 确定诊断域名
    domain = args.domain
    if not domain:
        domain = get_default_domain()

    timeout = args.timeout
    conclusions = []  # 收集诊断结论

    print("================ SysAK DNS 一键诊断 ================\n")

    # ===== [1/3] 配置检查 =====
    print("[1/3] 配置检查")
    config = read_resolv_conf()

    # 检查 nameserver
    if config['errors']:
        print(format_status(False, "resolv.conf 错误: %s" % config['errors'][0]))
        conclusions.append(("[FAIL]", "resolv.conf 配置异常"))
    elif not config['nameservers']:
        print(format_status(False, "resolv.conf 中没有 nameserver"))
        conclusions.append(("[FAIL]", "没有配置 DNS 服务器"))
        _print_conclusion(conclusions)
        return 1
    else:
        print(format_status(True, "resolv.conf: %d 个 nameserver (%s)" % (
            len(config['nameservers']), ', '.join(config['nameservers']))))

    # 检查 options
    opts = config['options']
    print(format_status(True, "ndots=%s, timeout=%s, attempts=%s" % (
        opts.get('ndots', 1), opts.get('timeout', 5), opts.get('attempts', 2))))

    # 检查 nsswitch
    nsswitch = read_nsswitch_conf()
    if nsswitch['has_dns']:
        print(format_status(True, "nsswitch: %s" % ' '.join(nsswitch['hosts_order'])))
    elif nsswitch['errors']:
        print(format_status(True, "nsswitch: 文件不可读 (非致命)"))
    else:
        print(format_status(False, "nsswitch hosts 行缺少 dns"))
        conclusions.append(("[WARNING]", "nsswitch.conf 未配置 dns"))

    print("")

    # ===== [2/3] 延迟测试 =====
    print("[2/3] 延迟测试 (目标: %s)" % domain)
    servers = config['nameservers']
    latency_results = latency_check(domain, servers, count=3, timeout=timeout)

    for r in latency_results:
        if r['success'] == 0:
            print(format_status(False, "%-15s 不可达 (丢包率 100%%)" % r['server']))
            conclusions.append(("[FAIL]", "nameserver %s 不可达" % r['server']))
        elif r['status'] == 'SLOW':
            print(format_status(False, "%-15s avg=%.1fms  loss=%.0f%%    [SLOW]" % (
                r['server'], r['avg'], r['loss_pct'])))
            conclusions.append(("[WARNING]", "nameserver %s 延迟异常 (%.0fms), 丢包率 %.0f%%" % (
                r['server'], r['avg'], r['loss_pct'])))
        else:
            print(format_status(True, "%-15s avg=%.1fms  loss=%.0f%%" % (
                r['server'], r['avg'], r['loss_pct'])))

    print("")

    # ===== [3/3] 解析验证 =====
    print("[3/3] 解析验证")
    # 对第一个可达的 server 进行实际查询
    working_server = None
    for r in latency_results:
        if r['success'] > 0:
            working_server = r['server']
            break

    if working_server:
        resp = send_dns_query(working_server, domain, qtype='A', timeout=timeout)
        if resp['error']:
            print(format_status(False, "查询 %s 失败: %s" % (domain, resp['error'])))
            conclusions.append(("[FAIL]", "DNS 查询失败: %s" % resp['error']))
        elif resp['response']:
            rcode = resp['response']['rcode']
            if rcode == 0:
                answers = resp['response']['answers']
                if answers:
                    answer_str = answers[0]['data']
                    ttl = answers[0]['ttl']
                    print(format_status(True, "A记录: %s (TTL=%d)" % (answer_str, ttl)))
                else:
                    print(format_status(True, "NOERROR 但无 A 记录 (可能是其他记录类型)"))
            elif rcode == 3:
                print(format_status(False, "域名 %s 不存在 (NXDOMAIN)" % domain))
                conclusions.append(("[FAIL]", "域名不存在 (NXDOMAIN)"))
            else:
                rcode_name = resp['response']['rcode_name']
                print(format_status(False, "查询返回 %s" % rcode_name))
                conclusions.append(("[FAIL]", "DNS 服务器返回 %s" % rcode_name))
    else:
        print(format_status(False, "所有 nameserver 不可达，无法验证解析"))
        conclusions.append(("[FAIL]", "无可用的 DNS 服务器"))

    # ===== 输出结论 =====
    _print_conclusion(conclusions)

    return 1 if any(c[0] == '[FAIL]' for c in conclusions) else 0


def _print_conclusion(conclusions):
    """输出诊断结论"""
    print("\n================== 诊断结论 ==================")
    if not conclusions:
        print("  [OK] DNS 配置正常，解析正常")
    else:
        for tag, msg in conclusions:
            print("  %s %s" % (tag, msg))
        # 输出建议
        fail_items = [c for c in conclusions if c[0] == '[FAIL]']
        if fail_items:
            print("")
            print("建议:")
            for _, msg in fail_items:
                if '不可达' in msg:
                    print("  - 检查到 DNS 服务器的网络连通性和防火墙规则")
                elif 'NXDOMAIN' in msg:
                    print("  - 检查域名拼写是否正确，或 resolv.conf 的 search 域配置")
                elif 'SERVFAIL' in msg:
                    print("  - 联系 DNS 管理员检查上游权威服务器状态")
                elif '没有配置' in msg:
                    print("  - 检查 /etc/resolv.conf 是否正确配置了 nameserver")


# =============================================================================
# 精细控制模式
# =============================================================================

def run_subcmd_mode():
    """子命令精细控制模式"""
    parser = argparse.ArgumentParser(
        prog='dnsdiag',
        description='SysAK DNS 综合诊断工具 (精细控制模式)')
    subparsers = parser.add_subparsers(dest='command')

    # --- latency 子命令 ---
    lat_parser = subparsers.add_parser('latency', help='DNS 解析延迟诊断')
    lat_parser.add_argument('-d', '--domain', required=True, help='查询域名')
    lat_parser.add_argument('-n', '--count', type=int, default=5, help='重复查询次数 (默认5)')
    lat_parser.add_argument('-s', '--server', action='append', help='指定 DNS 服务器 (可多次)')
    lat_parser.add_argument('--no-cache', action='store_true', help='穿透缓存 (随机子域名)')

    # --- failure 子命令 ---
    fail_parser = subparsers.add_parser('failure', help='DNS 解析失败根因分析')
    fail_parser.add_argument('-d', '--domain', required=True, help='查询域名')
    fail_parser.add_argument('--full', action='store_true', help='深度检查 (含劫持检测)')

    # --- trace 子命令 ---
    trace_parser = subparsers.add_parser('trace', help='进程级 DNS 查询追踪')
    trace_parser.add_argument('-d', '--duration', type=int, default=10, help='捕获时长(秒, 默认10)')
    trace_parser.add_argument('-p', '--pid', type=int, help='只追踪指定 PID')
    trace_parser.add_argument('--top', action='store_true', help='汇总模式 (按进程统计)')
    trace_parser.add_argument('--no-proc-map', action='store_true', help='禁用进程映射')

    args = parser.parse_args()

    if args.command == 'latency':
        from dnsdiag_latency import run_latency_mode
        return run_latency_mode(args)
    elif args.command == 'failure':
        from dnsdiag_failure import run_failure_mode
        return run_failure_mode(args)
    elif args.command == 'trace':
        from dnsdiag_trace import run_trace_mode
        return run_trace_mode(args)
    else:
        parser.print_help()
        return 0


# =============================================================================
# 主入口
# =============================================================================

def main():
    """
    双模式入口分发。

    判断逻辑：
    - 如果第一个参数是已知子命令（latency/failure/trace）→ 精细控制模式
    - 否则 → 一键诊断模式（零参数也走这里）
    """
    if len(sys.argv) > 1 and sys.argv[1] in SUBCMDS:
        ret = run_subcmd_mode()
    else:
        # 一键诊断模式
        parser = argparse.ArgumentParser(
            prog='dnsdiag',
            description='SysAK DNS 一键诊断')
        parser.add_argument('-d', '--domain', default=None,
                            help='诊断域名 (默认自动检测)')
        parser.add_argument('-t', '--timeout', type=int, default=5,
                            help='超时时间(秒, 默认5)')
        args = parser.parse_args()
        ret = run_oneclick(args)

    sys.exit(ret)


if __name__ == '__main__':
    main()
