chore: adiciona .gitignore e commit.command
This commit is contained in:
+611
@@ -0,0 +1,611 @@
|
||||
import AppKit
|
||||
import ApplicationServices
|
||||
import Carbon
|
||||
import ComposableArchitecture
|
||||
import CoreGraphics
|
||||
import Dependencies
|
||||
import DependenciesMacros
|
||||
import Foundation
|
||||
import HexCore
|
||||
import IOKit
|
||||
import IOKit.hidsystem
|
||||
import Sauce
|
||||
|
||||
private let logger = HexLog.keyEvent
|
||||
|
||||
struct KeyEventMonitorToken: Sendable {
|
||||
private let cancelHandler: @Sendable () -> Void
|
||||
|
||||
init(cancel: @escaping @Sendable () -> Void) {
|
||||
self.cancelHandler = cancel
|
||||
}
|
||||
|
||||
func cancel() {
|
||||
cancelHandler()
|
||||
}
|
||||
|
||||
static let noop = KeyEventMonitorToken(cancel: {})
|
||||
}
|
||||
|
||||
public extension KeyEvent {
|
||||
init(cgEvent: CGEvent, type: CGEventType, isFnPressed: Bool) {
|
||||
let keyCode = Int(cgEvent.getIntegerValueField(.keyboardEventKeycode))
|
||||
// Accessing keyboard layout / input source via Sauce must be on main thread.
|
||||
let key: Key?
|
||||
if cgEvent.type == .keyDown {
|
||||
if Thread.isMainThread {
|
||||
key = Sauce.shared.key(for: keyCode)
|
||||
} else {
|
||||
key = DispatchQueue.main.sync { Sauce.shared.key(for: keyCode) }
|
||||
}
|
||||
} else {
|
||||
key = nil
|
||||
}
|
||||
|
||||
var modifiers = Modifiers.from(carbonFlags: cgEvent.flags)
|
||||
if !isFnPressed {
|
||||
modifiers = modifiers.removing(kind: .fn)
|
||||
}
|
||||
self.init(key: key, modifiers: modifiers)
|
||||
}
|
||||
}
|
||||
|
||||
@DependencyClient
|
||||
struct KeyEventMonitorClient {
|
||||
var listenForKeyPress: @Sendable () async -> AsyncThrowingStream<KeyEvent, Error> = {
|
||||
AsyncThrowingStream { _ in }
|
||||
}
|
||||
var handleKeyEvent: @Sendable (@Sendable @escaping (KeyEvent) -> Bool) -> KeyEventMonitorToken = { _ in .noop }
|
||||
var handleInputEvent: @Sendable (@Sendable @escaping (InputEvent) -> Bool) -> KeyEventMonitorToken = { _ in .noop }
|
||||
var startMonitoring: @Sendable () async -> Void = {}
|
||||
var stopMonitoring: @Sendable () -> Void = {}
|
||||
}
|
||||
|
||||
extension KeyEventMonitorClient: DependencyKey {
|
||||
static var liveValue: KeyEventMonitorClient {
|
||||
let live = KeyEventMonitorClientLive()
|
||||
return KeyEventMonitorClient(
|
||||
listenForKeyPress: {
|
||||
live.listenForKeyPress()
|
||||
},
|
||||
handleKeyEvent: { handler in
|
||||
live.handleKeyEvent(handler)
|
||||
},
|
||||
handleInputEvent: { handler in
|
||||
live.handleInputEvent(handler)
|
||||
},
|
||||
startMonitoring: {
|
||||
live.startMonitoring()
|
||||
},
|
||||
stopMonitoring: {
|
||||
live.stopMonitoring()
|
||||
}
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
extension DependencyValues {
|
||||
var keyEventMonitor: KeyEventMonitorClient {
|
||||
get { self[KeyEventMonitorClient.self] }
|
||||
set { self[KeyEventMonitorClient.self] = newValue }
|
||||
}
|
||||
}
|
||||
|
||||
class KeyEventMonitorClientLive {
|
||||
private var eventTapPort: CFMachPort?
|
||||
private var runLoopSource: CFRunLoopSource?
|
||||
private var continuations: [UUID: @Sendable (KeyEvent) -> Bool] = [:]
|
||||
private var inputContinuations: [UUID: @Sendable (InputEvent) -> Bool] = [:]
|
||||
private let queue = DispatchQueue(label: "com.kitlangton.Hex.KeyEventMonitor", attributes: .concurrent)
|
||||
private let queueSpecificKey = DispatchSpecificKey<Void>()
|
||||
private var isMonitoring = false
|
||||
private var wantsMonitoring = false
|
||||
private var accessibilityTrusted = false
|
||||
private var inputMonitoringTrusted = false
|
||||
/// Set when key events are observed arriving at the tap. Key events only flow when Input
|
||||
/// Monitoring is genuinely granted, so this overrides stale `IOHIDCheckAccess` denials (#250).
|
||||
private var inputMonitoringProvenByEvents = false
|
||||
private var trustMonitorTask: Task<Void, Never>?
|
||||
private var systemEventObservers: [NSObjectProtocol] = []
|
||||
private var isFnPressed = false
|
||||
private var hasPromptedForAccessibilityTrust = false
|
||||
private let accessibilityTrustProvider: @Sendable () -> Bool
|
||||
private let accessibilityTrustPrompt: @Sendable () -> Bool
|
||||
private let inputMonitoringTrustProvider: @Sendable () -> Bool
|
||||
@Shared(.hotkeyPermissionState) private var hotkeyPermissionState: HotkeyPermissionState
|
||||
|
||||
private let trustCheckIntervalNanoseconds: UInt64 = 100_000_000 // 100ms
|
||||
|
||||
init(
|
||||
accessibilityTrustProvider: @escaping @Sendable () -> Bool = {
|
||||
let promptKey = kAXTrustedCheckOptionPrompt.takeUnretainedValue() as String
|
||||
return AXIsProcessTrustedWithOptions([promptKey: false] as CFDictionary)
|
||||
},
|
||||
accessibilityTrustPrompt: @escaping @Sendable () -> Bool = {
|
||||
let promptKey = kAXTrustedCheckOptionPrompt.takeUnretainedValue() as String
|
||||
return AXIsProcessTrustedWithOptions([promptKey: true] as CFDictionary)
|
||||
},
|
||||
inputMonitoringTrustProvider: @escaping @Sendable () -> Bool = {
|
||||
IOHIDCheckAccess(kIOHIDRequestTypeListenEvent) == kIOHIDAccessTypeGranted
|
||||
}
|
||||
) {
|
||||
self.accessibilityTrustProvider = accessibilityTrustProvider
|
||||
self.accessibilityTrustPrompt = accessibilityTrustPrompt
|
||||
self.inputMonitoringTrustProvider = inputMonitoringTrustProvider
|
||||
queue.setSpecific(key: queueSpecificKey, value: ())
|
||||
logger.info("Initializing HotKeyClient with CGEvent tap.")
|
||||
registerSystemEventObservers()
|
||||
}
|
||||
|
||||
deinit {
|
||||
let center = NSWorkspace.shared.notificationCenter
|
||||
for observer in systemEventObservers {
|
||||
center.removeObserver(observer)
|
||||
}
|
||||
self.stopMonitoring()
|
||||
}
|
||||
|
||||
private var hasHandlers: Bool {
|
||||
readState { !(continuations.isEmpty && inputContinuations.isEmpty) }
|
||||
}
|
||||
|
||||
func readState<Value>(_ operation: () -> Value) -> Value {
|
||||
// Handler registration performs permission checks from a barrier block on this queue.
|
||||
if DispatchQueue.getSpecific(key: queueSpecificKey) != nil {
|
||||
return operation()
|
||||
}
|
||||
return queue.sync(execute: operation)
|
||||
}
|
||||
|
||||
private func setMonitoringIntent(_ value: Bool) {
|
||||
queue.async(flags: .barrier) { [weak self] in
|
||||
self?.wantsMonitoring = value
|
||||
}
|
||||
}
|
||||
|
||||
private func desiredMonitoringState() -> Bool {
|
||||
// Intentionally not gated on `inputMonitoringTrusted`: `IOHIDCheckAccess` is notorious for
|
||||
// returning stale denials (after sleep, MDM re-logins, or OS updates) while events still
|
||||
// flow, and tearing the tap down on that signal killed working hotkeys (#250). The tap only
|
||||
// needs Accessibility to exist; without Input Monitoring, macOS simply withholds key events
|
||||
// (modifiers still arrive), and creating the tap is what triggers the permission prompt.
|
||||
readState {
|
||||
wantsMonitoring
|
||||
&& accessibilityTrusted
|
||||
&& !(continuations.isEmpty && inputContinuations.isEmpty)
|
||||
}
|
||||
}
|
||||
|
||||
/// Provide a stream of key events.
|
||||
func listenForKeyPress() -> AsyncThrowingStream<KeyEvent, Error> {
|
||||
AsyncThrowingStream { continuation in
|
||||
let uuid = UUID()
|
||||
|
||||
queue.async(flags: .barrier) { [weak self] in
|
||||
guard let self = self else { return }
|
||||
self.continuations[uuid] = { event in
|
||||
continuation.yield(event)
|
||||
return false
|
||||
}
|
||||
let shouldStart = self.continuations.count == 1 && self.inputContinuations.isEmpty
|
||||
|
||||
// Start monitoring if this is the first subscription
|
||||
if shouldStart {
|
||||
self.startMonitoring()
|
||||
}
|
||||
}
|
||||
|
||||
// Cleanup on cancellation
|
||||
continuation.onTermination = { [weak self] _ in
|
||||
self?.removeHandlerContinuation(uuid: uuid)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
private func removeHandlerContinuation(uuid: UUID) {
|
||||
queue.async(flags: .barrier) { [weak self] in
|
||||
guard let self = self else { return }
|
||||
self.continuations[uuid] = nil
|
||||
if self.continuations.isEmpty && self.inputContinuations.isEmpty {
|
||||
self.stopMonitoring()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
private func removeInputContinuation(uuid: UUID) {
|
||||
queue.async(flags: .barrier) { [weak self] in
|
||||
guard let self = self else { return }
|
||||
self.inputContinuations[uuid] = nil
|
||||
if self.continuations.isEmpty && self.inputContinuations.isEmpty {
|
||||
self.stopMonitoring()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func startMonitoring() {
|
||||
setMonitoringIntent(true)
|
||||
startTrustMonitorIfNeeded()
|
||||
refreshTrustedFlag(promptIfUntrusted: true)
|
||||
Task { [weak self] in
|
||||
await self?.refreshMonitoringState(reason: "startMonitoring")
|
||||
}
|
||||
}
|
||||
// TODO: Handle removing the handler from the continuations on deinit/cancellation
|
||||
func handleKeyEvent(_ handler: @Sendable @escaping (KeyEvent) -> Bool) -> KeyEventMonitorToken {
|
||||
let uuid = UUID()
|
||||
|
||||
queue.async(flags: .barrier) { [weak self] in
|
||||
guard let self = self else { return }
|
||||
self.continuations[uuid] = handler
|
||||
let shouldStart = self.continuations.count == 1 && self.inputContinuations.isEmpty
|
||||
|
||||
if shouldStart {
|
||||
self.startMonitoring()
|
||||
}
|
||||
}
|
||||
|
||||
return KeyEventMonitorToken { [weak self] in
|
||||
self?.removeHandlerContinuation(uuid: uuid)
|
||||
}
|
||||
}
|
||||
|
||||
func handleInputEvent(_ handler: @Sendable @escaping (InputEvent) -> Bool) -> KeyEventMonitorToken {
|
||||
let uuid = UUID()
|
||||
|
||||
queue.async(flags: .barrier) { [weak self] in
|
||||
guard let self = self else { return }
|
||||
self.inputContinuations[uuid] = handler
|
||||
let shouldStart = self.inputContinuations.count == 1 && self.continuations.isEmpty
|
||||
|
||||
if shouldStart {
|
||||
self.startMonitoring()
|
||||
}
|
||||
}
|
||||
|
||||
return KeyEventMonitorToken { [weak self] in
|
||||
self?.removeInputContinuation(uuid: uuid)
|
||||
}
|
||||
}
|
||||
|
||||
func stopMonitoring() {
|
||||
setMonitoringIntent(false)
|
||||
Task { [weak self] in
|
||||
await self?.refreshMonitoringState(reason: "stopMonitoring")
|
||||
}
|
||||
cancelTrustMonitorIfNeeded()
|
||||
}
|
||||
|
||||
private func startTrustMonitorIfNeeded() {
|
||||
queue.async(flags: .barrier) { [weak self] in
|
||||
guard let self else { return }
|
||||
guard self.trustMonitorTask == nil else { return }
|
||||
self.trustMonitorTask = Task { [weak self] in
|
||||
await self?.watchPermissions()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
private func cancelTrustMonitorIfNeeded() {
|
||||
queue.async(flags: .barrier) { [weak self] in
|
||||
guard let self else { return }
|
||||
guard !self.wantsMonitoring else { return }
|
||||
self.trustMonitorTask?.cancel()
|
||||
self.trustMonitorTask = nil
|
||||
}
|
||||
}
|
||||
|
||||
// no separate helper; handled inline above
|
||||
|
||||
private func watchPermissions() async {
|
||||
var last = (
|
||||
accessibility: currentAccessibilityTrust(),
|
||||
input: currentInputMonitoringTrust()
|
||||
)
|
||||
await handlePermissionChange(accessibility: last.accessibility, input: last.input, reason: "initial")
|
||||
|
||||
while !Task.isCancelled {
|
||||
try? await Task.sleep(nanoseconds: trustCheckIntervalNanoseconds)
|
||||
let current = (
|
||||
accessibility: currentAccessibilityTrust(),
|
||||
input: currentInputMonitoringTrust()
|
||||
)
|
||||
|
||||
if current.accessibility != last.accessibility || current.input != last.input {
|
||||
let combinedBefore = last.accessibility && last.input
|
||||
let combinedAfter = current.accessibility && current.input
|
||||
let reason: String
|
||||
if combinedAfter && !combinedBefore {
|
||||
reason = "regained"
|
||||
} else if !combinedAfter && combinedBefore {
|
||||
reason = "revoked"
|
||||
} else {
|
||||
reason = "updated"
|
||||
}
|
||||
await handlePermissionChange(accessibility: current.accessibility, input: current.input, reason: reason)
|
||||
last = current
|
||||
} else if current.accessibility {
|
||||
await ensureTapIsRunning()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
private func handlePermissionChange(accessibility: Bool, input: Bool, reason: String) async {
|
||||
setPermissionFlags(accessibility: accessibility, input: input)
|
||||
logger.notice("Permission update: accessibility=\(accessibility), inputMonitoring=\(input), reason=\(reason)")
|
||||
if accessibility && input {
|
||||
logger.notice("Keyboard monitoring permissions granted (\(reason)).")
|
||||
} else {
|
||||
if !accessibility {
|
||||
logger.error("Accessibility permission missing (\(reason)); suspending tap.")
|
||||
}
|
||||
if !input {
|
||||
logger.error("Input Monitoring permission missing (\(reason)); keyed hotkeys may not fire until it is granted. Tap stays alive in case the check is stale.")
|
||||
}
|
||||
}
|
||||
await refreshMonitoringState(reason: "trust_\(reason)")
|
||||
}
|
||||
|
||||
private func ensureTapIsRunning() async {
|
||||
guard desiredMonitoringState() else { return }
|
||||
await activateTapOnMain(reason: "watchdog_keepalive")
|
||||
}
|
||||
|
||||
private func refreshMonitoringState(reason: String) async {
|
||||
let shouldMonitor = desiredMonitoringState()
|
||||
if shouldMonitor {
|
||||
await activateTapOnMain(reason: reason)
|
||||
} else {
|
||||
await deactivateTapOnMain(reason: reason)
|
||||
}
|
||||
}
|
||||
|
||||
private func setPermissionFlags(accessibility: Bool, input: Bool) {
|
||||
queue.async(flags: .barrier) { [weak self] in
|
||||
self?.accessibilityTrusted = accessibility
|
||||
self?.inputMonitoringTrusted = input
|
||||
}
|
||||
recordSharedPermissionState(accessibility: accessibility, input: input)
|
||||
}
|
||||
|
||||
private func recordSharedPermissionState(accessibility: Bool, input: Bool) {
|
||||
$hotkeyPermissionState.withLock {
|
||||
$0.accessibility = accessibility ? .granted : .denied
|
||||
$0.inputMonitoring = input ? .granted : .denied
|
||||
$0.lastUpdated = Date()
|
||||
}
|
||||
}
|
||||
|
||||
private func activateTapOnMain(reason: String) async {
|
||||
await MainActor.run {
|
||||
self.activateTapIfNeeded(reason: reason)
|
||||
}
|
||||
}
|
||||
|
||||
private func deactivateTapOnMain(reason: String) async {
|
||||
await MainActor.run {
|
||||
self.deactivateTap(reason: reason)
|
||||
}
|
||||
}
|
||||
|
||||
@MainActor
|
||||
private func activateTapIfNeeded(reason: String) {
|
||||
if isMonitoring {
|
||||
// The 100ms permission watchdog lands here while healthy; use it to revive taps that
|
||||
// macOS disabled without sending a tapDisabled event (observed after sleep, #250).
|
||||
if let eventTapPort, !CGEvent.tapIsEnabled(tap: eventTapPort) {
|
||||
CGEvent.tapEnable(tap: eventTapPort, enable: true)
|
||||
logger.notice("Re-enabled event tap that was silently disabled (reason: \(reason)).")
|
||||
}
|
||||
return
|
||||
}
|
||||
guard hasHandlers else { return }
|
||||
|
||||
let accessibilityTrusted = currentAccessibilityTrust()
|
||||
let inputMonitoringTrusted = currentInputMonitoringTrust()
|
||||
setPermissionFlags(accessibility: accessibilityTrusted, input: inputMonitoringTrusted)
|
||||
guard accessibilityTrusted else {
|
||||
logger.error("Cannot start key event monitoring (reason: \(reason)); accessibility permission is not granted.")
|
||||
return
|
||||
}
|
||||
|
||||
if !inputMonitoringTrusted {
|
||||
logger.notice("Input Monitoring not yet granted; creating event tap will trigger permission prompt (reason: \(reason)).")
|
||||
}
|
||||
|
||||
let eventMask =
|
||||
((1 << CGEventType.keyDown.rawValue)
|
||||
| (1 << CGEventType.keyUp.rawValue)
|
||||
| (1 << CGEventType.flagsChanged.rawValue)
|
||||
| (1 << CGEventType.leftMouseDown.rawValue)
|
||||
| (1 << CGEventType.rightMouseDown.rawValue)
|
||||
| (1 << CGEventType.otherMouseDown.rawValue))
|
||||
|
||||
guard
|
||||
let eventTap = CGEvent.tapCreate(
|
||||
tap: .cghidEventTap,
|
||||
place: .headInsertEventTap,
|
||||
options: .defaultTap,
|
||||
eventsOfInterest: CGEventMask(eventMask),
|
||||
callback: { _, type, cgEvent, userInfo in
|
||||
guard
|
||||
let hotKeyClientLive = Unmanaged<KeyEventMonitorClientLive>
|
||||
.fromOpaque(userInfo!)
|
||||
.takeUnretainedValue() as KeyEventMonitorClientLive?
|
||||
else {
|
||||
return Unmanaged.passUnretained(cgEvent)
|
||||
}
|
||||
|
||||
if type == .tapDisabledByUserInput || type == .tapDisabledByTimeout {
|
||||
hotKeyClientLive.handleTapDisabledEvent(type)
|
||||
return Unmanaged.passUnretained(cgEvent)
|
||||
}
|
||||
|
||||
// An event arriving at the tap is authoritative proof the underlying permission is
|
||||
// granted. Never drop delivered events because a cached permission check went
|
||||
// stale (#250) — that turned recoverable TCC hiccups into dead hotkeys.
|
||||
hotKeyClientLive.noteEventDelivered(type: type)
|
||||
|
||||
if type == .leftMouseDown || type == .rightMouseDown || type == .otherMouseDown {
|
||||
_ = hotKeyClientLive.processInputEvent(.mouseClick)
|
||||
return Unmanaged.passUnretained(cgEvent)
|
||||
}
|
||||
|
||||
hotKeyClientLive.updateFnStateIfNeeded(type: type, cgEvent: cgEvent)
|
||||
|
||||
let keyEvent = KeyEvent(cgEvent: cgEvent, type: type, isFnPressed: hotKeyClientLive.isFnPressed)
|
||||
let handledByKeyHandler = hotKeyClientLive.processKeyEvent(keyEvent)
|
||||
let handledByInputHandler = hotKeyClientLive.processInputEvent(.keyboard(keyEvent))
|
||||
|
||||
return (handledByKeyHandler || handledByInputHandler) ? nil : Unmanaged.passUnretained(cgEvent)
|
||||
},
|
||||
userInfo: UnsafeMutableRawPointer(Unmanaged.passUnretained(self).toOpaque())
|
||||
)
|
||||
else {
|
||||
logger.error("Failed to create event tap (reason: \(reason)).")
|
||||
return
|
||||
}
|
||||
|
||||
eventTapPort = eventTap
|
||||
|
||||
let runLoopSource = CFMachPortCreateRunLoopSource(kCFAllocatorDefault, eventTap, 0)
|
||||
self.runLoopSource = runLoopSource
|
||||
|
||||
CFRunLoopAddSource(CFRunLoopGetMain(), runLoopSource, .commonModes)
|
||||
CGEvent.tapEnable(tap: eventTap, enable: true)
|
||||
isMonitoring = true
|
||||
logger.info("Started monitoring key events via CGEvent tap (reason: \(reason)).")
|
||||
}
|
||||
|
||||
@MainActor
|
||||
private func deactivateTap(reason: String) {
|
||||
guard isMonitoring || eventTapPort != nil else { return }
|
||||
|
||||
if let runLoopSource = runLoopSource {
|
||||
CFRunLoopRemoveSource(CFRunLoopGetMain(), runLoopSource, .commonModes)
|
||||
self.runLoopSource = nil
|
||||
}
|
||||
|
||||
if let eventTapPort = eventTapPort {
|
||||
CGEvent.tapEnable(tap: eventTapPort, enable: false)
|
||||
self.eventTapPort = nil
|
||||
}
|
||||
|
||||
isMonitoring = false
|
||||
clearInputMonitoringProof()
|
||||
logger.info("Suspended key event monitoring (reason: \(reason)).")
|
||||
}
|
||||
|
||||
private func clearInputMonitoringProof() {
|
||||
queue.async(flags: .barrier) { [weak self] in
|
||||
self?.inputMonitoringProvenByEvents = false
|
||||
}
|
||||
}
|
||||
|
||||
private func handleTapDisabledEvent(_ type: CGEventType) {
|
||||
let reason = type == .tapDisabledByTimeout ? "timeout" : "userInput"
|
||||
logger.error("Event tap disabled by \(reason); scheduling restart.")
|
||||
Task { [weak self] in
|
||||
guard let self else { return }
|
||||
await self.refreshMonitoringState(reason: "tap_disabled_\(reason)")
|
||||
}
|
||||
}
|
||||
|
||||
private func processEvent<T>(
|
||||
_ event: T,
|
||||
handlers: () -> [UUID: @Sendable (T) -> Bool]
|
||||
) -> Bool {
|
||||
let handlerList = readState { Array(handlers().values) }
|
||||
return handlerList.reduce(false) { handled, handler in
|
||||
handler(event) || handled
|
||||
}
|
||||
}
|
||||
|
||||
private func processKeyEvent(_ keyEvent: KeyEvent) -> Bool {
|
||||
processEvent(keyEvent, handlers: { continuations })
|
||||
}
|
||||
|
||||
private func processInputEvent(_ inputEvent: InputEvent) -> Bool {
|
||||
processEvent(inputEvent, handlers: { inputContinuations })
|
||||
}
|
||||
|
||||
/// Records that an event was delivered to the tap. Key events are proof that Input
|
||||
/// Monitoring is genuinely granted even when `IOHIDCheckAccess` reports otherwise (#250).
|
||||
fileprivate func noteEventDelivered(type: CGEventType) {
|
||||
guard type == .keyDown || type == .keyUp else { return }
|
||||
let alreadyProven = readState { inputMonitoringProvenByEvents }
|
||||
guard !alreadyProven else { return }
|
||||
queue.async(flags: .barrier) { [weak self] in
|
||||
self?.inputMonitoringProvenByEvents = true
|
||||
self?.inputMonitoringTrusted = true
|
||||
}
|
||||
if IOHIDCheckAccess(kIOHIDRequestTypeListenEvent) != kIOHIDAccessTypeGranted {
|
||||
logger.notice("Key events are flowing while IOHIDCheckAccess reports denied; treating Input Monitoring as granted (stale TCC cache, #250).")
|
||||
}
|
||||
}
|
||||
|
||||
/// Recreate the event tap after system transitions that are known to leave taps in a dead or
|
||||
/// stale state (wake from sleep, session reactivation after fast user switching / MDM logout).
|
||||
private func registerSystemEventObservers() {
|
||||
let center = NSWorkspace.shared.notificationCenter
|
||||
let events: [(Notification.Name, String)] = [
|
||||
(NSWorkspace.didWakeNotification, "system_wake"),
|
||||
(NSWorkspace.sessionDidBecomeActiveNotification, "session_active"),
|
||||
]
|
||||
for (name, reason) in events {
|
||||
let observer = center.addObserver(forName: name, object: nil, queue: .main) { [weak self] _ in
|
||||
self?.restartTapAfterSystemEvent(reason: reason)
|
||||
}
|
||||
systemEventObservers.append(observer)
|
||||
}
|
||||
}
|
||||
|
||||
private func restartTapAfterSystemEvent(reason: String) {
|
||||
logger.notice("System transition (\(reason)); recreating event tap to recover from any stale state.")
|
||||
Task { [weak self] in
|
||||
guard let self else { return }
|
||||
await self.deactivateTapOnMain(reason: reason)
|
||||
await self.refreshMonitoringState(reason: reason)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
extension KeyEventMonitorClientLive {
|
||||
private func updateFnStateIfNeeded(type: CGEventType, cgEvent: CGEvent) {
|
||||
guard type == .flagsChanged else { return }
|
||||
let keyCode = Int(cgEvent.getIntegerValueField(.keyboardEventKeycode))
|
||||
guard keyCode == kVK_Function else { return }
|
||||
isFnPressed = cgEvent.flags.contains(.maskSecondaryFn)
|
||||
}
|
||||
|
||||
private func refreshTrustedFlag(promptIfUntrusted: Bool) {
|
||||
var accessibilityTrusted = currentAccessibilityTrust()
|
||||
if !accessibilityTrusted && promptIfUntrusted && !hasPromptedForAccessibilityTrust {
|
||||
accessibilityTrusted = requestAccessibilityTrustPrompt()
|
||||
hasPromptedForAccessibilityTrust = true
|
||||
logger.notice("Prompted for accessibility trust")
|
||||
}
|
||||
|
||||
let inputMonitoringTrusted = currentInputMonitoringTrust()
|
||||
setPermissionFlags(accessibility: accessibilityTrusted, input: inputMonitoringTrusted)
|
||||
}
|
||||
|
||||
private func currentAccessibilityTrust() -> Bool {
|
||||
accessibilityTrustProvider()
|
||||
}
|
||||
|
||||
private func requestAccessibilityTrustPrompt() -> Bool {
|
||||
accessibilityTrustPrompt()
|
||||
}
|
||||
|
||||
private func currentInputMonitoringTrust() -> Bool {
|
||||
if inputMonitoringTrustProvider() {
|
||||
return true
|
||||
}
|
||||
// A stale TCC cache can report denied while key events demonstrably flow (#250);
|
||||
// trust the events over the check so the watchdog and settings UI stay honest.
|
||||
return readState { inputMonitoringProvenByEvents }
|
||||
}
|
||||
|
||||
// Intentionally no request helper: creating the event tap prompts macOS 15+ for Input Monitoring
|
||||
// the same way older versions did, while we still track status for UI.
|
||||
}
|
||||
@@ -0,0 +1,49 @@
|
||||
import Foundation
|
||||
|
||||
enum LegacyModelCacheMigrator {
|
||||
static func migrate(
|
||||
from legacyDirectory: URL,
|
||||
to currentDirectory: URL,
|
||||
fileManager: FileManager = .default,
|
||||
isValid: (URL) -> Bool
|
||||
) throws -> Bool {
|
||||
guard legacyDirectory.standardizedFileURL != currentDirectory.standardizedFileURL,
|
||||
fileManager.fileExists(atPath: legacyDirectory.path),
|
||||
!isValid(currentDirectory)
|
||||
else { return false }
|
||||
|
||||
let backupDirectory = currentDirectory
|
||||
.deletingLastPathComponent()
|
||||
.appendingPathComponent(".\(currentDirectory.lastPathComponent)-migration-\(UUID().uuidString)")
|
||||
let hadCurrentDirectory = fileManager.fileExists(atPath: currentDirectory.path)
|
||||
|
||||
if hadCurrentDirectory {
|
||||
try fileManager.moveItem(at: currentDirectory, to: backupDirectory)
|
||||
}
|
||||
|
||||
do {
|
||||
try fileManager.moveItem(at: legacyDirectory, to: currentDirectory)
|
||||
guard isValid(currentDirectory) else {
|
||||
throw MigrationError.invalidLegacyModel
|
||||
}
|
||||
if hadCurrentDirectory {
|
||||
try? fileManager.removeItem(at: backupDirectory)
|
||||
}
|
||||
return true
|
||||
} catch {
|
||||
if fileManager.fileExists(atPath: currentDirectory.path),
|
||||
!fileManager.fileExists(atPath: legacyDirectory.path)
|
||||
{
|
||||
try? fileManager.moveItem(at: currentDirectory, to: legacyDirectory)
|
||||
}
|
||||
if hadCurrentDirectory, fileManager.fileExists(atPath: backupDirectory.path) {
|
||||
try? fileManager.moveItem(at: backupDirectory, to: currentDirectory)
|
||||
}
|
||||
throw error
|
||||
}
|
||||
}
|
||||
|
||||
enum MigrationError: Error {
|
||||
case invalidLegacyModel
|
||||
}
|
||||
}
|
||||
+214
@@ -0,0 +1,214 @@
|
||||
import Foundation
|
||||
import HexCore
|
||||
|
||||
#if canImport(FluidAudio)
|
||||
import FluidAudio
|
||||
|
||||
actor ParakeetClient {
|
||||
private var asr: AsrManager?
|
||||
private var models: AsrModels?
|
||||
private var currentVariant: ParakeetModel?
|
||||
private let logger = HexLog.parakeet
|
||||
private let vendorDirs = [
|
||||
// Our app-specific cache path convention (under XDG or com.kitlangton.Hex/cache)
|
||||
"fluidaudio/Models",
|
||||
"FluidAudio/Models"
|
||||
]
|
||||
|
||||
func isModelAvailable(_ modelName: String) async -> Bool {
|
||||
guard let variant = ParakeetModel(rawValue: modelName) else {
|
||||
logger.error("Unknown Parakeet variant requested: \(modelName)")
|
||||
return false
|
||||
}
|
||||
if currentVariant == variant, asr != nil { return true }
|
||||
|
||||
let directory = AsrModels.defaultCacheDirectory(for: variant.asrVersion)
|
||||
migrateLegacyCacheIfNeeded(variant, to: directory)
|
||||
let available = AsrModels.modelsExist(
|
||||
at: directory,
|
||||
version: variant.asrVersion
|
||||
)
|
||||
if available {
|
||||
logger.notice("Found Parakeet cache at \(directory.path)")
|
||||
} else {
|
||||
logger.debug("No Parakeet cache detected variant=\(variant.identifier) path=\(directory.path)")
|
||||
}
|
||||
return available
|
||||
}
|
||||
|
||||
func ensureLoaded(modelName: String, progress: @escaping (Progress) -> Void) async throws {
|
||||
guard let variant = ParakeetModel(rawValue: modelName) else {
|
||||
throw NSError(
|
||||
domain: "Parakeet",
|
||||
code: -4,
|
||||
userInfo: [NSLocalizedDescriptionKey: "Unsupported Parakeet variant: \(modelName)"]
|
||||
)
|
||||
}
|
||||
if currentVariant == variant, asr != nil { return }
|
||||
if currentVariant != variant {
|
||||
asr = nil
|
||||
models = nil
|
||||
}
|
||||
migrateLegacyCacheIfNeeded(
|
||||
variant,
|
||||
to: AsrModels.defaultCacheDirectory(for: variant.asrVersion)
|
||||
)
|
||||
let t0 = Date()
|
||||
logger.notice("Starting Parakeet load variant=\(variant.identifier)")
|
||||
let p = Progress(totalUnitCount: 100)
|
||||
p.completedUnitCount = 1
|
||||
progress(p)
|
||||
|
||||
// Best-effort progress polling while FluidAudio downloads
|
||||
let fm = FileManager.default
|
||||
let support = try? fm.url(for: .applicationSupportDirectory, in: .userDomainMask, appropriateFor: nil, create: true)
|
||||
let faDir = support?.appendingPathComponent("FluidAudio/Models/\(variant.identifier)", isDirectory: true)
|
||||
let pollTask = Task {
|
||||
while p.completedUnitCount < 95 {
|
||||
try? await Task.sleep(nanoseconds: 250_000_000)
|
||||
if let dir = faDir, let size = directorySize(dir) {
|
||||
let target: Double = 650 * 1024 * 1024 // ~650MB
|
||||
let frac = max(0.0, min(1.0, Double(size) / target))
|
||||
p.completedUnitCount = Int64(5 + frac * 90)
|
||||
progress(p)
|
||||
}
|
||||
if Task.isCancelled { break }
|
||||
}
|
||||
}
|
||||
defer { pollTask.cancel() }
|
||||
|
||||
// Download + load the requested variant (returns when all assets are present)
|
||||
let models = try await AsrModels.downloadAndLoad(version: variant.asrVersion)
|
||||
self.models = models
|
||||
let manager = AsrManager(config: .init(), models: models)
|
||||
self.asr = manager
|
||||
self.currentVariant = variant
|
||||
p.completedUnitCount = 100
|
||||
progress(p)
|
||||
logger.notice("Parakeet ensureLoaded completed in \(String(format: "%.2f", Date().timeIntervalSince(t0)))s")
|
||||
}
|
||||
|
||||
private func directorySize(_ dir: URL) -> UInt64? {
|
||||
let fm = FileManager.default
|
||||
guard let en = fm.enumerator(at: dir, includingPropertiesForKeys: [.isRegularFileKey, .fileSizeKey], options: .skipsHiddenFiles) else { return nil }
|
||||
var total: UInt64 = 0
|
||||
for case let url as URL in en {
|
||||
if let vals = try? url.resourceValues(forKeys: [.isRegularFileKey, .fileSizeKey]), vals.isRegularFile == true {
|
||||
total &+= UInt64(vals.fileSize ?? 0)
|
||||
}
|
||||
}
|
||||
return total
|
||||
}
|
||||
|
||||
private func migrateLegacyCacheIfNeeded(_ variant: ParakeetModel, to directory: URL) {
|
||||
let legacyDirectory = directory
|
||||
.deletingLastPathComponent()
|
||||
.appendingPathComponent(variant.identifier, isDirectory: true)
|
||||
|
||||
do {
|
||||
if try LegacyModelCacheMigrator.migrate(
|
||||
from: legacyDirectory,
|
||||
to: directory,
|
||||
isValid: {
|
||||
AsrModels.modelsExist(at: $0, version: variant.asrVersion)
|
||||
}
|
||||
) {
|
||||
logger.notice("Migrated legacy Parakeet cache from \(legacyDirectory.path) to \(directory.path)")
|
||||
}
|
||||
} catch {
|
||||
logger.error("Failed to migrate legacy Parakeet cache: \(error.localizedDescription)")
|
||||
}
|
||||
}
|
||||
|
||||
func transcribe(_ url: URL) async throws -> String {
|
||||
guard let asr else { throw NSError(domain: "Parakeet", code: -1, userInfo: [NSLocalizedDescriptionKey: "Parakeet not initialized"]) }
|
||||
let t0 = Date()
|
||||
logger.notice("Transcribing with Parakeet file=\(url.lastPathComponent)")
|
||||
var decoderState = TdtDecoderState.make(decoderLayers: await asr.decoderLayerCount)
|
||||
let result = try await asr.transcribe(url, decoderState: &decoderState)
|
||||
logger.info("Parakeet transcription finished in \(String(format: "%.2f", Date().timeIntervalSince(t0)))s")
|
||||
return result.text
|
||||
}
|
||||
|
||||
// Delete cached Parakeet models from known locations and reset state
|
||||
func deleteCaches(modelName: String) async throws {
|
||||
guard let variant = ParakeetModel(rawValue: modelName) else { return }
|
||||
let fm = FileManager.default
|
||||
|
||||
var removedAny = false
|
||||
for dir in modelDirectories(variant) {
|
||||
if fm.fileExists(atPath: dir.path) {
|
||||
try? fm.removeItem(at: dir)
|
||||
removedAny = true
|
||||
}
|
||||
}
|
||||
|
||||
// Reset live objects so a future download can proceed cleanly
|
||||
if removedAny {
|
||||
self.asr = nil
|
||||
self.models = nil
|
||||
if currentVariant == variant {
|
||||
currentVariant = nil
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Returns all candidate directories where a Parakeet model might be cached.
|
||||
/// Includes both exact matches and prefixed directories (e.g. versioned folders).
|
||||
private func modelDirectories(_ variant: ParakeetModel) -> [URL] {
|
||||
let fm = FileManager.default
|
||||
var result: [URL] = [AsrModels.defaultCacheDirectory(for: variant.asrVersion)]
|
||||
|
||||
for root in candidateRoots() {
|
||||
for vendor in vendorDirs {
|
||||
let base = root.appendingPathComponent(vendor, isDirectory: true)
|
||||
// Exact match directory
|
||||
let direct = base.appendingPathComponent(variant.identifier, isDirectory: true)
|
||||
result.append(direct)
|
||||
// Prefixed directories (e.g. versioned folders)
|
||||
if let items = try? fm.contentsOfDirectory(at: base, includingPropertiesForKeys: [.isDirectoryKey], options: .skipsHiddenFiles) {
|
||||
for item in items where item.lastPathComponent.hasPrefix(variant.identifier) && item != direct {
|
||||
result.append(item)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
var seen = Set<String>()
|
||||
return result.filter { seen.insert($0.standardizedFileURL.path).inserted }
|
||||
}
|
||||
|
||||
private func candidateRoots() -> [URL] {
|
||||
let fm = FileManager.default
|
||||
let xdg = ProcessInfo.processInfo.environment["XDG_CACHE_HOME"].flatMap { URL(fileURLWithPath: $0, isDirectory: true) }
|
||||
let appSupport = try? fm.url(for: .applicationSupportDirectory, in: .userDomainMask, appropriateFor: nil, create: false)
|
||||
let appCache = try? URL.hexApplicationSupport.appendingPathComponent("cache", isDirectory: true)
|
||||
let userCache = FileManager.default.homeDirectoryForCurrentUser.appendingPathComponent(".cache", isDirectory: true)
|
||||
return [xdg, appCache, appSupport, userCache].compactMap { $0 }
|
||||
}
|
||||
}
|
||||
|
||||
private extension ParakeetModel {
|
||||
var asrVersion: AsrModelVersion {
|
||||
switch self {
|
||||
case .englishV2: return .v2
|
||||
case .multilingualV3: return .v3
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#else
|
||||
|
||||
actor ParakeetClient {
|
||||
func isModelAvailable(_ modelName: String) async -> Bool { false }
|
||||
func ensureLoaded(modelName: String, progress: @escaping (Progress) -> Void) async throws {
|
||||
throw NSError(
|
||||
domain: "Parakeet",
|
||||
code: -2,
|
||||
userInfo: [NSLocalizedDescriptionKey: "Parakeet support not linked. Add Swift Package: https://github.com/FluidInference/FluidAudio.git and link FluidAudio to Hex."]
|
||||
)
|
||||
}
|
||||
func transcribe(_ url: URL) async throws -> String { throw NSError(domain: "Parakeet", code: -3, userInfo: [NSLocalizedDescriptionKey: "Parakeet not available"]) }
|
||||
func deleteCaches(modelName: String) async throws {}
|
||||
}
|
||||
|
||||
#endif
|
||||
+118
@@ -0,0 +1,118 @@
|
||||
import AVFoundation
|
||||
import Foundation
|
||||
import HexCore
|
||||
import os.log
|
||||
|
||||
struct ParakeetClipPreparationResult {
|
||||
let url: URL
|
||||
private let cleanupURL: URL?
|
||||
|
||||
init(url: URL, cleanupURL: URL?) {
|
||||
self.url = url
|
||||
self.cleanupURL = cleanupURL
|
||||
}
|
||||
|
||||
func cleanup() {
|
||||
guard let cleanupURL else { return }
|
||||
try? FileManager.default.removeItem(at: cleanupURL)
|
||||
}
|
||||
}
|
||||
|
||||
enum ParakeetClipPreparer {
|
||||
private enum Error: LocalizedError {
|
||||
case unsupportedFormat
|
||||
case bufferAllocationFailed
|
||||
|
||||
var errorDescription: String? {
|
||||
switch self {
|
||||
case .unsupportedFormat:
|
||||
return "Parakeet can only pad mono Float32 PCM recordings."
|
||||
case .bufferAllocationFailed:
|
||||
return "Unable to allocate buffer while preparing Parakeet audio clip."
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// FluidAudio's LastChunkHandling guidance recommends chunk_duration 1.5s,
|
||||
// so pad to at least that window to avoid decoder errors.
|
||||
static let defaultMinimumDuration: TimeInterval = 1.5
|
||||
|
||||
static func ensureMinimumDuration(
|
||||
url: URL,
|
||||
minimumDuration: TimeInterval = defaultMinimumDuration,
|
||||
logger: os.Logger = HexLog.parakeet
|
||||
) throws -> ParakeetClipPreparationResult {
|
||||
let audioFile = try AVAudioFile(forReading: url)
|
||||
let format = audioFile.processingFormat
|
||||
let duration = Double(audioFile.length) / format.sampleRate
|
||||
|
||||
logger.debug(
|
||||
"Parakeet clip check file=\(url.lastPathComponent) duration=\(String(format: "%.3f", duration))s sampleRate=\(String(format: "%.0f", format.sampleRate))Hz channels=\(format.channelCount)"
|
||||
)
|
||||
|
||||
guard duration < minimumDuration else {
|
||||
return ParakeetClipPreparationResult(url: url, cleanupURL: nil)
|
||||
}
|
||||
|
||||
guard format.commonFormat == .pcmFormatFloat32 else {
|
||||
throw Error.unsupportedFormat
|
||||
}
|
||||
|
||||
let minimumFrames = AVAudioFrameCount((minimumDuration * format.sampleRate).rounded(.up))
|
||||
let existingFrames64 = max(AVAudioFramePosition(0), audioFile.length)
|
||||
let sourceCapacity = max(AVAudioFrameCount(min(existingFrames64, AVAudioFramePosition(AVAudioFrameCount.max))), 1)
|
||||
|
||||
guard
|
||||
let readBuffer = AVAudioPCMBuffer(pcmFormat: format, frameCapacity: sourceCapacity),
|
||||
let paddedBuffer = AVAudioPCMBuffer(pcmFormat: format, frameCapacity: minimumFrames)
|
||||
else {
|
||||
throw Error.bufferAllocationFailed
|
||||
}
|
||||
|
||||
try audioFile.read(into: readBuffer)
|
||||
let framesRead = min(readBuffer.frameLength, minimumFrames)
|
||||
|
||||
guard
|
||||
let sourceChannels = readBuffer.floatChannelData,
|
||||
let paddedChannels = paddedBuffer.floatChannelData
|
||||
else {
|
||||
throw Error.unsupportedFormat
|
||||
}
|
||||
|
||||
let channelCount = Int(format.channelCount)
|
||||
for channel in 0..<channelCount {
|
||||
let dest = paddedChannels[channel]
|
||||
let src = sourceChannels[channel]
|
||||
if framesRead > 0 {
|
||||
dest.update(from: src, count: Int(framesRead))
|
||||
}
|
||||
let padCount = Int(minimumFrames - framesRead)
|
||||
if padCount > 0 {
|
||||
dest.advanced(by: Int(framesRead)).initialize(repeating: 0, count: padCount)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
paddedBuffer.frameLength = minimumFrames
|
||||
|
||||
let paddedURL = makePaddedURL(from: url)
|
||||
if FileManager.default.fileExists(atPath: paddedURL.path) {
|
||||
try FileManager.default.removeItem(at: paddedURL)
|
||||
}
|
||||
|
||||
let paddedFile = try AVAudioFile(forWriting: paddedURL, settings: audioFile.fileFormat.settings)
|
||||
try paddedFile.write(from: paddedBuffer)
|
||||
|
||||
logger.notice(
|
||||
"Padded clip for Parakeet file=\(url.lastPathComponent) original=\(String(format: "%.3f", duration))s paddedTo=\(String(format: "%.3f", minimumDuration))s output=\(paddedURL.lastPathComponent)"
|
||||
)
|
||||
|
||||
return ParakeetClipPreparationResult(url: paddedURL, cleanupURL: paddedURL)
|
||||
}
|
||||
|
||||
private static func makePaddedURL(from url: URL) -> URL {
|
||||
let base = url.deletingLastPathComponent()
|
||||
let stem = url.deletingPathExtension().lastPathComponent
|
||||
return base.appendingPathComponent("\(stem)-parakeet-padded.wav")
|
||||
}
|
||||
}
|
||||
+380
@@ -0,0 +1,380 @@
|
||||
//
|
||||
// PasteboardClient.swift
|
||||
// Hex
|
||||
//
|
||||
// Created by Kit Langton on 1/24/25.
|
||||
//
|
||||
|
||||
import ComposableArchitecture
|
||||
import Dependencies
|
||||
import DependenciesMacros
|
||||
import Foundation
|
||||
import HexCore
|
||||
import Sauce
|
||||
import SwiftUI
|
||||
|
||||
private let pasteboardLogger = HexLog.pasteboard
|
||||
|
||||
@DependencyClient
|
||||
struct PasteboardClient {
|
||||
var paste: @Sendable (String) async -> Void
|
||||
var copy: @Sendable (String) async -> Void
|
||||
var sendKeyboardCommand: @Sendable (KeyboardCommand) async -> Void
|
||||
}
|
||||
|
||||
extension PasteboardClient: DependencyKey {
|
||||
static var liveValue: Self {
|
||||
let live = PasteboardClientLive()
|
||||
return .init(
|
||||
paste: { text in
|
||||
await live.paste(text: text)
|
||||
},
|
||||
copy: { text in
|
||||
await live.copy(text: text)
|
||||
},
|
||||
sendKeyboardCommand: { command in
|
||||
await live.sendKeyboardCommand(command)
|
||||
}
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
extension DependencyValues {
|
||||
var pasteboard: PasteboardClient {
|
||||
get { self[PasteboardClient.self] }
|
||||
set { self[PasteboardClient.self] = newValue }
|
||||
}
|
||||
}
|
||||
|
||||
struct PasteboardClientLive {
|
||||
@Shared(.hexSettings) var hexSettings: HexSettings
|
||||
|
||||
private struct PasteboardSnapshot {
|
||||
let items: [[String: Any]]
|
||||
|
||||
init(pasteboard: NSPasteboard) {
|
||||
var saved: [[String: Any]] = []
|
||||
for item in pasteboard.pasteboardItems ?? [] {
|
||||
var itemDict: [String: Any] = [:]
|
||||
for type in item.types {
|
||||
if let data = item.data(forType: type) {
|
||||
itemDict[type.rawValue] = data
|
||||
}
|
||||
}
|
||||
saved.append(itemDict)
|
||||
}
|
||||
self.items = saved
|
||||
}
|
||||
|
||||
func restore(to pasteboard: NSPasteboard) {
|
||||
pasteboard.clearContents()
|
||||
for itemDict in items {
|
||||
let item = NSPasteboardItem()
|
||||
for (type, data) in itemDict {
|
||||
if let data = data as? Data {
|
||||
item.setData(data, forType: NSPasteboard.PasteboardType(rawValue: type))
|
||||
}
|
||||
}
|
||||
pasteboard.writeObjects([item])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@MainActor
|
||||
func paste(text: String) async {
|
||||
if hexSettings.useClipboardPaste {
|
||||
await pasteWithClipboard(text)
|
||||
} else {
|
||||
simulateTypingWithAppleScript(text)
|
||||
}
|
||||
}
|
||||
|
||||
@MainActor
|
||||
func copy(text: String) async {
|
||||
let pasteboard = NSPasteboard.general
|
||||
pasteboard.clearContents()
|
||||
pasteboard.setString(text, forType: .string)
|
||||
}
|
||||
|
||||
@MainActor
|
||||
func sendKeyboardCommand(_ command: KeyboardCommand) async {
|
||||
let source = CGEventSource(stateID: .combinedSessionState)
|
||||
|
||||
// Convert modifiers to CGEventFlags and key codes for modifier keys
|
||||
var modifierKeyCodes: [CGKeyCode] = []
|
||||
var flags = CGEventFlags()
|
||||
|
||||
for modifier in command.modifiers.sorted {
|
||||
switch modifier.kind {
|
||||
case .command:
|
||||
flags.insert(.maskCommand)
|
||||
modifierKeyCodes.append(55) // Left Cmd
|
||||
case .shift:
|
||||
flags.insert(.maskShift)
|
||||
modifierKeyCodes.append(56) // Left Shift
|
||||
case .option:
|
||||
flags.insert(.maskAlternate)
|
||||
modifierKeyCodes.append(58) // Left Option
|
||||
case .control:
|
||||
flags.insert(.maskControl)
|
||||
modifierKeyCodes.append(59) // Left Control
|
||||
case .fn:
|
||||
flags.insert(.maskSecondaryFn)
|
||||
// Fn key doesn't need explicit key down/up
|
||||
}
|
||||
}
|
||||
|
||||
// Press modifiers down
|
||||
for keyCode in modifierKeyCodes {
|
||||
let modDown = CGEvent(keyboardEventSource: source, virtualKey: keyCode, keyDown: true)
|
||||
modDown?.post(tap: .cghidEventTap)
|
||||
}
|
||||
|
||||
// Press main key if present
|
||||
if let key = command.key {
|
||||
let keyCode = Sauce.shared.keyCode(for: key)
|
||||
|
||||
let keyDown = CGEvent(keyboardEventSource: source, virtualKey: keyCode, keyDown: true)
|
||||
keyDown?.flags = flags
|
||||
keyDown?.post(tap: .cghidEventTap)
|
||||
|
||||
let keyUp = CGEvent(keyboardEventSource: source, virtualKey: keyCode, keyDown: false)
|
||||
keyUp?.flags = flags
|
||||
keyUp?.post(tap: .cghidEventTap)
|
||||
}
|
||||
|
||||
// Release modifiers in reverse order
|
||||
for keyCode in modifierKeyCodes.reversed() {
|
||||
let modUp = CGEvent(keyboardEventSource: source, virtualKey: keyCode, keyDown: false)
|
||||
modUp?.post(tap: .cghidEventTap)
|
||||
}
|
||||
|
||||
pasteboardLogger.debug("Sent keyboard command: \(command.displayName)")
|
||||
}
|
||||
|
||||
/// Pastes current clipboard content to the frontmost application
|
||||
static func pasteToFrontmostApp() -> Bool {
|
||||
let script = """
|
||||
if application "System Events" is not running then
|
||||
tell application "System Events" to launch
|
||||
delay 0.1
|
||||
end if
|
||||
tell application "System Events"
|
||||
tell process (name of first application process whose frontmost is true)
|
||||
tell (menu item "Paste" of menu of menu item "Paste" of menu "Edit" of menu bar item "Edit" of menu bar 1)
|
||||
if exists then
|
||||
log (get properties of it)
|
||||
if enabled then
|
||||
click it
|
||||
return true
|
||||
else
|
||||
return false
|
||||
end if
|
||||
end if
|
||||
end tell
|
||||
tell (menu item "Paste" of menu "Edit" of menu bar item "Edit" of menu bar 1)
|
||||
if exists then
|
||||
if enabled then
|
||||
click it
|
||||
return true
|
||||
else
|
||||
return false
|
||||
end if
|
||||
else
|
||||
return false
|
||||
end if
|
||||
end tell
|
||||
end tell
|
||||
end tell
|
||||
"""
|
||||
|
||||
var error: NSDictionary?
|
||||
if let scriptObject = NSAppleScript(source: script) {
|
||||
let result = scriptObject.executeAndReturnError(&error)
|
||||
if let error = error {
|
||||
pasteboardLogger.error("AppleScript paste failed: \(error)")
|
||||
return false
|
||||
}
|
||||
return result.booleanValue
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
@MainActor
|
||||
func pasteWithClipboard(_ text: String) async {
|
||||
let pasteboard = NSPasteboard.general
|
||||
let snapshot = PasteboardSnapshot(pasteboard: pasteboard)
|
||||
let targetChangeCount = writeAndTrackChangeCount(pasteboard: pasteboard, text: text)
|
||||
_ = await waitForPasteboardCommit(targetChangeCount: targetChangeCount)
|
||||
let pasteSucceeded = await performPaste(text)
|
||||
|
||||
// Only restore original pasteboard contents if:
|
||||
// 1. Copying to clipboard is disabled AND
|
||||
// 2. The paste operation succeeded
|
||||
if !hexSettings.copyToClipboard && pasteSucceeded {
|
||||
let savedSnapshot = snapshot
|
||||
Task { @MainActor in
|
||||
// Give slower apps a short window to read the plain-text entry
|
||||
// before we repopulate the clipboard with the user's previous rich data.
|
||||
try? await Task.sleep(for: .milliseconds(500))
|
||||
pasteboard.clearContents()
|
||||
savedSnapshot.restore(to: pasteboard)
|
||||
}
|
||||
}
|
||||
|
||||
// If we failed to paste AND user doesn't want clipboard retention,
|
||||
// show a notification that text is available in clipboard
|
||||
if !pasteSucceeded && !hexSettings.copyToClipboard {
|
||||
// Keep the transcribed text in clipboard regardless of setting
|
||||
pasteboardLogger.notice("Paste operation failed; text remains in clipboard as fallback.")
|
||||
|
||||
// TODO: Could add a notification here to inform user
|
||||
// that text is available in clipboard
|
||||
}
|
||||
}
|
||||
|
||||
@MainActor
|
||||
private func writeAndTrackChangeCount(pasteboard: NSPasteboard, text: String) -> Int {
|
||||
let before = pasteboard.changeCount
|
||||
pasteboard.clearContents()
|
||||
pasteboard.setString(text, forType: .string)
|
||||
let after = pasteboard.changeCount
|
||||
if after == before {
|
||||
// Ensure we always advance by at least one to avoid infinite waits if the system
|
||||
// coalesces writes (seen on Sonoma betas with zero-length strings).
|
||||
return after + 1
|
||||
}
|
||||
return after
|
||||
}
|
||||
|
||||
@MainActor
|
||||
private func waitForPasteboardCommit(
|
||||
targetChangeCount: Int,
|
||||
timeout: Duration = .milliseconds(150),
|
||||
pollInterval: Duration = .milliseconds(5)
|
||||
) async -> Bool {
|
||||
guard targetChangeCount > NSPasteboard.general.changeCount else { return true }
|
||||
|
||||
let deadline = ContinuousClock.now + timeout
|
||||
while ContinuousClock.now < deadline {
|
||||
if NSPasteboard.general.changeCount >= targetChangeCount {
|
||||
return true
|
||||
}
|
||||
try? await Task.sleep(for: pollInterval)
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// MARK: - Paste Orchestration
|
||||
|
||||
@MainActor
|
||||
private enum PasteStrategy: CaseIterable {
|
||||
case cmdV
|
||||
case menuItem
|
||||
case accessibility
|
||||
}
|
||||
|
||||
@MainActor
|
||||
private func performPaste(_ text: String) async -> Bool {
|
||||
for strategy in PasteStrategy.allCases {
|
||||
if await attemptPaste(text, using: strategy) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
@MainActor
|
||||
private func attemptPaste(_ text: String, using strategy: PasteStrategy) async -> Bool {
|
||||
switch strategy {
|
||||
case .cmdV:
|
||||
return await postCmdV(delayMs: 0)
|
||||
case .menuItem:
|
||||
return PasteboardClientLive.pasteToFrontmostApp()
|
||||
case .accessibility:
|
||||
return (try? Self.insertTextAtCursor(text)) != nil
|
||||
}
|
||||
}
|
||||
|
||||
// MARK: - Helpers
|
||||
|
||||
@MainActor
|
||||
private func postCmdV(delayMs: Int) async -> Bool {
|
||||
// Optional tiny wait before keystrokes
|
||||
try? await wait(milliseconds: delayMs)
|
||||
let source = CGEventSource(stateID: .combinedSessionState)
|
||||
let vKey = vKeyCode()
|
||||
let cmdKey: CGKeyCode = 55
|
||||
let cmdDown = CGEvent(keyboardEventSource: source, virtualKey: cmdKey, keyDown: true)
|
||||
let vDown = CGEvent(keyboardEventSource: source, virtualKey: vKey, keyDown: true)
|
||||
vDown?.flags = .maskCommand
|
||||
let vUp = CGEvent(keyboardEventSource: source, virtualKey: vKey, keyDown: false)
|
||||
vUp?.flags = .maskCommand
|
||||
let cmdUp = CGEvent(keyboardEventSource: source, virtualKey: cmdKey, keyDown: false)
|
||||
cmdDown?.post(tap: .cghidEventTap)
|
||||
vDown?.post(tap: .cghidEventTap)
|
||||
vUp?.post(tap: .cghidEventTap)
|
||||
cmdUp?.post(tap: .cghidEventTap)
|
||||
return true
|
||||
}
|
||||
|
||||
@MainActor
|
||||
private func vKeyCode() -> CGKeyCode {
|
||||
if Thread.isMainThread { return Sauce.shared.keyCode(for: .v) }
|
||||
return DispatchQueue.main.sync { Sauce.shared.keyCode(for: .v) }
|
||||
}
|
||||
|
||||
@MainActor
|
||||
private func wait(milliseconds: Int) async throws {
|
||||
try Task.checkCancellation()
|
||||
try await Task.sleep(nanoseconds: UInt64(milliseconds) * 1_000_000)
|
||||
}
|
||||
|
||||
func simulateTypingWithAppleScript(_ text: String) {
|
||||
let escapedText = text.replacingOccurrences(of: "\"", with: "\\\"")
|
||||
let script = NSAppleScript(source: "tell application \"System Events\" to keystroke \"\(escapedText)\"")
|
||||
var error: NSDictionary?
|
||||
script?.executeAndReturnError(&error)
|
||||
if let error = error {
|
||||
pasteboardLogger.error("Error executing AppleScript typing fallback: \(error)")
|
||||
}
|
||||
}
|
||||
|
||||
enum PasteError: Error {
|
||||
case systemWideElementCreationFailed
|
||||
case focusedElementNotFound
|
||||
case elementDoesNotSupportTextEditing
|
||||
case failedToInsertText
|
||||
}
|
||||
|
||||
static func insertTextAtCursor(_ text: String) throws {
|
||||
// Get the system-wide accessibility element
|
||||
let systemWideElement = AXUIElementCreateSystemWide()
|
||||
|
||||
// Get the focused element
|
||||
var focusedElementRef: CFTypeRef?
|
||||
let axError = AXUIElementCopyAttributeValue(systemWideElement, kAXFocusedUIElementAttribute as CFString, &focusedElementRef)
|
||||
|
||||
guard axError == .success, let focusedElementRef = focusedElementRef else {
|
||||
throw PasteError.focusedElementNotFound
|
||||
}
|
||||
|
||||
let focusedElement = focusedElementRef as! AXUIElement
|
||||
|
||||
// Verify if the focused element supports text insertion
|
||||
var value: CFTypeRef?
|
||||
let supportsText = AXUIElementCopyAttributeValue(focusedElement, kAXValueAttribute as CFString, &value) == .success
|
||||
let supportsSelectedText = AXUIElementCopyAttributeValue(focusedElement, kAXSelectedTextAttribute as CFString, &value) == .success
|
||||
|
||||
if !supportsText && !supportsSelectedText {
|
||||
throw PasteError.elementDoesNotSupportTextEditing
|
||||
}
|
||||
|
||||
// Insert text at cursor position by replacing selected text (or empty selection)
|
||||
let insertResult = AXUIElementSetAttributeValue(focusedElement, kAXSelectedTextAttribute as CFString, text as CFTypeRef)
|
||||
|
||||
if insertResult != .success {
|
||||
throw PasteError.failedToInsertText
|
||||
}
|
||||
}
|
||||
}
|
||||
+1701
File diff suppressed because it is too large
Load Diff
Executable
+196
@@ -0,0 +1,196 @@
|
||||
//
|
||||
// SoundEffect.swift
|
||||
// Hex
|
||||
//
|
||||
// Created by Kit Langton on 1/26/25.
|
||||
//
|
||||
|
||||
import AVFoundation
|
||||
import ComposableArchitecture
|
||||
import Dependencies
|
||||
import DependenciesMacros
|
||||
import Foundation
|
||||
import HexCore
|
||||
import SwiftUI
|
||||
|
||||
// Thank you. Never mind then.What a beautiful idea.
|
||||
public enum SoundEffect: String, CaseIterable {
|
||||
case pasteTranscript
|
||||
case startRecording
|
||||
case stopRecording
|
||||
case cancel
|
||||
|
||||
public var fileName: String {
|
||||
self.rawValue
|
||||
}
|
||||
|
||||
var fileExtension: String {
|
||||
"mp3"
|
||||
}
|
||||
}
|
||||
|
||||
@DependencyClient
|
||||
public struct SoundEffectsClient {
|
||||
public var play: @Sendable (SoundEffect) -> Void
|
||||
public var stop: @Sendable (SoundEffect) -> Void
|
||||
public var stopAll: @Sendable () -> Void
|
||||
public var preloadSounds: @Sendable () async -> Void
|
||||
public var setEnabled: @Sendable (Bool) async -> Void
|
||||
}
|
||||
|
||||
extension SoundEffectsClient: DependencyKey {
|
||||
public static var liveValue: SoundEffectsClient {
|
||||
let live = SoundEffectsClientLive()
|
||||
return SoundEffectsClient(
|
||||
play: { soundEffect in
|
||||
Task { await live.play(soundEffect) }
|
||||
},
|
||||
stop: { soundEffect in
|
||||
Task { await live.stop(soundEffect) }
|
||||
},
|
||||
stopAll: {
|
||||
Task { await live.stopAll() }
|
||||
},
|
||||
preloadSounds: {
|
||||
await live.preloadSounds()
|
||||
},
|
||||
setEnabled: { enabled in
|
||||
await live.setEnabled(enabled)
|
||||
}
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
public extension DependencyValues {
|
||||
var soundEffects: SoundEffectsClient {
|
||||
get { self[SoundEffectsClient.self] }
|
||||
set { self[SoundEffectsClient.self] = newValue }
|
||||
}
|
||||
}
|
||||
|
||||
actor SoundEffectsClientLive {
|
||||
private let logger = HexLog.sound
|
||||
private let baselineVolume = HexSettings.baseSoundEffectsVolume
|
||||
|
||||
private let engine = AVAudioEngine()
|
||||
@Shared(.hexSettings) var hexSettings: HexSettings
|
||||
private var playerNodes: [SoundEffect: AVAudioPlayerNode] = [:]
|
||||
private var audioBuffers: [SoundEffect: AVAudioPCMBuffer] = [:]
|
||||
private var idleShutdownTask: Task<Void, Never>?
|
||||
/// Comfortably longer than any sound effect, short enough that an idle Hex doesn't keep
|
||||
/// an output IOProc (and coreaudiod) running around the clock (#209).
|
||||
private static let idleShutdownDelay: Duration = .seconds(10)
|
||||
|
||||
func play(_ soundEffect: SoundEffect) {
|
||||
guard hexSettings.soundEffectsEnabled else { return }
|
||||
guard let player = playerNodes[soundEffect], let buffer = audioBuffers[soundEffect] else {
|
||||
logger.error("Requested sound \(soundEffect.rawValue) not preloaded")
|
||||
return
|
||||
}
|
||||
prepareEngineIfNeeded()
|
||||
let clampedVolume = min(max(hexSettings.soundEffectsVolume, 0), baselineVolume)
|
||||
player.volume = Float(clampedVolume)
|
||||
player.stop()
|
||||
player.scheduleBuffer(buffer, at: nil, options: [], completionHandler: nil)
|
||||
player.play()
|
||||
scheduleIdleShutdown()
|
||||
}
|
||||
|
||||
/// Stops the output engine shortly after playback so it doesn't run while idle.
|
||||
/// Restarting it on the next play costs only a few milliseconds.
|
||||
private func scheduleIdleShutdown() {
|
||||
idleShutdownTask?.cancel()
|
||||
idleShutdownTask = Task {
|
||||
try? await Task.sleep(for: Self.idleShutdownDelay)
|
||||
guard !Task.isCancelled else { return }
|
||||
stopEngineIfNeeded()
|
||||
}
|
||||
}
|
||||
|
||||
func stop(_ soundEffect: SoundEffect) {
|
||||
playerNodes[soundEffect]?.stop()
|
||||
}
|
||||
|
||||
func stopAll() {
|
||||
playerNodes.values.forEach { $0.stop() }
|
||||
}
|
||||
|
||||
func preloadSounds() async {
|
||||
guard !isSetup else { return }
|
||||
|
||||
for soundEffect in SoundEffect.allCases {
|
||||
loadSound(soundEffect)
|
||||
}
|
||||
|
||||
isSetup = true
|
||||
}
|
||||
|
||||
func setEnabled(_: Bool) async {
|
||||
await preloadSounds()
|
||||
|
||||
// No prewarm on enable: play() starts the engine lazily, and an idle prewarm would
|
||||
// just be shut down again by the idle timer.
|
||||
if !hexSettings.soundEffectsEnabled {
|
||||
stopAll()
|
||||
idleShutdownTask?.cancel()
|
||||
stopEngineIfNeeded()
|
||||
}
|
||||
}
|
||||
|
||||
private var isSetup = false
|
||||
|
||||
private func loadSound(_ soundEffect: SoundEffect) {
|
||||
guard let url = Bundle.main.url(
|
||||
forResource: soundEffect.fileName,
|
||||
withExtension: soundEffect.fileExtension
|
||||
) else {
|
||||
logger.error("Missing sound resource \(soundEffect.fileName).\(soundEffect.fileExtension)")
|
||||
return
|
||||
}
|
||||
|
||||
do {
|
||||
let file = try AVAudioFile(forReading: url)
|
||||
let frameCount = AVAudioFrameCount(file.length)
|
||||
guard let buffer = AVAudioPCMBuffer(pcmFormat: file.processingFormat, frameCapacity: frameCount) else {
|
||||
logger.error("Failed to allocate buffer for \(soundEffect.rawValue)")
|
||||
return
|
||||
}
|
||||
try file.read(into: buffer)
|
||||
audioBuffers[soundEffect] = buffer
|
||||
|
||||
let player = AVAudioPlayerNode()
|
||||
engine.attach(player)
|
||||
engine.connect(player, to: engine.mainMixerNode, format: buffer.format)
|
||||
playerNodes[soundEffect] = player
|
||||
} catch {
|
||||
logger.error("Failed to load sound \(soundEffect.rawValue): \(error.localizedDescription)")
|
||||
}
|
||||
}
|
||||
|
||||
private func prepareEngineIfNeeded() {
|
||||
guard !engine.isRunning else { return }
|
||||
engine.prepare()
|
||||
if #available(macOS 13.0, *) {
|
||||
engine.isAutoShutdownEnabled = false
|
||||
}
|
||||
do {
|
||||
try engine.start()
|
||||
} catch {
|
||||
logger.error("Failed to start AVAudioEngine: \(error.localizedDescription)")
|
||||
}
|
||||
}
|
||||
|
||||
private func stopEngineIfNeeded() {
|
||||
guard engine.isRunning else { return }
|
||||
engine.stop()
|
||||
logger.debug("Sound effects engine stopped")
|
||||
}
|
||||
|
||||
deinit {
|
||||
playerNodes.values.forEach {
|
||||
$0.stop()
|
||||
engine.detach($0)
|
||||
}
|
||||
engine.stop()
|
||||
}
|
||||
}
|
||||
+527
@@ -0,0 +1,527 @@
|
||||
import AVFoundation
|
||||
import Foundation
|
||||
import HexCore
|
||||
|
||||
private final class FloatRingBuffer {
|
||||
private let lock = NSLock()
|
||||
private var buffer: [Float]
|
||||
private var writeIndex = 0
|
||||
private var validSampleCount = 0
|
||||
|
||||
init(capacity: Int) {
|
||||
buffer = Array(repeating: 0, count: max(1, capacity))
|
||||
}
|
||||
|
||||
func append(_ samples: UnsafeBufferPointer<Float>) {
|
||||
guard !samples.isEmpty else { return }
|
||||
|
||||
lock.lock()
|
||||
defer { lock.unlock() }
|
||||
|
||||
for sample in samples {
|
||||
buffer[writeIndex] = sample
|
||||
writeIndex = (writeIndex + 1) % buffer.count
|
||||
}
|
||||
|
||||
validSampleCount = min(buffer.count, validSampleCount + samples.count)
|
||||
}
|
||||
|
||||
func recentSamples(count requestedCount: Int) -> [Float] {
|
||||
lock.lock()
|
||||
defer { lock.unlock() }
|
||||
|
||||
let sampleCount = min(max(0, requestedCount), validSampleCount)
|
||||
guard sampleCount > 0 else { return [] }
|
||||
|
||||
let startIndex = (writeIndex - sampleCount + buffer.count) % buffer.count
|
||||
if startIndex + sampleCount <= buffer.count {
|
||||
return Array(buffer[startIndex ..< startIndex + sampleCount])
|
||||
}
|
||||
|
||||
let firstChunk = Array(buffer[startIndex ..< buffer.count])
|
||||
let secondChunk = Array(buffer[0 ..< (sampleCount - firstChunk.count)])
|
||||
return firstChunk + secondChunk
|
||||
}
|
||||
|
||||
func clear() {
|
||||
lock.lock()
|
||||
defer { lock.unlock() }
|
||||
|
||||
writeIndex = 0
|
||||
validSampleCount = 0
|
||||
}
|
||||
}
|
||||
|
||||
private struct SuperFastCaptureConstants {
|
||||
static let sampleRate: Double = 16_000
|
||||
static let ringBufferDuration: TimeInterval = 1.0
|
||||
static let defaultPreRollDuration: TimeInterval = 0.45
|
||||
static let tapBufferSize: AVAudioFrameCount = 2_048
|
||||
static let fallbackStopGracePeriod: TimeInterval = 0.05
|
||||
static let minimumStopGracePeriod: TimeInterval = 0.02
|
||||
static let maximumStopGracePeriod: TimeInterval = 0.08
|
||||
static let stopGraceSafetyMargin: TimeInterval = 0.008
|
||||
static let callbackTimingWindowSize = 8
|
||||
}
|
||||
|
||||
enum CaptureRecordingMode: String {
|
||||
case standard = "standard"
|
||||
case superFast = "super-fast"
|
||||
|
||||
var preRollDuration: TimeInterval {
|
||||
switch self {
|
||||
case .standard:
|
||||
0
|
||||
case .superFast:
|
||||
SuperFastCaptureConstants.defaultPreRollDuration
|
||||
}
|
||||
}
|
||||
|
||||
var keepsWarmBuffer: Bool {
|
||||
self == .superFast
|
||||
}
|
||||
}
|
||||
|
||||
final class SuperFastCaptureController {
|
||||
enum FinishRecordingResult {
|
||||
case captured(URL)
|
||||
case failed(RecordingFailure)
|
||||
case idle
|
||||
}
|
||||
|
||||
struct StopTimingEstimate {
|
||||
let gracePeriod: TimeInterval
|
||||
let callbackInterval: TimeInterval
|
||||
let bufferDuration: TimeInterval
|
||||
}
|
||||
|
||||
private struct ActiveRecording {
|
||||
let url: URL
|
||||
let file: AVAudioFile
|
||||
let requestedAt: Date
|
||||
let prependedDuration: TimeInterval
|
||||
var didLogFirstBuffer: Bool
|
||||
}
|
||||
|
||||
private let logger = HexLog.recording
|
||||
private let processingQueue = DispatchQueue(label: "com.kitlangton.Hex.SuperFastCapture")
|
||||
private let meterContinuation: AsyncStream<Meter>.Continuation
|
||||
private let ringBuffer = FloatRingBuffer(
|
||||
capacity: Int(SuperFastCaptureConstants.sampleRate * SuperFastCaptureConstants.ringBufferDuration)
|
||||
)
|
||||
private let targetFormat = AVAudioFormat(
|
||||
commonFormat: .pcmFormatFloat32,
|
||||
sampleRate: SuperFastCaptureConstants.sampleRate,
|
||||
channels: 1,
|
||||
interleaved: false
|
||||
)!
|
||||
|
||||
private var engine: AVAudioEngine?
|
||||
private var converter: AVAudioConverter?
|
||||
private var configurationChangeObserver: NSObjectProtocol?
|
||||
private var activeRecording: ActiveRecording?
|
||||
private var captureGeneration = 0
|
||||
private var recordingFailure: RecordingFailure?
|
||||
private var keepWarmBuffer = false
|
||||
private var lastProcessedBufferAt: Date?
|
||||
private var recentCallbackIntervals: [TimeInterval] = []
|
||||
private var recentBufferDurations: [TimeInterval] = []
|
||||
private let onEngineConfigurationChange: @Sendable (Int) -> Void
|
||||
|
||||
init(
|
||||
meterContinuation: AsyncStream<Meter>.Continuation,
|
||||
onEngineConfigurationChange: @escaping @Sendable (Int) -> Void
|
||||
) {
|
||||
self.meterContinuation = meterContinuation
|
||||
self.onEngineConfigurationChange = onEngineConfigurationChange
|
||||
}
|
||||
|
||||
deinit {
|
||||
stop()
|
||||
}
|
||||
|
||||
var isRunning: Bool {
|
||||
engine?.isRunning == true
|
||||
}
|
||||
|
||||
var isRecording: Bool {
|
||||
processingQueue.sync { activeRecording != nil }
|
||||
}
|
||||
|
||||
var stopTimingEstimate: StopTimingEstimate {
|
||||
processingQueue.sync {
|
||||
let callbackInterval = recentCallbackIntervals.max() ?? 0
|
||||
let bufferDuration = recentBufferDurations.max() ?? 0
|
||||
let observedCadence = max(callbackInterval, bufferDuration)
|
||||
let gracePeriod = min(
|
||||
max(
|
||||
observedCadence > 0
|
||||
? observedCadence + SuperFastCaptureConstants.stopGraceSafetyMargin
|
||||
: SuperFastCaptureConstants.fallbackStopGracePeriod,
|
||||
SuperFastCaptureConstants.minimumStopGracePeriod
|
||||
),
|
||||
SuperFastCaptureConstants.maximumStopGracePeriod
|
||||
)
|
||||
return StopTimingEstimate(
|
||||
gracePeriod: gracePeriod,
|
||||
callbackInterval: callbackInterval,
|
||||
bufferDuration: bufferDuration
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
func startIfNeeded(reason: String = "unknown", keepWarmBuffer: Bool = false) throws {
|
||||
processingQueue.sync {
|
||||
let didDisableWarmBuffer = self.keepWarmBuffer && !keepWarmBuffer
|
||||
self.keepWarmBuffer = keepWarmBuffer
|
||||
if didDisableWarmBuffer, activeRecording == nil {
|
||||
ringBuffer.clear()
|
||||
}
|
||||
}
|
||||
|
||||
if engine?.isRunning == true {
|
||||
logger.debug("Capture engine already armed reason=\(reason)")
|
||||
return
|
||||
}
|
||||
|
||||
stop(reason: "restart-before-arm")
|
||||
try armEngine(reason: reason)
|
||||
}
|
||||
|
||||
/// Tears down and recreates the engine while keeping the active recording file open, so
|
||||
/// capture resumes onto the same file after a device/route change mid-recording
|
||||
/// (#251, #252, #218, #226). The ring buffer, timing metrics, and active recording survive;
|
||||
/// only the engine, tap, and converter are rebuilt.
|
||||
func restartPreservingRecording(reason: String) throws {
|
||||
logger.notice("Restarting capture engine preserving active recording reason=\(reason)")
|
||||
detachEngine()
|
||||
try armEngine(reason: reason)
|
||||
}
|
||||
|
||||
private func armEngine(reason: String) throws {
|
||||
let engine = AVAudioEngine()
|
||||
let inputNode = engine.inputNode
|
||||
let inputFormat = inputNode.inputFormat(forBus: 0)
|
||||
guard let converter = AVAudioConverter(from: inputFormat, to: targetFormat) else {
|
||||
throw NSError(
|
||||
domain: "SuperFastCapture",
|
||||
code: -1,
|
||||
userInfo: [NSLocalizedDescriptionKey: "Unable to create the capture engine audio converter."]
|
||||
)
|
||||
}
|
||||
if inputFormat.channelCount > 1 {
|
||||
converter.channelMap = [NSNumber(value: 0)]
|
||||
}
|
||||
|
||||
let generation = processingQueue.sync {
|
||||
captureGeneration += 1
|
||||
self.converter = converter
|
||||
recordingFailure = nil
|
||||
return captureGeneration
|
||||
}
|
||||
|
||||
inputNode.installTap(onBus: 0, bufferSize: SuperFastCaptureConstants.tapBufferSize, format: inputFormat) {
|
||||
[weak self] buffer, _ in
|
||||
self?.enqueue(buffer, generation: generation)
|
||||
}
|
||||
|
||||
engine.prepare()
|
||||
do {
|
||||
try engine.start()
|
||||
} catch {
|
||||
inputNode.removeTap(onBus: 0)
|
||||
processingQueue.sync {
|
||||
captureGeneration += 1
|
||||
self.converter = nil
|
||||
}
|
||||
throw error
|
||||
}
|
||||
self.engine = engine
|
||||
configurationChangeObserver = NotificationCenter.default.addObserver(
|
||||
forName: .AVAudioEngineConfigurationChange,
|
||||
object: engine,
|
||||
queue: .main
|
||||
) { [weak self] _ in
|
||||
self?.handleConfigurationChange(generation: generation)
|
||||
}
|
||||
logger.notice(
|
||||
"Capture engine armed reason=\(reason) sampleRate=\(String(format: "%.0f", inputFormat.sampleRate))Hz channels=\(inputFormat.channelCount) ringBuffer=\(String(format: "%.2f", SuperFastCaptureConstants.ringBufferDuration))s defaultPreRoll=\(String(format: "%.2f", SuperFastCaptureConstants.defaultPreRollDuration))s"
|
||||
)
|
||||
}
|
||||
|
||||
func stop(reason: String = "unknown") {
|
||||
if engine != nil {
|
||||
logger.notice("Capture engine stopped reason=\(reason)")
|
||||
}
|
||||
detachEngine(clearingRecordingState: true)
|
||||
}
|
||||
|
||||
/// Removes the tap, observer, converter, and engine. Bumps the capture generation so
|
||||
/// in-flight tap callbacks from the old engine are ignored. Recording state (active file,
|
||||
/// ring buffer, timing metrics) is preserved unless `clearingRecordingState` is set, which
|
||||
/// is what lets restartPreservingRecording resume capture onto the same file.
|
||||
private func detachEngine(clearingRecordingState: Bool = false) {
|
||||
if let inputNode = engine?.inputNode {
|
||||
inputNode.removeTap(onBus: 0)
|
||||
}
|
||||
if let configurationChangeObserver {
|
||||
NotificationCenter.default.removeObserver(configurationChangeObserver)
|
||||
self.configurationChangeObserver = nil
|
||||
}
|
||||
processingQueue.sync {
|
||||
captureGeneration += 1
|
||||
converter = nil
|
||||
if clearingRecordingState {
|
||||
activeRecording = nil
|
||||
recordingFailure = nil
|
||||
ringBuffer.clear()
|
||||
lastProcessedBufferAt = nil
|
||||
recentCallbackIntervals.removeAll(keepingCapacity: false)
|
||||
recentBufferDurations.removeAll(keepingCapacity: false)
|
||||
}
|
||||
}
|
||||
engine?.stop()
|
||||
engine = nil
|
||||
}
|
||||
|
||||
private func handleConfigurationChange(generation: Int) {
|
||||
guard processingQueue.sync(execute: { Self.shouldProcessCallback(callbackGeneration: generation, currentGeneration: captureGeneration) }) else {
|
||||
return
|
||||
}
|
||||
logger.notice("Capture engine configuration changed")
|
||||
onEngineConfigurationChange(generation)
|
||||
}
|
||||
|
||||
static func shouldProcessCallback(callbackGeneration: Int, currentGeneration: Int) -> Bool {
|
||||
callbackGeneration == currentGeneration
|
||||
}
|
||||
|
||||
func isCurrentGeneration(_ generation: Int) -> Bool {
|
||||
processingQueue.sync { generation == captureGeneration }
|
||||
}
|
||||
|
||||
func beginRecording(to url: URL, requestedAt: Date = Date(), mode: CaptureRecordingMode) throws {
|
||||
try startIfNeeded(reason: "begin-recording", keepWarmBuffer: mode.keepsWarmBuffer)
|
||||
|
||||
var startError: Error?
|
||||
processingQueue.sync {
|
||||
do {
|
||||
recordingFailure = nil
|
||||
let file = try AVAudioFile(
|
||||
forWriting: url,
|
||||
settings: [
|
||||
AVFormatIDKey: Int(kAudioFormatLinearPCM),
|
||||
AVSampleRateKey: SuperFastCaptureConstants.sampleRate,
|
||||
AVNumberOfChannelsKey: 1,
|
||||
AVLinearPCMBitDepthKey: 32,
|
||||
AVLinearPCMIsFloatKey: true,
|
||||
AVLinearPCMIsBigEndianKey: false,
|
||||
AVLinearPCMIsNonInterleaved: true,
|
||||
],
|
||||
commonFormat: .pcmFormatFloat32,
|
||||
interleaved: false
|
||||
)
|
||||
|
||||
let preRollDuration = mode.preRollDuration
|
||||
let preRollFrameCount = Int(preRollDuration * SuperFastCaptureConstants.sampleRate)
|
||||
let preRollSamples = ringBuffer.recentSamples(count: preRollFrameCount)
|
||||
let prependedDuration = Double(preRollSamples.count) / SuperFastCaptureConstants.sampleRate
|
||||
if !preRollSamples.isEmpty {
|
||||
try write(samples: preRollSamples, to: file)
|
||||
}
|
||||
|
||||
logger.notice(
|
||||
"Capture engine recording file opened prepended=\(String(format: "%.3f", prependedDuration))s requestedPreRoll=\(String(format: "%.3f", preRollDuration))s"
|
||||
)
|
||||
activeRecording = ActiveRecording(
|
||||
url: url,
|
||||
file: file,
|
||||
requestedAt: requestedAt,
|
||||
prependedDuration: prependedDuration,
|
||||
didLogFirstBuffer: false
|
||||
)
|
||||
} catch {
|
||||
startError = error
|
||||
}
|
||||
}
|
||||
|
||||
if let startError {
|
||||
throw startError
|
||||
}
|
||||
}
|
||||
|
||||
func finishRecording(clearBuffer: Bool = true) -> FinishRecordingResult {
|
||||
processingQueue.sync {
|
||||
let result: FinishRecordingResult
|
||||
if let recordingFailure {
|
||||
result = .failed(recordingFailure)
|
||||
} else if let url = activeRecording?.url {
|
||||
result = .captured(url)
|
||||
} else {
|
||||
result = .idle
|
||||
}
|
||||
activeRecording = nil
|
||||
recordingFailure = nil
|
||||
if clearBuffer {
|
||||
ringBuffer.clear()
|
||||
}
|
||||
return result
|
||||
}
|
||||
}
|
||||
|
||||
private func enqueue(_ buffer: AVAudioPCMBuffer, generation: Int) {
|
||||
guard let copy = clone(buffer) else { return }
|
||||
processingQueue.async { [weak self] in
|
||||
self?.process(copy, generation: generation)
|
||||
}
|
||||
}
|
||||
|
||||
private func process(_ buffer: AVAudioPCMBuffer, generation: Int) {
|
||||
guard Self.shouldProcessCallback(callbackGeneration: generation, currentGeneration: captureGeneration) else {
|
||||
return
|
||||
}
|
||||
let now = Date()
|
||||
if let lastProcessedBufferAt {
|
||||
appendRecentMetric(now.timeIntervalSince(lastProcessedBufferAt), to: &recentCallbackIntervals)
|
||||
}
|
||||
lastProcessedBufferAt = now
|
||||
appendRecentMetric(Double(buffer.frameLength) / buffer.format.sampleRate, to: &recentBufferDurations)
|
||||
|
||||
guard let converted = convert(buffer),
|
||||
converted.frameLength > 0,
|
||||
let samples = converted.floatChannelData?[0]
|
||||
else {
|
||||
return
|
||||
}
|
||||
|
||||
let sampleCount = Int(converted.frameLength)
|
||||
if keepWarmBuffer, activeRecording == nil {
|
||||
ringBuffer.append(UnsafeBufferPointer(start: samples, count: sampleCount))
|
||||
}
|
||||
|
||||
if activeRecording != nil {
|
||||
meterContinuation.yield(meter(for: samples, count: sampleCount))
|
||||
}
|
||||
|
||||
guard var recording = activeRecording else { return }
|
||||
if !recording.didLogFirstBuffer {
|
||||
let timeToFirstBuffer = Date().timeIntervalSince(recording.requestedAt)
|
||||
logger.notice(
|
||||
"Capture engine first buffer latency=\(String(format: "%.3f", timeToFirstBuffer))s prepended=\(String(format: "%.3f", recording.prependedDuration))s frames=\(sampleCount)"
|
||||
)
|
||||
recording.didLogFirstBuffer = true
|
||||
activeRecording = recording
|
||||
}
|
||||
|
||||
do {
|
||||
try recording.file.write(from: converted)
|
||||
} catch {
|
||||
logger.error("Failed to write capture engine audio: \(error.localizedDescription)")
|
||||
activeRecording = nil
|
||||
recordingFailure = .captureWriteFailed(error.localizedDescription)
|
||||
FileManager.default.removeItemIfExists(at: recording.url)
|
||||
}
|
||||
}
|
||||
|
||||
private func convert(_ inputBuffer: AVAudioPCMBuffer) -> AVAudioPCMBuffer? {
|
||||
guard let converter else { return nil }
|
||||
|
||||
let sampleRateRatio = targetFormat.sampleRate / inputBuffer.format.sampleRate
|
||||
let frameCapacity = AVAudioFrameCount(
|
||||
max(1, (Double(inputBuffer.frameLength) * sampleRateRatio).rounded(.up) + 32)
|
||||
)
|
||||
|
||||
guard let outputBuffer = AVAudioPCMBuffer(pcmFormat: targetFormat, frameCapacity: frameCapacity) else {
|
||||
return nil
|
||||
}
|
||||
|
||||
var error: NSError?
|
||||
var consumedInput = false
|
||||
let status = converter.convert(to: outputBuffer, error: &error) { _, outStatus in
|
||||
if consumedInput {
|
||||
outStatus.pointee = .noDataNow
|
||||
return nil
|
||||
}
|
||||
|
||||
consumedInput = true
|
||||
outStatus.pointee = .haveData
|
||||
return inputBuffer
|
||||
}
|
||||
|
||||
if let error {
|
||||
logger.error("Failed to convert capture engine audio: \(error.localizedDescription)")
|
||||
return nil
|
||||
}
|
||||
|
||||
switch status {
|
||||
case .haveData, .inputRanDry, .endOfStream:
|
||||
return outputBuffer.frameLength > 0 ? outputBuffer : nil
|
||||
case .error:
|
||||
return nil
|
||||
@unknown default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
private func write(samples: [Float], to file: AVAudioFile) throws {
|
||||
guard !samples.isEmpty,
|
||||
let buffer = AVAudioPCMBuffer(pcmFormat: targetFormat, frameCapacity: AVAudioFrameCount(samples.count)),
|
||||
let channelData = buffer.floatChannelData?[0]
|
||||
else {
|
||||
return
|
||||
}
|
||||
|
||||
buffer.frameLength = AVAudioFrameCount(samples.count)
|
||||
samples.withUnsafeBufferPointer { sampleBuffer in
|
||||
guard let baseAddress = sampleBuffer.baseAddress else { return }
|
||||
channelData.update(from: baseAddress, count: sampleBuffer.count)
|
||||
}
|
||||
try file.write(from: buffer)
|
||||
}
|
||||
|
||||
private func meter(for samples: UnsafePointer<Float>, count: Int) -> Meter {
|
||||
guard count > 0 else {
|
||||
return Meter(averagePower: 0, peakPower: 0)
|
||||
}
|
||||
|
||||
var sumOfSquares: Float = 0
|
||||
var peak: Float = 0
|
||||
for index in 0 ..< count {
|
||||
let sample = samples[index]
|
||||
let magnitude = abs(sample)
|
||||
sumOfSquares += sample * sample
|
||||
peak = max(peak, magnitude)
|
||||
}
|
||||
|
||||
let rms = sqrt(sumOfSquares / Float(count))
|
||||
return Meter(averagePower: Double(rms), peakPower: Double(peak))
|
||||
}
|
||||
|
||||
private func clone(_ buffer: AVAudioPCMBuffer) -> AVAudioPCMBuffer? {
|
||||
guard let copy = AVAudioPCMBuffer(pcmFormat: buffer.format, frameCapacity: buffer.frameLength) else {
|
||||
return nil
|
||||
}
|
||||
|
||||
copy.frameLength = buffer.frameLength
|
||||
|
||||
let sourceBuffers = UnsafeMutableAudioBufferListPointer(buffer.mutableAudioBufferList)
|
||||
let destinationBuffers = UnsafeMutableAudioBufferListPointer(copy.mutableAudioBufferList)
|
||||
for index in sourceBuffers.indices {
|
||||
let source = sourceBuffers[index]
|
||||
let destination = destinationBuffers[index]
|
||||
guard let sourceData = source.mData, let destinationData = destination.mData else { continue }
|
||||
memcpy(destinationData, sourceData, Int(source.mDataByteSize))
|
||||
destinationBuffers[index].mDataByteSize = source.mDataByteSize
|
||||
}
|
||||
|
||||
return copy
|
||||
}
|
||||
|
||||
private func appendRecentMetric(_ value: TimeInterval, to metrics: inout [TimeInterval]) {
|
||||
guard value.isFinite, value > 0 else { return }
|
||||
metrics.append(value)
|
||||
if metrics.count > SuperFastCaptureConstants.callbackTimingWindowSize {
|
||||
metrics.removeFirst(metrics.count - SuperFastCaptureConstants.callbackTimingWindowSize)
|
||||
}
|
||||
}
|
||||
}
|
||||
+428
@@ -0,0 +1,428 @@
|
||||
//
|
||||
// TranscriptionClient.swift
|
||||
// Hex
|
||||
//
|
||||
// Created by Kit Langton on 1/24/25.
|
||||
//
|
||||
|
||||
import AVFoundation
|
||||
import Dependencies
|
||||
import DependenciesMacros
|
||||
import Foundation
|
||||
import HexCore
|
||||
import WhisperKit
|
||||
|
||||
private let transcriptionLogger = HexLog.transcription
|
||||
private let modelsLogger = HexLog.models
|
||||
private let parakeetLogger = HexLog.parakeet
|
||||
|
||||
/// A client that downloads and loads WhisperKit models, then transcribes audio files using the loaded model.
|
||||
/// Exposes progress callbacks to report overall download-and-load percentage and transcription progress.
|
||||
@DependencyClient
|
||||
struct TranscriptionClient {
|
||||
/// Transcribes an audio file at the specified `URL` using the named `model`.
|
||||
/// Reports transcription progress via `progressCallback`.
|
||||
var transcribe: @Sendable (URL, String, DecodingOptions, @escaping (Progress) -> Void) async throws -> String
|
||||
|
||||
/// Ensures a model is downloaded (if missing) and loaded into memory, reporting progress via `progressCallback`.
|
||||
var downloadModel: @Sendable (String, @escaping (Progress) -> Void) async throws -> Void
|
||||
|
||||
/// Deletes a model from disk if it exists
|
||||
var deleteModel: @Sendable (String) async throws -> Void
|
||||
|
||||
/// Checks if a named model is already downloaded on this system.
|
||||
var isModelDownloaded: @Sendable (String) async -> Bool = { _ in false }
|
||||
|
||||
/// Fetches a recommended set of models for the user's hardware from Hugging Face's `argmaxinc/whisperkit-coreml`.
|
||||
var getRecommendedModels: @Sendable () async throws -> ModelSupport
|
||||
|
||||
/// Lists all model variants found in `argmaxinc/whisperkit-coreml`.
|
||||
var getAvailableModels: @Sendable () async throws -> [String]
|
||||
}
|
||||
|
||||
extension TranscriptionClient: DependencyKey {
|
||||
static var liveValue: Self {
|
||||
let live = TranscriptionClientLive()
|
||||
return Self(
|
||||
transcribe: { try await live.transcribe(url: $0, model: $1, options: $2, progressCallback: $3) },
|
||||
downloadModel: { try await live.downloadAndLoadModel(variant: $0, progressCallback: $1) },
|
||||
deleteModel: { try await live.deleteModel(variant: $0) },
|
||||
isModelDownloaded: { await live.isModelDownloaded($0) },
|
||||
getRecommendedModels: { await live.getRecommendedModels() },
|
||||
getAvailableModels: { try await live.getAvailableModels() }
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
extension DependencyValues {
|
||||
var transcription: TranscriptionClient {
|
||||
get { self[TranscriptionClient.self] }
|
||||
set { self[TranscriptionClient.self] = newValue }
|
||||
}
|
||||
}
|
||||
|
||||
/// An `actor` that manages WhisperKit models by downloading (from Hugging Face),
|
||||
// loading them into memory, and then performing transcriptions.
|
||||
|
||||
actor TranscriptionClientLive {
|
||||
// MARK: - Stored Properties
|
||||
|
||||
/// The current in-memory `WhisperKit` instance, if any.
|
||||
private var whisperKit: WhisperKit?
|
||||
|
||||
/// The name of the currently loaded model, if any.
|
||||
private var currentModelName: String?
|
||||
private var parakeet: ParakeetClient = ParakeetClient()
|
||||
|
||||
/// The base folder under which we store model data (e.g., ~/Library/Application Support/...).
|
||||
private lazy var modelsBaseFolder: URL = {
|
||||
do {
|
||||
return try URL.hexModelsDirectory
|
||||
} catch {
|
||||
fatalError("Could not create Application Support folder: \(error)")
|
||||
}
|
||||
}()
|
||||
|
||||
// MARK: - Public Methods
|
||||
|
||||
/// Ensures the given `variant` model is downloaded and loaded, reporting
|
||||
/// overall progress (0%–50% for downloading, 50%–100% for loading).
|
||||
func downloadAndLoadModel(variant: String, progressCallback: @escaping (Progress) -> Void) async throws {
|
||||
// If Parakeet, use Parakeet client path
|
||||
if isParakeet(variant) {
|
||||
try await parakeet.ensureLoaded(modelName: variant, progress: progressCallback)
|
||||
currentModelName = variant
|
||||
return
|
||||
}
|
||||
// Resolve wildcard patterns (e.g., "distil*large-v3") to a concrete variant
|
||||
let variant = await resolveVariant(variant)
|
||||
// Special handling for corrupted or malformed variant names
|
||||
if variant.isEmpty {
|
||||
throw NSError(
|
||||
domain: "TranscriptionClient",
|
||||
code: -3,
|
||||
userInfo: [
|
||||
NSLocalizedDescriptionKey: "Cannot download model: Empty model name",
|
||||
]
|
||||
)
|
||||
}
|
||||
|
||||
let overallProgress = Progress(totalUnitCount: 100)
|
||||
overallProgress.completedUnitCount = 0
|
||||
progressCallback(overallProgress)
|
||||
|
||||
modelsLogger.info("Preparing model download and load for \(variant)")
|
||||
|
||||
// 1) Model download phase (0-50% progress)
|
||||
if !(await isModelDownloaded(variant)) {
|
||||
try await downloadModelIfNeeded(variant: variant) { downloadProgress in
|
||||
let fraction = downloadProgress.fractionCompleted * 0.5
|
||||
overallProgress.completedUnitCount = Int64(fraction * 100)
|
||||
progressCallback(overallProgress)
|
||||
}
|
||||
} else {
|
||||
// Skip download phase if already downloaded
|
||||
overallProgress.completedUnitCount = 50
|
||||
progressCallback(overallProgress)
|
||||
}
|
||||
|
||||
// 2) Model loading phase (50-100% progress)
|
||||
try await loadWhisperKitModel(variant) { loadingProgress in
|
||||
let fraction = 0.5 + (loadingProgress.fractionCompleted * 0.5)
|
||||
overallProgress.completedUnitCount = Int64(fraction * 100)
|
||||
progressCallback(overallProgress)
|
||||
}
|
||||
|
||||
// Final progress update
|
||||
overallProgress.completedUnitCount = 100
|
||||
progressCallback(overallProgress)
|
||||
}
|
||||
|
||||
/// Deletes a model from disk if it exists
|
||||
func deleteModel(variant: String) async throws {
|
||||
if isParakeet(variant) {
|
||||
try await parakeet.deleteCaches(modelName: variant)
|
||||
if currentModelName == variant { unloadCurrentModel() }
|
||||
return
|
||||
}
|
||||
let modelFolder = modelPath(for: variant)
|
||||
|
||||
// Check if the model exists
|
||||
guard FileManager.default.fileExists(atPath: modelFolder.path) else {
|
||||
// Model doesn't exist, nothing to delete
|
||||
return
|
||||
}
|
||||
|
||||
// If this is the currently loaded model, unload it first
|
||||
if currentModelName == variant {
|
||||
unloadCurrentModel()
|
||||
}
|
||||
|
||||
// Delete the model directory
|
||||
try FileManager.default.removeItem(at: modelFolder)
|
||||
|
||||
modelsLogger.info("Deleted model \(variant)")
|
||||
}
|
||||
|
||||
/// Returns `true` if the model is already downloaded to the local folder.
|
||||
/// Performs a thorough check to ensure the model files are actually present and usable.
|
||||
func isModelDownloaded(_ modelName: String) async -> Bool {
|
||||
if isParakeet(modelName) {
|
||||
let available = await parakeet.isModelAvailable(modelName)
|
||||
parakeetLogger.debug("Parakeet available? \(available)")
|
||||
return available
|
||||
}
|
||||
let modelFolderPath = modelPath(for: modelName).path
|
||||
let fileManager = FileManager.default
|
||||
|
||||
// First, check if the basic model directory exists
|
||||
guard fileManager.fileExists(atPath: modelFolderPath) else {
|
||||
// Don't print logs that would spam the console
|
||||
return false
|
||||
}
|
||||
|
||||
do {
|
||||
// Check if the directory has actual model files in it
|
||||
let contents = try fileManager.contentsOfDirectory(atPath: modelFolderPath)
|
||||
|
||||
// Model should have multiple files and certain key components
|
||||
guard !contents.isEmpty else {
|
||||
return false
|
||||
}
|
||||
|
||||
// Check for specific model structure - need both tokenizer and model files
|
||||
let hasModelFiles = contents.contains { $0.hasSuffix(".mlmodelc") || $0.contains("model") }
|
||||
let tokenizerFolderPath = tokenizerPath(for: modelName).path
|
||||
let hasTokenizer = fileManager.fileExists(atPath: tokenizerFolderPath)
|
||||
|
||||
// Both conditions must be true for a model to be considered downloaded
|
||||
return hasModelFiles && hasTokenizer
|
||||
} catch {
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
/// Returns a list of recommended models based on current device hardware.
|
||||
func getRecommendedModels() async -> ModelSupport {
|
||||
await WhisperKit.recommendedRemoteModels()
|
||||
}
|
||||
|
||||
/// Lists all model variants available in the `argmaxinc/whisperkit-coreml` repository.
|
||||
func getAvailableModels() async throws -> [String] {
|
||||
var names = try await WhisperKit.fetchAvailableModels()
|
||||
#if canImport(FluidAudio)
|
||||
for model in ParakeetModel.allCases.reversed() {
|
||||
if !names.contains(model.identifier) { names.insert(model.identifier, at: 0) }
|
||||
}
|
||||
#endif
|
||||
return names
|
||||
}
|
||||
|
||||
/// Transcribes the audio file at `url` using a `model` name.
|
||||
/// If the model is not yet loaded (or if it differs from the current model), it is downloaded and loaded first.
|
||||
/// Transcription progress can be monitored via `progressCallback`.
|
||||
func transcribe(
|
||||
url: URL,
|
||||
model: String,
|
||||
options: DecodingOptions,
|
||||
progressCallback: @escaping (Progress) -> Void
|
||||
) async throws -> String {
|
||||
let startAll = Date()
|
||||
if isParakeet(model) {
|
||||
transcriptionLogger.notice("Transcribing with Parakeet model=\(model) file=\(url.lastPathComponent)")
|
||||
let startLoad = Date()
|
||||
try await downloadAndLoadModel(variant: model) { p in
|
||||
progressCallback(p)
|
||||
}
|
||||
transcriptionLogger.info("Parakeet ensureLoaded took \(String(format: "%.2f", Date().timeIntervalSince(startLoad)))s")
|
||||
let preparedClip = try ParakeetClipPreparer.ensureMinimumDuration(url: url, logger: parakeetLogger)
|
||||
defer { preparedClip.cleanup() }
|
||||
let startTx = Date()
|
||||
let text = try await parakeet.transcribe(preparedClip.url)
|
||||
transcriptionLogger.info("Parakeet transcription took \(String(format: "%.2f", Date().timeIntervalSince(startTx)))s")
|
||||
transcriptionLogger.info("Parakeet request total elapsed \(String(format: "%.2f", Date().timeIntervalSince(startAll)))s")
|
||||
return text
|
||||
}
|
||||
let model = await resolveVariant(model)
|
||||
// Load or switch to the required model if needed.
|
||||
if whisperKit == nil || model != currentModelName {
|
||||
unloadCurrentModel()
|
||||
let startLoad = Date()
|
||||
try await downloadAndLoadModel(variant: model) { p in
|
||||
// Debug logging, or scale as desired:
|
||||
progressCallback(p)
|
||||
}
|
||||
let loadDuration = Date().timeIntervalSince(startLoad)
|
||||
transcriptionLogger.info("WhisperKit ensureLoaded model=\(model) took \(String(format: "%.2f", loadDuration))s")
|
||||
}
|
||||
|
||||
guard let whisperKit = whisperKit else {
|
||||
throw NSError(
|
||||
domain: "TranscriptionClient",
|
||||
code: -1,
|
||||
userInfo: [
|
||||
NSLocalizedDescriptionKey: "Failed to initialize WhisperKit for model: \(model)",
|
||||
]
|
||||
)
|
||||
}
|
||||
|
||||
// Perform the transcription.
|
||||
transcriptionLogger.notice("Transcribing with WhisperKit model=\(model) file=\(url.lastPathComponent)")
|
||||
let startTx = Date()
|
||||
let results = try await whisperKit.transcribe(audioPath: url.path, decodeOptions: options)
|
||||
transcriptionLogger.info("WhisperKit transcription took \(String(format: "%.2f", Date().timeIntervalSince(startTx)))s")
|
||||
transcriptionLogger.info("WhisperKit request total elapsed \(String(format: "%.2f", Date().timeIntervalSince(startAll)))s")
|
||||
|
||||
// Concatenate results from all segments.
|
||||
let text = results.map(\.text).joined(separator: " ")
|
||||
return text
|
||||
}
|
||||
|
||||
// MARK: - Private Helpers
|
||||
|
||||
/// Resolve wildcard patterns (e.g. "distil*large-v3") to a concrete model name.
|
||||
/// Preference: downloaded > non-turbo > any match.
|
||||
private func resolveVariant(_ variant: String) async -> String {
|
||||
guard variant.contains("*") || variant.contains("?") else { return variant }
|
||||
|
||||
let names: [String]
|
||||
do { names = try await WhisperKit.fetchAvailableModels() } catch { return variant }
|
||||
|
||||
// Build tuple array with download status for matching models
|
||||
var models: [(name: String, isDownloaded: Bool)] = []
|
||||
for name in names where ModelPatternMatcher.matches(variant, name) {
|
||||
models.append((name, await isModelDownloaded(name)))
|
||||
}
|
||||
|
||||
return ModelPatternMatcher.resolvePattern(variant, from: models) ?? variant
|
||||
}
|
||||
|
||||
private func isParakeet(_ name: String) -> Bool {
|
||||
ParakeetModel(rawValue: name) != nil
|
||||
}
|
||||
|
||||
/// Creates or returns the local folder (on disk) for a given `variant` model.
|
||||
private func modelPath(for variant: String) -> URL {
|
||||
// Remove any possible path traversal or invalid characters from variant name
|
||||
let sanitizedVariant = variant.components(separatedBy: CharacterSet(charactersIn: "./\\")).joined(separator: "_")
|
||||
|
||||
return modelsBaseFolder
|
||||
.appendingPathComponent("argmaxinc")
|
||||
.appendingPathComponent("whisperkit-coreml")
|
||||
.appendingPathComponent(sanitizedVariant, isDirectory: true)
|
||||
}
|
||||
|
||||
/// Creates or returns the local folder for the tokenizer files of a given `variant`.
|
||||
private func tokenizerPath(for variant: String) -> URL {
|
||||
modelPath(for: variant).appendingPathComponent("tokenizer", isDirectory: true)
|
||||
}
|
||||
|
||||
// Unloads any currently loaded model (clears `whisperKit` and `currentModelName`).
|
||||
private func unloadCurrentModel() {
|
||||
whisperKit = nil
|
||||
currentModelName = nil
|
||||
}
|
||||
|
||||
/// Downloads the model to a temporary folder (if it isn't already on disk),
|
||||
/// then moves it into its final folder in `modelsBaseFolder`.
|
||||
private func downloadModelIfNeeded(
|
||||
variant: String,
|
||||
progressCallback: @escaping (Progress) -> Void
|
||||
) async throws {
|
||||
let modelFolder = modelPath(for: variant)
|
||||
|
||||
// If the model folder exists but isn't a complete model, clean it up
|
||||
let isDownloaded = await isModelDownloaded(variant)
|
||||
if FileManager.default.fileExists(atPath: modelFolder.path), !isDownloaded {
|
||||
try FileManager.default.removeItem(at: modelFolder)
|
||||
}
|
||||
|
||||
// If model is already fully downloaded, we're done
|
||||
if isDownloaded {
|
||||
return
|
||||
}
|
||||
|
||||
modelsLogger.info("Downloading model \(variant)")
|
||||
|
||||
// Create parent directories
|
||||
let parentDir = modelFolder.deletingLastPathComponent()
|
||||
try FileManager.default.createDirectory(at: parentDir, withIntermediateDirectories: true)
|
||||
|
||||
do {
|
||||
// Download directly using the exact variant name provided
|
||||
// WhisperKit 0.15.0 changed downloader params: passing
|
||||
// "argmaxinc/whisperkit-coreml" to a parameter interpreted as a host
|
||||
// yields NSURLErrorCannotFindHost in production builds that need
|
||||
// to fetch models for the first time. Let WhisperKit use its
|
||||
// default repo/host (Hugging Face) by omitting the repo/host arg.
|
||||
let tempFolder = try await WhisperKit.download(
|
||||
variant: variant,
|
||||
downloadBase: nil,
|
||||
useBackgroundSession: false,
|
||||
progressCallback: { progress in
|
||||
progressCallback(progress)
|
||||
}
|
||||
)
|
||||
|
||||
// Ensure target folder exists
|
||||
try FileManager.default.createDirectory(at: modelFolder, withIntermediateDirectories: true)
|
||||
|
||||
// Move the downloaded snapshot to the final location
|
||||
try moveContents(of: tempFolder, to: modelFolder)
|
||||
|
||||
modelsLogger.info("Downloaded model to \(modelFolder.path)")
|
||||
} catch {
|
||||
// Clean up any partial download if an error occurred
|
||||
FileManager.default.removeItemIfExists(at: modelFolder)
|
||||
|
||||
// Rethrow the original error
|
||||
modelsLogger.error("Error downloading model \(variant): \(error.localizedDescription)")
|
||||
throw error
|
||||
}
|
||||
}
|
||||
|
||||
/// Loads a local model folder via `WhisperKitConfig`, optionally reporting load progress.
|
||||
private func loadWhisperKitModel(
|
||||
_ modelName: String,
|
||||
progressCallback: @escaping (Progress) -> Void
|
||||
) async throws {
|
||||
let loadingProgress = Progress(totalUnitCount: 100)
|
||||
loadingProgress.completedUnitCount = 0
|
||||
progressCallback(loadingProgress)
|
||||
|
||||
let modelFolder = modelPath(for: modelName)
|
||||
let tokenizerFolder = tokenizerPath(for: modelName)
|
||||
|
||||
// Use WhisperKit's config to load the model
|
||||
let config = WhisperKitConfig(
|
||||
model: modelName,
|
||||
modelFolder: modelFolder.path,
|
||||
tokenizerFolder: tokenizerFolder,
|
||||
// verbose: true,
|
||||
// logLevel: .debug,
|
||||
prewarm: false,
|
||||
load: true
|
||||
)
|
||||
|
||||
// The initializer automatically calls `loadModels`.
|
||||
whisperKit = try await WhisperKit(config)
|
||||
currentModelName = modelName
|
||||
|
||||
// Finalize load progress
|
||||
loadingProgress.completedUnitCount = 100
|
||||
progressCallback(loadingProgress)
|
||||
|
||||
modelsLogger.info("Loaded WhisperKit model \(modelName)")
|
||||
}
|
||||
|
||||
/// Moves all items from `sourceFolder` into `destFolder` (shallow move of directory contents).
|
||||
private func moveContents(of sourceFolder: URL, to destFolder: URL) throws {
|
||||
let fileManager = FileManager.default
|
||||
let items = try fileManager.contentsOfDirectory(atPath: sourceFolder.path)
|
||||
for item in items {
|
||||
let src = sourceFolder.appendingPathComponent(item)
|
||||
let dst = destFolder.appendingPathComponent(item)
|
||||
try fileManager.moveItem(at: src, to: dst)
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user