225 lines
5.6 KiB
Go
225 lines
5.6 KiB
Go
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, `<!doctype html><html><head><meta name="viewport" content="width=device-width, initial-scale=1"><title>Drops</title></head><body><h1>Drops</h1><ul>`)
|
|
for _, e := range entries {
|
|
if e.IsDir() {
|
|
continue
|
|
}
|
|
info, _ := e.Info()
|
|
name := e.Name()
|
|
fmt.Fprintf(w, `<li><a href="/download?name=%s">%s</a> (%s)</li>`,
|
|
url.QueryEscape(name), html.EscapeString(name), formatBytes(info.Size()))
|
|
}
|
|
fmt.Fprint(w, `</ul></body></html>`)
|
|
}
|
|
}
|
|
|
|
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"
|
|
}
|