Skip to content

Commit 10bf57b

Browse files
committed
fix: add automatic checking to runtime_cal
1 parent 133a7ec commit 10bf57b

7 files changed

Lines changed: 198 additions & 73 deletions

File tree

doc/releases/changelog-dev.md

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -21,6 +21,9 @@
2121

2222
<h3>Improvements 🛠</h3>
2323

24+
* Add checking of the pointer and status code from `runtime_call`.
25+
[(#XXXX)](https://github.com/PennyLaneAI/catalyst/pull/XXXX)
26+
2427
* a PennyLane `Backline` is serialized to the `catalyst.backline` module attribute and compiled
2528
through the transport passes.
2629
[(#3068)](https://github.com/PennyLaneAI/catalyst/pull/3068)

mlir/lib/Transport/Transforms/TransportToLLVM.cpp

Lines changed: 51 additions & 27 deletions
Original file line numberDiff line numberDiff line change
@@ -88,6 +88,25 @@ Value globalStr(ConversionPatternRewriter &rewriter, Location loc, ModuleOp mod,
8888
ArrayRef<LLVM::GEPArg>{0, 0}, LLVM::GEPNoWrapFlags::inbounds);
8989
}
9090

91+
void emitCheckedStatus(ConversionPatternRewriter &rewriter, Location loc, ModuleOp mod,
92+
StringRef name, ArrayRef<Type> paramTys, ValueRange args, StringRef what) {
93+
auto *ctx = rewriter.getContext();
94+
Value rc = emitCall(rewriter, loc, mod, name, paramTys, i32Ty(ctx), args);
95+
Value msg = globalStr(rewriter, loc, mod, "transport_check_", what);
96+
emitCall(rewriter, loc, mod, "__catalyst__transport__check", {i32Ty(ctx), ptrTy(ctx)}, Type(),
97+
{rc, msg});
98+
}
99+
100+
Value emitCheckedSession(ConversionPatternRewriter &rewriter, Location loc, ModuleOp mod,
101+
StringRef name, ArrayRef<Type> paramTys, ValueRange args, StringRef what) {
102+
auto *ctx = rewriter.getContext();
103+
Value s = emitCall(rewriter, loc, mod, name, paramTys, ptrTy(ctx), args);
104+
Value msg = globalStr(rewriter, loc, mod, "transport_session_", what);
105+
emitCall(rewriter, loc, mod, "__catalyst__transport__check_session", {ptrTy(ctx), ptrTy(ctx)},
106+
Type(), {s, msg});
107+
return s;
108+
}
109+
91110
Value constInt(ConversionPatternRewriter &rewriter, Location loc, Type ty, int64_t v) {
92111
return LLVM::ConstantOp::create(rewriter, loc, ty, rewriter.getIntegerAttr(ty, v));
93112
}
@@ -121,9 +140,9 @@ struct CreateLowering : public OpConversionPattern<CreateOp> {
121140
Value key = globalStr(rewriter, op.getLoc(), mod, "transport_key_", op.getKey());
122141
Value role =
123142
constInt(rewriter, op.getLoc(), i32Ty(ctx), static_cast<int64_t>(sessTy.getRole()));
124-
Value s = emitCall(rewriter, op.getLoc(), mod, "__catalyst__transport__create",
125-
{ptrTy(ctx), ptrTy(ctx), i32Ty(ctx), ptrTy(ctx)}, ptrTy(ctx),
126-
{lib, cfg, role, key});
143+
Value s = emitCheckedSession(rewriter, op.getLoc(), mod, "__catalyst__transport__create",
144+
{ptrTy(ctx), ptrTy(ctx), i32Ty(ctx), ptrTy(ctx)},
145+
{lib, cfg, role, key}, "create");
127146
rewriter.replaceOp(op, s);
128147
return success();
129148
}
@@ -143,9 +162,9 @@ template <typename OpT, bool Async> struct ConnectLoweringBase : public OpConver
143162
{adaptor.getSession(), peer, port});
144163
rewriter.replaceOp(op, r);
145164
} else {
146-
emitCall(rewriter, op.getLoc(), mod, "__catalyst__transport__connect",
147-
{ptrTy(ctx), ptrTy(ctx), IntegerType::get(ctx, 16)}, i32Ty(ctx),
148-
{adaptor.getSession(), peer, port});
165+
emitCheckedStatus(rewriter, op.getLoc(), mod, "__catalyst__transport__connect",
166+
{ptrTy(ctx), ptrTy(ctx), IntegerType::get(ctx, 16)},
167+
{adaptor.getSession(), peer, port}, "connect");
149168
rewriter.eraseOp(op);
150169
}
151170
return success();
@@ -167,8 +186,8 @@ struct ExchangeKeysLoweringBase : public OpConversionPattern<OpT> {
167186
{ptrTy(ctx)}, i64Ty(ctx), {adaptor.getSession()});
168187
rewriter.replaceOp(op, r);
169188
} else {
170-
emitCall(rewriter, op.getLoc(), mod, "__catalyst__transport__exchange_keys",
171-
{ptrTy(ctx)}, i32Ty(ctx), {adaptor.getSession()});
189+
emitCheckedStatus(rewriter, op.getLoc(), mod, "__catalyst__transport__exchange_keys",
190+
{ptrTy(ctx)}, {adaptor.getSession()}, "exchange_keys");
172191
rewriter.eraseOp(op);
173192
}
174193
return success();
@@ -182,8 +201,8 @@ struct AwaitLowering : public OpConversionPattern<AwaitOp> {
182201
LogicalResult matchAndRewrite(AwaitOp op, OpAdaptor adaptor,
183202
ConversionPatternRewriter &rewriter) const override {
184203
auto *ctx = op.getContext();
185-
emitCall(rewriter, op.getLoc(), moduleOf(op), "__catalyst__transport__await", {i64Ty(ctx)},
186-
i32Ty(ctx), {adaptor.getToken()});
204+
emitCheckedStatus(rewriter, op.getLoc(), moduleOf(op), "__catalyst__transport__await",
205+
{i64Ty(ctx)}, {adaptor.getToken()}, "await");
187206
rewriter.eraseOp(op);
188207
return success();
189208
}
@@ -196,8 +215,9 @@ struct EstablishChannelLowering : public OpConversionPattern<EstablishChannelOp>
196215
auto *ctx = op.getContext();
197216
Value transport =
198217
globalStr(rewriter, op.getLoc(), moduleOf(op), "transport_kind_", op.getTransport());
199-
emitCall(rewriter, op.getLoc(), moduleOf(op), "__catalyst__transport__establish_channel",
200-
{ptrTy(ctx), ptrTy(ctx)}, i32Ty(ctx), {adaptor.getSession(), transport});
218+
emitCheckedStatus(rewriter, op.getLoc(), moduleOf(op),
219+
"__catalyst__transport__establish_channel", {ptrTy(ctx), ptrTy(ctx)},
220+
{adaptor.getSession(), transport}, "establish_channel");
201221
rewriter.eraseOp(op);
202222
return success();
203223
}
@@ -210,8 +230,9 @@ struct SetCoprocessorFnLowering : public OpConversionPattern<SetCoprocessorFnOp>
210230
auto *ctx = op.getContext();
211231
ModuleOp mod = moduleOf(op);
212232
Value sym = globalStr(rewriter, op.getLoc(), mod, "transport_coproc_fn_", op.getSymbol());
213-
emitCall(rewriter, op.getLoc(), mod, "__catalyst__transport__set_coprocessor_fn",
214-
{ptrTy(ctx), ptrTy(ctx)}, i32Ty(ctx), {adaptor.getSession(), sym});
233+
emitCheckedStatus(rewriter, op.getLoc(), mod, "__catalyst__transport__set_coprocessor_fn",
234+
{ptrTy(ctx), ptrTy(ctx)}, {adaptor.getSession(), sym},
235+
"set_coprocessor_fn");
215236
rewriter.eraseOp(op);
216237
return success();
217238
}
@@ -225,9 +246,10 @@ struct SetMessageSizesLowering : public OpConversionPattern<SetMessageSizesOp> {
225246
Value idx = constInt(rewriter, op.getLoc(), i32Ty(ctx), op.getWorkItemIdx());
226247
Value inB = constInt(rewriter, op.getLoc(), i64Ty(ctx), op.getInBytes());
227248
Value outB = constInt(rewriter, op.getLoc(), i64Ty(ctx), op.getOutBytes());
228-
emitCall(rewriter, op.getLoc(), moduleOf(op), "__catalyst__transport__set_message_sizes",
229-
{ptrTy(ctx), i32Ty(ctx), i64Ty(ctx), i64Ty(ctx)}, i32Ty(ctx),
230-
{adaptor.getSession(), idx, inB, outB});
249+
emitCheckedStatus(rewriter, op.getLoc(), moduleOf(op),
250+
"__catalyst__transport__set_message_sizes",
251+
{ptrTy(ctx), i32Ty(ctx), i64Ty(ctx), i64Ty(ctx)},
252+
{adaptor.getSession(), idx, inB, outB}, "set_message_sizes");
231253
rewriter.eraseOp(op);
232254
return success();
233255
}
@@ -281,9 +303,10 @@ struct StagePayloadLowering : public OpConversionPattern<StagePayloadOp> {
281303
auto [srcPtr, bytes] =
282304
memrefPtrAndBytes(rewriter, op.getLoc(), adaptor.getPayload(), memTy);
283305
Value decoderId = constInt(rewriter, op.getLoc(), i32Ty(ctx), op.getDecoderId());
284-
emitCall(rewriter, op.getLoc(), moduleOf(op), "__catalyst__transport__stage_payload",
285-
{ptrTy(ctx), ptrTy(ctx), i64Ty(ctx), i32Ty(ctx)}, i32Ty(ctx),
286-
{adaptor.getSession(), srcPtr, bytes, decoderId});
306+
emitCheckedStatus(rewriter, op.getLoc(), moduleOf(op),
307+
"__catalyst__transport__stage_payload",
308+
{ptrTy(ctx), ptrTy(ctx), i64Ty(ctx), i32Ty(ctx)},
309+
{adaptor.getSession(), srcPtr, bytes, decoderId}, "stage_payload");
287310
rewriter.eraseOp(op);
288311
return success();
289312
}
@@ -295,8 +318,8 @@ struct PostLowering : public OpConversionPattern<PostOp> {
295318
ConversionPatternRewriter &rewriter) const override {
296319
auto *ctx = op.getContext();
297320
Value idx = constInt(rewriter, op.getLoc(), i32Ty(ctx), op.getWorkItemIdx());
298-
emitCall(rewriter, op.getLoc(), moduleOf(op), "__catalyst__transport__post",
299-
{ptrTy(ctx), i32Ty(ctx)}, i32Ty(ctx), {adaptor.getSession(), idx});
321+
emitCheckedStatus(rewriter, op.getLoc(), moduleOf(op), "__catalyst__transport__post",
322+
{ptrTy(ctx), i32Ty(ctx)}, {adaptor.getSession(), idx}, "post");
300323
rewriter.eraseOp(op);
301324
return success();
302325
}
@@ -319,9 +342,9 @@ struct CollectLowering : public OpConversionPattern<CollectOp> {
319342
return rewriter.notifyMatchFailure(op, "collect dest must have identity layout");
320343
}
321344
auto [dstPtr, bytes] = memrefPtrAndBytes(rewriter, op.getLoc(), adaptor.getDest(), memTy);
322-
emitCall(rewriter, op.getLoc(), mod, "__catalyst__transport__collect",
323-
{ptrTy(ctx), ptrTy(ctx), i64Ty(ctx)}, i32Ty(ctx),
324-
{adaptor.getSession(), dstPtr, bytes});
345+
emitCheckedStatus(rewriter, op.getLoc(), mod, "__catalyst__transport__collect",
346+
{ptrTy(ctx), ptrTy(ctx), i64Ty(ctx)},
347+
{adaptor.getSession(), dstPtr, bytes}, "collect");
325348
rewriter.eraseOp(op);
326349
return success();
327350
}
@@ -364,8 +387,9 @@ struct GetSessionLowering : public OpConversionPattern<GetSessionOp> {
364387
Value role =
365388
constInt(rewriter, op.getLoc(), i32Ty(ctx), static_cast<int64_t>(sessTy.getRole()));
366389
Value key = globalStr(rewriter, op.getLoc(), mod, "transport_key_", op.getKey());
367-
Value s = emitCall(rewriter, op.getLoc(), mod, "__catalyst__transport__get_session",
368-
{i32Ty(ctx), ptrTy(ctx)}, ptrTy(ctx), {role, key});
390+
Value s =
391+
emitCheckedSession(rewriter, op.getLoc(), mod, "__catalyst__transport__get_session",
392+
{i32Ty(ctx), ptrTy(ctx)}, {role, key}, "get_session");
369393
rewriter.replaceOp(op, s);
370394
return success();
371395
}

mlir/test/Transport/ConvertTransportToLLVM.mlir

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -15,6 +15,8 @@
1515
// RUN: quantum-opt %s --convert-transport-to-llvm --split-input-file | FileCheck %s
1616

1717
// CHECK-DAG: llvm.func @__catalyst__transport__create(!llvm.ptr, !llvm.ptr, i32, !llvm.ptr) -> !llvm.ptr
18+
// CHECK-DAG: llvm.func @__catalyst__transport__check_session(!llvm.ptr, !llvm.ptr)
19+
// CHECK-DAG: llvm.func @__catalyst__transport__check(i32, !llvm.ptr)
1820
// CHECK-DAG: llvm.func @__catalyst__transport__connect(!llvm.ptr, !llvm.ptr, i16) -> i32
1921
// CHECK-DAG: llvm.func @__catalyst__transport__exchange_keys(!llvm.ptr) -> i32
2022
// CHECK-DAG: llvm.func @__catalyst__transport__establish_channel(!llvm.ptr, !llvm.ptr) -> i32
@@ -30,22 +32,30 @@
3032
// CHECK-LABEL: func.func @controller
3133
func.func @controller(%syndrome: memref<?xi8>, %correction: memref<?xi8>) {
3234
// CHECK: %[[S:.*]] = llvm.call @__catalyst__transport__create({{.*}}) : (!llvm.ptr, !llvm.ptr, i32, !llvm.ptr) -> !llvm.ptr
35+
// CHECK: llvm.call @__catalyst__transport__check_session(%[[S]]
3336
%s = transport.create {backend_lib = "libbackend.so", config = "cfg"} -> !transport.session<controller>
3437
// CHECK: llvm.call @__catalyst__transport__connect(%[[S]]
38+
// CHECK: llvm.call @__catalyst__transport__check(
3539
transport.connect %s {peer = "127.0.0.1", oob_port = 18560 : i32} : !transport.session<controller>
3640
// CHECK: llvm.call @__catalyst__transport__exchange_keys(%[[S]])
41+
// CHECK: llvm.call @__catalyst__transport__check(
3742
transport.exchange_keys %s : !transport.session<controller>
3843
// CHECK: llvm.call @__catalyst__transport__establish_channel(%[[S]]
44+
// CHECK: llvm.call @__catalyst__transport__check(
3945
transport.establish_channel %s "rdma" : !transport.session<controller>
4046
// CHECK: llvm.call @__catalyst__transport__set_message_sizes(%[[S]]
47+
// CHECK: llvm.call @__catalyst__transport__check(
4148
transport.set_message_sizes %s {in_bytes = 8 : i64, out_bytes = 8 : i64} : !transport.session<controller>
4249
// CHECK: llvm.call @__catalyst__transport__start(%[[S]])
4350
transport.start %s : !transport.session<controller>
4451
// CHECK: llvm.call @__catalyst__transport__stage_payload(%[[S]]
52+
// CHECK: llvm.call @__catalyst__transport__check(
4553
// CHECK: llvm.call @__catalyst__transport__post(%[[S]]
54+
// CHECK: llvm.call @__catalyst__transport__check(
4655
transport.stage_payload %s, %syndrome : !transport.session<controller>, memref<?xi8>
4756
transport.post %s : !transport.session<controller>
4857
// CHECK: llvm.call @__catalyst__transport__collect(%[[S]]
58+
// CHECK: llvm.call @__catalyst__transport__check(
4959
transport.collect %s, %correction : !transport.session<controller>, memref<?xi8>
5060
// CHECK: llvm.call @__catalyst__transport__stop(%[[S]])
5161
transport.stop %s : !transport.session<controller>
@@ -60,12 +70,15 @@ func.func @controller(%syndrome: memref<?xi8>, %correction: memref<?xi8>) {
6070
// CHECK-LABEL: func.func @coprocessor
6171
func.func @coprocessor() {
6272
// CHECK: %[[C:.*]] = llvm.call @__catalyst__transport__create({{.*}}) : (!llvm.ptr, !llvm.ptr, i32, !llvm.ptr) -> !llvm.ptr
73+
// CHECK: llvm.call @__catalyst__transport__check_session(%[[C]]
6374
%c = transport.create {backend_lib = "libbackend.so", config = "cfg"} -> !transport.session<coprocessor>
6475
// CHECK: llvm.call @__catalyst__transport__connect_async(%[[C]]
6576
%t = transport.connect_async %c {peer = "127.0.0.1", oob_port = 18560 : i32} : !transport.session<coprocessor> -> !transport.token
6677
// CHECK: llvm.call @__catalyst__transport__await
78+
// CHECK: llvm.call @__catalyst__transport__check(
6779
transport.await %t : !transport.token
6880
// CHECK: llvm.call @__catalyst__transport__set_coprocessor_fn(%[[C]], %{{.*}}) : (!llvm.ptr, !llvm.ptr) -> i32
81+
// CHECK: llvm.call @__catalyst__transport__check(
6982
transport.set_coprocessor_fn %c {symbol = "foo"} : !transport.session<coprocessor>
7083
// CHECK: llvm.call @__catalyst__transport__destroy(%[[C]])
7184
transport.destroy %c : !transport.session<coprocessor>
@@ -80,11 +93,13 @@ func.func @coprocessor() {
8093
func.func @resolve(%syndrome: memref<?xi8>, %correction: memref<?xi8>) {
8194
// CHECK: %[[R:.*]] = llvm.mlir.constant(0 : i32) : i32
8295
// CHECK: %[[S:.*]] = llvm.call @__catalyst__transport__get_session(%[[R]], {{.*}}) : (i32, !llvm.ptr) -> !llvm.ptr
96+
// CHECK: llvm.call @__catalyst__transport__check_session(%[[S]]
8397
%s = transport.get_session {key = "cop0"} : !transport.session<controller>
8498
// CHECK: llvm.call @__catalyst__transport__post(%[[S]]
8599
transport.stage_payload %s, %syndrome : !transport.session<controller>, memref<?xi8>
86100
transport.post %s : !transport.session<controller>
87101
// CHECK: llvm.call @__catalyst__transport__collect(%[[S]]
102+
// CHECK: llvm.call @__catalyst__transport__check(
88103
transport.collect %s, %correction : !transport.session<controller>, memref<?xi8>
89104
return
90105
}

runtime/include/TransportCAPI.h

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -46,6 +46,13 @@ enum {
4646
CATALYST_TRANSPORT_ROLE_COPROCESSOR = 1,
4747
};
4848

49+
// Abort if `rc` is not CATALYST_TRANSPORT_OK. Compiled-program adapters call this so a failed
50+
// round cannot be mistaken for success. The C functions themselves still return `rc`.
51+
void __catalyst__transport__check(int rc, const char *what);
52+
53+
// Abort if `s` is `nullptr`. Compiled-program adapters call this after create / get_session.
54+
void __catalyst__transport__check_session(CatalystTransportSession *s, const char *what);
55+
4956
// Create a session from a named backend plugin `.so` (dlopen'd by the runtime). `role` selects
5057
// which factory symbol is looked up (controller vs coprocessor). `config` is the backend's string.
5158
// Returns NULL on failure.

runtime/lib/transport/TransportCAPI.cpp

Lines changed: 21 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -30,6 +30,7 @@
3030
#include <vector>
3131

3232
#include "DynamicLibraryLoader.hpp"
33+
#include "Exception.hpp"
3334
#include "Transport.hpp"
3435
#include "TransportBackend.h"
3536
#include "WireProtocol.hpp"
@@ -220,6 +221,24 @@ void *resolve_coprocessor_fn_symbol(CatalystTransportSession *s, const char *sym
220221

221222
extern "C" {
222223

224+
void __catalyst__transport__check(int rc, const char *what) {
225+
if (rc == ::CATALYST_TRANSPORT_OK) {
226+
return;
227+
}
228+
const std::string msg = std::string{"[transport] "} + ((what != nullptr) ? what : "call") +
229+
" failed (" + collect_error_name(rc) + ")";
230+
RT_FAIL(msg.c_str());
231+
}
232+
233+
void __catalyst__transport__check_session(CatalystTransportSession *s, const char *what) {
234+
if (s != nullptr) {
235+
return;
236+
}
237+
const std::string msg =
238+
std::string{"[transport] "} + ((what != nullptr) ? what : "session") + ": null session";
239+
RT_FAIL(msg.c_str());
240+
}
241+
223242
CatalystTransportSession *__catalyst__transport__create(const char *backend_lib, const char *config,
224243
std::int32_t role, const char *key) {
225244
try {
@@ -433,8 +452,8 @@ int __catalyst__transport__collect(CatalystTransportSession *s, void *reply,
433452
return s->sess->collect(replies, replies_bytes, 1);
434453
});
435454
if (rc != CATALYST_TRANSPORT_OK) {
436-
// The generated code discards this return value, so a failed round is otherwise silent:
437-
// `reply` keeps whatever it held, and the caller consumes that as a valid result.
455+
// C callers see the return code. Compiled adapters call __catalyst__transport__check so a
456+
// failed round cannot be consumed as a valid reply. Log here so the C path is not silent.
438457
std::cerr << "[transport] collect failed (rc=" << rc << ": " << collect_error_name(rc)
439458
<< "); the reply buffer was not written\n";
440459
}

0 commit comments

Comments
 (0)