162 lines
3 KiB
Go
162 lines
3 KiB
Go
package crypto
|
|
|
|
import (
|
|
"cachefs/pkg/chunk"
|
|
"crypto/cipher"
|
|
"errors"
|
|
"fmt"
|
|
"io"
|
|
|
|
"golang.org/x/crypto/chacha20poly1305"
|
|
)
|
|
|
|
type cryptoWriter interface {
|
|
io.Writer
|
|
io.WriterAt
|
|
}
|
|
|
|
const lastChunkFlag = 0x01
|
|
|
|
func setLastChunkFlag(nonce *[NonceSize]byte) {
|
|
nonce[len(nonce)-1] = lastChunkFlag
|
|
}
|
|
|
|
type writer struct {
|
|
a cipher.AEAD
|
|
dst cryptoWriter
|
|
unwritten []byte // backed by buf
|
|
buf [FullSize]byte
|
|
nonce [NonceSize]byte
|
|
err error
|
|
cm map[int64]*chunkWriter
|
|
last *chunkWriter
|
|
flush bool
|
|
}
|
|
|
|
func newWriter(key []byte, dst cryptoWriter) (*writer, error) {
|
|
aead, err := chacha20poly1305.New(key)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
w := &writer{
|
|
a: aead,
|
|
dst: dst,
|
|
cm: make(map[int64]*chunkWriter),
|
|
}
|
|
w.unwritten = w.buf[:0]
|
|
return w, nil
|
|
}
|
|
|
|
type chunkWriter struct {
|
|
offset int64
|
|
cn int64
|
|
w *writer
|
|
buf [ChunkSize]byte
|
|
nonce [NonceSize]byte
|
|
chunks chunk.Chunks
|
|
}
|
|
|
|
func (c *chunkWriter) writeAt(p []byte, offset int64) (n int, err error) {
|
|
n = copy(c.buf[offset:], p)
|
|
c.chunks.Add(offset, n)
|
|
return
|
|
}
|
|
|
|
func (c *chunkWriter) full() bool {
|
|
if len(c.chunks) == 1 {
|
|
return c.chunks[0][1] == ChunkSize
|
|
}
|
|
return false
|
|
}
|
|
|
|
func (c *chunkWriter) flush(last bool) (n int, err error) {
|
|
wb := c.buf[:0]
|
|
if last {
|
|
fmt.Println("last nonce", c.nonce)
|
|
setLastChunkFlag(&c.nonce)
|
|
fmt.Println("last nonce flag", c.nonce)
|
|
s := c.chunks.Size()
|
|
wb = c.buf[:s]
|
|
n = int(s)
|
|
} else {
|
|
wb = c.buf[:]
|
|
n = len(c.buf)
|
|
}
|
|
buf := c.w.a.Seal(wb[:0], c.nonce[:], wb, nil)
|
|
_, err = c.w.dst.WriteAt(buf, c.offset)
|
|
return
|
|
}
|
|
|
|
func (w *writer) Write(p []byte) (n int, err error) {
|
|
// TODO: consider refactoring with a bytes.Buffer.
|
|
if w.err != nil {
|
|
return 0, w.err
|
|
}
|
|
if len(p) == 0 {
|
|
return 0, nil
|
|
}
|
|
|
|
total := len(p)
|
|
for len(p) > 0 {
|
|
freeBuf := w.buf[len(w.unwritten):ChunkSize]
|
|
n := copy(freeBuf, p)
|
|
p = p[n:]
|
|
w.unwritten = w.unwritten[:len(w.unwritten)+n]
|
|
w.flush = true
|
|
|
|
if len(w.unwritten) == ChunkSize && len(p) > 0 {
|
|
if err := w.flushChunk(notLastChunk); err != nil {
|
|
w.err = err
|
|
return 0, err
|
|
}
|
|
}
|
|
}
|
|
return total, nil
|
|
}
|
|
|
|
func (w *writer) Close() error {
|
|
fmt.Println("writer closed", w.err)
|
|
if w.err != nil {
|
|
return w.err
|
|
}
|
|
|
|
if w.last != nil {
|
|
fmt.Println("flushing last chunk")
|
|
_, w.err = w.last.flush(lastChunk)
|
|
if w.err != nil {
|
|
fmt.Println("error flushing last chunk", w.err)
|
|
return w.err
|
|
}
|
|
w.flush = false
|
|
}
|
|
|
|
if w.flush {
|
|
w.err = w.flushChunk(lastChunk)
|
|
if w.err != nil {
|
|
return w.err
|
|
}
|
|
}
|
|
|
|
w.err = errors.New("stream.Writer is already closed")
|
|
return nil
|
|
}
|
|
|
|
const (
|
|
lastChunk = true
|
|
notLastChunk = false
|
|
)
|
|
|
|
func (w *writer) flushChunk(last bool) error {
|
|
if !last && len(w.unwritten) != ChunkSize {
|
|
panic("stream: internal error: flush called with partial chunk")
|
|
}
|
|
|
|
if last {
|
|
setLastChunkFlag(&w.nonce)
|
|
}
|
|
buf := w.a.Seal(w.buf[:0], w.nonce[:], w.unwritten, nil)
|
|
_, err := w.dst.Write(buf)
|
|
w.unwritten = w.buf[:0]
|
|
incNonce(&w.nonce)
|
|
return err
|
|
}
|