️ perf(download): 启用八路 HTTP Range 并发下载

- 探测服务端 Range 支持并以八个分片并发写入文件
- 保留单路回退、下载进度汇总与 SHA-256 完整性校验
- 补充更新包及驱动资产并发下载回归测试
This commit is contained in:
Syngnat
2026-08-11 09:03:08 +08:00
parent 3adedff09c
commit e25d89759c
3 changed files with 373 additions and 24 deletions

View File

@@ -2053,8 +2053,9 @@ func TestDownloadOptionalDriverAgentFromBundleSharesConcurrentDownload(t *testin
t.Fatalf("bundle install failed: %v", err)
}
}
if got := atomic.LoadInt32(&requestCount); got != 1 {
t.Fatalf("expected one shared bundle download, got %d requests", got)
expectedRequests := int32(1 + updateDownloadParallelism) // 1 次 Range 探测 + 8 个并发分片
if got := atomic.LoadInt32(&requestCount); got != expectedRequests {
t.Fatalf("expected one shared parallel bundle download with %d requests, got %d", expectedRequests, got)
}
}

View File

@@ -1,6 +1,7 @@
package app
import (
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
@@ -37,6 +38,7 @@ const (
updateReleaseCacheTTL = 10 * time.Minute
updateGitHubAPIVersion = "2022-11-28"
updateHTTPBodySnippetLimit = 240
updateDownloadParallelism = 8
)
type cachedGitHubRelease struct {
@@ -63,6 +65,7 @@ var (
)
var errUpdateChecksumMismatch = errors.New("update package checksum mismatch")
var errUpdateRangeUnsupported = errors.New("update server does not support byte ranges")
type updateState struct {
lastCheck *UpdateInfo
@@ -1196,6 +1199,7 @@ func parseSHA256Sums(content string) map[string]string {
}
type downloadProgressWriter struct {
mu sync.Mutex
total int64
written int64
lastEmit time.Time
@@ -1208,6 +1212,8 @@ func (w *downloadProgressWriter) Write(p []byte) (int, error) {
if n == 0 {
return 0, nil
}
w.mu.Lock()
defer w.mu.Unlock()
w.written += int64(n)
if w.onProgress == nil {
return n, nil
@@ -1220,6 +1226,16 @@ func (w *downloadProgressWriter) Write(p []byte) (int, error) {
return n, nil
}
func (w *downloadProgressWriter) finish() int64 {
w.mu.Lock()
defer w.mu.Unlock()
if w.onProgress != nil {
w.lastEmit = time.Now()
w.onProgress(w.written, w.total)
}
return w.written
}
func downloadFileWithHash(url, filePath string, onProgress func(downloaded, total int64)) (string, error) {
return downloadFileWithHashWithTimeout(url, filePath, onProgress, 10*time.Minute)
}
@@ -1229,34 +1245,55 @@ func downloadFileWithHashWithTimeout(url, filePath string, onProgress func(downl
timeout = 10 * time.Minute
}
client := newHTTPClientWithGlobalProxy(timeout)
probeResp, err := doGitHubDownloadRange(client, url, "bytes=0-0", context.Background())
if err != nil {
return "", err
}
if probeResp.StatusCode == http.StatusPartialContent {
start, end, total, valid := parseDownloadContentRange(probeResp.Header.Get("Content-Range"))
_, _ = io.Copy(io.Discard, probeResp.Body)
_ = probeResp.Body.Close()
if valid && start == 0 && end == 0 && total >= updateDownloadParallelism {
hash, parallelErr := downloadFileWithParallelRanges(client, url, filePath, total, onProgress)
if parallelErr == nil {
return hash, nil
}
if !errors.Is(parallelErr, errUpdateRangeUnsupported) {
return "", parallelErr
}
} else {
return downloadFileWithSingleRequest(client, url, filePath, onProgress)
}
} else if probeResp.StatusCode == http.StatusOK {
defer probeResp.Body.Close()
return downloadResponseWithHash(probeResp, filePath, onProgress)
} else {
body, _ := io.ReadAll(io.LimitReader(probeResp.Body, 64<<10))
_ = probeResp.Body.Close()
return "", classifyGitHubUpdateHTTPError(probeResp.StatusCode, body, probeResp.Header, false)
}
return downloadFileWithSingleRequest(client, url, filePath, onProgress)
}
func downloadFileWithSingleRequest(client *http.Client, url, filePath string, onProgress func(downloaded, total int64)) (string, error) {
resp, err := doGitHubDownload(client, url)
if err != nil {
return "", err
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusOK {
body, _ := io.ReadAll(io.LimitReader(resp.Body, 64<<10))
return "", classifyGitHubUpdateHTTPError(resp.StatusCode, body, resp.Header, false)
}
return downloadResponseWithHash(resp, filePath, onProgress)
}
// Windows 上旧文件可能被杀毒软件/索引服务占用,先尝试删除并重试
_ = os.Remove(filePath)
var out *os.File
for retry := 0; retry < 5; retry++ {
out, err = os.Create(filePath)
if err == nil {
break
}
if retry < 4 {
time.Sleep(time.Duration(retry+1) * 500 * time.Millisecond)
}
}
func downloadResponseWithHash(resp *http.Response, filePath string, onProgress func(downloaded, total int64)) (string, error) {
out, err := createUpdateDownloadFile(filePath)
if err != nil {
return "", localizedUpdateError{
key: "app.update.backend.error.package_file_busy",
params: map[string]any{"detail": err.Error()},
}
return "", err
}
hasher := sha256.New()
@@ -1271,16 +1308,14 @@ func downloadFileWithHashWithTimeout(url, filePath string, onProgress func(downl
onProgress(0, total)
}
if _, err := io.Copy(io.MultiWriter(writers...), resp.Body); err != nil {
out.Close()
_ = out.Close()
return "", wrapUpdateNetworkError(err)
}
if onProgress != nil {
onProgress(progressWriter.written, total)
}
progressWriter.finish()
// 显式 Sync + Close确保数据落盘且文件句柄释放
if err := out.Sync(); err != nil {
out.Close()
_ = out.Close()
return "", err
}
if err := out.Close(); err != nil {
@@ -1290,6 +1325,196 @@ func downloadFileWithHashWithTimeout(url, filePath string, onProgress func(downl
return hex.EncodeToString(hasher.Sum(nil)), nil
}
func downloadFileWithParallelRanges(
client *http.Client,
url string,
filePath string,
total int64,
onProgress func(downloaded, total int64),
) (string, error) {
out, err := createUpdateDownloadFile(filePath)
if err != nil {
return "", err
}
if err := out.Truncate(total); err != nil {
_ = out.Close()
return "", err
}
progressWriter := &downloadProgressWriter{
total: total,
emitEvery: 120 * time.Millisecond,
onProgress: onProgress,
}
if onProgress != nil {
onProgress(0, total)
}
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
errCh := make(chan error, updateDownloadParallelism)
for index := 0; index < updateDownloadParallelism; index++ {
start := total * int64(index) / updateDownloadParallelism
end := total*int64(index+1)/updateDownloadParallelism - 1
go func() {
errCh <- downloadUpdateRange(ctx, client, url, out, progressWriter, start, end, total)
}()
}
var firstErr error
rangeUnsupported := false
for range updateDownloadParallelism {
rangeErr := <-errCh
if rangeErr == nil {
continue
}
if errors.Is(rangeErr, errUpdateRangeUnsupported) {
rangeUnsupported = true
}
if firstErr == nil {
firstErr = rangeErr
cancel()
}
}
close(errCh)
if firstErr != nil {
_ = out.Close()
if rangeUnsupported {
return "", errUpdateRangeUnsupported
}
return "", firstErr
}
progressWriter.finish()
if err := out.Sync(); err != nil {
_ = out.Close()
return "", err
}
if err := out.Close(); err != nil {
return "", err
}
return hashDownloadedFile(filePath)
}
func downloadUpdateRange(
ctx context.Context,
client *http.Client,
url string,
out *os.File,
progressWriter *downloadProgressWriter,
start, end, total int64,
) error {
resp, err := doGitHubDownloadRange(client, url, fmt.Sprintf("bytes=%d-%d", start, end), ctx)
if err != nil {
return err
}
defer resp.Body.Close()
if resp.StatusCode != http.StatusPartialContent {
_, _ = io.Copy(io.Discard, io.LimitReader(resp.Body, 64<<10))
if resp.StatusCode == http.StatusOK || resp.StatusCode == http.StatusRequestedRangeNotSatisfiable {
return errUpdateRangeUnsupported
}
return classifyGitHubUpdateHTTPError(resp.StatusCode, nil, resp.Header, false)
}
actualStart, actualEnd, actualTotal, valid := parseDownloadContentRange(resp.Header.Get("Content-Range"))
if !valid || actualStart != start || actualEnd != end || actualTotal != total {
return errUpdateRangeUnsupported
}
expected := end - start + 1
writer := io.NewOffsetWriter(out, start)
written, err := io.Copy(io.MultiWriter(writer, progressWriter), io.LimitReader(resp.Body, expected+1))
if err != nil {
return wrapUpdateNetworkError(err)
}
if written != expected {
return wrapUpdateNetworkError(io.ErrUnexpectedEOF)
}
return nil
}
func doGitHubDownloadRange(client *http.Client, rawURL, byteRange string, ctx context.Context) (*http.Response, error) {
rawURL = strings.TrimSpace(rawURL)
if rawURL == "" {
return nil, localizedUpdateError{
key: "app.update.backend.error.package_download_http_failed",
params: map[string]any{"status": 0},
}
}
req, err := http.NewRequestWithContext(ctx, http.MethodGet, rawURL, nil)
if err != nil {
return nil, err
}
applyGitHubDownloadRequestHeaders(req, isGitHubReleaseAssetAPIURL(rawURL))
req.Header.Set("Range", byteRange)
req.Header.Set("Accept-Encoding", "identity")
return doUpdateRequest(client, req)
}
func parseDownloadContentRange(value string) (start, end, total int64, valid bool) {
fields := strings.Fields(strings.TrimSpace(value))
if len(fields) != 2 || !strings.EqualFold(fields[0], "bytes") {
return 0, 0, 0, false
}
rangePart, totalPart, ok := strings.Cut(fields[1], "/")
if !ok || totalPart == "*" {
return 0, 0, 0, false
}
startPart, endPart, ok := strings.Cut(rangePart, "-")
if !ok {
return 0, 0, 0, false
}
start, err := strconv.ParseInt(startPart, 10, 64)
if err != nil {
return 0, 0, 0, false
}
end, err = strconv.ParseInt(endPart, 10, 64)
if err != nil {
return 0, 0, 0, false
}
total, err = strconv.ParseInt(totalPart, 10, 64)
if err != nil || start < 0 || end < start || total <= end {
return 0, 0, 0, false
}
return start, end, total, true
}
func createUpdateDownloadFile(filePath string) (*os.File, error) {
// Windows 上旧文件可能被杀毒软件/索引服务占用,先尝试删除并重试。
_ = os.Remove(filePath)
var (
out *os.File
err error
)
for retry := 0; retry < 5; retry++ {
out, err = os.Create(filePath)
if err == nil {
return out, nil
}
if retry < 4 {
time.Sleep(time.Duration(retry+1) * 500 * time.Millisecond)
}
}
return nil, localizedUpdateError{
key: "app.update.backend.error.package_file_busy",
params: map[string]any{"detail": err.Error()},
}
}
func hashDownloadedFile(filePath string) (string, error) {
file, err := os.Open(filePath)
if err != nil {
return "", err
}
defer file.Close()
hasher := sha256.New()
if _, err := io.Copy(hasher, file); err != nil {
return "", err
}
return hex.EncodeToString(hasher.Sum(nil)), nil
}
func doUpdateRequest(client *http.Client, req *http.Request) (*http.Response, error) {
resp, err := client.Do(req)
if err == nil {

View File

@@ -1,6 +1,7 @@
package app
import (
"bytes"
"crypto/sha256"
"encoding/json"
"errors"
@@ -117,6 +118,128 @@ func TestDownloadUpdateAssetWithFallbackRetriesChecksumMismatch(t *testing.T) {
}
}
func TestDownloadFileWithHashUsesEightParallelRanges(t *testing.T) {
configureUpdateManifestHTTPTest(t)
payload := make([]byte, 100_003)
for index := range payload {
payload[index] = byte(index % 251)
}
expectedHash := fmt.Sprintf("%x", sha256.Sum256(payload))
var probeHits atomic.Int32
var rangeHits atomic.Int32
started := make(chan struct{}, updateDownloadParallelism)
release := make(chan struct{})
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
rawRange := req.Header.Get("Range")
var start, end int64
if count, err := fmt.Sscanf(rawRange, "bytes=%d-%d", &start, &end); err != nil || count != 2 || start < 0 || end < start || end >= int64(len(payload)) {
http.Error(w, "invalid range", http.StatusBadRequest)
return
}
w.Header().Set("Content-Range", fmt.Sprintf("bytes %d-%d/%d", start, end, len(payload)))
w.Header().Set("Content-Length", fmt.Sprintf("%d", end-start+1))
w.WriteHeader(http.StatusPartialContent)
if rawRange == "bytes=0-0" {
probeHits.Add(1)
_, _ = w.Write(payload[:1])
return
}
rangeHits.Add(1)
started <- struct{}{}
select {
case <-release:
case <-req.Context().Done():
return
}
_, _ = w.Write(payload[start : end+1])
}))
defer server.Close()
type downloadOutcome struct {
hash string
err error
}
assetPath := filepath.Join(t.TempDir(), "GoNavi.bin")
var downloaded atomic.Int64
var total atomic.Int64
done := make(chan downloadOutcome, 1)
go func() {
hash, err := downloadFileWithHashWithTimeout(server.URL, assetPath, func(current, expected int64) {
downloaded.Store(current)
total.Store(expected)
}, 5*time.Second)
done <- downloadOutcome{hash: hash, err: err}
}()
for index := 0; index < updateDownloadParallelism; index++ {
select {
case <-started:
case <-time.After(2 * time.Second):
close(release)
t.Fatalf("only %d parallel range requests started", index)
}
}
close(release)
outcome := <-done
if outcome.err != nil {
t.Fatalf("parallel range download failed: %v", outcome.err)
}
if outcome.hash != expectedHash {
t.Fatalf("hash = %q, want %q", outcome.hash, expectedHash)
}
if probeHits.Load() != 1 || rangeHits.Load() != updateDownloadParallelism {
t.Fatalf("request counts: probe=%d ranges=%d", probeHits.Load(), rangeHits.Load())
}
if downloaded.Load() != int64(len(payload)) || total.Load() != int64(len(payload)) {
t.Fatalf("progress = %d/%d, want %d/%d", downloaded.Load(), total.Load(), len(payload), len(payload))
}
actual, err := os.ReadFile(assetPath)
if err != nil {
t.Fatalf("read downloaded asset: %v", err)
}
if !bytes.Equal(actual, payload) {
t.Fatal("parallel range download reconstructed different content")
}
}
func TestDownloadFileWithHashFallsBackWhenRangeIsUnsupported(t *testing.T) {
configureUpdateManifestHTTPTest(t)
payload := []byte("single request fallback payload")
expectedHash := fmt.Sprintf("%x", sha256.Sum256(payload))
var hits atomic.Int32
var rangeHits atomic.Int32
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, req *http.Request) {
hits.Add(1)
if req.Header.Get("Range") != "" {
rangeHits.Add(1)
}
w.WriteHeader(http.StatusOK)
_, _ = w.Write(payload)
}))
defer server.Close()
assetPath := filepath.Join(t.TempDir(), "GoNavi.bin")
actualHash, err := downloadFileWithHashWithTimeout(server.URL, assetPath, nil, 5*time.Second)
if err != nil {
t.Fatalf("single request fallback failed: %v", err)
}
if actualHash != expectedHash {
t.Fatalf("hash = %q, want %q", actualHash, expectedHash)
}
if hits.Load() != 1 || rangeHits.Load() != 1 {
t.Fatalf("request counts: total=%d range=%d, want 1/1", hits.Load(), rangeHits.Load())
}
actual, err := os.ReadFile(assetPath)
if err != nil {
t.Fatalf("read downloaded asset: %v", err)
}
if !bytes.Equal(actual, payload) {
t.Fatal("single request fallback wrote different content")
}
}
func TestReleaseFromUpdateManifestMapsAssets(t *testing.T) {
release := releaseFromUpdateManifest(&updateReleaseManifest{
TagName: "v1.2.3",