diff options
| author | Rose Hogenson <rosehogenson@posteo.net> | 2026-06-07 21:31:35 -0700 |
|---|---|---|
| committer | Rose Hogenson <rosehogenson@posteo.net> | 2026-06-07 21:31:35 -0700 |
| commit | 027919fac08e33fe41d44bb5dedf4cbd02a8b9e5 (patch) | |
| tree | 415f7ee43eae923966e8d88d1ed834ab26fa57fa /roseh.moe.go | |
| parent | 8d1eff61416040a9d77982650200cbe7b245f833 (diff) | |
| download | roseh.moe-027919fac08e33fe41d44bb5dedf4cbd02a8b9e5.tar.zst | |
Allow uploading multiple files to the wormhole
Diffstat (limited to 'roseh.moe.go')
| -rw-r--r-- | roseh.moe.go | 201 |
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) } |
