diff --git a/ociref/reference.go b/ociref/reference.go index d3a68ee..fea511a 100644 --- a/ociref/reference.go +++ b/ociref/reference.go @@ -142,7 +142,7 @@ func IsValidHost(s string) bool { // IsValidRepository reports whether s is a valid repository part // of a reference string. func IsValidRepository(s string) bool { - return repoPat().MatchString(s) + return len(s) <= 255 && repoPat().MatchString(s) } // IsValidTag reports whether s is a valid reference tag. @@ -232,6 +232,9 @@ func parse(refStr string) (Reference, error) { } func checkTag(s string) error { + if len(s) == 0 { + return fmt.Errorf("tag is empty") + } if len(s) > 128 { return fmt.Errorf("tag too long") } diff --git a/ociref/reference_test.go b/ociref/reference_test.go index f8287b5..edf7f6b 100644 --- a/ociref/reference_test.go +++ b/ociref/reference_test.go @@ -557,6 +557,9 @@ var isValidTagTests = []struct { tag string want bool }{{ + tag: "", + want: false, +}, { tag: "hello", want: true, }, { diff --git a/ociserver/blobuploads.go b/ociserver/blobuploads.go index f5d74aa..30eeb6c 100644 --- a/ociserver/blobuploads.go +++ b/ociserver/blobuploads.go @@ -74,17 +74,30 @@ func (s *Server) blobUploadPost() http.HandlerFunc { mount := r.URL.Query().Get("mount") from := r.URL.Query().Get("from") - if mount != "" && from != "" { - if !ociref.IsValidRepository(from) { - returnError(w, ErrBlobUploadInvalid("invalid from parameter")) + var mountDigest, dgst ocidigest.Digest + var err error + if mount != "" { + mountDigest, err = ocidigest.Parse(mount) + if err != nil { + returnError(w, ErrBlobUploadInvalid("invalid mount digest")) return } - dgst, err := ocidigest.Parse(mount) + } + if dgstString != "" { + dgst, err = ocidigest.Parse(dgstString) if err != nil { returnError(w, ErrBlobUploadInvalid("invalid digest")) return } - blob, err := s.db.MountBlob(r.Context(), from, name, dgst) + } + + if mountDigest != "" && from != "" { + if !ociref.IsValidRepository(from) { + returnError(w, ErrBlobUploadInvalid("invalid from parameter")) + return + } + + blob, err := s.db.MountBlob(r.Context(), from, name, mountDigest) if err != nil { goto FALLBACK } @@ -94,12 +107,7 @@ func (s *Server) blobUploadPost() http.HandlerFunc { w.Header().Set("Docker-Content-Digest", blob.Digest.String()) w.WriteHeader(http.StatusCreated) return - } else if dgstString != "" { - dgst, err := ocidigest.Parse(dgstString) - if err != nil { - returnError(w, ErrBlobUploadInvalid("invalid digest")) - return - } + } else if dgst != "" { contentLength := r.Header.Get("Content-Length") if contentLength == "" { contentLength = "0" diff --git a/ociserver/blobuploads_test.go b/ociserver/blobuploads_test.go index b7b4100..01a8b55 100644 --- a/ociserver/blobuploads_test.go +++ b/ociserver/blobuploads_test.go @@ -12,8 +12,40 @@ import ( "github.com/docker/oci" "github.com/docker/oci/ocidigest" + "github.com/stretchr/testify/require" ) +func TestBlobUploadPostValidatesMountParameters(t *testing.T) { + t.Parallel() + + dgst := ocidigest.FromBytes([]byte("blob")) + tests := []struct { + name string + query string + }{ + {name: "from is not a local repository", query: "mount=" + dgst.String() + "&from=UPPERCASE"}, + {name: "from contains a tag", query: "mount=" + dgst.String() + "&from=repo%3Alatest"}, + {name: "from is too long", query: "mount=" + dgst.String() + "&from=" + strings.Repeat("a", 256)}, + {name: "mount is not a digest", query: "mount=another%2Frepository&from=repo"}, + {name: "mount without from is still validated", query: "mount=not-a-digest"}, + {name: "digest with mount is still validated", query: "digest=not-a-digest&mount=" + dgst.String() + "&from=repo"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + s := &Server{db: (*oci.Funcs)(nil)} + req := httptest.NewRequest(http.MethodPost, "/v2/repo/blobs/uploads/?"+tt.query, nil) + rec := httptest.NewRecorder() + + serveTestRoute(t, `/v2/*name/blobs/uploads/`, s.blobUploadPost(), rec, req) + + require.Equal(t, http.StatusBadRequest, rec.Code) + require.Contains(t, rec.Body.String(), `"code":"BLOB_UPLOAD_INVALID"`) + }) + } +} + func TestParseRange(t *testing.T) { t.Parallel() diff --git a/ociserver/input_validation_test.go b/ociserver/input_validation_test.go new file mode 100644 index 0000000..f03e44e --- /dev/null +++ b/ociserver/input_validation_test.go @@ -0,0 +1,60 @@ +package ociserver + +import ( + "context" + "iter" + "net/http" + "net/http/httptest" + "net/url" + "strings" + "testing" + + "github.com/docker/oci" + "github.com/docker/oci/ocidigest" + "github.com/stretchr/testify/require" +) + +func TestTagsGetValidatesLast(t *testing.T) { + t.Parallel() + + s := &Server{db: &oci.Funcs{ + Tags_: func(context.Context, string, *oci.TagsParameters) iter.Seq2[string, error] { + t.Fatal("invalid last parameter reached storage") + return nil + }, + }} + req := httptest.NewRequest(http.MethodGet, "/v2/repo/tags/list?last=-bad", nil) + rec := httptest.NewRecorder() + + serveTestRoute(t, `/v2/*name/tags/list`, s.tagsGet(), rec, req) + + require.Equal(t, http.StatusBadRequest, rec.Code) + require.Contains(t, rec.Body.String(), `"code":"BAD_REQUEST"`) +} + +func TestReferrersGetValidatesArtifactType(t *testing.T) { + t.Parallel() + + dgst := ocidigest.FromBytes([]byte("manifest")) + s := &Server{db: &oci.Funcs{ + Referrers_: func(context.Context, string, oci.Digest, *oci.ReferrersParameters) iter.Seq2[oci.Descriptor, error] { + t.Fatal("invalid artifactType reached storage") + return nil + }, + }} + tests := []string{ + "not a media type", + "*/*", + "application/example; charset=utf-8", + strings.Repeat("a", oci.MaxArtifactTypeLen+1) + "/x", + } + for _, artifactType := range tests { + req := httptest.NewRequest(http.MethodGet, "/v2/repo/referrers/"+dgst.String()+"?artifactType="+url.QueryEscape(artifactType), nil) + rec := httptest.NewRecorder() + + serveTestRoute(t, `/v2/*name/referrers/:digest`, s.referrersGet(), rec, req) + + require.Equal(t, http.StatusBadRequest, rec.Code) + require.Contains(t, rec.Body.String(), `"code":"BAD_REQUEST"`) + } +} diff --git a/ociserver/manifests.go b/ociserver/manifests.go index 055d45a..17e1b4f 100644 --- a/ociserver/manifests.go +++ b/ociserver/manifests.go @@ -33,6 +33,10 @@ func (s *Server) manifestHeadGet() http.HandlerFunc { } desc, err = s.db.ResolveManifest(r.Context(), name, dgst) } else { + if !ociref.IsValidTag(reference) { + returnError(w, ErrManifestInvalid("invalid tag name")) + return + } desc, err = s.db.ResolveTag(r.Context(), name, reference) } if err != nil { @@ -135,6 +139,12 @@ func (s *Server) manifestPut() http.HandlerFunc { name := mux.URLParam(r, "name") reference := mux.URLParam(r, "reference") tags := r.URL.Query()["tag"] + for _, tag := range tags { + if !ociref.IsValidTag(tag) { + returnError(w, ErrManifestInvalid("invalid tag name")) + return + } + } defer func() { err := r.Body.Close() @@ -182,7 +192,12 @@ func (s *Server) manifestPut() http.HandlerFunc { } contentType := r.Header.Get("Content-Type") if contentType != "" { - contentType, _, _ = strings.Cut(contentType, ";") // strip any parameters + var err error + contentType, _, err = mime.ParseMediaType(contentType) + if err != nil { + returnError(w, ErrManifestInvalid("invalid Content-Type")) + return + } } if mani.MediaType != "" && contentType != "" && mani.MediaType != contentType { returnError(w, ErrManifestInvalid("mediaType does not match Content-Type")) @@ -266,6 +281,10 @@ func (s *Server) manifestDelete() http.HandlerFunc { } err = s.db.DeleteManifest(r.Context(), name, dgst) } else { + if !ociref.IsValidTag(reference) { + returnError(w, ErrManifestInvalid("invalid tag name")) + return + } err = s.db.DeleteTag(r.Context(), name, reference) } if err != nil { diff --git a/ociserver/manifests_test.go b/ociserver/manifests_test.go index 4750240..613df0a 100644 --- a/ociserver/manifests_test.go +++ b/ociserver/manifests_test.go @@ -1,6 +1,14 @@ package ociserver -import "testing" +import ( + "bytes" + "net/http" + "net/http/httptest" + "testing" + + "github.com/docker/oci" + "github.com/stretchr/testify/require" +) func TestAcceptsMediaType(t *testing.T) { t.Parallel() @@ -76,3 +84,53 @@ func TestAcceptsMediaType(t *testing.T) { }) } } + +func TestManifestHandlersValidateTags(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + method string + target string + handler func(*Server) http.HandlerFunc + body []byte + }{ + {name: "get reference", method: http.MethodGet, target: "/v2/repo/manifests/-bad", handler: (*Server).manifestHeadGet}, + {name: "delete reference", method: http.MethodDelete, target: "/v2/repo/manifests/-bad", handler: (*Server).manifestDelete}, + { + name: "put tag query", + method: http.MethodPut, + target: "/v2/repo/manifests/latest?tag=-bad", + handler: (*Server).manifestPut, + body: []byte(`{}`), + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + t.Parallel() + s := &Server{db: (*oci.Funcs)(nil)} + req := httptest.NewRequest(tt.method, tt.target, bytes.NewReader(tt.body)) + rec := httptest.NewRecorder() + + serveTestRoute(t, `/v2/*name/manifests/:reference`, tt.handler(s), rec, req) + + require.Equal(t, http.StatusBadRequest, rec.Code) + require.Contains(t, rec.Body.String(), `"code":"MANIFEST_INVALID"`) + }) + } +} + +func TestManifestPutValidatesContentType(t *testing.T) { + t.Parallel() + + s := &Server{db: (*oci.Funcs)(nil)} + req := httptest.NewRequest(http.MethodPut, "/v2/repo/manifests/latest", bytes.NewReader([]byte(`{}`))) + req.Header.Set("Content-Type", `application/vnd.oci.image.manifest.v1+json; broken`) + rec := httptest.NewRecorder() + + serveTestRoute(t, `/v2/*name/manifests/:reference`, s.manifestPut(), rec, req) + + require.Equal(t, http.StatusBadRequest, rec.Code) + require.Contains(t, rec.Body.String(), `"code":"MANIFEST_INVALID"`) +} diff --git a/ociserver/referrers.go b/ociserver/referrers.go index 6ef5560..108eb28 100644 --- a/ociserver/referrers.go +++ b/ociserver/referrers.go @@ -3,7 +3,9 @@ package ociserver import ( "encoding/json" "errors" + "mime" "net/http" + "strings" "github.com/docker/oci" "github.com/docker/oci/ocidigest" @@ -23,6 +25,10 @@ func (s *Server) referrersGet() http.HandlerFunc { name := mux.URLParam(r, "name") dgstString := mux.URLParam(r, "digest") artifactType := r.URL.Query().Get("artifactType") + if artifactType != "" && !isValidArtifactType(artifactType) { + returnError(w, ErrBadRequest("invalid artifactType")) + return + } dgst, err := ocidigest.Parse(dgstString) if err != nil { @@ -71,3 +77,15 @@ func (s *Server) referrersGet() http.HandlerFunc { } } } + +func isValidArtifactType(artifactType string) bool { + if len(artifactType) > oci.MaxArtifactTypeLen { + return false + } + mediaType, params, err := mime.ParseMediaType(artifactType) + if err != nil || len(params) != 0 { + return false + } + typeName, subtype, ok := strings.Cut(mediaType, "/") + return ok && typeName != "*" && subtype != "*" +} diff --git a/ociserver/server_test.go b/ociserver/server_test.go index e7fd409..941e15e 100644 --- a/ociserver/server_test.go +++ b/ociserver/server_test.go @@ -146,3 +146,15 @@ func TestServerInvalidRepositoryNameReturnsOCIError(t *testing.T) { }] }`, rec.Body.String()) } + +func TestServerRejectsOverlongRepositoryName(t *testing.T) { + srv, err := New((*oci.Funcs)(nil), nil) + require.NoError(t, err) + + name := strings.Repeat("a", 256) + rec := httptest.NewRecorder() + srv.ServeHTTP(rec, httptest.NewRequest(http.MethodGet, "/v2/"+name+"/tags/list", nil)) + + require.Equal(t, http.StatusBadRequest, rec.Code) + require.Contains(t, rec.Body.String(), `"code":"NAME_INVALID"`) +} diff --git a/ociserver/tags.go b/ociserver/tags.go index efbbdc9..92bb53f 100644 --- a/ociserver/tags.go +++ b/ociserver/tags.go @@ -9,6 +9,7 @@ import ( "strconv" "github.com/docker/oci" + "github.com/docker/oci/ociref" "github.com/docker/oci/ociserver/mux" ) @@ -16,6 +17,10 @@ func (s *Server) tagsGet() http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { name := mux.URLParam(r, "name") last := r.URL.Query().Get("last") + if last != "" && !ociref.IsValidTag(last) { + returnError(w, ErrBadRequest("invalid last")) + return + } limit := 0 if n := r.URL.Query().Get("n"); n != "" { i, err := strconv.Atoi(n)