diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..b3da701 --- /dev/null +++ b/.gitignore @@ -0,0 +1,5 @@ +see.db-shm +see +see.db +see.db-wal +go.sum diff --git a/favicon.png b/favicon.png new file mode 100644 index 0000000..864f3f5 Binary files /dev/null and b/favicon.png differ diff --git a/go.mod b/go.mod new file mode 100644 index 0000000..d2775ff --- /dev/null +++ b/go.mod @@ -0,0 +1,21 @@ +module see + +go 1.21 + +require modernc.org/sqlite v1.29.5 + +require ( + github.com/dustin/go-humanize v1.0.1 // indirect + github.com/google/uuid v1.3.0 // indirect + github.com/hashicorp/golang-lru/v2 v2.0.7 // indirect + github.com/mattn/go-isatty v0.0.16 // indirect + github.com/ncruces/go-strftime v0.1.9 // indirect + github.com/remyoudompheng/bigfft v0.0.0-20230129092748-24d4a6f8daec // indirect + golang.org/x/sys v0.16.0 // indirect + modernc.org/gc/v3 v3.0.0-20240107210532-573471604cb6 // indirect + modernc.org/libc v1.41.0 // indirect + modernc.org/mathutil v1.6.0 // indirect + modernc.org/memory v1.7.2 // indirect + modernc.org/strutil v1.2.0 // indirect + modernc.org/token v1.1.0 // indirect +) diff --git a/index.html b/index.html new file mode 100644 index 0000000..41911ba --- /dev/null +++ b/index.html @@ -0,0 +1,278 @@ + + + + + +see + + + + + + +
+ +
+ +
+ +
+ + +
+ +
+ +
+ +
Short URL
+ + + +
+ + + +
+ + + +
+ + + + + diff --git a/main.go b/main.go new file mode 100644 index 0000000..193f5ac --- /dev/null +++ b/main.go @@ -0,0 +1,343 @@ +package main + +import ( + "crypto/rand" + "crypto/sha256" + "database/sql" + _ "embed" + "encoding/hex" + "encoding/json" + "flag" + "fmt" + "log" + "math" + "math/big" + "net" + "net/http" + "net/url" + "strings" + "sync" + "time" + + _ "modernc.org/sqlite" +) + +const ( + alphabet = "0123456789abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ" + codeLen = 6 +) + +var ( + db *sql.DB + limiter = newRateLimiter(10, 10) +) + +//go:embed index.html +var indexHTML string + +//go:embed favicon.png +var faviconPNG []byte + +type resp struct { + Error string `json:"error,omitempty"` +} + +type visitor struct { + tokens float64 + lastSeen time.Time +} + +type rateLimiter struct { + mu sync.RWMutex + visitors map[string]*visitor + rate float64 + burst int +} + +func newRateLimiter(perMinute int, burst int) *rateLimiter { + rl := &rateLimiter{ + visitors: make(map[string]*visitor), + rate: float64(perMinute) / 60.0, + burst: burst, + } + go rl.cleanup() + return rl +} + +func (rl *rateLimiter) allow(ip string) bool { + rl.mu.Lock() + defer rl.mu.Unlock() + + v, ok := rl.visitors[ip] + if !ok { + rl.visitors[ip] = &visitor{tokens: float64(rl.burst) - 1, lastSeen: time.Now()} + return true + } + + now := time.Now() + v.tokens = math.Min(float64(rl.burst), v.tokens+now.Sub(v.lastSeen).Seconds()*rl.rate) + v.lastSeen = now + + if v.tokens >= 1 { + v.tokens-- + return true + } + return false +} + +func (rl *rateLimiter) cleanup() { + for { + time.Sleep(time.Minute) + rl.mu.Lock() + for ip, v := range rl.visitors { + if time.Since(v.lastSeen) > 3*time.Minute { + delete(rl.visitors, ip) + } + } + rl.mu.Unlock() + } +} + +func clientIP(r *http.Request) string { + ip := r.Header.Get("X-Forwarded-For") + if ip == "" { + ip = r.Header.Get("X-Real-Ip") + } + if ip == "" { + ip = r.RemoteAddr + } + if i := strings.LastIndex(ip, ":"); i != -1 { + ip = ip[:i] + } + return ip +} + +func rateLimit(next http.HandlerFunc) http.HandlerFunc { + return func(w http.ResponseWriter, r *http.Request) { + if !limiter.allow(clientIP(r)) { + w.WriteHeader(http.StatusTooManyRequests) + json.NewEncoder(w).Encode(resp{Error: "rate limit exceeded"}) + return + } + next(w, r) + } +} + +func hashURL(s string) string { + h := sha256.Sum256([]byte(s)) + return hex.EncodeToString(h[:]) +} + +func generateCode() string { + b := make([]byte, codeLen) + for i := range b { + n, _ := rand.Int(rand.Reader, big.NewInt(int64(len(alphabet)))) + b[i] = alphabet[n.Int64()] + } + return string(b) +} + +func baseURL(r *http.Request) string { + scheme := "http" + if r.TLS != nil || r.Header.Get("X-Forwarded-Proto") == "https" { + scheme = "https" + } + host := r.Host + if h := r.Header.Get("X-Forwarded-Host"); h != "" { + host = h + } + return scheme + "://" + host +} + +func writeShort(w http.ResponseWriter, r *http.Request, code string) { + short := baseURL(r) + "/" + code + w.Header().Set("Content-Type", "text/plain") + fmt.Fprintln(w, short) +} + +func initDB(path string) { + var err error + db, err = sql.Open("sqlite", path+"?_pragma=busy_timeout(5000)&_pragma=journal_mode(WAL)&_pragma=synchronous(NORMAL)") + if err != nil { + log.Fatal(err) + } + + if _, err := db.Exec(`CREATE TABLE IF NOT EXISTS urls ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + code TEXT UNIQUE NOT NULL, + url TEXT NOT NULL, + url_hash TEXT, + created_at DATETIME DEFAULT CURRENT_TIMESTAMP + )`); err != nil { + log.Fatal(err) + } + + if _, err := db.Exec("ALTER TABLE urls ADD COLUMN url_hash TEXT"); err != nil { + if !strings.Contains(err.Error(), "duplicate column name") { + log.Fatal(err) + } + } + + if _, err := db.Exec("CREATE INDEX IF NOT EXISTS idx_hash ON urls(url_hash)"); err != nil { + log.Fatal(err) + } + if _, err := db.Exec("CREATE INDEX IF NOT EXISTS idx_code ON urls(code)"); err != nil { + log.Fatal(err) + } + + rows, err := db.Query("SELECT id, url FROM urls WHERE url_hash IS NULL") + if err != nil { + log.Fatal(err) + } + for rows.Next() { + var id int + var u string + if err := rows.Scan(&id, &u); err != nil { + continue + } + db.Exec("UPDATE urls SET url_hash = ? WHERE id = ?", hashURL(u), id) + } + rows.Close() +} + +func normalizeURL(u *url.URL) string { + if u.Path == "/" { + u.Path = "" + } + if (u.Scheme == "http" && u.Port() == "80") || (u.Scheme == "https" && u.Port() == "443") { + u.Host = u.Hostname() + } + return u.String() +} + +func shorten(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + http.Error(w, "method not allowed", http.StatusMethodNotAllowed) + return + } + + var raw string + ct := r.Header.Get("Content-Type") + if strings.Contains(ct, "application/json") { + var req struct{ URL string `json:"url"` } + if json.NewDecoder(r.Body).Decode(&req) != nil || req.URL == "" { + w.WriteHeader(http.StatusBadRequest) + json.NewEncoder(w).Encode(resp{Error: "url required"}) + return + } + raw = req.URL + } else { + if err := r.ParseForm(); err == nil { + raw = r.FormValue("url") + } + } + + if raw == "" { + w.WriteHeader(http.StatusBadRequest) + json.NewEncoder(w).Encode(resp{Error: "url required"}) + return + } + + if !strings.HasPrefix(raw, "http://") && !strings.HasPrefix(raw, "https://") { + raw = "https://" + raw + } + u, err := url.Parse(raw) + if err != nil { + w.WriteHeader(http.StatusBadRequest) + json.NewEncoder(w).Encode(resp{Error: "invalid url"}) + return + } + + host := u.Hostname() + if host == "" || host == "." || (net.ParseIP(host) == nil && !strings.Contains(host, ".")) { + w.WriteHeader(http.StatusBadRequest) + json.NewEncoder(w).Encode(resp{Error: "invalid url"}) + return + } + + normalized := normalizeURL(u) + urlHash := hashURL(normalized) + + var existingCode string + if err := db.QueryRow("SELECT code FROM urls WHERE url_hash = ? AND url = ?", urlHash, normalized).Scan(&existingCode); err == nil { + writeShort(w, r, existingCode) + return + } + + var code string + for i := 0; i < 10; i++ { + code = generateCode() + if _, err := db.Exec("INSERT INTO urls (code, url, url_hash) VALUES (?, ?, ?)", code, normalized, urlHash); err == nil { + break + } + if !strings.Contains(err.Error(), "UNIQUE constraint failed") { + log.Printf("db error: %v", err) + w.WriteHeader(http.StatusInternalServerError) + json.NewEncoder(w).Encode(resp{Error: "internal error"}) + return + } + code = "" + } + if code == "" { + w.WriteHeader(http.StatusInternalServerError) + json.NewEncoder(w).Encode(resp{Error: "could not generate code"}) + return + } + + writeShort(w, r, code) +} + +func redirect(w http.ResponseWriter, r *http.Request) { + code := strings.TrimPrefix(r.URL.Path, "/") + if code == "" { + http.Redirect(w, r, "/", http.StatusTemporaryRedirect) + return + } + + var target string + if err := db.QueryRow("SELECT url FROM urls WHERE code = ?", code).Scan(&target); err != nil { + if err == sql.ErrNoRows { + http.NotFound(w, r) + return + } + log.Printf("db error: %v", err) + http.Error(w, "internal error", http.StatusInternalServerError) + return + } + http.Redirect(w, r, target, http.StatusMovedPermanently) +} + +func favicon(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "image/png") + w.Write(faviconPNG) +} + +func index(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "text/html; charset=utf-8") + fmt.Fprint(w, indexHTML) +} + +func main() { + port := flag.String("p", "8080", "port") + dbPath := flag.String("db", "see.db", "sqlite path") + flag.Parse() + + initDB(*dbPath) + defer db.Close() + + http.HandleFunc("/", func(w http.ResponseWriter, r *http.Request) { + switch { + case r.URL.Path == "/": + index(w, r) + case r.URL.Path == "/favicon.ico" || r.URL.Path == "/favicon.png": + favicon(w, r) + case r.URL.Path == "/api/": + rateLimit(shorten)(w, r) + default: + redirect(w, r) + } + }) + + log.Printf("see running on :%s", *port) + log.Fatal(http.ListenAndServe(":"+*port, nil)) +}