file-share/receiver/main.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"
}