// Compare sync.Map implementations by allocation (clock-free) and, as a
// secondary signal, by wall time.
//
//	go build -o syncmap_trie .                                 # Go 1.24+ default: HashTrieMap
//	GOEXPERIMENT=nosynchashtriemap go build -o syncmap_old .   # read/dirty implementation
package main

import (
	"flag"
	"fmt"
	"runtime"
	"sync"
	"time"
)

type kv interface {
	Load(k any) (any, bool)
	Store(k, v any)
}

type rwMap struct {
	mu sync.RWMutex
	m  map[any]any
}

func (r *rwMap) Load(k any) (any, bool) {
	r.mu.RLock()
	v, ok := r.m[k]
	r.mu.RUnlock()
	return v, ok
}

func (r *rwMap) Store(k, v any) {
	r.mu.Lock()
	r.m[k] = v
	r.mu.Unlock()
}

func newMap(name string) kv {
	if name == "rwmutex" {
		return &rwMap{m: map[any]any{}}
	}
	return &sync.Map{}
}

// keys are boxed once up front so that boxing is not counted as map cost
func makeKeys(n int) []any {
	ks := make([]any, n)
	for i := range ks {
		ks[i] = i + 1<<20
	}
	return ks
}

type result struct {
	bytesPerOp, allocsPerOp float64
	nsPerOp               float64
}

func measure(ops int, f func()) result {
	runtime.GC()
	var a, b runtime.MemStats
	runtime.ReadMemStats(&a)
	t0 := time.Now()
	f()
	el := time.Since(t0)
	runtime.ReadMemStats(&b)
	return result{
		float64(b.TotalAlloc-a.TotalAlloc) / float64(ops),
		float64(b.Mallocs-a.Mallocs) / float64(ops),
		float64(el.Nanoseconds()) / float64(ops),
	}
}

func main() {
	impl := flag.String("impl", "syncmap", "syncmap or rwmutex")
	n := flag.Int("n", 1<<18, "number of keys")
	flag.Parse()
	keys := makeKeys(*n)
	val := any(struct{}{})

	// W1: grow. Store a new key, then load it once (cache fill pattern).
	m := newMap(*impl)
	r := measure(*n, func() {
		for _, k := range keys {
			m.Store(k, val)
			m.Load(k)
		}
	})
	fmt.Printf("impl=%s workload=grow n=%d bytes/op=%.1f allocs/op=%.3f ns/op=%.1f\n", *impl, *n, r.bytesPerOp, r.allocsPerOp, r.nsPerOp)

	// W2: read-only loads of existing keys on the filled map.
	loads := 4 * *n
	r = measure(loads, func() {
		for i := 0; i < loads; i++ {
			m.Load(keys[(i*7919)%*n])
		}
	})
	fmt.Printf("impl=%s workload=load n=%d bytes/op=%.1f allocs/op=%.3f ns/op=%.1f\n", *impl, *n, r.bytesPerOp, r.allocsPerOp, r.nsPerOp)

	// W3: two goroutines overwrite disjoint halves of the existing keys.
	r = measure(2**n, func() {
		var wg sync.WaitGroup
		for g := 0; g < 2; g++ {
			wg.Add(1)
			go func(g int) {
				defer wg.Done()
				for i := g; i < *n; i += 2 {
					m.Store(keys[i], val)
					m.Store(keys[i], val)
				}
			}(g)
		}
		wg.Wait()
	})
	fmt.Printf("impl=%s workload=overwrite-disjoint n=%d bytes/op=%.1f allocs/op=%.3f ns/op=%.1f\n", *impl, *n, r.bytesPerOp, r.allocsPerOp, r.nsPerOp)
}
