Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 17 additions & 5 deletions std/compress/io.go
Original file line number Diff line number Diff line change
Expand Up @@ -187,11 +187,23 @@ func (nr *NumReader) SetNumNbBits(numNbBits int) {
panic("decreasing wordsPerNum not supported")
}

if nr.last != nil { // nothing to compensate for if no values have yet been read
nbToRead := min(len(nr.toRead), wordsPerNum-nr.wordsPerNum)
delta := ReadNum(nr.api, nr.toRead[:nbToRead], nr.radix)
nr.toRead = nr.toRead[:nbToRead]
nr.last = nr.api.Add(nr.api.Mul(nr.last, twoPow(wordsPerNum-nr.wordsPerNum)), delta)
if nr.last != nil && wordsPerNum > nr.wordsPerNum { // nothing to compensate for if no values have yet been read
// nr.last currently holds (b₀ b₁ ... bₙ₋₁)ᵣ where n = nr.wordsPerNum and r = nr.radix.
// it must become (b₀ b₁ ... bₙ'₋₁)ᵣ where n' = wordsPerNum, i.e. the window is widened
// in place: the reader head stays put and the words bₙ ... bₙ'₋₁ are the ones sitting
// right after the current window. Words past the end of toRead are taken to be zero,
// as in next(). Note that toRead is not consumed here: the window it defines always
// starts at index 0, and next() is what shifts it.
nbNewWords := wordsPerNum - nr.wordsPerNum
nr.last = nr.api.Mul(nr.last, twoPow(nbNewWords*wordNbBits)) // × rⁿ'⁻ⁿ
nbToRead := min(max(len(nr.toRead)-nr.wordsPerNum, 0), nbNewWords)
if nbToRead > 0 {
delta := ReadNum(nr.api, nr.toRead[nr.wordsPerNum:nr.wordsPerNum+nbToRead], nr.radix)
if nbToRead < nbNewWords { // pad the missing least significant words with zeros
delta = nr.api.Mul(delta, twoPow((nbNewWords-nbToRead)*wordNbBits))
}
nr.last = nr.api.Add(nr.last, delta)
}
}

nr.wordsPerNum, nr.numBound = wordsPerNum, twoPow(numNbBits)
Expand Down
101 changes: 86 additions & 15 deletions std/compress/io_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -127,15 +127,17 @@ func TestSetNumNbBits(t *testing.T) {
runTest := func(words, increases []byte, nums []uint64) {
test.NewAssert(t).CheckCircuit(
&testSetNumNbBitsCircuit{
increases: increases,
Words: make([]frontend.Variable, len(words)),
Nums: make([]frontend.Variable, len(nums)),
increases: increases,
wordNbBits: 1,
Words: make([]frontend.Variable, len(words)),
Nums: make([]frontend.Variable, len(nums)),
},
test.WithCurves(ecc.BLS12_377), test.WithBackends(backend.PLONK),
test.WithValidAssignment(&testSetNumNbBitsCircuit{
increases: increases,
Words: internal.ToVariableSlice(words),
Nums: internal.ToVariableSlice(nums),
increases: increases,
wordNbBits: 1,
Words: internal.ToVariableSlice(words),
Nums: internal.ToVariableSlice(nums),
}))
}

Expand All @@ -149,7 +151,7 @@ func TestSetNumNbBits(t *testing.T) {
nbWordsUsed := 0 // starting with one word per num
nbWordsPerNum := 1
increases := make([]byte, 0)
for nbWordsUsed < maxNbWords && nbWordsPerNum < 64 {
for nbWordsUsed < maxNbWords && nbWordsPerNum < 64 && len(increases) < len(buf) {
inc := buf[len(increases)] % 3 % 2 // double mod to make 0 more likely TODO try other increase values too
nbWordsPerNum += int(inc)
if nbWordsPerNum >= 64 {
Expand All @@ -170,9 +172,6 @@ func TestSetNumNbBits(t *testing.T) {
}
}

words = []byte{1, 0}
increases = []byte{1}

nbWordsPerNum = 1
nums := make([]uint64, len(increases))
for i := range nums {
Expand All @@ -186,18 +185,90 @@ func TestSetNumNbBits(t *testing.T) {
}
}

// numReaderNums is a native reference implementation of NumReader: it returns the
// numbers read by a reader that starts with numNbBits bits per number and is resized
// by increases[i] bits (a multiple of wordNbBits) before each read.
// The words are assumed to lie in [0, 2^wordNbBits) and to be numerous enough that the
// sliding window never runs past the end of the slice.
func numReaderNums(words []uint64, wordNbBits, numNbBits int, increases []byte) []uint64 {
radix := uint64(1) << wordNbBits
nums := make([]uint64, len(increases))
head := -1 // the reader head is moved to 0 by the very first Next
for i, inc := range increases {
numNbBits += int(inc)
wordsPerNum := numNbBits / wordNbBits
head++
for j := 0; j < wordsPerNum; j++ {
nums[i] = nums[i]*radix + words[head+j]
}
}
return nums
}

// TestSetNumNbBitsAfterNext covers NumReader.SetNumNbBits called once the reader is
// already running, i.e. after Next has been called at least once. In that case the
// number currently held must be widened in place: (b₀ ... bₙ₋₁)ᵣ becomes
// (b₀ ... bₙ'₋₁)ᵣ, the extra words bₙ ... bₙ'₋₁ being the ones sitting right after the
// current window, and the reader head must not move.
func TestSetNumNbBitsAfterNext(t *testing.T) {
// for each read, by how many words the window is widened
schedules := [][]int{
{0, 1, 0, 0, 1, 0},
{0, 2, 0, 1, 0, 0},
{1, 0, 0, 2, 0, 1},
{0, 0, 1, 1, 1, 0},
{0, 3, 0, 0, 2, 0},
}

for _, wordNbBits := range []int{1, 2, 4, 8} {
for s, schedule := range schedules {
t.Run(fmt.Sprintf("wordNbBits=%d/schedule=%d", wordNbBits, s), func(t *testing.T) {
words := make([]uint64, len(schedule)+8)
for i := range words {
w, err := rand.Int(rand.Reader, big.NewInt(int64(1)<<wordNbBits))
assert.NoError(t, err)
words[i] = w.Uint64()
}

increases := make([]byte, len(schedule))
for i := range increases {
increases[i] = byte(schedule[i] * wordNbBits)
}

nums := numReaderNums(words, wordNbBits, wordNbBits, increases)

test.NewAssert(t).CheckCircuit(
&testSetNumNbBitsCircuit{
increases: increases,
wordNbBits: wordNbBits,
Words: make([]frontend.Variable, len(words)),
Nums: make([]frontend.Variable, len(nums)),
},
test.WithCurves(ecc.BLS12_377), test.WithBackends(backend.PLONK),
test.WithValidAssignment(&testSetNumNbBitsCircuit{
increases: increases,
wordNbBits: wordNbBits,
Words: internal.ToVariableSlice(words),
Nums: internal.ToVariableSlice(nums),
}))
})
}
}
}

type testSetNumNbBitsCircuit struct {
increases []byte
Words []frontend.Variable
Nums []frontend.Variable
increases []byte
wordNbBits int
Words []frontend.Variable
Nums []frontend.Variable
}

func (c *testSetNumNbBitsCircuit) Define(api frontend.API) error {
if len(c.increases) != len(c.Nums) {
return errors.New("must have as many steps as read values")
}
l := 1
nr := compress.NewNumReader(api, c.Words, l, 1)
l := c.wordNbBits
nr := compress.NewNumReader(api, c.Words, l, c.wordNbBits)
for i := range c.increases {
l += int(c.increases[i])
nr.SetNumNbBits(l)
Expand Down