diff --git a/db.go b/db.go new file mode 100644 index 0000000..3fac84c --- /dev/null +++ b/db.go @@ -0,0 +1,105 @@ +package main + +import ( + "encoding/json" + "fmt" + "net/url" + "os" + "path/filepath" + "strings" + "time" +) + +// Link is the structure stored for each go short link. +type Link struct { + Short string // the "foo" part of http://go/foo + Long string // the target URL or text/template pattern to run + Created time.Time + LastEdit time.Time // when the link was last edited + Owner string // user@domain +} + +// DB provides storage for Links. +type DB interface { + // List the short name of all stored Links. + List() ([]string, error) + + // Load a Link by its short name. It returns fs.ErrNotExist if the link does not exist. + Load(short string) (*Link, error) + + // Save a Link. + Save(link *Link) error +} + +// FileDB stores Links in JSON files on disk. +type FileDB struct { + // dir is the directory to store one JSON file per link. + dir string +} + +// NewFileDB returns a new FileDB which will store links in individual JSON +// files in the specified directory. If mkdir is true, the directory will be +// created if it does not exist. +func NewFileDB(dir string, mkdir bool) (*FileDB, error) { + if mkdir { + if err := os.MkdirAll(dir, 0755); err != nil { + return nil, err + } + } + if fi, err := os.Stat(dir); err != nil { + return nil, err + } else if !fi.IsDir() { + return nil, fmt.Errorf("%q is not a directory", dir) + } + return &FileDB{dir: dir}, nil +} + +// linkPath returns the path to the file on disk for the specified link. Short +// name is normalized to be case insensitive, remove dashes, and escape some +// characters. +// +// TODO(willnorris): some of this normalization is not unique to FileDB and +// should be moved elsewhere +func (f *FileDB) linkPath(short string) string { + name := url.PathEscape(strings.ToLower(short)) + name = strings.ReplaceAll(name, "-", "") + name = strings.ReplaceAll(name, ".", "%2e") + return filepath.Join(f.dir, name) +} + +func (f *FileDB) List() ([]string, error) { + d, err := os.Open(f.dir) + if err != nil { + return nil, err + } + defer d.Close() + + names, err := d.Readdirnames(0) + if err != nil { + return nil, err + } + return names, nil +} + +func (f *FileDB) Load(short string) (*Link, error) { + data, err := os.ReadFile(f.linkPath(short)) + if err != nil { + return nil, err + } + link := new(Link) + if err := json.Unmarshal(data, link); err != nil { + return nil, err + } + return link, nil +} + +func (f *FileDB) Save(link *Link) error { + j, err := json.MarshalIndent(link, "", " ") + if err != nil { + return err + } + if err := os.WriteFile(f.linkPath(link.Short), j, 0600); err != nil { + return err + } + return nil +} diff --git a/golink.go b/golink.go index 40ce0ee..8e533b9 100644 --- a/golink.go +++ b/golink.go @@ -1,4 +1,4 @@ -// The golink server runs http://go/, a company shortlink service. +// The golink server runs http://go/, a private shortlink service for tailnets. package main import ( @@ -7,15 +7,15 @@ import ( "embed" _ "embed" "encoding/json" + "errors" "flag" "fmt" "html" + "io/fs" "io/ioutil" "log" "net/http" "net/url" - "os" - "path/filepath" "regexp" "sort" "strings" @@ -45,16 +45,10 @@ var lastSnapshot []byte //go:embed *.html static var embeddedFS embed.FS -var localClient *tailscale.LocalClient +// db stores short links. +var db DB -// DiskLink is the JSON structure stored on disk in a file for each go short link. -type DiskLink struct { - Short string // the "foo" part of http://go/foo - Long string // the target URL - Created time.Time - LastEdit time.Time // when the link was created - Owner string // foo@tailscale.com -} +var localClient *tailscale.LocalClient func main() { flag.Parse() @@ -72,16 +66,10 @@ func main() { } } - if *doMkdir { - if err := os.MkdirAll(*linkDir, 0755); err != nil { - log.Fatal(err) - } - } - - if fi, err := os.Stat(*linkDir); err != nil { - log.Fatal(err) - } else if !fi.IsDir() { - log.Fatalf("--linkdir %q is not a directory", *linkDir) + var err error + db, err = NewFileDB(*linkDir, *doMkdir) + if err != nil { + log.Fatalf("NewFileDB(%q): %v", *linkDir, err) } restoreLastSnapshot() @@ -97,7 +85,7 @@ func main() { srv := &tsnet.Server{ Hostname: "go", - Logf: func(format string, args ...interface{}) {}, + Logf: func(format string, args ...any) {}, } if *verbose { srv.Logf = log.Printf @@ -137,25 +125,6 @@ func init() { homeCreate = template.Must(template.ParseFS(embeddedFS, "home.html")) } -func linkPath(short string) string { - name := url.PathEscape(strings.ToLower(short)) - name = strings.ReplaceAll(name, "-", "") - name = strings.ReplaceAll(name, ".", "%2e") - return filepath.Join(*linkDir, name) -} - -func loadLink(short string) (*DiskLink, error) { - data, err := os.ReadFile(linkPath(short)) - if err != nil { - return nil, err - } - dl := new(DiskLink) - if err := json.Unmarshal(data, dl); err != nil { - return nil, err - } - return dl, nil -} - func serveHome(w http.ResponseWriter, short string) { var clicks []visitData @@ -195,10 +164,10 @@ func serveGo(w http.ResponseWriter, r *http.Request) { return } - short, remainder, _ := strings.Cut(strings.ToLower(strings.TrimPrefix(r.RequestURI, "/")), "/") + short, remainder, _ := strings.Cut(strings.TrimPrefix(r.RequestURI, "/"), "/") - dl, err := loadLink(short) - if os.IsNotExist(err) { + link, err := db.Load(short) + if errors.Is(err, fs.ErrNotExist) { serveHome(w, short) return } @@ -212,20 +181,18 @@ func serveGo(w http.ResponseWriter, r *http.Request) { if stats.clicks == nil { stats.clicks = make(map[string]int) } - stats.clicks[dl.Short]++ + stats.clicks[link.Short]++ stats.mu.Unlock() - target, err := expandLink(dl.Long, expandEnv{Now: time.Now().UTC(), Path: remainder}) + target, err := expandLink(link.Long, expandEnv{Now: time.Now().UTC(), Path: remainder}) if err != nil { - log.Printf("expanding %q: %v", dl.Long, err) + log.Printf("expanding %q: %v", link.Long, err) http.Error(w, err.Error(), http.StatusInternalServerError) return } http.Redirect(w, r, target, http.StatusFound) } -var reVarExpand = regexp.MustCompile(`\$\{\w+\}`) - type expandEnv struct { Now time.Time @@ -270,8 +237,26 @@ func expandLink(long string, env expandEnv) (string, error) { func devMode() bool { return *dev != "" } +func currentUser(r *http.Request) (string, error) { + login := "" + if devMode() { + login = "foo@example.com" + } else { + res, err := localClient.WhoIs(r.Context(), r.RemoteAddr) + if err != nil { + return "", err + } + login = res.UserProfile.LoginName + } + return login, nil + +} + var reShortName = regexp.MustCompile(`^[\w\-\.]+$`) +// serveSave handles requests to save or update a Link. Both short name and +// long URL are validated for proper format. Existing links may only be updated +// by their owner. func serveSave(w http.ResponseWriter, r *http.Request) { short, long := r.FormValue("short"), r.FormValue("long") if short == "" || long == "" { @@ -287,75 +272,55 @@ func serveSave(w http.ResponseWriter, r *http.Request) { return } - login := "" - if devMode() { - login = "foo@example.com" - } else { - res, err := localClient.WhoIs(r.Context(), r.RemoteAddr) - if err != nil { - http.Error(w, err.Error(), http.StatusInternalServerError) - return - } - login = res.UserProfile.LoginName - } - - dl, err := loadLink(short) - if err == nil && dl.Owner != login { - http.Error(w, "not your link; owned by "+dl.Owner, http.StatusForbidden) - return - } - - now := time.Now().UTC() - if dl == nil { - dl = &DiskLink{ - Short: short, - Created: now, - } - } - dl.Short = short - dl.Long = long - dl.LastEdit = now - dl.Owner = login - j, err := json.MarshalIndent(dl, "", " ") + login, err := currentUser(r) if err != nil { http.Error(w, err.Error(), http.StatusInternalServerError) return } - if err := os.WriteFile(linkPath(short), j, 0600); err != nil { + + link, err := db.Load(short) + if err == nil && link.Owner != login { + http.Error(w, "not your link; owned by "+link.Owner, http.StatusForbidden) + return + } + + now := time.Now().UTC() + if link == nil { + link = &Link{ + Short: short, + Created: now, + } + } + link.Short = short + link.Long = long + link.LastEdit = now + link.Owner = login + if err := db.Save(link); err != nil { http.Error(w, err.Error(), http.StatusInternalServerError) return } fmt.Fprintf(w, "