Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 12 additions & 0 deletions Sources/NIOSSH/Child Channels/ChildChannelUserEvents.swift
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -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
}
Expand Down
7 changes: 7 additions & 0 deletions Sources/NIOSSH/Child Channels/SSHChannelType.swift
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -129,6 +132,8 @@ extension SSHChannelType {
originatorAddress: message.originatorAddress
)
)
case .authAgent:
self = .authAgent
}
}
}
Expand All @@ -154,6 +159,8 @@ extension SSHMessage.ChannelOpenMessage.ChannelType {
originatorAddress: data.originatorAddress
)
)
case .authAgent:
self = .authAgent
}
}
}
17 changes: 17 additions & 0 deletions Sources/NIOSSH/SSHMessages.swift
Original file line number Diff line number Diff line change
Expand Up @@ -221,6 +221,7 @@ extension SSHMessage {
case session
case forwardedTCPIP(ForwardedTCPIP)
case directTCPIP(DirectTCPIP)
case authAgent
}

struct ForwardedTCPIP: Equatable {
Expand Down Expand Up @@ -319,6 +320,7 @@ extension SSHMessage {
case windowChange(WindowChange)
case xonXoff(Bool)
case signal(String)
case agentForwarding
case unknown
}

Expand Down Expand Up @@ -902,6 +904,9 @@ extension ByteBuffer {
)
)

case "auth-agent@openssh.com":
type = .authAgent

default:
throw NIOSSHError.unknownPacketType(diagnostic: "Channel request with \(typeRawValue)")
}
Expand Down Expand Up @@ -1127,6 +1132,8 @@ extension ByteBuffer {
return nil
}
type = .signal(signalName)
case "auth-agent-req@openssh.com":
type = .agentForwarding
default:
type = .unknown
}
Expand Down Expand Up @@ -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)
Expand All @@ -1477,6 +1487,9 @@ extension ByteBuffer {
writtenBytes += self.writeInteger(UInt32(data.portToConnectTo))
writtenBytes += self.writeSSHString((data.originatorAddress.ipAddress ?? "<nio-error>").utf8)
writtenBytes += self.writeInteger(UInt32(data.originatorAddress.port ?? -1))

case .authAgent:
break
}

return writtenBytes
Expand Down Expand Up @@ -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()
}
Expand Down Expand Up @@ -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()
}
Expand Down
2 changes: 2 additions & 0 deletions Sources/NIOSSHServer/main.swift
Original file line number Diff line number Diff line change
Expand Up @@ -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)
}
}

Expand Down
2 changes: 2 additions & 0 deletions Tests/NIOSSHTests/ChildChannelMultiplexerTests.swift
Original file line number Diff line number Diff line change
Expand Up @@ -1901,6 +1901,7 @@ final class ChildChannelMultiplexerTests: XCTestCase {
originatorAddress: try! .init(ipAddress: "fe80::1", port: 70)
)
),
SSHChannelType.authAgent,
]

for channelType in channelTypes {
Expand Down Expand Up @@ -1948,6 +1949,7 @@ final class ChildChannelMultiplexerTests: XCTestCase {
originatorAddress: try! .init(ipAddress: "fe80::1", port: 70)
)
),
SSHChannelType.authAgent,
]

for (channelID, channelType) in channelTypes.enumerated() {
Expand Down
35 changes: 35 additions & 0 deletions Tests/NIOSSHTests/SSHMessagesTests.swift
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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())
}
Expand All @@ -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))
Expand Down
Loading