137 lines
3.8 KiB
Go
137 lines
3.8 KiB
Go
package serve
|
|
|
|
import (
|
|
"context"
|
|
"encoding/json"
|
|
"errors"
|
|
"fmt"
|
|
"io"
|
|
"net/http"
|
|
"strconv"
|
|
"strings"
|
|
"time"
|
|
|
|
"gitea.dddbg.com/youbin/bdrclone/internal/baidu"
|
|
)
|
|
|
|
type Server struct {
|
|
client *baidu.Client
|
|
server *http.Server
|
|
}
|
|
|
|
func New(client *baidu.Client, address string) *Server {
|
|
s := &Server{client: client}
|
|
s.server = &http.Server{Addr: address, Handler: s, ReadHeaderTimeout: 10 * time.Second}
|
|
return s
|
|
}
|
|
|
|
func (s *Server) ListenAndServe(ctx context.Context) error {
|
|
errCh := make(chan error, 1)
|
|
go func() { errCh <- s.server.ListenAndServe() }()
|
|
select {
|
|
case err := <-errCh:
|
|
if errors.Is(err, http.ErrServerClosed) {
|
|
return nil
|
|
}
|
|
return err
|
|
case <-ctx.Done():
|
|
shutdownCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
|
defer cancel()
|
|
return errors.Join(s.server.Shutdown(shutdownCtx), ctx.Err())
|
|
}
|
|
}
|
|
|
|
func (s *Server) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
|
if r.Method != http.MethodGet && r.Method != http.MethodHead {
|
|
w.Header().Set("Allow", "GET, HEAD")
|
|
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
|
|
return
|
|
}
|
|
entry, err := s.client.Stat(r.Context(), r.URL.Path)
|
|
if err != nil {
|
|
if errors.Is(err, baidu.ErrNotFound) {
|
|
http.NotFound(w, r)
|
|
return
|
|
}
|
|
http.Error(w, err.Error(), http.StatusBadGateway)
|
|
return
|
|
}
|
|
if entry.IsDirectory() {
|
|
s.serveDirectory(w, r)
|
|
return
|
|
}
|
|
s.serveFile(w, r, entry)
|
|
}
|
|
|
|
func (s *Server) serveDirectory(w http.ResponseWriter, r *http.Request) {
|
|
entries, err := s.client.List(r.Context(), r.URL.Path)
|
|
if err != nil {
|
|
http.Error(w, err.Error(), http.StatusBadGateway)
|
|
return
|
|
}
|
|
w.Header().Set("Content-Type", "application/json; charset=utf-8")
|
|
if r.Method == http.MethodHead {
|
|
return
|
|
}
|
|
_ = json.NewEncoder(w).Encode(entries)
|
|
}
|
|
|
|
func (s *Server) serveFile(w http.ResponseWriter, r *http.Request, file baidu.File) {
|
|
offset, length, partial, err := parseRange(r.Header.Get("Range"), file.Size)
|
|
if err != nil {
|
|
w.Header().Set("Content-Range", fmt.Sprintf("bytes */%d", file.Size))
|
|
http.Error(w, err.Error(), http.StatusRequestedRangeNotSatisfiable)
|
|
return
|
|
}
|
|
w.Header().Set("Accept-Ranges", "bytes")
|
|
w.Header().Set("Last-Modified", file.ModTime().UTC().Format(http.TimeFormat))
|
|
if r.Method == http.MethodHead || length == 0 {
|
|
w.Header().Set("Content-Length", strconv.FormatInt(length, 10))
|
|
if partial {
|
|
w.Header().Set("Content-Range", fmt.Sprintf("bytes %d-%d/%d", offset, offset+length-1, file.Size))
|
|
w.WriteHeader(http.StatusPartialContent)
|
|
}
|
|
return
|
|
}
|
|
body, err := s.client.Open(r.Context(), file, offset, length)
|
|
if err != nil {
|
|
http.Error(w, err.Error(), http.StatusBadGateway)
|
|
return
|
|
}
|
|
defer body.Close()
|
|
w.Header().Set("Content-Length", strconv.FormatInt(length, 10))
|
|
if partial {
|
|
w.Header().Set("Content-Range", fmt.Sprintf("bytes %d-%d/%d", offset, offset+length-1, file.Size))
|
|
w.WriteHeader(http.StatusPartialContent)
|
|
}
|
|
_, _ = io.CopyN(w, body, length)
|
|
}
|
|
|
|
func parseRange(value string, size int64) (offset, length int64, partial bool, err error) {
|
|
if value == "" {
|
|
return 0, size, false, nil
|
|
}
|
|
if !strings.HasPrefix(value, "bytes=") || strings.Contains(value, ",") {
|
|
return 0, 0, false, errors.New("only one bytes range is supported")
|
|
}
|
|
parts := strings.SplitN(strings.TrimPrefix(value, "bytes="), "-", 2)
|
|
if len(parts) != 2 || parts[0] == "" {
|
|
return 0, 0, false, errors.New("suffix ranges are not supported")
|
|
}
|
|
start, parseErr := strconv.ParseInt(parts[0], 10, 64)
|
|
if parseErr != nil || start < 0 || start >= size {
|
|
return 0, 0, false, errors.New("invalid range start")
|
|
}
|
|
end := size - 1
|
|
if parts[1] != "" {
|
|
end, parseErr = strconv.ParseInt(parts[1], 10, 64)
|
|
if parseErr != nil || end < start {
|
|
return 0, 0, false, errors.New("invalid range end")
|
|
}
|
|
if end >= size {
|
|
end = size - 1
|
|
}
|
|
}
|
|
return start, end - start + 1, true, nil
|
|
}
|