summaryrefslogtreecommitdiffstats
path: root/roseh.moe.go
diff options
context:
space:
mode:
authorRose Hogenson <rosehogenson@posteo.net>2026-06-07 21:31:35 -0700
committerRose Hogenson <rosehogenson@posteo.net>2026-06-07 21:31:35 -0700
commit027919fac08e33fe41d44bb5dedf4cbd02a8b9e5 (patch)
tree415f7ee43eae923966e8d88d1ed834ab26fa57fa /roseh.moe.go
parent8d1eff61416040a9d77982650200cbe7b245f833 (diff)
downloadroseh.moe-027919fac08e33fe41d44bb5dedf4cbd02a8b9e5.tar.zst
Allow uploading multiple files to the wormhole
Diffstat (limited to 'roseh.moe.go')
-rw-r--r--roseh.moe.go201
1 files changed, 151 insertions, 50 deletions
diff --git a/roseh.moe.go b/roseh.moe.go
index 40bc076..075748c 100644
--- a/roseh.moe.go
+++ b/roseh.moe.go
@@ -1,6 +1,9 @@
package main
import (
+ "archive/tar"
+ "cmp"
+ "compress/gzip"
"crypto/hkdf"
"crypto/hmac"
"crypto/rand"
@@ -249,34 +252,40 @@ func makeHole() string {
const blockSize = 4 * 1024 * 1024
-func writeHoleFile(part *multipart.Part) (string, error) {
+type holeFileWriter struct {
+ *oae2.Writer
+ f *os.File
+}
+
+func (w *holeFileWriter) Close() error {
+ return cmp.Or(w.Writer.Close(), w.f.Close())
+}
+
+func writeHoleFile(filename string) (_ string, _ *holeFileWriter, err error) {
hole := makeHole()
primaryKey, err := hkdf.Key(sha256.New, []byte(hole), nil, "", 32)
if err != nil {
- return "", err
+ return "", nil, err
}
f, err := os.Create(fmt.Sprintf("%s/%x", *holeTempDir, primaryKey))
if err != nil {
- return "", err
+ return "", nil, err
}
- defer f.Close()
+ defer func() {
+ if err != nil {
+ f.Close()
+ }
+ }()
w := oae2.NewWriter(f, []byte(hole), blockSize, nil)
- fileName := part.FileName()
buf := make([]byte, 8)
- binary.BigEndian.PutUint64(buf, uint64(len(fileName)))
+ binary.BigEndian.PutUint64(buf, uint64(len(filename)))
if _, err := w.Write(buf); err != nil {
- return "", err
- }
- if _, err := io.WriteString(w, fileName); err != nil {
- return "", err
- }
- if _, err := io.Copy(w, part); err != nil {
- return "", err
+ return "", nil, err
}
- if err := w.Close(); err != nil {
- return "", err
+ if _, err := io.WriteString(w, filename); err != nil {
+ return "", nil, err
}
- return hole, f.Close()
+ return hole, &holeFileWriter{w, f}, nil
}
type offsetReader struct {
@@ -292,38 +301,59 @@ func (r *offsetReader) Seek(offset int64, whence int) (int64, error) {
return n - int64(r.offset), err
}
+type holeFileReader struct {
+ *offsetReader
+ f *os.File
+}
+
+func (r *holeFileReader) Close() error {
+ return r.f.Close()
+}
+
var errNoHole = errors.New("no such hole")
-func readHoleFile(w http.ResponseWriter, req *http.Request, hole string) error {
+func openHoleFile(hole string) (_ string, _ *holeFileReader, err error) {
primaryKey, err := hkdf.Key(sha256.New, []byte(hole), nil, "", 32)
if err != nil {
- return err
+ return "", nil, err
}
f, err := os.Open(fmt.Sprintf("%s/%x", *holeTempDir, primaryKey))
if err != nil {
if errors.Is(err, fs.ErrNotExist) {
- return errNoHole
+ return "", nil, errNoHole
}
- return err
+ return "", nil, err
}
- defer f.Close()
+ defer func() {
+ if err != nil {
+ f.Close()
+ }
+ }()
r := oae2.NewReader(f, []byte(hole), blockSize, nil)
buf := make([]byte, 8)
if _, err := io.ReadFull(r, buf); err != nil {
- return err
+ return "", nil, err
}
fileNameLen := binary.BigEndian.Uint64(buf)
- buf = make([]byte, fileNameLen)
- if _, err := io.ReadFull(r, buf); err != nil {
+ fileName := make([]byte, fileNameLen)
+ if _, err := io.ReadFull(r, fileName); err != nil {
+ return "", nil, err
+ }
+ return string(fileName), &holeFileReader{&offsetReader{r, 8 + int(fileNameLen)}, f}, nil
+}
+
+func readHoleFile(w http.ResponseWriter, req *http.Request, hole string) error {
+ fileName, r, err := openHoleFile(hole)
+ if err != nil {
return err
}
- fileName := string(buf)
+ defer r.Close()
w.Header().Set("Content-Disposition", "attachment; filename*=UTF-8''"+fileName)
var modTime time.Time
- if stat, err := f.Stat(); err == nil {
+ if stat, err := r.f.Stat(); err == nil {
modTime = stat.ModTime()
}
- http.ServeContent(w, req, fileName, modTime, &offsetReader{r, 8 + int(fileNameLen)})
+ http.ServeContent(w, req, fileName, modTime, r)
return nil
}
@@ -333,6 +363,84 @@ func wormhole(w http.ResponseWriter, r *http.Request) {
}
}
+func writeHoleFileToTarFile(tarWriter *tar.Writer, hole string, now time.Time) error {
+ fileName, f, err := openHoleFile(hole)
+ if err != nil {
+ return fmt.Errorf("read back hole file: %s", err)
+ }
+ defer f.Close()
+ size, err := f.Seek(0, io.SeekEnd)
+ if err != nil {
+ return fmt.Errorf("read back hole file: %s", err)
+ }
+ if _, err := f.Seek(0, io.SeekStart); err != nil {
+ return fmt.Errorf("read back hole file: %s", err)
+ }
+ if err := tarWriter.WriteHeader(&tar.Header{Name: "wormhole-files/" + fileName, Mode: 0644, Size: size, ModTime: now}); err != nil {
+ return err
+ }
+ if _, err := io.Copy(tarWriter, f); err != nil {
+ return fmt.Errorf("write file to tar file: %s", err)
+ }
+ return nil
+}
+
+var errNoFile = errors.New("no file")
+
+func uploadWormhole(reader *multipart.Reader) (string, error) {
+ var holes []string
+ for {
+ part, err := reader.NextPart()
+ if err != nil {
+ break
+ }
+ if part.FormName() != "file" || part.FileName() == "" {
+ continue
+ }
+ hole, holeWriter, err := writeHoleFile(part.FileName())
+ if err != nil {
+ return "", fmt.Errorf("open hole file: %s", err)
+ }
+ _, err = io.Copy(holeWriter, part)
+ if err := cmp.Or(err, holeWriter.Close()); err != nil {
+ return "", fmt.Errorf("write hole file: %s", err)
+ }
+ holes = append(holes, hole)
+ }
+ if len(holes) == 0 {
+ return "", errNoFile
+ }
+ if len(holes) == 1 {
+ return holes[0], nil
+ }
+ tarFileHole, tarHoleWriter, err := writeHoleFile("wormhole-files.tar.gz")
+ if err != nil {
+ return "", fmt.Errorf("make tar file: %s", err)
+ }
+ defer tarHoleWriter.Close()
+ gzipWriter := gzip.NewWriter(tarHoleWriter)
+ tarWriter := tar.NewWriter(gzipWriter)
+ now := time.Now()
+ if err := tarWriter.WriteHeader(&tar.Header{Name: "wormhole-files/", Mode: 0755, ModTime: now}); err != nil {
+ return "", fmt.Errorf("create tar file main directory: %s", err)
+ }
+ for _, hole := range holes {
+ if err := writeHoleFileToTarFile(tarWriter, hole, now); err != nil {
+ return "", err
+ }
+ }
+ if err := tarWriter.Close(); err != nil {
+ return "", fmt.Errorf("write tar file: %s", err)
+ }
+ if err := gzipWriter.Close(); err != nil {
+ return "", fmt.Errorf("gzip tar file: %s", err)
+ }
+ if err := tarHoleWriter.Close(); err != nil {
+ return "", fmt.Errorf("write tar hole file: %s", err)
+ }
+ return tarFileHole, nil
+}
+
var (
//go:embed templates/wormhole-success.html.template
wormholeSuccessTemplateString string
@@ -350,32 +458,25 @@ func wormholeUpload(w http.ResponseWriter, r *http.Request) {
http.Error(w, fmt.Sprintf("Not a multipart/form-data request: %s", err), http.StatusBadRequest)
return
}
- for {
- part, err := reader.NextPart()
- if err != nil {
- break
- }
- if part.FormName() != "file" || part.FileName() == "" {
- continue
- }
- hole, err := writeHoleFile(part)
- if err != nil {
- log.Printf("Warning: failed to open hole file: %s", err)
- http.Error(w, "Internal error", http.StatusInternalServerError)
+ hole, err := uploadWormhole(reader)
+ if err != nil {
+ if err == errNoFile {
+ w.WriteHeader(http.StatusBadRequest)
+ if err := wormholeTemplate.Execute(w, wormholeTemplateArgs{Err: "don't forget to pick a file"}); err != nil {
+ log.Printf("Warning: wormholeTemplate: %s", err)
+ }
return
}
- downloadURL := *selfURL + "/wormhole/" + hole
- if err := wormholeSuccessTemplate.Execute(w, wormholeSuccessTemplateArgs{
- URL: downloadURL,
- QR: "/qr.png?t=" + url.QueryEscape(downloadURL) + "&s=" + base64.RawURLEncoding.EncodeToString(mac([]byte(downloadURL), "qr")),
- }); err != nil {
- log.Printf("Warning: wormholeSend: %s", err)
- }
+ log.Printf("Warning: failed to upload hole file: %s", err)
+ http.Error(w, "Internal error", http.StatusInternalServerError)
return
}
- w.WriteHeader(http.StatusBadRequest)
- if err := wormholeTemplate.Execute(w, wormholeTemplateArgs{Err: "don't forget to pick a file"}); err != nil {
- log.Printf("Warning: wormholeTemplate: %s", err)
+ downloadURL := *selfURL + "/wormhole/" + hole
+ if err := wormholeSuccessTemplate.Execute(w, wormholeSuccessTemplateArgs{
+ URL: downloadURL,
+ QR: "/qr.png?t=" + url.QueryEscape(downloadURL) + "&s=" + base64.RawURLEncoding.EncodeToString(mac([]byte(downloadURL), "qr")),
+ }); err != nil {
+ log.Printf("Warning: wormholeSend: %s", err)
}
}
@@ -705,7 +806,7 @@ func main() {
}
go func() {
- for range time.NewTicker(time.Hour).C {
+ for range time.NewTicker(2 * time.Hour).C {
if err := cleanHole(); err != nil {
log.Printf("Warning: clean hole: %s", err)
}