343 lines
7.5 KiB
Go
343 lines
7.5 KiB
Go
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))
|
|
}
|