secio: buffer remainders in calls to Read()
used to return io.ErrShortBuffer, but this makes client code much more complicated. we're already allocating buffers when it's too large, so might as well just keep it for later.
Juan Batiz-Benet committed
Jan 1, 2015 at 07:00 UTC
293ea03eb034cb354359b641f9cde805ead75b45
1 file changed
+32
-9
crypto/secio/rw.go
+32
-9
@@ -76,6 +76,9 @@ type etmReader struct {
76
msgio.Reader
77
io.Closer
78
79
+ // buffer
80
+ buf []byte
81
+
82
// params
83
msg msgio.ReadCloser // msgio for knowing where boundaries lie
84
str cipher.Stream // the stream cipher to encrypt with
@@ -91,20 +94,35 @@ func (r *etmReader) NextMsgLen() (int, error) {
94
return r.msg.NextMsgLen()
95
}
96
97
+func (r *etmReader) drainBuf(buf []byte) int {
98
+ if r.buf == nil {
99
+ return 0
100
+ }
101
+
102
+ n := copy(buf, r.buf)
103
+ r.buf = r.buf[n:]
104
+ return n
105
+}
106
+
107
func (r *etmReader) Read(buf []byte) (int, error) {
95
- // first, check the buffer has enough space.
108
+ // first, check if we have anything in the buffer
109
+ copied := r.drainBuf(buf)
110
+ buf = buf[copied:]
111
+ if copied > 0 {
112
+ return copied, nil
113
+ // return here to avoid complicating the rest...
114
+ // user can call io.ReadFull.
115
+ }
116
+
117
+ // check the buffer has enough space for the next msg
118
fullLen, err := r.msg.NextMsgLen()
119
if err != nil {
120
return 0, err
121
}
122
101
- dataLen := fullLen - r.mac.size
102
- if cap(buf) < dataLen {
103
- return 0, io.ErrShortBuffer
104
- }
105
-
123
buf2 := buf
124
changed := false
125
+ // if not enough space, allocate a new buffer.
126
if cap(buf) < fullLen {
127
buf2 = make([]byte, fullLen)
128
changed = true
@@ -121,10 +139,15 @@ func (r *etmReader) Read(buf []byte) (int, error) {
139
return 0, err
140
}
141
buf2 = buf2[:m]
124
- if changed {
125
- return copy(buf, buf2), nil
142
+ if !changed {
143
+ return m, nil
144
+ }
145
+
146
+ n = copy(buf, buf2)
147
+ if len(buf2) > len(buf) {
148
+ r.buf = buf2[len(buf):] // had some left over? save it.
149
}
127
- return m, nil
150
+ return n, nil
151
}
152
153
func (r *etmReader) ReadMsg() ([]byte, error) {