1
2
3
4
5 package walk
6
7 import (
8 "cmp"
9 "fmt"
10 "go/constant"
11 "go/token"
12 "math"
13 "math/bits"
14 "slices"
15 "sort"
16 "strings"
17
18 "cmd/compile/internal/base"
19 "cmd/compile/internal/ir"
20 "cmd/compile/internal/objw"
21 "cmd/compile/internal/reflectdata"
22 "cmd/compile/internal/rttype"
23 "cmd/compile/internal/ssagen"
24 "cmd/compile/internal/staticdata"
25 "cmd/compile/internal/typecheck"
26 "cmd/compile/internal/types"
27 "cmd/internal/obj"
28 "cmd/internal/src"
29 )
30
31
32 func walkSwitch(sw *ir.SwitchStmt) {
33
34 if sw.Walked() {
35 return
36 }
37 sw.SetWalked(true)
38
39 if sw.Tag != nil && sw.Tag.Op() == ir.OTYPESW {
40 walkSwitchType(sw)
41 } else {
42 walkSwitchExpr(sw)
43 }
44 }
45
46
47
48 func walkSwitchExpr(sw *ir.SwitchStmt) {
49 lno := ir.SetPos(sw)
50
51 cond := sw.Tag
52 sw.Tag = nil
53
54
55 if cond == nil {
56 cond = ir.NewBool(base.Pos, true)
57 cond = typecheck.Expr(cond)
58 cond = typecheck.DefaultLit(cond, nil)
59 }
60
61
62
63
64
65
66
67
68 if cond.Op() == ir.OBYTES2STR && allCaseExprsAreSideEffectFree(sw) {
69 cond := cond.(*ir.ConvExpr)
70 cond.SetOp(ir.OBYTES2STRTMP)
71 }
72
73 cond = walkExpr(cond, sw.PtrInit())
74 if cond.Op() != ir.OLITERAL && cond.Op() != ir.ONIL {
75 cond = copyExpr(cond, cond.Type(), &sw.Compiled)
76 }
77
78 base.Pos = lno
79
80 tryLookupTable(sw, cond)
81
82 s := exprSwitch{
83 pos: lno,
84 exprname: cond,
85 }
86
87 var defaultGoto ir.Node
88 var body ir.Nodes
89 for _, ncase := range sw.Cases {
90 label := typecheck.AutoLabel(".s")
91 jmp := ir.NewBranchStmt(ncase.Pos(), ir.OGOTO, label)
92
93
94 if len(ncase.List) == 0 {
95 if defaultGoto != nil {
96 base.Fatalf("duplicate default case not detected during typechecking")
97 }
98 defaultGoto = jmp
99 }
100
101 for i, n1 := range ncase.List {
102 var rtype ir.Node
103 if i < len(ncase.RTypes) {
104 rtype = ncase.RTypes[i]
105 }
106 s.Add(ncase.Pos(), n1, rtype, jmp)
107 }
108
109
110 body.Append(ir.NewLabelStmt(ncase.Pos(), label))
111 body.Append(ncase.Body...)
112 if fall, pos := endsInFallthrough(ncase.Body); !fall {
113 br := ir.NewBranchStmt(base.Pos, ir.OBREAK, nil)
114 br.SetPos(pos)
115 body.Append(br)
116 }
117 }
118 sw.Cases = nil
119
120 if defaultGoto == nil {
121 br := ir.NewBranchStmt(base.Pos, ir.OBREAK, nil)
122 br.SetPos(br.Pos().WithNotStmt())
123 defaultGoto = br
124 }
125
126 s.Emit(&sw.Compiled)
127 sw.Compiled.Append(defaultGoto)
128 sw.Compiled.Append(body.Take()...)
129 walkStmtList(sw.Compiled)
130 }
131
132
133 type exprSwitch struct {
134 pos src.XPos
135 exprname ir.Node
136
137 done ir.Nodes
138 clauses []exprClause
139 }
140
141 type exprClause struct {
142 pos src.XPos
143 lo, hi ir.Node
144 rtype ir.Node
145 jmp ir.Node
146 }
147
148 func (s *exprSwitch) Add(pos src.XPos, expr, rtype, jmp ir.Node) {
149 c := exprClause{pos: pos, lo: expr, hi: expr, rtype: rtype, jmp: jmp}
150 if types.IsOrdered[s.exprname.Type().Kind()] && expr.Op() == ir.OLITERAL {
151 s.clauses = append(s.clauses, c)
152 return
153 }
154
155 s.flush()
156 s.clauses = append(s.clauses, c)
157 s.flush()
158 }
159
160 func (s *exprSwitch) Emit(out *ir.Nodes) {
161 s.flush()
162 out.Append(s.done.Take()...)
163 }
164
165 func (s *exprSwitch) flush() {
166 cc := s.clauses
167 s.clauses = nil
168 if len(cc) == 0 {
169 return
170 }
171
172
173
174
175
176
177 if s.exprname.Type().IsString() && len(cc) >= 2 {
178
179
180
181
182 slices.SortFunc(cc, func(a, b exprClause) int {
183 si := ir.StringVal(a.lo)
184 sj := ir.StringVal(b.lo)
185 if len(si) != len(sj) {
186 return cmp.Compare(len(si), len(sj))
187 }
188 return strings.Compare(si, sj)
189 })
190
191
192
193 runLen := func(run []exprClause) int64 { return int64(len(ir.StringVal(run[0].lo))) }
194
195
196 var runs [][]exprClause
197 start := 0
198 for i := 1; i < len(cc); i++ {
199 if runLen(cc[start:]) != runLen(cc[i:]) {
200 runs = append(runs, cc[start:i])
201 start = i
202 }
203 }
204 runs = append(runs, cc[start:])
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227 outerLabel := typecheck.AutoLabel(".s")
228 endLabel := typecheck.AutoLabel(".s")
229
230
231 s.done.Append(ir.NewBranchStmt(s.pos, ir.OGOTO, outerLabel))
232
233 var outer exprSwitch
234 outer.exprname = ir.NewUnaryExpr(s.pos, ir.OLEN, s.exprname)
235 outer.exprname.SetType(types.Types[types.TINT])
236
237 for _, run := range runs {
238
239 label := typecheck.AutoLabel(".s")
240
241
242 pos := run[0].pos
243 s.done.Append(ir.NewLabelStmt(pos, label))
244 stringSearch(s.exprname, run, &s.done)
245 s.done.Append(ir.NewBranchStmt(pos, ir.OGOTO, endLabel))
246
247
248 cas := ir.NewInt(pos, runLen(run))
249 jmp := ir.NewBranchStmt(pos, ir.OGOTO, label)
250 outer.Add(pos, cas, nil, jmp)
251 }
252 s.done.Append(ir.NewLabelStmt(s.pos, outerLabel))
253 outer.Emit(&s.done)
254 s.done.Append(ir.NewLabelStmt(s.pos, endLabel))
255 return
256 }
257
258 sort.Slice(cc, func(i, j int) bool {
259 return constant.Compare(cc[i].lo.Val(), token.LSS, cc[j].lo.Val())
260 })
261
262
263 if s.exprname.Type().IsInteger() {
264 consecutive := func(last, next constant.Value) bool {
265 delta := constant.BinaryOp(next, token.SUB, last)
266 return constant.Compare(delta, token.EQL, constant.MakeInt64(1))
267 }
268
269 merged := cc[:1]
270 for _, c := range cc[1:] {
271 last := &merged[len(merged)-1]
272 if last.jmp == c.jmp && consecutive(last.hi.Val(), c.lo.Val()) {
273 last.hi = c.lo
274 } else {
275 merged = append(merged, c)
276 }
277 }
278 cc = merged
279 }
280
281 s.search(cc, &s.done)
282 }
283
284 func (s *exprSwitch) search(cc []exprClause, out *ir.Nodes) {
285 if s.tryJumpTable(cc, out) {
286 return
287 }
288 binarySearch(len(cc), out,
289 func(i int) ir.Node {
290 return ir.NewBinaryExpr(base.Pos, ir.OLE, s.exprname, cc[i-1].hi)
291 },
292 func(i int, nif *ir.IfStmt) {
293 c := &cc[i]
294 nif.Cond = c.test(s.exprname)
295 nif.Body = []ir.Node{c.jmp}
296 },
297 )
298 }
299
300
301 func (s *exprSwitch) tryJumpTable(cc []exprClause, out *ir.Nodes) bool {
302 const minCases = 8
303 const minDensity = 4
304
305 if base.Flag.N != 0 || !ssagen.Arch.LinkArch.CanJumpTable || base.Ctxt.Retpoline {
306 return false
307 }
308 if len(cc) < minCases {
309 return false
310 }
311 if cc[0].lo.Val().Kind() != constant.Int {
312 return false
313 }
314 if s.exprname.Type().Size() > int64(types.PtrSize) {
315 return false
316 }
317 min := cc[0].lo.Val()
318 max := cc[len(cc)-1].hi.Val()
319 width := constant.BinaryOp(constant.BinaryOp(max, token.SUB, min), token.ADD, constant.MakeInt64(1))
320 limit := constant.MakeInt64(int64(len(cc)) * minDensity)
321 if constant.Compare(width, token.GTR, limit) {
322
323
324 return false
325 }
326 jt := ir.NewJumpTableStmt(base.Pos, s.exprname)
327 for _, c := range cc {
328 jmp := c.jmp.(*ir.BranchStmt)
329 if jmp.Op() != ir.OGOTO || jmp.Label == nil {
330 panic("bad switch case body")
331 }
332 for i := c.lo.Val(); constant.Compare(i, token.LEQ, c.hi.Val()); i = constant.BinaryOp(i, token.ADD, constant.MakeInt64(1)) {
333 jt.Cases = append(jt.Cases, i)
334 jt.Targets = append(jt.Targets, jmp.Label)
335 }
336 }
337 out.Append(jt)
338 return true
339 }
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378 func tryLookupTable(sw *ir.SwitchStmt, cond ir.Node) {
379 const minCases = 4
380
381 if base.Flag.N != 0 {
382 return
383 }
384 if !cond.Type().IsInteger() {
385 return
386 }
387 if cond.Type().Size() > int64(types.PtrSize) {
388 return
389 }
390
391 fn := ir.CurFunc
392 if fn == nil || fn.Type().NumResults() != 1 {
393 return
394 }
395 resultType := fn.Type().Results()[0].Type
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416 constSet := make(map[int64]ir.Node)
417 constCaseSet := make(map[int]bool)
418 excludeSet := make(map[int64]bool)
419 var defaultVal ir.Node
420 minVal, maxVal := int64(math.MaxInt64), int64(math.MinInt64)
421 var excludeNextCase bool
422
423 for i, ncase := range sw.Cases {
424
425
426 isFallthroughTarget := excludeNextCase
427 excludeNextCase, _ = endsInFallthrough(ncase.Body)
428
429 if len(ncase.List) == 0 {
430
431 if isConstReturn(ncase) && !isFallthroughTarget {
432 defaultVal = ncase.Body[0].(*ir.ReturnStmt).Results[0]
433 }
434 continue
435 }
436
437 vals, ok := constIntCaseVals(ncase)
438 if !ok {
439
440
441
442
443
444
445
446
447
448 return
449 }
450
451 if !isConstReturn(ncase) || isFallthroughTarget || excludeNextCase {
452
453
454
455
456 for _, v := range vals {
457 excludeSet[v] = true
458 }
459 continue
460 }
461
462 retVal := ncase.Body[0].(*ir.ReturnStmt).Results[0]
463 for _, v := range vals {
464 constSet[v] = retVal
465 minVal = min(minVal, v)
466 maxVal = max(maxVal, v)
467 }
468 constCaseSet[i] = true
469 }
470
471 if len(constSet) < minCases {
472 return
473 }
474
475 tableSize := maxVal - minVal + 1
476 if tableSize <= 0 || !isSwitchDense(int64(len(constSet)), tableSize) {
477 return
478 }
479
480
481
482 tabType := types.NewArray(resultType, tableSize)
483 tabName := readonlystaticname(tabType)
484 elemSize := int(resultType.Size())
485 maxBitmaskSize := int64(types.PtrSize * 8)
486
487 var needMask bool
488 var bitmask uint64
489 validSlots := make([]bool, tableSize)
490
491 tabName.Linksym().WriteInt(base.Ctxt, tableSize*int64(elemSize)-1, 1, 0)
492 for i := range tableSize {
493 caseVal := minVal + i
494 if excludeSet[caseVal] {
495
496 needMask = true
497 } else {
498 val := constSet[caseVal]
499 if val == nil {
500 if defaultVal == nil {
501
502 needMask = true
503 continue
504 }
505 val = defaultVal
506 }
507 staticdata.InitConst(tabName, i*int64(elemSize), val, elemSize)
508 validSlots[i] = true
509 bitmask |= 1 << uint(i)
510 }
511 }
512
513
514
515
516
517 var maskName *ir.Name
518 useBitmask := needMask && tableSize <= maxBitmaskSize
519 if needMask && !useBitmask {
520 maskType := types.NewArray(types.Types[types.TUINT8], tableSize)
521 maskName = readonlystaticname(maskType)
522 maskSym := maskName.Linksym()
523 for i := range tableSize {
524 var v uint8
525 if validSlots[i] {
526 v = 1
527 }
528 maskSym.WriteInt(base.Ctxt, i, 1, int64(v))
529 }
530 }
531
532
533
534
535 pos := sw.Pos()
536
537
538 intType := types.Types[types.TINT]
539 wideCond := typecheck.Conv(cond, intType)
540
541
542 var idx ir.Node
543 if minVal != 0 {
544 minLit := ir.NewBasicLit(pos, intType, constant.MakeInt64(minVal))
545 idx = typecheck.Expr(ir.NewBinaryExpr(pos, ir.OSUB, wideCond, minLit))
546 } else {
547 idx = wideCond
548 }
549
550
551
552 uintType := types.Types[types.TUINT]
553 uidx := typecheck.Conv(idx, uintType)
554 uidx = copyExpr(uidx, uintType, &sw.Compiled)
555
556
557 rangeLit := ir.NewBasicLit(pos, uintType, constant.MakeUint64(uint64(maxVal-minVal)))
558 boundsCheck := typecheck.Expr(ir.NewBinaryExpr(pos, ir.OLE, uidx, rangeLit))
559 boundsCheck = typecheck.DefaultLit(boundsCheck, nil)
560
561
562 lookup := ir.NewIndexExpr(pos, tabName, uidx)
563 lookup.SetBounded(true)
564 lookup = typecheck.Expr(lookup).(*ir.IndexExpr)
565
566 retStmt := ir.NewReturnStmt(pos, []ir.Node{lookup})
567
568 var ifBody []ir.Node
569 if needMask {
570 var maskCheck ir.Node
571 if useBitmask {
572
573
574 bitmaskType := types.Types[types.TUINTPTR]
575 bitmaskLit := ir.NewBasicLit(pos, bitmaskType, constant.MakeUint64(bitmask))
576 shifted := typecheck.Expr(ir.NewBinaryExpr(pos, ir.ORSH, bitmaskLit, uidx))
577 one := ir.NewBasicLit(pos, bitmaskType, constant.MakeUint64(1))
578 masked := typecheck.Expr(ir.NewBinaryExpr(pos, ir.OAND, shifted, one))
579 zero := ir.NewBasicLit(pos, bitmaskType, constant.MakeUint64(0))
580 maskCheck = typecheck.Expr(ir.NewBinaryExpr(pos, ir.ONE, masked, zero))
581 } else {
582
583 maskLookup := ir.NewIndexExpr(pos, maskName, uidx)
584 maskLookup.SetBounded(true)
585 maskLookup = typecheck.Expr(maskLookup).(*ir.IndexExpr)
586 zero := ir.NewBasicLit(pos, types.Types[types.TUINT8], constant.MakeInt64(0))
587 maskCheck = typecheck.Expr(ir.NewBinaryExpr(pos, ir.ONE, maskLookup, zero))
588 }
589 maskCheck = typecheck.DefaultLit(maskCheck, nil)
590
591 innerIf := ir.NewIfStmt(pos, maskCheck, []ir.Node{retStmt}, nil)
592 ifBody = []ir.Node{innerIf}
593 } else {
594 ifBody = []ir.Node{retStmt}
595 }
596
597 outerIf := ir.NewIfStmt(pos, boundsCheck, ifBody, nil)
598 sw.Compiled.Append(outerIf)
599
600
601
602 newCases := make([]*ir.CaseClause, 0, len(sw.Cases)-len(constCaseSet))
603 for i, ncase := range sw.Cases {
604 if !constCaseSet[i] {
605 newCases = append(newCases, ncase)
606 }
607 }
608 sw.Cases = newCases
609 }
610
611
612
613
614
615 func isSwitchDense(numCases, tableSize int64) bool {
616 const minDensity = 40
617 if tableSize >= math.MaxInt64/100 {
618 return false
619 }
620 return numCases*100 >= tableSize*minDensity
621 }
622
623
624
625 func isConstReturn(ncase *ir.CaseClause) bool {
626 if len(ncase.Body) != 1 {
627 return false
628 }
629 ret, ok := ncase.Body[0].(*ir.ReturnStmt)
630 if !ok || len(ret.Results) != 1 {
631 return false
632 }
633 return ret.Results[0].Op() == ir.OLITERAL
634 }
635
636
637
638
639 func constIntCaseVals(ncase *ir.CaseClause) (vals []int64, ok bool) {
640 for _, n1 := range ncase.List {
641 if n1.Op() != ir.OLITERAL || n1.Val().Kind() != constant.Int {
642 return nil, false
643 }
644 v, fit := constant.Int64Val(n1.Val())
645 if !fit {
646 return nil, false
647 }
648 vals = append(vals, v)
649 }
650 return vals, true
651 }
652
653 func (c *exprClause) test(exprname ir.Node) ir.Node {
654
655 if c.hi != c.lo {
656 low := ir.NewBinaryExpr(c.pos, ir.OGE, exprname, c.lo)
657 high := ir.NewBinaryExpr(c.pos, ir.OLE, exprname, c.hi)
658 return ir.NewLogicalExpr(c.pos, ir.OANDAND, low, high)
659 }
660
661
662 if ir.IsConst(exprname, constant.Bool) && !c.lo.Type().IsInterface() {
663 if ir.BoolVal(exprname) {
664 return c.lo
665 } else {
666 return ir.NewUnaryExpr(c.pos, ir.ONOT, c.lo)
667 }
668 }
669
670 n := ir.NewBinaryExpr(c.pos, ir.OEQ, exprname, c.lo)
671 n.RType = c.rtype
672 return n
673 }
674
675 func allCaseExprsAreSideEffectFree(sw *ir.SwitchStmt) bool {
676
677
678
679
680
681
682
683 for _, ncase := range sw.Cases {
684 for _, v := range ncase.List {
685 if v.Op() != ir.OLITERAL {
686 return false
687 }
688 }
689 }
690 return true
691 }
692
693
694 func endsInFallthrough(stmts []ir.Node) (bool, src.XPos) {
695 if len(stmts) == 0 {
696 return false, src.NoXPos
697 }
698 i := len(stmts) - 1
699 return stmts[i].Op() == ir.OFALL, stmts[i].Pos()
700 }
701
702
703
704 func walkSwitchType(sw *ir.SwitchStmt) {
705 var s typeSwitch
706 origSrc := sw.Tag.(*ir.TypeSwitchGuard).X
707 s.srcName = origSrc
708 s.srcName = walkExpr(s.srcName, sw.PtrInit())
709 s.srcName = copyExpr(s.srcName, s.srcName.Type(), &sw.Compiled)
710 s.okName = typecheck.TempAt(base.Pos, ir.CurFunc, types.Types[types.TBOOL])
711 s.itabName = typecheck.TempAt(base.Pos, ir.CurFunc, types.Types[types.TUINT8].PtrTo())
712
713
714
715
716 srcItab := ir.NewUnaryExpr(base.Pos, ir.OITAB, s.srcName)
717 srcData := ir.NewUnaryExpr(base.Pos, ir.OIDATA, s.srcName)
718 srcData.SetType(types.Types[types.TUINT8].PtrTo())
719 srcData.SetTypecheck(1)
720
721
722
723
724
725
726
727 ifNil := ir.NewIfStmt(base.Pos, nil, nil, nil)
728 ifNil.Cond = ir.NewBinaryExpr(base.Pos, ir.OEQ, srcItab, typecheck.NodNil())
729 base.Pos = base.Pos.WithNotStmt()
730 ifNil.Cond = typecheck.Expr(ifNil.Cond)
731 ifNil.Cond = typecheck.DefaultLit(ifNil.Cond, nil)
732
733 sw.Compiled.Append(ifNil)
734
735
736 dotHash := typeHashFieldOf(base.Pos, srcItab)
737 s.hashName = copyExpr(dotHash, dotHash.Type(), &sw.Compiled)
738
739
740 labels := make([]*types.Sym, len(sw.Cases))
741 for i := range sw.Cases {
742 labels[i] = typecheck.AutoLabel(".s")
743 }
744
745
746 br := ir.NewBranchStmt(base.Pos, ir.OBREAK, nil)
747
748
749
750
751 type oneCase struct {
752 pos src.XPos
753 jmp ir.Node
754
755
756
757
758
759 typ ir.Node
760
761
762
763
764
765 val ir.Node
766 idx int
767 }
768 var cases []oneCase
769 var defaultGoto, nilGoto ir.Node
770 for i, ncase := range sw.Cases {
771 jmp := ir.NewBranchStmt(ncase.Pos(), ir.OGOTO, labels[i])
772 if len(ncase.List) == 0 {
773 if defaultGoto != nil {
774 base.Fatalf("duplicate default case not detected during typechecking")
775 }
776 defaultGoto = jmp
777 }
778 for _, n1 := range ncase.List {
779 if ir.IsNil(n1) {
780 if nilGoto != nil {
781 base.Fatalf("duplicate nil case not detected during typechecking")
782 }
783 nilGoto = jmp
784 continue
785 }
786 idx := -1
787 var val ir.Node
788
789 if len(ncase.List) == 1 && ncase.List[0].Op() == ir.ODYNAMICTYPE && ncase.Var != nil {
790 val = typecheck.TempAt(ncase.Pos(), ir.CurFunc, ncase.Var.Type())
791 idx = i
792 }
793 cases = append(cases, oneCase{
794 pos: ncase.Pos(),
795 typ: n1,
796 jmp: jmp,
797 val: val,
798 idx: idx,
799 })
800 }
801 }
802 if defaultGoto == nil {
803 defaultGoto = br
804 }
805 if nilGoto == nil {
806 nilGoto = defaultGoto
807 }
808 ifNil.Body = []ir.Node{nilGoto}
809
810
811 var concreteCases []oneCase
812 var interfaceCases []oneCase
813 flush := func() {
814
815
816
817
818 if len(concreteCases) > 0 {
819 var clauses []typeClause
820 for _, c := range concreteCases {
821 as := ir.NewAssignListStmt(c.pos, ir.OAS2,
822 []ir.Node{ir.BlankNode, s.okName},
823 []ir.Node{ir.NewTypeAssertExpr(c.pos, s.srcName, c.typ.Type())})
824 nif := ir.NewIfStmt(c.pos, s.okName, []ir.Node{c.jmp}, nil)
825 clauses = append(clauses, typeClause{
826 hash: types.TypeHash(c.typ.Type()),
827 body: []ir.Node{typecheck.Stmt(as), typecheck.Stmt(nif)},
828 })
829 }
830 s.flush(clauses, &sw.Compiled)
831 concreteCases = concreteCases[:0]
832 }
833
834
835
836
837 var anyGoto ir.Node
838 if len(interfaceCases) > 0 && interfaceCases[len(interfaceCases)-1].typ.Type().IsEmptyInterface() {
839 anyGoto = interfaceCases[len(interfaceCases)-1].jmp
840 interfaceCases = interfaceCases[:len(interfaceCases)-1]
841 }
842
843
844 if len(interfaceCases) > 0 {
845
846
847 lsym := types.LocalPkg.Lookup(fmt.Sprintf(".interfaceSwitch.%d", interfaceSwitchGen)).LinksymABI(obj.ABI0)
848 interfaceSwitchGen++
849 c := rttype.NewCursor(lsym, 0, rttype.InterfaceSwitch)
850 c.Field("Cache").WritePtr(typecheck.LookupRuntimeVar("emptyInterfaceSwitchCache"))
851 c.Field("NCases").WriteInt(int64(len(interfaceCases)))
852 array, sizeDelta := c.Field("Cases").ModifyArray(len(interfaceCases))
853 for i, c := range interfaceCases {
854 array.Elem(i).WritePtr(reflectdata.TypeLinksym(c.typ.Type()))
855 }
856 objw.Global(lsym, int32(rttype.InterfaceSwitch.Size()+sizeDelta), obj.LOCAL)
857
858
859
860 lsym.Gotype = reflectdata.TypeLinksym(rttype.InterfaceSwitch)
861
862
863
864 var typeArg ir.Node
865 if s.srcName.Type().IsEmptyInterface() {
866 typeArg = ir.NewConvExpr(base.Pos, ir.OCONVNOP, types.Types[types.TUINT8].PtrTo(), srcItab)
867 } else {
868 typeArg = itabType(srcItab)
869 }
870 caseVar := typecheck.TempAt(base.Pos, ir.CurFunc, types.Types[types.TINT])
871 isw := ir.NewInterfaceSwitchStmt(base.Pos, caseVar, s.itabName, typeArg, dotHash, lsym)
872 sw.Compiled.Append(isw)
873
874
875 var newCases []*ir.CaseClause
876 for i, c := range interfaceCases {
877 newCases = append(newCases, &ir.CaseClause{
878 List: []ir.Node{ir.NewInt(base.Pos, int64(i))},
879 Body: []ir.Node{c.jmp},
880 })
881 }
882
883 sw2 := ir.NewSwitchStmt(base.Pos, caseVar, newCases)
884 sw.Compiled.Append(typecheck.Stmt(sw2))
885 interfaceCases = interfaceCases[:0]
886 }
887
888 if anyGoto != nil {
889
890
891 sw.Compiled.Append(anyGoto)
892 }
893 }
894 caseLoop:
895 for _, c := range cases {
896 if c.typ.Op() == ir.ODYNAMICTYPE {
897 flush()
898 dt := c.typ.(*ir.DynamicType)
899 dot := ir.NewDynamicTypeAssertExpr(c.pos, ir.ODYNAMICDOTTYPE, s.srcName, dt.RType)
900 dot.ITab = dt.ITab
901 dot.SetType(c.typ.Type())
902 dot.SetTypecheck(1)
903
904 as := ir.NewAssignListStmt(c.pos, ir.OAS2, nil, nil)
905 as.Lhs = []ir.Node{ir.BlankNode, s.okName}
906 if c.val != nil {
907 as.Lhs[0] = c.val
908 }
909 as.Rhs = []ir.Node{dot}
910 typecheck.Stmt(as)
911
912 nif := ir.NewIfStmt(c.pos, s.okName, []ir.Node{c.jmp}, nil)
913 sw.Compiled.Append(as, nif)
914 continue
915 }
916
917
918
919
920
921 for _, ic := range interfaceCases {
922
923
924 if typecheck.Implements(c.typ.Type(), ic.typ.Type()) {
925 continue caseLoop
926 }
927
928
929
930
931
932
933
934
935 }
936
937 if shapeTypeAssertImpossible(origSrc, c.typ.Type()) {
938 continue
939 }
940
941 if c.typ.Type().IsInterface() {
942 interfaceCases = append(interfaceCases, c)
943 } else {
944 concreteCases = append(concreteCases, c)
945 }
946 }
947 flush()
948
949 sw.Compiled.Append(defaultGoto)
950
951
952 for i, ncase := range sw.Cases {
953 sw.Compiled.Append(ir.NewLabelStmt(ncase.Pos(), labels[i]))
954 if caseVar := ncase.Var; caseVar != nil {
955 val := s.srcName
956 if len(ncase.List) == 1 {
957
958 if ncase.List[0].Op() == ir.OTYPE {
959 t := ncase.List[0].Type()
960 if t.IsInterface() {
961
962
963 if t.IsEmptyInterface() {
964 var typ ir.Node
965 if s.srcName.Type().IsEmptyInterface() {
966
967 typ = srcItab
968 } else {
969
970 typ = itabType(srcItab)
971 typ.SetPos(ncase.Pos())
972 }
973 val = ir.NewBinaryExpr(ncase.Pos(), ir.OMAKEFACE, typ, srcData)
974 } else {
975
976 val = ir.NewBinaryExpr(ncase.Pos(), ir.OMAKEFACE, s.itabName, srcData)
977 }
978 } else {
979
980 val = ifaceData(ncase.Pos(), s.srcName, t)
981 }
982 } else if ncase.List[0].Op() == ir.ODYNAMICTYPE {
983 var found bool
984 for _, c := range cases {
985 if c.idx == i {
986 val = c.val
987 found = val != nil
988 break
989 }
990 }
991
992 if !found {
993 base.Fatalf("an error occurred when processing type switch case %v", ncase.List[0])
994 }
995 } else if ir.IsNil(ncase.List[0]) {
996 } else {
997 base.Fatalf("unhandled type switch case %v", ncase.List[0])
998 }
999 val.SetType(caseVar.Type())
1000 val.SetTypecheck(1)
1001 }
1002 l := []ir.Node{
1003 ir.NewDecl(ncase.Pos(), ir.ODCL, caseVar),
1004 ir.NewAssignStmt(ncase.Pos(), caseVar, val),
1005 }
1006 typecheck.Stmts(l)
1007 sw.Compiled.Append(l...)
1008 }
1009 sw.Compiled.Append(ncase.Body...)
1010 sw.Compiled.Append(br)
1011 }
1012
1013 walkStmtList(sw.Compiled)
1014 sw.Tag = nil
1015 sw.Cases = nil
1016 }
1017
1018 var interfaceSwitchGen int
1019
1020
1021
1022
1023 func typeHashFieldOf(pos src.XPos, itab *ir.UnaryExpr) *ir.SelectorExpr {
1024 if itab.Op() != ir.OITAB {
1025 base.Fatalf("expected OITAB, got %v", itab.Op())
1026 }
1027 var hashField *types.Field
1028 if itab.X.Type().IsEmptyInterface() {
1029
1030 if rtypeHashField == nil {
1031 rtypeHashField = runtimeField("hash", rttype.Type.OffsetOf("Hash"), types.Types[types.TUINT32])
1032 }
1033 hashField = rtypeHashField
1034 } else {
1035
1036 if itabHashField == nil {
1037 itabHashField = runtimeField("hash", rttype.ITab.OffsetOf("Hash"), types.Types[types.TUINT32])
1038 }
1039 hashField = itabHashField
1040 }
1041 return boundedDotPtr(pos, itab, hashField)
1042 }
1043
1044 var rtypeHashField, itabHashField *types.Field
1045
1046
1047 type typeSwitch struct {
1048
1049 srcName ir.Node
1050 hashName ir.Node
1051 okName ir.Node
1052 itabName ir.Node
1053 }
1054
1055 type typeClause struct {
1056 hash uint32
1057 body ir.Nodes
1058 }
1059
1060 func (s *typeSwitch) flush(cc []typeClause, compiled *ir.Nodes) {
1061 if len(cc) == 0 {
1062 return
1063 }
1064
1065 slices.SortFunc(cc, func(a, b typeClause) int { return cmp.Compare(a.hash, b.hash) })
1066
1067
1068 merged := cc[:1]
1069 for _, c := range cc[1:] {
1070 last := &merged[len(merged)-1]
1071 if last.hash == c.hash {
1072 last.body.Append(c.body.Take()...)
1073 } else {
1074 merged = append(merged, c)
1075 }
1076 }
1077 cc = merged
1078
1079 if s.tryJumpTable(cc, compiled) {
1080 return
1081 }
1082 binarySearch(len(cc), compiled,
1083 func(i int) ir.Node {
1084 return ir.NewBinaryExpr(base.Pos, ir.OLE, s.hashName, ir.NewInt(base.Pos, int64(cc[i-1].hash)))
1085 },
1086 func(i int, nif *ir.IfStmt) {
1087
1088
1089 c := cc[i]
1090 nif.Cond = ir.NewBinaryExpr(base.Pos, ir.OEQ, s.hashName, ir.NewInt(base.Pos, int64(c.hash)))
1091 nif.Body.Append(c.body.Take()...)
1092 },
1093 )
1094 }
1095
1096
1097 func (s *typeSwitch) tryJumpTable(cc []typeClause, out *ir.Nodes) bool {
1098 const minCases = 5
1099 if base.Flag.N != 0 || !ssagen.Arch.LinkArch.CanJumpTable || base.Ctxt.Retpoline {
1100 return false
1101 }
1102 if len(cc) < minCases {
1103 return false
1104 }
1105 hashes := make([]uint32, len(cc))
1106
1107
1108 b0 := bits.Len(uint(len(cc) - 1))
1109 for b := b0; b < b0+3; b++ {
1110 pickI:
1111 for i := 0; i <= 32-b; i++ {
1112
1113
1114 hashes = hashes[:0]
1115 for _, c := range cc {
1116 h := c.hash >> i & (1<<b - 1)
1117 hashes = append(hashes, h)
1118 }
1119
1120 slices.Sort(hashes)
1121 for j := 1; j < len(hashes); j++ {
1122 if hashes[j] == hashes[j-1] {
1123
1124 continue pickI
1125 }
1126 }
1127
1128
1129 h := s.hashName
1130 if i != 0 {
1131 h = ir.NewBinaryExpr(base.Pos, ir.ORSH, h, ir.NewInt(base.Pos, int64(i)))
1132 }
1133 h = ir.NewBinaryExpr(base.Pos, ir.OAND, h, ir.NewInt(base.Pos, int64(1<<b-1)))
1134 h = typecheck.Expr(h)
1135
1136
1137 jt := ir.NewJumpTableStmt(base.Pos, h)
1138 jt.Cases = make([]constant.Value, 1<<b)
1139 jt.Targets = make([]*types.Sym, 1<<b)
1140 out.Append(jt)
1141
1142
1143 noMatch := typecheck.AutoLabel(".s")
1144 for j := 0; j < 1<<b; j++ {
1145 jt.Cases[j] = constant.MakeInt64(int64(j))
1146 jt.Targets[j] = noMatch
1147 }
1148
1149
1150 out.Append(ir.NewBranchStmt(base.Pos, ir.OGOTO, noMatch))
1151
1152
1153 for _, c := range cc {
1154 h := c.hash >> i & (1<<b - 1)
1155 label := typecheck.AutoLabel(".s")
1156 jt.Targets[h] = label
1157 out.Append(ir.NewLabelStmt(base.Pos, label))
1158 out.Append(c.body...)
1159
1160 out.Append(ir.NewBranchStmt(base.Pos, ir.OGOTO, noMatch))
1161 }
1162
1163 out.Append(ir.NewLabelStmt(base.Pos, noMatch))
1164 return true
1165 }
1166 }
1167
1168 return false
1169 }
1170
1171
1172
1173
1174
1175
1176
1177
1178
1179
1180 func binarySearch(n int, out *ir.Nodes, less func(i int) ir.Node, leaf func(i int, nif *ir.IfStmt)) {
1181 const binarySearchMin = 4
1182
1183 var do func(lo, hi int, out *ir.Nodes)
1184 do = func(lo, hi int, out *ir.Nodes) {
1185 n := hi - lo
1186 if n < binarySearchMin {
1187 for i := lo; i < hi; i++ {
1188 nif := ir.NewIfStmt(base.Pos, nil, nil, nil)
1189 leaf(i, nif)
1190 base.Pos = base.Pos.WithNotStmt()
1191 nif.Cond = typecheck.Expr(nif.Cond)
1192 nif.Cond = typecheck.DefaultLit(nif.Cond, nil)
1193 out.Append(nif)
1194 out = &nif.Else
1195 }
1196 return
1197 }
1198
1199 half := lo + n/2
1200 nif := ir.NewIfStmt(base.Pos, nil, nil, nil)
1201 nif.Cond = less(half)
1202 base.Pos = base.Pos.WithNotStmt()
1203 nif.Cond = typecheck.Expr(nif.Cond)
1204 nif.Cond = typecheck.DefaultLit(nif.Cond, nil)
1205 do(lo, half, &nif.Body)
1206 do(half, hi, &nif.Else)
1207 out.Append(nif)
1208 }
1209
1210 do(0, n, out)
1211 }
1212
1213 func stringSearch(expr ir.Node, cc []exprClause, out *ir.Nodes) {
1214 if len(cc) < 4 {
1215
1216 for _, c := range cc {
1217 nif := ir.NewIfStmt(base.Pos.WithNotStmt(), typecheck.DefaultLit(typecheck.Expr(c.test(expr)), nil), []ir.Node{c.jmp}, nil)
1218 out.Append(nif)
1219 out = &nif.Else
1220 }
1221 return
1222 }
1223
1224
1225
1226
1227
1228
1229
1230
1231
1232
1233
1234
1235
1236
1237
1238
1239
1240
1241
1242
1243 n := len(ir.StringVal(cc[0].lo))
1244 bestScore := int64(0)
1245 bestIdx := 0
1246 bestByte := int8(0)
1247 for idx := 0; idx < n; idx++ {
1248 for b := int8(-128); b < 127; b++ {
1249 le := 0
1250 for _, c := range cc {
1251 s := ir.StringVal(c.lo)
1252 if int8(s[idx]) <= b {
1253 le++
1254 }
1255 }
1256 score := int64(le) * int64(len(cc)-le)
1257 if score > bestScore {
1258 bestScore = score
1259 bestIdx = idx
1260 bestByte = b
1261 }
1262 }
1263 }
1264
1265
1266
1267
1268 if bestScore == 0 {
1269 base.Fatalf("unable to split string set")
1270 }
1271
1272
1273 slice := ir.NewConvExpr(base.Pos, ir.OSTR2BYTESTMP, types.NewSlice(types.Types[types.TINT8]), expr)
1274 slice.SetTypecheck(1)
1275 slice.MarkNonNil()
1276
1277 load := ir.NewIndexExpr(base.Pos, slice, ir.NewInt(base.Pos, int64(bestIdx)))
1278
1279 cmp := ir.Node(ir.NewBinaryExpr(base.Pos, ir.OLE, load, ir.NewInt(base.Pos, int64(bestByte))))
1280 cmp = typecheck.DefaultLit(typecheck.Expr(cmp), nil)
1281 nif := ir.NewIfStmt(base.Pos, cmp, nil, nil)
1282
1283 var le []exprClause
1284 var gt []exprClause
1285 for _, c := range cc {
1286 s := ir.StringVal(c.lo)
1287 if int8(s[bestIdx]) <= bestByte {
1288 le = append(le, c)
1289 } else {
1290 gt = append(gt, c)
1291 }
1292 }
1293 stringSearch(expr, le, &nif.Body)
1294 stringSearch(expr, gt, &nif.Else)
1295 out.Append(nif)
1296
1297
1298 }
1299
View as plain text