mirror of
https://github.com/Rain-kl/OpenFlare.git
synced 2026-09-29 05:56:38 +08:00
fix(pages): 收紧部署包与 Agent 同步边界
完成 V2 Phase 0 安全与一致性前置:统一真实归档限额、流式拉取、候选裁剪、保留上传删除语义及 Pages 路由引用锁。
This commit is contained in:
@@ -30,6 +30,10 @@ sidebar: false
|
||||
- WAF 规则编辑器支持为节点自定义显示名称,并从节点库拖放到画布指定位置添加节点。
|
||||
- WAF 规则画布支持右键删除节点或连线,并屏蔽浏览器默认右键菜单。
|
||||
|
||||
### 修复
|
||||
|
||||
- 修复 Pages 部署包路径校验、归档展开限额、历史版本裁剪、代理路由绑定与 Agent 下载过程中的安全和一致性问题;大包改为流式处理,部署入口、旧版目录切换、保留版本及上传记录在并发场景下更加可靠。
|
||||
|
||||
## [v3.4.0] - 2026-07-19
|
||||
|
||||
### 新增
|
||||
|
||||
+43
-2
@@ -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
@@ -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": {
|
||||
|
||||
@@ -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: 删除我的文件
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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[:])
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -34,4 +34,6 @@ const (
|
||||
errPagesPackageNotInActiveConfig = "pages 部署尚未进入激活配置"
|
||||
errPagesDeploymentHashMissing = "pages 部署包哈希缺失"
|
||||
errPagesInvalidSnapshotFormat = "配置快照格式无效"
|
||||
errPagesActorMissing = "无法识别当前用户"
|
||||
errPagesEntryFileMissing = "当前激活部署中不存在指定入口文件"
|
||||
)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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) {
|
||||
|
||||
@@ -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"])
|
||||
}
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
},
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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"
|
||||
)
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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"`
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user