@@ -10,14 +10,13 @@ import (
1010 "fmt"
1111 "net"
1212 "net/url"
13- "slices"
1413 "strings"
1514
1615 "golang.org/x/net/idna"
1716)
1817
18+ // Errors
1919var (
20- ErrInvalidURL = fmt .Errorf ("invalid url" )
2120 ErrSchemeMismatch = fmt .Errorf ("scheme mismatch" )
2221 ErrPortMismatch = fmt .Errorf ("port mismatch" )
2322 ErrHostnameMismatch = fmt .Errorf ("hostname mismatch" )
@@ -29,8 +28,7 @@ type DomainValidatorOptions struct {
2928 WithScheme bool
3029 // Ensure domains have the same port.
3130 WithPort bool
32- // Specify a list of allowed schemes IF WithScheme is set to true.
33- // Leave empty to allow any scheme.
31+ // Specify a list of allowed schemes if WithScheme is set to true.
3432 AllowedSchemes []string
3533}
3634
@@ -48,53 +46,74 @@ func NewDomainValidator(opts DomainValidatorOptions) *DomainValidator {
4846 }
4947}
5048
51- func (v * DomainValidator ) getURL (i string ) (* url.URL , error ) {
52- u , err := url .Parse (i )
53-
54- if ! v .opts .WithScheme && (err != nil || u .Host == "" ) {
55- u , err = url .Parse ("tinyauth://" + i )
49+ func (v * DomainValidator ) checkScheme (rawURL string ) error {
50+ if ! v .opts .WithScheme {
51+ return nil
5652 }
5753
58- if err != nil {
59- return nil , fmt .Errorf ("failed to parse input url: %w" , err )
54+ if len ( v . opts . AllowedSchemes ) == 0 {
55+ return fmt .Errorf ("allowed schemes must be specified" )
6056 }
6157
62- if u .Host == "" {
63- return nil , ErrInvalidURL
58+ for _ , scheme := range v .opts .AllowedSchemes {
59+ if strings .HasPrefix (strings .ToLower (rawURL ), strings .ToLower (scheme )+ "://" ) {
60+ return nil
61+ }
6462 }
6563
66- if v .opts .WithPort && u .Port () == "" && (u .Scheme != "http" && u .Scheme != "https" ) {
67- return nil , fmt .Errorf ("port validation is enabled but port is missing in input url and schemes are not enabled" )
64+ return fmt .Errorf ("invalid scheme" )
65+
66+ }
67+
68+ func (v * DomainValidator ) getURL (i string ) (* url.URL , error ) {
69+ if i == "" {
70+ return nil , fmt .Errorf ("url cannot be empty" )
6871 }
6972
7073 if v .opts .WithScheme {
71- // Empty scheme means that we parsed the url with the tinyauth:// placeholder
72- if u .Scheme == "tinyauth" {
73- return nil , fmt .Errorf ("input url is missing scheme" )
74+ err := v .checkScheme (i )
75+
76+ if err != nil {
77+ return nil , fmt .Errorf ("invalid scheme: %w" , err )
7478 }
75- if len (v .opts .AllowedSchemes ) > 0 && ! slices .Contains (v .opts .AllowedSchemes , u .Scheme ) {
76- return nil , fmt .Errorf ("scheme %s not allowed" , u .Scheme )
79+
80+ u , err := url .Parse (i )
81+
82+ if err != nil {
83+ return nil , fmt .Errorf ("failed to parse input url: %w" , err )
84+ }
85+
86+ if u .Host == "" || u .Scheme == "" {
87+ return nil , fmt .Errorf ("missing host or scheme in url: %s" , i )
7788 }
89+
90+ return u , nil
7891 }
7992
80- return u , nil
81- }
93+ rawURL := i
8294
83- func (v * DomainValidator ) getEffectivePort (u * url.URL ) (string , bool ) {
84- if u .Port () != "" {
85- return u .Port (), true
95+ if ! strings .Contains (i , "://" ) {
96+ // From godoc: [scheme:][//[userinfo@]host][/]path[?query][#fragment]
97+ // So, we can omit the colon and tell the Go URL lib that we want
98+ // to parse the URL without the scheme. If we don't do this,
99+ // the URL lib will parse our entire domain as the path.
100+ rawURL = "//" + i
86101 }
87- switch u .Scheme {
88- case "http" :
89- return "80" , true
90- case "https" :
91- return "443" , true
92- default :
93- return "" , false
102+
103+ u , err := url .Parse (rawURL )
104+
105+ if err != nil {
106+ return nil , fmt .Errorf ("failed to parse host: %w" , err )
94107 }
108+
109+ if u .Host == "" {
110+ return nil , fmt .Errorf ("missing host in url: %s" , i )
111+ }
112+
113+ return u , nil
95114}
96115
97- func (v * DomainValidator ) formatHostname (hostname string ) (string , error ) {
116+ func (v * DomainValidator ) getHostname (hostname string ) (string , error ) {
98117 hostname = strings .ToLower (hostname )
99118 hostname = strings .TrimSuffix (hostname , "." )
100119 if net .ParseIP (hostname ) != nil {
@@ -133,26 +152,18 @@ func (v *DomainValidator) Validate(expected, actual string) error {
133152 }
134153
135154 if v .opts .WithPort {
136- eup , ok := v .getEffectivePort (eu )
137- if ! ok {
138- return fmt .Errorf ("failed to get effective port for url: %s" , eu .String ())
139- }
140- aup , ok := v .getEffectivePort (au )
141- if ! ok {
142- return fmt .Errorf ("failed to get effective port for url: %s" , au .String ())
143- }
144- if eup != aup {
155+ if eu .Port () != au .Port () {
145156 return ErrPortMismatch
146157 }
147158 }
148159
149- euf , err := v .formatHostname (eu .Hostname ())
160+ euf , err := v .getHostname (eu .Hostname ())
150161
151162 if err != nil {
152163 return err
153164 }
154165
155- auf , err := v .formatHostname (au .Hostname ())
166+ auf , err := v .getHostname (au .Hostname ())
156167
157168 if err != nil {
158169 return err
@@ -165,7 +176,7 @@ func (v *DomainValidator) Validate(expected, actual string) error {
165176 return nil
166177}
167178
168- // SafeHostname uses the internal validation for domains that Validator uses
179+ // SafeHostname uses the internal validation for domains that the validator uses
169180// to parse a hostname. It ensures the input URL is a valid URL, that a host
170181// is present and that the hostname is lowercased and without a trailing dot.
171182func (v * DomainValidator ) SafeHostname (input string ) (string , error ) {
@@ -175,5 +186,5 @@ func (v *DomainValidator) SafeHostname(input string) (string, error) {
175186 return "" , err
176187 }
177188
178- return v .formatHostname (u .Hostname ())
189+ return v .getHostname (u .Hostname ())
179190}
0 commit comments