diff --git a/native/computer-use-macos/Sources/OrcaComputerUseMacOS/main.swift b/native/computer-use-macos/Sources/OrcaComputerUseMacOS/main.swift index da6c4d34cda..d7cd9204623 100644 --- a/native/computer-use-macos/Sources/OrcaComputerUseMacOS/main.swift +++ b/native/computer-use-macos/Sources/OrcaComputerUseMacOS/main.swift @@ -2530,9 +2530,12 @@ private enum KeyMap { } private final class AgentRuntime: NSObject, NSApplicationDelegate { + private static let unclaimedSessionDeadline: TimeInterval = 30 + private let socketPath: String private let token: String? private var listener: SocketListener? + private var unclaimedSessionTimeout: DispatchWorkItem? init(socketPath: String, token: String?) { self.socketPath = socketPath @@ -2541,9 +2544,31 @@ private final class AgentRuntime: NSObject, NSApplicationDelegate { func applicationDidFinishLaunching(_ notification: Notification) { do { - let listener = try SocketListener(socketPath: socketPath, token: token) + let timeout = DispatchWorkItem { + fputs("computer-use agent received no authenticated session before its deadline\n", stderr) + NSApp.terminate(nil) + } + unclaimedSessionTimeout = timeout + let listener = try SocketListener( + socketPath: socketPath, + token: token, + onSessionClaimed: { + DispatchQueue.main.async { + timeout.cancel() + } + }, + onSessionClosed: { + DispatchQueue.main.async { + NSApp.terminate(nil) + } + } + ) self.listener = listener listener.start() + DispatchQueue.main.asyncAfter( + deadline: .now() + Self.unclaimedSessionDeadline, + execute: timeout + ) } catch { fputs("failed to start computer-use socket: \(error)\n", stderr) NSApp.terminate(nil) @@ -2551,6 +2576,8 @@ private final class AgentRuntime: NSObject, NSApplicationDelegate { } func applicationWillTerminate(_ notification: Notification) { + unclaimedSessionTimeout?.cancel() + unclaimedSessionTimeout = nil listener?.stop() } } @@ -3456,14 +3483,25 @@ private final class ButtonTarget: NSObject { private final class SocketListener: @unchecked Sendable { private let socketPath: String private let token: String? + private let onSessionClaimed: () -> Void + private let onSessionClosed: () -> Void private let provider = Provider() private let providerLock = NSLock() + private let sessionLock = NSLock() + private var sessionOwnership = AgentSessionOwnership() private var socketFd: Int32 = -1 private var isStopped = false - init(socketPath: String, token: String?) throws { + init( + socketPath: String, + token: String?, + onSessionClaimed: @escaping () -> Void, + onSessionClosed: @escaping () -> Void + ) throws { self.socketPath = socketPath self.token = token + self.onSessionClaimed = onSessionClaimed + self.onSessionClosed = onSessionClosed try bindSocket() } @@ -3545,7 +3583,15 @@ private final class SocketListener: @unchecked Sendable { } private func handleConnection(_ fd: Int32) { - defer { close(fd) } + defer { + sessionLock.lock() + let shouldTerminate = sessionOwnership.disconnect(fd) + sessionLock.unlock() + close(fd) + if shouldTerminate { + onSessionClosed() + } + } let authorizedPeer = peerProcessId(fd).map(isAuthorizedAgentPeer) == true let decoder = JSONDecoder() while let line = readLine(from: fd) { @@ -3554,6 +3600,15 @@ private final class SocketListener: @unchecked Sendable { else { continue } + sessionLock.lock() + let claimed = sessionOwnership.registerConnection( + fd, + authenticated: isAuthenticated(request: request, authorizedPeer: authorizedPeer) + ) + sessionLock.unlock() + if claimed { + onSessionClaimed() + } let response = handleRequest( provider: provider, lock: providerLock, @@ -3564,6 +3619,11 @@ private final class SocketListener: @unchecked Sendable { writeJSON(response, to: fd) } } + + private func isAuthenticated(request: Request, authorizedPeer: Bool) -> Bool { + guard let token else { return true } + return request.token == token && authorizedPeer + } } private func existingPathMode(_ path: String) -> mode_t? { diff --git a/native/computer-use-macos/Sources/OrcaComputerUseMacOSCore/AgentSessionOwnership.swift b/native/computer-use-macos/Sources/OrcaComputerUseMacOSCore/AgentSessionOwnership.swift new file mode 100644 index 00000000000..0fb5258163a --- /dev/null +++ b/native/computer-use-macos/Sources/OrcaComputerUseMacOSCore/AgentSessionOwnership.swift @@ -0,0 +1,19 @@ +public struct AgentSessionOwnership: Sendable { + private var authenticatedConnections: Set = [] + private var wasClaimed = false + + public init() {} + + public mutating func registerConnection(_ connection: Int32, authenticated: Bool) -> Bool { + guard authenticated else { return false } + let inserted = authenticatedConnections.insert(connection).inserted + guard inserted, !wasClaimed else { return false } + wasClaimed = true + return true + } + + public mutating func disconnect(_ connection: Int32) -> Bool { + guard authenticatedConnections.remove(connection) != nil else { return false } + return wasClaimed && authenticatedConnections.isEmpty + } +} diff --git a/native/computer-use-macos/Tests/OrcaComputerUseMacOSTests/AgentSessionOwnershipTests.swift b/native/computer-use-macos/Tests/OrcaComputerUseMacOSTests/AgentSessionOwnershipTests.swift new file mode 100644 index 00000000000..f050ac5a5e4 --- /dev/null +++ b/native/computer-use-macos/Tests/OrcaComputerUseMacOSTests/AgentSessionOwnershipTests.swift @@ -0,0 +1,41 @@ +import OrcaComputerUseMacOSCore +import XCTest + +final class AgentSessionOwnershipTests: XCTestCase { + func testUnclaimedDisconnectDoesNotTerminateAgent() { + var ownership = AgentSessionOwnership() + + XCTAssertFalse(ownership.disconnect(12)) + } + + func testUnauthenticatedConnectionCannotClaimOrRetainAgent() { + var ownership = AgentSessionOwnership() + + XCTAssertFalse(ownership.registerConnection(12, authenticated: false)) + XCTAssertFalse(ownership.disconnect(12)) + } + + func testLastAuthenticatedDisconnectTerminatesAgent() { + var ownership = AgentSessionOwnership() + + XCTAssertTrue(ownership.registerConnection(12, authenticated: true)) + XCTAssertTrue(ownership.disconnect(12)) + } + + func testAgentWaitsForEveryAuthenticatedConnectionToClose() { + var ownership = AgentSessionOwnership() + + XCTAssertTrue(ownership.registerConnection(12, authenticated: true)) + XCTAssertFalse(ownership.registerConnection(13, authenticated: true)) + XCTAssertFalse(ownership.disconnect(12)) + XCTAssertTrue(ownership.disconnect(13)) + } + + func testDuplicateRegistrationDoesNotRetainAgent() { + var ownership = AgentSessionOwnership() + + XCTAssertTrue(ownership.registerConnection(12, authenticated: true)) + XCTAssertFalse(ownership.registerConnection(12, authenticated: true)) + XCTAssertTrue(ownership.disconnect(12)) + } +}