package main import ( "crypto/rand" "encoding/hex" "fmt" "html" "io" "log" "net" "net/http" "net/url" "os" "path/filepath" "strconv" "strings" ) func main() { home, _ := os.UserHomeDir() cfgDir := filepath.Join(home, ".file-share") _ = os.MkdirAll(cfgDir, 0o755) cfgPath := filepath.Join(cfgDir, "receiver.env") dropDir := env("DROP_DIR", filepath.Join(home, "Drops")) port := env("PORT", "8787") host := env("HOST", lanAddress()) authToken := env("AUTH_TOKEN", "") if authToken == "" { authToken = loadOrGenerateToken(cfgPath) } _ = os.MkdirAll(dropDir, 0o755) log.Printf("drop dir: %s", dropDir) log.Printf("auth token: %s", authToken) log.Printf("listening on http://%s:%s", host, port) mux := http.NewServeMux() mux.HandleFunc("GET /{$}", func(w http.ResponseWriter, r *http.Request) { fmt.Fprint(w, "ok") }) mux.HandleFunc("POST /upload", handleUpload(dropDir, authToken)) mux.HandleFunc("GET /files", handleList(dropDir)) mux.HandleFunc("GET /download", handleDownload(dropDir)) log.Fatal(http.ListenAndServe(host+":"+port, mux)) } func handleUpload(dropDir, authToken string) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { if authToken != "" && r.Header.Get("X-Auth") != authToken { http.Error(w, "unauthorized", http.StatusUnauthorized) return } if err := r.ParseMultipartForm(32 << 20); err != nil { http.Error(w, err.Error(), http.StatusBadRequest) return } files := r.MultipartForm.File["file"] if len(files) == 0 { http.Error(w, "no file", http.StatusBadRequest) return } type savedFile struct { Saved string `json:"saved"` } saved := make([]savedFile, 0, len(files)) for _, fh := range files { src, err := fh.Open() if err != nil { http.Error(w, err.Error(), http.StatusInternalServerError) return } dest := uniquePath(dropDir, filepath.Base(fh.Filename)) dst, err := os.Create(dest) if err != nil { _ = src.Close() http.Error(w, err.Error(), http.StatusInternalServerError) return } n, err := io.Copy(dst, src) _ = src.Close() _ = dst.Close() if err != nil { http.Error(w, err.Error(), http.StatusInternalServerError) return } name := filepath.Base(dest) log.Printf("saved %s (%d bytes) from %s", name, n, r.RemoteAddr) saved = append(saved, savedFile{Saved: name}) } w.Header().Set("Content-Type", "application/json") if len(saved) == 1 { fmt.Fprintf(w, `{"saved":%q}`, saved[0].Saved) return } fmt.Fprint(w, "[") for i, s := range saved { if i > 0 { fmt.Fprint(w, ",") } fmt.Fprintf(w, `{"saved":%q}`, s.Saved) } fmt.Fprint(w, "]") } } func handleList(dropDir string) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { entries, err := os.ReadDir(dropDir) if err != nil { http.Error(w, err.Error(), http.StatusInternalServerError) return } w.Header().Set("Content-Type", "text/html; charset=utf-8") fmt.Fprint(w, `Drops

Drops

`) } } func handleDownload(dropDir string) http.HandlerFunc { return func(w http.ResponseWriter, r *http.Request) { name := r.URL.Query().Get("name") if name == "" { http.Error(w, "missing name", http.StatusBadRequest) return } path := filepath.Join(dropDir, filepath.Base(name)) http.ServeFile(w, r, path) } } func uniquePath(dir, filename string) string { if filename == "" { filename = "unnamed" } ext := filepath.Ext(filename) base := strings.TrimSuffix(filename, ext) dest := filepath.Join(dir, filename) if _, err := os.Stat(dest); os.IsNotExist(err) { return dest } for i := 1; ; i++ { candidate := filepath.Join(dir, base+" ("+strconv.Itoa(i)+")"+ext) if _, err := os.Stat(candidate); os.IsNotExist(err) { return candidate } } } func loadOrGenerateToken(path string) string { if data, err := os.ReadFile(path); err == nil { for _, line := range strings.Split(string(data), "\n") { if strings.HasPrefix(line, "AUTH_TOKEN=") { return strings.TrimPrefix(line, "AUTH_TOKEN=") } } } b := make([]byte, 4) if _, err := rand.Read(b); err != nil { log.Fatal(err) } token := hex.EncodeToString(b) _ = os.WriteFile(path, []byte("AUTH_TOKEN="+token+"\n"), 0o600) log.Printf("generated token, saved to %s", path) return token } func lanAddress() string { addrs, err := net.InterfaceAddrs() if err != nil { return "127.0.0.1" } var fallback string for _, a := range addrs { ipnet, ok := a.(*net.IPNet) if !ok || ipnet.IP.IsLoopback() || ipnet.IP.To4() == nil { continue } ip := ipnet.IP.To4() if ip[0] == 192 && ip[1] == 168 { return ip.String() } if fallback == "" && (ip[0] == 10 || (ip[0] == 172 && ip[1] >= 16 && ip[1] <= 31)) { fallback = ip.String() } } if fallback != "" { return fallback } return "127.0.0.1" } func env(key, def string) string { if v := os.Getenv(key); v != "" { return v } return def } func formatBytes(n int64) string { if n < 1024 { return strconv.FormatInt(n, 10) + " B" } units := []string{"KB", "MB", "GB", "TB"} f := float64(n) / 1024 for _, u := range units { if f < 1024 || u == units[len(units)-1] { return fmt.Sprintf("%.1f %s", f, u) } f /= 1024 } return strconv.FormatInt(n, 10) + " B" }