summaryrefslogtreecommitdiffstats
path: root/katzenpost/dispatcher/reassembler.go
diff options
context:
space:
mode:
Diffstat (limited to 'katzenpost/dispatcher/reassembler.go')
-rw-r--r--katzenpost/dispatcher/reassembler.go91
1 files changed, 91 insertions, 0 deletions
diff --git a/katzenpost/dispatcher/reassembler.go b/katzenpost/dispatcher/reassembler.go
new file mode 100644
index 0000000..87af98c
--- /dev/null
+++ b/katzenpost/dispatcher/reassembler.go
@@ -0,0 +1,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)
+ }
+ }
+}