Skip to content

Commit 8bbc15d

Browse files
Merge pull request #13 from GraphQLSwift/experiment/parallel-execution
GraphQL executions are handled in parallel
2 parents 1fb0dbd + 8dd79b2 commit 8bbc15d

5 files changed

Lines changed: 141 additions & 75 deletions

File tree

Sources/GraphQLTransportWS/Client.swift

Lines changed: 22 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -50,43 +50,52 @@ public actor Client<InitPayload: Equatable & Codable> {
5050
do {
5151
response = try decoder.decode(Response.self, from: message)
5252
} catch {
53-
try await self.error(.noType())
53+
try await messenger.error(.noType())
5454
return
5555
}
5656

5757
switch response.type {
5858
case .connectionAck:
59-
guard
60-
let connectionAckResponse = try? decoder.decode(
59+
let connectionAckResponse: ConnectionAckResponse
60+
do {
61+
connectionAckResponse = try decoder.decode(
6162
ConnectionAckResponse.self,
6263
from: message
6364
)
64-
else {
65-
try await error(.invalidResponseFormat(messageType: .connectionAck))
65+
} catch {
66+
try await messenger.error(.invalidResponseFormat(messageType: .connectionAck, error: error))
6667
return
6768
}
6869
try await onConnectionAck(connectionAckResponse, self)
6970
case .next:
70-
guard let nextResponse = try? decoder.decode(NextResponse.self, from: message) else {
71-
try await error(.invalidResponseFormat(messageType: .next))
71+
let nextResponse: NextResponse
72+
do {
73+
nextResponse = try decoder.decode(NextResponse.self, from: message)
74+
} catch {
75+
try await messenger.error(.invalidResponseFormat(messageType: .next, error: error))
7276
return
7377
}
7478
try await onNext(nextResponse, self)
7579
case .error:
76-
guard let errorResponse = try? decoder.decode(ErrorResponse.self, from: message) else {
77-
try await error(.invalidResponseFormat(messageType: .error))
80+
let errorResponse: ErrorResponse
81+
do {
82+
errorResponse = try decoder.decode(ErrorResponse.self, from: message)
83+
} catch {
84+
try await messenger.error(.invalidResponseFormat(messageType: .error, error: error))
7885
return
7986
}
8087
try await onError(errorResponse, self)
8188
case .complete:
82-
guard let completeResponse = try? decoder.decode(CompleteResponse.self, from: message)
83-
else {
84-
try await error(.invalidResponseFormat(messageType: .complete))
89+
let completeResponse: CompleteResponse
90+
do {
91+
completeResponse = try decoder.decode(CompleteResponse.self, from: message)
92+
} catch {
93+
try await messenger.error(.invalidResponseFormat(messageType: .complete, error: error))
8594
return
8695
}
8796
try await onComplete(completeResponse, self)
8897
default:
89-
try await error(.invalidType())
98+
try await messenger.error(.invalidType())
9099
}
91100
}
92101

@@ -123,9 +132,4 @@ public actor Client<InitPayload: Equatable & Codable> {
123132
)
124133
)
125134
}
126-
127-
/// Send an error through the messenger and close the connection
128-
private func error(_ error: GraphQLTransportWSError) async throws {
129-
try await messenger.error(error.message, code: error.code.rawValue)
130-
}
131135
}

Sources/GraphQLTransportWS/GraphqlTransportWSError.swift

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -58,16 +58,16 @@ struct GraphQLTransportWSError: Error {
5858
)
5959
}
6060

61-
static func invalidRequestFormat(messageType: RequestMessageType) -> Self {
61+
static func invalidRequestFormat(messageType: RequestMessageType, error: Error) -> Self {
6262
return self.init(
63-
"Request message doesn't match '\(messageType.type.rawValue)' JSON format",
63+
"Request message doesn't match '\(messageType.type.rawValue)' JSON format: \(error)",
6464
code: .miscellaneous
6565
)
6666
}
6767

68-
static func invalidResponseFormat(messageType: ResponseMessageType) -> Self {
68+
static func invalidResponseFormat(messageType: ResponseMessageType, error: Error) -> Self {
6969
return self.init(
70-
"Response message doesn't match '\(messageType.type.rawValue)' JSON format",
70+
"Response message doesn't match '\(messageType.type.rawValue)' JSON format: \(error)",
7171
code: .miscellaneous
7272
)
7373
}

Sources/GraphQLTransportWS/Messenger.swift

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -15,3 +15,10 @@ public protocol Messenger: Sendable {
1515
/// - code: An error code
1616
func error(_ message: String, code: Int) async throws
1717
}
18+
19+
extension Messenger {
20+
/// Send an error through the messenger and close the connection
21+
func error(_ error: GraphQLTransportWSError) async throws {
22+
try await self.error(error.message, code: error.code.rawValue)
23+
}
24+
}

Sources/GraphQLTransportWS/Server.swift

Lines changed: 52 additions & 53 deletions
Original file line numberDiff line numberDiff line change
@@ -25,7 +25,7 @@ where
2525

2626
private var initialized = false
2727
private var initResult: InitPayloadResult?
28-
private var subscriptionTasks = [String: Task<Void, any Error>]()
28+
private var executionTasks = [String: Task<Void, any Error>]()
2929

3030
/// Create a new server
3131
///
@@ -53,7 +53,7 @@ where
5353
}
5454

5555
deinit {
56-
subscriptionTasks.values.forEach { $0.cancel() }
56+
executionTasks.values.forEach { $0.cancel() }
5757
}
5858

5959
/// Listen and react to the provided async sequence of client messages. This function will block until the stream is completed.
@@ -70,39 +70,44 @@ where
7070
do {
7171
request = try decoder.decode(Request.self, from: message)
7272
} catch {
73-
try await self.error(.noType())
73+
try await messenger.error(.noType())
7474
return
7575
}
7676

7777
// handle incoming message
7878
switch request.type {
7979
case .connectionInit:
80-
guard
81-
let connectionInitRequest = try? decoder.decode(
80+
let connectionInitRequest: ConnectionInitRequest<InitPayload>
81+
do {
82+
connectionInitRequest = try decoder.decode(
8283
ConnectionInitRequest<InitPayload>.self,
8384
from: message
8485
)
85-
else {
86-
try await error(.invalidRequestFormat(messageType: .connectionInit))
86+
} catch {
87+
try await messenger.error(.invalidRequestFormat(messageType: .connectionInit, error: error))
8788
return
8889
}
8990
try await onConnectionInit(connectionInitRequest, messenger)
9091
case .subscribe:
91-
guard let subscribeRequest = try? decoder.decode(SubscribeRequest.self, from: message)
92-
else {
93-
try await error(.invalidRequestFormat(messageType: .subscribe))
92+
let subscribeRequest: SubscribeRequest
93+
do {
94+
subscribeRequest = try decoder.decode(SubscribeRequest.self, from: message)
95+
} catch {
96+
try await messenger.error(.invalidRequestFormat(messageType: .subscribe, error: error))
9497
return
9598
}
9699
try await onSubscribe(subscribeRequest)
97100
case .complete:
98-
guard let completeRequest = try? decoder.decode(CompleteRequest.self, from: message)
99-
else {
100-
try await error(.invalidRequestFormat(messageType: .complete))
101+
let completeRequest: CompleteRequest
102+
do {
103+
completeRequest = try decoder.decode(CompleteRequest.self, from: message)
104+
} catch {
105+
try await messenger.error(.invalidRequestFormat(messageType: .complete, error: error))
101106
return
102107
}
103108
try await onOperationComplete(completeRequest)
104109
default:
105-
try await error(.invalidType())
110+
try await messenger.error(.invalidType())
106111
}
107112
}
108113

@@ -111,14 +116,14 @@ where
111116
_: Messenger
112117
) async throws {
113118
guard !initialized else {
114-
try await error(.tooManyInitializations())
119+
try await messenger.error(.tooManyInitializations())
115120
return
116121
}
117122

118123
do {
119124
initResult = try await onInit(connectionInitRequest.payload)
120125
} catch {
121-
try await self.error(.forbidden())
126+
try await messenger.error(.forbidden())
122127
return
123128
}
124129
initialized = true
@@ -128,62 +133,66 @@ where
128133

129134
private func onSubscribe(_ subscribeRequest: SubscribeRequest) async throws {
130135
guard initialized, let initResult else {
131-
try await error(.notInitialized())
136+
try await messenger.error(.notInitialized())
132137
return
133138
}
134139

135140
let id = subscribeRequest.id
136-
if subscriptionTasks[id] != nil {
137-
try await error(.subscriberAlreadyExists(id: id))
138-
}
139-
140141
let graphQLRequest = subscribeRequest.payload
141142

142-
var isStreaming = false
143+
let isStreaming: Bool
143144
do {
144145
isStreaming = try graphQLRequest.isSubscription()
145146
} catch {
146147
try await sendError(error, id: id)
147148
return
148149
}
149150

150-
if isStreaming {
151-
subscriptionTasks[id] = Task {
151+
guard executionTasks[id] == nil else {
152+
try await messenger.error(.subscriberAlreadyExists(id: id))
153+
return
154+
}
155+
executionTasks[id] = Task {
156+
defer {
157+
executionTasks.removeValue(forKey: id)
158+
}
159+
160+
if isStreaming {
161+
let stream: SubscriptionSequenceType
152162
do {
153-
let stream = try await onSubscribe(graphQLRequest, initResult)
154-
for try await event in stream {
155-
try Task.checkCancellation()
156-
try await self.sendNext(event, id: id)
157-
}
163+
stream = try await onSubscribe(graphQLRequest, initResult)
158164
} catch {
159165
try await sendError(error, id: id)
160-
subscriptionTasks.removeValue(forKey: id)
161-
throw error
166+
return
167+
}
168+
for try await event in stream {
169+
try await self.sendNext(event, id: id)
170+
}
171+
executionTasks.removeValue(forKey: id)
172+
} else {
173+
let result: GraphQLResult
174+
do {
175+
result = try await onExecute(graphQLRequest, initResult)
176+
} catch {
177+
try await sendError(error, id: id)
178+
return
162179
}
163-
try await self.sendComplete(id: id)
164-
subscriptionTasks.removeValue(forKey: id)
165-
}
166-
} else {
167-
do {
168-
let result = try await onExecute(graphQLRequest, initResult)
169180
try await sendNext(result, id: id)
170-
try await sendComplete(id: id)
171-
} catch {
172-
try await sendError(error, id: id)
173181
}
182+
try await sendComplete(id: id)
174183
}
175184
}
176185

177186
private func onOperationComplete(_ completeRequest: CompleteRequest) async throws {
178187
guard initialized else {
179-
try await error(.notInitialized())
188+
try await messenger.error(.notInitialized())
180189
return
181190
}
182191

183192
let id = completeRequest.id
184-
if let task = subscriptionTasks[id] {
193+
if let task = executionTasks[id] {
185194
task.cancel()
186-
subscriptionTasks.removeValue(forKey: id)
195+
executionTasks.removeValue(forKey: id)
187196
}
188197
try await onOperationComplete(id)
189198
}
@@ -238,14 +247,4 @@ where
238247
private func sendError(_ error: Error, id: String) async throws {
239248
try await sendError([error], id: id)
240249
}
241-
242-
/// Send an `error` response through the messenger
243-
private func sendError(_ errorMessage: String, id: String) async throws {
244-
try await sendError(GraphQLError(message: errorMessage), id: id)
245-
}
246-
247-
/// Send an error through the messenger and close the connection
248-
private func error(_ error: GraphQLTransportWSError) async throws {
249-
try await messenger.error(error.message, code: error.code.rawValue)
250-
}
251250
}

Tests/GraphQLTransportWSTests/GraphQLTransportWSTests.swift

Lines changed: 56 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -268,6 +268,62 @@ struct GraphqlTransportWSTests {
268268
)
269269
}
270270

271+
/// Tests malformed requests include decoder details in the transport error
272+
@Test func malformedRequestIncludesDecodingDetails() async throws {
273+
let api = TestAPI()
274+
let context = TestContext()
275+
let server = Server<TokenInitPayload, Void, AsyncThrowingStream<GraphQLResult, Error>>(
276+
messenger: serverMessenger,
277+
onInit: { _ in },
278+
onExecute: { graphQLRequest, _ in
279+
try await api.execute(
280+
request: graphQLRequest.query,
281+
context: context
282+
)
283+
},
284+
onSubscribe: { graphQLRequest, _ in
285+
try await api.subscribe(
286+
request: graphQLRequest.query,
287+
context: context
288+
).get()
289+
}
290+
)
291+
let (incoming, continuation) = AsyncThrowingStream<Data, any Error>.makeStream()
292+
293+
continuation.yield(Data(#"{"type":"complete"}"#.utf8))
294+
continuation.finish()
295+
296+
try await server.listen(to: incoming)
297+
298+
let error = await #expect(throws: TestMessengerError.self) {
299+
for try await _ in serverMessenger.stream {}
300+
}
301+
#expect(error?.code == 4400)
302+
#expect(error?.message.contains("Request message doesn't match 'complete' JSON format") == true)
303+
#expect(error?.message.contains("keyNotFound") == true)
304+
#expect(error?.message.contains(#""id""#) == true)
305+
}
306+
307+
/// Tests malformed responses include decoder details in the transport error
308+
@Test func malformedResponseIncludesDecodingDetails() async throws {
309+
let messenger = TestMessenger()
310+
let client = Client<TokenInitPayload>(messenger: messenger)
311+
let (incoming, continuation) = AsyncThrowingStream<Data, any Error>.makeStream()
312+
313+
continuation.yield(Data(#"{"type":"next"}"#.utf8))
314+
continuation.finish()
315+
316+
try await client.listen(to: incoming)
317+
318+
let error = await #expect(throws: TestMessengerError.self) {
319+
for try await _ in messenger.stream {}
320+
}
321+
#expect(error?.code == 4400)
322+
#expect(error?.message.contains("Response message doesn't match 'next' JSON format") == true)
323+
#expect(error?.message.contains("keyNotFound") == true)
324+
#expect(error?.message.contains(#""id""#) == true)
325+
}
326+
271327
enum TestError: Error {
272328
case couldBeAnything
273329
}

0 commit comments

Comments
 (0)