diff --git a/internal/format/format.go b/internal/format/format.go index 794b48a..b274653 100644 --- a/internal/format/format.go +++ b/internal/format/format.go @@ -112,7 +112,10 @@ func (r *Stanza) Marshal(w io.Writer) error { if _, err := w.Write(stanzaPrefix); err != nil { return err } - for _, a := range append([]string{r.Type}, r.Args...) { + if _, err := io.WriteString(w, " "+r.Type); err != nil { + return err + } + for _, a := range r.Args { if _, err := io.WriteString(w, " "+a); err != nil { return err } @@ -154,7 +157,9 @@ func (h *Header) Marshal(w io.Writer) error { } type StanzaReader struct { - r *bufio.Reader + r interface { + ReadBytes(delim byte) ([]byte, error) + } err error } @@ -182,11 +187,6 @@ func (r *StanzaReader) ReadStanza() (s *Stanza, err error) { if prefix != string(stanzaPrefix) || len(args) < 1 { return nil, fmt.Errorf("malformed stanza: %q", line) } - for _, a := range args { - if !isValidString(a) { - return nil, fmt.Errorf("malformed stanza: %q", line) - } - } s.Type = args[0] s.Args = args[1:] @@ -218,6 +218,44 @@ type ParseError struct { err error } +const ( + maxHeaderBytes = 2 << 20 + maxRecipientStanzas = 1024 + maxStanzaArgs = 128 +) + +type headerReader struct { + r *bufio.Reader + n int +} + +func (r *headerReader) ReadBytes(delim byte) ([]byte, error) { + var line []byte + for { + frag, err := r.r.ReadSlice(delim) + if len(frag) > maxHeaderBytes-r.n { + return nil, errorf("header exceeds 2 MiB") + } + r.n += len(frag) + line = append(line, frag...) + if err != bufio.ErrBufferFull { + return line, err + } + } +} + +func (r *headerReader) ReadString(delim byte) (string, error) { + line, err := r.ReadBytes(delim) + return string(line), err +} + +func (r *headerReader) Peek(n int) ([]byte, error) { + if r.n+n > maxHeaderBytes { + return nil, errorf("header exceeds 2 MiB") + } + return r.r.Peek(n) +} + func (e *ParseError) Error() string { return "parsing age header: " + e.err.Error() } @@ -235,26 +273,35 @@ func errorf(format string, a ...any) error { func Parse(input io.Reader) (*Header, io.Reader, error) { h := &Header{} rr := bufio.NewReader(input) + hr := &headerReader{r: rr} - line, err := rr.ReadString('\n') + line, err := hr.ReadString('\n') if err == io.EOF { return nil, nil, errorf("file is empty") } else if err != nil { + // headerReader errors are already ParseErrors; don't nest the prefix. + if _, ok := err.(*ParseError); ok { + return nil, nil, err + } return nil, nil, errorf("failed to read intro: %w", err) } if line != intro { return nil, nil, errorf("unexpected intro: %q", line) } - sr := NewStanzaReader(rr) + sr := &StanzaReader{r: hr} for { - peek, err := rr.Peek(len(footerPrefix)) + peek, err := hr.Peek(len(footerPrefix)) if err != nil { + // headerReader errors are already ParseErrors; don't nest the prefix. + if _, ok := err.(*ParseError); ok { + return nil, nil, err + } return nil, nil, errorf("failed to read header: %w", err) } if bytes.Equal(peek, footerPrefix) { - line, err := rr.ReadBytes('\n') + line, err := hr.ReadBytes('\n') if err != nil { return nil, nil, fmt.Errorf("failed to read header: %w", err) } @@ -269,6 +316,9 @@ func Parse(input io.Reader) (*Header, io.Reader, error) { } break } + if len(h.Recipients) == maxRecipientStanzas { + return nil, nil, errorf("header contains more than %d recipient stanzas", maxRecipientStanzas) + } s, err := sr.ReadStanza() if err != nil { @@ -294,8 +344,19 @@ func Parse(input io.Reader) (*Header, io.Reader, error) { func splitArgs(line []byte) (string, []string) { l := strings.TrimSuffix(string(line), "\n") - parts := strings.Split(l, " ") - return parts[0], parts[1:] + prefix, rest, ok := strings.Cut(l, " ") + if !ok { + return l, nil + } + + var args []string + for arg := range strings.SplitSeq(rest, " ") { + if !isValidString(arg) || len(args) > maxStanzaArgs { + return l, nil + } + args = append(args, arg) + } + return prefix, args } func isValidString(s string) bool { diff --git a/internal/format/format_test.go b/internal/format/format_test.go index dca10e6..74344db 100644 --- a/internal/format/format_test.go +++ b/internal/format/format_test.go @@ -7,10 +7,13 @@ package format_test import ( + "bufio" "bytes" + "errors" "io" "os" "path/filepath" + "strings" "testing" "filippo.io/age/internal/format" @@ -43,6 +46,105 @@ func TestStanzaMarshal(t *testing.T) { } } +func TestParseLimits(t *testing.T) { + const ( + intro = "age-encryption.org/v1\n" + maxHeaderBytes = 2 << 20 + ) + footer := "--- " + format.EncodeToString(make([]byte, 32)) + "\n" + + t.Run("header size", func(t *testing.T) { + makeHeader := func(size int) []byte { + const opening = "-> test " + fixed := len(intro) + len(opening) + len("\n\n") + len(footer) + return []byte(intro + opening + strings.Repeat("a", size-fixed) + "\n\n" + footer) + } + + if _, _, err := format.Parse(bytes.NewReader(makeHeader(maxHeaderBytes))); err != nil { + t.Fatalf("maximum-size header was rejected: %v", err) + } + _, _, err := format.Parse(bytes.NewReader(makeHeader(maxHeaderBytes + 1))) + if err == nil || !strings.Contains(err.Error(), "header exceeds 2 MiB") { + t.Fatalf("unexpected oversized-header error: %v", err) + } + var parseError *format.ParseError + if !errors.As(err, &parseError) { + t.Errorf("oversized-header error is not a ParseError: %T", err) + } + if strings.Count(err.Error(), "parsing age header") != 1 { + t.Errorf("oversized-header error nests ParseError prefixes: %v", err) + } + + // An intro line that exceeds the limit before any newline. + _, _, err = format.Parse(bytes.NewReader(bytes.Repeat([]byte{'a'}, maxHeaderBytes+1))) + if err == nil || !strings.Contains(err.Error(), "header exceeds 2 MiB") { + t.Fatalf("unexpected oversized-intro error: %v", err) + } + if !errors.As(err, &parseError) { + t.Errorf("oversized-intro error is not a ParseError: %T", err) + } + if strings.Count(err.Error(), "parsing age header") != 1 { + t.Errorf("oversized-intro error nests ParseError prefixes: %v", err) + } + }) + + t.Run("recipient stanzas", func(t *testing.T) { + makeHeader := func(stanzas int) []byte { + return []byte(intro + strings.Repeat("-> test\n\n", stanzas) + footer) + } + + if _, _, err := format.Parse(bytes.NewReader(makeHeader(1024))); err != nil { + t.Fatalf("header with 1024 stanzas was rejected: %v", err) + } + _, _, err := format.Parse(bytes.NewReader(makeHeader(1025))) + if err == nil || !strings.Contains(err.Error(), "more than 1024 recipient stanzas") { + t.Fatalf("unexpected stanza-limit error: %v", err) + } + }) + + t.Run("recipient stanza arguments", func(t *testing.T) { + makeHeader := func(args int) []byte { + return []byte(intro + "-> test" + strings.Repeat(" a", args) + "\n\n" + footer) + } + + if _, _, err := format.Parse(bytes.NewReader(makeHeader(128))); err != nil { + t.Fatalf("stanza with 128 arguments was rejected: %v", err) + } + _, _, err := format.Parse(bytes.NewReader(makeHeader(129))) + if err == nil { + t.Fatal("stanza with 129 arguments was accepted") + } + }) + + t.Run("payload is not limited", func(t *testing.T) { + header := []byte(intro + "-> test\n\n" + footer) + want := bytes.Repeat([]byte("p"), maxHeaderBytes+1) + _, payload, err := format.Parse(io.MultiReader(bytes.NewReader(header), bytes.NewReader(want))) + if err != nil { + t.Fatal(err) + } + got, err := io.ReadAll(payload) + if err != nil { + t.Fatal(err) + } + if !bytes.Equal(got, want) { + t.Error("payload was truncated or modified") + } + }) + + t.Run("bufio reader is preserved", func(t *testing.T) { + header := []byte(intro + "-> test\n\n" + footer) + input := bufio.NewReader(bytes.NewReader(append(header, "payload"...))) + _, payload, err := format.Parse(input) + if err != nil { + t.Fatal(err) + } + if payload != input { + t.Error("Parse did not return the input bufio.Reader") + } + }) +} + func FuzzMalleability(f *testing.F) { tests, err := filepath.Glob("../../testdata/testkit/*") if err != nil {