Selaa lähdekoodia

ensure full FLAP payload is read

Mike 2 vuotta sitten
vanhempi
commit
155b2b59bb
4 muutettua tiedostoa jossa 18 lisäystä ja 10 poistoa
  1. 5 3
      server/oscar/connection.go
  2. 8 4
      wire/decode.go
  3. 1 1
      wire/decode_test.go
  4. 4 2
      wire/frames.go

+ 5 - 3
server/oscar/connection.go

@@ -147,9 +147,11 @@ func consumeFLAPFrames(r io.Reader, msgCh chan incomingMessage, errCh chan error
 
 		if in.flap.FrameType == wire.FLAPFrameData {
 			buf := make([]byte, in.flap.PayloadLength)
-			if _, err := r.Read(buf); err != nil {
-				errCh <- err
-				return
+			if in.flap.PayloadLength > 0 {
+				if _, err := io.ReadFull(r, buf); err != nil {
+					errCh <- err
+					return
+				}
 			}
 			in.payload = bytes.NewBuffer(buf)
 		}

+ 8 - 4
wire/decode.go

@@ -47,8 +47,10 @@ func unmarshal(t reflect.Type, v reflect.Value, tag reflect.StructTag, r io.Read
 			return fmt.Errorf("%w: missing len_prefix tag", ErrUnmarshalFailure)
 		}
 		buf := make([]byte, bufLen)
-		if _, err := r.Read(buf); err != nil {
-			return err
+		if bufLen > 0 {
+			if _, err := io.ReadFull(r, buf); err != nil {
+				return err
+			}
 		}
 		// todo is there a more efficient way?
 		v.SetString(string(buf))
@@ -102,8 +104,10 @@ func unmarshal(t reflect.Type, v reflect.Value, tag reflect.StructTag, r io.Read
 			}
 
 			buf := make([]byte, bufLen)
-			if _, err := r.Read(buf); err != nil {
-				return err
+			if bufLen > 0 {
+				if _, err := io.ReadFull(r, buf); err != nil {
+					return err
+				}
 			}
 			b := bytes.NewBuffer(buf)
 			slice := reflect.New(v.Type()).Elem()

+ 1 - 1
wire/decode_test.go

@@ -190,7 +190,7 @@ func TestUnmarshal(t *testing.T) {
 				Val []int `len_prefix:"uint8"`
 			}{},
 			wantErr: ErrUnmarshalFailure,
-			given:   []byte{0x68, 0x65, 0x6c, 0x6c, 0x6f},
+			given:   []byte{0x04, 0x65, 0x6c, 0x6c, 0x6f},
 		},
 		{
 			name: "byte slice with uint8 len_prefix with read error",

+ 4 - 2
wire/frames.go

@@ -27,8 +27,10 @@ type FLAPFrame struct {
 
 func (f FLAPFrame) ReadBody(r io.Reader) (*bytes.Buffer, error) {
 	b := make([]byte, f.PayloadLength)
-	if _, err := r.Read(b); err != nil {
-		return nil, err
+	if f.PayloadLength > 0 {
+		if _, err := io.ReadFull(r, b); err != nil {
+			return nil, err
+		}
 	}
 	return bytes.NewBuffer(b), nil
 }