diff --git a/Sources/NIOSSH/Child Channels/ChildChannelUserEvents.swift b/Sources/NIOSSH/Child Channels/ChildChannelUserEvents.swift index 7c1a1613..b9a7abb3 100644 --- a/Sources/NIOSSH/Child Channels/ChildChannelUserEvents.swift +++ b/Sources/NIOSSH/Child Channels/ChildChannelUserEvents.swift @@ -309,6 +309,16 @@ public enum SSHChannelRequestEvent: Sendable { self.signal = signal } } + + /// A request from the client to enable SSH agent forwarding for this session. + public struct AgentForwardingRequest: Hashable, Sendable { + /// Whether a reply to this request is desired. + public var wantReply: Bool + + public init(wantReply: Bool) { + self.wantReply = wantReply + } + } } extension SSHChannelRequestEvent { @@ -350,6 +360,8 @@ extension SSHChannelRequestEvent { return LocalFlowControlRequest(clientCanDo: clientCanDo) case .signal(let signalName): return SignalRequest(signal: signalName) + case .agentForwarding: + return AgentForwardingRequest(wantReply: message.wantReply) case .unknown: return nil } diff --git a/Sources/NIOSSH/Child Channels/SSHChannelType.swift b/Sources/NIOSSH/Child Channels/SSHChannelType.swift index 0a537f20..b705882f 100644 --- a/Sources/NIOSSH/Child Channels/SSHChannelType.swift +++ b/Sources/NIOSSH/Child Channels/SSHChannelType.swift @@ -35,6 +35,9 @@ public enum SSHChannelType: Equatable, Sendable { /// "Forwarded TCP/IP" is a connection that was accepted from a listening socket and is being forwarded to the client. case forwardedTCPIP(ForwardedTCPIP) + + /// "Auth Agent" is a channel for forwarding SSH agent requests to the client's agent. + case authAgent } extension SSHChannelType { @@ -129,6 +132,8 @@ extension SSHChannelType { originatorAddress: message.originatorAddress ) ) + case .authAgent: + self = .authAgent } } } @@ -154,6 +159,8 @@ extension SSHMessage.ChannelOpenMessage.ChannelType { originatorAddress: data.originatorAddress ) ) + case .authAgent: + self = .authAgent } } } diff --git a/Sources/NIOSSH/SSHMessages.swift b/Sources/NIOSSH/SSHMessages.swift index a127d336..6d5921ef 100644 --- a/Sources/NIOSSH/SSHMessages.swift +++ b/Sources/NIOSSH/SSHMessages.swift @@ -221,6 +221,7 @@ extension SSHMessage { case session case forwardedTCPIP(ForwardedTCPIP) case directTCPIP(DirectTCPIP) + case authAgent } struct ForwardedTCPIP: Equatable { @@ -319,6 +320,7 @@ extension SSHMessage { case windowChange(WindowChange) case xonXoff(Bool) case signal(String) + case agentForwarding case unknown } @@ -902,6 +904,9 @@ extension ByteBuffer { ) ) + case "auth-agent@openssh.com": + type = .authAgent + default: throw NIOSSHError.unknownPacketType(diagnostic: "Channel request with \(typeRawValue)") } @@ -1127,6 +1132,8 @@ extension ByteBuffer { return nil } type = .signal(signalName) + case "auth-agent-req@openssh.com": + type = .agentForwarding default: type = .unknown } @@ -1454,6 +1461,9 @@ extension ByteBuffer { case .directTCPIP: writtenBytes += self.writeSSHString("direct-tcpip".utf8) + + case .authAgent: + writtenBytes += self.writeSSHString("auth-agent@openssh.com".utf8) } writtenBytes += self.writeInteger(message.senderChannel) @@ -1477,6 +1487,9 @@ extension ByteBuffer { writtenBytes += self.writeInteger(UInt32(data.portToConnectTo)) writtenBytes += self.writeSSHString((data.originatorAddress.ipAddress ?? "").utf8) writtenBytes += self.writeInteger(UInt32(data.originatorAddress.port ?? -1)) + + case .authAgent: + break } return writtenBytes @@ -1571,6 +1584,8 @@ extension ByteBuffer { writtenBytes += self.writeSSHString("xon-xoff".utf8) case .signal: writtenBytes += self.writeSSHString("signal".utf8) + case .agentForwarding: + writtenBytes += self.writeSSHString("auth-agent-req@openssh.com".utf8) case .unknown: preconditionFailure() } @@ -1610,6 +1625,8 @@ extension ByteBuffer { writtenBytes += self.writeSSHBoolean(clientCanDo) case .signal(let name): writtenBytes += self.writeSSHString(name.utf8) + case .agentForwarding: + break case .unknown: preconditionFailure() } diff --git a/Sources/NIOSSHServer/main.swift b/Sources/NIOSSHServer/main.swift index 28c9f6f3..f77f5bf2 100644 --- a/Sources/NIOSSHServer/main.swift +++ b/Sources/NIOSSHServer/main.swift @@ -85,6 +85,8 @@ func sshChildChannelInitializer(_ channel: Channel, _ channelType: SSHChannelTyp } case .forwardedTCPIP: return channel.eventLoop.makeFailedFuture(SSHServerError.invalidChannelType) + case .authAgent: + return channel.eventLoop.makeFailedFuture(SSHServerError.invalidChannelType) } } diff --git a/Tests/NIOSSHTests/ChildChannelMultiplexerTests.swift b/Tests/NIOSSHTests/ChildChannelMultiplexerTests.swift index f870cff8..b379811e 100644 --- a/Tests/NIOSSHTests/ChildChannelMultiplexerTests.swift +++ b/Tests/NIOSSHTests/ChildChannelMultiplexerTests.swift @@ -1901,6 +1901,7 @@ final class ChildChannelMultiplexerTests: XCTestCase { originatorAddress: try! .init(ipAddress: "fe80::1", port: 70) ) ), + SSHChannelType.authAgent, ] for channelType in channelTypes { @@ -1948,6 +1949,7 @@ final class ChildChannelMultiplexerTests: XCTestCase { originatorAddress: try! .init(ipAddress: "fe80::1", port: 70) ) ), + SSHChannelType.authAgent, ] for (channelID, channelType) in channelTypes.enumerated() { diff --git a/Tests/NIOSSHTests/SSHMessagesTests.swift b/Tests/NIOSSHTests/SSHMessagesTests.swift index 442ab36a..b4ba1407 100644 --- a/Tests/NIOSSHTests/SSHMessagesTests.swift +++ b/Tests/NIOSSHTests/SSHMessagesTests.swift @@ -557,6 +557,13 @@ final class SSHMessagesTests: XCTestCase { XCTAssertEqual(try buffer.readSSHMessage(), message) try self.assertCorrectlyManagesPartialRead(message) + message = SSHMessage.channelOpen( + .init(type: .authAgent, senderChannel: 0, initialWindowSize: 42, maximumPacketSize: 24) + ) + buffer.writeSSHMessage(message) + XCTAssertEqual(try buffer.readSSHMessage(), message) + try self.assertCorrectlyManagesPartialRead(message) + func writeBadMessage(into buffer: inout ByteBuffer, type: String, firstPort: UInt32, secondPort: UInt32) { buffer.writeInteger(SSHMessage.ChannelOpenMessage.id) buffer.writeSSHString(type.utf8) @@ -713,6 +720,16 @@ final class SSHMessagesTests: XCTestCase { XCTAssertEqual(try buffer.readSSHMessage(), message) try self.assertCorrectlyManagesPartialRead(message) + message = SSHMessage.channelRequest(.init(recipientChannel: 0, type: .agentForwarding, wantReply: true)) + buffer.writeSSHMessage(message) + XCTAssertEqual(try buffer.readSSHMessage(), message) + try self.assertCorrectlyManagesPartialRead(message) + + message = SSHMessage.channelRequest(.init(recipientChannel: 0, type: .agentForwarding, wantReply: false)) + buffer.writeSSHMessage(message) + XCTAssertEqual(try buffer.readSSHMessage(), message) + try self.assertCorrectlyManagesPartialRead(message) + buffer.writeBytes([SSHMessage.ChannelRequestMessage.id, 0, 0]) XCTAssertNil(try buffer.readSSHMessage()) } @@ -730,6 +747,24 @@ final class SSHMessagesTests: XCTestCase { try self.assertCorrectlyManagesPartialRead(message) } + func testAgentForwardingChannelRequestEvent() throws { + let wantReplyMessage = SSHMessage.ChannelRequestMessage( + recipientChannel: 0, + type: .agentForwarding, + wantReply: true + ) + let event = SSHChannelRequestEvent.fromMessage(wantReplyMessage) + XCTAssertEqual(event as? SSHChannelRequestEvent.AgentForwardingRequest, .init(wantReply: true)) + + let noReplyMessage = SSHMessage.ChannelRequestMessage( + recipientChannel: 0, + type: .agentForwarding, + wantReply: false + ) + let noReplyEvent = SSHChannelRequestEvent.fromMessage(noReplyMessage) + XCTAssertEqual(noReplyEvent as? SSHChannelRequestEvent.AgentForwardingRequest, .init(wantReply: false)) + } + func testChannelFailure() throws { var buffer = ByteBufferAllocator().buffer(capacity: 100) let message = SSHMessage.channelFailure(.init(recipientChannel: 0))