#!/usr/bin/env swift
// SPDX-License-Identifier: MIT
//
// Per-user lease coordinator for the Pi Caps Lock LED indicator.
// This source is intentionally self-contained so it can be run by `swift`.

import CoreFoundation
import Dispatch
import Foundation
import IOKit.hidsystem
import Darwin

// This HID constructor is exported by IOKit but is not in the public SDK
// module. Client type 2 is the passive event-system client.
@_silgen_name("IOHIDEventSystemClientCreateWithType")
private func createPassiveHIDClient(
    _ allocator: CFAllocator?,
    _ clientType: UInt32,
    _ attributes: CFDictionary?
) -> IOHIDEventSystemClient?

private let runtimeLeaf = "pi-caps-blink-v1"
private let lockLeaf = "l"
private let socketLeaf = "s"
private let queue = DispatchQueue(label: "dev.pi.caps-blink.daemon")

private enum DaemonError: Error, CustomStringConvertible {
    case usage
    case unsupportedOS
    case invalidRuntimeDirectory
    case unsafeRuntimeDirectory(String)
    case lockUnavailable
    case socketPathTooLong
    case systemCall(String)
    case noBuiltInKeyboard
    case ledWriteRejected

    var description: String {
        switch self {
        case .usage:
            return "usage: CapsBlinkDaemon.swift --runtime-dir <DARWIN_USER_TEMP_DIR>/pi-caps-blink-v1 [--fake-led-file <path>]"
        case .unsupportedOS:
            return "macOS 26 or later is required"
        case .invalidRuntimeDirectory:
            return "runtime directory does not match DARWIN_USER_TEMP_DIR"
        case let .unsafeRuntimeDirectory(reason):
            return "unsafe runtime directory: \(reason)"
        case .lockUnavailable:
            return "another daemon owns the runtime lock"
        case .socketPathTooLong:
            return "Unix socket path is too long"
        case let .systemCall(name):
            return "\(name): \(String(cString: strerror(errno)))"
        case .noBuiltInKeyboard:
            return "no built-in keyboard HID service is available"
        case .ledWriteRejected:
            return "macOS rejected the Caps Lock LED update"
        }
    }
}

private enum LEDMode: String {
    case on = "On"
    case off = "Off"
    case auto = "Auto"
}

private protocol LEDBackend {
    func apply(_ mode: LEDMode) throws
}

/// Uses the HID keyboard filter's private override property. It deliberately
/// selects only a built-in Generic Desktop Keyboard service.
private final class HIDLEDBackend: LEDBackend {
    private let client: IOHIDEventSystemClient
    private let ledKey = "HIDCapsLockLED" as CFString
    private let builtInKey = "Built-In" as CFString

    init() throws {
        guard let client = createPassiveHIDClient(kCFAllocatorDefault, 2, nil) else {
            throw DaemonError.noBuiltInKeyboard
        }
        self.client = client
    }

    func apply(_ mode: LEDMode) throws {
        guard let services = IOHIDEventSystemClientCopyServices(client) else {
            throw DaemonError.noBuiltInKeyboard
        }

        var foundBuiltInKeyboard = false
        var acceptedWrite = false
        let count = CFArrayGetCount(services)
        for index in 0..<count {
            guard let rawService = CFArrayGetValueAtIndex(services, index) else { continue }
            let service = unsafeBitCast(rawService, to: IOHIDServiceClient.self)
            guard IOHIDServiceClientConformsTo(service, 0x01, 0x06) != 0 else { continue }
            guard isBuiltIn(service) else { continue }

            foundBuiltInKeyboard = true
            let value = mode.rawValue as NSString
            if IOHIDServiceClientSetProperty(service, ledKey, value) {
                acceptedWrite = true
            }
        }

        guard foundBuiltInKeyboard else { throw DaemonError.noBuiltInKeyboard }
        guard acceptedWrite else { throw DaemonError.ledWriteRejected }
    }

    private func isBuiltIn(_ service: IOHIDServiceClient) -> Bool {
        guard let value = IOHIDServiceClientCopyProperty(service, builtInKey) else { return false }
        return (value as? NSNumber)?.boolValue == true
    }
}

/// Explicitly opt-in test backend. It writes only the latest requested state.
private final class FakeLEDBackend: LEDBackend {
    private let path: String

    init(path: String) {
        self.path = path
    }

    func apply(_ mode: LEDMode) throws {
        let fd = open(path, O_WRONLY | O_CREAT | O_TRUNC | O_CLOEXEC | O_NOFOLLOW, S_IRUSR | S_IWUSR)
        guard fd >= 0 else { throw DaemonError.systemCall("open fake LED file") }
        defer { _ = close(fd) }

        var info = stat()
        guard fstat(fd, &info) == 0 else { throw DaemonError.systemCall("fstat fake LED file") }
        guard (info.st_mode & S_IFMT) == S_IFREG, info.st_uid == getuid() else {
            throw DaemonError.unsafeRuntimeDirectory("fake LED file is not a regular file owned by this user")
        }

        let bytes = Array((mode.rawValue + "\n").utf8)
        var written = 0
        while written < bytes.count {
            let result = bytes.withUnsafeBytes { buffer in
                write(fd, buffer.baseAddress!.advanced(by: written), bytes.count - written)
            }
            if result > 0 {
                written += result
            } else if result < 0 && errno == EINTR {
                continue
            } else {
                throw DaemonError.systemCall("write fake LED file")
            }
        }
    }
}

private struct Arguments {
    let runtimeDirectory: String
    let fakeLEDFile: String?

    static func parse(_ arguments: [String]) throws -> Arguments {
        var runtimeDirectory: String?
        var fakeLEDFile: String?
        var index = 1
        while index < arguments.count {
            let argument = arguments[index]
            guard index + 1 < arguments.count else { throw DaemonError.usage }
            let value = arguments[index + 1]
            switch argument {
            case "--runtime-dir":
                guard runtimeDirectory == nil else { throw DaemonError.usage }
                runtimeDirectory = value
            case "--fake-led-file":
                guard fakeLEDFile == nil else { throw DaemonError.usage }
                fakeLEDFile = value
            default:
                throw DaemonError.usage
            }
            index += 2
        }
        guard let runtimeDirectory else { throw DaemonError.usage }
        return Arguments(runtimeDirectory: runtimeDirectory, fakeLEDFile: fakeLEDFile)
    }
}

private func darwinUserTemporaryDirectory() throws -> String {
    let required = confstr(_CS_DARWIN_USER_TEMP_DIR, nil, 0)
    guard required > 1 else { throw DaemonError.systemCall("confstr(DARWIN_USER_TEMP_DIR)") }
    var buffer = [CChar](repeating: 0, count: Int(required))
    guard confstr(_CS_DARWIN_USER_TEMP_DIR, &buffer, buffer.count) == required else {
        throw DaemonError.systemCall("confstr(DARWIN_USER_TEMP_DIR)")
    }
    let root = String(cString: buffer)
    guard root.hasPrefix("/") else { throw DaemonError.invalidRuntimeDirectory }
    return root
}

private func canonicalDARWINUserTemporaryDirectory() throws -> String {
    let rawRoot = try darwinUserTemporaryDirectory()
    var resolved = [CChar](repeating: 0, count: Int(PATH_MAX))
    guard rawRoot.withCString({ realpath($0, &resolved) != nil }) else {
        throw DaemonError.systemCall("realpath(DARWIN_USER_TEMP_DIR)")
    }

    let canonicalRoot = String(cString: resolved)
    guard canonicalRoot.hasPrefix("/") else { throw DaemonError.invalidRuntimeDirectory }
    if canonicalRoot == "/" { return canonicalRoot }
    return canonicalRoot.hasSuffix("/") ? String(canonicalRoot.dropLast()) : canonicalRoot
}

private func checkOwnedDirectory(_ path: String) throws {
    var info = stat()
    guard lstat(path, &info) == 0 else { throw DaemonError.systemCall("lstat runtime directory") }
    guard (info.st_mode & S_IFMT) == S_IFDIR else {
        throw DaemonError.unsafeRuntimeDirectory("not a directory or is a symlink")
    }
    guard info.st_uid == getuid() else { throw DaemonError.unsafeRuntimeDirectory("wrong owner") }
    guard (info.st_mode & 0o777) == 0o700 else { throw DaemonError.unsafeRuntimeDirectory("permissions are not 0700") }
}

private func prepareRuntimeDirectory(passedPath: String) throws -> String {
    // The TypeScript launcher passes the realpath form. Canonicalize the
    // independently obtained confstr root too, because on this system confstr
    // begins with /var while realpath produces /private/var.
    let canonicalRoot = try canonicalDARWINUserTemporaryDirectory()
    let expected = canonicalRoot == "/" ? "/" + runtimeLeaf : canonicalRoot + "/" + runtimeLeaf
    guard passedPath == expected else { throw DaemonError.invalidRuntimeDirectory }

    if mkdir(expected, 0o700) != 0 && errno != EEXIST {
        throw DaemonError.systemCall("mkdir runtime directory")
    }
    try checkOwnedDirectory(expected)
    return expected
}

private func setNonBlockingAndCloseOnExec(_ fd: Int32) throws {
    let flags = fcntl(fd, F_GETFL)
    guard flags >= 0, fcntl(fd, F_SETFL, flags | O_NONBLOCK) == 0 else {
        throw DaemonError.systemCall("fcntl(O_NONBLOCK)")
    }
    let descriptorFlags = fcntl(fd, F_GETFD)
    guard descriptorFlags >= 0, fcntl(fd, F_SETFD, descriptorFlags | FD_CLOEXEC) == 0 else {
        throw DaemonError.systemCall("fcntl(FD_CLOEXEC)")
    }
}

private final class Daemon {
    private let runtimeDirectory: String
    private let lockPath: String
    private let socketPath: String
    private let backend: LEDBackend
    private let lockFD: Int32

    private var listenerFD: Int32 = -1
    private var listenerSource: DispatchSourceRead?
    private var clients: [Int32: DispatchSourceRead] = [:]
    private var blinkTimer: DispatchSourceTimer?
    private var idleTimer: DispatchSourceTimer?
    private var signalSources: [DispatchSourceSignal] = []
    private var ledIsOn = false
    private var blinkGeneration: UInt64 = 0
    private var shuttingDown = false
    private var pendingFDCloses = Set<Int32>()

    init(runtimeDirectory: String, backend: LEDBackend) throws {
        self.runtimeDirectory = runtimeDirectory
        self.lockPath = runtimeDirectory + "/" + lockLeaf
        self.socketPath = runtimeDirectory + "/" + socketLeaf
        self.backend = backend

        let fd = open(lockPath, O_RDWR | O_CREAT | O_CLOEXEC | O_NOFOLLOW, S_IRUSR | S_IWUSR)
        guard fd >= 0 else { throw DaemonError.systemCall("open lock") }
        self.lockFD = fd

        var info = stat()
        guard fstat(fd, &info) == 0 else {
            _ = close(fd)
            throw DaemonError.systemCall("fstat lock")
        }
        guard (info.st_mode & S_IFMT) == S_IFREG, info.st_uid == getuid(), (info.st_mode & 0o777) == 0o600 else {
            _ = close(fd)
            throw DaemonError.unsafeRuntimeDirectory("lock is not a 0600 regular file owned by this user")
        }
        guard flock(fd, LOCK_EX | LOCK_NB) == 0 else {
            _ = close(fd)
            if errno == EWOULDBLOCK { throw DaemonError.lockUnavailable }
            throw DaemonError.systemCall("flock lock")
        }
    }

    deinit {
        if lockFD >= 0 { _ = close(lockFD) }
    }

    func start() throws {
        // A prior catchable shutdown may not have reached Auto; reset before
        // accepting any lease.
        try backend.apply(.auto)
        try removeStaleSocketAsLockOwner()
        try bindListener()
        installSignalHandlers()
        armIdleExit(startSuspended: true)

        // All sources are constructed suspended. Resume them together on the
        // serial state queue only after every mutable field is initialized.
        queue.sync { [self] in activateInitialSources() }
    }

    private func activateInitialSources() {
        listenerSource?.resume()
        for source in signalSources { source.resume() }
        idleTimer?.resume()
    }

    private func removeStaleSocketAsLockOwner() throws {
        var info = stat()
        if lstat(socketPath, &info) != 0 {
            if errno == ENOENT { return }
            throw DaemonError.systemCall("lstat socket")
        }
        guard (info.st_mode & S_IFMT) == S_IFSOCK else {
            throw DaemonError.unsafeRuntimeDirectory("socket path is not a socket")
        }
        let permissions = info.st_mode & 0o777
        guard info.st_uid == getuid(), permissions == 0o600 || permissions == 0o700 else {
            throw DaemonError.unsafeRuntimeDirectory("socket has wrong owner or permissions")
        }
        guard unlink(socketPath) == 0 else { throw DaemonError.systemCall("unlink stale socket") }
    }

    private func bindListener() throws {
        let capacity = MemoryLayout.size(ofValue: sockaddr_un().sun_path)
        let encodedPath = Array(socketPath.utf8CString)
        guard encodedPath.count <= capacity else { throw DaemonError.socketPathTooLong }

        let fd = socket(AF_UNIX, SOCK_STREAM, 0)
        guard fd >= 0 else { throw DaemonError.systemCall("socket") }
        // chmod after bind leaves a SIGKILL window. Create the socket with its
        // final mode instead, then immediately restore this process's umask.
        let priorUmask = umask(0o177)
        defer { _ = umask(priorUmask) }
        do {
            try setNonBlockingAndCloseOnExec(fd)
            var address = sockaddr_un()
            address.sun_family = sa_family_t(AF_UNIX)
            let offset = MemoryLayout<sockaddr_un>.offset(of: \sockaddr_un.sun_path)!
            withUnsafeMutableBytes(of: &address) { raw in
                encodedPath.withUnsafeBytes { pathBytes in
                    raw.baseAddress!.advanced(by: offset).copyMemory(from: pathBytes.baseAddress!, byteCount: encodedPath.count)
                }
            }
            let addressLength = socklen_t(offset + encodedPath.count)
            let bindResult = withUnsafePointer(to: &address) { pointer in
                pointer.withMemoryRebound(to: sockaddr.self, capacity: 1) {
                    bind(fd, $0, addressLength)
                }
            }
            guard bindResult == 0 else { throw DaemonError.systemCall("bind socket") }
            guard chmod(socketPath, 0o600) == 0 else { throw DaemonError.systemCall("chmod socket") }
            guard listen(fd, SOMAXCONN) == 0 else { throw DaemonError.systemCall("listen") }
        } catch {
            _ = close(fd)
            _ = unlink(socketPath)
            throw error
        }

        listenerFD = fd
        let source = DispatchSource.makeReadSource(fileDescriptor: fd, queue: queue)
        source.setEventHandler { [weak self] in self?.drainAcceptedConnections() }
        source.setCancelHandler { [weak self] in
            _ = close(fd)
            self?.fdDidClose(fd)
        }
        listenerSource = source
    }

    private func installSignalHandlers() {
        for signalNumber in [SIGINT, SIGTERM] {
            signal(signalNumber, SIG_IGN)
            let source = DispatchSource.makeSignalSource(signal: signalNumber, queue: queue)
            source.setEventHandler { [weak self] in self?.beginShutdown() }
            signalSources.append(source)
        }
    }

    private func drainAcceptedConnections() {
        guard !shuttingDown else { return }
        while true {
            let fd = accept(listenerFD, nil, nil)
            if fd >= 0 {
                do {
                    try setNonBlockingAndCloseOnExec(fd)
                    var peerUID: uid_t = 0
                    var peerGID: gid_t = 0
                    guard getpeereid(fd, &peerUID, &peerGID) == 0, peerUID == getuid() else {
                        _ = close(fd)
                        continue
                    }
                    addClient(fd)
                } catch {
                    _ = close(fd)
                }
                continue
            }
            if errno == EINTR { continue }
            if errno == EAGAIN || errno == EWOULDBLOCK { return }
            beginShutdown()
            return
        }
    }

    private func addClient(_ fd: Int32) {
        let hadNoClients = clients.isEmpty
        let source = DispatchSource.makeReadSource(fileDescriptor: fd, queue: queue)
        source.setEventHandler { [weak self] in self?.drainClient(fd) }
        source.setCancelHandler { [weak self] in
            _ = close(fd)
            self?.fdDidClose(fd)
        }
        clients[fd] = source
        source.resume()
        if hadNoClients { startBlinking() }
    }

    private func drainClient(_ fd: Int32) {
        guard clients[fd] != nil else { return }
        var bytes = [UInt8](repeating: 0, count: 1024)
        let byteCount = bytes.count
        while true {
            let result = bytes.withUnsafeMutableBytes { read(fd, $0.baseAddress!, byteCount) }
            if result > 0 { continue }
            if result == 0 {
                removeClient(fd)
                return
            }
            if errno == EINTR { continue }
            if errno == EAGAIN || errno == EWOULDBLOCK { return }
            removeClient(fd)
            return
        }
    }

    private func removeClient(_ fd: Int32) {
        guard let source = clients.removeValue(forKey: fd) else { return }
        source.cancel()
        if clients.isEmpty { stopBlinkingAndArmIdleExit() }
    }

    private func startBlinking() {
        idleTimer?.cancel()
        idleTimer = nil
        cancelBlinkTimer()
        blinkGeneration &+= 1
        let generation = blinkGeneration
        setLED(.on)
        ledIsOn = true

        let timer = DispatchSource.makeTimerSource(queue: queue)
        timer.setEventHandler { [weak self] in self?.blinkTick(generation: generation) }
        blinkTimer = timer
        timer.schedule(deadline: .now() + .milliseconds(250), repeating: .never)
        timer.resume()
    }

    private func cancelBlinkTimer() {
        // Dispatch may already have queued a canceled timer's handler. Changing
        // this token makes that handler inert before a replacement is installed.
        blinkGeneration &+= 1
        blinkTimer?.cancel()
        blinkTimer = nil
    }

    private func blinkTick(generation: UInt64) {
        guard generation == blinkGeneration,
              !shuttingDown,
              !clients.isEmpty,
              let timer = blinkTimer else { return }
        if ledIsOn {
            setLED(.off)
            ledIsOn = false
            timer.schedule(deadline: .now() + .milliseconds(750), repeating: .never)
        } else {
            setLED(.on)
            ledIsOn = true
            timer.schedule(deadline: .now() + .milliseconds(250), repeating: .never)
        }
    }

    private func stopBlinkingAndArmIdleExit() {
        cancelBlinkTimer()
        ledIsOn = false
        setLED(.auto)
        armIdleExit()
    }

    private func armIdleExit(startSuspended: Bool = false) {
        guard !shuttingDown, clients.isEmpty else { return }
        idleTimer?.cancel()
        let timer = DispatchSource.makeTimerSource(queue: queue)
        timer.setEventHandler { [weak self] in
            guard let self, self.clients.isEmpty else { return }
            self.beginShutdown()
        }
        idleTimer = timer
        timer.schedule(deadline: .now() + .seconds(300), repeating: .never)
        if !startSuspended { timer.resume() }
    }

    private func setLED(_ mode: LEDMode) {
        // A failed HID write must not destabilize IPC. A later edge or daemon
        // restart gets another chance to restore the system-managed value.
        try? backend.apply(mode)
    }

    private func beginShutdown() {
        guard !shuttingDown else { return }
        shuttingDown = true

        cancelBlinkTimer()
        idleTimer?.cancel()
        idleTimer = nil
        ledIsOn = false
        setLED(.auto)

        for source in signalSources { source.cancel() }
        signalSources.removeAll()

        if let source = listenerSource {
            pendingFDCloses.insert(listenerFD)
            source.cancel()
            listenerSource = nil
        }
        for (fd, source) in clients {
            pendingFDCloses.insert(fd)
            source.cancel()
        }
        clients.removeAll()
        finishShutdownIfReady()
    }

    private func fdDidClose(_ fd: Int32) {
        guard pendingFDCloses.remove(fd) != nil else { return }
        finishShutdownIfReady()
    }

    private func finishShutdownIfReady() {
        guard shuttingDown, pendingFDCloses.isEmpty else { return }
        // This process owns the flock, so it alone removes the active pathname.
        // Re-validate before unlinking rather than deleting an unexpected file.
        var info = stat()
        if lstat(socketPath, &info) == 0,
           (info.st_mode & S_IFMT) == S_IFSOCK,
           info.st_uid == getuid(),
           (info.st_mode & 0o777) == 0o600 {
            _ = unlink(socketPath)
        }
        _ = flock(lockFD, LOCK_UN)
        _ = close(lockFD)
        exit(0)
    }
}

private func run() throws {
    guard ProcessInfo.processInfo.operatingSystemVersion.majorVersion >= 26 else {
        throw DaemonError.unsupportedOS
    }
    _ = umask(0o077)
    let arguments = try Arguments.parse(CommandLine.arguments)
    let runtimeDirectory = try prepareRuntimeDirectory(passedPath: arguments.runtimeDirectory)
    let backend: LEDBackend
    if let fakeLEDFile = arguments.fakeLEDFile {
        backend = FakeLEDBackend(path: fakeLEDFile)
    } else {
        backend = try HIDLEDBackend()
    }
    let daemon = try Daemon(runtimeDirectory: runtimeDirectory, backend: backend)
    try daemon.start()
    dispatchMain()
}

do {
    try run()
} catch DaemonError.lockUnavailable {
    // Spawn races are expected. The process holding the lock is authoritative.
    exit(0)
} catch {
    fputs("CapsBlinkDaemon: \(error)\n", stderr)
    exit(1)
}
