Skip to content
Merged
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
5 changes: 4 additions & 1 deletion FlyingFox/Sources/HTTPConnection.swift
Original file line number Diff line number Diff line change
Expand Up @@ -92,7 +92,10 @@ struct HTTPConnection: Sendable {
func switchToWebSocket(with handler: some WSHandler, response: Data) async throws {
// Reuse the connection-wide buffered stream so any bytes already
// pulled past the upgrade request remain available to the WS framer.
let client = AsyncThrowingStream.decodingFrames(from: bytes)
// Client frames must be masked (RFC 6455 §5.1); an unmasked frame
// fails the stream. MessageFrameWSHandler responds with a 1002 close
// frame; custom handlers own their error handling.
let client = AsyncThrowingStream.decodingClientFrames(from: bytes)
let server = try await handler.makeFrames(for: client)
try await socket.write(response)
logger.logSwitchProtocol(self, to: "websocket")
Expand Down
15 changes: 15 additions & 0 deletions FlyingFox/Sources/WebSocket/AsyncStream+WSFrame.swift
Original file line number Diff line number Diff line change
Expand Up @@ -44,6 +44,21 @@ extension AsyncThrowingStream<WSFrame, any Error> {
}
}
}

/// Decodes client → server frames via `WSFrameEncoder.decodeClientFrame(from:)`,
/// throwing when a frame is unmasked (RFC 6455 §5.1). A clean disconnect
/// ends the stream without error.
static func decodingClientFrames(from bytes: some AsyncBufferedSequence<UInt8>) -> Self {
AsyncThrowingStream<WSFrame, any Error> {
do {
return try await WSFrameEncoder.decodeClientFrame(from: bytes)
} catch SocketError.disconnected, is SequenceTerminationError {
return nil
} catch {
throw error
}
}
}
}

extension AsyncStream<WSFrame> {
Expand Down
26 changes: 26 additions & 0 deletions FlyingFox/Sources/WebSocket/WSFrameEncoder.swift
Original file line number Diff line number Diff line change
Expand Up @@ -70,10 +70,36 @@ struct WSFrameEncoder {
static func decodeFrame(from bytes: some AsyncBufferedSequence<UInt8>) async throws -> WSFrame {
var frame = try await decodeFrame(from: bytes.take())
let (length, mask) = try await decodeLengthMask(from: bytes)
// The payload is stored unmasked; the mask is preserved so servers can
// enforce RFC 6455 §5.1 — "a client MUST mask all frames that it sends
// to the server."
frame.mask = mask
frame.payload = try await decodePayload(from: bytes, length: length, mask: mask)
return frame
}

/// Decodes a single client → server frame, enforcing RFC 6455 §5.1:
/// "a client MUST mask all frames that it sends to the server. ...
/// The server MUST close the connection upon receiving a frame that is
/// not masked."
///
/// The error thrown for an unmasked frame surfaces through the client
/// frame stream received by `WSHandler`, which is responsible for
/// terminating the connection — `MessageFrameWSHandler` responds with a
/// 1002 (protocol error) close frame.
///
/// The mask is cleared on the returned frame so handlers can safely echo
/// frames (e.g. ping → pong) — "A server MUST NOT mask any frames that it
/// sends to the client." (§5.1)
static func decodeClientFrame(from bytes: some AsyncBufferedSequence<UInt8>) async throws -> WSFrame {
var frame = try await decodeFrame(from: bytes)
guard frame.mask != nil else {
throw Error("Incoming client frames must be masked")
}
frame.mask = nil
return frame
}

static func encodeFrame0(_ frame: WSFrame) -> UInt8 {
var byte: UInt8 = frame.opcode.rawValue
byte |= frame.fin.byte << 7
Expand Down
97 changes: 97 additions & 0 deletions FlyingFox/Tests/HTTPConnectionTests.swift
Original file line number Diff line number Diff line change
Expand Up @@ -171,6 +171,103 @@ struct HTTPConnectionTests {
HTTPConnection.makeIdentifier(from: .unix("/var/sock/fox")) == "/var/sock/fox"
)
}

@Test
func webSocket_UnmaskedClientFrame_ClosesWithProtocolError() async throws {
// RFC 6455 §5.1: "The server MUST close the connection upon receiving
// a frame that is not masked. In this case, a server MAY send a Close
// frame with a status code of 1002 (protocol error)."
// Frame equality also asserts the close frame leaves the wire
// unmasked; `sendResponse` returning proves the response loop
// terminates — the owning `HTTPServer` then closes the socket.
let (s1, s2) = try await AsyncSocket.makePair()
let connection = HTTPConnection(socket: s1)

let response = Task {
try await connection.sendResponse(HTTPResponse(webSocket: MessageFrameWSHandler.make()))
}

_ = try await s2.readResponse()
try await s2.writeFrame(.fish)

// .close(message:) carries WSCloseCode.protocolError (1002).
#expect(
try await s2.readFrame() == .close(message: "Protocol Error")
)
try await response.value

try s1.close()
try s2.close()
}

@Test
func webSocket_MaskedClientFrames_AreDeliveredToHandlerUnmasked() async throws {
// Wire masks are consumed by `decodeClientFrame`; handlers receive
// frames with `mask == nil` and the payload already unmasked.
let (s1, s2) = try await AsyncSocket.makePair()
let connection = HTTPConnection(socket: s1)

let response = Task {
try await connection.sendResponse(HTTPResponse(webSocket: MaskReportingWSHandler()))
}

_ = try await s2.readResponse()
try await s2.writeFrame(.fish.masked())

#expect(
try await s2.readFrame() == .make(
opcode: .binary,
payload: Data([1]) + "Fish".data(using: .utf8)!
)
)

response.cancel()
try s1.close()
try s2.close()
}

@Test
func webSocket_ClientDisconnect_EndsConnectionWithoutError() async throws {
// A peer that closes TCP without sending a Close frame ends the client
// stream (SocketError.disconnected → nil); the handler's output then
// finishes and the response loop completes cleanly.
let (s1, s2) = try await AsyncSocket.makePair()
let connection = HTTPConnection(socket: s1)

let response = Task {
try await connection.sendResponse(HTTPResponse(webSocket: MessageFrameWSHandler.make()))
}

_ = try await s2.readResponse()
try s2.close()

try await response.value

try s1.close()
}
}

private struct MaskReportingWSHandler: WSHandler {
// Echoes each frame as binary: first byte 1 when the received frame had
// no mask, followed by the received payload.
func makeFrames(for client: AsyncThrowingStream<WSFrame, any Error>) async throws -> AsyncStream<WSFrame> {
AsyncStream { continuation in
let task = Task {
do {
for try await frame in client {
continuation.yield(
WSFrame(fin: true,
opcode: .binary,
mask: nil,
payload: Data([frame.mask == nil ? 1 : 0]) + frame.payload)
)
}
} catch { }
continuation.finish()
}
continuation.onTermination = { _ in task.cancel() }
}
}
}

private extension HTTPConnection {
Expand Down
34 changes: 34 additions & 0 deletions FlyingFox/Tests/WebSocket/AsyncStream+WSFrameTests.swift
Original file line number Diff line number Diff line change
Expand Up @@ -46,11 +46,40 @@ struct WSFrameSequenceTests {
#expect(
try await AsyncThrowingStream.make([.close]).collectAll() == [.close]
)
// Decoding preserves the mask of a masked frame (RFC 6455 §5.1) while
// storing the payload unmasked.
#expect(
try await AsyncThrowingStream.make([.fish.masked()]).collectAll() == [.fish.masked()]
)
#expect(
try await AsyncThrowingStream.make([]).collectAll() == []
)
}

@Test
func clientSequence_DeliversMaskedFramesUnmasked() async throws {
// Masked client frames reach handlers with the payload decoded and the
// mask cleared (RFC 6455 §5.1); a clean disconnect ends the stream
// without error.
#expect(
try await AsyncThrowingStream.makeClient([.fish.masked(), .chips.masked()]).collectAll() == [
.fish, .chips
]
)
#expect(
try await AsyncThrowingStream.makeClient([]).collectAll() == []
)
}

@Test
func clientSequence_RejectsUnmaskedFrames() async {
// RFC 6455 §5.1: an unmasked client frame fails the stream — the error
// must not be normalized away like a disconnect.
await #expect(throws: WSFrameEncoder.Error.self) {
try await AsyncThrowingStream.makeClient([.fish]).collectAll()
}
}

@Test
func protocolFrames_CatchErrors_AndCloseStream() async throws {
let (stream, continuation) = AsyncThrowingStream<WSFrame, any Error>.makeStream()
Expand Down Expand Up @@ -80,6 +109,11 @@ extension AsyncThrowingStream where Element == WSFrame, Failure == any Error {
let bytes = ConsumingAsyncSequence(frames.flatMap(WSFrameEncoder.encodeFrame))
return AsyncThrowingStream.decodingFrames(from: bytes)
}

static func makeClient(_ frames: [WSFrame]) -> Self {
let bytes = ConsumingAsyncSequence(frames.flatMap(WSFrameEncoder.encodeFrame))
return AsyncThrowingStream.decodingClientFrames(from: bytes)
}
}

extension AsyncSequence {
Expand Down
41 changes: 41 additions & 0 deletions FlyingFox/Tests/WebSocket/WSFrameEncoderTests.swift
Original file line number Diff line number Diff line change
Expand Up @@ -152,6 +152,43 @@ struct WSFrameEncoderTests {
)
}

@Test
func decodeFrame_PreservesMask() async throws {
// RFC 6455 §5.1 — the mask is preserved so servers can detect unmasked
// client frames; the payload is stored unmasked.
// "Abc" masked with 0x1 0x2 0x3 0x4 → 0x40 0x60 0x60.
#expect(
try await WSFrameEncoder.decodeFrame(0b10000001, 0b10000011, 0x1, 0x2, 0x3, 0x4, 0x40, 0x60, 0x60) == .make(
fin: true,
opcode: .text,
mask: .mock,
payload: "Abc".data(using: .utf8)!
)
)
}

@Test
func decodeClientFrame_ThrowsWhenUnmasked() async {
// RFC 6455 §5.1: "The server MUST close the connection upon receiving
// a frame that is not masked."
await #expect(throws: WSFrameEncoder.Error.self) {
try await WSFrameEncoder.decodeClientFrame(0b10000001, 3, .ascii("A"), .ascii("b"), .ascii("c"))
}
}

@Test
func decodeClientFrame_ClearsMask() async throws {
// Handlers observe unmasked frames so they can safely echo them —
// "A server MUST NOT mask any frames that it sends to the client." (§5.1)
#expect(
try await WSFrameEncoder.decodeClientFrame(0b10000001, 0b10000011, 0x1, 0x2, 0x3, 0x4, 0x40, 0x60, 0x60) == .make(
fin: true,
opcode: .text,
payload: "Abc".data(using: .utf8)!
)
)
}

@Test
func decodeFrame0() {
#expect(
Expand Down Expand Up @@ -418,6 +455,10 @@ private extension WSFrameEncoder {
try await decodeFrame(from: ConsumingAsyncSequence(bytes))
}

static func decodeClientFrame(_ bytes: UInt8...) async throws -> WSFrame {
try await decodeClientFrame(from: ConsumingAsyncSequence(bytes))
}

static func decodeLength(_ bytes: UInt8...) async throws -> Int {
try await decodeLengthMask(bytes).length
}
Expand Down
8 changes: 8 additions & 0 deletions FlyingFox/Tests/WebSocket/WSFrameTests.swift
Original file line number Diff line number Diff line change
Expand Up @@ -128,6 +128,14 @@ extension WSFrame {
payload: text.data(using: .utf8)!)
}

// Copy of the frame carrying a client masking key, as sent client → server.
// RFC 6455 §5.1: "a client MUST mask all frames that it sends to the server."
func masked(_ mask: Mask = .mock) -> Self {
var frame = self
frame.mask = mask
return frame
}

static func makeTextFrames(_ payload: String, maxCharacters: Int) -> [WSFrame] {
var messages = payload.chunked(size: maxCharacters).enumerated().map { idx, substring in
WSFrame.make(fin: false, isContinuation: idx != 0, text: String(substring))
Expand Down
Loading