summaryrefslogtreecommitdiffstats
path: root/katzenpost/dispatcher/reassembler.go
blob: 87af98c5252043966f4861ed4f43b70bb3d54d6d (plain) (blame)
1
2
3
4
5
6
7
8
9
10
11
12
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
package dispatcher

import (
	"bytes"
	"crypto/sha256"
	"errors"
	"sync"
	"time"
)

const MaxPendingMessages = 64

type pendingMessage struct {
	digest  [32]byte
	total   uint16
	parts   map[uint16][]byte
	created time.Time
}

type Reassembler struct {
	mu      sync.Mutex
	pending map[[16]byte]*pendingMessage
	ttl     time.Duration
	now     func() time.Time
}

func NewReassembler(ttl time.Duration) *Reassembler {
	return &Reassembler{
		pending: make(map[[16]byte]*pendingMessage),
		ttl:     ttl,
		now:     time.Now,
	}
}

func (r *Reassembler) Add(frame Frame) ([]byte, bool, error) {
	r.mu.Lock()
	defer r.mu.Unlock()

	r.expireLocked()
	pending, ok := r.pending[frame.ID]
	if !ok {
		if len(r.pending) >= MaxPendingMessages {
			return nil, false, errors.New("too many pending messages")
		}
		pending = &pendingMessage{
			digest:  frame.Digest,
			total:   frame.Total,
			parts:   make(map[uint16][]byte, frame.Total),
			created: r.now(),
		}
		r.pending[frame.ID] = pending
	}
	if pending.total != frame.Total || pending.digest != frame.Digest {
		delete(r.pending, frame.ID)
		return nil, false, errors.New("inconsistent frame metadata")
	}
	if _, duplicate := pending.parts[frame.Index]; !duplicate {
		pending.parts[frame.Index] = bytes.Clone(frame.Data)
	}
	if len(pending.parts) != int(pending.total) {
		return nil, false, nil
	}

	var assembled bytes.Buffer
	for index := uint16(0); index < pending.total; index++ {
		part, exists := pending.parts[index]
		if !exists {
			return nil, false, nil
		}
		if assembled.Len()+len(part) > MaxEnvelopeBytes {
			delete(r.pending, frame.ID)
			return nil, false, errors.New("reassembled envelope is too large")
		}
		assembled.Write(part)
	}
	delete(r.pending, frame.ID)
	result := assembled.Bytes()
	if sha256.Sum256(result) != pending.digest {
		return nil, false, errors.New("reassembled envelope digest mismatch")
	}
	return result, true, nil
}

func (r *Reassembler) expireLocked() {
	cutoff := r.now().Add(-r.ttl)
	for id, pending := range r.pending {
		if pending.created.Before(cutoff) {
			delete(r.pending, id)
		}
	}
}