pass url.URL rather than string

return a *url.URL value from expandLink rather than a string, and both
accept and return a *url.URL value from resolveLink rather than a
string. These are both unexported funcs, so this has no changes to the
current behavior. It does prevent a few unnecessary conversions back and
forth between url.URL and string values, and will make it simpler to
retain request query strings.

Updates #77

Signed-off-by: Will Norris <will@tailscale.com>
This commit is contained in:
Will Norris
2023-05-16 10:17:35 -07:00
committed by Will Norris
parent f00de63b45
commit ae5d8c9b6c
2 changed files with 52 additions and 28 deletions
+25 -20
View File
@@ -124,11 +124,15 @@ func Run() error {
// if link specified on command line, resolve and exit
if flag.NArg() > 0 {
destination, err := resolveLink(flag.Arg(0))
u, err := url.Parse(flag.Arg(0))
if err != nil {
log.Fatal(err)
}
fmt.Println(destination)
dst, err := resolveLink(u)
if err != nil {
log.Fatal(err)
}
fmt.Println(dst.String())
os.Exit(0)
}
@@ -403,7 +407,7 @@ func serveGo(w http.ResponseWriter, r *http.Request) {
http.Error(w, err.Error(), http.StatusInternalServerError)
return
}
http.Redirect(w, r, target, http.StatusFound)
http.Redirect(w, r, target.String(), http.StatusFound)
}
// acceptHTML returns whether the request can accept a text/html response.
@@ -494,7 +498,7 @@ var expandFuncMap = texttemplate.FuncMap{
//
// If long does not include templates, the default behavior is to append
// env.Path to long.
func expandLink(long string, env expandEnv) (string, error) {
func expandLink(long string, env expandEnv) (*url.URL, error) {
if !strings.Contains(long, "{{") {
// default behavior is to append remaining path to long URL
if strings.HasSuffix(long, "/") {
@@ -505,19 +509,19 @@ func expandLink(long string, env expandEnv) (string, error) {
}
tmpl, err := texttemplate.New("").Funcs(expandFuncMap).Parse(long)
if err != nil {
return "", err
return nil, err
}
buf := new(bytes.Buffer)
if err := tmpl.Execute(buf, env); err != nil {
return "", err
return nil, err
}
long = buf.String()
_, err = url.Parse(long)
u, err := url.Parse(buf.String())
if err != nil {
return "", err
return nil, err
}
return long, nil
return u, nil
}
func devMode() bool { return *dev != "" }
@@ -740,22 +744,23 @@ func restoreLastSnapshot() error {
return bs.Err()
}
func resolveLink(link string) (string, error) {
// if link specified as "go/name", trim "go" prefix.
// Remainder will parse as URL with no scheme or host
link = strings.TrimPrefix(link, *hostname)
u, err := url.Parse(link)
if err != nil {
return "", err
func resolveLink(link *url.URL) (*url.URL, error) {
path := link.Path
// if link was specified as "go/name", it will parse with no scheme or host.
// Trim "go" prefix from beginning of path.
if link.Host == "" {
path = strings.TrimPrefix(path, *hostname)
}
short, remainder, _ := strings.Cut(strings.TrimPrefix(u.RequestURI(), "/"), "/")
short, remainder, _ := strings.Cut(strings.TrimPrefix(path, "/"), "/")
l, err := db.Load(short)
if err != nil {
return "", err
return nil, err
}
dst, err := expandLink(l.Long, expandEnv{Now: time.Now().UTC(), Path: remainder})
if err == nil {
if u, uErr := url.Parse(dst); uErr == nil && (u.Hostname() == "" || u.Hostname() == *hostname) {
if dst.Host == "" || dst.Host == *hostname {
dst, err = resolveLink(dst)
}
}
+27 -8
View File
@@ -13,6 +13,7 @@ import (
"time"
"golang.org/x/net/xsrftoken"
"tailscale.com/util/must"
)
func init() {
@@ -344,14 +345,14 @@ func TestExpandLink(t *testing.T) {
{
name: "template-with-pathescape-func",
long: "http://host.com/{{PathEscape .Path}}",
remainder: "a/b",
want: "http://host.com/a%2Fb",
remainder: "a/b+c",
want: "http://host.com/a%2Fb+c",
},
{
name: "template-with-queryescape-func",
long: "http://host.com/{{QueryEscape .Path}}",
remainder: "a+b",
want: "http://host.com/a%2Bb",
remainder: "a/b+c",
want: "http://host.com/a%2Fb%2Bc",
},
{
name: "template-with-trimsuffix-func",
@@ -359,13 +360,30 @@ func TestExpandLink(t *testing.T) {
remainder: "a/",
want: "http://host.com/a",
},
{
name: "relative-link",
long: `rel`,
remainder: "a",
want: "rel/a",
},
{
name: "relative-link-with-slash",
long: `/rel`,
remainder: "a",
want: "/rel/a",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
got, err := expandLink(tt.long, expandEnv{Now: tt.now, Path: tt.remainder, user: tt.user})
link, err := expandLink(tt.long, expandEnv{Now: tt.now, Path: tt.remainder, user: tt.user})
if (err != nil) != tt.wantErr {
t.Fatalf("expandLink(%q) returned error %v; want %v", tt.long, err, tt.wantErr)
}
var got string
if link != nil {
got = link.String()
}
if got != tt.want {
t.Errorf("expandLink(%q) = %q; want %q", tt.long, got, tt.want)
}
@@ -431,12 +449,13 @@ func TestResolveLink(t *testing.T) {
for _, tt := range tests {
name := "golink " + tt.link
t.Run(name, func(t *testing.T) {
got, err := resolveLink(tt.link)
u := must.Get(url.Parse(tt.link))
got, err := resolveLink(u)
if err != nil {
t.Error(err)
}
if got != tt.want {
t.Errorf("ResolveLink(%q) = %q; want %q", tt.link, got, tt.want)
if got.String() != tt.want {
t.Errorf("ResolveLink(%q) = %q; want %q", tt.link, got.String(), tt.want)
}
})
}