diff --git a/compiler/channel.go b/compiler/channel.go index 82139de8ba..59b34cc69e 100644 --- a/compiler/channel.go +++ b/compiler/channel.go @@ -104,6 +104,108 @@ func (b *builder) createChanClose(ch llvm.Value) { b.createRuntimeInvoke("chanClose", []llvm.Value{ch}, "") } +// createNonBlockingSelect emits IR for a non-blocking select with one channel case. +func (b *builder) createNonBlockingSelect(expr *ssa.Select) llvm.Value { + state := expr.States[0] + ch := b.getValue(state.Chan, state.Pos) + + resultType := b.ctx.StructType([]llvm.Type{ + b.ctx.Int32Type(), + b.ctx.Int1Type(), + }, false) + + var selected llvm.Value + selectOk := llvm.ConstInt(b.ctx.Int1Type(), 1, false) + + switch state.Dir { + case types.SendOnly: + sendValue := b.getValue(state.Send, state.Pos) + valueType := b.getLLVMType(state.Send.Type()) + isZeroSize := b.targetData.TypeAllocSize(valueType) == 0 + + valuePtr := llvm.ConstNull(b.dataPtrType) + var valueAlloca, valueAllocaSize llvm.Value + + if !isZeroSize { + valueAlloca, valueAllocaSize = b.createTemporaryAlloca( + valueType, + "select.send.value", + ) + b.CreateStore(sendValue, valueAlloca) + valuePtr = valueAlloca + } + + selected = b.createRuntimeCall( + "chanTrySend", + []llvm.Value{ch, valuePtr}, + "select.sent", + ) + + if !isZeroSize { + b.emitLifetimeEnd(valueAlloca, valueAllocaSize) + } + + case types.RecvOnly: + valueType := b.getLLVMType( + state.Chan.Type().Underlying().(*types.Chan).Elem(), + ) + isZeroSize := b.targetData.TypeAllocSize(valueType) == 0 + + // getChanSelectResult loads the received value later. + recvbuf := llvm.Undef(b.dataPtrType) + runtimeRecvbuf := llvm.ConstNull(b.dataPtrType) + + if !isZeroSize { + recvbuf, _ = b.createTemporaryAlloca( + valueType, + "select.recvbuf", + ) + runtimeRecvbuf = recvbuf + } + results := b.createRuntimeCall( + "chanTryRecv", + []llvm.Value{ch, runtimeRecvbuf}, + "select.recv", + ) + + selected = b.CreateExtractValue(results, 0, "select.received") + recvOk := b.CreateExtractValue(results, 1, "select.recv.ok") + selectOk = b.CreateSelect( + selected, + recvOk, + llvm.ConstInt(b.ctx.Int1Type(), 1, false), + "select.ok", + ) + + if b.selectRecvBuf == nil { + b.selectRecvBuf = make(map[*ssa.Select]llvm.Value) + } + b.selectRecvBuf[expr] = recvbuf + + default: + panic("unreachable") + } + + selectedIndex := llvm.ConstInt(b.ctx.Int32Type(), 0, false) + defaultIndex := llvm.ConstInt( + b.ctx.Int32Type(), + math.MaxUint32, + false, + ) + + selectIndex := b.CreateSelect( + selected, + selectedIndex, + defaultIndex, + "select.index", + ) + + result := llvm.Undef(resultType) + result = b.CreateInsertValue(result, selectIndex, 0, "") + result = b.CreateInsertValue(result, selectOk, 1, "") + return result +} + // createSelect emits all IR necessary for a select statements. That's a // non-trivial amount of code because select is very complex to implement. func (b *builder) createSelect(expr *ssa.Select) llvm.Value { @@ -126,6 +228,10 @@ func (b *builder) createSelect(expr *ssa.Select) llvm.Value { } } + if !expr.Blocking && len(expr.States) == 1 { + return b.createNonBlockingSelect(expr) + } + const maxSelectStates = math.MaxUint32 >> 2 if len(expr.States) > maxSelectStates { // The runtime code assumes that the number of state must fit in 30 bits diff --git a/compiler/testdata/channel.go b/compiler/testdata/channel.go index ecc837a71a..564185dd9e 100644 --- a/compiler/testdata/channel.go +++ b/compiler/testdata/channel.go @@ -23,3 +23,39 @@ func selectZeroRecv(ch1 chan int, ch2 chan struct{}) { default: } } + +func selectNonBlockingSend(ch chan int, value int) bool { + select { + case ch <- value: + return true + default: + return false + } +} + +func selectNonBlockingRecv(ch chan int) (int, bool, bool) { + select { + case value, ok := <-ch: + return value, ok, true + default: + return 0, false, false + } +} + +func selectNonBlockingZeroSend(ch chan struct{}) bool { + select { + case ch <- struct{}{}: + return true + default: + return false + } +} + +func selectNonBlockingZeroRecv(ch chan struct{}) (bool, bool) { + select { + case _, ok := <-ch: + return ok, true + default: + return false, false + } +} diff --git a/compiler/testdata/channel.ll b/compiler/testdata/channel.ll index bd91a57736..cd27246eb0 100644 --- a/compiler/testdata/channel.ll +++ b/compiler/testdata/channel.ll @@ -6,7 +6,7 @@ target triple = "wasm32-unknown-wasi" %runtime.channelOp = type { ptr, ptr, i32, ptr } %runtime.chanSelectState = type { ptr, ptr } -declare void @runtime.trackPointer(ptr nocapture readonly, ptr, ptr) #0 +declare void @runtime.trackPointer(ptr readonly captures(none), ptr, ptr) #0 ; Function Attrs: nounwind define hidden void @main.init(ptr %context) unnamed_addr #1 { @@ -19,33 +19,33 @@ define hidden void @main.chanIntSend(ptr dereferenceable_or_null(36) %ch, ptr %c entry: %chan.op = alloca %runtime.channelOp, align 8 %chan.value = alloca i32, align 4 - call void @llvm.lifetime.start.p0(i64 4, ptr nonnull %chan.value) + call void @llvm.lifetime.start.p0(ptr nonnull %chan.value) store i32 3, ptr %chan.value, align 4 - call void @llvm.lifetime.start.p0(i64 16, ptr nonnull %chan.op) + call void @llvm.lifetime.start.p0(ptr nonnull %chan.op) call void @runtime.chanSend(ptr %ch, ptr nonnull %chan.value, ptr nonnull %chan.op, ptr undef) #3 - call void @llvm.lifetime.end.p0(i64 16, ptr nonnull %chan.op) - call void @llvm.lifetime.end.p0(i64 4, ptr nonnull %chan.value) + call void @llvm.lifetime.end.p0(ptr nonnull %chan.op) + call void @llvm.lifetime.end.p0(ptr nonnull %chan.value) ret void } ; Function Attrs: nocallback nofree nosync nounwind willreturn memory(argmem: readwrite) -declare void @llvm.lifetime.start.p0(i64 immarg, ptr nocapture) #2 +declare void @llvm.lifetime.start.p0(ptr captures(none)) #2 declare void @runtime.chanSend(ptr dereferenceable_or_null(36), ptr, ptr dereferenceable_or_null(16), ptr) #0 ; Function Attrs: nocallback nofree nosync nounwind willreturn memory(argmem: readwrite) -declare void @llvm.lifetime.end.p0(i64 immarg, ptr nocapture) #2 +declare void @llvm.lifetime.end.p0(ptr captures(none)) #2 ; Function Attrs: nounwind define hidden void @main.chanIntRecv(ptr dereferenceable_or_null(36) %ch, ptr %context) unnamed_addr #1 { entry: %chan.op = alloca %runtime.channelOp, align 8 %chan.value = alloca i32, align 4 - call void @llvm.lifetime.start.p0(i64 4, ptr nonnull %chan.value) - call void @llvm.lifetime.start.p0(i64 16, ptr nonnull %chan.op) + call void @llvm.lifetime.start.p0(ptr nonnull %chan.value) + call void @llvm.lifetime.start.p0(ptr nonnull %chan.op) %0 = call i1 @runtime.chanRecv(ptr %ch, ptr nonnull %chan.value, ptr nonnull %chan.op, ptr undef) #3 - call void @llvm.lifetime.end.p0(i64 4, ptr nonnull %chan.value) - call void @llvm.lifetime.end.p0(i64 16, ptr nonnull %chan.op) + call void @llvm.lifetime.end.p0(ptr nonnull %chan.value) + call void @llvm.lifetime.end.p0(ptr nonnull %chan.op) ret void } @@ -55,9 +55,9 @@ declare i1 @runtime.chanRecv(ptr dereferenceable_or_null(36), ptr, ptr dereferen define hidden void @main.chanZeroSend(ptr dereferenceable_or_null(36) %ch, ptr %context) unnamed_addr #1 { entry: %chan.op = alloca %runtime.channelOp, align 8 - call void @llvm.lifetime.start.p0(i64 16, ptr nonnull %chan.op) + call void @llvm.lifetime.start.p0(ptr nonnull %chan.op) call void @runtime.chanSend(ptr %ch, ptr null, ptr nonnull %chan.op, ptr undef) #3 - call void @llvm.lifetime.end.p0(i64 16, ptr nonnull %chan.op) + call void @llvm.lifetime.end.p0(ptr nonnull %chan.op) ret void } @@ -65,9 +65,9 @@ entry: define hidden void @main.chanZeroRecv(ptr dereferenceable_or_null(36) %ch, ptr %context) unnamed_addr #1 { entry: %chan.op = alloca %runtime.channelOp, align 8 - call void @llvm.lifetime.start.p0(i64 16, ptr nonnull %chan.op) + call void @llvm.lifetime.start.p0(ptr nonnull %chan.op) %0 = call i1 @runtime.chanRecv(ptr %ch, ptr null, ptr nonnull %chan.op, ptr undef) #3 - call void @llvm.lifetime.end.p0(i64 16, ptr nonnull %chan.op) + call void @llvm.lifetime.end.p0(ptr nonnull %chan.op) ret void } @@ -77,7 +77,7 @@ entry: %select.states.alloca = alloca [2 x %runtime.chanSelectState], align 8 %select.send.value = alloca i32, align 4 store i32 1, ptr %select.send.value, align 4 - call void @llvm.lifetime.start.p0(i64 16, ptr nonnull %select.states.alloca) + call void @llvm.lifetime.start.p0(ptr nonnull %select.states.alloca) store ptr %ch1, ptr %select.states.alloca, align 4 %select.states.alloca.repack1 = getelementptr inbounds nuw i8, ptr %select.states.alloca, i32 4 store ptr %select.send.value, ptr %select.states.alloca.repack1, align 4 @@ -86,7 +86,7 @@ entry: %.repack3 = getelementptr inbounds nuw i8, ptr %select.states.alloca, i32 12 store ptr null, ptr %.repack3, align 4 %select.result = call { i32, i1 } @runtime.chanSelect(ptr undef, ptr nonnull %select.states.alloca, i32 2, i32 2, ptr null, i32 0, i32 0, ptr undef) #3 - call void @llvm.lifetime.end.p0(i64 16, ptr nonnull %select.states.alloca) + call void @llvm.lifetime.end.p0(ptr nonnull %select.states.alloca) %1 = extractvalue { i32, i1 } %select.result, 0 %2 = icmp eq i32 %1, 0 br i1 %2, label %select.done, label %select.next @@ -104,6 +104,80 @@ select.body: ; preds = %select.next declare { i32, i1 } @runtime.chanSelect(ptr, ptr, i32, i32, ptr, i32, i32, ptr) #0 +; Function Attrs: nounwind +define hidden i1 @main.selectNonBlockingSend(ptr dereferenceable_or_null(36) %ch, i32 %value, ptr %context) unnamed_addr #1 { +entry: + %select.send.value = alloca i32, align 4 + call void @llvm.lifetime.start.p0(ptr nonnull %select.send.value) + store i32 %value, ptr %select.send.value, align 4 + %select.sent = call i1 @runtime.chanTrySend(ptr %ch, ptr nonnull %select.send.value, ptr undef) #3 + call void @llvm.lifetime.end.p0(ptr nonnull %select.send.value) + br i1 %select.sent, label %select.body, label %select.next + +select.body: ; preds = %entry + ret i1 true + +select.next: ; preds = %entry + ret i1 false +} + +declare i1 @runtime.chanTrySend(ptr dereferenceable_or_null(36), ptr, ptr) #0 + +; Function Attrs: nounwind +define hidden { i32, i1, i1 } @main.selectNonBlockingRecv(ptr dereferenceable_or_null(36) %ch, ptr %context) unnamed_addr #1 { +entry: + %select.recvbuf = alloca i32, align 4 + %stackalloc = alloca i8, align 1 + call void @llvm.lifetime.start.p0(ptr nonnull %select.recvbuf) + %select.recv = call { i1, i1 } @runtime.chanTryRecv(ptr %ch, ptr nonnull %select.recvbuf, ptr undef) #3 + %select.received = extractvalue { i1, i1 } %select.recv, 0 + call void @runtime.trackPointer(ptr nonnull %select.recvbuf, ptr nonnull %stackalloc, ptr undef) #3 + br i1 %select.received, label %select.body, label %select.next + +select.body: ; preds = %entry + %select.recv.ok = extractvalue { i1, i1 } %select.recv, 1 + %0 = load i32, ptr %select.recvbuf, align 4 + %1 = insertvalue { i32, i1, i1 } zeroinitializer, i32 %0, 0 + %2 = insertvalue { i32, i1, i1 } %1, i1 %select.recv.ok, 1 + %3 = insertvalue { i32, i1, i1 } %2, i1 true, 2 + ret { i32, i1, i1 } %3 + +select.next: ; preds = %entry + ret { i32, i1, i1 } zeroinitializer +} + +declare { i1, i1 } @runtime.chanTryRecv(ptr dereferenceable_or_null(36), ptr, ptr) #0 + +; Function Attrs: nounwind +define hidden i1 @main.selectNonBlockingZeroSend(ptr dereferenceable_or_null(36) %ch, ptr %context) unnamed_addr #1 { +entry: + %select.sent = call i1 @runtime.chanTrySend(ptr %ch, ptr null, ptr undef) #3 + br i1 %select.sent, label %select.body, label %select.next + +select.body: ; preds = %entry + ret i1 true + +select.next: ; preds = %entry + ret i1 false +} + +; Function Attrs: nounwind +define hidden { i1, i1 } @main.selectNonBlockingZeroRecv(ptr dereferenceable_or_null(36) %ch, ptr %context) unnamed_addr #1 { +entry: + %select.recv = call { i1, i1 } @runtime.chanTryRecv(ptr %ch, ptr null, ptr undef) #3 + %select.received = extractvalue { i1, i1 } %select.recv, 0 + br i1 %select.received, label %select.body, label %select.next + +select.body: ; preds = %entry + %select.recv.ok = extractvalue { i1, i1 } %select.recv, 1 + %0 = insertvalue { i1, i1 } zeroinitializer, i1 %select.recv.ok, 0 + %1 = insertvalue { i1, i1 } %0, i1 true, 1 + ret { i1, i1 } %1 + +select.next: ; preds = %entry + ret { i1, i1 } zeroinitializer +} + attributes #0 = { "target-features"="+bulk-memory,+bulk-memory-opt,+call-indirect-overlong,+mutable-globals,+nontrapping-fptoint,+sign-ext,-multivalue,-reference-types" } attributes #1 = { nounwind "target-features"="+bulk-memory,+bulk-memory-opt,+call-indirect-overlong,+mutable-globals,+nontrapping-fptoint,+sign-ext,-multivalue,-reference-types" } attributes #2 = { nocallback nofree nosync nounwind willreturn memory(argmem: readwrite) } diff --git a/src/runtime/chan.go b/src/runtime/chan.go index a85e9b6617..2b7332d4ac 100644 --- a/src/runtime/chan.go +++ b/src/runtime/chan.go @@ -344,6 +344,46 @@ func chanRecv(ch *channel, value unsafe.Pointer, op *channelOp) bool { return t.DataUint32() != chanOperationClosed } +// chanTrySend attempts a non-blocking send. +func chanTrySend(ch *channel, value unsafe.Pointer) bool { + if ch == nil { + return false + } + + mask := interrupt.Disable() + ch.lock.Lock() + + sent, wake := ch.trySend(value) + + ch.lock.Unlock() + if wake != nil { + scheduleTask(wake) + } + interrupt.Restore(mask) + + return sent +} + +// chanTryRecv attempts a non-blocking receive. +func chanTryRecv(ch *channel, value unsafe.Pointer) (received, ok bool) { + if ch == nil { + return false, true + } + + mask := interrupt.Disable() + ch.lock.Lock() + + received, ok, wake := ch.tryRecv(value) + + ch.lock.Unlock() + if wake != nil { + scheduleTask(wake) + } + interrupt.Restore(mask) + + return received, ok +} + // chanClose closes the given channel. If this channel has a receiver or is // empty, it closes the channel. Else, it panics. func chanClose(ch *channel) {