1
2
3
4
5 package unify
6
7 import (
8 "fmt"
9 "iter"
10 "reflect"
11 "strings"
12 "sync/atomic"
13 )
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94 type envSet struct {
95 root *envExpr
96 }
97
98 type envExpr struct {
99
100
101
102
103
104
105
106 kind envExprKind
107 mask [2]uint64
108
109
110 id *ident
111 val *Value
112
113
114
115 operands []*envExpr
116 }
117
118 type envExprKind byte
119
120 const (
121 envZero envExprKind = iota
122 envUnit
123 envProduct
124 envSum
125 envBinding
126 )
127
128 var (
129
130 topEnv = envSet{envExprUnit}
131
132 bottomEnv = envSet{envExprZero}
133
134 envExprZero = &envExpr{kind: envZero}
135 envExprUnit = &envExpr{kind: envUnit}
136 )
137
138
139
140
141
142
143
144 func (e envSet) bind(id *ident, vals ...*Value) envSet {
145 if e.isEmpty() {
146 return bottomEnv
147 }
148
149
150
151
152
153
154 for range e.root.bindings(id) {
155 panic("id " + id.name + " already present in environment")
156 }
157
158
159 bindings := make([]*envExpr, 0, 1)
160 for _, val := range vals {
161 bindings = append(bindings, &envExpr{kind: envBinding, mask: id.mask, id: id, val: val})
162 }
163
164
165 return envSet{newEnvExprProduct(e.root, newEnvExprSum(bindings...))}
166 }
167
168 func (e envSet) isEmpty() bool {
169 return e.root.kind == envZero
170 }
171
172
173
174 func (e *envExpr) bindings(id *ident) iter.Seq[*envExpr] {
175
176
177 return func(yield func(*envExpr) bool) {
178 var rec func(e *envExpr) bool
179 rec = func(e *envExpr) bool {
180 if id != nil && (e.mask[0]&id.mask[0] == 0 && e.mask[1]&id.mask[1] == 0) {
181 return true
182 }
183 if e.kind == envBinding && (id == nil || e.id == id) {
184 if !yield(e) {
185 return false
186 }
187 }
188 for _, o := range e.operands {
189 if !rec(o) {
190 return false
191 }
192 }
193 return true
194 }
195 rec(e)
196 }
197 }
198
199
200
201 func newEnvExprProduct(exprs ...*envExpr) *envExpr {
202 factors := make([]*envExpr, 0, 2)
203 for _, expr := range exprs {
204 switch expr.kind {
205 case envZero:
206 return envExprZero
207 case envUnit:
208
209 case envProduct:
210 factors = append(factors, expr.operands...)
211 default:
212 factors = append(factors, expr)
213 }
214 }
215
216 if len(factors) == 0 {
217 return envExprUnit
218 } else if len(factors) == 1 {
219 return factors[0]
220 }
221 var mask [2]uint64
222 for _, f := range factors {
223 mask[0] |= f.mask[0]
224 mask[1] |= f.mask[1]
225 }
226 return &envExpr{kind: envProduct, mask: mask, operands: factors}
227 }
228
229
230 func newEnvExprSum(exprs ...*envExpr) *envExpr {
231
232
233
234
235 var have smallSet[*envExpr]
236 terms := make([]*envExpr, 0, 2)
237 for _, expr := range exprs {
238 switch expr.kind {
239 case envZero:
240
241 case envSum:
242 for _, expr1 := range expr.operands {
243 if have.Add(expr1) {
244 terms = append(terms, expr1)
245 }
246 }
247 default:
248 if have.Add(expr) {
249 terms = append(terms, expr)
250 }
251 }
252 }
253
254 if len(terms) == 0 {
255 return envExprZero
256 } else if len(terms) == 1 {
257 return terms[0]
258 }
259 var mask [2]uint64
260 for _, t := range terms {
261 mask[0] |= t.mask[0]
262 mask[1] |= t.mask[1]
263 }
264 return &envExpr{kind: envSum, mask: mask, operands: terms}
265 }
266
267 func crossEnvs(env1, env2 envSet) envSet {
268
269 var ids1 smallSet[*ident]
270 for e := range env1.root.bindings(nil) {
271 ids1.Add(e.id)
272 }
273 for e := range env2.root.bindings(nil) {
274 if ids1.Has(e.id) {
275 panic(fmt.Sprintf("%s bound on both sides of cross-product", e.id.name))
276 }
277 }
278
279 return envSet{newEnvExprProduct(env1.root, env2.root)}
280 }
281
282 func unionEnvs(envs ...envSet) envSet {
283 exprs := make([]*envExpr, len(envs))
284 for i := range envs {
285 exprs[i] = envs[i].root
286 }
287 return envSet{newEnvExprSum(exprs...)}
288 }
289
290
291
292 type envPartition struct {
293 id *ident
294 value *Value
295 env envSet
296 }
297
298
299
300
301
302
303
304
305
306 func (e envSet) partitionBy(id *ident) []envPartition {
307 if e.isEmpty() {
308
309
310 panic("cannot partition empty environment set")
311 }
312
313
314 var seen smallSet[*Value]
315 var parts []envPartition
316 for n := range e.root.bindings(id) {
317 if !seen.Add(n.val) {
318
319 continue
320 }
321
322 parts = append(parts, envPartition{
323 id: id,
324 value: n.val,
325 env: envSet{e.root.substitute(id, n.val)},
326 })
327 }
328
329 return parts
330 }
331
332
333
334 func (e *envExpr) substitute(id *ident, val *Value) *envExpr {
335 if e.mask[0]&id.mask[0] == 0 && e.mask[1]&id.mask[1] == 0 {
336 return e
337 }
338 switch e.kind {
339 default:
340 panic("bad kind")
341
342 case envZero, envUnit:
343 return e
344
345 case envBinding:
346 if e.id != id {
347 return e
348 } else if e.val != val {
349 return envExprZero
350 } else {
351 return envExprUnit
352 }
353
354 case envProduct, envSum:
355
356
357 var nOperands []*envExpr
358 for i, op := range e.operands {
359 nOp := op.substitute(id, val)
360 if nOperands == nil && op != nOp {
361
362 nOperands = make([]*envExpr, 0, len(e.operands))
363 nOperands = append(nOperands, e.operands[:i]...)
364 }
365 if nOperands != nil {
366 nOperands = append(nOperands, nOp)
367 }
368 }
369 if nOperands == nil {
370
371 return e
372 }
373 if e.kind == envProduct {
374 return newEnvExprProduct(nOperands...)
375 } else {
376 return newEnvExprSum(nOperands...)
377 }
378 }
379 }
380
381
382 type smallSet[T comparable] struct {
383 array [32]T
384 n int
385
386 m map[T]struct{}
387 }
388
389
390 func (s *smallSet[T]) Has(val T) bool {
391 arr := s.array[:s.n]
392 for i := range arr {
393 if arr[i] == val {
394 return true
395 }
396 }
397 _, ok := s.m[val]
398 return ok
399 }
400
401
402
403 func (s *smallSet[T]) Add(val T) bool {
404
405 if s.Has(val) {
406 return false
407 }
408
409
410 if s.n < len(s.array) {
411 s.array[s.n] = val
412 s.n++
413 } else {
414 if s.m == nil {
415 s.m = make(map[T]struct{})
416 }
417 s.m[val] = struct{}{}
418 }
419 return true
420 }
421
422 type ident struct {
423 _ [0]func()
424 mask [2]uint64
425 name string
426 }
427
428 var identCounter atomic.Uint64
429
430 func newIdent(name string) *ident {
431 c := identCounter.Add(1)
432 var mask [2]uint64
433 bit := c % 128
434 mask[bit/64] = 1 << (bit % 64)
435 return &ident{
436 mask: mask,
437 name: name,
438 }
439 }
440
441 type Var struct {
442 id *ident
443 }
444
445 func (d Var) Exact() bool {
446
447 panic("Exact called on non-concrete Value")
448 }
449
450 func (d Var) WhyNotExact() string {
451
452 return "WhyNotExact called on non-concrete Value"
453 }
454
455 func (d Var) decode(rv reflect.Value) error {
456 return &inexactError{"var", rv.Type().String()}
457 }
458
459 func (d Var) unify(w *Value, e envSet, swap bool, uf *unifier) (Domain, envSet, error) {
460
461
462
463
464
465
466
467
468 if vd, ok := w.Domain.(Var); ok && d.id == vd.id {
469
470
471
472
473
474 return vd, e, nil
475 }
476
477
478
479
480 var nEnvs []envSet
481 envParts := e.partitionBy(d.id)
482 for i, envPart := range envParts {
483 exit := uf.enterVar(d.id, i)
484
485
486
487 res, e2, err := w.unify(envPart.value, envPart.env, swap, uf)
488 exit.exit()
489 if err != nil {
490 return nil, envSet{}, err
491 }
492 if res.Domain == nil {
493
494 continue
495 }
496 nEnv := e2.bind(d.id, res)
497 nEnvs = append(nEnvs, nEnv)
498 }
499
500 if len(nEnvs) == 0 {
501
502 return nil, bottomEnv, nil
503 }
504
505
506
507 return d, unionEnvs(nEnvs...), nil
508 }
509
510
511 type identPrinter struct {
512 ids map[*ident]string
513 idGen map[string]int
514 }
515
516 func (p *identPrinter) unique(id *ident) string {
517 if p.ids == nil {
518 p.ids = make(map[*ident]string)
519 p.idGen = make(map[string]int)
520 }
521
522 name, ok := p.ids[id]
523 if !ok {
524 gen := p.idGen[id.name]
525 p.idGen[id.name]++
526 if gen == 0 {
527 name = id.name
528 } else {
529 name = fmt.Sprintf("%s#%d", id.name, gen)
530 }
531 p.ids[id] = name
532 }
533
534 return name
535 }
536
537 func (p *identPrinter) slice(ids []*ident) string {
538 var strs []string
539 for _, id := range ids {
540 strs = append(strs, p.unique(id))
541 }
542 return fmt.Sprintf("[%s]", strings.Join(strs, ", "))
543 }
544
View as plain text