From b030f04671d33ee4b2c0c4568f19c44273f16a6d Mon Sep 17 00:00:00 2001 From: Filippo Valsorda Date: Sat, 29 Aug 2026 18:15:31 +0200 Subject: [PATCH] armor: enforce writer lifecycle Reported by Joe Doyle of Trail of Bits. --- armor/armor.go | 24 +++++++++++++++++++----- armor/armor_test.go | 22 ++++++++++++++++++++++ 2 files changed, 41 insertions(+), 5 deletions(-) diff --git a/armor/armor.go b/armor/armor.go index c6b5fb0..cc4dd7a 100644 --- a/armor/armor.go +++ b/armor/armor.go @@ -31,13 +31,24 @@ type armoredWriter struct { dst io.Writer } -func (a *armoredWriter) Write(p []byte) (int, error) { - if !a.started { - if _, err := io.WriteString(a.dst, Header+"\n"); err != nil { - return 0, err - } +func (a *armoredWriter) writeHeader() error { + if a.started { + return nil + } + if _, err := io.WriteString(a.dst, Header+"\n"); err != nil { + return err } a.started = true + return nil +} + +func (a *armoredWriter) Write(p []byte) (int, error) { + if a.closed { + return 0, errors.New("ArmoredWriter already closed") + } + if err := a.writeHeader(); err != nil { + return 0, err + } return a.encoder.Write(p) } @@ -46,6 +57,9 @@ func (a *armoredWriter) Close() error { return errors.New("ArmoredWriter already closed") } a.closed = true + if err := a.writeHeader(); err != nil { + return err + } if err := a.encoder.Close(); err != nil { return err } diff --git a/armor/armor_test.go b/armor/armor_test.go index 6500d39..ca25e2b 100644 --- a/armor/armor_test.go +++ b/armor/armor_test.go @@ -94,6 +94,28 @@ func TestArmor(t *testing.T) { t.Run("FullLine", func(t *testing.T) { testArmor(t, 10*format.BytesPerLine) }) } +func TestWriterLifecycle(t *testing.T) { + t.Run("CloseWithoutWrite", func(t *testing.T) { + buf := &bytes.Buffer{} + w := armor.NewWriter(buf) + if err := w.Close(); err != nil { + t.Fatal(err) + } + if want := armor.Header + "\n" + armor.Footer + "\n"; buf.String() != want { + t.Errorf("output = %q, want %q", buf.String(), want) + } + }) + t.Run("WriteAfterClose", func(t *testing.T) { + w := armor.NewWriter(io.Discard) + if err := w.Close(); err != nil { + t.Fatal(err) + } + if n, err := w.Write([]byte("x")); n != 0 || err == nil { + t.Errorf("Write after Close = %d, %v", n, err) + } + }) +} + func testArmor(t *testing.T, size int) { buf := &bytes.Buffer{} w := armor.NewWriter(buf)