package main

import (
	"bytes"
	"crypto/sha256"
	"fmt"
	"math/big"
	"os"
	"time"

	"github.com/consensys/gnark-crypto/ecc"
	"github.com/consensys/gnark-crypto/ecc/bn254/fr"
	fmimc "github.com/consensys/gnark-crypto/ecc/bn254/fr/mimc"
	"github.com/consensys/gnark/backend/groth16"
	"github.com/consensys/gnark/backend/plonk"
	"github.com/consensys/gnark/frontend"
	"github.com/consensys/gnark/frontend/cs/r1cs"
	"github.com/consensys/gnark/frontend/cs/scs"
	"github.com/consensys/gnark/std/hash/mimc"
	"github.com/consensys/gnark/std/hash/sha2"
	"github.com/consensys/gnark/std/math/uints"
	"github.com/consensys/gnark/test/unsafekzg"
)

// 1. Idade: nascimento secreto, ano atual público, prova de ano-nasc >= 18
type Idade struct {
	Nasc  frontend.Variable
	Ano   frontend.Variable `gnark:",public"`
	Salt  frontend.Variable
	Compr frontend.Variable `gnark:",public"` // compromisso do documento: MiMC(nasc, salt)
}

func (c *Idade) Define(api frontend.API) error {
	h, _ := mimc.NewMiMC(api)
	h.Write(c.Nasc, c.Salt)
	api.AssertIsEqual(h.Sum(), c.Compr)
	api.AssertIsLessOrEqual(api.Add(c.Nasc, 18), c.Ano)
	return nil
}

// 2. Saldo >= X, com o saldo comprometido
type Saldo struct {
	S     frontend.Variable
	Salt  frontend.Variable
	X     frontend.Variable `gnark:",public"`
	Compr frontend.Variable `gnark:",public"`
}

func (c *Saldo) Define(api frontend.API) error {
	h, _ := mimc.NewMiMC(api)
	h.Write(c.S, c.Salt)
	api.AssertIsEqual(h.Sum(), c.Compr)
	api.ToBinary(c.S, 64)
	api.AssertIsLessOrEqual(c.X, c.S)
	return nil
}

// 3. Pertence a uma lista de 2^20 (árvore de Merkle), sem dizer qual item
const D = 20

type Lista struct {
	Folha frontend.Variable
	Irmao [D]frontend.Variable
	Lado  [D]frontend.Variable
	Raiz  frontend.Variable `gnark:",public"`
}

func (c *Lista) Define(api frontend.API) error {
	cur := c.Folha
	for i := 0; i < D; i++ {
		api.AssertIsBoolean(c.Lado[i])
		l := api.Select(c.Lado[i], c.Irmao[i], cur)
		r := api.Select(c.Lado[i], cur, c.Irmao[i])
		h, _ := mimc.NewMiMC(api)
		h.Write(l, r)
		cur = h.Sum()
	}
	api.AssertIsEqual(cur, c.Raiz)
	return nil
}

// 4. Conheço um texto de 64 bytes cujo SHA-256 é este
type Sha struct {
	In  [64]uints.U8
	Out [32]uints.U8 `gnark:",public"`
}

func (c *Sha) Define(api frontend.API) error {
	h, err := sha2.New(api)
	if err != nil {
		return err
	}
	u, _ := uints.New[uints.U32](api)
	h.Write(c.In[:])
	res := h.Sum()
	for i := range c.Out {
		u.ByteAssertEq(c.Out[i], res[i])
	}
	return nil
}

func mimcHash(xs ...*big.Int) *big.Int {
	h := fmimc.NewMiMC()
	for _, x := range xs {
		var e fr.Element
		e.SetBigInt(x)
		b := e.Bytes()
		h.Write(b[:])
	}
	return new(big.Int).SetBytes(h.Sum(nil))
}

type caso struct {
	nome    string
	circ    frontend.Circuit
	witness frontend.Circuit
}

func casos() []caso {
	var out []caso
	// idade
	nasc, salt := big.NewInt(1990), big.NewInt(987654321)
	out = append(out, caso{"idade", &Idade{}, &Idade{Nasc: nasc, Ano: 2026, Salt: salt, Compr: mimcHash(nasc, salt)}})
	// saldo
	s, s2 := big.NewInt(12500), big.NewInt(55555)
	out = append(out, caso{"saldo", &Saldo{}, &Saldo{S: s, Salt: s2, X: 10000, Compr: mimcHash(s, s2)}})
	// lista
	var w Lista
	cur := big.NewInt(424242)
	w.Folha = cur
	for i := 0; i < D; i++ {
		sib := big.NewInt(int64(1000 + i))
		w.Irmao[i] = sib
		if i%2 == 0 {
			w.Lado[i] = 0
			cur = mimcHash(cur, sib)
		} else {
			w.Lado[i] = 1
			cur = mimcHash(sib, cur)
		}
	}
	w.Raiz = cur
	out = append(out, caso{"lista", &Lista{}, &w})
	// sha
	msg := bytes.Repeat([]byte("o segredo "), 7)[:64]
	dig := sha256.Sum256(msg)
	var sw Sha
	for i := range msg {
		sw.In[i] = uints.NewU8(msg[i])
	}
	for i := range dig {
		sw.Out[i] = uints.NewU8(dig[i])
	}
	out = append(out, caso{"sha256", &Sha{}, &sw})
	return out
}

func best(n int, f func()) time.Duration {
	var b time.Duration
	for i := 0; i < n; i++ {
		t := time.Now()
		f()
		d := time.Since(t)
		if i == 0 || d < b {
			b = d
		}
	}
	return b
}

func main() {
	for _, c := range casos() {
		wit, err := frontend.NewWitness(c.witness, ecc.BN254.ScalarField())
		if err != nil {
			panic(err)
		}
		pub, _ := wit.Public()
		// Groth16
		cs, err := frontend.Compile(ecc.BN254.ScalarField(), r1cs.NewBuilder, c.circ)
		if err != nil {
			panic(err)
		}
		var pk groth16.ProvingKey
		var vk groth16.VerifyingKey
		tSetup := best(1, func() { pk, vk, _ = groth16.Setup(cs) })
		var proof groth16.Proof
		tProve := best(5, func() { proof, err = groth16.Prove(cs, pk, wit); if err != nil { panic(err) } })
		tVer := best(20, func() { if e := groth16.Verify(proof, vk, pub); e != nil { panic(e) } })
		var buf bytes.Buffer
		proof.WriteTo(&buf)
		var pkb bytes.Buffer
		pk.WriteTo(&pkb)
		fmt.Printf("%s\tgroth16\trestricoes=%d\tsetup=%v\tprova=%v\tverifica=%v\ttam_prova=%dB\tchave_prova=%dB\n", c.nome, cs.GetNbConstraints(), tSetup.Round(time.Millisecond), tProve.Round(time.Millisecond), tVer.Round(10*time.Microsecond), buf.Len(), pkb.Len())
		// PLONK
		scsCS, err := frontend.Compile(ecc.BN254.ScalarField(), scs.NewBuilder, c.circ)
		if err != nil {
			panic(err)
		}
		srs, srsL, err := unsafekzg.NewSRS(scsCS)
		if err != nil {
			panic(err)
		}
		var ppk plonk.ProvingKey
		var pvk plonk.VerifyingKey
		tSetup2 := best(1, func() { ppk, pvk, err = plonk.Setup(scsCS, srs, srsL); if err != nil { panic(err) } })
		var pproof plonk.Proof
		tProve2 := best(3, func() { pproof, err = plonk.Prove(scsCS, ppk, wit); if err != nil { panic(err) } })
		tVer2 := best(20, func() { if e := plonk.Verify(pproof, pvk, pub); e != nil { panic(e) } })
		var buf2 bytes.Buffer
		pproof.WriteTo(&buf2)
		fmt.Printf("%s\tplonk\trestricoes=%d\tsetup=%v\tprova=%v\tverifica=%v\ttam_prova=%dB\n", c.nome, scsCS.GetNbConstraints(), tSetup2.Round(time.Millisecond), tProve2.Round(time.Millisecond), tVer2.Round(10*time.Microsecond), buf2.Len())
		// prova falsa: testemunha errada tem de falhar
	}
	fmt.Fprintln(os.Stderr, "ok")
}
