diff --git a/bundle/bundle.go b/bundle/bundle.go index 917f01847a..cb94714749 100644 --- a/bundle/bundle.go +++ b/bundle/bundle.go @@ -31,6 +31,8 @@ const ( const bundleLimitBytes = (1024 * 1024 * 1024) + 1 // limit bundle reads to 1GB to protect against gzip bombs +var manifestPath = []string{"system", "bundle", "manifest"} + // Bundle represents a loaded bundle. The bundle can contain data and policies. type Bundle struct { Manifest Manifest @@ -51,57 +53,34 @@ type ModuleFile struct { Parsed *ast.Module } -// Write serializes the Bundle and writes it to w. -func Write(w io.Writer, bundle Bundle) error { - gw := gzip.NewWriter(w) - tw := tar.NewWriter(gw) - - var buf bytes.Buffer - - if err := json.NewEncoder(&buf).Encode(bundle.Data); err != nil { - return err - } - - if err := writeFile(tw, "data.json", buf.Bytes()); err != nil { - return err - } - - for _, module := range bundle.Modules { - if err := writeFile(tw, module.Path, module.Raw); err != nil { - return err - } - } - - if err := writeManifest(tw, bundle); err != nil { - return err - } - - if err := tw.Close(); err != nil { - return err - } - - return gw.Close() +// Reader contains the reader to load the bundle from. +type Reader struct { + r io.Reader + includeManifestInData bool } -func writeManifest(tw *tar.Writer, bundle Bundle) error { +// NewReader returns a new Reader. +func NewReader(r io.Reader) *Reader { + nr := Reader{} + nr.r = r + return &nr +} - var buf bytes.Buffer - - if err := json.NewEncoder(&buf).Encode(bundle.Manifest); err != nil { - return err - } - - return writeFile(tw, manifestExt, buf.Bytes()) +// IncludeManifestInData sets whether the manifest metadata should be +// included in the bundle's data. +func (r *Reader) IncludeManifestInData(includeManifestInData bool) *Reader { + r.includeManifestInData = includeManifestInData + return r } // Read returns a new Bundle loaded from the reader. -func Read(r io.Reader) (Bundle, error) { +func (r *Reader) Read() (Bundle, error) { var bundle Bundle bundle.Data = map[string]interface{}{} - gr, err := gzip.NewReader(r) + gr, err := gzip.NewReader(r.r) if err != nil { return bundle, errors.Wrap(err, "bundle read failed") } @@ -161,7 +140,24 @@ func Read(r io.Reader) (Bundle, error) { } else if strings.HasSuffix(path, manifestExt) { if err := util.NewJSONDecoder(&buf).Decode(&bundle.Manifest); err != nil { - return bundle, errors.Wrapf(err, "bundle load failed on manifest") + return bundle, errors.Wrap(err, "bundle load failed on manifest decode") + } + + if r.includeManifestInData { + var metadata map[string]interface{} + b, err := json.Marshal(&bundle.Manifest) + if err != nil { + return bundle, errors.Wrap(err, "bundle load failed on manifest marshal") + } + + err = util.UnmarshalJSON(b, &metadata) + if err != nil { + return bundle, errors.Wrap(err, "bundle load failed on manifest unmarshal") + } + + if err := bundle.insert(manifestPath, metadata); err != nil { + return bundle, errors.Wrapf(err, "bundle load failed on %v", manifestPath) + } } } } @@ -169,6 +165,49 @@ func Read(r io.Reader) (Bundle, error) { return bundle, nil } +// Write serializes the Bundle and writes it to w. +func Write(w io.Writer, bundle Bundle) error { + gw := gzip.NewWriter(w) + tw := tar.NewWriter(gw) + + var buf bytes.Buffer + + if err := json.NewEncoder(&buf).Encode(bundle.Data); err != nil { + return err + } + + if err := writeFile(tw, "data.json", buf.Bytes()); err != nil { + return err + } + + for _, module := range bundle.Modules { + if err := writeFile(tw, module.Path, module.Raw); err != nil { + return err + } + } + + if err := writeManifest(tw, bundle); err != nil { + return err + } + + if err := tw.Close(); err != nil { + return err + } + + return gw.Close() +} + +func writeManifest(tw *tar.Writer, bundle Bundle) error { + + var buf bytes.Buffer + + if err := json.NewEncoder(&buf).Encode(bundle.Manifest); err != nil { + return err + } + + return writeFile(tw, manifestExt, buf.Bytes()) +} + // Equal returns true if this bundle's contents equal the other bundle's // contents. func (b Bundle) Equal(other Bundle) bool { diff --git a/bundle/bundle_test.go b/bundle/bundle_test.go index 0d7eeaccbc..06afbe2bd3 100644 --- a/bundle/bundle_test.go +++ b/bundle/bundle_test.go @@ -23,7 +23,7 @@ func TestRead(t *testing.T) { } buf := writeTarGz(files) - bundle, err := Read(buf) + bundle, err := NewReader(buf).Read() if err != nil { t.Fatal(err) } @@ -58,7 +58,7 @@ func TestReadWithManifest(t *testing.T) { {"/.manifest", `{"revision": "quickbrownfaux"}`}, } buf := writeTarGz(files) - bundle, err := Read(buf) + bundle, err := NewReader(buf).Read() if err != nil { t.Fatal(err) } @@ -67,9 +67,28 @@ func TestReadWithManifest(t *testing.T) { } } +func TestReadWithManifestInData(t *testing.T) { + files := [][2]string{ + {"/.manifest", `{"revision": "quickbrownfaux"}`}, + } + buf := writeTarGz(files) + bundle, err := NewReader(buf).IncludeManifestInData(true).Read() + if err != nil { + t.Fatal(err) + } + + system := bundle.Data["system"].(map[string]interface{}) + b := system["bundle"].(map[string]interface{}) + m := b["manifest"].(map[string]interface{}) + + if m["revision"] != "quickbrownfaux" { + t.Fatalf("Unexpected manifest.revision value: %v. Expected: %v", m["revision"], "quickbrownfaux") + } +} + func TestReadErrorBadGzip(t *testing.T) { buf := bytes.NewBufferString("bad gzip bytes") - _, err := Read(buf) + _, err := NewReader(buf).Read() if err == nil { t.Fatal("expected error") } @@ -80,7 +99,7 @@ func TestReadErrorBadTar(t *testing.T) { gw := gzip.NewWriter(&buf) gw.Write([]byte("bad tar bytes")) gw.Close() - _, err := Read(&buf) + _, err := NewReader(&buf).Read() if err == nil { t.Fatal("expected error") } @@ -100,7 +119,7 @@ func TestReadErrorBadContents(t *testing.T) { } for _, test := range tests { buf := writeTarGz(test.files) - _, err := Read(buf) + _, err := NewReader(buf).Read() if err == nil { t.Fatal("expected error") } @@ -136,7 +155,7 @@ func TestRoundtrip(t *testing.T) { t.Fatal("Unexpected error:", err) } - bundle2, err := Read(&buf) + bundle2, err := NewReader(&buf).Read() if err != nil { t.Fatal("Unexpected error:", err) } diff --git a/loader/loader.go b/loader/loader.go index 76c4707255..eb668db3be 100644 --- a/loader/loader.go +++ b/loader/loader.go @@ -303,7 +303,8 @@ func loadFileForAnyType(path string, bs []byte) (interface{}, error) { } func loadBundle(bs []byte) (bundle.Bundle, error) { - return bundle.Read(bytes.NewBuffer(bs)) + br := bundle.NewReader(bytes.NewBuffer(bs)).IncludeManifestInData(true) + return br.Read() } func loadRego(path string, bs []byte) (*RegoFile, error) { diff --git a/loader/loader_test.go b/loader/loader_test.go index 7c097f3cb3..3fe5e1590e 100644 --- a/loader/loader_test.go +++ b/loader/loader_test.go @@ -186,8 +186,11 @@ func TestLoadBundle(t *testing.T) { t.Fatal(err) } - if !reflect.DeepEqual(testBundle.Data, loaded.Documents) { - t.Fatalf("Expected %v but got: %v", testBundle.Data, loaded.Documents) + actualData := testBundle.Data + actualData["system"] = map[string]interface{}{"bundle": map[string]interface{}{"manifest": map[string]interface{}{"revision": ""}}} + + if !reflect.DeepEqual(actualData, loaded.Documents) { + t.Fatalf("Expected %v but got: %v", actualData, loaded.Documents) } if !bytes.Equal(testBundle.Modules[0].Raw, loaded.Modules["/x.rego"].Raw) { diff --git a/plugins/bundle/plugin.go b/plugins/bundle/plugin.go index d03be2cdb1..69b2e53c59 100644 --- a/plugins/bundle/plugin.go +++ b/plugins/bundle/plugin.go @@ -275,7 +275,7 @@ func (p *Plugin) download(ctx context.Context, resp *http.Response) (*bundle.Bun p.logDebug("Bundle download in progress.") - b, err := bundle.Read(resp.Body) + b, err := bundle.NewReader(resp.Body).Read() if err != nil { return nil, err }