feat(pages): 支持 GitHub Release 部署源

增加 latest/tag 手动检查与同步、ETag 与限流退避、资源替换确认,以及对应的前端来源管理和部署来源展示。
This commit is contained in:
deqiying
2026-07-19 18:31:42 +08:00
parent 38b0516937
commit c39a3edcc3
34 changed files with 5751 additions and 278 deletions
@@ -0,0 +1,725 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package githubrelease resolves and downloads public GitHub Release assets.
// It deliberately does not know about Pages projects, deployments or runtime
// state so other callers can reuse the same constrained HTTP contract.
package githubrelease
import (
"bytes"
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"io"
"math"
"net/http"
"net/url"
"os"
"path"
"regexp"
"strconv"
"strings"
"time"
"unicode"
"unicode/utf8"
)
const (
// APIVersion is the GitHub REST API contract used by this package.
APIVersion = "2026-03-10"
// SelectorLatest uses GitHub's repository latest-release endpoint.
SelectorLatest Selector = "latest"
// SelectorTag resolves one exact GitHub release tag.
SelectorTag Selector = "tag"
defaultAPIBaseURL = "https://api.github.com"
defaultUserAgent = "OpenFlare-GitHubRelease/1.0"
metadataAccept = "application/vnd.github+json"
assetAccept = "application/octet-stream"
maxMetadataBytes = 4 << 20
maxAssetErrorNames = 10
maxSafeTextBytes = 255
maxSafeAssetNameLen = 96
maxDigestBytes = 96
maxETagBytes = 512
safePartsCapacity = 6
)
var (
errInvalidRequest = errors.New("GitHub Release 请求参数无效")
errMetadata = errors.New("GitHub Release 元数据响应无效")
errAssetMissing = errors.New("GitHub Release 中未找到指定的已上传 asset")
errDownload = errors.New("GitHub Release asset 下载失败")
errTooLarge = errors.New("GitHub Release asset 超过大小限制")
errEmptyAsset = errors.New("GitHub Release asset 内容为空")
errDigest = errors.New("GitHub Release asset digest 无效或校验失败")
errCleanup = errors.New("GitHub Release 临时文件清理失败")
ownerPattern = regexp.MustCompile(`^[A-Za-z0-9](?:[A-Za-z0-9-]{0,37}[A-Za-z0-9])?$`)
repoPattern = regexp.MustCompile(`^[A-Za-z0-9._-]+$`)
hexPattern = regexp.MustCompile(`^[0-9a-fA-F]{64}$`)
)
var (
// ErrInvalidRequest identifies caller configuration errors.
ErrInvalidRequest = errInvalidRequest
// ErrMetadata identifies malformed, unavailable or failed Release metadata requests.
ErrMetadata = errMetadata
// ErrAssetNotFound identifies an otherwise valid Release without the exact uploaded asset.
ErrAssetNotFound = errAssetMissing
// ErrDownload identifies network or HTTP failures while downloading an asset.
ErrDownload = errDownload
// ErrAssetTooLarge identifies assets that exceed the caller's hard byte limit.
ErrAssetTooLarge = errTooLarge
// ErrEmptyAsset identifies an empty downloaded asset.
ErrEmptyAsset = errEmptyAsset
// ErrDigestMismatch identifies malformed or mismatched declared SHA-256 digests.
ErrDigestMismatch = errDigest
)
// Selector identifies GitHub's own latest endpoint or one exact tag.
type Selector string
// ResolveRequest describes one public repository release asset lookup.
type ResolveRequest struct {
Repository string
Selector Selector
Tag string
AssetName string
ETag string
}
// Release contains only metadata safe and necessary for source resolution.
type Release struct {
ID string `json:"release_id"`
Tag string `json:"tag"`
Name string `json:"name,omitempty"`
Draft bool `json:"draft"`
Prerelease bool `json:"prerelease"`
PublishedAt time.Time `json:"published_at,omitempty"`
}
// Asset contains the immutable target metadata returned by a resolve call.
type Asset struct {
ID string `json:"asset_id"`
Name string `json:"asset_name"`
State string `json:"state"`
Size int64 `json:"size"`
UpdatedAt time.Time `json:"updated_at,omitempty"`
Digest string `json:"digest,omitempty"`
}
// ResolveResult is either a selected uploaded asset or a not-modified marker.
type ResolveResult struct {
NotModified bool `json:"not_modified"`
ETag string `json:"etag,omitempty"`
Release Release `json:"release,omitempty"`
Asset Asset `json:"asset,omitempty"`
RetryAt *time.Time `json:"retry_at,omitempty"`
}
// DownloadRequest identifies an already resolved asset. Asset IDs never come
// from an untrusted URL and the download endpoint is built locally.
type DownloadRequest struct {
Repository string
Asset Asset
MaxBytes int64
}
// DownloadResult owns a temporary file. Call Cleanup after ingestion.
type DownloadResult struct {
Path string
Size int64
SHA256 string
DeclaredDigest string
}
// Cleanup removes the temporary file and is safe to call more than once.
func (result *DownloadResult) Cleanup() error {
if result == nil || result.Path == "" {
return nil
}
name := result.Path
err := os.Remove(name)
if err == nil || errors.Is(err, os.ErrNotExist) {
result.Path = ""
return nil
}
return errCleanup
}
// Error is a safe provider error. It never retains a response body, request
// URL, redirect location or request headers.
type Error struct {
Kind error
StatusCode int
RequestID string
Repository string
Tag string
AssetName string
AvailableAssets []string
RetryAt *time.Time
}
func (providerError *Error) Error() string {
if providerError == nil {
return "GitHub Release 请求失败"
}
message := "GitHub Release 请求失败"
if providerError.Kind != nil {
message = providerError.Kind.Error()
}
parts := make([]string, 0, safePartsCapacity)
if providerError.StatusCode != 0 {
parts = append(parts, "status="+strconv.Itoa(providerError.StatusCode))
}
if providerError.RequestID != "" {
parts = append(parts, "request_id="+providerError.RequestID)
}
if providerError.Repository != "" {
parts = append(parts, "repo="+providerError.Repository)
}
if providerError.Tag != "" {
parts = append(parts, "tag="+providerError.Tag)
}
if providerError.AssetName != "" {
parts = append(parts, "asset="+providerError.AssetName)
}
if len(providerError.AvailableAssets) > 0 {
parts = append(parts, "available="+strings.Join(providerError.AvailableAssets, ","))
}
if len(parts) == 0 {
return message
}
return message + " (" + strings.Join(parts, " ") + ")"
}
func (providerError *Error) Unwrap() error {
if providerError == nil {
return nil
}
return providerError.Kind
}
// RetryAt extracts the server-directed retry deadline from an error.
func RetryAt(err error) (time.Time, bool) {
var providerError *Error
if !errors.As(err, &providerError) || providerError.RetryAt == nil {
return time.Time{}, false
}
return *providerError.RetryAt, true
}
// RetryTime is retained as a compatibility alias for early callers.
//
// Deprecated: use RetryAt.
func RetryTime(err error) (time.Time, bool) {
return RetryAt(err)
}
// IsNotFound reports both a missing Release endpoint and a Release that lacks
// the exact uploaded asset requested by the caller.
func IsNotFound(err error) bool {
if errors.Is(err, ErrAssetNotFound) {
return true
}
var providerError *Error
return errors.As(err, &providerError) && providerError.StatusCode == http.StatusNotFound
}
// IsDigestError reports malformed or mismatched declared asset digests.
func IsDigestError(err error) bool {
return errors.Is(err, ErrDigestMismatch)
}
// IsRetryable classifies provider failures without relying on localized error
// strings. Configuration, not-found, size, empty-content and digest failures
// are permanent. Network failures, 408/425/429 and 5xx responses are retryable.
func IsRetryable(err error) bool {
if err == nil || errors.Is(err, ErrInvalidRequest) || IsNotFound(err) ||
errors.Is(err, ErrAssetTooLarge) || errors.Is(err, ErrEmptyAsset) || IsDigestError(err) {
return false
}
var providerError *Error
if !errors.As(err, &providerError) {
return false
}
if providerError.StatusCode == 0 {
return errors.Is(err, ErrMetadata) || errors.Is(err, ErrDownload) || errors.Is(err, errCleanup)
}
if providerError.StatusCode < http.StatusBadRequest {
return errors.Is(err, ErrMetadata) || errors.Is(err, ErrDownload) || errors.Is(err, errCleanup)
}
if providerError.RetryAt != nil {
return true
}
return providerError.StatusCode == http.StatusRequestTimeout ||
providerError.StatusCode == http.StatusTooEarly ||
providerError.StatusCode == http.StatusTooManyRequests ||
providerError.StatusCode >= http.StatusInternalServerError
}
// Client accesses public GitHub Releases using a fixed, constrained transport.
type Client struct {
httpClient *http.Client
baseURL string
createTemp func(string, string) (*os.File, error)
now func() time.Time
}
// NewClient constructs a production client for api.github.com. Public
// repositories do not require or send a token.
func NewClient() *Client {
return newClient(defaultClientOptions())
}
// Resolve calls GitHub's latest or exact-tag endpoint and selects one exact,
// case-sensitive uploaded asset. It never falls back to source archives.
func (client *Client) Resolve(ctx context.Context, request ResolveRequest) (ResolveResult, error) {
repository, tag, endpoint, err := normalizeResolveRequest(client.baseURL, request)
if err != nil {
return ResolveResult{}, safeError(errInvalidRequest, 0, "", repository, tag, validErrorAssetName(request.AssetName), nil, nil)
}
httpRequest, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil)
if err != nil {
return ResolveResult{}, safeError(errInvalidRequest, 0, "", repository, tag, validErrorAssetName(request.AssetName), nil, nil)
}
applyMetadataHeaders(httpRequest, request.ETag)
response, err := client.httpClient.Do(httpRequest) //nolint:gosec // endpoint and every dial target are constrained
if err != nil {
return ResolveResult{}, safeError(errMetadata, 0, "", repository, tag, request.AssetName, nil, nil)
}
defer func() { _ = response.Body.Close() }()
retryAt := responseRetryAt(response, client.now())
etag := safeETag(response.Header.Get("ETag"))
if response.StatusCode == http.StatusNotModified {
if etag == "" {
etag = safeETag(request.ETag)
}
return ResolveResult{NotModified: true, ETag: etag, RetryAt: retryAt}, nil
}
if response.StatusCode != http.StatusOK {
return ResolveResult{}, safeHTTPError(errMetadata, response, repository, tag, request.AssetName, retryAt)
}
body, readErr := io.ReadAll(io.LimitReader(response.Body, maxMetadataBytes+1))
if readErr != nil || len(body) > maxMetadataBytes || !utf8.Valid(body) {
return ResolveResult{}, safeHTTPError(errMetadata, response, repository, tag, request.AssetName, retryAt)
}
var payload releasePayload
decoder := json.NewDecoder(bytes.NewReader(body))
decoder.UseNumber()
if err := decoder.Decode(&payload); err != nil {
return ResolveResult{}, safeHTTPError(errMetadata, response, repository, tag, request.AssetName, retryAt)
}
if err := ensureJSONEOF(decoder); err != nil {
return ResolveResult{}, safeHTTPError(errMetadata, response, repository, tag, request.AssetName, retryAt)
}
release, assets, err := convertRelease(payload)
if err != nil {
return ResolveResult{}, safeHTTPError(errMetadata, response, repository, tag, request.AssetName, retryAt)
}
for _, asset := range assets {
if asset.State == "uploaded" && asset.Name == request.AssetName {
return ResolveResult{
ETag: etag,
Release: release,
Asset: asset,
RetryAt: retryAt,
}, nil
}
}
available := safeAssetNames(assets)
return ResolveResult{}, safeError(
errAssetMissing,
response.StatusCode,
response.Header.Get("X-GitHub-Request-Id"),
repository,
release.Tag,
request.AssetName,
available,
retryAt,
)
}
// Download streams an asset into a package-owned temporary file while
// enforcing a hard byte limit and verifying GitHub's declared sha256 digest.
func (client *Client) Download(ctx context.Context, request DownloadRequest) (*DownloadResult, error) {
repository, err := normalizeRepository(request.Repository)
if err != nil || request.MaxBytes <= 0 || !validPositiveID(request.Asset.ID) ||
!validAssetName(request.Asset.Name) || request.Asset.Size < 0 {
return nil, safeError(errInvalidRequest, 0, "", repository, "", validErrorAssetName(request.Asset.Name), nil, nil)
}
if request.Asset.Size > request.MaxBytes {
return nil, safeError(errTooLarge, 0, "", repository, "", validErrorAssetName(request.Asset.Name), nil, nil)
}
endpoint := strings.TrimRight(client.baseURL, "/") + "/repos/" + repository + "/releases/assets/" + request.Asset.ID
httpRequest, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil)
if err != nil {
return nil, safeError(errInvalidRequest, 0, "", repository, "", request.Asset.Name, nil, nil)
}
applyAssetHeaders(httpRequest)
response, err := client.httpClient.Do(httpRequest) //nolint:gosec // endpoint and every dial target are constrained
if err != nil {
return nil, safeError(errDownload, 0, "", repository, "", request.Asset.Name, nil, nil)
}
defer func() { _ = response.Body.Close() }()
retryAt := responseRetryAt(response, client.now())
if response.StatusCode != http.StatusOK {
return nil, safeHTTPError(errDownload, response, repository, "", request.Asset.Name, retryAt)
}
if response.ContentLength > request.MaxBytes {
return nil, safeHTTPError(errTooLarge, response, repository, "", request.Asset.Name, retryAt)
}
result, err := client.streamAsset(response.Body, request.MaxBytes, request.Asset.Digest)
if err != nil {
return nil, safeError(err, response.StatusCode, response.Header.Get("X-GitHub-Request-Id"), repository, "", request.Asset.Name, nil, retryAt)
}
return result, nil
}
func (client *Client) streamAsset(body io.Reader, maxBytes int64, declaredDigest string) (result *DownloadResult, resultErr error) {
tempFile, err := client.createTemp("", "openflare-github-release-*")
if err != nil {
return nil, errDownload
}
tempPath := tempFile.Name()
defer func() {
closeErr := tempFile.Close()
if resultErr == nil && closeErr != nil {
resultErr = errDownload
}
if resultErr != nil {
if removeErr := os.Remove(tempPath); removeErr != nil && !errors.Is(removeErr, os.ErrNotExist) {
resultErr = errCleanup
}
}
}()
hasher := sha256.New()
readLimit := maxBytes
if readLimit < math.MaxInt64 {
readLimit++
}
size, err := io.Copy(io.MultiWriter(tempFile, hasher), io.LimitReader(body, readLimit))
if err != nil {
return nil, errDownload
}
if size > maxBytes {
return nil, errTooLarge
}
if size == 0 {
return nil, errEmptyAsset
}
checksum := hex.EncodeToString(hasher.Sum(nil))
if err := verifyDeclaredDigest(declaredDigest, checksum); err != nil {
return nil, err
}
return &DownloadResult{
Path: tempPath,
Size: size,
SHA256: checksum,
DeclaredDigest: strings.ToLower(strings.TrimSpace(declaredDigest)),
}, nil
}
type releasePayload struct {
ID json.Number `json:"id"`
Tag string `json:"tag_name"`
Name string `json:"name"`
Draft bool `json:"draft"`
Prerelease bool `json:"prerelease"`
PublishedAt string `json:"published_at"`
Assets []assetPayload `json:"assets"`
}
type assetPayload struct {
ID json.Number `json:"id"`
Name string `json:"name"`
State string `json:"state"`
Size int64 `json:"size"`
UpdatedAt string `json:"updated_at"`
Digest string `json:"digest"`
}
func convertRelease(payload releasePayload) (Release, []Asset, error) {
releaseID, err := positiveJSONID(payload.ID)
if err != nil {
return Release{}, nil, err
}
if !validReleaseDisplayTag(payload.Tag) {
return Release{}, nil, errMetadata
}
publishedAt, err := parseOptionalTime(payload.PublishedAt)
if err != nil {
return Release{}, nil, err
}
release := Release{
ID: releaseID,
Tag: payload.Tag,
Name: safeText(payload.Name, maxSafeTextBytes),
Draft: payload.Draft,
Prerelease: payload.Prerelease,
PublishedAt: publishedAt,
}
assets := make([]Asset, 0, len(payload.Assets))
for _, rawAsset := range payload.Assets {
assetID, assetErr := positiveJSONID(rawAsset.ID)
if assetErr != nil || rawAsset.Size < 0 {
return Release{}, nil, errMetadata
}
updatedAt, assetErr := parseOptionalTime(rawAsset.UpdatedAt)
if assetErr != nil {
return Release{}, nil, errMetadata
}
assets = append(assets, Asset{
ID: assetID,
Name: rawAsset.Name,
State: rawAsset.State,
Size: rawAsset.Size,
UpdatedAt: updatedAt,
Digest: safeText(rawAsset.Digest, maxDigestBytes),
})
}
return release, assets, nil
}
func normalizeResolveRequest(baseURL string, request ResolveRequest) (string, string, string, error) {
repository, err := normalizeRepository(request.Repository)
if err != nil || !validAssetName(request.AssetName) {
return repository, validErrorTag(request.Tag), "", errInvalidRequest
}
baseURL = strings.TrimRight(baseURL, "/")
switch request.Selector {
case SelectorLatest:
if strings.TrimSpace(request.Tag) != "" {
return repository, "", "", errInvalidRequest
}
return repository, "latest", baseURL + "/repos/" + repository + "/releases/latest", nil
case SelectorTag:
if !validTag(request.Tag) {
return repository, validErrorTag(request.Tag), "", errInvalidRequest
}
return repository, request.Tag, baseURL + "/repos/" + repository + "/releases/tags/" + url.PathEscape(request.Tag), nil
default:
return repository, validErrorTag(request.Tag), "", errInvalidRequest
}
}
func normalizeRepository(repository string) (string, error) {
repository = strings.TrimSpace(repository)
parts := strings.Split(repository, "/")
if len(parts) != 2 || !ownerPattern.MatchString(parts[0]) || !repoPattern.MatchString(parts[1]) ||
len(parts[1]) > 100 || parts[1] == "." || parts[1] == ".." {
return "", errInvalidRequest
}
return parts[0] + "/" + parts[1], nil
}
func validAssetName(assetName string) bool {
return validLogText(assetName, maxSafeTextBytes, false) && path.Base(assetName) == assetName &&
assetName != "." && assetName != ".." && !strings.ContainsAny(assetName, `/\`)
}
func validTag(tag string) bool {
if !validLogText(tag, maxSafeTextBytes, false) || strings.ContainsAny(tag, " ~^:?*[\\") ||
strings.Contains(tag, "..") || strings.Contains(tag, "@{") || strings.Contains(tag, "//") ||
strings.HasPrefix(tag, "/") || strings.HasSuffix(tag, "/") || strings.HasSuffix(tag, ".") {
return false
}
for _, component := range strings.Split(tag, "/") {
if component == "" || strings.HasPrefix(component, ".") || strings.HasSuffix(component, ".lock") {
return false
}
}
return true
}
func validReleaseDisplayTag(tag string) bool {
return validLogText(tag, maxSafeTextBytes, false)
}
func validLogText(value string, maxBytes int, allowEmpty bool) bool {
if (!allowEmpty && value == "") || len(value) > maxBytes || !utf8.ValidString(value) {
return false
}
for _, character := range value {
if isLogControl(character) {
return false
}
}
return true
}
func isLogControl(character rune) bool {
if unicode.IsControl(character) || character == '\u2028' || character == '\u2029' {
return true
}
switch character {
case '\u061c', '\u200e', '\u200f',
'\u202a', '\u202b', '\u202c', '\u202d', '\u202e',
'\u2066', '\u2067', '\u2068', '\u2069':
return true
default:
return false
}
}
func validErrorTag(tag string) string {
if !validTag(tag) || containsSecretDelimiter(tag) {
return ""
}
return tag
}
func validErrorAssetName(assetName string) string {
if !validAssetName(assetName) || containsSecretDelimiter(assetName) {
return ""
}
return assetName
}
func containsSecretDelimiter(value string) bool {
return strings.ContainsAny(value, "?&=#") || strings.Contains(value, "://")
}
func validPositiveID(id string) bool {
parsed, err := strconv.ParseInt(id, 10, 64)
return err == nil && parsed > 0 && strconv.FormatInt(parsed, 10) == id
}
func positiveJSONID(id json.Number) (string, error) {
parsed, err := strconv.ParseInt(id.String(), 10, 64)
if err != nil || parsed <= 0 {
return "", errMetadata
}
return strconv.FormatInt(parsed, 10), nil
}
func parseOptionalTime(value string) (time.Time, error) {
if value == "" {
return time.Time{}, nil
}
parsed, err := time.Parse(time.RFC3339, value)
if err != nil {
return time.Time{}, errMetadata
}
return parsed, nil
}
func ensureJSONEOF(decoder *json.Decoder) error {
var trailing any
if err := decoder.Decode(&trailing); errors.Is(err, io.EOF) {
return nil
}
return errMetadata
}
func verifyDeclaredDigest(declaredDigest string, checksum string) error {
declaredDigest = strings.TrimSpace(declaredDigest)
if declaredDigest == "" {
return nil
}
algorithm, digest, ok := strings.Cut(declaredDigest, ":")
if !ok || !strings.EqualFold(algorithm, "sha256") || !hexPattern.MatchString(digest) ||
!strings.EqualFold(digest, checksum) {
return errDigest
}
return nil
}
func safeAssetNames(assets []Asset) []string {
count := len(assets)
if count > maxAssetErrorNames {
count = maxAssetErrorNames
}
names := make([]string, 0, count)
for _, asset := range assets[:count] {
name := safeText(asset.Name, maxSafeAssetNameLen)
if containsSecretDelimiter(name) {
name = "<redacted>"
}
names = append(names, name)
}
return names
}
func safeText(value string, maxBytes int) string {
var builder strings.Builder
for _, character := range value {
if isLogControl(character) {
builder.WriteByte('?')
continue
}
builder.WriteRune(character)
if builder.Len() >= maxBytes {
break
}
}
result := builder.String()
for len(result) > maxBytes {
_, size := utf8.DecodeLastRuneInString(result)
result = result[:len(result)-size]
}
return result
}
func safeETag(value string) string {
value = strings.TrimSpace(value)
if len(value) > maxETagBytes || safeText(value, maxETagBytes) != value {
return ""
}
return value
}
func safeHTTPError(kind error, response *http.Response, repository string, tag string, assetName string, retryAt *time.Time) error {
return safeError(
kind,
response.StatusCode,
response.Header.Get("X-GitHub-Request-Id"),
repository,
tag,
assetName,
nil,
retryAt,
)
}
func safeError(
kind error,
statusCode int,
requestID string,
repository string,
tag string,
assetName string,
availableAssets []string,
retryAt *time.Time,
) error {
return &Error{
Kind: kind,
StatusCode: statusCode,
RequestID: safeErrorToken(requestID, maxSafeTextBytes),
Repository: safeErrorToken(repository, maxSafeTextBytes),
Tag: safeErrorToken(tag, maxSafeTextBytes),
AssetName: safeErrorToken(assetName, maxSafeAssetNameLen),
AvailableAssets: availableAssets,
RetryAt: retryAt,
}
}
func safeErrorToken(value string, maxBytes int) string {
if !validLogText(value, maxBytes, true) {
return ""
}
value = safeText(value, maxBytes)
if containsSecretDelimiter(value) {
return ""
}
return value
}
@@ -0,0 +1,748 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package githubrelease
import (
"context"
"crypto/sha256"
"crypto/tls"
"encoding/hex"
"errors"
"fmt"
"net"
"net/http"
"net/http/httptest"
"net/netip"
"net/url"
"os"
"path/filepath"
"strconv"
"strings"
"sync"
"sync/atomic"
"testing"
"time"
)
type resolverFunc func(context.Context, string, string) ([]netip.Addr, error)
func (resolve resolverFunc) LookupNetIP(ctx context.Context, network string, host string) ([]netip.Addr, error) {
return resolve(ctx, network, host)
}
func TestResolveLatestUsesGitHubContractAndSelectsUploadedAsset(t *testing.T) {
t.Parallel()
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
if request.URL.Path != "/repos/acme/site/releases/latest" {
t.Errorf("path = %q", request.URL.Path)
}
assertHeader(t, request, "Accept", metadataAccept)
assertHeader(t, request, "User-Agent", defaultUserAgent)
assertHeader(t, request, "X-GitHub-Api-Version", APIVersion)
assertHeader(t, request, "If-None-Match", `W/"old"`)
writer.Header().Set("ETag", `W/"new"`)
writer.Header().Set("X-RateLimit-Remaining", "0")
writer.Header().Set("X-RateLimit-Reset", "1800000000")
_, _ = writer.Write([]byte(`{
"id": 9007199254740991,
"tag_name": "v1.2.3",
"name": "Stable",
"published_at": "2026-07-18T12:00:00Z",
"assets": [
{"id": 11, "name": "dist.zip", "state": "new", "size": 1},
{"id": 9007199254740990, "name": "dist.zip", "state": "uploaded", "size": 42,
"updated_at": "2026-07-18T12:10:00Z", "digest": "sha256:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa"}
]
}`))
}))
defer server.Close()
client := newTestClient(t, server.URL, nil)
result, err := client.Resolve(context.Background(), ResolveRequest{
Repository: "acme/site",
Selector: SelectorLatest,
AssetName: "dist.zip",
ETag: `W/"old"`,
})
if err != nil {
t.Fatalf("Resolve() error = %v", err)
}
if result.Release.ID != "9007199254740991" || result.Asset.ID != "9007199254740990" {
t.Fatalf("IDs lost precision: release=%q asset=%q", result.Release.ID, result.Asset.ID)
}
if result.ETag != `W/"new"` || result.Asset.Name != "dist.zip" || result.Asset.State != "uploaded" {
t.Fatalf("Resolve() = %+v", result)
}
if result.RetryAt == nil || result.RetryAt.Unix() != 1800000000 {
t.Fatalf("RetryAt = %v", result.RetryAt)
}
}
func TestResolveTagEscapesPathAndHandlesNotModified(t *testing.T) {
t.Parallel()
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
if request.RequestURI != "/repos/acme/site/releases/tags/release%2Fcandidate" {
t.Errorf("RequestURI = %q", request.RequestURI)
}
writer.WriteHeader(http.StatusNotModified)
}))
defer server.Close()
client := newTestClient(t, server.URL, nil)
result, err := client.Resolve(context.Background(), ResolveRequest{
Repository: "acme/site",
Selector: SelectorTag,
Tag: "release/candidate",
AssetName: "dist.zip",
ETag: `"cached"`,
})
if err != nil {
t.Fatalf("Resolve() error = %v", err)
}
if !result.NotModified || result.ETag != `"cached"` {
t.Fatalf("Resolve() = %+v", result)
}
}
func TestResolveAssetMissingTruncatesSafeNamesAndNeverIncludesBody(t *testing.T) {
t.Parallel()
assets := make([]string, 0, 12)
for index := 0; index < 12; index++ {
assets = append(assets, fmt.Sprintf(`{"id":%d,"name":"asset-%02d.zip","state":"uploaded","size":1}`, index+1, index))
}
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
_, _ = writer.Write([]byte(`{"id":1,"tag_name":"v1","message":"body-token","assets":[` + strings.Join(assets, ",") + `]}`))
}))
defer server.Close()
client := newTestClient(t, server.URL, nil)
_, err := client.Resolve(context.Background(), ResolveRequest{
Repository: "acme/site", Selector: SelectorLatest, AssetName: "dist.zip",
})
if !errors.Is(err, errAssetMissing) {
t.Fatalf("Resolve() error = %v", err)
}
message := err.Error()
if !strings.Contains(message, "asset-00.zip") || !strings.Contains(message, "asset-09.zip") {
t.Fatalf("error misses safe truncated names: %s", message)
}
if strings.Contains(message, "asset-10.zip") || strings.Contains(message, "asset-11.zip") || strings.Contains(message, "body-token") {
t.Fatalf("error leaked/truncation failed: %s", message)
}
}
func TestResolveHTTPErrorParsesRateLimitWithoutBodyLeak(t *testing.T) {
t.Parallel()
now := time.Date(2026, time.July, 19, 10, 0, 0, 0, time.UTC)
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
writer.Header().Set("Retry-After", "90")
writer.Header().Set("X-GitHub-Request-Id", "request-123")
writer.WriteHeader(http.StatusTooManyRequests)
_, _ = writer.Write([]byte(`{"message":"signed_url=https://secret.example/a?token=hidden"}`))
}))
defer server.Close()
client := newTestClient(t, server.URL, func(options *clientOptions) { options.now = func() time.Time { return now } })
_, err := client.Resolve(context.Background(), ResolveRequest{
Repository: "acme/site", Selector: SelectorLatest, AssetName: "dist.zip",
})
if err == nil || !strings.Contains(err.Error(), "status=429") || !strings.Contains(err.Error(), "request_id=request-123") {
t.Fatalf("Resolve() error = %v", err)
}
if strings.Contains(err.Error(), "secret.example") || strings.Contains(err.Error(), "hidden") {
t.Fatalf("error leaked body: %s", err)
}
retryAt, ok := RetryTime(err)
if !ok || !retryAt.Equal(now.Add(90*time.Second)) {
t.Fatalf("RetryTime() = %v, %v", retryAt, ok)
}
}
func TestDownloadStreamsVerifiesDigestAndCleansUp(t *testing.T) {
t.Parallel()
payload := []byte("package bytes")
digest := sha256.Sum256(payload)
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
if request.URL.Path != "/repos/acme/site/releases/assets/42" {
t.Errorf("path = %q", request.URL.Path)
}
assertHeader(t, request, "Accept", assetAccept)
assertHeader(t, request, "Accept-Encoding", "identity")
_, _ = writer.Write(payload)
}))
defer server.Close()
client := newTestClient(t, server.URL, nil)
result, err := client.Download(context.Background(), DownloadRequest{
Repository: "acme/site",
Asset: Asset{
ID: "42", Name: "dist.zip", Digest: "sha256:" + hex.EncodeToString(digest[:]),
},
MaxBytes: 1024,
})
if err != nil {
t.Fatalf("Download() error = %v", err)
}
if result.Size != int64(len(payload)) || result.SHA256 != hex.EncodeToString(digest[:]) {
t.Fatalf("Download() = %+v", result)
}
if _, err := os.Stat(result.Path); err != nil {
t.Fatalf("temp file stat: %v", err)
}
if err := result.Cleanup(); err != nil {
t.Fatalf("Cleanup() error = %v", err)
}
if err := result.Cleanup(); err != nil {
t.Fatalf("second Cleanup() error = %v", err)
}
}
func TestDownloadFollows302AndStripsCrossHostSensitiveHeaders(t *testing.T) {
t.Parallel()
var targetHost string
target := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
for _, header := range []string{"Authorization", "Cookie", "Proxy-Authorization", "Referer", "If-None-Match", "If-Modified-Since", "X-GitHub-Api-Version"} {
if value := request.Header.Get(header); value != "" {
t.Errorf("redirect leaked %s=%q", header, value)
}
}
_, _ = writer.Write([]byte("redirected package"))
}))
defer target.Close()
targetURL, _ := url.Parse(target.URL)
targetHost = "asset.example.test:" + targetURL.Port()
api := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
writer.Header().Set("Location", "http://"+targetHost+"/signed/package.zip?token=must-not-leak")
writer.WriteHeader(http.StatusFound)
}))
defer api.Close()
apiURL, _ := url.Parse(api.URL)
baseURL := "http://api.example.test:" + apiURL.Port()
client := newMappedTestClient(t, baseURL, nil)
result, err := client.Download(context.Background(), DownloadRequest{
Repository: "acme/site", Asset: Asset{ID: "42", Name: "dist.zip"}, MaxBytes: 1024,
})
if err != nil {
t.Fatalf("Download() redirect error = %v", err)
}
if cleanupErr := result.Cleanup(); cleanupErr != nil {
t.Fatalf("Cleanup() error = %v", cleanupErr)
}
request, _ := http.NewRequestWithContext(context.Background(), http.MethodGet, baseURL+"/repos/acme/site/releases/assets/42", nil)
applyAssetHeaders(request)
request.Header.Set("Authorization", "Bearer secret")
request.Header.Set("Cookie", "session=secret")
request.Header.Set("Proxy-Authorization", "proxy-secret")
request.Header.Set("Referer", "https://secret.example/path?token=x")
request.Header.Set("If-None-Match", `"secret-etag"`)
request.Header.Set("If-Modified-Since", time.Now().Format(http.TimeFormat))
response, err := client.httpClient.Do(request)
if err != nil {
t.Fatalf("Do() error = %v", err)
}
_ = response.Body.Close()
}
func TestRedirectSSRFAndDNSRebindingAreRejectedWithoutURLLeak(t *testing.T) {
t.Parallel()
t.Run("literal private redirect", func(t *testing.T) {
api := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
writer.Header().Set("Location", "http://127.0.0.1/private?token=secret-query")
writer.WriteHeader(http.StatusFound)
}))
defer api.Close()
client := newTestClient(t, api.URL, nil)
_, err := client.Download(context.Background(), DownloadRequest{
Repository: "acme/site", Asset: Asset{ID: "1", Name: "dist.zip"}, MaxBytes: 100,
})
if err == nil || strings.Contains(err.Error(), "secret-query") || strings.Contains(err.Error(), "127.0.0.1") {
t.Fatalf("Download() error = %v", err)
}
})
t.Run("DNS rebind between redirect and dial", func(t *testing.T) {
api := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
writer.Header().Set("Location", "http://rebind.example.test/package.zip")
writer.WriteHeader(http.StatusFound)
}))
defer api.Close()
var lock sync.Mutex
calls := map[string]int{}
resolve := resolverFunc(func(_ context.Context, _ string, host string) ([]netip.Addr, error) {
lock.Lock()
defer lock.Unlock()
calls[host]++
if host == "rebind.example.test" && calls[host] > 1 {
return []netip.Addr{netip.MustParseAddr("127.0.0.1")}, nil
}
return []netip.Addr{netip.MustParseAddr("8.8.8.8")}, nil
})
client := newTestClient(t, api.URL, func(options *clientOptions) { options.resolver = resolve })
_, err := client.Download(context.Background(), DownloadRequest{
Repository: "acme/site", Asset: Asset{ID: "1", Name: "dist.zip"}, MaxBytes: 100,
})
if err == nil {
t.Fatal("Download() error = nil")
}
lock.Lock()
defer lock.Unlock()
if calls["rebind.example.test"] != 2 {
t.Fatalf("rebind lookup calls = %d", calls["rebind.example.test"])
}
})
}
func TestDownloadFailureRemovesTemporaryFile(t *testing.T) {
t.Parallel()
payload := []byte("package bytes")
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
_, _ = writer.Write(payload)
}))
defer server.Close()
tempDir := t.TempDir()
client := newTestClient(t, server.URL, func(options *clientOptions) {
options.createTemp = func(_ string, pattern string) (*os.File, error) {
return os.CreateTemp(tempDir, pattern)
}
})
_, err := client.Download(context.Background(), DownloadRequest{
Repository: "acme/site",
Asset: Asset{
ID: "42", Name: "dist.zip", Digest: "sha256:" + strings.Repeat("0", 64),
},
MaxBytes: 1024,
})
if !errors.Is(err, errDigest) {
t.Fatalf("Download() error = %v", err)
}
files, readErr := filepath.Glob(filepath.Join(tempDir, "*"))
if readErr != nil || len(files) != 0 {
t.Fatalf("temporary files after failure = %v, err=%v", files, readErr)
}
}
func TestResolveRejectsInvalidRepositoryAndAssetWithoutRequest(t *testing.T) {
t.Parallel()
client := NewClient()
for _, request := range []ResolveRequest{
{Repository: "https://github.com/acme/site", Selector: SelectorLatest, AssetName: "dist.zip"},
{Repository: "acme/site/extra", Selector: SelectorLatest, AssetName: "dist.zip"},
{Repository: "acme/site", Selector: SelectorLatest, AssetName: "../dist.zip"},
{Repository: "acme/site", Selector: SelectorLatest, AssetName: `dir\dist.zip`},
{Repository: "acme/site", Selector: SelectorLatest, AssetName: string([]byte{'d', 'i', 's', 't', 0xff})},
{Repository: "acme/site", Selector: SelectorLatest, AssetName: "dist\nsecret.zip"},
{Repository: "acme/site", Selector: SelectorLatest, AssetName: "dist\u2028secret.zip"},
{Repository: "acme/site", Selector: SelectorLatest, AssetName: "dist\u202esecret.zip"},
{Repository: "acme/site", Selector: SelectorTag, AssetName: "dist.zip"},
} {
_, err := client.Resolve(context.Background(), request)
if !errors.Is(err, errInvalidRequest) {
t.Errorf("Resolve(%+v) error = %v", request, err)
}
}
}
func TestResolveAndDownloadAssetNameWithDelimiters(t *testing.T) {
t.Parallel()
assetName := "dist?channel=stable#1&x.zip"
payload := []byte("package with delimiter name")
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
switch request.URL.Path {
case "/repos/acme/site/releases/latest":
if request.Header.Get("If-None-Match") == "missing" {
_, _ = writer.Write([]byte(`{"id":1,"tag_name":"v1","assets":[]}`))
return
}
_, _ = fmt.Fprintf(writer, `{"id":1,"tag_name":"release/v1","assets":[{"id":42,"name":%q,"state":"uploaded","size":%d}]}`,
assetName, len(payload))
case "/repos/acme/site/releases/assets/42":
_, _ = writer.Write(payload)
default:
http.NotFound(writer, request)
}
}))
defer server.Close()
client := newTestClient(t, server.URL, nil)
resolved, err := client.Resolve(context.Background(), ResolveRequest{
Repository: "acme/site", Selector: SelectorLatest, AssetName: assetName,
})
if err != nil {
t.Fatalf("Resolve() error = %v", err)
}
if resolved.Asset.Name != assetName || resolved.Release.Tag != "release/v1" {
t.Fatalf("Resolve() = %+v", resolved)
}
download, err := client.Download(context.Background(), DownloadRequest{
Repository: "acme/site", Asset: resolved.Asset, MaxBytes: 1024,
})
if err != nil {
t.Fatalf("Download() error = %v", err)
}
if cleanupErr := download.Cleanup(); cleanupErr != nil {
t.Fatalf("Cleanup() error = %v", cleanupErr)
}
_, err = client.Resolve(context.Background(), ResolveRequest{
Repository: "acme/site", Selector: SelectorLatest, AssetName: assetName, ETag: "missing",
})
if !errors.Is(err, ErrAssetNotFound) {
t.Fatalf("missing Resolve() error = %v", err)
}
if strings.Contains(err.Error(), assetName) || strings.Contains(err.Error(), "channel=stable") {
t.Fatalf("missing error leaked delimiter-bearing name: %v", err)
}
}
func TestFixedTagGitRefRulesAndEscaping(t *testing.T) {
t.Parallel()
valid := []string{"@", "release/v1#stable&channel=prod", "foo.LOCK", "中文/发布=稳定"}
for _, tag := range valid {
if !validTag(tag) {
t.Errorf("validTag(%q) = false", tag)
}
}
invalid := []string{
"", "release v1", "release~v1", "release^v1", "release:v1", "release?v1", "release*v1",
"release[v1", `release\v1`, "release..v1", "release@{v1", "release//v1", "/release", "release/",
"release.", ".release", "release/.candidate", "release.lock", "release/v1.lock", "release\nsecret",
"release\u2028secret", "release\u202esecret", string([]byte{'v', '1', 0xff}), strings.Repeat("a", 256),
}
for _, tag := range invalid {
if validTag(tag) {
t.Errorf("validTag(%q) = true", tag)
}
}
tag := "release/v1#stable&channel=prod"
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
wantURI := "/repos/acme/site/releases/tags/" + url.PathEscape(tag)
if request.RequestURI != wantURI || request.URL.RawQuery != "" || request.URL.Fragment != "" {
t.Errorf("tag request = %q query=%q fragment=%q, want %q", request.RequestURI, request.URL.RawQuery, request.URL.Fragment, wantURI)
}
_, _ = fmt.Fprintf(writer, `{"id":1,"tag_name":%q,"assets":[{"id":2,"name":"dist.zip","state":"uploaded","size":1}]}`, tag)
}))
defer server.Close()
client := newTestClient(t, server.URL, nil)
result, err := client.Resolve(context.Background(), ResolveRequest{
Repository: "acme/site", Selector: SelectorTag, Tag: tag, AssetName: "dist.zip",
})
if err != nil {
t.Fatalf("Resolve() error = %v", err)
}
if result.Release.Tag != tag {
t.Fatalf("Release.Tag = %q", result.Release.Tag)
}
}
func TestResolveDoesNotMatchSanitizedRemoteAssetName(t *testing.T) {
t.Parallel()
tests := []struct {
name string
remote string
requested string
}{
{name: "unicode line separator", remote: "dist\u2028.zip", requested: "dist?.zip"},
{name: "overlong", remote: strings.Repeat("a", 256), requested: strings.Repeat("a", 255)},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
t.Parallel()
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
_, _ = fmt.Fprintf(writer, `{"id":1,"tag_name":"release/v1","assets":[{"id":2,"name":%q,"state":"uploaded","size":1}]}`, test.remote)
}))
defer server.Close()
client := newTestClient(t, server.URL, nil)
_, err := client.Resolve(context.Background(), ResolveRequest{
Repository: "acme/site", Selector: SelectorLatest, AssetName: test.requested,
})
if !errors.Is(err, ErrAssetNotFound) {
t.Fatalf("Resolve() error = %v", err)
}
})
}
}
func TestReleaseDisplayTagValidation(t *testing.T) {
t.Parallel()
for _, tag := range []string{"release/v1", "release v1", "release/v1#stable&channel=prod"} {
if !validReleaseDisplayTag(tag) {
t.Errorf("validReleaseDisplayTag(%q) = false", tag)
}
}
for _, tag := range []string{"", strings.Repeat("a", 256), "release\nsecret", "release\u2028secret", "release\u202esecret"} {
if validReleaseDisplayTag(tag) {
t.Errorf("validReleaseDisplayTag(%q) = true", tag)
}
}
}
func TestResolveRejectsInvalidUTF8Metadata(t *testing.T) {
t.Parallel()
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
_, _ = writer.Write(append([]byte(`{"id":1,"tag_name":"v1","assets":[{"id":2,"name":"dist`),
append([]byte{0xff}, []byte(`.zip","state":"uploaded","size":1}]}`)...)...))
}))
defer server.Close()
client := newTestClient(t, server.URL, nil)
_, err := client.Resolve(context.Background(), ResolveRequest{
Repository: "acme/site", Selector: SelectorLatest, AssetName: "dist�.zip",
})
if !errors.Is(err, ErrMetadata) {
t.Fatalf("Resolve() error = %v", err)
}
}
func TestDownloadRejectsImpossibleMetadataBeforeNetwork(t *testing.T) {
t.Parallel()
var requests atomic.Int64
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
requests.Add(1)
writer.WriteHeader(http.StatusInternalServerError)
}))
defer server.Close()
client := newTestClient(t, server.URL, nil)
tests := []struct {
name string
size int64
kind error
limit int64
}{
{name: "negative", size: -1, kind: ErrInvalidRequest, limit: 100},
{name: "declared too large", size: 101, kind: ErrAssetTooLarge, limit: 100},
}
for _, test := range tests {
_, err := client.Download(context.Background(), DownloadRequest{
Repository: "acme/site",
Asset: Asset{ID: "1", Name: "dist?token=hidden#asset.zip", Size: test.size},
MaxBytes: test.limit,
})
if !errors.Is(err, test.kind) {
t.Errorf("%s Download() error = %v", test.name, err)
}
if strings.Contains(err.Error(), "token=hidden") {
t.Errorf("%s error leaked asset name: %v", test.name, err)
}
}
if got := requests.Load(); got != 0 {
t.Fatalf("HTTP requests = %d, want 0", got)
}
}
func TestLogControlCharactersNeverEnterSafeErrors(t *testing.T) {
t.Parallel()
controls := []string{"\u2028", "\u2029", "\u061c", "\u200e", "\u200f", "\u202e", "\u2066", "\u2069"}
for _, control := range controls {
secret := "before" + control + "after"
err := safeError(errInvalidRequest, 0, secret, secret, secret, secret, nil, nil)
message := err.Error()
if strings.Contains(message, secret) || strings.Contains(message, control) || strings.Contains(message, "before") {
t.Errorf("safe error retained control %U: %q", []rune(control)[0], message)
}
}
}
func TestResolveRejectsMetadataOverHardLimit(t *testing.T) {
t.Parallel()
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
_, _ = writer.Write([]byte(`{"id":1,"tag_name":"v1","assets":[]}` + strings.Repeat(" ", maxMetadataBytes)))
}))
defer server.Close()
client := newTestClient(t, server.URL, nil)
_, err := client.Resolve(context.Background(), ResolveRequest{
Repository: "acme/site", Selector: SelectorLatest, AssetName: "dist.zip",
})
if !errors.Is(err, ErrMetadata) {
t.Fatalf("Resolve() error = %v", err)
}
}
func TestProductionTransportRejectsSelfSignedTLS(t *testing.T) {
t.Parallel()
server := httptest.NewTLSServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
_, _ = writer.Write([]byte("package"))
}))
defer server.Close()
parsed, _ := url.Parse(server.URL)
dialer := &net.Dialer{Timeout: time.Second}
client := newClient(clientOptions{
baseURL: "https://api.example.test:" + parsed.Port(),
resolver: resolverFunc(func(_ context.Context, _ string, _ string) ([]netip.Addr, error) {
return []netip.Addr{netip.MustParseAddr("8.8.8.8")}, nil
}),
dialContext: func(ctx context.Context, network string, _ string) (net.Conn, error) {
return dialer.DialContext(ctx, network, server.Listener.Addr().String())
},
tlsConfig: &tls.Config{MinVersion: tls.VersionTLS12},
createTemp: os.CreateTemp,
now: time.Now,
clientTimeout: 5 * time.Second,
})
_, err := client.Download(context.Background(), DownloadRequest{
Repository: "acme/site", Asset: Asset{ID: "1", Name: "dist.zip"}, MaxBytes: 100,
})
if !errors.Is(err, ErrDownload) {
t.Fatalf("Download() error = %v", err)
}
if strings.Contains(err.Error(), "api.example.test") || strings.Contains(err.Error(), server.URL) {
t.Fatalf("TLS error leaked URL: %v", err)
}
}
func TestStableErrorClassification(t *testing.T) {
t.Parallel()
now := time.Now()
assetMissing := safeError(errAssetMissing, http.StatusOK, "", "acme/site", "v1", "dist.zip", nil, nil)
if !IsNotFound(assetMissing) || IsRetryable(assetMissing) {
t.Fatalf("asset missing classification failed: %v", assetMissing)
}
metadata404 := safeError(errMetadata, http.StatusNotFound, "", "acme/site", "v1", "dist.zip", nil, nil)
if !IsNotFound(metadata404) || IsRetryable(metadata404) {
t.Fatalf("metadata 404 classification failed: %v", metadata404)
}
digest := safeError(errDigest, http.StatusOK, "", "acme/site", "", "dist.zip", nil, nil)
if !IsDigestError(digest) || IsRetryable(digest) {
t.Fatalf("digest classification failed: %v", digest)
}
for _, retryable := range []error{
safeError(errMetadata, 0, "", "acme/site", "", "dist.zip", nil, nil),
safeError(errMetadata, http.StatusOK, "", "acme/site", "", "dist.zip", nil, nil),
safeError(errDownload, http.StatusOK, "", "acme/site", "", "dist.zip", nil, nil),
safeError(errMetadata, http.StatusInternalServerError, "", "acme/site", "", "dist.zip", nil, nil),
safeError(errMetadata, http.StatusForbidden, "", "acme/site", "", "dist.zip", nil, &now),
} {
if !IsRetryable(retryable) {
t.Errorf("IsRetryable(%v) = false", retryable)
}
}
}
func TestDownloadRedirectLimitIsSafe(t *testing.T) {
t.Parallel()
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
step, _ := strconv.Atoi(request.URL.Query().Get("step"))
writer.Header().Set("Location", fmt.Sprintf("/repos/acme/site/releases/assets/1?step=%d&token=redirect-secret", step+1))
writer.WriteHeader(http.StatusFound)
}))
defer server.Close()
client := newTestClient(t, server.URL, nil)
_, err := client.Download(context.Background(), DownloadRequest{
Repository: "acme/site", Asset: Asset{ID: "1", Name: "dist.zip"}, MaxBytes: 100,
})
if err == nil {
t.Fatal("Download() error = nil")
}
if strings.Contains(err.Error(), "redirect-secret") || strings.Contains(err.Error(), "step=") {
t.Fatalf("redirect error leaked Location: %v", err)
}
}
func TestInvalidRequestDoesNotEchoURLQueryTagOrAsset(t *testing.T) {
t.Parallel()
client := NewClient()
requests := []ResolveRequest{
{Repository: "https://github.com/acme/site?token=repo-secret", Selector: SelectorLatest, AssetName: "dist.zip"},
{Repository: "acme/site", Selector: SelectorTag, Tag: "?token=tag-secret", AssetName: "dist.zip"},
{Repository: "acme/site", Selector: SelectorLatest, AssetName: "../dist.zip?token=asset-secret"},
}
for _, request := range requests {
_, err := client.Resolve(context.Background(), request)
if err == nil {
t.Fatalf("Resolve(%+v) error = nil", request)
}
for _, secret := range []string{"repo-secret", "tag-secret", "asset-secret", "https://github.com"} {
if strings.Contains(err.Error(), secret) {
t.Fatalf("Resolve(%+v) leaked %q: %v", request, secret, err)
}
}
}
}
func newTestClient(t *testing.T, rawBaseURL string, customize func(*clientOptions)) *Client {
t.Helper()
parsed, err := url.Parse(rawBaseURL)
if err != nil {
t.Fatal(err)
}
baseURL := "http://api.example.test:" + parsed.Port()
return newMappedTestClient(t, baseURL, customize)
}
func newMappedTestClient(t *testing.T, baseURL string, customize func(*clientOptions)) *Client {
t.Helper()
resolve := resolverFunc(func(_ context.Context, _ string, _ string) ([]netip.Addr, error) {
return []netip.Addr{netip.MustParseAddr("8.8.8.8")}, nil
})
dialer := &net.Dialer{Timeout: time.Second}
options := clientOptions{
baseURL: baseURL,
resolver: resolve,
allowHTTP: true,
tlsConfig: &tls.Config{MinVersion: tls.VersionTLS12},
dialContext: func(ctx context.Context, network string, address string) (net.Conn, error) {
_, port, splitErr := net.SplitHostPort(address)
if splitErr != nil {
return nil, splitErr
}
return dialer.DialContext(ctx, network, net.JoinHostPort("127.0.0.1", port))
},
createTemp: os.CreateTemp,
now: time.Now,
clientTimeout: 5 * time.Second,
}
if customize != nil {
customize(&options)
}
return newClient(options)
}
func assertHeader(t *testing.T, request *http.Request, name string, expected string) {
t.Helper()
if actual := request.Header.Get(name); actual != expected {
t.Errorf("%s = %q, want %q", name, actual, expected)
}
}
func TestResponseRetryAtHTTPDate(t *testing.T) {
t.Parallel()
want := time.Date(2026, time.July, 19, 12, 30, 0, 0, time.UTC)
response := &http.Response{Header: make(http.Header)}
response.Header.Set("Retry-After", want.Format(http.TimeFormat))
if got := responseRetryAt(response, time.Time{}); got == nil || !got.Equal(want) {
t.Fatalf("responseRetryAt() = %v", got)
}
}
func TestResponseRetryAtRejectsDurationOverflow(t *testing.T) {
t.Parallel()
response := &http.Response{Header: make(http.Header)}
response.Header.Set("Retry-After", strconv.FormatInt(maxRetryAfterSeconds+1, 10))
if got := responseRetryAt(response, time.Now()); got != nil {
t.Fatalf("responseRetryAt(overflow) = %v", got)
}
}
func TestSafeETagDropsOversizedOrControlValue(t *testing.T) {
t.Parallel()
if got := safeETag(strings.Repeat("x", 513)); got != "" {
t.Fatalf("safeETag(overlong) = %q", got)
}
if got := safeETag("ok\nsecret"); got != "" {
t.Fatalf("safeETag(control) = %q", got)
}
}
func TestDownloadSizeLimit(t *testing.T) {
t.Parallel()
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
writer.Header().Set("Content-Length", strconv.Itoa(20))
_, _ = writer.Write([]byte(strings.Repeat("x", 20)))
}))
defer server.Close()
client := newTestClient(t, server.URL, nil)
_, err := client.Download(context.Background(), DownloadRequest{
Repository: "acme/site", Asset: Asset{ID: "1", Name: "dist.zip"}, MaxBytes: 10,
})
if !errors.Is(err, errTooLarge) {
t.Fatalf("Download() error = %v", err)
}
}
@@ -0,0 +1,312 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package githubrelease
import (
"context"
"crypto/tls"
"errors"
"math"
"net"
"net/http"
"net/netip"
"net/url"
"os"
"strconv"
"strings"
"time"
"github.com/Rain-kl/Wavelet/pkg/httppool"
)
const (
clientTimeout = 10 * time.Minute
dialTimeout = 30 * time.Second
dialKeepAlive = 30 * time.Second
responseHeaderTimeout = 30 * time.Second
maxRedirects = 5
maxRetryAfterSeconds = math.MaxInt64 / int64(time.Second)
)
var (
errBlockedTarget = errors.New("GitHub Release 请求目标不是公网地址")
errResolveTarget = errors.New("GitHub Release 请求目标解析失败")
errRedirectLimit = errors.New("GitHub Release asset 重定向次数过多")
publicIPv6Prefix = netip.MustParsePrefix("2000::/3")
nonPublicPrefixes = []netip.Prefix{
netip.MustParsePrefix("0.0.0.0/8"),
netip.MustParsePrefix("10.0.0.0/8"),
netip.MustParsePrefix("100.64.0.0/10"),
netip.MustParsePrefix("127.0.0.0/8"),
netip.MustParsePrefix("169.254.0.0/16"),
netip.MustParsePrefix("172.16.0.0/12"),
netip.MustParsePrefix("192.0.0.0/24"),
netip.MustParsePrefix("192.0.2.0/24"),
netip.MustParsePrefix("192.168.0.0/16"),
netip.MustParsePrefix("198.18.0.0/15"),
netip.MustParsePrefix("198.51.100.0/24"),
netip.MustParsePrefix("203.0.113.0/24"),
netip.MustParsePrefix("224.0.0.0/4"),
netip.MustParsePrefix("240.0.0.0/4"),
netip.MustParsePrefix("::/128"),
netip.MustParsePrefix("::1/128"),
netip.MustParsePrefix("::ffff:0:0/96"),
netip.MustParsePrefix("64:ff9b::/96"),
netip.MustParsePrefix("100::/64"),
netip.MustParsePrefix("2001:db8::/32"),
netip.MustParsePrefix("fc00::/7"),
netip.MustParsePrefix("fe80::/10"),
netip.MustParsePrefix("ff00::/8"),
}
)
type resolver interface {
LookupNetIP(context.Context, string, string) ([]netip.Addr, error)
}
type clientOptions struct {
baseURL string
resolver resolver
dialContext func(context.Context, string, string) (net.Conn, error)
tlsConfig *tls.Config
allowHTTP bool
createTemp func(string, string) (*os.File, error)
now func() time.Time
clientTimeout time.Duration
}
func defaultClientOptions() clientOptions {
dialer := &net.Dialer{Timeout: dialTimeout, KeepAlive: dialKeepAlive}
return clientOptions{
baseURL: defaultAPIBaseURL,
resolver: net.DefaultResolver,
dialContext: dialer.DialContext,
tlsConfig: &tls.Config{MinVersion: tls.VersionTLS12},
createTemp: os.CreateTemp,
now: time.Now,
clientTimeout: clientTimeout,
}
}
func newClient(options clientOptions) *Client {
if options.baseURL == "" {
options.baseURL = defaultAPIBaseURL
}
if options.resolver == nil {
options.resolver = net.DefaultResolver
}
if options.dialContext == nil {
dialer := &net.Dialer{Timeout: dialTimeout, KeepAlive: dialKeepAlive}
options.dialContext = dialer.DialContext
}
if options.tlsConfig == nil {
options.tlsConfig = &tls.Config{MinVersion: tls.VersionTLS12}
}
if options.createTemp == nil {
options.createTemp = os.CreateTemp
}
if options.now == nil {
options.now = time.Now
}
if options.clientTimeout <= 0 {
options.clientTimeout = clientTimeout
}
secureDial := publicDialer(options.resolver, options.dialContext)
transport := httppool.NewTransport(httppool.TransportOptions{
Proxy: nil,
DialContext: secureDial,
TLSClientConfig: options.tlsConfig,
ResponseHeaderTimeout: responseHeaderTimeout,
TraceFilter: func(request *http.Request) bool {
return request.URL == nil || request.URL.RawQuery == ""
},
})
httpClient := &http.Client{Timeout: options.clientTimeout, Transport: transport}
httpClient.CheckRedirect = func(next *http.Request, previous []*http.Request) error {
if len(previous) > maxRedirects {
return errRedirectLimit
}
if err := validateTarget(next.Context(), next.URL, options.resolver, options.allowHTTP); err != nil {
return err
}
if len(previous) > 0 && !sameHost(previous[len(previous)-1].URL, next.URL) {
stripCrossHostHeaders(next)
}
return nil
}
return &Client{
httpClient: httpClient,
baseURL: strings.TrimRight(options.baseURL, "/"),
createTemp: options.createTemp,
now: options.now,
}
}
func applyMetadataHeaders(request *http.Request, etag string) {
request.Header.Set("Accept", metadataAccept)
request.Header.Set("User-Agent", defaultUserAgent)
request.Header.Set("X-GitHub-Api-Version", APIVersion)
if etag = safeETag(etag); etag != "" {
request.Header.Set("If-None-Match", etag)
}
}
func applyAssetHeaders(request *http.Request) {
request.Header.Set("Accept", assetAccept)
request.Header.Set("Accept-Encoding", "identity")
request.Header.Set("User-Agent", defaultUserAgent)
request.Header.Set("X-GitHub-Api-Version", APIVersion)
}
func stripCrossHostHeaders(request *http.Request) {
for _, header := range []string{
"Authorization",
"Cookie",
"Proxy-Authorization",
"Referer",
"If-None-Match",
"If-Modified-Since",
"X-GitHub-Api-Version",
} {
request.Header.Del(header)
}
}
func sameHost(left *url.URL, right *url.URL) bool {
if left == nil || right == nil {
return false
}
return strings.EqualFold(left.Hostname(), right.Hostname()) && effectivePort(left) == effectivePort(right)
}
func effectivePort(target *url.URL) string {
if port := target.Port(); port != "" {
return port
}
if strings.EqualFold(target.Scheme, "https") {
return "443"
}
return "80"
}
func validateTarget(ctx context.Context, target *url.URL, targetResolver resolver, allowHTTP bool) error {
if target == nil || target.User != nil || target.Fragment != "" || target.Opaque != "" || target.Hostname() == "" {
return errBlockedTarget
}
isHTTPS := strings.EqualFold(target.Scheme, "https")
isAllowedHTTP := allowHTTP && strings.EqualFold(target.Scheme, "http")
if !isHTTPS && !isAllowedHTTP {
return errBlockedTarget
}
_, err := resolvePublicIPs(ctx, targetResolver, target.Hostname())
return err
}
func publicDialer(
targetResolver resolver,
directDial func(context.Context, string, string) (net.Conn, error),
) func(context.Context, string, string) (net.Conn, error) {
return func(ctx context.Context, network string, address string) (net.Conn, error) {
host, port, err := net.SplitHostPort(address)
if err != nil {
return nil, errResolveTarget
}
addresses, err := resolvePublicIPs(ctx, targetResolver, host)
if err != nil {
return nil, err
}
for _, resolved := range addresses {
if !ipMatchesNetwork(resolved, network) {
continue
}
connection, dialErr := directDial(ctx, network, net.JoinHostPort(resolved.String(), port))
if dialErr == nil {
return connection, nil
}
}
return nil, errDownload
}
}
func resolvePublicIPs(ctx context.Context, targetResolver resolver, host string) ([]netip.Addr, error) {
if strings.Contains(host, "%") {
return nil, errBlockedTarget
}
if literal, err := netip.ParseAddr(host); err == nil {
if !isPublicIP(literal) {
return nil, errBlockedTarget
}
return []netip.Addr{literal}, nil
}
if targetResolver == nil {
return nil, errResolveTarget
}
addresses, err := targetResolver.LookupNetIP(ctx, "ip", host)
if err != nil || len(addresses) == 0 {
return nil, errResolveTarget
}
for _, address := range addresses {
if !isPublicIP(address) {
return nil, errBlockedTarget
}
}
return addresses, nil
}
func isPublicIP(address netip.Addr) bool {
if !address.IsValid() || address.Zone() != "" {
return false
}
address = address.Unmap()
if !address.IsGlobalUnicast() {
return false
}
if address.Is6() && !publicIPv6Prefix.Contains(address) {
return false
}
for _, prefix := range nonPublicPrefixes {
if prefix.Contains(address) {
return false
}
}
return true
}
func ipMatchesNetwork(address netip.Addr, network string) bool {
switch network {
case "tcp4":
return address.Unmap().Is4()
case "tcp6":
return address.Unmap().Is6()
default:
return true
}
}
func responseRetryAt(response *http.Response, now time.Time) *time.Time {
if response == nil {
return nil
}
if retryAfter := strings.TrimSpace(response.Header.Get("Retry-After")); retryAfter != "" {
if seconds, err := strconv.ParseInt(retryAfter, 10, 64); err == nil && seconds >= 0 && seconds <= maxRetryAfterSeconds {
retryAt := now.Add(time.Duration(seconds) * time.Second)
return &retryAt
}
if retryAt, err := http.ParseTime(retryAfter); err == nil {
retryAt = retryAt.UTC()
return &retryAt
}
}
if strings.TrimSpace(response.Header.Get("X-RateLimit-Remaining")) != "0" {
return nil
}
reset, err := strconv.ParseInt(strings.TrimSpace(response.Header.Get("X-RateLimit-Reset")), 10, 64)
if err != nil || reset <= 0 {
return nil
}
retryAt := time.Unix(reset, 0).UTC()
return &retryAt
}