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 }