fix(pages): 收紧部署包与 Agent 同步边界

完成 V2 Phase 0 安全与一致性前置:统一真实归档限额、流式拉取、候选裁剪、保留上传删除语义及 Pages 路由引用锁。
This commit is contained in:
deqiying
2026-07-19 16:42:45 +08:00
parent f386674464
commit 4e8ec23264
47 changed files with 4005 additions and 759 deletions
+4
View File
@@ -30,6 +30,10 @@ sidebar: false
- WAF 规则编辑器支持为节点自定义显示名称,并从节点库拖放到画布指定位置添加节点。
- WAF 规则画布支持右键删除节点或连线,并屏蔽浏览器默认右键菜单。
### 修复
- 修复 Pages 部署包路径校验、归档展开限额、历史版本裁剪、代理路由绑定与 Agent 下载过程中的安全和一致性问题;大包改为流式处理,部署入口、旧版目录切换、保留版本及上传记录在并发场景下更加可靠。
## [v3.4.0] - 2026-07-19
### 新增
+43 -2
View File
@@ -3738,6 +3738,12 @@ const docTemplate = `{
"schema": {
"$ref": "#/definitions/response.Any"
}
},
"409": {
"description": "系统保留类型或存储只读",
"schema": {
"$ref": "#/definitions/response.Any"
}
}
}
}
@@ -12535,6 +12541,12 @@ const docTemplate = `{
"$ref": "#/definitions/response.Any"
}
},
"409": {
"description": "系统保留类型或存储只读",
"schema": {
"$ref": "#/definitions/response.Any"
}
},
"500": {
"description": "内部错误",
"schema": {
@@ -12729,6 +12741,12 @@ const docTemplate = `{
"schema": {
"$ref": "#/definitions/response.Any"
}
},
"409": {
"description": "系统保留类型或存储只读",
"schema": {
"$ref": "#/definitions/response.Any"
}
}
}
}
@@ -14574,11 +14592,20 @@ const docTemplate = `{
"deployment_id": {
"type": "integer"
},
"file_count": {
"type": "integer"
},
"hash": {
"type": "string"
},
"package_size": {
"type": "integer"
},
"project_id": {
"type": "integer"
},
"total_size": {
"type": "integer"
}
}
},
@@ -16568,9 +16595,15 @@ const docTemplate = `{
"observability.AccessLogView": {
"type": "object",
"properties": {
"bytes_sent": {
"type": "integer"
},
"cache_status": {
"type": "string"
},
"created_at": {
"type": "string"
},
"host": {
"type": "string"
},
@@ -16595,6 +16628,12 @@ const docTemplate = `{
"remote_addr": {
"type": "string"
},
"request_length": {
"type": "integer"
},
"request_time_ms": {
"type": "integer"
},
"status_code": {
"type": "integer"
},
@@ -19310,7 +19349,8 @@ const docTemplate = `{
"block",
"ip_match",
"geo_match",
"pow"
"pow",
"ua_check"
],
"x-enum-varnames": [
"RuleNodeStart",
@@ -19318,7 +19358,8 @@ const docTemplate = `{
"RuleNodeBlock",
"RuleNodeIPMatch",
"RuleNodeGeoMatch",
"RuleNodePoW"
"RuleNodePoW",
"RuleNodeUACheck"
]
},
"waf.RulePosition": {
+43 -2
View File
@@ -3731,6 +3731,12 @@
"schema": {
"$ref": "#/definitions/response.Any"
}
},
"409": {
"description": "系统保留类型或存储只读",
"schema": {
"$ref": "#/definitions/response.Any"
}
}
}
}
@@ -12528,6 +12534,12 @@
"$ref": "#/definitions/response.Any"
}
},
"409": {
"description": "系统保留类型或存储只读",
"schema": {
"$ref": "#/definitions/response.Any"
}
},
"500": {
"description": "内部错误",
"schema": {
@@ -12722,6 +12734,12 @@
"schema": {
"$ref": "#/definitions/response.Any"
}
},
"409": {
"description": "系统保留类型或存储只读",
"schema": {
"$ref": "#/definitions/response.Any"
}
}
}
}
@@ -14567,11 +14585,20 @@
"deployment_id": {
"type": "integer"
},
"file_count": {
"type": "integer"
},
"hash": {
"type": "string"
},
"package_size": {
"type": "integer"
},
"project_id": {
"type": "integer"
},
"total_size": {
"type": "integer"
}
}
},
@@ -16561,9 +16588,15 @@
"observability.AccessLogView": {
"type": "object",
"properties": {
"bytes_sent": {
"type": "integer"
},
"cache_status": {
"type": "string"
},
"created_at": {
"type": "string"
},
"host": {
"type": "string"
},
@@ -16588,6 +16621,12 @@
"remote_addr": {
"type": "string"
},
"request_length": {
"type": "integer"
},
"request_time_ms": {
"type": "integer"
},
"status_code": {
"type": "integer"
},
@@ -19303,7 +19342,8 @@
"block",
"ip_match",
"geo_match",
"pow"
"pow",
"ua_check"
],
"x-enum-varnames": [
"RuleNodeStart",
@@ -19311,7 +19351,8 @@
"RuleNodeBlock",
"RuleNodeIPMatch",
"RuleNodeGeoMatch",
"RuleNodePoW"
"RuleNodePoW",
"RuleNodeUACheck"
]
},
"waf.RulePosition": {
+28
View File
@@ -760,10 +760,16 @@ definitions:
properties:
deployment_id:
type: integer
file_count:
type: integer
hash:
type: string
package_size:
type: integer
project_id:
type: integer
total_size:
type: integer
type: object
github_com_Rain-kl_Wavelet_pkg_protocol.WAFIPGroup:
properties:
@@ -2085,8 +2091,12 @@ definitions:
type: object
observability.AccessLogView:
properties:
bytes_sent:
type: integer
cache_status:
type: string
created_at:
type: string
host:
type: string
id:
@@ -2103,6 +2113,10 @@ definitions:
type: string
remote_addr:
type: string
request_length:
type: integer
request_time_ms:
type: integer
status_code:
type: integer
user_agent:
@@ -3908,6 +3922,7 @@ definitions:
- ip_match
- geo_match
- pow
- ua_check
type: string
x-enum-varnames:
- RuleNodeStart
@@ -3916,6 +3931,7 @@ definitions:
- RuleNodeIPMatch
- RuleNodeGeoMatch
- RuleNodePoW
- RuleNodeUACheck
waf.RulePosition:
properties:
x:
@@ -6157,6 +6173,10 @@ paths:
description: 文件不存在
schema:
$ref: '#/definitions/response.Any'
"409":
description: 系统保留类型或存储只读
schema:
$ref: '#/definitions/response.Any'
security:
- SessionCookie: []
summary: 删除文件
@@ -11589,6 +11609,10 @@ paths:
description: 未登录
schema:
$ref: '#/definitions/response.Any'
"409":
description: 系统保留类型或存储只读
schema:
$ref: '#/definitions/response.Any'
"500":
description: 内部错误
schema:
@@ -11622,6 +11646,10 @@ paths:
description: 文件不存在
schema:
$ref: '#/definitions/response.Any'
"409":
description: 系统保留类型或存储只读
schema:
$ref: '#/definitions/response.Any'
security:
- SessionCookie: []
summary: 删除我的文件
+113 -18
View File
@@ -1,9 +1,13 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package httpclient provides an authenticated HTTP client for the agent.
package httpclient
import (
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
@@ -13,6 +17,8 @@ import (
edgehttp "github.com/Rain-kl/Wavelet/internal/apps/edge/httpclient"
)
const pagesControlResponseMaxBytes = int64(64 * 1024)
// Client is a HTTP client used by the agent to communicate with the control plane server.
type Client struct {
base *edgehttp.Client
@@ -98,23 +104,43 @@ func (c *Client) GetPagesDeploymentHash(ctx context.Context, deploymentID uint)
return resp.Data.Hash, nil
}
// DownloadPagesDeploymentPackage downloads the deployment package for the given Pages deployment ID.
func (c *Client) DownloadPagesDeploymentPackage(ctx context.Context, deploymentID uint) ([]byte, error) {
res, err := c.base.DoRaw(ctx, http.MethodGet, fmt.Sprintf("/api/v1/agent/pages/deployments/%d/package", deploymentID), nil)
if err != nil {
return nil, err
}
defer func() { _ = res.Body.Close() }()
if res.StatusCode != http.StatusOK {
return nil, edgehttp.ReadHTTPError(res)
}
return io.ReadAll(res.Body)
// DownloadPagesDeploymentPackage streams the deployment package into dst while
// enforcing maxBytes against both advertised and actual response sizes.
func (c *Client) DownloadPagesDeploymentPackage(
ctx context.Context,
deploymentID uint,
dst io.Writer,
maxBytes int64,
) (int64, error) {
return c.downloadPagesPackage(
ctx,
fmt.Sprintf("/api/v1/agent/pages/deployments/%d/package", deploymentID),
dst,
maxBytes,
)
}
// GetPagesProjectLatestHash returns the active deployment package hash for a Pages project.
func (c *Client) GetPagesProjectLatestHash(ctx context.Context, projectID uint) (*protocol.PagesProjectLatestHashResponse, error) {
res, err := c.base.DoRaw(
ctx,
http.MethodGet,
fmt.Sprintf("/api/v1/agent/pages/projects/%d/latest/hash", projectID),
nil,
)
if err != nil {
return nil, err
}
defer func() { _ = res.Body.Close() }()
body, err := readPagesControlResponse(res, pagesControlResponseMaxBytes)
if err != nil {
return nil, err
}
if res.StatusCode != http.StatusOK {
return nil, edgehttp.ReadBodyError(body, res.Status)
}
resp := protocol.APIResponse[protocol.PagesProjectLatestHashResponse]{}
if err := c.base.GetJSON(ctx, fmt.Sprintf("/api/v1/agent/pages/projects/%d/latest/hash", projectID), &resp); err != nil {
if err := json.Unmarshal(body, &resp); err != nil {
return nil, err
}
if err := edgehttp.APIError(resp.ErrorMsg); err != nil {
@@ -123,17 +149,86 @@ func (c *Client) GetPagesProjectLatestHash(ctx context.Context, projectID uint)
return &resp.Data, nil
}
// DownloadPagesProjectLatestPackage downloads the active deployment package for a Pages project.
func (c *Client) DownloadPagesProjectLatestPackage(ctx context.Context, projectID uint) ([]byte, error) {
res, err := c.base.DoRaw(ctx, http.MethodGet, fmt.Sprintf("/api/v1/agent/pages/projects/%d/latest/package", projectID), nil)
// DownloadPagesProjectLatestPackage streams the active deployment package into
// dst while enforcing maxBytes against both advertised and actual sizes.
func (c *Client) DownloadPagesProjectLatestPackage(
ctx context.Context,
projectID uint,
dst io.Writer,
maxBytes int64,
) (int64, error) {
return c.downloadPagesPackage(
ctx,
fmt.Sprintf("/api/v1/agent/pages/projects/%d/latest/package", projectID),
dst,
maxBytes,
)
}
func (c *Client) downloadPagesPackage(
ctx context.Context,
path string,
dst io.Writer,
maxBytes int64,
) (int64, error) {
if dst == nil {
return 0, errors.New("pages package destination is required")
}
if maxBytes <= 0 {
return 0, errors.New("pages package byte limit must be positive")
}
res, err := c.base.DoRaw(ctx, http.MethodGet, path, nil)
if err != nil {
return nil, err
return 0, err
}
defer func() { _ = res.Body.Close() }()
if res.StatusCode != http.StatusOK {
return nil, edgehttp.ReadHTTPError(res)
body, readErr := readPagesControlResponse(res, pagesControlResponseMaxBytes)
if readErr != nil {
return 0, readErr
}
return 0, edgehttp.ReadBodyError(body, res.Status)
}
return io.ReadAll(res.Body)
return copyPagesPackageResponse(dst, res, maxBytes)
}
func readPagesControlResponse(res *http.Response, maxBytes int64) ([]byte, error) {
if res.ContentLength > maxBytes {
return nil, fmt.Errorf(
"pages control response Content-Length %d exceeds limit %d",
res.ContentLength,
maxBytes,
)
}
limited := &io.LimitedReader{R: res.Body, N: maxBytes + 1}
body, err := io.ReadAll(limited)
if err != nil {
return nil, fmt.Errorf("read pages control response: %w", err)
}
if int64(len(body)) > maxBytes {
return nil, fmt.Errorf("pages control response body exceeds limit %d", maxBytes)
}
return body, nil
}
func copyPagesPackageResponse(dst io.Writer, res *http.Response, maxBytes int64) (int64, error) {
if res.ContentLength > maxBytes {
return 0, fmt.Errorf(
"pages package Content-Length %d exceeds limit %d",
res.ContentLength,
maxBytes,
)
}
limited := &io.LimitedReader{R: res.Body, N: maxBytes + 1}
written, err := io.Copy(dst, limited)
if err != nil {
return written, fmt.Errorf("stream pages package: %w", err)
}
if written > maxBytes {
return written, fmt.Errorf("pages package body exceeds limit %d", maxBytes)
}
return written, nil
}
// SetToken updates the authentication token used for API requests.
@@ -0,0 +1,109 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package httpclient
import (
"bytes"
"context"
"io"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
)
func TestDownloadPagesProjectLatestPackageRejectsChunkedBodyOverLimit(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusOK)
if flusher, ok := w.(http.Flusher); ok {
flusher.Flush()
}
_, _ = io.WriteString(w, "123456")
}))
defer server.Close()
client := New(server.URL, "test-token", time.Second)
var dst bytes.Buffer
written, err := client.DownloadPagesProjectLatestPackage(
context.Background(),
7,
&dst,
4,
)
if err == nil || !strings.Contains(err.Error(), "body exceeds limit") {
t.Fatalf("DownloadPagesProjectLatestPackage(chunked, limit=4) error = %v, want body limit error", err)
}
if written != 5 {
t.Errorf("DownloadPagesProjectLatestPackage(chunked, limit=4) written = %d, want 5", written)
}
}
func TestCopyPagesPackageResponseRejectsAdvertisedContentLengthBeforeWrite(t *testing.T) {
response := &http.Response{
Body: io.NopCloser(strings.NewReader("123456")),
ContentLength: 6,
}
var dst bytes.Buffer
written, err := copyPagesPackageResponse(&dst, response, 4)
if err == nil || !strings.Contains(err.Error(), "Content-Length") {
t.Fatalf("copyPagesPackageResponse(Content-Length=6, limit=4) error = %v, want Content-Length limit error", err)
}
if written != 0 || dst.Len() != 0 {
t.Errorf("copyPagesPackageResponse(Content-Length=6, limit=4) wrote (%d, %d buffered), want no writes", written, dst.Len())
}
}
func TestCopyPagesPackageResponseRejectsForgedSmallContentLength(t *testing.T) {
response := &http.Response{
Body: io.NopCloser(strings.NewReader("123456")),
ContentLength: 2,
}
var dst bytes.Buffer
written, err := copyPagesPackageResponse(&dst, response, 4)
if err == nil || !strings.Contains(err.Error(), "body exceeds limit") {
t.Fatalf("copyPagesPackageResponse(forged Content-Length=2, limit=4) error = %v, want body limit error", err)
}
if written != 5 {
t.Errorf("copyPagesPackageResponse(forged Content-Length=2, limit=4) written = %d, want 5", written)
}
}
func TestDownloadPagesProjectLatestPackageBoundsChunkedErrorResponse(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusBadRequest)
if flusher, ok := w.(http.Flusher); ok {
flusher.Flush()
}
_, _ = io.WriteString(w, strings.Repeat("x", int(pagesControlResponseMaxBytes+1)))
}))
defer server.Close()
client := New(server.URL, "test-token", time.Second)
var dst bytes.Buffer
_, err := client.DownloadPagesProjectLatestPackage(context.Background(), 7, &dst, 1024)
if err == nil || !strings.Contains(err.Error(), "control response body exceeds limit") {
t.Fatalf("DownloadPagesProjectLatestPackage(large chunked 400) error = %v, want bounded response error", err)
}
if dst.Len() != 0 {
t.Errorf("DownloadPagesProjectLatestPackage(large chunked 400) wrote %d package bytes, want 0", dst.Len())
}
}
func TestGetPagesProjectLatestHashBoundsChunkedMetadataResponse(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusOK)
if flusher, ok := w.(http.Flusher); ok {
flusher.Flush()
}
_, _ = io.WriteString(w, strings.Repeat("x", int(pagesControlResponseMaxBytes+1)))
}))
defer server.Close()
client := New(server.URL, "test-token", time.Second)
_, err := client.GetPagesProjectLatestHash(context.Background(), 7)
if err == nil || !strings.Contains(err.Error(), "control response body exceeds limit") {
t.Fatalf("GetPagesProjectLatestHash(large chunked metadata) error = %v, want bounded response error", err)
}
}
+565 -94
View File
@@ -1,3 +1,6 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package sync applies control-plane configuration to the local agent runtime.
package sync
@@ -20,9 +23,13 @@ import (
)
const (
pagesDirPerm = 0o755
pagesFilePerm = 0o644
pagesManifestFilePerm = 0o644
pagesDirPerm = 0o755
pagesFilePerm = 0o644
pagesManifestFilePerm = 0o644
agentPagesMaxPackageBytes = int64(2 * 1024 * 1024 * 1024)
agentPagesMaxFiles = 1000
agentPagesMaxFileBytes = int64(8 * 1024 * 1024 * 1024)
agentPagesMaxTotalBytes = int64(8 * 1024 * 1024 * 1024)
// pagesLatestPullAttempts covers a race where the active deployment changes
// between the hash probe and the package download.
pagesLatestPullAttempts = 2
@@ -45,6 +52,11 @@ type pagesProjectRef struct {
Checksum string
}
type pagesPackageLimits struct {
PackageBytes int64
Extraction pagesarchive.Limits
}
type pagesDeploymentMarker struct {
ProjectID uint `json:"project_id"`
DeploymentID uint `json:"deployment_id,omitempty"`
@@ -191,10 +203,11 @@ func (s *Service) ensurePagesProject(ctx context.Context, snapshot *state.Snapsh
if err != nil {
return fmt.Errorf("fetch Pages project %d latest hash: %w", projectID, err)
}
hash := strings.TrimSpace(latest.Hash)
if hash == "" {
return fmt.Errorf("pages project %d latest hash is empty", projectID)
limits, err := validatePagesPackageMetadata(projectID, latest)
if err != nil {
return err
}
hash := strings.TrimSpace(latest.Hash)
effective := pagesProjectRef{
ProjectID: projectID,
DeploymentID: latest.DeploymentID,
@@ -211,44 +224,65 @@ func (s *Service) ensurePagesProject(ctx context.Context, snapshot *state.Snapsh
return nil
}
packageBytes, err := s.client.DownloadPagesProjectLatestPackage(ctx, projectID)
packagePath, got, err := s.downloadPagesProjectPackage(ctx, projectID, latest, limits.PackageBytes)
if err != nil {
return fmt.Errorf("download Pages project %d latest package: %w", projectID, err)
}
got := checksumBytes(packageBytes)
// Re-probe latest after download to detect activation races.
// Accept the package only when its content hash still matches latest.
// A deployment-id-only change is still a latest-pointer race even when
// deduplication makes both deployments share the same package hash.
verify, err := s.client.GetPagesProjectLatestHash(ctx, projectID)
if err != nil {
_ = os.Remove(packagePath)
return fmt.Errorf("re-fetch Pages project %d latest hash: %w", projectID, err)
}
verifyHash := strings.TrimSpace(verify.Hash)
if verifyHash == "" {
return fmt.Errorf("pages project %d latest hash is empty", projectID)
if _, err := validatePagesPackageMetadata(projectID, verify); err != nil {
_ = os.Remove(packagePath)
return err
}
if got != verifyHash {
if !samePagesPackageMetadata(latest, verify) {
_ = os.Remove(packagePath)
lastErr = fmt.Errorf(
"pages project %d package/hash race: downloaded %s, latest now %s (attempt %d/%d)",
projectID, got, verifyHash, attempt+1, pagesLatestPullAttempts,
"pages project %d latest metadata changed during download: deployment %d/%s -> %d/%s (attempt %d/%d)",
projectID,
latest.DeploymentID,
strings.TrimSpace(latest.Hash),
verify.DeploymentID,
strings.TrimSpace(verify.Hash),
attempt+1,
pagesLatestPullAttempts,
)
slog.Warn("pages latest package race, retrying",
slog.Warn("pages latest metadata race, retrying",
"project_id", projectID,
"before_deployment_id", latest.DeploymentID,
"before_hash", strings.TrimSpace(latest.Hash),
"after_deployment_id", verify.DeploymentID,
"after_hash", strings.TrimSpace(verify.Hash),
"attempt", attempt+1,
)
continue
}
if got != hash {
_ = os.Remove(packagePath)
lastErr = fmt.Errorf(
"pages project %d package hash mismatch: downloaded %s, expected %s (attempt %d/%d)",
projectID, got, hash, attempt+1, pagesLatestPullAttempts,
)
slog.Warn("pages latest package hash mismatch, retrying",
"project_id", projectID,
"downloaded_hash", got,
"latest_hash", verifyHash,
"expected_hash", hash,
"attempt", attempt+1,
)
continue
}
effective = pagesProjectRef{
ProjectID: projectID,
DeploymentID: verify.DeploymentID,
Checksum: got,
}
releaseDir = pagesProjectReleaseDir(s.pagesDir, projectID, got)
if err := extractPagesPackage(packageBytes, releaseDir, effective); err != nil {
return err
extractErr := extractPagesPackageFile(packagePath, releaseDir, effective, limits.Extraction, latest)
_ = os.Remove(packagePath)
if extractErr != nil {
return extractErr
}
if err := switchPagesProjectCurrentDir(s.pagesDir, projectID, releaseDir); err != nil {
return err
@@ -264,6 +298,145 @@ func (s *Service) ensurePagesProject(ctx context.Context, snapshot *state.Snapsh
return fmt.Errorf("pages project %d latest pull failed", projectID)
}
func validatePagesPackageMetadata(
projectID uint,
metadata *protocol.PagesProjectLatestHashResponse,
) (pagesPackageLimits, error) {
if metadata == nil {
return pagesPackageLimits{}, fmt.Errorf("pages project %d latest metadata is missing", projectID)
}
if metadata.ProjectID != projectID {
return pagesPackageLimits{}, fmt.Errorf(
"pages project %d latest metadata has project id %d",
projectID,
metadata.ProjectID,
)
}
if metadata.DeploymentID == 0 {
return pagesPackageLimits{}, fmt.Errorf("pages project %d latest deployment id is missing", projectID)
}
if strings.TrimSpace(metadata.Hash) == "" {
return pagesPackageLimits{}, fmt.Errorf("pages project %d latest hash is empty", projectID)
}
if metadata.PackageSize <= 0 {
return pagesPackageLimits{}, fmt.Errorf("pages project %d package size must be positive", projectID)
}
if metadata.PackageSize > agentPagesMaxPackageBytes {
return pagesPackageLimits{}, fmt.Errorf(
"pages project %d package size %d exceeds agent limit %d",
projectID,
metadata.PackageSize,
agentPagesMaxPackageBytes,
)
}
if metadata.FileCount <= 0 {
return pagesPackageLimits{}, fmt.Errorf("pages project %d file count must be positive", projectID)
}
if metadata.FileCount > agentPagesMaxFiles {
return pagesPackageLimits{}, fmt.Errorf(
"pages project %d file count %d exceeds agent limit %d",
projectID,
metadata.FileCount,
agentPagesMaxFiles,
)
}
if metadata.TotalSize < 0 {
return pagesPackageLimits{}, fmt.Errorf("pages project %d total size cannot be negative", projectID)
}
if metadata.TotalSize > agentPagesMaxTotalBytes {
return pagesPackageLimits{}, fmt.Errorf(
"pages project %d total size %d exceeds agent limit %d",
projectID,
metadata.TotalSize,
agentPagesMaxTotalBytes,
)
}
// pagesarchive treats zero limits as defaults. A one-byte extraction guard
// plus the exact post-extraction manifest check below preserves the valid
// case of one or more zero-byte files while still enforcing total_size=0.
extractedBytes := metadata.TotalSize
if extractedBytes == 0 {
extractedBytes = 1
}
maxFileBytes := extractedBytes
if maxFileBytes > agentPagesMaxFileBytes {
maxFileBytes = agentPagesMaxFileBytes
}
return pagesPackageLimits{
PackageBytes: metadata.PackageSize,
Extraction: pagesarchive.Limits{
MaxFiles: metadata.FileCount,
MaxFileBytes: maxFileBytes,
MaxTotalBytes: extractedBytes,
},
}, nil
}
func samePagesPackageMetadata(
before *protocol.PagesProjectLatestHashResponse,
after *protocol.PagesProjectLatestHashResponse,
) bool {
if before == nil || after == nil {
return false
}
return before.ProjectID == after.ProjectID &&
before.DeploymentID == after.DeploymentID &&
strings.TrimSpace(before.Hash) == strings.TrimSpace(after.Hash) &&
before.PackageSize == after.PackageSize &&
before.FileCount == after.FileCount &&
before.TotalSize == after.TotalSize
}
func (s *Service) downloadPagesProjectPackage(
ctx context.Context,
projectID uint,
metadata *protocol.PagesProjectLatestHashResponse,
maxBytes int64,
) (packagePath string, hash string, err error) {
releasesRoot := filepath.Join(s.pagesDir, "projects", fmt.Sprintf("%d", projectID), "releases")
if err := os.MkdirAll(releasesRoot, pagesDirPerm); err != nil {
return "", "", err
}
packageFile, err := os.CreateTemp(releasesRoot, ".package-*.tmp")
if err != nil {
return "", "", err
}
packagePath = packageFile.Name()
keep := false
defer func() {
if closeErr := packageFile.Close(); err == nil && closeErr != nil {
err = closeErr
}
if !keep || err != nil {
_ = os.Remove(packagePath)
packagePath = ""
}
}()
hasher := sha256.New()
written, err := s.client.DownloadPagesProjectLatestPackage(
ctx,
projectID,
io.MultiWriter(packageFile, hasher),
maxBytes,
)
if err != nil {
return "", "", err
}
if written != metadata.PackageSize {
return "", "", fmt.Errorf(
"pages project %d package size %d does not match metadata %d",
projectID,
written,
metadata.PackageSize,
)
}
keep = true
return packagePath, hex.EncodeToString(hasher.Sum(nil)), nil
}
// cleanupPagesProjectStaleReleases keeps only keepHash under projects/{id}/releases.
// Must be called only after the keepHash release is ready and current points at it.
func cleanupPagesProjectStaleReleases(baseDir string, projectID uint, keepHash string) error {
@@ -372,40 +545,258 @@ type pagesDeploymentSource struct {
Checksum string `json:"checksum"`
}
func extractPagesPackage(packageBytes []byte, releaseDir string, project pagesProjectRef) error {
tmpDir := releaseDir + ".tmp"
_ = os.RemoveAll(tmpDir)
if err := os.MkdirAll(tmpDir, pagesDirPerm); err != nil {
func extractPagesPackageFile(
packagePath string,
releaseDir string,
project pagesProjectRef,
limits pagesarchive.Limits,
expected *protocol.PagesProjectLatestHashResponse,
) error {
if err := os.MkdirAll(filepath.Dir(releaseDir), pagesDirPerm); err != nil {
return err
}
format, err := pagesarchive.DetectFormat("", packageBytes)
stagingDir, err := os.MkdirTemp(
filepath.Dir(releaseDir),
"."+filepath.Base(releaseDir)+"-*.tmp",
)
if err != nil {
_ = os.RemoveAll(tmpDir)
return fmt.Errorf("detect Pages package format: %w", err)
return err
}
// Control plane already inspected and accepted this package.
if err := pagesarchive.ExtractBytes(packageBytes, format, tmpDir, pagesarchive.ExtractOptions{
cleanupStaging := true
defer func() {
if cleanupStaging {
removePagesStagingUnlessCurrent(stagingDir, pagesCurrentDirFromRelease(releaseDir))
}
}()
if err := pagesarchive.ExtractFile(packagePath, "", stagingDir, pagesarchive.ExtractOptions{
StripCommonRoot: true,
EnforceLimits: false,
EnforceLimits: true,
Limits: limits,
}); err != nil {
_ = os.RemoveAll(tmpDir)
return fmt.Errorf("extract Pages package: %w", err)
}
if err := writePagesMarker(tmpDir, project); err != nil {
_ = os.RemoveAll(tmpDir)
if err := validateExtractedPagesMetadata(stagingDir, expected); err != nil {
return err
}
_ = os.RemoveAll(releaseDir)
return os.Rename(tmpDir, releaseDir)
if err := writePagesMarker(stagingDir, project); err != nil {
return err
}
if err := promotePagesRelease(stagingDir, releaseDir, project); err != nil {
return err
}
cleanupStaging = false
return nil
}
func switchPagesProjectCurrentDir(baseDir string, projectID uint, releaseDir string) error {
currentDir := pagesProjectCurrentDir(baseDir, projectID)
previousDir := currentDir + ".previous"
_ = os.RemoveAll(previousDir)
func validateExtractedPagesMetadata(
dir string,
expected *protocol.PagesProjectLatestHashResponse,
) error {
if expected == nil {
return nil
}
fileCount := 0
totalSize := int64(0)
err := filepath.WalkDir(dir, func(path string, entry os.DirEntry, walkErr error) error {
if walkErr != nil {
return walkErr
}
if entry.IsDir() {
return nil
}
info, err := entry.Info()
if err != nil {
return err
}
if !info.Mode().IsRegular() {
return fmt.Errorf("pages extracted entry is not a regular file: %s", path)
}
fileCount++
if fileCount > agentPagesMaxFiles {
return fmt.Errorf("pages extracted file count exceeds agent limit %d", agentPagesMaxFiles)
}
if info.Size() < 0 || info.Size() > agentPagesMaxTotalBytes-totalSize {
return fmt.Errorf("pages extracted size exceeds agent limit %d", agentPagesMaxTotalBytes)
}
totalSize += info.Size()
return nil
})
if err != nil {
return fmt.Errorf("validate extracted Pages package: %w", err)
}
if fileCount != expected.FileCount || totalSize != expected.TotalSize {
return fmt.Errorf(
"pages extracted metadata mismatch: got %d files/%d bytes, expected %d files/%d bytes",
fileCount,
totalSize,
expected.FileCount,
expected.TotalSize,
)
}
return nil
}
func promotePagesRelease(stagingDir string, releaseDir string, project pagesProjectRef) error {
return promotePagesReleaseWithCopy(stagingDir, releaseDir, project, copyPagesDir)
}
func promotePagesReleaseWithCopy(
stagingDir string,
releaseDir string,
project pagesProjectRef,
copyDir func(string, string) error,
) error {
currentDir := pagesCurrentDirFromRelease(releaseDir)
defer removePagesStagingUnlessCurrent(stagingDir, currentDir)
currentUsesRelease, err := pagesCurrentTargetsRelease(currentDir, releaseDir)
if err != nil {
return err
}
if !currentUsesRelease {
if err := os.RemoveAll(releaseDir); err != nil {
return err
}
return os.Rename(stagingDir, releaseDir)
}
// A same-hash repair cannot remove releaseDir while current still resolves
// through it. Keep traffic on the fully validated staging tree, rebuild the
// canonical release, then atomically point current back to the canonical path.
if err := switchPagesCurrentDir(currentDir, stagingDir, os.Rename); err != nil {
return fmt.Errorf("switch Pages current to repair staging: %w", err)
}
backupDir := stagingDir + ".previous"
if err := os.Rename(releaseDir, backupDir); err != nil {
restoreErr := switchPagesCurrentDir(currentDir, releaseDir, os.Rename)
return errors.Join(
fmt.Errorf("move previous Pages release aside: %w", err),
restoreErr,
)
}
rollback := func(cause error) error {
var rollbackErrors []error
rollbackErrors = append(rollbackErrors, cause)
if err := os.RemoveAll(releaseDir); err != nil {
rollbackErrors = append(rollbackErrors, fmt.Errorf("remove failed Pages release repair: %w", err))
}
if err := os.Rename(backupDir, releaseDir); err != nil {
rollbackErrors = append(rollbackErrors, fmt.Errorf("restore previous Pages release: %w", err))
return errors.Join(rollbackErrors...)
}
if err := switchPagesCurrentDir(currentDir, releaseDir, os.Rename); err != nil {
rollbackErrors = append(rollbackErrors, fmt.Errorf("restore previous Pages current target: %w", err))
}
return errors.Join(rollbackErrors...)
}
if err := copyDir(stagingDir, releaseDir); err != nil {
return rollback(fmt.Errorf("copy repaired Pages release: %w", err))
}
if !pagesProjectReleaseReady(releaseDir, project) {
return rollback(errors.New("repaired Pages release is not ready"))
}
if err := switchPagesCurrentDir(currentDir, releaseDir, os.Rename); err != nil {
return rollback(fmt.Errorf("switch Pages current to repaired release: %w", err))
}
if err := os.RemoveAll(backupDir); err != nil {
slog.Warn("failed to remove previous Pages release", "path", backupDir, "error", err)
}
if err := os.RemoveAll(stagingDir); err != nil {
slog.Warn("failed to remove Pages repair staging", "path", stagingDir, "error", err)
}
return nil
}
func pagesCurrentDirFromRelease(releaseDir string) string {
return filepath.Join(filepath.Dir(filepath.Dir(releaseDir)), "current")
}
func pagesCurrentTargetsRelease(currentDir string, releaseDir string) (bool, error) {
if _, err := os.Lstat(currentDir); err != nil {
if os.IsNotExist(err) {
return false, nil
}
return false, err
}
currentInfo, err := os.Stat(currentDir)
if err != nil {
if os.IsNotExist(err) {
return false, nil
}
return false, fmt.Errorf("stat Pages current target: %w", err)
}
releaseInfo, err := os.Stat(releaseDir)
if err != nil {
if os.IsNotExist(err) {
return false, nil
}
return false, err
}
return os.SameFile(currentInfo, releaseInfo), nil
}
func removePagesStagingUnlessCurrent(stagingDir string, currentDir string) {
currentUsesStaging, err := pagesCurrentTargetsRelease(currentDir, stagingDir)
if err == nil && currentUsesStaging {
slog.Error("preserving Pages staging because current still references it", "path", stagingDir)
return
}
if removeErr := os.RemoveAll(stagingDir); removeErr != nil {
slog.Warn("failed to remove Pages staging", "path", stagingDir, "error", removeErr)
}
}
func verifyPagesCurrentTarget(currentDir string, releaseDir string) error {
currentInfo, err := os.Stat(currentDir)
if err != nil {
return fmt.Errorf("stat Pages current target: %w", err)
}
releaseInfo, err := os.Stat(releaseDir)
if err != nil {
return fmt.Errorf("stat Pages release target: %w", err)
}
if !os.SameFile(currentInfo, releaseInfo) {
return fmt.Errorf("pages current target does not resolve to release %s", releaseDir)
}
return nil
}
func switchPagesCurrentDir(
currentDir string,
releaseDir string,
rename func(string, string) error,
) error {
return switchPagesCurrentDirWithOps(currentDir, releaseDir, rename, os.Symlink)
}
func switchPagesCurrentDirWithOps(
currentDir string,
releaseDir string,
rename func(string, string) error,
symlink func(string, string) error,
) error {
if err := os.MkdirAll(filepath.Dir(currentDir), pagesDirPerm); err != nil {
return err
}
currentInfo, currentErr := os.Lstat(currentDir)
if currentErr != nil && !os.IsNotExist(currentErr) {
return currentErr
}
if currentErr == nil && currentInfo.Mode()&os.ModeSymlink == 0 {
return fallbackCopyPagesCurrentDir(currentDir, releaseDir, rename)
}
previousTarget := ""
hadPrevious := currentErr == nil
if hadPrevious {
var err error
previousTarget, err = os.Readlink(currentDir)
if err != nil {
return err
}
}
relTarget, err := filepath.Rel(filepath.Dir(currentDir), releaseDir)
if err != nil {
@@ -413,44 +804,123 @@ func switchPagesProjectCurrentDir(baseDir string, projectID uint, releaseDir str
}
tmpSymlink := currentDir + ".tmp"
_ = os.Remove(tmpSymlink)
symlinkErr := os.Symlink(relTarget, tmpSymlink)
if symlinkErr != nil {
return fallbackCopyPagesCurrentDir(currentDir, previousDir, releaseDir)
}
_ = os.Remove(tmpSymlink)
if _, err := os.Lstat(currentDir); err == nil {
if err := os.Rename(currentDir, previousDir); err != nil {
return err
}
}
if err := os.Symlink(relTarget, currentDir); err != nil {
if _, restoreErr := os.Lstat(previousDir); restoreErr == nil {
_ = os.Rename(previousDir, currentDir)
}
if err := os.Remove(tmpSymlink); err != nil && !os.IsNotExist(err) {
return err
}
_ = os.RemoveAll(previousDir)
if err := symlink(relTarget, tmpSymlink); err != nil {
_ = os.Remove(tmpSymlink)
return fallbackCopyPagesCurrentDir(currentDir, releaseDir, rename)
}
defer func() { _ = os.Remove(tmpSymlink) }()
if err := verifyPagesCurrentTarget(tmpSymlink, releaseDir); err != nil {
return err
}
if err := rename(tmpSymlink, currentDir); err != nil {
return err
}
if err := verifyPagesCurrentTarget(currentDir, releaseDir); err != nil {
rollbackErr := rollbackPagesCurrentSymlink(
currentDir,
previousTarget,
hadPrevious,
rename,
)
return errors.Join(err, rollbackErr)
}
return nil
}
func fallbackCopyPagesCurrentDir(currentDir, previousDir, releaseDir string) error {
if _, err := os.Lstat(currentDir); err == nil {
if err := os.Rename(currentDir, previousDir); err != nil {
return err
}
}
if err := copyPagesDir(releaseDir, currentDir); err != nil {
_ = os.RemoveAll(currentDir)
if _, restoreErr := os.Lstat(previousDir); restoreErr == nil {
_ = os.Rename(previousDir, currentDir)
}
func fallbackCopyPagesCurrentDir(
currentDir string,
releaseDir string,
rename func(string, string) error,
) error {
stagingDir := currentDir + ".copy.tmp"
previousDir := currentDir + ".previous"
if err := os.RemoveAll(stagingDir); err != nil {
return err
}
_ = os.RemoveAll(previousDir)
if err := copyPagesDir(releaseDir, stagingDir); err != nil {
_ = os.RemoveAll(stagingDir)
return err
}
if err := os.RemoveAll(previousDir); err != nil {
_ = os.RemoveAll(stagingDir)
return err
}
hadPrevious := false
if _, err := os.Lstat(currentDir); err == nil {
if err := rename(currentDir, previousDir); err != nil {
_ = os.RemoveAll(stagingDir)
return err
}
hadPrevious = true
} else if !os.IsNotExist(err) {
_ = os.RemoveAll(stagingDir)
return err
}
if err := rename(stagingDir, currentDir); err != nil {
var restoreErr error
if hadPrevious {
restoreErr = rename(previousDir, currentDir)
}
_ = os.RemoveAll(stagingDir)
return errors.Join(err, restoreErr)
}
if err := os.RemoveAll(previousDir); err != nil {
slog.Warn("failed to remove previous Pages current directory", "path", previousDir, "error", err)
}
return nil
}
func switchPagesProjectCurrentDir(baseDir string, projectID uint, releaseDir string) error {
return switchPagesProjectCurrentDirWithRename(baseDir, projectID, releaseDir, os.Rename)
}
func switchPagesProjectCurrentDirWithRename(
baseDir string,
projectID uint,
releaseDir string,
rename func(string, string) error,
) error {
return switchPagesCurrentDir(pagesProjectCurrentDir(baseDir, projectID), releaseDir, rename)
}
func rollbackPagesCurrentSymlink(
currentDir string,
previousTarget string,
hadPrevious bool,
rename func(string, string) error,
) error {
if !hadPrevious {
if err := os.Remove(currentDir); err != nil && !os.IsNotExist(err) {
return fmt.Errorf("remove unverified Pages current symlink: %w", err)
}
return nil
}
rollbackSymlink := currentDir + ".rollback.tmp"
if err := os.Remove(rollbackSymlink); err != nil && !os.IsNotExist(err) {
return err
}
if err := os.Symlink(previousTarget, rollbackSymlink); err != nil {
return err
}
defer func() { _ = os.Remove(rollbackSymlink) }()
if err := rename(rollbackSymlink, currentDir); err != nil {
return fmt.Errorf("restore previous Pages current symlink: %w", err)
}
gotTarget, err := os.Readlink(currentDir)
if err != nil {
return fmt.Errorf("verify restored Pages current symlink: %w", err)
}
if gotTarget != previousTarget {
return fmt.Errorf(
"restored Pages current symlink target %q does not match %q",
gotTarget,
previousTarget,
)
}
return nil
}
@@ -467,24 +937,30 @@ func copyPagesDir(sourceDir string, targetDir string) error {
if entry.IsDir() {
return os.MkdirAll(targetPath, pagesDirPerm)
}
input, err := os.Open(sourcePath) //nolint:gosec // sourcePath is under managed PagesDir walk root
if err != nil {
return err
}
defer func() { _ = input.Close() }()
if err := os.MkdirAll(filepath.Dir(targetPath), pagesDirPerm); err != nil {
return err
}
output, err := os.OpenFile(targetPath, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, pagesFilePerm) //nolint:gosec // targetPath is under managed PagesDir walk root
if err != nil {
return err
}
defer func() { _ = output.Close() }()
_, err = io.Copy(output, input)
return err
return copyPagesFile(sourcePath, targetPath)
})
}
func copyPagesFile(sourcePath string, targetPath string) error {
input, err := os.Open(sourcePath) //nolint:gosec // sourcePath is under managed PagesDir walk root
if err != nil {
return err
}
if err := os.MkdirAll(filepath.Dir(targetPath), pagesDirPerm); err != nil {
_ = input.Close()
return err
}
output, err := os.OpenFile(targetPath, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, pagesFilePerm) //nolint:gosec // targetPath is under managed PagesDir walk root
if err != nil {
_ = input.Close()
return err
}
_, copyErr := io.Copy(output, input)
outputCloseErr := output.Close()
inputCloseErr := input.Close()
return errors.Join(copyErr, outputCloseErr, inputCloseErr)
}
func markerMatches(dir string, project pagesProjectRef) bool {
data, err := os.ReadFile(filepath.Join(dir, ".openflare-pages.json")) //nolint:gosec // dir is managed PagesDir
if err != nil {
@@ -515,8 +991,3 @@ func pagesProjectCurrentDir(baseDir string, projectID uint) string {
func pagesProjectReleaseDir(baseDir string, projectID uint, checksum string) string {
return filepath.Join(baseDir, "projects", fmt.Sprintf("%d", projectID), "releases", checksum)
}
func checksumBytes(data []byte) string {
sum := sha256.Sum256(data)
return hex.EncodeToString(sum[:])
}
@@ -0,0 +1,388 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package sync
import (
"context"
"errors"
"os"
"path/filepath"
"strings"
"testing"
"github.com/Rain-kl/Wavelet/internal/apps/agent/protocol"
"github.com/Rain-kl/Wavelet/internal/apps/agent/state"
)
func TestEnsurePagesProjectRejectsMetadataBeyondAgentCapsBeforeDownload(t *testing.T) {
packageBytes := testPagesPackage(t, map[string]string{"index.html": "x"})
base := protocol.PagesProjectLatestHashResponse{
ProjectID: 1,
DeploymentID: 1,
Hash: testBytesChecksum(packageBytes),
PackageSize: int64(len(packageBytes)),
FileCount: 1,
TotalSize: 1,
}
tests := []struct {
name string
mutate func(*protocol.PagesProjectLatestHashResponse)
}{
{
name: "package size",
mutate: func(metadata *protocol.PagesProjectLatestHashResponse) {
metadata.PackageSize = agentPagesMaxPackageBytes + 1
},
},
{
name: "file count",
mutate: func(metadata *protocol.PagesProjectLatestHashResponse) {
metadata.FileCount = agentPagesMaxFiles + 1
},
},
{
name: "total size",
mutate: func(metadata *protocol.PagesProjectLatestHashResponse) {
metadata.TotalSize = agentPagesMaxTotalBytes + 1
},
},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
metadata := base
test.mutate(&metadata)
client := &fakeClient{
pagesPackages: map[uint][]byte{1: packageBytes},
pagesMetadata: map[uint]protocol.PagesProjectLatestHashResponse{1: metadata},
}
service := New(client, &fakeManager{}, nil)
service.SetPagesDir(t.TempDir())
err := service.ensurePagesProject(context.Background(), &state.Snapshot{}, 1)
if err == nil || !strings.Contains(err.Error(), "agent limit") {
t.Fatalf("ensurePagesProject(%s metadata) error = %v, want agent limit error", test.name, err)
}
if client.pagesPackageDownloads != 0 {
t.Errorf("ensurePagesProject(%s metadata) downloads = %d, want 0", test.name, client.pagesPackageDownloads)
}
})
}
}
func TestEnsurePagesProjectRetriesSameHashDifferentDeployment(t *testing.T) {
packageBytes := testPagesPackage(t, map[string]string{"index.html": "same"})
hash := testBytesChecksum(packageBytes)
client := &racingLatestClient{
pkgA: packageBytes,
pkgB: packageBytes,
hashA: hash,
hashB: hash,
}
service := New(client, &fakeManager{}, nil)
pagesDir := t.TempDir()
service.SetPagesDir(pagesDir)
snapshot := &state.Snapshot{PagesDeployments: []state.PagesDeployment{{ProjectID: 42}}}
if err := service.ensurePagesProject(context.Background(), snapshot, 42); err != nil {
t.Fatalf("ensurePagesProject(same hash deployment race) error = %v", err)
}
if client.downloadCalls != 2 {
t.Errorf("ensurePagesProject(same hash deployment race) downloads = %d, want 2", client.downloadCalls)
}
if snapshot.PagesDeployments[0].DeploymentID != 2 || snapshot.PagesDeployments[0].Hash != hash {
t.Errorf("snapshot Pages deployment = %+v, want deployment 2/hash %s", snapshot.PagesDeployments[0], hash)
}
}
func TestEnsurePagesProjectExtractionFailureCleansTempAndPreservesCurrent(t *testing.T) {
projectID := uint(9)
oldPackage := testPagesPackage(t, map[string]string{"index.html": "old"})
oldHash := testBytesChecksum(oldPackage)
newPackage := testPagesPackage(t, map[string]string{"index.html": "new"})
newHash := testBytesChecksum(newPackage)
pagesDir := t.TempDir()
oldRelease := pagesProjectReleaseDir(pagesDir, projectID, oldHash)
if err := extractTestPagesPackage(t, oldPackage, oldRelease, pagesProjectRef{
ProjectID: projectID,
DeploymentID: 1,
Checksum: oldHash,
}); err != nil {
t.Fatalf("extractTestPagesPackage(old) error = %v", err)
}
if err := switchPagesProjectCurrentDir(pagesDir, projectID, oldRelease); err != nil {
t.Fatalf("switchPagesProjectCurrentDir(old) error = %v", err)
}
client := &fakeClient{
pagesPackages: map[uint][]byte{projectID: newPackage},
pagesMetadata: map[uint]protocol.PagesProjectLatestHashResponse{
projectID: {
ProjectID: projectID,
DeploymentID: 2,
Hash: newHash,
PackageSize: int64(len(newPackage)),
FileCount: 1,
TotalSize: 2, // Smaller than the actual three-byte file.
},
},
}
service := New(client, &fakeManager{}, nil)
service.SetPagesDir(pagesDir)
err := service.ensurePagesProject(context.Background(), &state.Snapshot{}, projectID)
if err == nil {
t.Fatal("ensurePagesProject(metadata-tightened extraction) error = nil, want error")
}
current, readErr := os.ReadFile(pagesProjectCurrentDir(pagesDir, projectID) + "/index.html")
if readErr != nil {
t.Fatalf("read old current after failed extraction error = %v", readErr)
}
if string(current) != "old" {
t.Errorf("current content after failed extraction = %q, want %q", current, "old")
}
entries, readErr := os.ReadDir(filepath.Join(pagesDir, "projects", "9", "releases"))
if readErr != nil {
t.Fatalf("read releases after failed extraction error = %v", readErr)
}
if len(entries) != 1 || entries[0].Name() != oldHash {
t.Errorf("releases after failed extraction = %v, want only %s", entries, oldHash)
}
}
func TestEnsurePagesProjectAcceptsAllZeroByteFiles(t *testing.T) {
packageBytes := testPagesPackage(t, map[string]string{
"index.html": "",
".gitkeep": "",
})
client := &fakeClient{pagesPackages: map[uint][]byte{5: packageBytes}}
service := New(client, &fakeManager{}, nil)
pagesDir := t.TempDir()
service.SetPagesDir(pagesDir)
if err := service.ensurePagesProject(context.Background(), &state.Snapshot{}, 5); err != nil {
t.Fatalf("ensurePagesProject(all-zero files) error = %v", err)
}
for _, name := range []string{"index.html", ".gitkeep"} {
info, err := os.Stat(filepath.Join(pagesProjectCurrentDir(pagesDir, 5), name))
if err != nil {
t.Errorf("stat all-zero file %q error = %v", name, err)
continue
}
if info.Size() != 0 {
t.Errorf("all-zero file %q size = %d, want 0", name, info.Size())
}
}
}
func TestSwitchPagesProjectCurrentDirRenameFailureKeepsPreviousCurrent(t *testing.T) {
pagesDir := t.TempDir()
projectID := uint(21)
oldRelease := pagesProjectReleaseDir(pagesDir, projectID, "old")
newRelease := pagesProjectReleaseDir(pagesDir, projectID, "new")
for path, content := range map[string]string{
oldRelease: "old",
newRelease: "new",
} {
if err := os.MkdirAll(path, pagesDirPerm); err != nil {
t.Fatalf("mkdir release %q error = %v", path, err)
}
if err := os.WriteFile(filepath.Join(path, "index.html"), []byte(content), pagesFilePerm); err != nil {
t.Fatalf("write release %q error = %v", path, err)
}
}
if err := switchPagesProjectCurrentDir(pagesDir, projectID, oldRelease); err != nil {
t.Fatalf("seed previous current error = %v", err)
}
currentDir := pagesProjectCurrentDir(pagesDir, projectID)
renameErr := errors.New("injected current rename failure")
err := switchPagesProjectCurrentDirWithRename(
pagesDir,
projectID,
newRelease,
func(oldPath string, newPath string) error {
if oldPath == currentDir+".tmp" && newPath == currentDir {
return renameErr
}
return os.Rename(oldPath, newPath)
},
)
if !errors.Is(err, renameErr) {
t.Fatalf("switchPagesProjectCurrentDirWithRename() error = %v, want injected rename error", err)
}
current, err := os.ReadFile(filepath.Join(currentDir, "index.html"))
if err != nil {
t.Fatalf("read previous current after rename failure error = %v", err)
}
if string(current) != "old" {
t.Errorf("current after rename failure = %q, want %q", current, "old")
}
if _, err := os.Lstat(currentDir + ".tmp"); !os.IsNotExist(err) {
t.Errorf("temporary current symlink remains after rename failure: %v", err)
}
}
func TestPromoteSameHashReleaseFailureRestoresPreviousCurrent(t *testing.T) {
pagesDir := t.TempDir()
projectID := uint(22)
hash := "same-hash"
project := pagesProjectRef{ProjectID: projectID, DeploymentID: 2, Checksum: hash}
releaseDir := pagesProjectReleaseDir(pagesDir, projectID, hash)
if err := os.MkdirAll(releaseDir, pagesDirPerm); err != nil {
t.Fatalf("mkdir previous same-hash release error = %v", err)
}
if err := os.WriteFile(filepath.Join(releaseDir, "index.html"), []byte("old"), pagesFilePerm); err != nil {
t.Fatalf("write previous same-hash release error = %v", err)
}
if err := writePagesMarker(releaseDir, project); err != nil {
t.Fatalf("write previous same-hash marker error = %v", err)
}
if err := switchPagesProjectCurrentDir(pagesDir, projectID, releaseDir); err != nil {
t.Fatalf("seed same-hash current error = %v", err)
}
stagingDir, err := os.MkdirTemp(filepath.Dir(releaseDir), ".same-hash-*.tmp")
if err != nil {
t.Fatalf("create same-hash staging error = %v", err)
}
if err := os.WriteFile(filepath.Join(stagingDir, "index.html"), []byte("new"), pagesFilePerm); err != nil {
t.Fatalf("write repaired same-hash release error = %v", err)
}
if err := writePagesMarker(stagingDir, project); err != nil {
t.Fatalf("write repaired same-hash marker error = %v", err)
}
copyErr := errors.New("injected same-hash copy failure")
err = promotePagesReleaseWithCopy(
stagingDir,
releaseDir,
project,
func(_ string, targetDir string) error {
if err := os.MkdirAll(targetDir, pagesDirPerm); err != nil {
return err
}
if err := os.WriteFile(filepath.Join(targetDir, "index.html"), []byte("partial"), pagesFilePerm); err != nil {
return err
}
return copyErr
},
)
if !errors.Is(err, copyErr) {
t.Fatalf("promotePagesReleaseWithCopy() error = %v, want injected copy error", err)
}
for name, path := range map[string]string{
"current": filepath.Join(pagesProjectCurrentDir(pagesDir, projectID), "index.html"),
"release": filepath.Join(releaseDir, "index.html"),
} {
content, readErr := os.ReadFile(path)
if readErr != nil {
t.Fatalf("read restored %s after same-hash repair failure error = %v", name, readErr)
}
if string(content) != "old" {
t.Errorf("restored %s after same-hash repair failure = %q, want %q", name, content, "old")
}
}
if _, err := os.Stat(stagingDir); !os.IsNotExist(err) {
t.Errorf("same-hash staging remains after successful rollback: %v", err)
}
}
func TestPromotePagesReleaseRepairsDanglingCurrent(t *testing.T) {
pagesDir := t.TempDir()
projectID := uint(23)
releaseDir := pagesProjectReleaseDir(pagesDir, projectID, "new-hash")
currentDir := pagesProjectCurrentDir(pagesDir, projectID)
requireTestMkdirAll(t, filepath.Dir(currentDir))
relTarget, err := filepath.Rel(filepath.Dir(currentDir), releaseDir)
if err != nil {
t.Fatalf("relative release target error = %v", err)
}
if err := os.Symlink(relTarget, currentDir); err != nil {
t.Skipf("symlink unsupported: %v", err)
}
requireTestMkdirAll(t, filepath.Dir(releaseDir))
stagingDir, err := os.MkdirTemp(filepath.Dir(releaseDir), ".dangling-*.tmp")
if err != nil {
t.Fatalf("create dangling repair staging error = %v", err)
}
if err := os.WriteFile(filepath.Join(stagingDir, "index.html"), []byte("repaired"), pagesFilePerm); err != nil {
t.Fatalf("write dangling repair staging error = %v", err)
}
project := pagesProjectRef{ProjectID: projectID, DeploymentID: 1, Checksum: "new-hash"}
if err := writePagesMarker(stagingDir, project); err != nil {
t.Fatalf("write dangling repair marker error = %v", err)
}
if err := promotePagesRelease(stagingDir, releaseDir, project); err != nil {
t.Fatalf("promotePagesRelease(dangling current) error = %v", err)
}
content, err := os.ReadFile(filepath.Join(currentDir, "index.html"))
if err != nil {
t.Fatalf("read repaired dangling current error = %v", err)
}
if string(content) != "repaired" {
t.Errorf("repaired dangling current = %q, want %q", content, "repaired")
}
}
func TestSwitchPagesCurrentDirCopiesOverLegacyDirectory(t *testing.T) {
pagesDir := t.TempDir()
currentDir := filepath.Join(pagesDir, "current")
releaseDir := filepath.Join(pagesDir, "releases", "new")
requireTestMkdirAll(t, currentDir)
requireTestMkdirAll(t, releaseDir)
if err := os.WriteFile(filepath.Join(currentDir, "index.html"), []byte("old"), pagesFilePerm); err != nil {
t.Fatalf("write legacy current error = %v", err)
}
if err := os.WriteFile(filepath.Join(releaseDir, "index.html"), []byte("new"), pagesFilePerm); err != nil {
t.Fatalf("write new release error = %v", err)
}
if err := switchPagesCurrentDir(currentDir, releaseDir, os.Rename); err != nil {
t.Fatalf("switchPagesCurrentDir(legacy directory) error = %v", err)
}
content, err := os.ReadFile(filepath.Join(currentDir, "index.html"))
if err != nil {
t.Fatalf("read copied legacy current error = %v", err)
}
if string(content) != "new" {
t.Errorf("copied legacy current = %q, want %q", content, "new")
}
}
func TestSwitchPagesCurrentDirFallsBackWhenSymlinkUnavailable(t *testing.T) {
pagesDir := t.TempDir()
currentDir := filepath.Join(pagesDir, "current")
releaseDir := filepath.Join(pagesDir, "releases", "new")
requireTestMkdirAll(t, releaseDir)
if err := os.WriteFile(filepath.Join(releaseDir, "index.html"), []byte("new"), pagesFilePerm); err != nil {
t.Fatalf("write fallback release error = %v", err)
}
symlinkErr := errors.New("injected symlink unavailable")
if err := switchPagesCurrentDirWithOps(
currentDir,
releaseDir,
os.Rename,
func(string, string) error { return symlinkErr },
); err != nil {
t.Fatalf("switchPagesCurrentDirWithOps(symlink unavailable) error = %v", err)
}
info, err := os.Lstat(currentDir)
if err != nil {
t.Fatalf("lstat copied current error = %v", err)
}
if !info.IsDir() {
t.Errorf("copied current mode = %v, want directory", info.Mode())
}
content, err := os.ReadFile(filepath.Join(currentDir, "index.html"))
if err != nil {
t.Fatalf("read fallback current error = %v", err)
}
if string(content) != "new" {
t.Errorf("fallback current = %q, want %q", content, "new")
}
}
func requireTestMkdirAll(t *testing.T, dir string) {
t.Helper()
if err := os.MkdirAll(dir, pagesDirPerm); err != nil {
t.Fatalf("mkdir %q error = %v", dir, err)
}
}
+6 -2
View File
@@ -1,3 +1,6 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package sync
import (
@@ -7,6 +10,7 @@ import (
"encoding/json"
"errors"
"fmt"
"io"
"log/slog"
"sort"
"strconv"
@@ -31,9 +35,9 @@ const (
type ConfigClient interface {
GetActiveConfig(ctx context.Context) (*protocol.ActiveConfigResponse, error)
GetPagesDeploymentHash(ctx context.Context, deploymentID uint) (string, error)
DownloadPagesDeploymentPackage(ctx context.Context, deploymentID uint) ([]byte, error)
DownloadPagesDeploymentPackage(ctx context.Context, deploymentID uint, dst io.Writer, maxBytes int64) (int64, error)
GetPagesProjectLatestHash(ctx context.Context, projectID uint) (*protocol.PagesProjectLatestHashResponse, error)
DownloadPagesProjectLatestPackage(ctx context.Context, projectID uint) ([]byte, error)
DownloadPagesProjectLatestPackage(ctx context.Context, projectID uint, dst io.Writer, maxBytes int64) (int64, error)
ReportApplyLog(ctx context.Context, payload protocol.ApplyLogPayload) error
SyncWAFIPGroups(ctx context.Context, payload protocol.WAFIPGroupSyncRequest) (*protocol.WAFIPGroupSyncResponse, error)
}
+128 -18
View File
@@ -1,3 +1,6 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package sync
import (
@@ -7,6 +10,7 @@ import (
"crypto/sha256"
"encoding/hex"
"fmt"
"io"
"os"
"path/filepath"
"strings"
@@ -16,6 +20,7 @@ import (
"github.com/Rain-kl/Wavelet/internal/apps/agent/nginx"
"github.com/Rain-kl/Wavelet/internal/apps/agent/protocol"
"github.com/Rain-kl/Wavelet/internal/apps/agent/state"
"github.com/Rain-kl/Wavelet/pkg/pagesarchive"
)
type fakeExecutor struct {
@@ -28,15 +33,17 @@ func testPagesSourceConfigJSON(projectID, deploymentID uint, checksum string) st
}
type fakeClient struct {
config protocol.ActiveConfigResponse
reports []protocol.ApplyLogPayload
wafSyncCalls []protocol.WAFIPGroupSyncRequest
pagesPackages map[uint][]byte // key: project_id (latest package)
pagesHashes map[uint]string // key: project_id
pagesLatestDeployIDs map[uint]uint // key: project_id → deployment_id
wafSyncResult protocol.WAFIPGroupSyncResponse
fetchCalls int
hashCalls int
config protocol.ActiveConfigResponse
reports []protocol.ApplyLogPayload
wafSyncCalls []protocol.WAFIPGroupSyncRequest
pagesPackages map[uint][]byte // key: project_id (latest package)
pagesHashes map[uint]string // key: project_id
pagesLatestDeployIDs map[uint]uint // key: project_id → deployment_id
pagesMetadata map[uint]protocol.PagesProjectLatestHashResponse
pagesPackageDownloads int
wafSyncResult protocol.WAFIPGroupSyncResponse
fetchCalls int
hashCalls int
}
type fakeManager struct {
@@ -98,17 +105,28 @@ func (f *fakeClient) GetPagesDeploymentHash(ctx context.Context, deploymentID ui
return f.projectHash(deploymentID)
}
func (f *fakeClient) DownloadPagesDeploymentPackage(ctx context.Context, deploymentID uint) ([]byte, error) {
func (f *fakeClient) DownloadPagesDeploymentPackage(
ctx context.Context,
deploymentID uint,
dst io.Writer,
maxBytes int64,
) (int64, error) {
for projectID, depID := range f.pagesLatestDeployIDs {
if depID == deploymentID {
return f.projectPackage(projectID)
return f.writeProjectPackage(projectID, dst, maxBytes)
}
}
return f.projectPackage(deploymentID)
return f.writeProjectPackage(deploymentID, dst, maxBytes)
}
func (f *fakeClient) GetPagesProjectLatestHash(ctx context.Context, projectID uint) (*protocol.PagesProjectLatestHashResponse, error) {
f.hashCalls++
if f.pagesMetadata != nil {
if metadata, ok := f.pagesMetadata[projectID]; ok {
result := metadata
return &result, nil
}
}
hash, err := f.projectHash(projectID)
if err != nil {
return nil, err
@@ -119,15 +137,32 @@ func (f *fakeClient) GetPagesProjectLatestHash(ctx context.Context, projectID ui
deploymentID = id
}
}
packageBytes, err := f.projectPackage(projectID)
if err != nil {
return nil, err
}
fileCount, totalSize, err := testPagesPackageStats(packageBytes)
if err != nil {
return nil, err
}
return &protocol.PagesProjectLatestHashResponse{
ProjectID: projectID,
DeploymentID: deploymentID,
Hash: hash,
PackageSize: int64(len(packageBytes)),
FileCount: fileCount,
TotalSize: totalSize,
}, nil
}
func (f *fakeClient) DownloadPagesProjectLatestPackage(ctx context.Context, projectID uint) ([]byte, error) {
return f.projectPackage(projectID)
func (f *fakeClient) DownloadPagesProjectLatestPackage(
ctx context.Context,
projectID uint,
dst io.Writer,
maxBytes int64,
) (int64, error) {
f.pagesPackageDownloads++
return f.writeProjectPackage(projectID, dst, maxBytes)
}
func (f *fakeClient) projectHash(projectID uint) (string, error) {
@@ -155,6 +190,22 @@ func (f *fakeClient) projectPackage(projectID uint) ([]byte, error) {
return packageBytes, nil
}
func (f *fakeClient) writeProjectPackage(projectID uint, dst io.Writer, maxBytes int64) (int64, error) {
packageBytes, err := f.projectPackage(projectID)
if err != nil {
return 0, err
}
limited := &io.LimitedReader{R: bytes.NewReader(packageBytes), N: maxBytes + 1}
written, err := io.Copy(dst, limited)
if err != nil {
return written, err
}
if written > maxBytes {
return written, fmt.Errorf("test pages package exceeds limit %d", maxBytes)
}
return written, nil
}
func (f *fakeClient) ReportApplyLog(ctx context.Context, payload protocol.ApplyLogPayload) error {
f.reports = append(f.reports, payload)
return nil
@@ -572,7 +623,9 @@ func TestSyncOnceRejectsPagesZipSlipBeforeApply(t *testing.T) {
service.SetPagesDir(t.TempDir())
err := service.SyncOnce(context.Background(), &protocol.ActiveConfigMeta{Version: "20260309-102", Checksum: "pages-config-checksum"})
if err == nil || (!strings.Contains(err.Error(), "escapes deployment root") && !strings.Contains(err.Error(), "escapes directory")) {
if err == nil || (!strings.Contains(err.Error(), "escapes deployment root") &&
!strings.Contains(err.Error(), "escapes directory") &&
!strings.Contains(err.Error(), "dot segment")) {
t.Fatalf("expected zip-slip rejection, got %v", err)
}
if len(manager.applyRouteContents) != 0 {
@@ -1316,7 +1369,7 @@ func TestSyncOnceRedownloadsPagesDeploymentWhenServerHashChanges(t *testing.T) {
}
pagesDir := t.TempDir()
releaseDir := pagesProjectReleaseDir(pagesDir, projectID, initialHash)
if err = extractPagesPackage(initialPackage, releaseDir, pagesProjectRef{
if err = extractTestPagesPackage(t, initialPackage, releaseDir, pagesProjectRef{
ProjectID: projectID,
Checksum: initialHash,
}); err != nil {
@@ -1385,20 +1438,42 @@ func (r *racingLatestClient) GetPagesProjectLatestHash(ctx context.Context, proj
// 2: verify after downloading B → B (race)
// 3+: stable on B for retry
hash, dep := r.hashA, uint(1)
packageBytes := r.pkgA
if r.hashCall >= 2 {
hash, dep = r.hashB, 2
packageBytes = r.pkgB
}
fileCount, totalSize, err := testPagesPackageStats(packageBytes)
if err != nil {
return nil, err
}
return &protocol.PagesProjectLatestHashResponse{
ProjectID: projectID,
DeploymentID: dep,
Hash: hash,
PackageSize: int64(len(packageBytes)),
FileCount: fileCount,
TotalSize: totalSize,
}, nil
}
func (r *racingLatestClient) DownloadPagesProjectLatestPackage(ctx context.Context, projectID uint) ([]byte, error) {
func (r *racingLatestClient) DownloadPagesProjectLatestPackage(
ctx context.Context,
projectID uint,
dst io.Writer,
maxBytes int64,
) (int64, error) {
r.downloadCalls++
// Always return package B (what "latest download" would stream mid-race / after).
return r.pkgB, nil
limited := &io.LimitedReader{R: bytes.NewReader(r.pkgB), N: maxBytes + 1}
written, err := io.Copy(dst, limited)
if err != nil {
return written, err
}
if written > maxBytes {
return written, fmt.Errorf("test pages package exceeds limit %d", maxBytes)
}
return written, nil
}
func TestEnsurePagesProjectSurvivesHashPackageRace(t *testing.T) {
@@ -1641,6 +1716,41 @@ func testPagesPackage(t *testing.T, files map[string]string) []byte {
return buffer.Bytes()
}
func extractTestPagesPackage(
t *testing.T,
packageBytes []byte,
releaseDir string,
project pagesProjectRef,
) error {
t.Helper()
packagePath := filepath.Join(t.TempDir(), "pages-package.zip")
if err := os.WriteFile(packagePath, packageBytes, pagesFilePerm); err != nil {
t.Fatalf("write test Pages package error = %v", err)
}
return extractPagesPackageFile(packagePath, releaseDir, project, pagesarchive.Limits{
MaxFiles: agentPagesMaxFiles,
MaxFileBytes: agentPagesMaxFileBytes,
MaxTotalBytes: agentPagesMaxTotalBytes,
}, nil)
}
func testPagesPackageStats(packageBytes []byte) (int, int64, error) {
reader, err := zip.NewReader(bytes.NewReader(packageBytes), int64(len(packageBytes)))
if err != nil {
return 0, 0, err
}
fileCount := 0
totalSize := int64(0)
for _, file := range reader.File {
if file.FileInfo().IsDir() {
continue
}
fileCount++
totalSize += int64(file.UncompressedSize64) //nolint:gosec // test packages are memory-bounded
}
return fileCount, totalSize, nil
}
func testBytesChecksum(data []byte) string {
sum := sha256.Sum256(data)
return hex.EncodeToString(sum[:])
+6 -3
View File
@@ -224,14 +224,17 @@ func GetPagesProjectLatestHashHandler(c *gin.Context) {
if !ok {
return
}
deploymentID, hash, err := pages.GetProjectLatestPackageHash(c.Request.Context(), projectID)
metadata, err := pages.GetProjectLatestPackageMetadata(c.Request.Context(), projectID)
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(protocol.PagesProjectLatestHashResponse{
ProjectID: projectID,
DeploymentID: deploymentID,
Hash: hash,
DeploymentID: metadata.DeploymentID,
Hash: metadata.Hash,
PackageSize: metadata.PackageSize,
FileCount: metadata.FileCount,
TotalSize: metadata.TotalSize,
}))
}
@@ -7,9 +7,11 @@ import (
"context"
"errors"
"fmt"
"path"
"strings"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/pkg/pagesarchive"
openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty"
"gorm.io/gorm"
)
@@ -58,23 +60,41 @@ func buildPagesRouteSnapshot(
}
pagesProjectID = route.PagesProjectID
deployment = buildSnapshotPagesDeployment(project, activeDeployment)
deployment, err = buildSnapshotPagesDeployment(project, activeDeployment)
if err != nil {
return "", nil, nil, nil, fmt.Errorf("路由 %s Pages 配置无效: %w", route.SiteName, err)
}
originURL = fmt.Sprintf("openflare-pages://project/%d", project.ID)
return originURL, []string{originURL}, pagesProjectID, deployment, nil
}
func buildSnapshotPagesDeployment(project *model.PagesProject, activeDeployment *model.PagesDeployment) *openrestyrender.PagesDeployment {
func buildSnapshotPagesDeployment(
project *model.PagesProject,
activeDeployment *model.PagesDeployment,
) (*openrestyrender.PagesDeployment, error) {
if project == nil || activeDeployment == nil {
return nil
return nil, errors.New("pages 项目或部署为空")
}
rootDir, err := pagesarchive.NormalizeLogicalPath(strings.TrimSpace(project.RootDir), true)
if err != nil {
return nil, fmt.Errorf("pages 根目录不合法: %w", err)
}
entryFile := strings.TrimSpace(project.EntryFile)
if entryFile == "" {
entryFile = defaultPagesSnapshotEntryFile
}
entryFile, err = pagesarchive.NormalizeLogicalPath(entryFile, false)
if err != nil {
return nil, fmt.Errorf("pages 入口文件不合法: %w", err)
}
fallbackPath := strings.TrimSpace(project.SPAFallbackPath)
if fallbackPath == "" {
fallbackPath = defaultPagesSnapshotFallbackPath
}
localRoot := openrestyrender.PagesProjectLocalRoot(project.ID)
if rootDir != "" {
localRoot = path.Join(localRoot, rootDir)
}
return &openrestyrender.PagesDeployment{
ProjectID: project.ID,
ProjectSlug: strings.TrimSpace(project.Slug),
@@ -90,6 +110,6 @@ func buildSnapshotPagesDeployment(project *model.PagesProject, activeDeployment
APIProxyRewrite: strings.TrimSpace(project.APIProxyRewrite),
// Root is project-scoped so Agents can swap active packages without
// re-publishing main config (nginx root stays stable).
LocalRoot: openrestyrender.PagesProjectLocalRoot(project.ID),
}
LocalRoot: localRoot,
}, nil
}
@@ -28,6 +28,7 @@ func TestBuildSnapshotRoutesPages(t *testing.T) {
Enabled: true,
SPAFallbackEnabled: true,
SPAFallbackPath: "/index.html",
RootDir: "public/site",
EntryFile: "index.html",
}
require.NoError(t, conn.Create(project).Error)
@@ -64,7 +65,7 @@ func TestBuildSnapshotRoutesPages(t *testing.T) {
require.NotNil(t, snapshotRoute.PagesDeployment)
assert.Equal(t, deployment.ID, snapshotRoute.PagesDeployment.DeploymentID)
assert.Equal(t, deployment.Checksum, snapshotRoute.PagesDeployment.Checksum)
assert.Equal(t, "__OPENFLARE_PAGES_DIR__/projects/1/current", snapshotRoute.PagesDeployment.LocalRoot)
assert.Equal(t, "__OPENFLARE_PAGES_DIR__/projects/1/current/public/site", snapshotRoute.PagesDeployment.LocalRoot)
_, err = renderSnapshotConfig(bundle.SnapshotJSON, nil)
require.NoError(t, err)
@@ -78,6 +79,17 @@ func TestBuildSnapshotRoutesPages(t *testing.T) {
require.NotNil(t, decoded.Routes[0].PagesDeployment)
}
func TestBuildSnapshotPagesDeploymentRejectsUnsafeStoredPaths(t *testing.T) {
deployment := &model.PagesDeployment{ID: 1, ProjectID: 1, Checksum: "checksum"}
for _, project := range []*model.PagesProject{
{ID: 1, RootDir: "../escape", EntryFile: "index.html"},
{ID: 1, RootDir: "public", EntryFile: "/index.html"},
} {
_, err := buildSnapshotPagesDeployment(project, deployment)
require.Error(t, err)
}
}
func requireDB(t *testing.T, ctx context.Context) *gorm.DB {
t.Helper()
conn := db.DB(ctx)
@@ -27,7 +27,7 @@ import (
const (
pagesURLDownloadTimeout = 10 * time.Minute
pagesURLMaxRedirects = 5
pagesMagicSniffBytes = 16
pagesMagicSniffBytes = 512
pagesURLDialTimeout = 30 * time.Second
pagesURLTLSHandshake = 15 * time.Second
pagesBrowserUserAgent = "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/131.0.0.0 Safari/537.36"
+2
View File
@@ -34,4 +34,6 @@ const (
errPagesPackageNotInActiveConfig = "pages 部署尚未进入激活配置"
errPagesDeploymentHashMissing = "pages 部署包哈希缺失"
errPagesInvalidSnapshotFormat = "配置快照格式无效"
errPagesActorMissing = "无法识别当前用户"
errPagesEntryFileMissing = "当前激活部署中不存在指定入口文件"
)
+80 -35
View File
@@ -15,13 +15,17 @@ import (
"path"
"path/filepath"
"regexp"
"strconv"
"strings"
"github.com/Rain-kl/Wavelet/internal/apps/upload"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/Rain-kl/Wavelet/pkg/pagesarchive"
"gorm.io/gorm"
"gorm.io/gorm/clause"
)
const (
@@ -31,11 +35,14 @@ const (
defaultPagesMaxHistoryCount = 20
defaultPagesEntryFile = "index.html"
defaultPagesFallbackPath = "/index.html"
pagesDeploymentUploadType = "openflare_pages_deployment"
pagesIngestMarkerKey = "pages_ingest_marker"
pagesIngestMarkerV2 = "pages_deployment_v2"
pagesProjectIDMetadataKey = "pages_project_id"
pagesMaxPathLength = 512
bytesPerMiB = 1024 * 1024
pagesExtractedSizeMultiplier = 4
pagesMinExtractedSizeBytes = 100 * bytesPerMiB
pagesRowLockStrength = "UPDATE"
)
var pagesSlugPattern = regexp.MustCompile(`^[a-z0-9][a-z0-9-]{0,126}[a-z0-9]$|^[a-z0-9]$`)
@@ -115,30 +122,14 @@ func normalizePagesSlug(raw string) string {
func validateAndNormalizePagesRootDir(raw string) (string, error) {
value := strings.TrimSpace(raw)
if value == "" {
return "", nil
}
if len(value) > pagesMaxPathLength {
return "", errors.New("pages 根目录长度不能超过 512")
}
if strings.Contains(value, "\\") || strings.ContainsAny(value, "\"';") {
return "", errors.New("pages 根目录包含不支持的字符")
normalized, err := pagesarchive.NormalizeLogicalPath(value, true)
if err != nil {
return "", fmt.Errorf("pages 根目录不合法: %w", err)
}
for _, r := range value {
if r <= 0x20 || r == 0x7f {
return "", errors.New("pages 根目录不能包含空白或控制字符")
}
}
cleaned := path.Clean(filepath.ToSlash(value))
if cleaned == "." || cleaned == "/" {
return "", nil
}
for _, segment := range strings.Split(cleaned, "/") {
if segment == "." || segment == ".." {
return "", errors.New("pages 根目录不能包含 . 或 .. 路径段")
}
}
return strings.TrimPrefix(cleaned, "/"), nil
return normalized, nil
}
func normalizePagesFallbackPath(raw string) (string, error) {
@@ -186,12 +177,19 @@ func normalizeStoredPagesFallbackPath(value string) string {
return normalized
}
func normalizePagesEntryFile(raw string) string {
value := path.Clean(strings.TrimSpace(filepath.ToSlash(raw)))
if value == "." || value == "/" {
return defaultPagesEntryFile
func validateAndNormalizePagesEntryFile(raw string) (string, error) {
value := strings.TrimSpace(raw)
if value == "" {
value = defaultPagesEntryFile
}
return strings.TrimPrefix(value, "/")
if len(value) > pagesMaxPathLength {
return "", errors.New("pages 入口文件长度不能超过 512")
}
normalized, err := pagesarchive.NormalizeLogicalPath(value, false)
if err != nil {
return "", fmt.Errorf("pages 入口文件不合法: %w", err)
}
return normalized, nil
}
func persistPagesUploadTemp(fileHeader *multipart.FileHeader, maxPackageBytes int64) (string, string, int64, pagesarchive.Format, error) {
@@ -262,7 +260,7 @@ func ingestPagesDeploymentPackage(
ctx context.Context,
localPath string,
checksum string,
projectSlug string,
projectID uint,
fileName string,
format pagesarchive.Format,
) (upload.IngestResult, error) {
@@ -275,28 +273,27 @@ func ingestPagesDeploymentPackage(
MimeType: pagesarchive.MIMEType(format),
Extension: extension,
Hash: checksum,
Type: pagesDeploymentUploadType,
Type: upload.ReservedPagesDeploymentType,
AccessMode: &accessMode,
SkipExtensionCheck: true,
Policy: upload.PolicyDedupNewRecord,
Metadata: model.UploadMetadata{
Extra: map[string]any{
"project_slug": projectSlug,
"format": string(format),
pagesIngestMarkerKey: pagesIngestMarkerV2,
pagesProjectIDMetadataKey: strconv.FormatUint(uint64(projectID), 10),
},
},
})
}
func removeDeploymentArtifact(ctx context.Context, deployment *model.PagesDeployment) {
func removeDeploymentArtifact(ctx context.Context, projectID uint, deployment *model.PagesDeployment) {
if deployment == nil {
return
}
if deployment.UploadID == 0 {
return
}
if _, err := upload.Remove(ctx, deployment.UploadID); err != nil {
// Soft-delete / storage cleanup failure must not undo DB prune; log for ops.
if err := removePagesUploadIfUnreferenced(ctx, projectID, deployment.UploadID); err != nil {
logger.WarnF(ctx,
"[Pages] remove deployment artifact failed: deployment_id=%d upload_id=%d error=%v",
deployment.ID, deployment.UploadID, err,
@@ -304,10 +301,58 @@ func removeDeploymentArtifact(ctx context.Context, deployment *model.PagesDeploy
}
}
// removePagesUploadIfUnreferenced soft-deletes a reserved Pages upload only
// after locking its project (when present), locking the upload, and rechecking
// deployment references in the same transaction.
func removePagesUploadIfUnreferenced(ctx context.Context, projectID uint, uploadID uint64) error {
if uploadID == 0 {
return nil
}
err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
if projectID != 0 {
var project model.PagesProject
projectErr := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&project, projectID).Error
if projectErr != nil && !errors.Is(projectErr, gorm.ErrRecordNotFound) {
return projectErr
}
}
var uploadRecord model.Upload
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).
Where("id = ?", uploadID).
First(&uploadRecord).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil
}
return err
}
if uploadRecord.Type != upload.ReservedPagesDeploymentType {
return fmt.Errorf("pages 部署包上传类型不匹配: %s", uploadRecord.Type)
}
var references int64
if err := tx.Model(&model.PagesDeployment{}).
Where("upload_id = ?", uploadID).
Count(&references).Error; err != nil {
return err
}
if references > 0 {
return nil
}
_, err := upload.RemoveLockedTx(tx, &uploadRecord)
return err
})
// Always invalidate after transaction completion, including idempotent no-op,
// so a prior post-commit cache interruption can heal on retry.
upload.InvalidateUploadMetaCache(ctx, uploadID)
return err
}
func inspectPagesPackage(packagePath string, format pagesarchive.Format, rootDir string, entryFile string, limits pagesLimits) (*deploymentManifest, error) {
archiveManifest, err := pagesarchive.InspectFile(packagePath, format, pagesarchive.InspectOptions{
RootDir: rootDir,
EntryFile: entryFile,
RootDir: rootDir,
EntryFile: entryFile,
VerifySizes: true,
Limits: pagesarchive.Limits{
MaxFiles: limits.MaxFiles,
MaxFileBytes: limits.ExtractedBytes,
+315 -117
View File
@@ -12,6 +12,7 @@ import (
"mime/multipart"
"net/url"
"os"
"path"
"strings"
"time"
@@ -21,6 +22,7 @@ import (
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/Rain-kl/Wavelet/pkg/pagesarchive"
"gorm.io/gorm"
"gorm.io/gorm/clause"
)
// DeploymentPackage is a streamable Pages deployment artifact for agent download.
@@ -137,28 +139,40 @@ func CreateProject(ctx context.Context, input Input) (*View, error) {
// UpdateProject 更新 Pages 项目。
func UpdateProject(ctx context.Context, id uint, input Input) (*View, error) {
project, err := model.GetPagesProjectByID(ctx, id)
var project *model.PagesProject
err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
var existing model.PagesProject
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&existing, id).Error; err != nil {
return err
}
updated := existing
var err error
project, err = buildProject(&updated, input)
if err != nil {
return err
}
if (existing.RootDir != project.RootDir || existing.EntryFile != project.EntryFile) &&
existing.ActiveDeploymentID != nil && *existing.ActiveDeploymentID != 0 {
if err := ensureDeploymentEntry(tx, *existing.ActiveDeploymentID, project.RootDir, project.EntryFile); err != nil {
return err
}
}
return tx.Model(&existing).Updates(map[string]any{
"name": project.Name,
"slug": project.Slug,
"description": project.Description,
"enabled": project.Enabled,
"spa_fallback_enabled": project.SPAFallbackEnabled,
"spa_fallback_path": project.SPAFallbackPath,
"api_proxy_enabled": project.APIProxyEnabled,
"api_proxy_path": project.APIProxyPath,
"api_proxy_pass": project.APIProxyPass,
"api_proxy_rewrite": project.APIProxyRewrite,
"root_dir": project.RootDir,
"entry_file": project.EntryFile,
}).Error
})
if err != nil {
return nil, err
}
project, err = buildProject(project, input)
if err != nil {
return nil, err
}
if err = db.DB(ctx).Model(project).Updates(map[string]any{
"name": project.Name,
"slug": project.Slug,
"description": project.Description,
"enabled": project.Enabled,
"spa_fallback_enabled": project.SPAFallbackEnabled,
"spa_fallback_path": project.SPAFallbackPath,
"api_proxy_enabled": project.APIProxyEnabled,
"api_proxy_path": project.APIProxyPath,
"api_proxy_pass": project.APIProxyPass,
"api_proxy_rewrite": project.APIProxyRewrite,
"root_dir": project.RootDir,
"entry_file": project.EntryFile,
}).Error; err != nil {
if isUniqueConstraintError(err) {
return nil, errors.New(errPagesSlugExists)
}
@@ -167,24 +181,49 @@ func UpdateProject(ctx context.Context, id uint, input Input) (*View, error) {
return buildProjectView(ctx, project)
}
func ensureDeploymentEntry(conn *gorm.DB, deploymentID uint, rootDir, entryFile string) error {
targetPath := entryFile
if rootDir != "" {
targetPath = path.Join(rootDir, entryFile)
}
var count int64
if err := conn.Model(&model.PagesDeploymentFile{}).
Where("deployment_id = ? AND path = ?", deploymentID, targetPath).
Count(&count).Error; err != nil {
return err
}
if count == 0 {
return fmt.Errorf("%s: %s", errPagesEntryFileMissing, targetPath)
}
return nil
}
// DeleteProject 删除 Pages 项目。
func DeleteProject(ctx context.Context, id uint) error {
project, err := model.GetPagesProjectByID(ctx, id)
if err != nil {
return err
}
routeCount, err := model.CountProxyRoutesByPagesProjectID(ctx, project.ID)
if err != nil {
return err
}
if routeCount > 0 {
return errors.New(errPagesDeleteReferenced)
}
deployments, err := model.ListPagesDeployments(ctx, project.ID)
if err != nil {
return err
}
return db.DB(ctx).Transaction(func(tx *gorm.DB) error {
var deployments []model.PagesDeployment
err = db.DB(ctx).Transaction(func(tx *gorm.DB) error {
var lockedProject model.PagesProject
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&lockedProject, project.ID).Error; err != nil {
return err
}
if tx.Migrator().HasTable(&model.ProxyRoute{}) {
var routeCount int64
if err := tx.Model(&model.ProxyRoute{}).
Where("pages_project_id = ?", project.ID).
Count(&routeCount).Error; err != nil {
return err
}
if routeCount > 0 {
return errors.New(errPagesDeleteReferenced)
}
}
if err := tx.Where("project_id = ?", project.ID).Find(&deployments).Error; err != nil {
return err
}
if err := tx.Where(
"deployment_id IN (?)",
tx.Model(&model.PagesDeployment{}).Select("id").Where("project_id = ?", project.ID),
@@ -197,11 +236,15 @@ func DeleteProject(ctx context.Context, id uint) error {
if err := tx.Delete(project).Error; err != nil {
return err
}
for index := range deployments {
removeDeploymentArtifact(ctx, &deployments[index])
}
return nil
})
if err != nil {
return err
}
for index := range deployments {
removeDeploymentArtifact(ctx, project.ID, &deployments[index])
}
return nil
}
// ListProjectDeployments 列出项目的全部部署。
@@ -298,7 +341,10 @@ func createDeploymentFromTempPackage(
if err != nil {
return nil, err
}
entryFile := normalizePagesEntryFile(project.EntryFile)
entryFile, err := validateAndNormalizePagesEntryFile(project.EntryFile)
if err != nil {
return nil, err
}
manifest, err := inspectPagesPackage(tempPath, format, rootDir, entryFile, limits)
if err != nil {
return nil, err
@@ -307,7 +353,7 @@ func createDeploymentFromTempPackage(
ctx,
tempPath,
checksum,
project.Slug,
project.ID,
fileName,
format,
)
@@ -317,11 +363,20 @@ func createDeploymentFromTempPackage(
ingestCommitted := false
defer func() {
if !ingestCommitted && ingestResult.Created {
_, _ = upload.Remove(ctx, ingestResult.Upload.ID)
if removeErr := removePagesUploadIfUnreferenced(ctx, project.ID, ingestResult.Upload.ID); removeErr != nil {
logger.ErrorF(ctx,
"[Pages] compensate deployment upload failed: project_id=%d upload_id=%d error=%v",
project.ID, ingestResult.Upload.ID, removeErr,
)
}
}
}()
deployment := &model.PagesDeployment{}
err = db.DB(ctx).Transaction(func(tx *gorm.DB) error {
var lockedProject model.PagesProject
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&lockedProject, project.ID).Error; err != nil {
return err
}
var maxNumber int
if err := tx.Model(&model.PagesDeployment{}).
Where("project_id = ?", project.ID).
@@ -357,7 +412,7 @@ func createDeploymentFromTempPackage(
}
ingestCommitted = true
if pruneErr := pruneProjectDeploymentHistory(ctx, project.ID, limits.HistoryCount); pruneErr != nil {
if pruneErr := pruneProjectDeploymentHistory(ctx, project.ID, limits.HistoryCount, deployment.ID); pruneErr != nil {
logger.ErrorF(ctx,
"[Pages] prune deployment history failed: project_id=%d keep=%d error=%v",
project.ID, limits.HistoryCount, pruneErr,
@@ -381,7 +436,7 @@ func createDeploymentFromTempPackage(
// Concurrency: DB row deletes run in a single transaction after a consistent read
// of project + deployments. Concurrent uploads may briefly exceed keepCount; the
// next successful prune brings the project back within the limit (eventual).
func pruneProjectDeploymentHistory(ctx context.Context, projectID uint, keepCount int) error {
func pruneProjectDeploymentHistory(ctx context.Context, projectID uint, keepCount int, preserveCandidateID uint) error {
if keepCount <= 0 {
return nil
}
@@ -390,7 +445,7 @@ func pruneProjectDeploymentHistory(ctx context.Context, projectID uint, keepCoun
// that inserted another deployment between our list and delete.
var lastErr error
for pass := 0; pass < 2; pass++ {
deleted, err := pruneProjectDeploymentHistoryOnce(ctx, projectID, keepCount)
deleted, err := pruneProjectDeploymentHistoryOnce(ctx, projectID, keepCount, preserveCandidateID)
if err != nil {
lastErr = err
break
@@ -404,73 +459,79 @@ func pruneProjectDeploymentHistory(ctx context.Context, projectID uint, keepCoun
// pruneProjectDeploymentHistoryOnce performs one list → select → delete cycle.
// Returns the number of deployments deleted from the database.
func pruneProjectDeploymentHistoryOnce(ctx context.Context, projectID uint, keepCount int) (int, error) {
project, err := model.GetPagesProjectByID(ctx, projectID)
if err != nil {
return 0, fmt.Errorf("load pages project: %w", err)
}
deployments, err := model.ListPagesDeployments(ctx, projectID)
if err != nil {
return 0, fmt.Errorf("list pages deployments: %w", err)
}
if len(deployments) <= keepCount {
return 0, nil
}
var activeID uint
if project.ActiveDeploymentID != nil {
activeID = *project.ActiveDeploymentID
}
toDelete := selectDeploymentsToPrune(deployments, activeID, keepCount)
if len(toDelete) == 0 {
return 0, nil
}
// Delete metadata in one transaction so partial prune does not leave
// orphan file-list rows without a parent deployment.
if err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
func pruneProjectDeploymentHistoryOnce(ctx context.Context, projectID uint, keepCount int, preserveCandidateID uint) (int, error) {
var deletedDeployments []model.PagesDeployment
err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
var project model.PagesProject
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&project, projectID).Error; err != nil {
return fmt.Errorf("load pages project: %w", err)
}
var deployments []model.PagesDeployment
if err := tx.Where("project_id = ?", projectID).Order("id desc").Find(&deployments).Error; err != nil {
return fmt.Errorf("list pages deployments: %w", err)
}
var activeID uint
if project.ActiveDeploymentID != nil {
activeID = *project.ActiveDeploymentID
}
// Preserve mode is signaled by a non-zero candidate ID, but the lock-time
// newest non-active deployment wins. This prevents concurrent upload A/B
// prune passes from deleting each other's newer candidate.
resolvedCandidateID := resolveLatestCandidateID(deployments, activeID, preserveCandidateID != 0)
toDelete := selectDeploymentsToPrune(deployments, activeID, resolvedCandidateID, keepCount)
for index := range toDelete {
deployment := toDelete[index]
// Never delete the active deployment even if project pointer raced.
if activeID != 0 && deployment.ID == activeID {
continue
}
if project.ActiveDeploymentID != nil && deployment.ID == *project.ActiveDeploymentID {
continue
}
if err := tx.Where("deployment_id = ?", deployment.ID).Delete(&model.PagesDeploymentFile{}).Error; err != nil {
return fmt.Errorf("delete deployment files id=%d: %w", deployment.ID, err)
}
if err := tx.Where("id = ? AND project_id = ?", deployment.ID, projectID).
Delete(&model.PagesDeployment{}).Error; err != nil {
return fmt.Errorf("delete deployment id=%d: %w", deployment.ID, err)
result := tx.Where("id = ? AND project_id = ?", deployment.ID, projectID).
Delete(&model.PagesDeployment{})
if result.Error != nil {
return fmt.Errorf("delete deployment id=%d: %w", deployment.ID, result.Error)
}
if result.RowsAffected == 1 {
deletedDeployments = append(deletedDeployments, deployment)
}
}
return nil
}); err != nil {
})
if err != nil {
return 0, err
}
// Artifacts are best-effort outside the transaction (object storage I/O).
for index := range toDelete {
deployment := toDelete[index]
if activeID != 0 && deployment.ID == activeID {
continue
}
removeDeploymentArtifact(ctx, &deployment)
for index := range deletedDeployments {
deployment := deletedDeployments[index]
removeDeploymentArtifact(ctx, projectID, &deployment)
}
logger.InfoF(ctx,
"[Pages] pruned deployment history: project_id=%d keep=%d deleted=%d",
projectID, keepCount, len(toDelete),
projectID, keepCount, len(deletedDeployments),
)
return len(toDelete), nil
return len(deletedDeployments), nil
}
func resolveLatestCandidateID(deployments []model.PagesDeployment, activeID uint, preserve bool) uint {
if !preserve {
return 0
}
for _, deployment := range deployments {
if deployment.ID != activeID {
return deployment.ID
}
}
return 0
}
// selectDeploymentsToPrune returns deployments that should be removed under the
// "at most keepCount, always keep active, fill with newest" policy.
// deployments must be ordered newest-first (id desc).
func selectDeploymentsToPrune(deployments []model.PagesDeployment, activeID uint, keepCount int) []model.PagesDeployment {
func selectDeploymentsToPrune(
deployments []model.PagesDeployment,
activeID uint,
preserveCandidateID uint,
keepCount int,
) []model.PagesDeployment {
if keepCount <= 0 || len(deployments) <= keepCount {
return nil
}
@@ -486,6 +547,17 @@ func selectDeploymentsToPrune(deployments []model.PagesDeployment, activeID uint
}
}
}
// A freshly uploaded manual candidate is temporarily protected in addition
// to the active deployment. This intentionally permits two rows when the
// configured history limit is one.
if preserveCandidateID != 0 {
for _, deployment := range deployments {
if deployment.ID == preserveCandidateID {
keepIDs[preserveCandidateID] = struct{}{}
break
}
}
}
// 2) Fill remaining slots from newest to oldest.
for _, deployment := range deployments {
if len(keepIDs) >= keepCount {
@@ -521,26 +593,69 @@ func ActivateDeployment(ctx context.Context, projectID uint, deploymentID uint)
if deployment.ProjectID != project.ID {
return nil, errors.New(errPagesDeploymentMismatch)
}
if deployment.UploadID == 0 {
if err = ensureDeploymentUploadRecord(ctx, deployment); err != nil {
return nil, err
}
}
now := time.Now()
if err = db.DB(ctx).Transaction(func(tx *gorm.DB) error {
var project model.PagesProject
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&project, projectID).Error; err != nil {
return err
}
var deployment model.PagesDeployment
if err := tx.First(&deployment, deploymentID).Error; err != nil {
return err
}
if deployment.ProjectID != project.ID {
return errors.New(errPagesDeploymentMismatch)
}
rootDir, err := validateAndNormalizePagesRootDir(project.RootDir)
if err != nil {
return err
}
entryFile, err := validateAndNormalizePagesEntryFile(project.EntryFile)
if err != nil {
return err
}
if err := ensureDeploymentEntry(tx, deployment.ID, rootDir, entryFile); err != nil {
return err
}
var uploadRecord model.Upload
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).
Where("id = ?", deployment.UploadID).
First(&uploadRecord).Error; err != nil {
return errors.New(errPagesPackageUploadMissing)
}
if uploadRecord.Status != model.UploadStatusUsed || uploadRecord.Type != upload.ReservedPagesDeploymentType {
return errors.New(errPagesPackageUploadMissing)
}
if err := tx.Model(&model.PagesDeployment{}).
Where("project_id = ?", project.ID).
Update("status", model.PagesDeploymentStatusUploaded).Error; err != nil {
return err
}
if err := tx.Model(deployment).Updates(map[string]any{
if err := tx.Model(&deployment).Updates(map[string]any{
"status": model.PagesDeploymentStatusActive,
"activated_at": &now,
}).Error; err != nil {
return err
}
return tx.Model(project).Updates(map[string]any{
return tx.Model(&project).Updates(map[string]any{
"active_deployment_id": deployment.ID,
}).Error
}); err != nil {
return nil, err
}
return GetProject(ctx, project.ID)
limits := resolvePagesLimits(ctx)
if pruneErr := pruneProjectDeploymentHistory(ctx, projectID, limits.HistoryCount, 0); pruneErr != nil {
logger.ErrorF(ctx,
"[Pages] strict prune after activation failed: project_id=%d keep=%d error=%v",
projectID, limits.HistoryCount, pruneErr,
)
}
return GetProject(ctx, projectID)
}
// GetDeploymentPackageHash returns the upload SHA-256 hash of the deployment package.
@@ -728,22 +843,97 @@ func hydrateLegacyDeploymentUpload(
ctx,
artifactPath,
deployment.Checksum,
project.Slug,
project.ID,
fmt.Sprintf("pages-deployment-%d.zip", deployment.ID),
pagesarchive.FormatZip,
)
if err != nil {
return nil, err
}
if err := db.DB(ctx).Model(deployment).Updates(map[string]any{
"upload_id": ingestResult.Upload.ID,
"artifact_path": "",
}).Error; err != nil {
winnerUploadID, err := attachLegacyDeploymentUpload(
ctx,
project.ID,
deployment.ID,
ingestResult.Upload.ID,
)
if ingestResult.Created && (err != nil || winnerUploadID != ingestResult.Upload.ID) {
if removeErr := removePagesUploadIfUnreferenced(ctx, project.ID, ingestResult.Upload.ID); removeErr != nil {
logger.ErrorF(ctx,
"[Pages] compensate legacy deployment upload failed: project_id=%d upload_id=%d error=%v",
project.ID, ingestResult.Upload.ID, removeErr,
)
}
}
if err != nil {
return nil, err
}
deployment.UploadID = ingestResult.Upload.ID
winner, err := upload.GetActiveUpload(ctx, winnerUploadID)
if err != nil {
return nil, err
}
deployment.UploadID = winnerUploadID
deployment.ArtifactPath = ""
return &ingestResult.Upload, nil
return &winner, nil
}
func attachLegacyDeploymentUpload(
ctx context.Context,
projectID uint,
deploymentID uint,
uploadID uint64,
) (uint64, error) {
winnerUploadID := uint64(0)
err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
var err error
winnerUploadID, err = attachLegacyDeploymentUploadTx(tx, projectID, deploymentID, uploadID)
return err
})
return winnerUploadID, err
}
func attachLegacyDeploymentUploadTx(
tx *gorm.DB,
projectID uint,
deploymentID uint,
uploadID uint64,
) (uint64, error) {
var lockedProject model.PagesProject
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&lockedProject, projectID).Error; err != nil {
return 0, err
}
var lockedDeployment model.PagesDeployment
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&lockedDeployment, deploymentID).Error; err != nil {
return 0, err
}
if lockedDeployment.ProjectID != lockedProject.ID {
return 0, errors.New(errPagesDeploymentMismatch)
}
if lockedDeployment.UploadID != 0 {
return lockedDeployment.UploadID, nil
}
var uploadRecord model.Upload
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).
Where("id = ?", uploadID).
First(&uploadRecord).Error; err != nil {
return 0, err
}
if uploadRecord.Status != model.UploadStatusUsed || uploadRecord.Type != upload.ReservedPagesDeploymentType {
return 0, errors.New(errPagesPackageUploadMissing)
}
result := tx.Model(&model.PagesDeployment{}).
Where("id = ? AND project_id = ? AND upload_id = 0", lockedDeployment.ID, lockedProject.ID).
Updates(map[string]any{
"upload_id": uploadRecord.ID,
"artifact_path": "",
})
if result.Error != nil {
return 0, result.Error
}
if result.RowsAffected != 1 {
return 0, errors.New(errPagesPackageUploadMissing)
}
return uploadRecord.ID, nil
}
// ensureDeploymentInActiveSnapshot allows download of a specific deployment when
@@ -842,30 +1032,34 @@ func parseSnapshotRoutes(snapshotJSON string) ([]snapshotRouteRef, error) {
// DeleteDeployment 删除 Pages 部署。
func DeleteDeployment(ctx context.Context, projectID uint, deploymentID uint) error {
project, err := model.GetPagesProjectByID(ctx, projectID)
if err != nil {
return err
}
deployment, err := model.GetPagesDeploymentByID(ctx, deploymentID)
if err != nil {
return err
}
if deployment.ProjectID != project.ID {
return errors.New(errPagesDeploymentMismatch)
}
if project.ActiveDeploymentID != nil && *project.ActiveDeploymentID == deployment.ID {
return errors.New(errPagesDeleteActiveDeploy)
}
return db.DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Where("deployment_id = ?", deployment.ID).Delete(&model.PagesDeploymentFile{}).Error; err != nil {
var removed model.PagesDeployment
err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
var project model.PagesProject
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&project, projectID).Error; err != nil {
return err
}
if err := tx.Delete(deployment).Error; err != nil {
if err := tx.First(&removed, deploymentID).Error; err != nil {
return err
}
if removed.ProjectID != project.ID {
return errors.New(errPagesDeploymentMismatch)
}
if project.ActiveDeploymentID != nil && *project.ActiveDeploymentID == removed.ID {
return errors.New(errPagesDeleteActiveDeploy)
}
if err := tx.Where("deployment_id = ?", removed.ID).Delete(&model.PagesDeploymentFile{}).Error; err != nil {
return err
}
if err := tx.Delete(&removed).Error; err != nil {
return err
}
removeDeploymentArtifact(ctx, deployment)
return nil
})
if err != nil {
return err
}
removeDeploymentArtifact(ctx, projectID, &removed)
return nil
}
func buildProject(existing *model.PagesProject, input Input) (*model.PagesProject, error) {
@@ -923,7 +1117,11 @@ func buildProject(existing *model.PagesProject, input Input) (*model.PagesProjec
return nil, err
}
existing.RootDir = rootDir
existing.EntryFile = normalizePagesEntryFile(input.EntryFile)
entryFile, err := validateAndNormalizePagesEntryFile(input.EntryFile)
if err != nil {
return nil, err
}
existing.EntryFile = entryFile
return existing, nil
}
+230 -6
View File
@@ -17,6 +17,7 @@ import (
"path/filepath"
"testing"
"github.com/Rain-kl/Wavelet/internal/apps/upload"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
@@ -163,6 +164,89 @@ func TestCreateProjectRejectsUnsafeFallbackPath(t *testing.T) {
assert.Contains(t, err.Error(), "回退路径")
}
func TestCreateProjectRejectsUnsafeContentPaths(t *testing.T) {
cleanup := setupPagesTestDB(t)
defer cleanup()
ctx := context.Background()
rootDirs := []string{"/public", "public/../dist", "C:/public", `public\\dist`, "./public", "public\x00dist"}
for index, rootDir := range rootDirs {
_, err := CreateProject(ctx, Input{
Name: fmt.Sprintf("Unsafe Root %d", index),
Slug: fmt.Sprintf("unsafe-root-%d", index),
RootDir: rootDir,
EntryFile: "index.html",
})
require.Error(t, err, rootDir)
}
entryFiles := []string{"/index.html", "../index.html", "C:/index.html", `public\\index.html`, "./index.html", "index.html;bad"}
for index, entryFile := range entryFiles {
_, err := CreateProject(ctx, Input{
Name: fmt.Sprintf("Unsafe Entry %d", index),
Slug: fmt.Sprintf("unsafe-entry-%d", index),
EntryFile: entryFile,
})
require.Error(t, err, entryFile)
}
}
func TestUpdateProjectValidatesActiveDeploymentEntry(t *testing.T) {
cleanup := setupPagesTestDB(t)
defer cleanup()
_, disableStorage := setupPagesStorageMock(t)
defer disableStorage()
ctx := context.Background()
project, err := CreateProject(ctx, Input{
Name: "Content Root",
Slug: "content-root",
Enabled: true,
EntryFile: "index.html",
})
require.NoError(t, err)
deployment, err := UploadDeployment(ctx, project.ID, testPagesMultipartFile(t, "site.zip", testPagesZip(t, map[string]string{
"index.html": "root",
"dist/index.html": "dist",
})), "user:1")
require.NoError(t, err)
staleCandidate, err := UploadDeployment(ctx, project.ID, testPagesMultipartFile(t, "stale.zip", testPagesZip(t, map[string]string{
"index.html": "stale",
})), "user:1")
require.NoError(t, err)
_, err = ActivateDeployment(ctx, project.ID, deployment.ID)
require.NoError(t, err)
updated, err := UpdateProject(ctx, project.ID, Input{
Name: project.Name,
Slug: project.Slug,
Enabled: true,
RootDir: "dist",
EntryFile: "index.html",
})
require.NoError(t, err)
assert.Equal(t, "dist", updated.RootDir)
_, err = ActivateDeployment(ctx, project.ID, staleCandidate.ID)
require.Error(t, err)
assert.Contains(t, err.Error(), errPagesEntryFileMissing)
_, err = UpdateProject(ctx, project.ID, Input{
Name: project.Name,
Slug: project.Slug,
Enabled: true,
RootDir: "missing",
EntryFile: "index.html",
})
require.Error(t, err)
assert.Contains(t, err.Error(), errPagesEntryFileMissing)
stored, err := model.GetPagesProjectByID(ctx, project.ID)
require.NoError(t, err)
assert.Equal(t, "dist", stored.RootDir)
assert.Equal(t, "index.html", stored.EntryFile)
}
func TestUploadDeploymentAcceptsZeroByteFiles(t *testing.T) {
cleanup := setupPagesTestDB(t)
defer cleanup()
@@ -213,6 +297,13 @@ func TestUploadDeploymentStoresPackageInUploadFramework(t *testing.T) {
var uploadCount int64
require.NoError(t, db.DB(ctx).Model(&model.Upload{}).Count(&uploadCount).Error)
assert.Equal(t, int64(1), uploadCount)
var uploadRecord model.Upload
require.NoError(t, db.DB(ctx).First(&uploadRecord, storedDeployment.UploadID).Error)
assert.Equal(t, upload.ReservedPagesDeploymentType, uploadRecord.Type)
assert.Equal(t, pagesIngestMarkerV2, uploadRecord.Metadata.Extra[pagesIngestMarkerKey])
assert.Equal(t, fmt.Sprint(project.ID), uploadRecord.Metadata.Extra[pagesProjectIDMetadataKey])
assert.NotContains(t, uploadRecord.Metadata.Extra, "project_slug")
assert.NotContains(t, uploadRecord.Metadata.Extra, "format")
}
func TestOpenDeploymentPackageHydratesLegacyArtifactPath(t *testing.T) {
@@ -448,34 +539,42 @@ func TestSelectDeploymentsToPruneKeepsActiveAndNewest(t *testing.T) {
{ID: 2, ProjectID: 1},
{ID: 1, ProjectID: 1},
}
toDelete := selectDeploymentsToPrune(deployments, 1, 2)
toDelete := selectDeploymentsToPrune(deployments, 1, 0, 2)
require.Len(t, toDelete, 2)
assert.Equal(t, uint(3), toDelete[0].ID)
assert.Equal(t, uint(2), toDelete[1].ID)
// active is newest; keep=2 → keep {4,3}, prune {2,1}
toDelete = selectDeploymentsToPrune(deployments, 4, 2)
toDelete = selectDeploymentsToPrune(deployments, 4, 0, 2)
require.Len(t, toDelete, 2)
assert.Equal(t, uint(2), toDelete[0].ID)
assert.Equal(t, uint(1), toDelete[1].ID)
// no active; keep=2 → keep {4,3}
toDelete = selectDeploymentsToPrune(deployments, 0, 2)
toDelete = selectDeploymentsToPrune(deployments, 0, 0, 2)
require.Len(t, toDelete, 2)
assert.Equal(t, uint(2), toDelete[0].ID)
assert.Equal(t, uint(1), toDelete[1].ID)
// keep=1 with active → only active, prune the rest
toDelete = selectDeploymentsToPrune(deployments, 2, 1)
toDelete = selectDeploymentsToPrune(deployments, 2, 0, 1)
require.Len(t, toDelete, 3)
for _, item := range toDelete {
assert.NotEqual(t, uint(2), item.ID)
}
// already within limit
assert.Nil(t, selectDeploymentsToPrune(deployments[:2], 4, 2))
assert.Nil(t, selectDeploymentsToPrune(deployments[:2], 4, 0, 2))
// unlimited
assert.Nil(t, selectDeploymentsToPrune(deployments, 1, 0))
assert.Nil(t, selectDeploymentsToPrune(deployments, 1, 0, 0))
// history=1 temporarily preserves active plus the freshly uploaded candidate.
toDelete = selectDeploymentsToPrune(deployments, 2, 4, 1)
require.Len(t, toDelete, 2)
assert.Equal(t, uint(3), toDelete[0].ID)
assert.Equal(t, uint(1), toDelete[1].ID)
assert.Equal(t, uint(4), resolveLatestCandidateID(deployments, 2, true))
assert.Zero(t, resolveLatestCandidateID(deployments, 2, false))
}
func TestPruneProjectDeploymentHistory(t *testing.T) {
@@ -540,6 +639,131 @@ func TestPruneProjectDeploymentHistory(t *testing.T) {
assert.True(t, hasLatest, "newest deployment must fill remaining slot")
}
func TestHistoryCountOnePreservesFreshCandidateUntilActivation(t *testing.T) {
cleanup := setupPagesTestDB(t)
defer cleanup()
_, disableStorage := setupPagesStorageMock(t)
defer disableStorage()
ctx := context.Background()
require.NoError(t, db.DB(ctx).Model(&model.SystemConfig{}).
Where("key = ?", model.ConfigKeyPagesMaxHistoryCount).
Update("value", "1").Error)
require.NoError(t, repository.InvalidateSystemConfigCache(ctx, model.ConfigKeyPagesMaxHistoryCount))
project, err := CreateProject(ctx, Input{Name: "Single History", Slug: "single-history", Enabled: true})
require.NoError(t, err)
active, err := UploadDeployment(ctx, project.ID, testPagesMultipartFile(t, "v1.zip", testPagesZip(t, map[string]string{
"index.html": "v1",
})), "user:1")
require.NoError(t, err)
_, err = ActivateDeployment(ctx, project.ID, active.ID)
require.NoError(t, err)
oldCandidate, err := UploadDeployment(ctx, project.ID, testPagesMultipartFile(t, "v2.zip", testPagesZip(t, map[string]string{
"index.html": "v2",
})), "user:1")
require.NoError(t, err)
deployments, err := model.ListPagesDeployments(ctx, project.ID)
require.NoError(t, err)
require.Len(t, deployments, 2)
newCandidate, err := UploadDeployment(ctx, project.ID, testPagesMultipartFile(t, "v3.zip", testPagesZip(t, map[string]string{
"index.html": "v3",
})), "user:1")
require.NoError(t, err)
deployments, err = model.ListPagesDeployments(ctx, project.ID)
require.NoError(t, err)
require.Len(t, deployments, 2)
kept := map[uint]bool{}
for _, deployment := range deployments {
kept[deployment.ID] = true
}
assert.True(t, kept[active.ID])
assert.True(t, kept[newCandidate.ID])
assert.False(t, kept[oldCandidate.ID])
var removedUpload model.Upload
require.NoError(t, db.DB(ctx).First(&removedUpload, oldCandidate.UploadID).Error)
assert.Equal(t, model.UploadStatusDeleted, removedUpload.Status)
_, err = ActivateDeployment(ctx, project.ID, newCandidate.ID)
require.NoError(t, err)
deployments, err = model.ListPagesDeployments(ctx, project.ID)
require.NoError(t, err)
require.Len(t, deployments, 1)
assert.Equal(t, newCandidate.ID, deployments[0].ID)
}
func TestPruneUsesLockTimeNewestCandidateInsteadOfStaleCaller(t *testing.T) {
cleanup := setupPagesTestDB(t)
defer cleanup()
_, disableStorage := setupPagesStorageMock(t)
defer disableStorage()
ctx := context.Background()
project, err := CreateProject(ctx, Input{Name: "Concurrent Candidate", Slug: "concurrent-candidate", Enabled: true})
require.NoError(t, err)
active, err := UploadDeployment(ctx, project.ID, testPagesMultipartFile(t, "v1.zip", testPagesZip(t, map[string]string{
"index.html": "v1",
})), "user:1")
require.NoError(t, err)
_, err = ActivateDeployment(ctx, project.ID, active.ID)
require.NoError(t, err)
staleCandidate, err := UploadDeployment(ctx, project.ID, testPagesMultipartFile(t, "v2.zip", testPagesZip(t, map[string]string{
"index.html": "v2",
})), "user:1")
require.NoError(t, err)
newCandidate, err := UploadDeployment(ctx, project.ID, testPagesMultipartFile(t, "v3.zip", testPagesZip(t, map[string]string{
"index.html": "v3",
})), "user:1")
require.NoError(t, err)
deleted, err := pruneProjectDeploymentHistoryOnce(ctx, project.ID, 1, staleCandidate.ID)
require.NoError(t, err)
assert.Equal(t, 1, deleted)
deployments, err := model.ListPagesDeployments(ctx, project.ID)
require.NoError(t, err)
require.Len(t, deployments, 2)
kept := map[uint]bool{}
for _, deployment := range deployments {
kept[deployment.ID] = true
}
assert.True(t, kept[active.ID])
assert.True(t, kept[newCandidate.ID])
assert.False(t, kept[staleCandidate.ID])
}
func TestDeleteDeploymentAndProjectSoftDeleteUnreferencedArtifacts(t *testing.T) {
cleanup := setupPagesTestDB(t)
defer cleanup()
_, disableStorage := setupPagesStorageMock(t)
defer disableStorage()
ctx := context.Background()
project, err := CreateProject(ctx, Input{Name: "Delete Artifacts", Slug: "delete-artifacts", Enabled: true})
require.NoError(t, err)
first, err := UploadDeployment(ctx, project.ID, testPagesMultipartFile(t, "first.zip", testPagesZip(t, map[string]string{
"index.html": "first",
})), "user:1")
require.NoError(t, err)
second, err := UploadDeployment(ctx, project.ID, testPagesMultipartFile(t, "second.zip", testPagesZip(t, map[string]string{
"index.html": "second",
})), "user:1")
require.NoError(t, err)
require.NoError(t, DeleteDeployment(ctx, project.ID, second.ID))
var secondUpload model.Upload
require.NoError(t, db.DB(ctx).First(&secondUpload, second.UploadID).Error)
assert.Equal(t, model.UploadStatusDeleted, secondUpload.Status)
require.NoError(t, DeleteProject(ctx, project.ID))
var firstUpload model.Upload
require.NoError(t, db.DB(ctx).First(&firstUpload, first.UploadID).Error)
assert.Equal(t, model.UploadStatusDeleted, firstUpload.Status)
_, err = model.GetPagesProjectByID(ctx, project.ID)
assert.Error(t, err)
}
func testPagesTarGz(t *testing.T, files map[string]string) []byte {
t.Helper()
@@ -0,0 +1,64 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pages
import (
"context"
"errors"
"fmt"
"strings"
"github.com/Rain-kl/Wavelet/internal/apps/upload"
"github.com/Rain-kl/Wavelet/internal/model"
)
// ProjectLatestPackageMetadata describes the active package limits published to Agents.
type ProjectLatestPackageMetadata struct {
DeploymentID uint
Hash string
PackageSize int64
FileCount int
TotalSize int64
}
// GetProjectLatestPackageMetadata returns one coherent metadata snapshot for a
// project's currently active deployment.
func GetProjectLatestPackageMetadata(ctx context.Context, projectID uint) (*ProjectLatestPackageMetadata, error) {
deployment, err := resolveProjectActiveDeploymentForAgent(ctx, projectID)
if err != nil {
return nil, err
}
if deployment.UploadID == 0 {
if err := ensureDeploymentUploadRecord(ctx, deployment); err != nil {
return nil, err
}
deployment, err = model.GetPagesDeploymentByID(ctx, deployment.ID)
if err != nil {
return nil, err
}
}
if deployment.UploadID == 0 {
return nil, errors.New(errPagesDeploymentNotFound)
}
uploadRecord, err := upload.GetActiveUpload(ctx, deployment.UploadID)
if err != nil {
return nil, fmt.Errorf("pages 部署包不存在: %w", err)
}
hash := strings.TrimSpace(uploadRecord.Hash)
if hash == "" {
hash = strings.TrimSpace(deployment.Checksum)
}
if hash == "" {
return nil, errors.New(errPagesDeploymentHashMissing)
}
return &ProjectLatestPackageMetadata{
DeploymentID: deployment.ID,
Hash: hash,
PackageSize: uploadRecord.FileSize,
FileCount: deployment.FileCount,
TotalSize: deployment.TotalSize,
}, nil
}
@@ -0,0 +1,81 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pages
import (
"context"
"crypto/sha256"
"encoding/hex"
"fmt"
"testing"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
)
func TestGetProjectLatestPackageMetadataPreservesZeroTotalSize(t *testing.T) {
cleanup := setupPagesTestDB(t)
defer cleanup()
_, disableStorage := setupPagesStorageMock(t)
defer disableStorage()
ctx := context.Background()
project, err := CreateProject(ctx, Input{
Name: "Empty Files",
Slug: "empty-files",
Enabled: true,
EntryFile: "index.html",
})
if err != nil {
t.Fatalf("CreateProject() error = %v", err)
}
packageBytes := testPagesZip(t, map[string]string{
"index.html": "",
".gitkeep": "",
})
deployment, err := UploadDeployment(
ctx,
project.ID,
testPagesMultipartFile(t, "empty-files.zip", packageBytes),
"test",
)
if err != nil {
t.Fatalf("UploadDeployment() error = %v", err)
}
if _, err := ActivateDeployment(ctx, project.ID, deployment.ID); err != nil {
t.Fatalf("ActivateDeployment() error = %v", err)
}
if err := db.DB(ctx).Create(&model.ConfigVersion{
Version: "v-package-metadata",
SnapshotJSON: fmt.Sprintf(
`{"routes":[{"upstream_type":"pages","pages_project_id":%d}]}`,
project.ID,
),
SupportFilesJSON: "[]",
Checksum: "package-metadata-config",
IsActive: true,
CreatedBy: "test",
}).Error; err != nil {
t.Fatalf("create active ConfigVersion error = %v", err)
}
got, err := GetProjectLatestPackageMetadata(ctx, project.ID)
if err != nil {
t.Fatalf("GetProjectLatestPackageMetadata(%d) error = %v", project.ID, err)
}
wantHashBytes := sha256.Sum256(packageBytes)
wantHash := hex.EncodeToString(wantHashBytes[:])
if got.DeploymentID != deployment.ID || got.Hash != wantHash {
t.Errorf("GetProjectLatestPackageMetadata(%d) identity = (%d, %q), want (%d, %q)",
project.ID, got.DeploymentID, got.Hash, deployment.ID, wantHash)
}
if got.PackageSize != int64(len(packageBytes)) {
t.Errorf("GetProjectLatestPackageMetadata(%d).PackageSize = %d, want %d",
project.ID, got.PackageSize, len(packageBytes))
}
if got.FileCount != 2 || got.TotalSize != 0 {
t.Errorf("GetProjectLatestPackageMetadata(%d) content = (%d files, %d bytes), want (2 files, 0 bytes)",
project.ID, got.FileCount, got.TotalSize)
}
}
+22 -7
View File
@@ -7,6 +7,7 @@ import (
"context"
"encoding/json"
"fmt"
"path"
"strings"
"github.com/Rain-kl/Wavelet/internal/model"
@@ -102,7 +103,10 @@ func rebindPagesRouteMaps(ctx context.Context, routes []map[string]json.RawMessa
if err != nil {
return false, err
}
deployment := buildLivePagesDeployment(project, activeDeployment)
deployment, err := buildLivePagesDeployment(project, activeDeployment)
if err != nil {
return false, err
}
projectIDCopy := project.ID
originURL := fmt.Sprintf("openflare-pages://project/%d", project.ID)
@@ -181,15 +185,26 @@ func loadActivePagesProject(ctx context.Context, projectID uint, siteName string
return project, activeDeployment, nil
}
func buildLivePagesDeployment(project *model.PagesProject, active *model.PagesDeployment) *openrestyrender.PagesDeployment {
entryFile := strings.TrimSpace(project.EntryFile)
if entryFile == "" {
entryFile = defaultPagesEntryFile
func buildLivePagesDeployment(
project *model.PagesProject,
active *model.PagesDeployment,
) (*openrestyrender.PagesDeployment, error) {
rootDir, err := validateAndNormalizePagesRootDir(project.RootDir)
if err != nil {
return nil, err
}
entryFile, err := validateAndNormalizePagesEntryFile(project.EntryFile)
if err != nil {
return nil, err
}
fallbackPath := strings.TrimSpace(project.SPAFallbackPath)
if fallbackPath == "" {
fallbackPath = defaultPagesFallbackPath
}
localRoot := openrestyrender.PagesProjectLocalRoot(project.ID)
if rootDir != "" {
localRoot = path.Join(localRoot, rootDir)
}
return &openrestyrender.PagesDeployment{
ProjectID: project.ID,
ProjectSlug: strings.TrimSpace(project.Slug),
@@ -203,8 +218,8 @@ func buildLivePagesDeployment(project *model.PagesProject, active *model.PagesDe
APIProxyPath: strings.TrimSpace(project.APIProxyPath),
APIProxyPass: strings.TrimSpace(project.APIProxyPass),
APIProxyRewrite: strings.TrimSpace(project.APIProxyRewrite),
LocalRoot: openrestyrender.PagesProjectLocalRoot(project.ID),
}
LocalRoot: localRoot,
}, nil
}
func rawJSONString(raw json.RawMessage) (string, bool) {
+6 -3
View File
@@ -20,9 +20,11 @@ func TestRebindSnapshotPagesToCurrentActive(t *testing.T) {
ctx := context.Background()
project, err := CreateProject(ctx, Input{
Name: "Rebind Site",
Slug: "rebind-site",
Enabled: true,
Name: "Rebind Site",
Slug: "rebind-site",
Enabled: true,
RootDir: "public/site",
EntryFile: "index.html",
})
require.NoError(t, err)
@@ -84,4 +86,5 @@ func TestRebindSnapshotPagesToCurrentActive(t *testing.T) {
deployment := route["pages_deployment"].(map[string]any)
assert.EqualValues(t, active.ID, deployment["deployment_id"])
assert.Equal(t, "new-checksum", deployment["checksum"])
assert.Equal(t, "__OPENFLARE_PAGES_DIR__/projects/1/current/public/site", deployment["local_root"])
}
+22 -2
View File
@@ -4,11 +4,14 @@
package pages
import (
"fmt"
"net/http"
"strconv"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/apiutil"
"github.com/Rain-kl/Wavelet/internal/common/response"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/gin-gonic/gin"
)
@@ -33,6 +36,15 @@ func deploymentIDParam(c *gin.Context) (uint, bool) {
return uint(id64), true
}
func currentPagesActor(c *gin.Context) (string, bool) {
user, ok := oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
if !ok || user == nil || user.ID == 0 {
response.AbortUnauthorized(c, errPagesActorMissing)
return "", false
}
return fmt.Sprintf("user:%d", user.ID), true
}
// ListProjectsHandler 列出全部 Pages 项目。
// @Summary 列出 Pages 项目
// @Description 返回全部 OpenFlare Pages 项目,需要管理员权限
@@ -214,7 +226,11 @@ func UploadDeploymentHandler(c *gin.Context) {
response.AbortBadRequest(c, errPagesPackageMissing)
return
}
deployment, err := UploadDeployment(c.Request.Context(), id, file, "")
actor, ok := currentPagesActor(c)
if !ok {
return
}
deployment, err := UploadDeployment(c.Request.Context(), id, file, actor)
if handleLogicError(c, err) {
return
}
@@ -247,7 +263,11 @@ func UploadDeploymentFromURLHandler(c *gin.Context) {
response.AbortBadRequest(c, errPagesPackageURLRequired)
return
}
deployment, err := UploadDeploymentFromURL(c.Request.Context(), id, req.URL, "")
actor, ok := currentPagesActor(c)
if !ok {
return
}
deployment, err := UploadDeploymentFromURL(c.Request.Context(), id, req.URL, actor)
if handleLogicError(c, err) {
return
}
@@ -0,0 +1,100 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pages
import (
"bytes"
"encoding/json"
"fmt"
"mime/multipart"
"net/http"
"net/http/httptest"
"strconv"
"testing"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
func TestUploadDeploymentHandlerRecordsCurrentUserActor(t *testing.T) {
cleanup := setupPagesTestDB(t)
defer cleanup()
_, disableStorage := setupPagesStorageMock(t)
defer disableStorage()
project, err := CreateProject(t.Context(), Input{Name: "Actor Upload", Slug: "actor-upload", Enabled: true})
require.NoError(t, err)
packageBytes := testPagesZip(t, map[string]string{"index.html": "ok"})
var requestBody bytes.Buffer
writer := multipart.NewWriter(&requestBody)
part, err := writer.CreateFormFile("package", "site.zip")
require.NoError(t, err)
_, err = part.Write(packageBytes)
require.NoError(t, err)
require.NoError(t, writer.Close())
req := httptest.NewRequest(http.MethodPost, "/api/v1/d/pages/1/deployments/upload", &requestBody)
req.Header.Set("Content-Type", writer.FormDataContentType())
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
c.Request = req
c.Params = gin.Params{{Key: "id", Value: strconv.FormatUint(uint64(project.ID), 10)}}
oauth.SetToContext(c, oauth.UserObjKey, &model.User{ID: 42})
UploadDeploymentHandler(c)
assert.Equal(t, http.StatusOK, recorder.Code)
deployments, err := model.ListPagesDeployments(t.Context(), project.ID)
require.NoError(t, err)
require.Len(t, deployments, 1)
assert.Equal(t, "user:42", deployments[0].CreatedBy)
}
func TestUploadDeploymentFromURLHandlerRecordsCurrentUserActor(t *testing.T) {
cleanup := setupPagesTestDB(t)
defer cleanup()
_, disableStorage := setupPagesStorageMock(t)
defer disableStorage()
project, err := CreateProject(t.Context(), Input{Name: "Actor URL", Slug: "actor-url", Enabled: true})
require.NoError(t, err)
packageBytes := testPagesZip(t, map[string]string{"index.html": "ok"})
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.Header().Set("Content-Type", "application/zip")
w.Header().Set("Content-Disposition", `attachment; filename="site.zip"`)
w.WriteHeader(http.StatusOK)
_, _ = w.Write(packageBytes)
}))
defer server.Close()
body, err := json.Marshal(UploadFromURLInput{URL: server.URL + "/site.zip"})
require.NoError(t, err)
req := httptest.NewRequest(http.MethodPost, "/api/v1/d/pages/1/deployments/upload-from-url", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
c.Request = req
c.Params = gin.Params{{Key: "id", Value: fmt.Sprint(project.ID)}}
oauth.SetToContext(c, oauth.UserObjKey, &model.User{ID: 77})
UploadDeploymentFromURLHandler(c)
assert.Equal(t, http.StatusOK, recorder.Code, recorder.Body.String())
deployments, err := model.ListPagesDeployments(t.Context(), project.ID)
require.NoError(t, err)
require.Len(t, deployments, 1)
assert.Equal(t, "user:77", deployments[0].CreatedBy)
}
func TestCurrentPagesActorRejectsMissingUser(t *testing.T) {
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
actor, ok := currentPagesActor(c)
assert.False(t, ok)
assert.Empty(t, actor)
assert.True(t, c.IsAborted())
}
@@ -6,12 +6,14 @@ package proxy_route
import (
"context"
"errors"
"sort"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"gorm.io/gorm"
"gorm.io/gorm/clause"
)
// CustomHeaderInput 自定义响应头。
@@ -122,6 +124,9 @@ func CreateProxyRoute(ctx context.Context, input Input) (*View, error) {
return nil, err
}
if err = db.DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := lockPagesProjectsForRouteMutation(tx, 0, route); err != nil {
return err
}
if err := tx.Create(route).Error; err != nil {
return err
}
@@ -141,11 +146,15 @@ func UpdateProxyRoute(ctx context.Context, id uint, input Input) (*View, error)
if err != nil {
return nil, err
}
previousPagesProjectID := pagesProjectIDForRoute(route)
route, _, err = buildProxyRoute(ctx, route, input)
if err != nil {
return nil, err
}
if err = db.DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := lockPagesProjectsForRouteMutation(tx, previousPagesProjectID, route); err != nil {
return err
}
if err := updateProxyRouteRecord(tx, route); err != nil {
return err
}
@@ -159,6 +168,58 @@ func UpdateProxyRoute(ctx context.Context, id uint, input Input) (*View, error)
return buildProxyRouteView(ctx, route)
}
func pagesProjectIDForRoute(route *model.ProxyRoute) uint {
if route == nil || route.UpstreamType != proxyRouteUpstreamTypePages || route.PagesProjectID == nil {
return 0
}
return *route.PagesProjectID
}
func lockPagesProjectsForRouteMutation(tx *gorm.DB, previousProjectID uint, route *model.ProxyRoute) error {
nextProjectID := pagesProjectIDForRoute(route)
var projectIDs []uint
if previousProjectID != 0 {
projectIDs = append(projectIDs, previousProjectID)
}
if nextProjectID != 0 && nextProjectID != previousProjectID {
projectIDs = append(projectIDs, nextProjectID)
}
sort.Slice(projectIDs, func(i int, j int) bool { return projectIDs[i] < projectIDs[j] })
for _, projectID := range projectIDs {
var project model.PagesProject
err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).First(&project, projectID).Error
if errors.Is(err, gorm.ErrRecordNotFound) && projectID != nextProjectID {
continue
}
if errors.Is(err, gorm.ErrRecordNotFound) {
return errors.New(errProxyRoutePagesNotFound)
}
if err != nil {
return err
}
if projectID == nextProjectID {
if err := validateLockedPagesRouteProject(&project); err != nil {
return err
}
}
}
return nil
}
func validateLockedPagesRouteProject(project *model.PagesProject) error {
if project == nil {
return errors.New(errProxyRoutePagesNotFound)
}
if !project.Enabled {
return errors.New(errProxyRoutePagesDisabled)
}
if project.ActiveDeploymentID == nil || *project.ActiveDeploymentID == 0 {
return errors.New(errProxyRoutePagesNoDeploy)
}
return nil
}
// DeleteProxyRoute 删除代理规则。
func DeleteProxyRoute(ctx context.Context, id uint) error {
if _, err := model.GetProxyRouteByID(ctx, id); err != nil {
@@ -19,7 +19,14 @@ func setupProxyRouteTestDB(t *testing.T) func() {
t.Helper()
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{DisableForeignKeyConstraintWhenMigrating: true})
require.NoError(t, err)
require.NoError(t, sqliteDB.AutoMigrate(&model.ProxyRoute{}, &model.Origin{}, &model.Zone{}, &model.ZoneDomain{}, &model.TLSCertificate{}))
require.NoError(t, sqliteDB.AutoMigrate(
&model.ProxyRoute{},
&model.Origin{},
&model.Zone{},
&model.ZoneDomain{},
&model.TLSCertificate{},
&model.PagesProject{},
))
db.SetDB(sqliteDB)
return func() { db.SetDB(nil) }
}
@@ -81,6 +88,52 @@ func TestCreateProxyRouteHTTPSRequiresCoveringCertificate(t *testing.T) {
require.EqualError(t, err, errProxyRouteCertRequired)
}
func TestPagesRouteLocksAndRevalidatesTargetProject(t *testing.T) {
cleanup := setupProxyRouteTestDB(t)
defer cleanup()
ctx := context.Background()
domain := createZoneDomain(t, ctx, "pages.example.com", nil)
activeDeploymentID := uint(99)
project := &model.PagesProject{
Name: "Pages Site",
Slug: "pages-site",
Enabled: true,
ActiveDeploymentID: &activeDeploymentID,
}
require.NoError(t, db.DB(ctx).Create(project).Error)
view, err := CreateProxyRoute(ctx, Input{
SiteName: "pages",
ZoneDomainIDs: []uint{domain.ID},
UpstreamType: proxyRouteUpstreamTypePages,
PagesProjectID: &project.ID,
Enabled: true,
})
require.NoError(t, err)
require.NotNil(t, view.PagesProjectID)
assert.Equal(t, project.ID, *view.PagesProjectID)
require.NoError(t, db.DB(ctx).Delete(&model.PagesProject{}, project.ID).Error)
route := &model.ProxyRoute{UpstreamType: proxyRouteUpstreamTypePages, PagesProjectID: &project.ID}
err = db.DB(ctx).Transaction(func(tx *gorm.DB) error {
return lockPagesProjectsForRouteMutation(tx, 0, route)
})
require.EqualError(t, err, errProxyRoutePagesNotFound)
}
func TestRouteCanMoveAwayFromAlreadyMissingPagesProject(t *testing.T) {
cleanup := setupProxyRouteTestDB(t)
defer cleanup()
ctx := context.Background()
missingProjectID := uint(404)
route := &model.ProxyRoute{UpstreamType: "direct"}
err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
return lockPagesProjectsForRouteMutation(tx, missingProjectID, route)
})
require.NoError(t, err)
}
func TestNormalizeCachePolicyDefaultsAndLegacy(t *testing.T) {
assert.Equal(t, "", normalizeCachePolicy(false, "static"))
// Empty/url on write = legacy all (compat); UI sends static explicitly for new default.
+15 -9
View File
@@ -8,6 +8,7 @@ import (
"github.com/Rain-kl/Wavelet/internal/apps/upload/filesrv"
"github.com/Rain-kl/Wavelet/internal/apps/upload/handler"
"github.com/Rain-kl/Wavelet/internal/apps/upload/ingest"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
uploadtask "github.com/Rain-kl/Wavelet/internal/apps/upload/task"
"github.com/Rain-kl/Wavelet/internal/apps/upload/util"
@@ -31,15 +32,17 @@ var (
// Programmatic ingest API
var (
Ingest = ingest.Ingest
Remove = ingest.Remove
RemoveOwned = ingest.RemoveOwned
FindByHash = ingest.FindByHash
GetActiveUpload = ingest.GetActive
OpenStoredUpload = ingest.OpenActiveObject
ActiveUploadHash = ingest.ActiveHash
ResolveLocalFile = ingest.ResolveLocalFile
IngestFromLocalPath = ingest.FromLocalPath
Ingest = ingest.Ingest
Remove = ingest.Remove
RemoveOwned = ingest.RemoveOwned
RemoveLockedTx = ingest.RemoveLockedTx
InvalidateUploadMetaCache = ingest.InvalidateUploadMetaCache
FindByHash = ingest.FindByHash
GetActiveUpload = ingest.GetActive
OpenStoredUpload = ingest.OpenActiveObject
ActiveUploadHash = ingest.ActiveHash
ResolveLocalFile = ingest.ResolveLocalFile
IngestFromLocalPath = ingest.FromLocalPath
)
type (
@@ -54,6 +57,8 @@ const (
PolicyCreate = ingest.PolicyCreate
PolicyDedupNewRecord = ingest.PolicyDedupNewRecord
PolicyResolveExisting = ingest.PolicyResolveExisting
// ReservedPagesDeploymentType is managed exclusively by the Pages domain.
ReservedPagesDeploymentType = shared.ReservedPagesDeploymentType
)
type (
@@ -69,6 +74,7 @@ type (
var (
ErrIngestForbidden = ingest.ErrForbidden
ErrIngestStorageReadOnly = ingest.ErrStorageReadOnly
ErrReservedUploadType = ingest.ErrReservedUploadType
)
// Cache management
@@ -4,6 +4,7 @@
package handler
import (
"errors"
"net/http"
"strconv"
@@ -96,6 +97,7 @@ func ListFiles(c *gin.Context) {
// @Success 200 {object} response.Any "删除成功"
// @Failure 403 {object} response.Any "无权操作"
// @Failure 404 {object} response.Any "文件不存在"
// @Failure 409 {object} response.Any "系统保留类型或存储只读"
// @Router /api/v1/admin/uploads/{id} [delete]
func DeleteFile(c *gin.Context) {
ctx := c.Request.Context()
@@ -111,6 +113,10 @@ func DeleteFile(c *gin.Context) {
}
if _, err := softDeleteUpload(ctx, uploadID); err != nil {
if errors.Is(err, ingest.ErrReservedUploadType) {
response.AbortConflict(c, shared.ErrReservedUploadType)
return
}
if isRecordNotFound(err) {
response.AbortNotFound(c, "文件记录未找到")
return
@@ -216,6 +222,7 @@ func ListMyFiles(c *gin.Context) {
// @Success 200 {object} response.Any "删除成功"
// @Failure 403 {object} response.Any "无权操作"
// @Failure 404 {object} response.Any "文件不存在"
// @Failure 409 {object} response.Any "系统保留类型或存储只读"
// @Router /api/v1/upload/{id} [delete]
func DeleteMyFile(c *gin.Context) {
currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
@@ -232,11 +239,15 @@ func DeleteMyFile(c *gin.Context) {
}
if _, err := softDeleteOwnedUpload(ctx, currUser.ID, uploadID); err != nil {
if errors.Is(err, ingest.ErrReservedUploadType) {
response.AbortConflict(c, shared.ErrReservedUploadType)
return
}
if isRecordNotFound(err) {
response.AbortNotFound(c, "文件记录未找到")
return
}
if err == ingest.ErrForbidden {
if errors.Is(err, ingest.ErrForbidden) {
response.AbortForbidden(c, "无权操作")
return
}
+5
View File
@@ -53,6 +53,7 @@ type batchDownloadRequest struct {
// @Success 200 {object} response.Any{data=model.Upload} "上传成功"
// @Failure 400 {object} response.Any "请求参数错误或文件受限"
// @Failure 401 {object} response.Any "未登录"
// @Failure 409 {object} response.Any "系统保留类型或存储只读"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/upload [post]
//
@@ -107,6 +108,10 @@ func UploadFile(c *gin.Context) {
}
uploadType := c.DefaultPostForm("type", "generic")
if uploadType == shared.ReservedPagesDeploymentType {
response.AbortConflict(c, shared.ErrReservedUploadType)
return
}
accessMode, errMsg := resolveUploadAccessMode(c, uploadType)
if errMsg != "" {
@@ -226,6 +226,42 @@ func TestUploadFile(t *testing.T) {
}
})
t.Run("upload rejects Pages reserved type", func(t *testing.T) {
putCountBefore := putCount
imgContent := []byte("\x89PNG\r\n\x1a\n\x00\x00\x00\rIHDR\x00\x00\x00\x01\x00\x00\x00\x01")
contentType, body := createMultipartRequest(t, "file", "pages.png", imgContent, map[string]string{
"type": shared.ReservedPagesDeploymentType,
})
req, _ := http.NewRequest("POST", "/api/v1/upload", body)
req.Header.Set("Content-Type", contentType)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusConflict {
t.Fatalf("expected status 409, got %d. Body: %s", w.Code, w.Body.String())
}
var resp testResponse
if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
t.Fatalf("unmarshal reserved type response: %v", err)
}
if resp.ErrorMsg != shared.ErrReservedUploadType {
t.Fatalf("reserved type error = %q, want %q", resp.ErrorMsg, shared.ErrReservedUploadType)
}
if putCount != putCountBefore {
t.Fatalf("reserved upload wrote storage object: put count %d -> %d", putCountBefore, putCount)
}
var count int64
if err := dbConn.Model(&model.Upload{}).
Where("type = ?", shared.ReservedPagesDeploymentType).
Count(&count).Error; err != nil {
t.Fatalf("count reserved uploads: %v", err)
}
if count != 0 {
t.Fatalf("reserved upload record count = %d, want 0", count)
}
})
t.Run("instant upload deduplication (秒传)", func(t *testing.T) {
putCount = 0
imgContent := []byte("\x89PNG\r\n\x1a\n\x00\x00\x00\rIHDR\x00\x00\x00\x01\x00\x00\x00\x01")
@@ -993,6 +1029,62 @@ func TestUserUploadManagement(t *testing.T) {
})
}
func TestDeleteReservedUploadType(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
authUser := &model.User{ID: 1001, Username: "reserved_owner"}
router := setupTestRouter(authUser)
reserved := model.Upload{
ID: 4101,
UserID: authUser.ID,
FileName: "pages.zip",
FilePath: "uploads/pages.zip",
FileSize: 128,
MimeType: "application/zip",
Extension: "zip",
Hash: "pages-reserved-hash",
Type: shared.ReservedPagesDeploymentType,
Status: model.UploadStatusUsed,
CreatedAt: time.Now(),
}
if err := dbConn.Create(&reserved).Error; err != nil {
t.Fatalf("seed reserved upload: %v", err)
}
for _, tc := range []struct {
name string
path string
}{
{name: "admin delete", path: "/api/v1/admin/uploads/4101"},
{name: "owner delete", path: "/api/v1/upload/4101"},
} {
t.Run(tc.name, func(t *testing.T) {
req, _ := http.NewRequest(http.MethodDelete, tc.path, nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusConflict {
t.Fatalf("expected status 409, got %d. Body: %s", w.Code, w.Body.String())
}
var resp testResponse
if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
t.Fatalf("unmarshal delete response: %v", err)
}
if resp.ErrorMsg != shared.ErrReservedUploadType {
t.Fatalf("reserved delete error = %q, want %q", resp.ErrorMsg, shared.ErrReservedUploadType)
}
})
}
var persisted model.Upload
if err := dbConn.First(&persisted, reserved.ID).Error; err != nil {
t.Fatalf("reload reserved upload: %v", err)
}
if persisted.Status != model.UploadStatusUsed {
t.Fatalf("reserved upload status = %s, want used", persisted.Status)
}
}
func configureLocalStorageRoot(t *testing.T, dbConn *gorm.DB, tempDir string) {
var sc model.SystemConfig
if err := dbConn.Where("key = ?", model.ConfigKeyStorageConfig).First(&sc).Error; err != nil {
+3
View File
@@ -12,5 +12,8 @@ import (
// ErrForbidden indicates the caller is not allowed to mutate the upload record.
var ErrForbidden = errors.New("upload forbidden")
// ErrReservedUploadType indicates that a generic mutation targeted a domain-reserved upload type.
var ErrReservedUploadType = errors.New(shared.ErrReservedUploadType)
// ErrStorageReadOnly indicates the storage backend is in migration read-only mode.
var ErrStorageReadOnly = errors.New(shared.ErrStorageReadOnly)
+18 -9
View File
@@ -100,13 +100,10 @@ func storeObject(ctx context.Context, objectKey string, reader io.Reader, size i
return result.Key, nil
}
func persistUploadRecord(ctx context.Context, upload *model.Upload, objectKey string) error {
func persistUploadRecord(ctx context.Context, upload *model.Upload, objectKey string, storedByRequest bool) error {
if err := createUploadWithStats(ctx, upload); err != nil {
_, backend, backendErr := storage.Active(ctx)
if backendErr == nil {
if deleteErr := backend.Delete(ctx, objectKey); deleteErr != nil {
logger.WarnF(ctx, "清理未写入数据库的上传对象失败: %v", deleteErr)
}
if storedByRequest {
cleanupUnpersistedObject(ctx, objectKey)
}
return err
}
@@ -114,6 +111,16 @@ func persistUploadRecord(ctx context.Context, upload *model.Upload, objectKey st
return nil
}
func cleanupUnpersistedObject(ctx context.Context, objectKey string) {
_, backend, err := storage.Active(ctx)
if err != nil {
return
}
if err := backend.Delete(ctx, objectKey); err != nil {
logger.WarnF(ctx, "清理未写入数据库的上传对象失败: %v", err)
}
}
func createUploadWithStats(ctx context.Context, upload *model.Upload) error {
return db.DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := repository.CreateUploadTx(tx, upload); err != nil {
@@ -125,6 +132,8 @@ func createUploadWithStats(ctx context.Context, upload *model.Upload) error {
func createDedupRecord(ctx context.Context, existing model.Upload, req Request) (Result, error) {
accessMode := resolveAccessMode(req.Type, req.AccessMode)
metadata := req.Metadata
metadata.Bucket = existing.Metadata.Bucket
newUpload := model.Upload{
ID: idgen.NextUint64ID(),
UserID: req.UserID,
@@ -137,9 +146,9 @@ func createDedupRecord(ctx context.Context, existing model.Upload, req Request)
Type: req.Type,
Status: req.Status,
AccessMode: accessMode,
Metadata: existing.Metadata,
Metadata: metadata,
}
if err := persistUploadRecord(ctx, &newUpload, existing.FilePath); err != nil {
if err := persistUploadRecord(ctx, &newUpload, existing.FilePath, false); err != nil {
return Result{}, err
}
logger.InfoF(ctx, "文件触发秒传成功! ID: %d, Path: %s", newUpload.ID, existing.FilePath)
@@ -186,7 +195,7 @@ func createNewUpload(ctx context.Context, req Request) (Result, error) {
AccessMode: accessMode,
Metadata: req.Metadata,
}
if err := persistUploadRecord(ctx, &upload, storedKey); err != nil {
if err := persistUploadRecord(ctx, &upload, storedKey, true); err != nil {
return Result{}, err
}
+233 -2
View File
@@ -8,15 +8,20 @@ import (
"context"
"crypto/sha256"
"encoding/hex"
"errors"
"io"
"os"
"sync"
"testing"
"time"
uploadcache "github.com/Rain-kl/Wavelet/internal/apps/upload/cache"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/storage"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"gorm.io/gorm"
)
func TestIngestPolicyCreateIncrementsStats(t *testing.T) {
@@ -141,7 +146,11 @@ func TestIngestPolicyDedupNewRecordCreatesSecondRecord(t *testing.T) {
Extension: "png",
Hash: hashStr,
Type: "avatar",
Policy: PolicyDedupNewRecord,
Metadata: model.UploadMetadata{
UserAgent: "first-agent",
Extra: map[string]any{"record": "first"},
},
Policy: PolicyDedupNewRecord,
})
if err != nil {
t.Fatalf("first Ingest returned error: %v", err)
@@ -149,6 +158,10 @@ func TestIngestPolicyDedupNewRecordCreatesSecondRecord(t *testing.T) {
if putCount != 1 {
t.Fatalf("putCount after first ingest = %d, want 1", putCount)
}
first.Upload.Metadata.Bucket = "shared-bucket"
if err := dbConn.Save(&first.Upload).Error; err != nil {
t.Fatalf("update first upload metadata failed: %v", err)
}
second, err := Ingest(ctx, Request{
UserID: 1002,
@@ -159,7 +172,12 @@ func TestIngestPolicyDedupNewRecordCreatesSecondRecord(t *testing.T) {
Extension: "png",
Hash: hashStr,
Type: "avatar",
Policy: PolicyDedupNewRecord,
Metadata: model.UploadMetadata{
UserAgent: "second-agent",
Bucket: "caller-bucket-must-not-survive",
Extra: map[string]any{"record": "second"},
},
Policy: PolicyDedupNewRecord,
})
if err != nil {
t.Fatalf("second Ingest returned error: %v", err)
@@ -173,6 +191,15 @@ func TestIngestPolicyDedupNewRecordCreatesSecondRecord(t *testing.T) {
if first.Upload.ID == second.Upload.ID {
t.Fatal("dedup records should have unique IDs")
}
if second.Upload.Metadata.Bucket != "shared-bucket" {
t.Fatalf("dedup bucket = %q, want inherited shared-bucket", second.Upload.Metadata.Bucket)
}
if second.Upload.Metadata.UserAgent != "second-agent" {
t.Fatalf("dedup user agent = %q, want caller metadata", second.Upload.Metadata.UserAgent)
}
if second.Upload.Metadata.Extra["record"] != "second" {
t.Fatalf("dedup extra metadata = %#v, want caller metadata", second.Upload.Metadata.Extra)
}
var count int64
if err := dbConn.Model(&model.Upload{}).Where("hash = ?", hashStr).Count(&count).Error; err != nil {
@@ -183,6 +210,77 @@ func TestIngestPolicyDedupNewRecordCreatesSecondRecord(t *testing.T) {
}
}
func TestDedupRecordFailureDoesNotDeleteSharedObject(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
ctx := context.Background()
content := []byte("\x89PNG\r\n\x1a\nshared-object")
hash := sha256.Sum256(content)
hashStr := hex.EncodeToString(hash[:])
deleteCount := 0
restoreStorage, disableStorage := setupMockStorageWithDeleteCount(t, nil, &deleteCount)
defer restoreStorage()
defer disableStorage()
first, err := Ingest(ctx, Request{
UserID: 1001,
Reader: bytes.NewReader(content),
Size: int64(len(content)),
FileName: "shared.png",
MimeType: "image/png",
Extension: "png",
Hash: hashStr,
Type: "avatar",
Policy: PolicyDedupNewRecord,
})
if err != nil {
t.Fatalf("first Ingest returned error: %v", err)
}
const callbackName = "test:reject_dedup_upload_record"
if err := dbConn.Callback().Create().Before("gorm:create").Register(callbackName, func(tx *gorm.DB) {
upload, ok := tx.Statement.Dest.(*model.Upload)
if ok && upload.FileName == "dedup-fail.png" {
tx.AddError(errors.New("injected upload create failure"))
}
}); err != nil {
t.Fatalf("register create failure callback: %v", err)
}
defer func() { _ = dbConn.Callback().Create().Remove(callbackName) }()
_, err = Ingest(ctx, Request{
UserID: 1002,
Reader: bytes.NewReader(content),
Size: int64(len(content)),
FileName: "dedup-fail.png",
MimeType: "image/png",
Extension: "png",
Hash: hashStr,
Type: "avatar",
Metadata: model.UploadMetadata{
Extra: map[string]any{"record": "dedup-failure"},
},
Policy: PolicyDedupNewRecord,
})
if err == nil {
t.Fatal("dedup Ingest expected injected persistence error")
}
if deleteCount != 0 {
t.Fatalf("shared object delete count = %d, want 0", deleteCount)
}
_, backend, err := storage.Active(ctx)
if err != nil {
t.Fatalf("load active storage: %v", err)
}
obj, err := backend.Get(ctx, first.Upload.FilePath)
if err != nil {
t.Fatalf("shared object became unreadable after dedup failure: %v", err)
}
_ = obj.Body.Close()
}
func TestCreateUploadWithStatsRollsBackOnCreateFailure(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
@@ -259,6 +357,18 @@ func TestRemoveDecrementsStats(t *testing.T) {
if _, err := Remove(ctx, result.Upload.ID); err != nil {
t.Fatalf("Remove(%d) returned error: %v", result.Upload.ID, err)
}
stale := result.Upload
uploadcache.SetUploadMetaCache(ctx, &stale)
removedAgain, err := Remove(ctx, result.Upload.ID)
if err != nil {
t.Fatalf("second Remove(%d) returned error: %v", result.Upload.ID, err)
}
if removedAgain.Status != model.UploadStatusDeleted {
t.Fatalf("second Remove status = %s, want deleted", removedAgain.Status)
}
if _, err := uploadcache.GetUploadByID(ctx, result.Upload.ID); !errors.Is(err, gorm.ErrRecordNotFound) {
t.Fatalf("cache lookup after idempotent Remove error = %v, want record not found", err)
}
stats, err := loadTotalStats(ctx)
if err != nil {
@@ -269,6 +379,120 @@ func TestRemoveDecrementsStats(t *testing.T) {
}
}
func TestConcurrentRemoveDecrementsStatsOnce(t *testing.T) {
_, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
ctx := context.Background()
content := []byte("\x89PNG\r\n\x1a\nconcurrent-remove")
hash := sha256.Sum256(content)
restoreStorage, disableStorage := setupMockStorage(t, nil)
defer restoreStorage()
defer disableStorage()
result, err := Ingest(ctx, Request{
UserID: 1001,
Reader: bytes.NewReader(content),
Size: int64(len(content)),
FileName: "concurrent.png",
MimeType: "image/png",
Extension: "png",
Hash: hex.EncodeToString(hash[:]),
Type: "generic",
Policy: PolicyCreate,
})
if err != nil {
t.Fatalf("Ingest returned error: %v", err)
}
const workers = 8
start := make(chan struct{})
errs := make(chan error, workers)
var wg sync.WaitGroup
for i := 0; i < workers; i++ {
wg.Add(1)
go func() {
defer wg.Done()
<-start
_, removeErr := Remove(ctx, result.Upload.ID)
errs <- removeErr
}()
}
close(start)
wg.Wait()
close(errs)
for removeErr := range errs {
if removeErr != nil {
t.Fatalf("concurrent Remove returned error: %v", removeErr)
}
}
stats, err := loadTotalStats(ctx)
if err != nil {
t.Fatalf("loadTotalStats returned error: %v", err)
}
if stats.TotalCount != 0 || stats.TotalSize != 0 {
t.Fatalf("stats after concurrent remove = count %d size %d, want zero", stats.TotalCount, stats.TotalSize)
}
}
func TestRemoveOwnedAndReservedTypeBoundaries(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
ctx := context.Background()
ordinary := model.Upload{
ID: 99101,
UserID: 1001,
FileName: "owned.txt",
FilePath: "uploads/owned.txt",
FileSize: 16,
MimeType: "text/plain",
Extension: "txt",
Hash: "owned-hash",
Type: "generic",
Status: model.UploadStatusUsed,
CreatedAt: time.Now(),
}
reserved := model.Upload{
ID: 99102,
UserID: 1001,
FileName: "pages.zip",
FilePath: "uploads/pages.zip",
FileSize: 32,
MimeType: "application/zip",
Extension: "zip",
Hash: "reserved-hash",
Type: shared.ReservedPagesDeploymentType,
Status: model.UploadStatusUsed,
CreatedAt: time.Now(),
}
if err := dbConn.Create(&ordinary).Error; err != nil {
t.Fatalf("seed ordinary upload: %v", err)
}
if err := dbConn.Create(&reserved).Error; err != nil {
t.Fatalf("seed reserved upload: %v", err)
}
if _, err := RemoveOwned(ctx, 2002, ordinary.ID); !errors.Is(err, ErrForbidden) {
t.Fatalf("RemoveOwned non-owner error = %v, want ErrForbidden", err)
}
if _, err := Remove(ctx, reserved.ID); !errors.Is(err, ErrReservedUploadType) {
t.Fatalf("Remove reserved error = %v, want ErrReservedUploadType", err)
}
if _, err := RemoveOwned(ctx, reserved.UserID, reserved.ID); !errors.Is(err, ErrReservedUploadType) {
t.Fatalf("RemoveOwned reserved error = %v, want ErrReservedUploadType", err)
}
var persisted model.Upload
if err := dbConn.First(&persisted, reserved.ID).Error; err != nil {
t.Fatalf("reload reserved upload: %v", err)
}
if persisted.Status != model.UploadStatusUsed {
t.Fatalf("reserved upload status = %s, want used", persisted.Status)
}
}
type totalStatsSnapshot struct {
TotalCount int64
TotalSize int64
@@ -289,6 +513,10 @@ func loadTotalStats(ctx context.Context) (totalStatsSnapshot, error) {
}
func setupMockStorage(t *testing.T, putCount *int) (restore func(), disable func()) {
return setupMockStorageWithDeleteCount(t, putCount, nil)
}
func setupMockStorageWithDeleteCount(t *testing.T, putCount, deleteCount *int) (restore func(), disable func()) {
t.Helper()
mockFiles := make(map[string][]byte)
restore = storage.MockStorage(
@@ -316,6 +544,9 @@ func setupMockStorage(t *testing.T, putCount *int) (restore func(), disable func
},
func(ctx context.Context, key string) error {
delete(mockFiles, key)
if deleteCount != nil {
*deleteCount++
}
return nil
},
)
+46 -22
View File
@@ -7,52 +7,76 @@ import (
"context"
uploadcache "github.com/Rain-kl/Wavelet/internal/apps/upload/cache"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"gorm.io/gorm"
"gorm.io/gorm/clause"
)
// Remove soft-deletes an upload and decrements incremental stats.
// Remove soft-deletes an ordinary upload and decrements incremental stats once.
func Remove(ctx context.Context, uploadID uint64) (model.Upload, error) {
upload, err := repository.GetActiveUploadByID(ctx, uploadID)
upload, err := remove(ctx, 0, uploadID, false)
if err != nil {
return model.Upload{}, err
}
if err := softDeleteUploadWithStats(ctx, &upload); err != nil {
return model.Upload{}, err
}
upload.Status = model.UploadStatusDeleted
return upload, nil
}
// RemoveOwned soft-deletes an upload owned by userID and decrements incremental stats.
// RemoveOwned soft-deletes an ordinary upload owned by userID and decrements incremental stats once.
func RemoveOwned(ctx context.Context, userID, uploadID uint64) (model.Upload, error) {
upload, err := repository.GetActiveUploadByID(ctx, uploadID)
upload, err := remove(ctx, userID, uploadID, true)
if err != nil {
return model.Upload{}, err
}
if upload.UserID != userID {
return model.Upload{}, ErrForbidden
}
if err := softDeleteUploadWithStats(ctx, &upload); err != nil {
return model.Upload{}, err
}
upload.Status = model.UploadStatusDeleted
return upload, nil
}
func softDeleteUploadWithStats(ctx context.Context, upload *model.Upload) error {
statsSnapshot := *upload
func remove(ctx context.Context, userID, uploadID uint64, owned bool) (model.Upload, error) {
var upload model.Upload
if err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := repository.SoftDeleteUploadTx(tx, upload); err != nil {
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).
Where("id = ?", uploadID).
First(&upload).Error; err != nil {
return err
}
return uploadstats.ApplyUploadStatsDeltaTx(tx, &statsSnapshot, -1)
}); err != nil {
if owned && upload.UserID != userID {
return ErrForbidden
}
if upload.Type == shared.ReservedPagesDeploymentType {
return ErrReservedUploadType
}
_, err := RemoveLockedTx(tx, &upload)
return err
}); err != nil {
return model.Upload{}, err
}
uploadcache.InvalidateUploadMetaCache(ctx, upload.ID)
return nil
InvalidateUploadMetaCache(ctx, uploadID)
upload.Status = model.UploadStatusDeleted
return upload, nil
}
// RemoveLockedTx performs the idempotent active-to-deleted transition for a row
// that the caller has already locked in its surrounding transaction.
func RemoveLockedTx(tx *gorm.DB, upload *model.Upload) (bool, error) {
rowsAffected, err := repository.SoftDeleteUploadTx(tx, upload)
if err != nil {
return false, err
}
if rowsAffected == 0 {
return false, nil
}
if err := uploadstats.ApplyUploadStatsDeltaTx(tx, upload, -1); err != nil {
return false, err
}
upload.Status = model.UploadStatusDeleted
return true, nil
}
// InvalidateUploadMetaCache invalidates upload metadata after the caller commits its transaction.
func InvalidateUploadMetaCache(ctx context.Context, uploadID uint64) {
uploadcache.InvalidateUploadMetaCache(ctx, uploadID)
}
+2
View File
@@ -17,4 +17,6 @@ const (
FileStatsTrendDays = 7
MaxS3KeyLength = 1024
AccessCacheTTL = 5 // seconds; multiplied by time.Second at use site
// ReservedPagesDeploymentType is managed exclusively by the Pages domain.
ReservedPagesDeploymentType = "openflare_pages_deployment"
)
+1
View File
@@ -27,6 +27,7 @@ const (
ErrQueryFileCountFailed = "查询文件数量失败"
ErrQueryFileListFailed = "查询文件列表失败"
ErrDeleteFileFailed = "删除文件失败"
ErrReservedUploadType = "系统保留的文件类型不能通过通用文件接口操作"
ErrStorageReadOnly = "存储迁移维护中,当前仅允许读取文件"
ErrS3KeyRequired = "s3 key must not be empty"
ErrS3KeyTooLongFormat = "s3 key exceeds maximum length of %d"
+16 -18
View File
@@ -10,16 +10,15 @@ import (
"fmt"
"time"
uploadcache "github.com/Rain-kl/Wavelet/internal/apps/upload/cache"
"github.com/Rain-kl/Wavelet/internal/apps/upload/ingest"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/storage"
"github.com/Rain-kl/Wavelet/internal/task"
"github.com/Rain-kl/Wavelet/pkg/logger"
"gorm.io/gorm"
"gorm.io/gorm/clause"
)
const (
@@ -77,32 +76,31 @@ func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.Tas
for _, u := range unusedUploads {
totalProcessed++
transitioned := false
if err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Model(&model.Upload{}).
Where("id = ? AND status = ?", u.ID, model.UploadStatusPending).
Update("status", model.UploadStatusDeleted).Error; err != nil {
var locked model.Upload
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).
Where("id = ?", u.ID).
First(&locked).Error; err != nil {
return err
}
_, backend, err := storage.Active(ctx)
if err != nil {
return err
if locked.Status != model.UploadStatusPending || !locked.CreatedAt.Before(oneHourAgo) {
return nil
}
if err := backend.Delete(ctx, u.FilePath); err != nil {
return err
}
return nil
var err error
transitioned, err = ingest.RemoveLockedTx(tx, &locked)
return err
}); err != nil {
task.AppendLog(ctx, "清理上传文件失败 [ID:%d]: %v", u.ID, err)
lastID = u.ID
continue
}
uploadstats.RecordUploadStatsRemove(ctx, &u)
uploadcache.InvalidateUploadMetaCache(ctx, u.ID)
totalDeleted++
ingest.InvalidateUploadMetaCache(ctx, u.ID)
if transitioned {
totalDeleted++
}
lastID = u.ID
}
}
+28 -2
View File
@@ -19,6 +19,7 @@ import (
"github.com/Rain-kl/Wavelet/internal/apps/upload/filesrv"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/diskcache"
"github.com/Rain-kl/Wavelet/internal/model"
@@ -33,13 +34,17 @@ func TestSystemCleanupHandler_Execute(t *testing.T) {
_, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
// Mock S3 存储(让 DeleteObject 总是成功)
deleteCount := 0
// Mock S3 存储并记录 Delete,cleanup 不应物理删除共享对象。
storageMock := storage.MockStorage(
func(ctx context.Context, key string, body io.Reader, size int64, contentType string) error {
return nil
},
func(ctx context.Context, key string) (*storage.Object, error) { return nil, nil },
func(ctx context.Context, key string) error { return nil },
func(ctx context.Context, key string) error {
deleteCount++
return nil
},
)
defer storageMock()
storage.IsEnabledFunc = func() bool { return true }
@@ -86,6 +91,7 @@ func TestSystemCleanupHandler_Execute(t *testing.T) {
for _, r := range records {
err := db.DB(ctx).Create(r).Error
require.NoError(t, err)
require.NoError(t, uploadstats.ApplyUploadStatsAdd(ctx, r))
}
// 准备推送历史测试数据:1个旧的(应删除),1个新的(应保留)
@@ -147,6 +153,26 @@ func TestSystemCleanupHandler_Execute(t *testing.T) {
var usedCount int64
db.DB(ctx).Model(&model.Upload{}).Where("status = ?", model.UploadStatusUsed).Count(&usedCount)
assert.Equal(t, int64(1), usedCount, "used 状态的文件不应受影响")
assert.Equal(t, 0, deleteCount, "记录级 cleanup 不应调用 storage backend Delete")
var totalStats model.UploadStat
err = db.DB(ctx).
Where("dimension = ? AND stat_key = ?", model.UploadStatDimensionTotal, "").
First(&totalStats).Error
require.NoError(t, err)
assert.Equal(t, int64(2), totalStats.FileCount, "cleanup 后统计只应保留 used 与最近 pending 记录")
assert.Equal(t, int64(768), totalStats.FileSize, "cleanup 后统计大小应只扣减一次")
_, err = handler.Execute(ctx, nil)
require.NoError(t, err)
var statsAfterSecondRun model.UploadStat
err = db.DB(ctx).
Where("dimension = ? AND stat_key = ?", model.UploadStatDimensionTotal, "").
First(&statsAfterSecondRun).Error
require.NoError(t, err)
assert.Equal(t, totalStats.FileCount, statsAfterSecondRun.FileCount, "重复 cleanup 不应再次扣减统计")
assert.Equal(t, totalStats.FileSize, statsAfterSecondRun.FileSize, "重复 cleanup 不应再次扣减统计大小")
assert.Equal(t, 0, deleteCount, "重复 cleanup 仍不应调用 storage backend Delete")
// 验证推送历史数据状态:10天前的应被删除,今天的应保留
var pushCount int64
+12 -5
View File
@@ -62,15 +62,22 @@ func GetActiveUploadByID(ctx context.Context, id uint64) (model.Upload, error) {
return upload, nil
}
// SoftDeleteUpload marks an upload as deleted.
// SoftDeleteUpload marks an active upload as deleted and reports whether the row transitioned.
// External modules must use upload.Remove or upload.RemoveOwned; only internal/apps/upload may call this.
func SoftDeleteUpload(ctx context.Context, upload *model.Upload) error {
func SoftDeleteUpload(ctx context.Context, upload *model.Upload) (int64, error) {
return SoftDeleteUploadTx(db.DB(ctx), upload)
}
// SoftDeleteUploadTx marks an upload as deleted within an existing transaction.
func SoftDeleteUploadTx(tx *gorm.DB, upload *model.Upload) error {
return tx.Model(upload).Update("status", model.UploadStatusDeleted).Error
// SoftDeleteUploadTx marks an active upload as deleted within an existing transaction.
// RowsAffected is one only for the single successful active-to-deleted transition.
func SoftDeleteUploadTx(tx *gorm.DB, upload *model.Upload) (int64, error) {
result := tx.Model(&model.Upload{}).
Where("id = ? AND status IN ?", upload.ID, []model.UploadStatus{
model.UploadStatusPending,
model.UploadStatusUsed,
}).
Update("status", model.UploadStatusDeleted)
return result.RowsAffected, result.Error
}
// UpdateUpload applies partial field updates to an upload record.
+40 -23
View File
@@ -4,6 +4,7 @@
package pagesarchive
import (
"errors"
"fmt"
"io"
"math"
@@ -45,39 +46,50 @@ type Entry struct {
IsDir bool
// IsSymlink marks symbolic links (unsupported for Pages).
IsSymlink bool
// Size is the declared uncompressed size when known; 0 means empty or unknown.
// IsHardlink marks hard links (unsupported for Pages).
IsHardlink bool
// IsSpecial marks device, FIFO, socket, and other non-regular entries.
IsSpecial bool
// Size is the archive-declared uncompressed size; 0 means an empty member.
Size uint64
// Open returns a reader for the entry body. Caller must Close it.
// May be unavailable for inspect-only tar listings (body not materialized).
Open func() (io.ReadCloser, error)
}
// copyLimited copies src to dst.
// When maxBytes <= 0, size limits are not enforced (trusted extract path).
func copyLimited(dst io.Writer, src io.Reader, declaredSize uint64, maxBytes int64) (int64, error) {
if maxBytes <= 0 {
if declaredSize > 0 {
if declaredSize > uint64(math.MaxInt64) {
return 0, fmt.Errorf("pages file size out of bounds")
}
//nolint:gosec // declaredSize is bounded to MaxInt64 above
return io.CopyN(dst, src, int64(declaredSize))
}
// copyLimited copies actual bytes from src. maxBytes < 0 disables the byte cap;
// maxBytes == 0 permits only an empty stream.
func copyLimited(dst io.Writer, src io.Reader, maxBytes int64) (int64, error) {
if maxBytes < 0 {
return io.Copy(dst, src)
}
if declaredSize > uint64(maxBytes) || declaredSize > uint64(math.MaxInt64) { //nolint:gosec // maxBytes positive
return 0, fmt.Errorf("pages file size out of bounds")
readLimit := maxBytes
if maxBytes < math.MaxInt64 {
readLimit++
}
if declaredSize > 0 {
//nolint:gosec // declaredSize is bounded to MaxInt64 above
return io.CopyN(dst, src, int64(declaredSize))
written, err := io.Copy(dst, io.LimitReader(src, readLimit))
if err != nil {
return written, err
}
limited := io.LimitReader(src, maxBytes+1)
written, err := io.Copy(dst, limited)
if written > maxBytes {
return written, fmt.Errorf("pages file size out of bounds")
}
return written, err
return written, nil
}
func copyAndVerifySize(dst io.Writer, src io.Reader, declaredSize uint64, maxBytes int64) (int64, error) {
if declaredSize > uint64(math.MaxInt64) {
return 0, fmt.Errorf("pages file size out of bounds")
}
written, err := copyLimited(dst, src, maxBytes)
if err != nil {
return written, err
}
//nolint:gosec // declaredSize is bounded to MaxInt64 above
if written != int64(declaredSize) {
return written, fmt.Errorf("pages declared size %d does not match actual %d", declaredSize, written)
}
return written, nil
}
func writeEntryFile(targetPath string, src io.Reader, declaredSize uint64, maxBytes int64, perm os.FileMode) (int64, error) {
@@ -88,6 +100,11 @@ func writeEntryFile(targetPath string, src io.Reader, declaredSize uint64, maxBy
if err != nil {
return 0, err
}
defer func() { _ = target.Close() }()
return copyLimited(target, src, declaredSize, maxBytes)
written, copyErr := copyAndVerifySize(target, src, declaredSize, maxBytes)
closeErr := target.Close()
if err := errors.Join(copyErr, closeErr); err != nil {
_ = os.Remove(targetPath)
return written, err
}
return written, nil
}
+141 -68
View File
@@ -4,7 +4,9 @@
package pagesarchive
import (
"archive/tar"
"bytes"
"errors"
"fmt"
"io"
"os"
@@ -16,14 +18,12 @@ const formatDetectHeadBytes = 512
// ExtractOptions controls package extraction.
type ExtractOptions struct {
// Limits bounds files and sizes during extraction when EnforceLimits is true.
// Limits bounds actual files and sizes during extraction when EnforceLimits is true.
Limits Limits
// StripCommonRoot strips a single shared top-level directory when present.
StripCommonRoot bool
// EnforceLimits enables MaxFiles / MaxFileBytes / MaxTotalBytes checks.
// When false, the caller is assumed to have already validated the package
// (e.g. Agent trusts control-plane inspection). Path-escape and symlink
// guards still apply so local extraction cannot leave destDir.
// Path, member type, and declared/actual-size validation always remain enabled.
EnforceLimits bool
}
@@ -36,16 +36,11 @@ func ExtractBytes(data []byte, format Format, destDir string, opts ExtractOption
return err
}
}
entries, err := listEntriesAt(bytes.NewReader(data), int64(len(data)), format, true)
if err != nil {
return err
}
return extractEntries(entries, destDir, opts)
return extractFromReaderAt(bytes.NewReader(data), int64(len(data)), format, destDir, opts)
}
// ExtractFile opens path and extracts it into destDir without buffering the
// whole archive as an intermediate []byte for zip/7z (ReaderAt). Tar-family
// formats still materialize member bodies so random Open works for extract.
// whole archive or tar member bodies in memory.
func ExtractFile(filePath string, format Format, destDir string, opts ExtractOptions) error {
file, err := os.Open(filePath) //nolint:gosec // controlled path
if err != nil {
@@ -68,7 +63,14 @@ func ExtractFile(filePath string, format Format, destDir string, opts ExtractOpt
return err
}
}
entries, err := listEntriesAt(file, info.Size(), format, true)
return extractFromReaderAt(file, info.Size(), format, destDir, opts)
}
func extractFromReaderAt(ra io.ReaderAt, size int64, format Format, destDir string, opts ExtractOptions) error {
if isTarFamily(format) {
return extractTarFamilyAt(ra, size, format, destDir, opts)
}
entries, err := listRandomAccessEntriesAt(ra, size, format)
if err != nil {
return err
}
@@ -80,88 +82,159 @@ func extractEntries(entries []Entry, destDir string, opts ExtractOptions) error
if opts.EnforceLimits {
limits = normalizeLimits(opts.Limits)
}
commonPrefix := ""
if opts.StripCommonRoot {
commonPrefix = FindCommonRootPrefix(collectFileNames(entries))
commonPrefix, err := commonRootForEntries(entries, opts.StripCommonRoot)
if err != nil {
return err
}
var totalSize int64
var fileCount int
measured := &measuredArchive{files: make([]measuredFile, 0, len(entries))}
for _, entry := range entries {
written, counted, err := extractSingleEntry(entry, destDir, commonPrefix, limits, opts.EnforceLimits)
normalizedPath, skip, err := validateArchiveEntry(entry)
if err != nil {
return err
}
if !counted {
if skip {
continue
}
fileCount++
if opts.EnforceLimits && fileCount > limits.MaxFiles {
return fmt.Errorf("pages deployment file count exceeds %d", limits.MaxFiles)
normalizedPath = StripPrefix(normalizedPath, commonPrefix)
if normalizedPath == "" {
continue
}
totalSize += written
if opts.EnforceLimits && totalSize > limits.MaxTotalBytes {
return fmt.Errorf("pages extracted size exceeds limit")
if err := prepareMeasuredFile(measured, normalizedPath, entry.Size, limits, opts.EnforceLimits); err != nil {
return err
}
if entry.Open == nil {
return fmt.Errorf("%s: pages archive member cannot be opened", normalizedPath)
}
src, err := entry.Open()
if err != nil {
return fmt.Errorf("%s: %w", normalizedPath, err)
}
maxBytes := effectiveFileLimit(limits, measured.totalSize, opts.EnforceLimits)
targetPath, err := safeExtractionTarget(destDir, normalizedPath)
if err != nil {
_ = src.Close()
return err
}
actual, writeErr := writeEntryFile(targetPath, src, entry.Size, maxBytes, filePerm)
closeErr := src.Close()
if err := errors.Join(writeErr, closeErr); err != nil {
return fmt.Errorf("%s: %w", normalizedPath, err)
}
appendMeasuredFile(measured, normalizedPath, actual)
}
if fileCount == 0 {
if measured.fileCount == 0 {
return fmt.Errorf("pages package is empty")
}
return nil
}
func extractSingleEntry(
entry Entry,
destDir, commonPrefix string,
func extractTarFamilyAt(ra io.ReaderAt, size int64, format Format, destDir string, opts ExtractOptions) error {
limits := Limits{}
if opts.EnforceLimits {
limits = normalizeLimits(opts.Limits)
}
firstPass, err := scanTarFamilyAt(ra, size, format, limits, opts.EnforceLimits)
if err != nil {
return err
}
if firstPass.fileCount == 0 {
return fmt.Errorf("pages package is empty")
}
commonPrefix := ""
if opts.StripCommonRoot {
paths := make([]string, 0, len(firstPass.files))
for _, file := range firstPass.files {
paths = append(paths, file.path)
}
commonPrefix = FindCommonRootPrefix(paths)
}
tarReader, closeReader, err := openTarFamilyReader(io.NewSectionReader(ra, 0, size), format)
if err != nil {
return err
}
secondPass, extractErr := extractTarReader(tarReader, destDir, commonPrefix, limits, opts.EnforceLimits)
if closeErr := closeReader(); closeErr != nil {
extractErr = errors.Join(extractErr, closeErr)
}
if extractErr != nil {
return extractErr
}
if secondPass.fileCount != firstPass.fileCount || secondPass.totalSize != firstPass.totalSize {
return fmt.Errorf("pages tar package changed between validation and extraction")
}
return nil
}
func extractTarReader(
tarReader *tar.Reader,
destDir string,
commonPrefix string,
limits Limits,
enforceLimits bool,
) (written int64, counted bool, err error) {
relativePath, skip, err := NormalizeEntryPath(entry.Name)
if err != nil {
return 0, false, err
) (*measuredArchive, error) {
measured := &measuredArchive{files: make([]measuredFile, 0)}
for {
header, err := tarReader.Next()
if err == io.EOF {
break
}
if err != nil {
return nil, fmt.Errorf("read tar pages package: %w", err)
}
entry := entryFromTarHeader(header)
normalizedPath, skip, err := validateArchiveEntry(entry)
if err != nil {
return nil, err
}
if skip {
continue
}
normalizedPath = StripPrefix(normalizedPath, commonPrefix)
if normalizedPath == "" {
continue
}
if err := prepareMeasuredFile(measured, normalizedPath, entry.Size, limits, enforceLimits); err != nil {
return nil, err
}
targetPath, err := safeExtractionTarget(destDir, normalizedPath)
if err != nil {
return nil, err
}
maxBytes := effectiveFileLimit(limits, measured.totalSize, enforceLimits)
actual, err := writeEntryFile(targetPath, tarReader, entry.Size, maxBytes, filePerm)
if err != nil {
return nil, fmt.Errorf("%s: %w", normalizedPath, err)
}
appendMeasuredFile(measured, normalizedPath, actual)
}
if skip {
return 0, false, nil
}
if commonPrefix != "" {
relativePath = StripPrefix(relativePath, commonPrefix)
if relativePath == "" {
return 0, false, nil
return measured, nil
}
func commonRootForEntries(entries []Entry, strip bool) (string, error) {
paths := make([]string, 0, len(entries))
for _, entry := range entries {
normalizedPath, skip, err := validateArchiveEntry(entry)
if err != nil {
return "", err
}
if !skip {
paths = append(paths, normalizedPath)
}
}
if entry.IsSymlink {
return 0, false, fmt.Errorf("pages package contains unsupported symlink: %s", relativePath)
if !strip {
return "", nil
}
return FindCommonRootPrefix(paths), nil
}
func safeExtractionTarget(destDir, relativePath string) (string, error) {
targetPath := filepath.Join(destDir, filepath.FromSlash(relativePath))
if !isWithinDir(destDir, targetPath) {
return 0, false, fmt.Errorf("pages package path escapes directory: %s", entry.Name)
return "", fmt.Errorf("pages package path escapes directory: %s", relativePath)
}
if entry.IsDir {
if err := os.MkdirAll(targetPath, dirPerm); err != nil {
return 0, false, err
}
return 0, false, nil
}
maxFileBytes := int64(0) // unlimited when not enforcing
if enforceLimits {
if exceedsFileByteLimit(entry.Size, limits.MaxFileBytes) {
return 0, false, fmt.Errorf("pages file too large: %s", relativePath)
}
maxFileBytes = limits.MaxFileBytes
}
src, err := entry.Open()
if err != nil {
return 0, false, fmt.Errorf("%s: %w", relativePath, err)
}
written, writeErr := writeEntryFile(targetPath, src, entry.Size, maxFileBytes, filePerm)
_ = src.Close()
if writeErr != nil {
return 0, false, fmt.Errorf("%s: %w", relativePath, writeErr)
}
return written, true, nil
return targetPath, nil
}
func isWithinDir(baseDir, targetPath string) bool {
+189 -93
View File
@@ -4,13 +4,14 @@
package pagesarchive
import (
"archive/tar"
"bytes"
"errors"
"fmt"
"io"
"math"
"os"
"path"
"strings"
)
// InspectOptions controls package inspection.
@@ -19,17 +20,26 @@ type InspectOptions struct {
RootDir string
// EntryFile is the required entry file name (e.g. index.html).
EntryFile string
// Limits bounds files and sizes.
// Limits bounds files and actual extracted sizes.
Limits Limits
// VerifySizes, when true, streams each regular file and compares the actual
// byte count against the archive-declared size (no content hashing).
// Default false: trust zip central directory / tar header sizes.
// VerifySizes is retained for source compatibility. Inspection now always
// streams regular members and verifies actual bytes against declared sizes.
VerifySizes bool
}
type measuredFile struct {
path string
size int64
}
type measuredArchive struct {
files []measuredFile
fileCount int
totalSize int64
}
// InspectFile opens path and inspects it as a Pages deployment package without
// loading the whole archive into memory. File inventory uses declared sizes;
// per-file content hashes are not computed.
// loading the whole archive or any tar member body into memory.
func InspectFile(filePath string, format Format, opts InspectOptions) (*Manifest, error) {
file, err := os.Open(filePath) //nolint:gosec // filePath is a controlled temp upload path
if err != nil {
@@ -68,55 +78,137 @@ func InspectBytes(data []byte, format Format, opts InspectOptions) (*Manifest, e
}
func inspectFromReaderAt(ra io.ReaderAt, size int64, format Format, opts InspectOptions) (*Manifest, error) {
// Default: zip/7z use central directory only; tar streams headers and discards bodies.
// VerifySizes needs openable tar bodies, so materialize only when requested.
entries, err := listEntriesAt(ra, size, format, opts.VerifySizes)
limits := normalizeLimits(opts.Limits)
var (
measured *measuredArchive
err error
)
if isTarFamily(format) {
measured, err = scanTarFamilyAt(ra, size, format, limits, true)
} else {
var entries []Entry
entries, err = listRandomAccessEntriesAt(ra, size, format)
if err == nil {
measured, err = inspectRandomAccessEntries(entries, limits)
}
}
if err != nil {
return nil, err
}
return buildManifest(entries, opts)
return buildMeasuredManifest(measured, opts)
}
func buildManifest(entries []Entry, opts InspectOptions) (*Manifest, error) {
limits := normalizeLimits(opts.Limits)
commonPrefix := FindCommonRootPrefix(collectFileNames(entries))
targetEntryPath := resolveTargetEntryPath(opts.RootDir, opts.EntryFile)
manifest := &Manifest{Files: make([]FileEntry, 0)}
entrySeen := false
func inspectRandomAccessEntries(entries []Entry, limits Limits) (*measuredArchive, error) {
measured := &measuredArchive{files: make([]measuredFile, 0, len(entries))}
for _, entry := range entries {
normalizedPath, skip, err := prepareEntryPath(entry, commonPrefix)
normalizedPath, skip, err := validateArchiveEntry(entry)
if err != nil {
return nil, err
}
if skip {
continue
}
if exceedsFileByteLimit(entry.Size, limits.MaxFileBytes) {
return nil, fmt.Errorf("pages file too large: %s", normalizedPath)
if err := prepareMeasuredFile(measured, normalizedPath, entry.Size, limits, true); err != nil {
return nil, err
}
if entry.Open == nil {
return nil, fmt.Errorf("%s: pages archive member cannot be opened", normalizedPath)
}
src, err := entry.Open()
if err != nil {
return nil, fmt.Errorf("%s: %w", normalizedPath, err)
}
maxBytes := effectiveFileLimit(limits, measured.totalSize, true)
actual, copyErr := copyAndVerifySize(io.Discard, src, entry.Size, maxBytes)
closeErr := src.Close()
if err := errors.Join(copyErr, closeErr); err != nil {
return nil, fmt.Errorf("%s: %w", normalizedPath, err)
}
appendMeasuredFile(measured, normalizedPath, actual)
}
return measured, nil
}
fileEntry, err := inspectRegularFile(entry, normalizedPath, limits, opts.VerifySizes)
func scanTarFamilyAt(
ra io.ReaderAt,
size int64,
format Format,
limits Limits,
enforceLimits bool,
) (*measuredArchive, error) {
if size < 0 {
return nil, fmt.Errorf("invalid pages package size")
}
tarReader, closeReader, err := openTarFamilyReader(io.NewSectionReader(ra, 0, size), format)
if err != nil {
return nil, err
}
measured, scanErr := scanTarReader(tarReader, limits, enforceLimits)
if closeErr := closeReader(); closeErr != nil {
scanErr = errors.Join(scanErr, closeErr)
}
return measured, scanErr
}
func scanTarReader(tarReader *tar.Reader, limits Limits, enforceLimits bool) (*measuredArchive, error) {
measured := &measuredArchive{files: make([]measuredFile, 0)}
for {
header, err := tarReader.Next()
if err == io.EOF {
break
}
if err != nil {
return nil, fmt.Errorf("read tar pages package: %w", err)
}
entry := entryFromTarHeader(header)
normalizedPath, skip, err := validateArchiveEntry(entry)
if err != nil {
return nil, err
}
manifest.FileCount++
if manifest.FileCount > limits.MaxFiles {
return nil, fmt.Errorf("pages deployment file count exceeds %d", limits.MaxFiles)
if skip {
continue
}
manifest.TotalSize += fileEntry.Size
if manifest.TotalSize > limits.MaxTotalBytes {
return nil, fmt.Errorf("pages extracted size exceeds limit")
if err := prepareMeasuredFile(measured, normalizedPath, entry.Size, limits, enforceLimits); err != nil {
return nil, err
}
maxBytes := effectiveFileLimit(limits, measured.totalSize, enforceLimits)
actual, err := copyAndVerifySize(io.Discard, tarReader, entry.Size, maxBytes)
if err != nil {
return nil, fmt.Errorf("%s: %w", normalizedPath, err)
}
appendMeasuredFile(measured, normalizedPath, actual)
}
return measured, nil
}
func buildMeasuredManifest(measured *measuredArchive, opts InspectOptions) (*Manifest, error) {
if measured == nil || measured.fileCount == 0 {
return nil, fmt.Errorf("pages package is empty")
}
targetEntryPath, err := resolveTargetEntryPath(opts.RootDir, opts.EntryFile)
if err != nil {
return nil, err
}
paths := make([]string, 0, len(measured.files))
for _, file := range measured.files {
paths = append(paths, file.path)
}
commonPrefix := FindCommonRootPrefix(paths)
manifest := &Manifest{
Files: make([]FileEntry, 0, measured.fileCount),
FileCount: measured.fileCount,
TotalSize: measured.totalSize,
}
entrySeen := false
for _, file := range measured.files {
normalizedPath := StripPrefix(file.path, commonPrefix)
if normalizedPath == targetEntryPath {
entrySeen = true
}
manifest.Files = append(manifest.Files, fileEntry)
}
if manifest.FileCount == 0 {
return nil, fmt.Errorf("pages package is empty")
manifest.Files = append(manifest.Files, FileEntry{
Path: normalizedPath,
Size: file.size,
})
}
if !entrySeen {
return nil, fmt.Errorf("pages package is missing entry file %s", targetEntryPath)
@@ -124,82 +216,86 @@ func buildManifest(entries []Entry, opts InspectOptions) (*Manifest, error) {
return manifest, nil
}
func collectFileNames(entries []Entry) []string {
names := make([]string, 0, len(entries))
for _, entry := range entries {
if entry.IsDir || entry.IsSymlink {
continue
}
names = append(names, entry.Name)
func prepareMeasuredFile(measured *measuredArchive, normalizedPath string, declaredSize uint64, limits Limits, enforceLimits bool) error {
if declaredSize > uint64(math.MaxInt64) {
return fmt.Errorf("%s: pages file size out of bounds", normalizedPath)
}
return names
if !enforceLimits {
return nil
}
if measured.fileCount >= limits.MaxFiles {
return fmt.Errorf("pages deployment file count exceeds %d", limits.MaxFiles)
}
if exceedsFileByteLimit(declaredSize, limits.MaxFileBytes) {
return fmt.Errorf("pages file too large: %s", normalizedPath)
}
remaining := limits.MaxTotalBytes - measured.totalSize
if remaining < 0 || declaredSize > uint64(remaining) { //nolint:gosec // remaining is checked non-negative
return fmt.Errorf("pages extracted size exceeds limit")
}
return nil
}
func resolveTargetEntryPath(rootDir, entryFile string) string {
normalizedEntry := strings.TrimSpace(entryFile)
if normalizedEntry == "" {
normalizedEntry = "index.html"
func appendMeasuredFile(measured *measuredArchive, normalizedPath string, actual int64) {
measured.files = append(measured.files, measuredFile{path: normalizedPath, size: actual})
measured.fileCount++
measured.totalSize += actual
}
func effectiveFileLimit(limits Limits, totalSize int64, enforceLimits bool) int64 {
if !enforceLimits {
return -1
}
remaining := limits.MaxTotalBytes - totalSize
if remaining < limits.MaxFileBytes {
return remaining
}
return limits.MaxFileBytes
}
func resolveTargetEntryPath(rootDir, entryFile string) (string, error) {
normalizedRoot, err := NormalizeLogicalPath(rootDir, true)
if err != nil {
return "", fmt.Errorf("invalid pages root directory: %w", err)
}
if entryFile == "" {
entryFile = "index.html"
}
normalizedEntry, err := NormalizeLogicalPath(entryFile, false)
if err != nil {
return "", fmt.Errorf("invalid pages entry file: %w", err)
}
normalizedRoot := strings.Trim(strings.TrimSpace(rootDir), "/")
if normalizedRoot == "" {
return normalizedEntry
return normalizedEntry, nil
}
return path.Join(normalizedRoot, normalizedEntry)
return path.Join(normalizedRoot, normalizedEntry), nil
}
func prepareEntryPath(entry Entry, commonPrefix string) (string, bool, error) {
func validateArchiveEntry(entry Entry) (string, bool, error) {
normalizedPath, skip, err := NormalizeEntryPath(entry.Name)
if err != nil {
return "", false, err
}
if skip || entry.IsDir {
return "", true, nil
}
normalizedPath = StripPrefix(normalizedPath, commonPrefix)
if entry.IsSymlink {
return "", false, fmt.Errorf("pages package contains unsupported symlink: %s", normalizedPath)
}
if entry.IsHardlink {
return "", false, fmt.Errorf("pages package contains unsupported hardlink: %s", normalizedPath)
}
if entry.IsSpecial {
return "", false, fmt.Errorf("pages package contains unsupported special entry: %s", normalizedPath)
}
if skip || entry.IsDir {
return normalizedPath, true, nil
}
return normalizedPath, false, nil
}
func inspectRegularFile(entry Entry, normalizedPath string, limits Limits, verifySizes bool) (FileEntry, error) {
if entry.Size > uint64(math.MaxInt64) {
return FileEntry{}, fmt.Errorf("%s: pages file size out of bounds", normalizedPath)
func isTarFamily(format Format) bool {
switch format {
case FormatTar, FormatTarGz, FormatTarXz, FormatTarBz2:
return true
default:
return false
}
//nolint:gosec // bounded to MaxInt64 above
declaredSize := int64(entry.Size)
if !verifySizes {
return FileEntry{
Path: normalizedPath,
Size: declaredSize,
Checksum: "",
}, nil
}
if entry.Open == nil {
return FileEntry{}, fmt.Errorf("%s: cannot verify size without entry open", normalizedPath)
}
src, err := entry.Open()
if err != nil {
return FileEntry{}, fmt.Errorf("%s: %w", normalizedPath, err)
}
actual, measureErr := measureReader(src, entry.Size, limits.MaxFileBytes)
_ = src.Close()
if measureErr != nil {
return FileEntry{}, fmt.Errorf("%s: %w", normalizedPath, measureErr)
}
if declaredSize > 0 && actual != declaredSize {
return FileEntry{}, fmt.Errorf("%s: declared size %d does not match actual %d", normalizedPath, declaredSize, actual)
}
return FileEntry{
Path: normalizedPath,
Size: actual,
Checksum: "",
}, nil
}
// measureReader counts bytes without hashing, enforcing maxBytes when positive.
func measureReader(src io.Reader, declaredSize uint64, maxBytes int64) (int64, error) {
return copyLimited(io.Discard, src, declaredSize, maxBytes)
}
+49 -172
View File
@@ -6,7 +6,6 @@ package pagesarchive
import (
"archive/tar"
"archive/zip"
"bytes"
"compress/bzip2"
"compress/gzip"
"fmt"
@@ -19,8 +18,8 @@ import (
type archiveFile interface {
Name() string
Mode() os.FileMode
IsDir() bool
IsSymlink() bool
Size() uint64
Open() (io.ReadCloser, error)
}
@@ -29,12 +28,10 @@ type zipArchiveFile struct {
file *zip.File
}
func (z zipArchiveFile) Name() string { return z.file.Name }
func (z zipArchiveFile) IsDir() bool { return z.file.FileInfo().IsDir() }
func (z zipArchiveFile) IsSymlink() bool {
return z.file.Mode()&os.ModeSymlink != 0
}
func (z zipArchiveFile) Size() uint64 { return z.file.UncompressedSize64 }
func (z zipArchiveFile) Name() string { return z.file.Name }
func (z zipArchiveFile) Mode() os.FileMode { return z.file.Mode() }
func (z zipArchiveFile) IsDir() bool { return z.file.FileInfo().IsDir() }
func (z zipArchiveFile) Size() uint64 { return z.file.UncompressedSize64 }
func (z zipArchiveFile) Open() (io.ReadCloser, error) {
return z.file.Open()
}
@@ -43,65 +40,27 @@ type sevenZipArchiveFile struct {
file *sevenzip.File
}
func (z sevenZipArchiveFile) Name() string { return z.file.Name }
func (z sevenZipArchiveFile) IsDir() bool { return z.file.FileInfo().IsDir() }
func (z sevenZipArchiveFile) IsSymlink() bool {
return z.file.Mode()&os.ModeSymlink != 0
}
func (z sevenZipArchiveFile) Size() uint64 { return z.file.UncompressedSize }
func (z sevenZipArchiveFile) Name() string { return z.file.Name }
func (z sevenZipArchiveFile) Mode() os.FileMode { return z.file.Mode() }
func (z sevenZipArchiveFile) IsDir() bool { return z.file.FileInfo().IsDir() }
func (z sevenZipArchiveFile) Size() uint64 { return z.file.UncompressedSize }
func (z sevenZipArchiveFile) Open() (io.ReadCloser, error) {
return z.file.Open()
}
// listEntriesAt lists archive members from a random-access source.
// When materializeBodies is true, tar-family streams buffer regular-file bodies so Entry.Open works.
// When false (inspect path), tar bodies are discarded after reading headers; zip/7z only use central directory metadata.
func listEntriesAt(ra io.ReaderAt, size int64, format Format, materializeBodies bool) ([]Entry, error) {
// listRandomAccessEntriesAt lists zip/7z members without reading their bodies.
// Tar-family archives use the sequential streaming paths in inspect.go/extract.go.
func listRandomAccessEntriesAt(ra io.ReaderAt, size int64, format Format) ([]Entry, error) {
if size < 0 {
return nil, fmt.Errorf("invalid pages package size")
}
switch format {
case FormatZip:
return listZipEntriesAt(ra, size)
case FormatTar:
return listTarFamily(io.NewSectionReader(ra, 0, size), FormatTar, materializeBodies)
case FormatTarGz:
return listTarFamily(io.NewSectionReader(ra, 0, size), FormatTarGz, materializeBodies)
case FormatTarXz:
return listTarFamily(io.NewSectionReader(ra, 0, size), FormatTarXz, materializeBodies)
case FormatTarBz2:
return listTarFamily(io.NewSectionReader(ra, 0, size), FormatTarBz2, materializeBodies)
case FormatSevenZip:
return listSevenZipEntriesAt(ra, size)
default:
return nil, fmt.Errorf("unsupported pages package format: %s", format)
}
}
func listTarFamily(r io.Reader, format Format, materializeBodies bool) ([]Entry, error) {
switch format {
case FormatTar:
if materializeBodies {
return listTarEntries(r, true)
}
return listTarEntries(r, false)
case FormatTarGz:
gzReader, err := gzip.NewReader(r)
if err != nil {
return nil, fmt.Errorf("open gzip pages package: %w", err)
}
defer func() { _ = gzReader.Close() }()
return listTarEntries(gzReader, materializeBodies)
case FormatTarXz:
xzReader, err := xz.NewReader(r)
if err != nil {
return nil, fmt.Errorf("open xz pages package: %w", err)
}
return listTarEntries(xzReader, materializeBodies)
case FormatTarBz2:
return listTarEntries(bzip2.NewReader(r), materializeBodies)
default:
return nil, fmt.Errorf("unsupported tar family format: %s", format)
return nil, fmt.Errorf("unsupported random-access pages package format: %s", format)
}
}
@@ -109,10 +68,15 @@ func entriesFromArchiveFiles(files []archiveFile) []Entry {
entries := make([]Entry, 0, len(files))
for _, item := range files {
file := item
mode := file.Mode()
isDir := file.IsDir()
isSymlink := mode&os.ModeSymlink != 0
isSpecial := !isDir && !isSymlink && !mode.IsRegular()
entries = append(entries, Entry{
Name: file.Name(),
IsDir: file.IsDir(),
IsSymlink: file.IsSymlink(),
IsDir: isDir,
IsSymlink: isSymlink,
IsSpecial: isSpecial,
Size: file.Size(),
Open: file.Open,
})
@@ -144,132 +108,45 @@ func listSevenZipEntriesAt(ra io.ReaderAt, size int64) ([]Entry, error) {
return entriesFromArchiveFiles(files), nil
}
func listTarEntries(r io.Reader, materializeBodies bool) ([]Entry, error) {
tarReader := tar.NewReader(r)
type materialised struct {
header *tar.Header
body []byte
}
items := make([]materialised, 0)
for {
header, err := tarReader.Next()
if err == io.EOF {
break
}
func openTarFamilyReader(r io.Reader, format Format) (*tar.Reader, func() error, error) {
switch format {
case FormatTar:
return tar.NewReader(r), func() error { return nil }, nil
case FormatTarGz:
gzReader, err := gzip.NewReader(r)
if err != nil {
return nil, fmt.Errorf("read tar pages package: %w", err)
return nil, nil, fmt.Errorf("open gzip pages package: %w", err)
}
item, skip, err := readTarHeader(tarReader, header, materializeBodies)
return tar.NewReader(gzReader), gzReader.Close, nil
case FormatTarXz:
xzReader, err := xz.NewReader(r)
if err != nil {
return nil, err
return nil, nil, fmt.Errorf("open xz pages package: %w", err)
}
if skip {
continue
}
items = append(items, item)
}
entries := make([]Entry, 0, len(items))
for _, item := range items {
entries = append(entries, tarEntryFromHeader(item.header, item.body, materializeBodies))
}
return entries, nil
}
func readTarHeader(tarReader *tar.Reader, header *tar.Header, materializeBodies bool) (item struct {
header *tar.Header
body []byte
}, skip bool, err error) {
switch header.Typeflag {
case tar.TypeDir, tar.TypeSymlink, tar.TypeLink:
return struct {
header *tar.Header
body []byte
}{header: header}, false, nil
case tar.TypeReg, tar.TypeRegA: //nolint:staticcheck // TypeRegA still appears in older archives
if !materializeBodies {
if err := discardTarBody(tarReader, header); err != nil {
return item, false, err
}
return struct {
header *tar.Header
body []byte
}{header: header}, false, nil
}
body, readErr := readTarBody(tarReader, header)
if readErr != nil {
return item, false, readErr
}
return struct {
header *tar.Header
body []byte
}{header: header, body: body}, false, nil
return tar.NewReader(xzReader), func() error { return nil }, nil
case FormatTarBz2:
return tar.NewReader(bzip2.NewReader(r)), func() error { return nil }, nil
default:
if header.Size > 0 {
if _, copyErr := io.CopyN(io.Discard, tarReader, header.Size); copyErr != nil {
return item, false, fmt.Errorf("skip tar entry %s: %w", header.Name, copyErr)
}
}
return item, true, nil
return nil, nil, fmt.Errorf("unsupported tar family format: %s", format)
}
}
func discardTarBody(tarReader *tar.Reader, header *tar.Header) error {
if header.Size <= 0 {
_, err := io.Copy(io.Discard, tarReader)
if err != nil {
return fmt.Errorf("discard tar entry %s: %w", header.Name, err)
}
return nil
}
if _, err := io.CopyN(io.Discard, tarReader, header.Size); err != nil {
return fmt.Errorf("discard tar entry %s: %w", header.Name, err)
}
return nil
}
func readTarBody(tarReader *tar.Reader, header *tar.Header) ([]byte, error) {
func entryFromTarHeader(header *tar.Header) Entry {
entry := Entry{Name: header.Name}
if header.Size > 0 {
body := make([]byte, header.Size)
if _, err := io.ReadFull(tarReader, body); err != nil {
return nil, fmt.Errorf("read tar entry %s: %w", header.Name, err)
}
return body, nil
entry.Size = uint64(header.Size) //nolint:gosec // archive/tar rejects negative sizes
}
body, err := io.ReadAll(tarReader)
if err != nil {
return nil, fmt.Errorf("read tar entry %s: %w", header.Name, err)
}
return body, nil
}
func tarEntryFromHeader(header *tar.Header, body []byte, materializeBodies bool) Entry {
size := header.Size
if materializeBodies && int64(len(body)) > size {
size = int64(len(body))
}
entry := Entry{
Name: header.Name,
IsDir: header.Typeflag == tar.TypeDir,
IsSymlink: header.Typeflag == tar.TypeSymlink || header.Typeflag == tar.TypeLink,
Size: uint64(size), //nolint:gosec // non-negative sizes
}
if entry.IsDir || entry.IsSymlink {
entry.Open = func() (io.ReadCloser, error) {
return io.NopCloser(bytes.NewReader(nil)), nil
}
return entry
}
if materializeBodies {
bodyCopy := body
entry.Open = func() (io.ReadCloser, error) {
return io.NopCloser(bytes.NewReader(bodyCopy)), nil
}
return entry
}
// Inspect path: body not retained; Open is unavailable.
entry.Open = func() (io.ReadCloser, error) {
return nil, fmt.Errorf("tar entry body not materialized: %s", header.Name)
switch header.Typeflag {
case tar.TypeReg, tar.TypeRegA: //nolint:staticcheck // TypeRegA appears in older archives
// Regular file.
case tar.TypeDir:
entry.IsDir = true
case tar.TypeSymlink:
entry.IsSymlink = true
case tar.TypeLink:
entry.IsHardlink = true
default:
entry.IsSpecial = true
}
return entry
}
+74 -18
View File
@@ -6,33 +6,89 @@ package pagesarchive
import (
"fmt"
"path"
"path/filepath"
"strings"
"unicode"
"unicode/utf8"
)
// NormalizeLogicalPath validates and normalizes a relative POSIX path.
// Empty input is returned unchanged only when allowEmpty is true.
func NormalizeLogicalPath(raw string, allowEmpty bool) (string, error) {
if raw == "" {
if allowEmpty {
return "", nil
}
return "", fmt.Errorf("pages path is required")
}
if err := validateLogicalPathText(raw); err != nil {
return "", err
}
cleaned := path.Clean(raw)
if cleaned == "." || cleaned == "" {
if allowEmpty {
return "", nil
}
return "", fmt.Errorf("pages path is required")
}
if strings.HasPrefix(cleaned, "/") || cleaned == ".." || strings.HasPrefix(cleaned, "../") {
return "", fmt.Errorf("pages path escapes directory: %s", raw)
}
return cleaned, nil
}
func validateLogicalPathText(raw string) error {
if !utf8.ValidString(raw) {
return fmt.Errorf("pages path is not valid UTF-8")
}
if strings.Contains(raw, "\\") {
return fmt.Errorf("pages path must use POSIX separators: %s", raw)
}
if strings.HasPrefix(raw, "/") || path.IsAbs(raw) {
return fmt.Errorf("pages path must be relative: %s", raw)
}
if err := validateLogicalPathRunes(raw); err != nil {
return err
}
return validateLogicalPathSegments(raw)
}
func validateLogicalPathRunes(raw string) error {
for _, r := range raw {
if r == 0 || unicode.IsControl(r) {
return fmt.Errorf("pages path contains a control character")
}
if r == '\'' || r == '"' || r == ';' {
return fmt.Errorf("pages path contains an unsupported character: %s", raw)
}
}
return nil
}
func validateLogicalPathSegments(raw string) error {
for _, segment := range strings.Split(raw, "/") {
if len(segment) >= 2 && segment[1] == ':' {
return fmt.Errorf("pages path contains a Windows drive: %s", raw)
}
if segment == "." || segment == ".." {
return fmt.Errorf("pages path escapes directory or contains a dot segment: %s", raw)
}
}
return nil
}
// NormalizeEntryPath cleans an archive entry path and rejects zip-slip / absolute paths.
// skip=true means the entry should be ignored (empty path or directory marker).
func NormalizeEntryPath(raw string) (cleaned string, skip bool, err error) {
name := strings.TrimSpace(filepath.ToSlash(raw))
if name == "" {
if raw == "" {
return "", true, nil
}
if strings.HasSuffix(name, "/") {
return "", true, nil
cleanedPath, normalizeErr := NormalizeLogicalPath(raw, false)
if normalizeErr != nil {
return "", false, fmt.Errorf("invalid pages package path %q: %w", raw, normalizeErr)
}
if strings.HasPrefix(name, "/") || path.IsAbs(name) {
return "", false, fmt.Errorf("pages package contains absolute path: %s", raw)
}
// Reject Windows drive / UNC-style paths that may appear after ToSlash.
if len(name) >= 2 && name[1] == ':' {
return "", false, fmt.Errorf("pages package contains absolute path: %s", raw)
}
cleanedPath := path.Clean(name)
if cleanedPath == "." {
return "", true, nil
}
if cleanedPath == ".." || strings.HasPrefix(cleanedPath, "../") || strings.Contains(cleanedPath, "/../") {
return "", false, fmt.Errorf("pages package path escapes directory: %s", raw)
if strings.HasSuffix(raw, "/") {
return cleanedPath, true, nil
}
return cleanedPath, false, nil
}
+499
View File
@@ -0,0 +1,499 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pagesarchive
import (
"archive/tar"
"archive/zip"
"bytes"
"encoding/base64"
"io"
"os"
"path/filepath"
"strings"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)
var testLimits = Limits{
MaxFiles: 100,
MaxFileBytes: 1 << 20,
MaxTotalBytes: 1 << 20,
}
func TestNormalizeLogicalPathStrict(t *testing.T) {
t.Parallel()
valid := []struct {
name string
raw string
allowEmpty bool
want string
}{
{name: "empty root", allowEmpty: true},
{name: "single file", raw: "index.html", want: "index.html"},
{name: "nested posix", raw: "public/assets/app.js", want: "public/assets/app.js"},
{name: "unicode", raw: "静态/首页.html", want: "静态/首页.html"},
{name: "repeated separator is normalized", raw: "public//app.js", want: "public/app.js"},
}
for _, tt := range valid {
t.Run(tt.name, func(t *testing.T) {
got, err := NormalizeLogicalPath(tt.raw, tt.allowEmpty)
require.NoError(t, err)
assert.Equal(t, tt.want, got)
})
}
invalidUTF8 := string([]byte{'a', '/', 0xff})
invalid := []struct {
name string
raw string
}{
{name: "empty entry"},
{name: "absolute", raw: "/etc/passwd"},
{name: "unc", raw: "//server/share"},
{name: "windows drive", raw: "C:/site/index.html"},
{name: "nested windows drive", raw: "site/C:/index.html"},
{name: "windows separator", raw: `site\index.html`},
{name: "windows unc", raw: `\\server\share`},
{name: "parent segment", raw: "../index.html"},
{name: "nested parent segment", raw: "site/../index.html"},
{name: "current segment", raw: "site/./index.html"},
{name: "nul", raw: "site/\x00index.html"},
{name: "newline", raw: "site/\nindex.html"},
{name: "delete control", raw: "site/\x7findex.html"},
{name: "single quote", raw: "site/'index.html"},
{name: "double quote", raw: `site/"index.html`},
{name: "semicolon", raw: "site/;index.html"},
{name: "invalid utf8", raw: invalidUTF8},
}
for _, tt := range invalid {
t.Run(tt.name, func(t *testing.T) {
_, err := NormalizeLogicalPath(tt.raw, false)
require.Error(t, err)
})
}
cleaned, skip, err := NormalizeEntryPath("")
require.NoError(t, err)
assert.Empty(t, cleaned)
assert.True(t, skip)
cleaned, skip, err = NormalizeEntryPath("assets/")
require.NoError(t, err)
assert.Equal(t, "assets", cleaned)
assert.True(t, skip)
}
func TestSupportedFormatsInspectAndExtract(t *testing.T) {
t.Parallel()
sevenZipData := decodeFixture(t, "N3q8ryccAASgR6WICAAAAAAAAABmAAAAAAAAAN2R8/FiYXIKZm9vCgEEBgACCQQEAAcLAgABAQABAQAMBAQACAoB6bOiBKhlMn4AAAUCGQUAAAAAABERAGIAYQByAAAAZgBvAG8AAAAZAgAAFBIBAACFM3PyY9YBAFgCcvJj1gEVCgEAIICkgSCApIEAAA==")
bzipTarData := decodeFixture(t, "QlpoOTFBWSZTWYp5f6EAAHV//P64A8RQAf/iOm/9cO/v/9AAAgBADlAABAADAAgwAU1RIZJpNCaammmnqbSGTI9Q0BoBpp6mjIaGmmRoaHGRpkxNBkyYTTIGQ0BoDTJoYATQGG1KCntExT01MhoAABoAAHqAAPU9QacVDtN45fA6MmuGVQlWowrijpZgwASITYSPUcpJpoQGMkKq69jMkUR6L86R5j0IySUaZEjazEqhQ9E8vuuxsmWZQLCA84jNsobYNEzuEB1eCPhw8nc2AOz+xrCY5hVxQW1IIokpfSRKi+McvXU+QoYuEg6BD4w8x3K0imi+bULpkLCylCZ4lzoGlTQgibvG67sQcrTCRBTbBCVL7zC0q0qULmK/WOneu94s9cs4s4K98SjY2YvpdZvl42kwtxvvPMheorYQ2pcxyF4sNQYvd4+bgqm5gKXElqnGF3jhxGTeXp9eCUxWVlbi9ikxAik4xxATl7cJrISVWnHwUFiLdhEnKWw0Lhm3ZyKlX7P5Wj7b9TLAmWBaAwH/F3JFOFCQinl/oQ==")
cases := []struct {
name string
format Format
data []byte
entryFile string
wantPath string
}{
{name: "zip", format: FormatZip, data: testZip(t, map[string]string{"bundle/index.html": "zip"}), entryFile: "index.html", wantPath: "index.html"},
{name: "tar", format: FormatTar, data: testTar(t, map[string]string{"bundle/index.html": "tar"}), entryFile: "index.html", wantPath: "index.html"},
{name: "tar gzip", format: FormatTarGz, data: testTarGz(t, map[string]string{"bundle/index.html": "gzip"}), entryFile: "index.html", wantPath: "index.html"},
{name: "tar xz", format: FormatTarXz, data: testTarXz(t, map[string]string{"bundle/index.html": "xz"}), entryFile: "index.html", wantPath: "index.html"},
{name: "tar bzip2", format: FormatTarBz2, data: bzipTarData, entryFile: "index.html", wantPath: "index.html"},
{name: "7z", format: FormatSevenZip, data: sevenZipData, entryFile: "foo", wantPath: "foo"},
}
for _, tt := range cases {
t.Run(tt.name, func(t *testing.T) {
manifest, err := InspectBytes(tt.data, tt.format, InspectOptions{
EntryFile: tt.entryFile,
Limits: testLimits,
})
require.NoError(t, err)
assert.Positive(t, manifest.FileCount)
assertManifestContains(t, manifest, tt.wantPath)
destDir := t.TempDir()
require.NoError(t, ExtractBytes(tt.data, tt.format, destDir, ExtractOptions{
Limits: testLimits,
StripCommonRoot: true,
EnforceLimits: true,
}))
_, err = os.Stat(filepath.Join(destDir, filepath.FromSlash(tt.wantPath)))
require.NoError(t, err)
})
}
}
func TestExtractFilePreservesCommonRootAndEnforcesLimits(t *testing.T) {
t.Parallel()
data := testTarGz(t, map[string]string{
"repository/dist/index.html": "pages",
"repository/dist/app.js": "javascript",
})
archivePath := filepath.Join(t.TempDir(), "site.tar.gz")
require.NoError(t, os.WriteFile(archivePath, data, 0o600))
manifest, err := InspectFile(archivePath, FormatTarGz, InspectOptions{
RootDir: "dist",
EntryFile: "index.html",
Limits: testLimits,
})
require.NoError(t, err)
assertManifestContains(t, manifest, "dist/index.html")
destDir := t.TempDir()
require.NoError(t, ExtractFile(archivePath, FormatTarGz, destDir, ExtractOptions{
Limits: testLimits,
StripCommonRoot: true,
EnforceLimits: true,
}))
body, err := os.ReadFile(filepath.Join(destDir, "dist", "index.html")) //nolint:gosec
require.NoError(t, err)
assert.Equal(t, "pages", string(body))
err = ExtractFile(archivePath, FormatTarGz, t.TempDir(), ExtractOptions{
Limits: Limits{
MaxFiles: 10,
MaxFileBytes: 5,
MaxTotalBytes: 1 << 20,
},
StripCommonRoot: true,
EnforceLimits: true,
})
require.ErrorContains(t, err, "file too large")
err = ExtractFile(archivePath, FormatTarGz, t.TempDir(), ExtractOptions{
Limits: Limits{
MaxFiles: 10,
MaxFileBytes: 1 << 20,
MaxTotalBytes: int64(len("pages") + len("javascript") - 1),
},
StripCommonRoot: true,
EnforceLimits: true,
})
require.ErrorContains(t, err, "extracted size exceeds limit")
}
func TestRandomAccessMembersVerifyActualSize(t *testing.T) {
t.Parallel()
entry := Entry{
Name: "index.html",
Size: 1,
Open: func() (io.ReadCloser, error) {
return io.NopCloser(strings.NewReader("actual-body")), nil
},
}
_, err := inspectRandomAccessEntries([]Entry{entry}, testLimits)
require.ErrorContains(t, err, "declared size 1 does not match actual 11")
destDir := t.TempDir()
err = extractEntries([]Entry{entry}, destDir, ExtractOptions{
Limits: testLimits,
EnforceLimits: true,
})
require.ErrorContains(t, err, "declared size 1 does not match actual 11")
_, statErr := os.Stat(filepath.Join(destDir, "index.html"))
assert.ErrorIs(t, statErr, os.ErrNotExist, "failed extraction must remove the partial file")
}
func TestActualByteLimitAbortsReaderEarly(t *testing.T) {
t.Parallel()
reader := &countingFillReader{remaining: 1 << 30}
entry := Entry{
Name: "index.html",
Size: 0,
Open: func() (io.ReadCloser, error) {
return io.NopCloser(reader), nil
},
}
_, err := inspectRandomAccessEntries([]Entry{entry}, Limits{
MaxFiles: 1,
MaxFileBytes: 32,
MaxTotalBytes: 32,
})
require.ErrorContains(t, err, "size out of bounds")
assert.LessOrEqual(t, reader.read, int64(33), "inspection must stop after limit+1 actual bytes")
tarData := tarWithDeclaredBodyOnly(t, "index.html", 1<<30)
_, err = InspectBytes(tarData, FormatTar, InspectOptions{
EntryFile: "index.html",
Limits: Limits{
MaxFiles: 1,
MaxFileBytes: 32,
MaxTotalBytes: 32,
},
})
require.ErrorContains(t, err, "file too large")
}
func TestArchiveLimitsUseFilesAndActualTotals(t *testing.T) {
t.Parallel()
data := testZip(t, map[string]string{
"index.html": "1234",
"app.js": "5678",
})
_, err := InspectBytes(data, FormatZip, InspectOptions{
EntryFile: "index.html",
Limits: Limits{
MaxFiles: 1,
MaxFileBytes: 8,
MaxTotalBytes: 16,
},
})
require.ErrorContains(t, err, "file count exceeds")
_, err = InspectBytes(data, FormatZip, InspectOptions{
EntryFile: "index.html",
Limits: Limits{
MaxFiles: 2,
MaxFileBytes: 8,
MaxTotalBytes: 7,
},
})
require.ErrorContains(t, err, "extracted size exceeds limit")
sevenZipData := decodeFixture(t, "N3q8ryccAASgR6WICAAAAAAAAABmAAAAAAAAAN2R8/FiYXIKZm9vCgEEBgACCQQEAAcLAgABAQABAQAMBAQACAoB6bOiBKhlMn4AAAUCGQUAAAAAABERAGIAYQByAAAAZgBvAG8AAAAZAgAAFBIBAACFM3PyY9YBAFgCcvJj1gEVCgEAIICkgSCApIEAAA==")
_, err = InspectBytes(sevenZipData, FormatSevenZip, InspectOptions{
EntryFile: "foo",
Limits: Limits{
MaxFiles: 10,
MaxFileBytes: 3,
MaxTotalBytes: 32,
},
})
require.ErrorContains(t, err, "file too large")
}
func TestRejectUnsupportedTarMemberTypes(t *testing.T) {
t.Parallel()
cases := []struct {
name string
header tar.Header
wantErr string
}{
{name: "symlink", header: tar.Header{Name: "link", Typeflag: tar.TypeSymlink, Linkname: "index.html"}, wantErr: "unsupported symlink"},
{name: "hardlink", header: tar.Header{Name: "hard", Typeflag: tar.TypeLink, Linkname: "index.html"}, wantErr: "unsupported hardlink"},
{name: "fifo", header: tar.Header{Name: "pipe", Typeflag: tar.TypeFifo, Mode: 0o600}, wantErr: "unsupported special entry"},
{name: "character device", header: tar.Header{Name: "tty", Typeflag: tar.TypeChar, Mode: 0o600, Devmajor: 1, Devminor: 3}, wantErr: "unsupported special entry"},
}
for _, tt := range cases {
t.Run(tt.name, func(t *testing.T) {
data := tarWithSpecialEntry(t, &tt.header)
_, err := InspectBytes(data, FormatTar, InspectOptions{
EntryFile: "index.html",
Limits: testLimits,
})
require.ErrorContains(t, err, tt.wantErr)
destDir := t.TempDir()
err = ExtractBytes(data, FormatTar, destDir, ExtractOptions{
Limits: testLimits,
EnforceLimits: true,
})
require.ErrorContains(t, err, tt.wantErr)
_, statErr := os.Stat(filepath.Join(destDir, "index.html"))
assert.ErrorIs(t, statErr, os.ErrNotExist, "tar validation pass must reject before writing files")
})
}
}
func TestRejectUnsupportedZipMemberTypes(t *testing.T) {
t.Parallel()
cases := []struct {
name string
mode os.FileMode
wantErr string
}{
{name: "symlink", mode: os.ModeSymlink | 0o777, wantErr: "unsupported symlink"},
{name: "named pipe", mode: os.ModeNamedPipe | 0o600, wantErr: "unsupported special entry"},
{name: "device", mode: os.ModeDevice | 0o600, wantErr: "unsupported special entry"},
}
for _, tt := range cases {
t.Run(tt.name, func(t *testing.T) {
data := zipWithSpecialEntry(t, tt.mode)
_, err := InspectBytes(data, FormatZip, InspectOptions{
EntryFile: "index.html",
Limits: testLimits,
})
require.ErrorContains(t, err, tt.wantErr)
err = ExtractBytes(data, FormatZip, t.TempDir(), ExtractOptions{
Limits: testLimits,
EnforceLimits: true,
})
require.ErrorContains(t, err, tt.wantErr)
})
}
}
func TestTarMetadataHeadersRemainTransparent(t *testing.T) {
t.Parallel()
for _, format := range []tar.Format{tar.FormatPAX, tar.FormatGNU} {
format := format
t.Run(format.String(), func(t *testing.T) {
data := tarWithLongMetadata(t, format)
manifest, err := InspectBytes(data, FormatTar, InspectOptions{
EntryFile: "index.html",
Limits: testLimits,
})
require.NoError(t, err)
assert.Equal(t, 2, manifest.FileCount)
assertManifestContains(t, manifest, "index.html")
destDir := t.TempDir()
require.NoError(t, ExtractBytes(data, FormatTar, destDir, ExtractOptions{
Limits: testLimits,
EnforceLimits: true,
}))
_, err = os.Stat(filepath.Join(destDir, "index.html"))
require.NoError(t, err)
})
}
}
type countingFillReader struct {
remaining int64
read int64
}
func (r *countingFillReader) Read(p []byte) (int, error) {
if r.remaining == 0 {
return 0, io.EOF
}
if int64(len(p)) > r.remaining {
p = p[:r.remaining]
}
for i := range p {
p[i] = 'x'
}
r.remaining -= int64(len(p))
r.read += int64(len(p))
return len(p), nil
}
func decodeFixture(t *testing.T, encoded string) []byte {
t.Helper()
data, err := base64.StdEncoding.DecodeString(encoded)
require.NoError(t, err)
return data
}
func assertManifestContains(t *testing.T, manifest *Manifest, path string) {
t.Helper()
for _, file := range manifest.Files {
if file.Path == path {
return
}
}
require.Failf(t, "manifest path missing", "path %q not found in %#v", path, manifest.Files)
}
func testTar(t *testing.T, files map[string]string) []byte {
t.Helper()
var buffer bytes.Buffer
writer := tar.NewWriter(&buffer)
for name, content := range files {
require.NoError(t, writer.WriteHeader(&tar.Header{
Name: name,
Mode: 0o644,
Size: int64(len(content)),
}))
_, err := writer.Write([]byte(content))
require.NoError(t, err)
}
require.NoError(t, writer.Close())
return buffer.Bytes()
}
func tarWithDeclaredBodyOnly(t *testing.T, name string, size int64) []byte {
t.Helper()
var buffer bytes.Buffer
writer := tar.NewWriter(&buffer)
require.NoError(t, writer.WriteHeader(&tar.Header{
Name: name,
Mode: 0o644,
Size: size,
}))
// Deliberately omit the body and trailer. The limit must reject from the
// header before archive/tar attempts to stream the declared body.
return buffer.Bytes()
}
func tarWithSpecialEntry(t *testing.T, special *tar.Header) []byte {
t.Helper()
var buffer bytes.Buffer
writer := tar.NewWriter(&buffer)
require.NoError(t, writer.WriteHeader(&tar.Header{
Name: "index.html",
Mode: 0o644,
Size: 2,
}))
_, err := writer.Write([]byte("ok"))
require.NoError(t, err)
require.NoError(t, writer.WriteHeader(special))
require.NoError(t, writer.Close())
return buffer.Bytes()
}
func zipWithSpecialEntry(t *testing.T, mode os.FileMode) []byte {
t.Helper()
var buffer bytes.Buffer
writer := zip.NewWriter(&buffer)
index, err := writer.Create("index.html")
require.NoError(t, err)
_, err = index.Write([]byte("ok"))
require.NoError(t, err)
header := &zip.FileHeader{Name: "special"}
header.SetMode(mode)
special, err := writer.CreateHeader(header)
require.NoError(t, err)
if mode&os.ModeSymlink != 0 {
_, err = special.Write([]byte("index.html"))
require.NoError(t, err)
}
require.NoError(t, writer.Close())
return buffer.Bytes()
}
func tarWithLongMetadata(t *testing.T, format tar.Format) []byte {
t.Helper()
var buffer bytes.Buffer
writer := tar.NewWriter(&buffer)
longName := strings.Repeat("long-segment-", 12) + "asset.js"
header := &tar.Header{
Name: longName,
Mode: 0o644,
Size: 1,
Format: format,
}
if format == tar.FormatPAX {
header.PAXRecords = map[string]string{"comment": "metadata is not a deployable member"}
}
require.NoError(t, writer.WriteHeader(header))
_, err := writer.Write([]byte("x"))
require.NoError(t, err)
require.NoError(t, writer.WriteHeader(&tar.Header{
Name: "index.html",
Mode: 0o644,
Size: 2,
Format: format,
}))
_, err = writer.Write([]byte("ok"))
require.NoError(t, err)
require.NoError(t, writer.Close())
return buffer.Bytes()
}
+6
View File
@@ -1,3 +1,6 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package protocol defines the communication protocol between OpenFlare server, agent, and relay components.
package protocol
@@ -218,4 +221,7 @@ type PagesProjectLatestHashResponse struct {
ProjectID uint `json:"project_id"`
DeploymentID uint `json:"deployment_id"`
Hash string `json:"hash"`
PackageSize int64 `json:"package_size"`
FileCount int `json:"file_count"`
TotalSize int64 `json:"total_size"`
}
+15
View File
@@ -1,3 +1,6 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package protocol
import (
@@ -47,6 +50,18 @@ func TestAgentProtocolJSONTags(t *testing.T) {
"AccessToken": "agent_token",
},
},
{
name: "PagesProjectLatestHashResponse",
value: PagesProjectLatestHashResponse{},
expected: map[string]string{
"ProjectID": "project_id",
"DeploymentID": "deployment_id",
"Hash": "hash",
"PackageSize": "package_size",
"FileCount": "file_count",
"TotalSize": "total_size",
},
},
}
for _, tc := range cases {