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