implement host validation for login to check for URL schemes and paths. Halt the execution if any is found with a dedicated error message.

Signed-off-by: Yannick Alexander <yannick@alexanderdev.io>
pull/31986/head
Yannick Alexander 6 months ago
parent e2a2ed5009
commit c59f26a659
No known key found for this signature in database

@ -24,10 +24,10 @@ import (
"errors" "errors"
"fmt" "fmt"
"io" "io"
"log/slog"
"net/http" "net/http"
"net/url" "net/url"
"os" "os"
"regexp"
"sort" "sort"
"strings" "strings"
@ -227,15 +227,29 @@ type (
} }
) )
// warnIfHostHasPath checks if the host contains a repository path and logs a warning if it does. var hostRegex = regexp.MustCompile(`^(?P<scheme>[a-z]*:\/\/)?(?P<host>[a-zA-Z0-9-\.\:]+)(?P<path>\/.*)?$`)
// Returns true if the host contains a path component (i.e., contains a '/').
func warnIfHostHasPath(host string) bool { // validateHost checks that the host matches some required pre-checks e.g. does not contain a scheme or path.
if strings.Contains(host, "/") { // While ORAS will also validate some of these things, the current errors are a bit opaque.
registryHost := strings.Split(host, "/")[0] // By validating these things upfront, we can provide clearer error messages to users when they attempt to login with an invalid host string.
slog.Warn("registry login currently only supports registry hostname, not a repository path", "host", host, "suggested", registryHost) func validateHost(host string) error {
return true matches := hostRegex.FindStringSubmatch(host)
if len(matches) == 0 {
return fmt.Errorf("invalid host: %q", host)
}
scheme := matches[1]
path := matches[3]
if scheme != "" {
return fmt.Errorf("host should not contain a scheme (e.g. http://), found %q", scheme)
}
if path != "" {
return fmt.Errorf("host should not contain a path, found %q", path)
} }
return false
return nil
} }
// Login authenticates the client with a remote OCI registry using the provided host and options. // Login authenticates the client with a remote OCI registry using the provided host and options.
@ -244,7 +258,9 @@ func (c *Client) Login(host string, options ...LoginOption) error {
option(&loginOperation{host, c}) option(&loginOperation{host, c})
} }
warnIfHostHasPath(host) if err := validateHost(host); err != nil {
return err
}
reg, err := remote.NewRegistry(host) reg, err := remote.NewRegistry(host)
if err != nil { if err != nil {

@ -121,47 +121,61 @@ func TestLogin_ResetsForceAttemptOAuth2_OnFailure(t *testing.T) {
} }
} }
// TestWarnIfHostHasPath verifies that warnIfHostHasPath correctly detects path components. func TestValidateHost(t *testing.T) {
func TestWarnIfHostHasPath(t *testing.T) {
t.Parallel() t.Parallel()
tests := []struct { tests := []struct {
name string name string
host string host string
wantWarn bool wantErr bool
}{ }{
{ {
name: "domain only", name: "domain only",
host: "ghcr.io", host: "ghcr.io",
wantWarn: false, wantErr: false,
}, },
{ {
name: "domain with port", name: "domain with port",
host: "localhost:8000", host: "localhost:8000",
wantWarn: false, wantErr: false,
}, },
{ {
name: "domain with repository path", name: "domain with repository path",
host: "ghcr.io/terryhowe", host: "ghcr.io/terryhowe",
wantWarn: true, wantErr: true,
}, },
{ {
name: "domain with nested path", name: "domain with nested path",
host: "ghcr.io/terryhowe/myrepo", host: "ghcr.io/terryhowe/myrepo",
wantWarn: true, wantErr: true,
}, },
{ {
name: "localhost with port and path", name: "localhost with port and path",
host: "localhost:8000/myrepo", host: "localhost:8000/myrepo",
wantWarn: true, wantErr: true,
},
{
name: "domain with http protocol",
host: "http://ghcr.io",
wantErr: true,
},
{
name: "domain with https protocol",
host: "https://ghcr.io",
wantErr: true,
},
{
name: "domain with oci protocol",
host: "oci://ghcr.io",
wantErr: true,
}, },
} }
for _, tt := range tests { for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) { t.Run(tt.name, func(t *testing.T) {
got := warnIfHostHasPath(tt.host) err := validateHost(tt.host)
if got != tt.wantWarn { if (err != nil) != tt.wantErr {
t.Errorf("warnIfHostHasPath(%q) = %v, want %v", tt.host, got, tt.wantWarn) t.Errorf("validateHost(%q) error = %v, wantErr %v", tt.host, err, tt.wantErr)
} }
}) })
} }

Loading…
Cancel
Save