ir_wasmspawn.mx raw
1 package main
2
3 import (
4 "git.smesh.lol/moxie/pkg/mxutil"
5 . "git.smesh.lol/moxie/pkg/types"
6 )
7
8 // wasm spawn: parent and child halves of a domain boundary that crosses a
9 // Worker instead of a fork.
10 //
11 // On wasm a spawned domain is a dedicated Worker running the same module.
12 // Channels that cross the boundary are SPSC rings over SharedArrayBuffer, and
13 // the module only ever holds an int32 handle for one: bridge.channel_create
14 // returns it, bridge.channel_send/recv/close take it, and bridge.spawn_domain
15 // starts the Worker with the target's dispatch index and the handle list.
16 //
17 // The parent side packs scalar arguments into a byte buffer and the handles
18 // into an i32 array; the child's __spawn_entry switches on the index, reloads
19 // the scalars at the offsets the parent used, turns each handle back into the
20 // pointer-shaped value the channel parameter holds, and calls the target.
21 //
22 // The enumeration of targets and their indices lives in cctx and is filled by
23 // scanWasmSpawnTargets, so both halves agree on the numbering.
24
25 // setupWasmSpawnTargetChannels records the SAB handle each channel parameter
26 // of the current spawn target carries. __spawn_entry delivered it as an int32
27 // inttoptr'd into the parameter, so ptrtoint recovers it. Every send, receive
28 // and close on that parameter then goes through the bridge instead of the
29 // runtime's channel object, which the child does not have.
30 func (e *irEmitter) setupWasmSpawnTargetChannels() {
31 if e.ptrBits != 32 || e.curFunc == nil {
32 return
33 }
34 t := e.wasmSpawnTargetForFunc(e.curFunc)
35 if t == nil {
36 return
37 }
38 mxutil.WriteStr(2, " wasm-spawn: target-channels\n")
39 ipt := e.intptrType()
40 for _, pi := range t.chanParamIdx {
41 if pi < 0 || pi >= int32(len(e.curFunc.Params)) {
42 continue
43 }
44 p := e.curFunc.Params[pi]
45 e.nextReg++
46 h := "%sabh" | irItoa(e.nextReg)
47 e.w(" ") ; e.w(h) ; e.w(" = ptrtoint ptr ") ; e.w(e.regName(p)) ; e.w(" to ") ; e.w(ipt) ; e.w("\n")
48 }
49 }
50
51 // emitWasmSpawnCall emits spawn(fn, args...) for the wasm target.
52 func (e *irEmitter) emitWasmSpawnCall(c *SSACall) {
53 mxutil.WriteStr(2, " wasm-spawn: call\n")
54 ipt := e.intptrType()
55 i32 := "i32"
56 reg := e.regName(c)
57
58 targetFn, ok := c.Call.Args[0].(*SSAFunction)
59 if !ok {
60 cctx.compileErrors = mxutil.Ensure(cctx.compileErrors, 1)
61 push(cctx.compileErrors, "spawn: the target must be a static top-level function")
62 e.emitZeroReg(reg, c.SSAType())
63 return
64 }
65 fnIdx := int32(-1)
66 ok2 := false
67 if e.wasmSpawnIndex != nil {
68 fnIdx, ok2 = e.wasmSpawnIndex[targetFn.name]
69 }
70 if !ok2 {
71 // The scan in the same package should have found it. Reaching here
72 // means the numbering would not match the child's switch, so refuse
73 // rather than hand spawn_domain an index that dispatches nowhere.
74 cctx.compileErrors = mxutil.Ensure(cctx.compileErrors, 1)
75 push(cctx.compileErrors, "spawn: " | targetFn.name | " is not in the wasm dispatch table; a wasm spawn target must be a static function in the same package as the spawn call")
76 e.emitZeroReg(reg, c.SSAType())
77 return
78 }
79 mxutil.WriteStr(2, " call: target=" | targetFn.name | "\n")
80 if targetFn.Signature == nil || targetFn.Signature.Params == nil {
81 e.emitZeroReg(reg, c.SSAType())
82 return
83 }
84 nParams := targetFn.Signature.Params.Len()
85 if len(c.Call.Args)-1 != nParams {
86 cctx.compileErrors = mxutil.Ensure(cctx.compileErrors, 1)
87 push(cctx.compileErrors, "spawn: " | targetFn.name | " argument count does not match its parameters")
88 e.emitZeroReg(reg, c.SSAType())
89 return
90 }
91
92
93 // Channel arguments become SAB handles: one ring per channel, sized from
94 // the element type. Both sides keep the handle, not the ring.
95 handles := []string{:0:nParams}
96 for i := int32(0); i < nParams; i++ {
97 arg := c.Call.Args[i+1]
98 isChan2 := false
99 if _, ic := SafeUnderlying(arg.SSAType()).(*TCChan); ic {
100 isChan2 = true
101 }
102 if !isChan2 {
103 continue
104 }
105 elemT := e.chanElemTType(arg.SSAType())
106 elemSz := int32(64)
107 if elemT != nil {
108 if tsz := TypeSize(elemT); tsz > 0 {
109 elemSz = int32(tsz)
110 }
111 }
112 slotSize := int32(64)
113 for slotSize < elemSz+12 {
114 slotSize = slotSize * 2
115 }
116 e.nextReg++
117 h := "%sabch" | irItoa(e.nextReg)
118 e.w(" ") ; e.w(h) ; e.w(" = call ") ; e.w(ipt) ; e.w(" @runtime.SpawnChannelCreate(") ; e.w(i32) ; e.w(" ") ; e.w(irItoa(slotSize)) ; e.w(", ") ; e.w(i32) ; e.w(" 16, ptr null)\n")
119 e.declareRuntime("runtime.SpawnChannelCreate", ipt, i32 | ", " | i32)
120 // Tag the parent's channel so its sends and receives reach the ring.
121 // Marking at run time is what keeps this correct through the copies
122 // and reloads stage4's alloca-based SSA introduces: the same source
123 // variable is a different SSA value at the spawn and at the send, so
124 // tracking SSA values misses it.
125 e.w(" call void @runtime.MarkChannelSAB(ptr ") ; e.w(e.operand(arg)) ; e.w(", ") ; e.w(ipt) ; e.w(" ") ; e.w(h) ; e.w(", ptr null)\n")
126 e.declareRuntime("runtime.MarkChannelSAB", "void", "ptr, " | ipt)
127 handles = mxutil.Ensure(handles, 1)
128 push(handles, h)
129 }
130
131 // Pack the scalar arguments in parameter order, skipping channels, at the
132 // same offsets the child will read them from.
133 scalarVals := []string{:0:nParams}
134 scalarTys := []string{:0:nParams}
135 scalarSizes := []int32{:0:nParams}
136 total := int32(0)
137 for i := int32(0); i < nParams; i++ {
138 arg := c.Call.Args[i+1]
139 if _, isChan := SafeUnderlying(arg.SSAType()).(*TCChan); isChan {
140 continue
141 }
142 lt := e.llvmType(arg.SSAType())
143 if lt == "" || lt == "void" {
144 lt = ipt
145 }
146 sz := int32(TypeSize(arg.SSAType()))
147 scalarVals = mxutil.Ensure(scalarVals, 1)
148 push(scalarVals, e.operand(arg))
149 scalarTys = mxutil.Ensure(scalarTys, 1)
150 push(scalarTys, lt)
151 scalarSizes = mxutil.Ensure(scalarSizes, 1)
152 push(scalarSizes, sz)
153 total += sz
154 }
155
156 mxutil.WriteStr(2, " call: scalars=" | simpleItoa(int32(len(scalarVals))) | " bytes=" | simpleItoa(total) | "\n")
157 argPtr := "null"
158 argLen := "0"
159 if total > 0 {
160 e.nextReg++
161 buf := "%sabargs" | irItoa(e.nextReg)
162 e.w(" ") ; e.w(buf) ; e.w(" = alloca [") ; e.w(irItoa(total)) ; e.w(" x i8], align 8\n")
163 off := int32(0)
164 for i := 0; i < len(scalarVals); i++ {
165 if scalarSizes[i] <= 0 {
166 continue
167 }
168 e.nextReg++
169 g := "%sabarg" | irItoa(e.nextReg)
170 e.w(" ") ; e.w(g) ; e.w(" = getelementptr inbounds [") ; e.w(irItoa(total)) ; e.w(" x i8], ptr ") ; e.w(buf) ; e.w(", i32 0, i32 ") ; e.w(irItoa(off)) ; e.w("\n")
171 e.w(" store ") ; e.w(scalarTys[i]) ; e.w(" ") ; e.w(scalarVals[i]) ; e.w(", ptr ") ; e.w(g) ; e.w("\n")
172 off += scalarSizes[i]
173 }
174 e.nextReg++
175 ap := "%sabargp" | irItoa(e.nextReg)
176 e.w(" ") ; e.w(ap) ; e.w(" = getelementptr inbounds [") ; e.w(irItoa(total)) ; e.w(" x i8], ptr ") ; e.w(buf) ; e.w(", i32 0, i32 0\n")
177 argPtr = ap
178 argLen = irItoa(total)
179 }
180
181 // The handle array mirrors the child's channel-parameter order.
182 chanPtr := "null"
183 nChans := int32(len(handles))
184 if nChans > 0 {
185 e.nextReg++
186 harr := "%sabharr" | irItoa(e.nextReg)
187 e.w(" ") ; e.w(harr) ; e.w(" = alloca [") ; e.w(irItoa(nChans)) ; e.w(" x ") ; e.w(ipt) ; e.w("], align 8\n")
188 for i := 0; i < len(handles); i++ {
189 e.nextReg++
190 g := "%sabhp" | irItoa(e.nextReg)
191 e.w(" ") ; e.w(g) ; e.w(" = getelementptr inbounds [") ; e.w(irItoa(nChans)) ; e.w(" x ") ; e.w(ipt) ; e.w("], ptr ") ; e.w(harr) ; e.w(", i32 0, i32 ") ; e.w(irItoa(int32(i))) ; e.w("\n")
192 e.w(" store ") ; e.w(ipt) ; e.w(" ") ; e.w(handles[i]) ; e.w(", ptr ") ; e.w(g) ; e.w("\n")
193 }
194 e.nextReg++
195 hp := "%sabhp0" | irItoa(e.nextReg)
196 e.w(" ") ; e.w(hp) ; e.w(" = getelementptr inbounds [") ; e.w(irItoa(nChans)) ; e.w(" x ") ; e.w(ipt) ; e.w("], ptr ") ; e.w(harr) ; e.w(", i32 0, i32 0\n")
197 chanPtr = hp
198 }
199
200 mxutil.WriteStr(2, " call: emitting spawn_domain\n")
201 e.w(" call void @runtime.SpawnDomain(") ; e.w(i32) ; e.w(" ") ; e.w(irItoa(fnIdx))
202 e.w(", ptr ") ; e.w(argPtr) ; e.w(", ") ; e.w(i32) ; e.w(" ") ; e.w(argLen)
203 e.w(", ptr ") ; e.w(chanPtr) ; e.w(", ") ; e.w(i32) ; e.w(" ") ; e.w(irItoa(nChans)) ; e.w(", ptr null)\n")
204 e.declareRuntime("runtime.SpawnDomain", "void", i32 | ", ptr, " | i32 | ", ptr, " | i32)
205
206 // spawn() yields the domain's lifecycle channel on native. A wasm domain
207 // has no such channel, and the legacy compiler returns an undefined value
208 // here too, so keep the two compilers agreeing on that.
209 e.emitZeroReg(reg, c.SSAType())
210 }
211
212 // emitWasmSpawnEntry emits the child-side dispatch the Worker calls instead of
213 // _start: __spawn_entry(fnIdx, argPtr, argLen, chanHandlesPtr, nChans).
214 //
215 // One block per target reloads that target's scalars from the same offsets
216 // emitWasmSpawnCall wrote, converts each handle back to the pointer-shaped
217 // value the channel parameter holds, and calls it. An unknown index traps,
218 // because silently returning would leave the parent waiting on a domain that
219 // never existed.
220 func (e *irEmitter) emitWasmSpawnEntry() {
221 if e.ptrBits != 32 || len(e.wasmSpawnTargets) == 0 {
222 return
223 }
224 mxutil.WriteStr(2, " wasm-spawn: entry\n")
225 ipt := e.intptrType()
226 i32 := "i32"
227 // Declared here rather than through declareRuntime: the declaration block
228 // is emitted before this runs, so anything registered now would be dropped.
229 e.w("\ndeclare ptr @runtime.SpawnChannelWrap(") ; e.w(ipt) ; e.w(", ") ; e.w(ipt) ; e.w(", ptr)\n")
230 e.w("declare void @llvm.trap()\n")
231 e.w("\ndefine void @__spawn_entry(") ; e.w(i32) ; e.w(" %fnIdx, ptr %argPtr, ") ; e.w(i32) ; e.w(" %argLen, ptr %chanHandlesPtr, ") ; e.w(i32) ; e.w(" %nChans) {\nentry:\n")
232 e.w(" switch ") ; e.w(i32) ; e.w(" %fnIdx, label %spawn.bad [\n")
233 for _, t := range e.wasmSpawnTargets {
234 e.w(" ") ; e.w(i32) ; e.w(" ") ; e.w(irItoa(t.fnIdx)) ; e.w(", label %spawn.case") ; e.w(irItoa(t.fnIdx)) ; e.w("\n")
235 }
236 e.w(" ]\n\n")
237 e.w("spawn.bad:\n call void @llvm.trap()\n unreachable\n")
238
239 for _, t := range e.wasmSpawnTargets {
240 e.w("\nspawn.case") ; e.w(irItoa(t.fnIdx)) ; e.w(":\n")
241 sig := t.fn.Signature
242 if sig == nil || sig.Params == nil {
243 e.w(" ret void\n")
244 continue
245 }
246 // Scalars arrive in parameter order at the offsets the parent used;
247 // channel parameters take the next handle from the handle array.
248 scalarOff := int32(0)
249 chanN := int32(0)
250 var callArgs []string
251 callArgs = []string{:0:sig.Params.Len()}
252 for i := int32(0); i < sig.Params.Len(); i++ {
253 pt := sig.Params.At(i).Typ
254 if _, isChan := SafeUnderlying(pt).(*TCChan); isChan {
255 e.nextReg++
256 hg := "%seh" | irItoa(e.nextReg)
257 e.w(" ") ; e.w(hg) ; e.w(" = getelementptr inbounds ") ; e.w(ipt) ; e.w(", ptr %chanHandlesPtr, i32 ") ; e.w(irItoa(chanN)) ; e.w("\n")
258 e.nextReg++
259 hv := "%shv" | irItoa(e.nextReg)
260 e.w(" ") ; e.w(hv) ; e.w(" = load ") ; e.w(ipt) ; e.w(", ptr ") ; e.w(hg) ; e.w("\n")
261 e.nextReg++
262 hp := "%shp" | irItoa(e.nextReg)
263 // The child gets an ordinary channel whose handle routes send,
264 // receive and close to the parent's ring.
265 e.w(" ") ; e.w(hp) ; e.w(" = call ptr @runtime.SpawnChannelWrap(") ; e.w(ipt) ; e.w(" ") ; e.w(hv) ; e.w(", ") ; e.w(ipt) ; e.w(" ") ; e.w(irItoa(int32(TypeSize(pt)))) ; e.w(", ptr null)\n")
266 push(callArgs, "ptr " | hp)
267 chanN++
268 continue
269 }
270 lt := e.llvmType(pt)
271 if lt == "" || lt == "void" {
272 lt = ipt
273 }
274 e.nextReg++
275 g := "%sag" | irItoa(e.nextReg)
276 e.w(" ") ; e.w(g) ; e.w(" = getelementptr inbounds i8, ptr %argPtr, i32 ") ; e.w(irItoa(scalarOff)) ; e.w("\n")
277 e.nextReg++
278 v := "%sav" | irItoa(e.nextReg)
279 e.w(" ") ; e.w(v) ; e.w(" = load ") ; e.w(lt) ; e.w(", ptr ") ; e.w(g) ; e.w("\n")
280 push(callArgs, lt | " " | v)
281 scalarOff += int32(TypeSize(pt))
282 }
283 e.w(" call void ") ; e.w(e.funcSymbol(t.fn)) ; e.w("(")
284 for i, a := range callArgs {
285 if i > 0 {
286 e.w(", ")
287 }
288 e.w(a)
289 }
290 if len(callArgs) > 0 {
291 e.w(", ")
292 }
293 e.w("ptr null)\n")
294 e.w(" ret void\n")
295 }
296 e.w("}\n")
297 }
298
299 // wasmSpawnTargetForFunc returns the dispatch entry for a function, or nil.
300 func (e *irEmitter) wasmSpawnTargetForFunc(f *SSAFunction) (t *wasmSpawnTarget) {
301 if f == nil {
302 return nil
303 }
304 for _, t2 := range e.wasmSpawnTargets {
305 if t2.fn == f {
306 return t2
307 }
308 }
309 return nil
310 }
311