import AppKit
import ApplicationServices
import Foundation

private let canvasRole = "AXWebArea"
private let maxDepth = 32
private let maxNodes = 20_000

enum Diagnosis: String, CaseIterable {
    case pass = "PASS"
    case screenModelInvalid = "SCREEN_MODEL_INVALID"
    case traversalLimit = "TRAVERSAL_LIMIT"
    case traversalAXError = "TRAVERSAL_AX_ERROR"
    case traversalDataInvalid = "TRAVERSAL_DATA_INVALID"
    case webAreaZero = "WEB_AREA_ZERO"
    case webAreaMultiple = "WEB_AREA_MULTIPLE"
    case hiddenValueInvalid = "HIDDEN_VALUE_INVALID"
    case hiddenTrue = "HIDDEN_TRUE"
    case geometryMissingOrInvalid = "GEOMETRY_MISSING_OR_INVALID"
    case outsideVisibleDisplay = "OUTSIDE_VISIBLE_DISPLAY"

    var line: String { "diagnosis=\(rawValue)" }
}

let readinessDeadlineNanoseconds: UInt64 = 5_000_000_000
let readinessMaxSamples = 50
let readinessIntervalNanoseconds: UInt64 = 100_000_000
let readinessRequiredConsecutivePasses = 3

enum ReadinessOutcome: Equatable {
    case pass
    case timeoutWebAreaZero
    case fatal(Diagnosis)

    var line: String {
        switch self {
        case .pass: return "readiness=PASS"
        case .timeoutWebAreaZero: return "readiness=TIMEOUT_WEB_AREA_ZERO"
        case .fatal(let diagnosis): return "readiness=\(diagnosis.rawValue)"
        }
    }
}

struct ReadinessResult {
    let outcome: ReadinessOutcome
    let sampleCount: Int
}

func runReadiness(
    nowNanoseconds: () -> UInt64,
    sample: () -> Diagnosis,
    sleep: (UInt64) -> Void
) -> ReadinessResult {
    let start = nowNanoseconds()
    var consecutivePasses = 0
    for index in 0..<readinessMaxSamples {
        if index > 0, nowNanoseconds() &- start >= readinessDeadlineNanoseconds {
            return ReadinessResult(outcome: .timeoutWebAreaZero, sampleCount: index)
        }
        let diagnosis = sample()
        let count = index + 1
        switch diagnosis {
        case .pass:
            consecutivePasses += 1
            if consecutivePasses == readinessRequiredConsecutivePasses {
                return ReadinessResult(outcome: .pass, sampleCount: count)
            }
        case .webAreaZero:
            consecutivePasses = 0
        default:
            return ReadinessResult(outcome: .fatal(diagnosis), sampleCount: count)
        }
        if count == readinessMaxSamples || nowNanoseconds() &- start >= readinessDeadlineNanoseconds {
            return ReadinessResult(outcome: .timeoutWebAreaZero, sampleCount: count)
        }
        sleep(readinessIntervalNanoseconds)
    }
    return ReadinessResult(outcome: .timeoutWebAreaZero, sampleCount: readinessMaxSamples)
}

enum ReadValue<Value> {
    case value(Value)
    case absent
    case invalid
}

struct CanvasDescriptor {
    let hidden: ReadValue<Bool>
    let rect: ReadValue<CGRect>
}

enum TraversalFailure: Error {
    case limit
    case axError
    case dataInvalid
}

func validRect(_ rect: CGRect) -> Bool {
    [rect.origin.x, rect.origin.y, rect.width, rect.height].allSatisfy(\.isFinite)
        && rect.width > 0 && rect.height > 0
}

func convertVisibleFrames(primaryFrame: CGRect?, visibleFrames: [CGRect]) -> [CGRect]? {
    guard let primaryFrame, validRect(primaryFrame), !visibleFrames.isEmpty,
          visibleFrames.allSatisfy(validRect) else { return nil }
    let primaryTop = primaryFrame.maxY
    return visibleFrames.map {
        CGRect(x: $0.minX, y: primaryTop - $0.maxY, width: $0.width, height: $0.height)
    }
}

func collect<Node>(
    _ node: Node,
    depth: Int = 0,
    visited: inout Int,
    into descriptors: inout [CanvasDescriptor],
    role: (Node) throws -> ReadValue<String>,
    children: (Node) throws -> ReadValue<[Node]>,
    canvas: (Node) throws -> CanvasDescriptor
) throws {
    guard depth <= maxDepth, visited < maxNodes else { throw TraversalFailure.limit }
    visited += 1
    switch try role(node) {
    case .value(let value):
        if value == canvasRole { descriptors.append(try canvas(node)) }
    case .absent:
        break
    case .invalid:
        throw TraversalFailure.dataInvalid
    }
    switch try children(node) {
    case .value(let nodes):
        for child in nodes {
            try collect(
                child,
                depth: depth + 1,
                visited: &visited,
                into: &descriptors,
                role: role,
                children: children,
                canvas: canvas
            )
        }
    case .absent:
        break
    case .invalid:
        throw TraversalFailure.dataInvalid
    }
}

func classify(
    screenFrames: [CGRect]?,
    traversal: Result<[CanvasDescriptor], TraversalFailure>
) -> Diagnosis {
    guard let screenFrames, !screenFrames.isEmpty, screenFrames.allSatisfy(validRect) else {
        return .screenModelInvalid
    }
    let descriptors: [CanvasDescriptor]
    switch traversal {
    case .failure(.limit): return .traversalLimit
    case .failure(.axError): return .traversalAXError
    case .failure(.dataInvalid): return .traversalDataInvalid
    case .success(let value): descriptors = value
    }
    guard !descriptors.isEmpty else { return .webAreaZero }
    guard descriptors.count == 1 else { return .webAreaMultiple }
    let descriptor = descriptors[0]
    switch descriptor.hidden {
    case .absent: break
    case .invalid: return .hiddenValueInvalid
    case .value(true): return .hiddenTrue
    case .value(false): break
    }
    let rect: CGRect
    switch descriptor.rect {
    case .absent, .invalid: return .geometryMissingOrInvalid
    case .value(let value):
        guard validRect(value) else { return .geometryMissingOrInvalid }
        rect = value
    }
    guard screenFrames.contains(where: { $0.contains(rect) }) else { return .outsideVisibleDisplay }
    return .pass
}

private struct FixtureNode {
    let role: ReadValue<String>
    let children: ReadValue<[Int]>
    let hidden: ReadValue<Bool>
    let rect: ReadValue<CGRect>
}

private func fixtureTraversal(
    _ nodes: [Int: FixtureNode],
    root: Int = 0,
    axErrorAt: Int? = nil
) -> Result<[CanvasDescriptor], TraversalFailure> {
    var visited = 0
    var descriptors: [CanvasDescriptor] = []
    do {
        try collect(
            root,
            visited: &visited,
            into: &descriptors,
            role: {
                if $0 == axErrorAt { throw TraversalFailure.axError }
                return nodes[$0]?.role ?? .invalid
            },
            children: { nodes[$0]?.children ?? .invalid },
            canvas: {
                guard let node = nodes[$0] else { throw TraversalFailure.dataInvalid }
                return CanvasDescriptor(hidden: node.hidden, rect: node.rect)
            }
        )
        return .success(descriptors)
    } catch let failure as TraversalFailure {
        return .failure(failure)
    } catch {
        return .failure(.axError)
    }
}

func selfTest() -> Int32 {
    let display = CGRect(x: 0, y: 0, width: 1920, height: 1040)
    let validCanvas = CanvasDescriptor(
        hidden: .value(false),
        rect: .value(CGRect(x: 100, y: 100, width: 800, height: 600))
    )
    func readinessFixture(
        _ samples: [Diagnosis],
        advancePerSleep: UInt64 = readinessIntervalNanoseconds
    ) -> (ReadinessResult, Int, Int) {
        var index = 0
        var now: UInt64 = 0
        var sleeps = 0
        let result = runReadiness(
            nowNanoseconds: { now },
            sample: {
                let diagnosis = index < samples.count ? samples[index] : samples.last!
                index += 1
                return diagnosis
            },
            sleep: { _ in sleeps += 1; now += advancePerSleep }
        )
        return (result, index, sleeps)
    }
    let zeroThenPass = readinessFixture([.webAreaZero, .pass, .pass, .pass])
    let zeroTimeout = readinessFixture([.webAreaZero])
    let resetThenPass = readinessFixture([.pass, .pass, .webAreaZero, .pass, .pass, .pass])
    let fatalFirst = readinessFixture([.webAreaMultiple, .pass, .pass, .pass])
    let fatalAfterZero = readinessFixture([.webAreaZero, .hiddenTrue, .pass, .pass, .pass])
    let deadline = readinessFixture([.webAreaZero], advancePerSleep: readinessDeadlineNanoseconds)
    guard zeroThenPass.0.outcome == .pass,
          zeroThenPass.0.sampleCount == 4, zeroThenPass.1 == 4, zeroThenPass.2 == 3,
          zeroTimeout.0.outcome == .timeoutWebAreaZero,
          zeroTimeout.0.sampleCount == readinessMaxSamples,
          zeroTimeout.1 == readinessMaxSamples,
          zeroTimeout.2 == readinessMaxSamples - 1,
          resetThenPass.0.outcome == .pass,
          resetThenPass.0.sampleCount == 6, resetThenPass.1 == 6, resetThenPass.2 == 5,
          fatalFirst.0.outcome == .fatal(.webAreaMultiple),
          fatalFirst.0.sampleCount == 1, fatalFirst.1 == 1, fatalFirst.2 == 0,
          fatalAfterZero.0.outcome == .fatal(.hiddenTrue),
          fatalAfterZero.0.sampleCount == 2, fatalAfterZero.1 == 2, fatalAfterZero.2 == 1,
          deadline.0.outcome == .timeoutWebAreaZero,
          deadline.0.sampleCount == 1, deadline.1 == 1, deadline.2 == 1,
          ReadinessOutcome.pass.line == "readiness=PASS",
          ReadinessOutcome.timeoutWebAreaZero.line == "readiness=TIMEOUT_WEB_AREA_ZERO",
          Diagnosis.allCases.filter({ $0 != .pass && $0 != .webAreaZero }).allSatisfy({
              ReadinessOutcome.fatal($0).line == "readiness=\($0.rawValue)"
          }) else { return 1 }
    let leaf = FixtureNode(
        role: .value(canvasRole), children: .absent,
        hidden: validCanvas.hidden, rect: validCanvas.rect
    )
    let root = FixtureNode(
        role: .value("AXApplication"), children: .value([1]),
        hidden: .absent, rect: .absent
    )
    let base = [0: root, 1: leaf]
    let noCanvas = [
        0: FixtureNode(role: .value("AXApplication"), children: .absent, hidden: .absent, rect: .absent)
    ]
    let twoCanvases = [
        0: FixtureNode(role: .value("AXApplication"), children: .value([1, 2]), hidden: .absent, rect: .absent),
        1: leaf,
        2: leaf,
    ]
    let allCases: [(Diagnosis, [CGRect]?, Result<[CanvasDescriptor], TraversalFailure>)] = [
        (.pass, [display], .success([validCanvas])),
        (.screenModelInvalid, nil, .success([validCanvas])),
        (.traversalLimit, [display], .failure(.limit)),
        (.traversalAXError, [display], .failure(.axError)),
        (.traversalDataInvalid, [display], .failure(.dataInvalid)),
        (.webAreaZero, [display], .success([])),
        (.webAreaMultiple, [display], .success([validCanvas, validCanvas])),
        (.pass, [display], .success([CanvasDescriptor(hidden: .absent, rect: validCanvas.rect)])),
        (.hiddenValueInvalid, [display], .success([CanvasDescriptor(hidden: .invalid, rect: validCanvas.rect)])),
        (.hiddenTrue, [display], .success([CanvasDescriptor(hidden: .value(true), rect: validCanvas.rect)])),
        (.geometryMissingOrInvalid, [display], .success([CanvasDescriptor(hidden: .value(false), rect: .absent)])),
        (.outsideVisibleDisplay, [display], .success([CanvasDescriptor(hidden: .value(false), rect: .value(CGRect(x: 1900, y: 100, width: 100, height: 100)))])),
    ]
    let invalidRects = [
        CGRect(x: 0, y: 0, width: 0, height: 1),
        CGRect(x: 0, y: 0, width: 1, height: 0),
        CGRect(x: CGFloat.nan, y: 0, width: 1, height: 1),
        CGRect(x: 0, y: CGFloat.infinity, width: 1, height: 1),
    ]
    let converted = convertVisibleFrames(
        primaryFrame: CGRect(x: 0, y: 0, width: 1920, height: 1080),
        visibleFrames: [
            CGRect(x: 0, y: 0, width: 1920, height: 1055),
            CGRect(x: 0, y: 1080, width: 1600, height: 1200),
            CGRect(x: -1280, y: 0, width: 1280, height: 1024),
        ]
    )
    var deepNodes: [Int: FixtureNode] = [:]
    for index in 0...maxDepth + 1 {
        deepNodes[index] = FixtureNode(
            role: .value("AXGroup"),
            children: index == maxDepth + 1 ? .absent : .value([index + 1]),
            hidden: .absent, rect: .absent
        )
    }
    var wideChildren: [Int] = []
    var wideNodes: [Int: FixtureNode] = [0: FixtureNode(
        role: .value("AXApplication"), children: .value([]), hidden: .absent, rect: .absent
    )]
    for index in 1...maxNodes {
        wideChildren.append(index)
        wideNodes[index] = FixtureNode(
            role: .value("AXGroup"), children: .absent, hidden: .absent, rect: .absent
        )
    }
    wideNodes[0] = FixtureNode(
        role: .value("AXApplication"), children: .value(wideChildren), hidden: .absent, rect: .absent
    )
    let invalidChildren = [
        0: FixtureNode(role: .value("AXApplication"), children: .invalid, hidden: .absent, rect: .absent)
    ]
    do {
        guard case .absent = try hiddenRead(error: .noValue, value: nil),
              case .absent = try hiddenRead(error: .attributeUnsupported, value: nil),
              case .invalid = try hiddenRead(error: .success, value: nil),
              case .invalid = try hiddenRead(error: .success, value: "invalid" as CFString),
              case .value(true) = try hiddenRead(error: .success, value: kCFBooleanTrue),
              case .value(false) = try hiddenRead(error: .success, value: kCFBooleanFalse) else {
            return 1
        }
        do {
            _ = try hiddenRead(error: .cannotComplete, value: nil)
            return 1
        } catch TraversalFailure.axError {
        }
    } catch {
        return 1
    }
    guard allCases.allSatisfy({ classify(screenFrames: $0.1, traversal: $0.2) == $0.0 }),
          Diagnosis.pass == classify(
              screenFrames: [display],
              traversal: .success([CanvasDescriptor(hidden: .absent, rect: validCanvas.rect)])
          ),
          Diagnosis(rawValue: "HIDDEN_VALUE_INVALID") == classify(
              screenFrames: [display],
              traversal: .success([CanvasDescriptor(hidden: .invalid, rect: validCanvas.rect)])
          ),
          Diagnosis.allCases.allSatisfy({ $0.line.split(separator: "\n").count == 1 && $0.line == "diagnosis=\($0.rawValue)" }),
          classify(screenFrames: [display], traversal: fixtureTraversal(base)) == .pass,
          classify(screenFrames: [display], traversal: fixtureTraversal(noCanvas)) == .webAreaZero,
          classify(screenFrames: [display], traversal: fixtureTraversal(twoCanvases)) == .webAreaMultiple,
          classify(screenFrames: [display], traversal: fixtureTraversal(base, axErrorAt: 1)) == .traversalAXError,
          classify(screenFrames: [display], traversal: fixtureTraversal([0: root])) == .traversalDataInvalid,
          classify(screenFrames: [display], traversal: fixtureTraversal(invalidChildren)) == .traversalDataInvalid,
          classify(screenFrames: [display], traversal: fixtureTraversal(deepNodes)) == .traversalLimit,
          classify(screenFrames: [display], traversal: fixtureTraversal(wideNodes)) == .traversalLimit,
          invalidRects.allSatisfy({
              classify(
                  screenFrames: [display],
                  traversal: .success([CanvasDescriptor(hidden: .value(false), rect: .value($0))])
              ) == .geometryMissingOrInvalid
          }),
          classify(
              screenFrames: [display],
              traversal: .success([CanvasDescriptor(hidden: .invalid, rect: validCanvas.rect)])
          ) == .hiddenValueInvalid,
          classify(
              screenFrames: [display],
              traversal: .success([CanvasDescriptor(hidden: .absent, rect: .absent)])
          ) == .geometryMissingOrInvalid,
          classify(
              screenFrames: [display],
              traversal: .success([CanvasDescriptor(hidden: .absent, rect: .invalid)])
          ) == .geometryMissingOrInvalid,
          classify(
              screenFrames: [display],
              traversal: .success([CanvasDescriptor(
                  hidden: .absent,
                  rect: .value(CGRect(x: 1900, y: 100, width: 100, height: 100))
              )])
          ) == .outsideVisibleDisplay,
          classify(
              screenFrames: [display],
              traversal: .success([CanvasDescriptor(hidden: .value(false), rect: .invalid)])
          ) == .geometryMissingOrInvalid,
          let converted,
          converted == [
              CGRect(x: 0, y: 25, width: 1920, height: 1055),
              CGRect(x: 0, y: -1200, width: 1600, height: 1200),
              CGRect(x: -1280, y: 56, width: 1280, height: 1024),
          ],
          convertVisibleFrames(primaryFrame: nil, visibleFrames: [display]) == nil,
          convertVisibleFrames(primaryFrame: display, visibleFrames: []) == nil,
          classify(screenFrames: nil, traversal: .failure(.axError)) == .screenModelInvalid else { return 1 }
    print(Diagnosis.pass.line)
    return 0
}

if CommandLine.arguments == [CommandLine.arguments[0], "--self-test"] {
    exit(selfTest())
}

let waitMode: Bool
let candidatePID: pid_t
if CommandLine.arguments.count == 2,
   let parsedPID = pid_t(CommandLine.arguments[1]), parsedPID > 0 {
    waitMode = false
    candidatePID = parsedPID
} else if CommandLine.arguments.count == 3,
          CommandLine.arguments[1] == "--wait",
          let parsedPID = pid_t(CommandLine.arguments[2]), parsedPID > 0 {
    waitMode = true
    candidatePID = parsedPID
} else {
    exit(64)
}

func attribute(_ element: AXUIElement, _ name: String) throws -> ReadValue<CFTypeRef> {
    var value: CFTypeRef?
    let error = AXUIElementCopyAttributeValue(element, name as CFString, &value)
    if error == .success {
        guard let value else { return .invalid }
        return .value(value)
    }
    if error == .noValue || error == .attributeUnsupported { return .absent }
    throw TraversalFailure.axError
}

func stringValue(_ element: AXUIElement, _ name: String) throws -> ReadValue<String> {
    switch try attribute(element, name) {
    case .absent: return .absent
    case .invalid: return .invalid
    case .value(let value):
        guard CFGetTypeID(value) == CFStringGetTypeID(), let string = value as? String else { return .invalid }
        return .value(string)
    }
}

func hiddenRead(error: AXError, value: CFTypeRef?) throws -> ReadValue<Bool> {
    if error == .noValue || error == .attributeUnsupported { return .absent }
    guard error == .success else { throw TraversalFailure.axError }
    guard let value, CFGetTypeID(value) == CFBooleanGetTypeID() else { return .invalid }
    return .value(CFBooleanGetValue((value as! CFBoolean)))
}

func hiddenValue(_ element: AXUIElement) throws -> ReadValue<Bool> {
    var value: CFTypeRef?
    let error = AXUIElementCopyAttributeValue(element, "AXHidden" as CFString, &value)
    return try hiddenRead(error: error, value: value)
}

func rectValue(_ element: AXUIElement) throws -> ReadValue<CGRect> {
    let position = try attribute(element, kAXPositionAttribute)
    let size = try attribute(element, kAXSizeAttribute)
    guard case .value(let positionValue) = position,
          case .value(let sizeValue) = size else {
        if case .invalid = position { return .invalid }
        if case .invalid = size { return .invalid }
        return .absent
    }
    guard CFGetTypeID(positionValue) == AXValueGetTypeID(),
          CFGetTypeID(sizeValue) == AXValueGetTypeID() else { return .invalid }
    var point = CGPoint.zero
    var dimensions = CGSize.zero
    guard AXValueGetValue(positionValue as! AXValue, .cgPoint, &point),
          AXValueGetValue(sizeValue as! AXValue, .cgSize, &dimensions) else { return .invalid }
    return .value(CGRect(origin: point, size: dimensions))
}

func childElements(_ element: AXUIElement) throws -> ReadValue<[AXUIElement]> {
    switch try attribute(element, kAXChildrenAttribute) {
    case .absent: return .absent
    case .invalid: return .invalid
    case .value(let value):
        guard CFGetTypeID(value) == CFArrayGetTypeID(),
              let children = value as? [AXUIElement] else { return .invalid }
        return .value(children)
    }
}

func screenFrames() -> [CGRect]? {
    let screens = NSScreen.screens
    guard let primary = screens.first else { return nil }
    return convertVisibleFrames(primaryFrame: primary.frame, visibleFrames: screens.map(\.visibleFrame))
}

func diagnose(_ candidatePID: pid_t) -> Diagnosis {
    let app = AXUIElementCreateApplication(candidatePID)
    var visited = 0
    var descriptors: [CanvasDescriptor] = []
    let traversal: Result<[CanvasDescriptor], TraversalFailure>
    do {
        try collect(
            app,
            visited: &visited,
            into: &descriptors,
            role: { try stringValue($0, kAXRoleAttribute) },
            children: childElements,
            canvas: { CanvasDescriptor(hidden: try hiddenValue($0), rect: try rectValue($0)) }
        )
        traversal = .success(descriptors)
    } catch let failure as TraversalFailure {
        traversal = .failure(failure)
    } catch {
        traversal = .failure(.axError)
    }
    return classify(screenFrames: screenFrames(), traversal: traversal)
}

if waitMode {
    let readiness = runReadiness(
        nowNanoseconds: { DispatchTime.now().uptimeNanoseconds },
        sample: { diagnose(candidatePID) },
        sleep: { usleep(useconds_t($0 / 1_000)) }
    )
    print(readiness.outcome.line)
    exit(readiness.outcome == .pass ? 0 : 67)
} else {
    let diagnosis = diagnose(candidatePID)
    print(diagnosis.line)
    exit(diagnosis == .pass ? 0 : 67)
}
