From 4e8ec232641c0210a28bfdbade97bc57383ef4d4 Mon Sep 17 00:00:00 2001 From: deqiying Date: Sun, 19 Jul 2026 16:42:45 +0800 Subject: [PATCH] =?UTF-8?q?fix(pages):=20=E6=94=B6=E7=B4=A7=E9=83=A8?= =?UTF-8?q?=E7=BD=B2=E5=8C=85=E4=B8=8E=20Agent=20=E5=90=8C=E6=AD=A5?= =?UTF-8?q?=E8=BE=B9=E7=95=8C?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 完成 V2 Phase 0 安全与一致性前置:统一真实归档限额、流式拉取、候选裁剪、保留上传删除语义及 Pages 路由引用锁。 --- docs/changelog/index.md | 4 + docs/docs.go | 45 +- docs/swagger.json | 45 +- docs/swagger.yaml | 28 + internal/apps/agent/httpclient/client.go | 131 +++- internal/apps/agent/httpclient/client_test.go | 109 +++ internal/apps/agent/sync/pages.go | 659 +++++++++++++++--- internal/apps/agent/sync/pages_stream_test.go | 388 +++++++++++ internal/apps/agent/sync/service.go | 8 +- internal/apps/agent/sync/service_test.go | 146 +++- internal/apps/openflare/agent/routers.go | 9 +- .../config_version/pages_snapshot.go | 30 +- .../config_version/pages_snapshot_test.go | 14 +- internal/apps/openflare/pages/download_url.go | 2 +- internal/apps/openflare/pages/errs.go | 2 + internal/apps/openflare/pages/helpers.go | 115 ++- internal/apps/openflare/pages/logics.go | 432 ++++++++---- internal/apps/openflare/pages/logics_test.go | 236 ++++++- .../apps/openflare/pages/package_metadata.go | 64 ++ .../openflare/pages/package_metadata_test.go | 81 +++ internal/apps/openflare/pages/rebind.go | 29 +- internal/apps/openflare/pages/rebind_test.go | 9 +- internal/apps/openflare/pages/routers.go | 24 +- internal/apps/openflare/pages/routers_test.go | 100 +++ internal/apps/openflare/proxy_route/logics.go | 61 ++ .../apps/openflare/proxy_route/logics_test.go | 55 +- internal/apps/upload/exports.go | 24 +- .../apps/upload/handler/file_management.go | 13 +- internal/apps/upload/handler/routers.go | 5 + internal/apps/upload/handler/routers_test.go | 92 +++ internal/apps/upload/ingest/errors.go | 3 + internal/apps/upload/ingest/helpers.go | 27 +- internal/apps/upload/ingest/ingest_test.go | 235 ++++++- internal/apps/upload/ingest/remove.go | 68 +- internal/apps/upload/shared/constants.go | 2 + internal/apps/upload/shared/errs.go | 1 + internal/apps/upload/task/cleanup.go | 34 +- internal/apps/upload/task/tasks_test.go | 30 +- internal/repository/upload.go | 17 +- pkg/pagesarchive/entry.go | 63 +- pkg/pagesarchive/extract.go | 209 ++++-- pkg/pagesarchive/inspect.go | 282 +++++--- pkg/pagesarchive/list.go | 221 ++---- pkg/pagesarchive/path.go | 92 ++- pkg/pagesarchive/security_test.go | 499 +++++++++++++ pkg/protocol/agent.go | 6 + pkg/protocol/agent_test.go | 15 + 47 files changed, 4005 insertions(+), 759 deletions(-) create mode 100644 internal/apps/agent/httpclient/client_test.go create mode 100644 internal/apps/agent/sync/pages_stream_test.go create mode 100644 internal/apps/openflare/pages/package_metadata.go create mode 100644 internal/apps/openflare/pages/package_metadata_test.go create mode 100644 internal/apps/openflare/pages/routers_test.go create mode 100644 pkg/pagesarchive/security_test.go diff --git a/docs/changelog/index.md b/docs/changelog/index.md index fec2320f..db529121 100644 --- a/docs/changelog/index.md +++ b/docs/changelog/index.md @@ -30,6 +30,10 @@ sidebar: false - WAF 规则编辑器支持为节点自定义显示名称,并从节点库拖放到画布指定位置添加节点。 - WAF 规则画布支持右键删除节点或连线,并屏蔽浏览器默认右键菜单。 +### 修复 + +- 修复 Pages 部署包路径校验、归档展开限额、历史版本裁剪、代理路由绑定与 Agent 下载过程中的安全和一致性问题;大包改为流式处理,部署入口、旧版目录切换、保留版本及上传记录在并发场景下更加可靠。 + ## [v3.4.0] - 2026-07-19 ### 新增 diff --git a/docs/docs.go b/docs/docs.go index 6e2473b2..df744b22 100644 --- a/docs/docs.go +++ b/docs/docs.go @@ -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": { diff --git a/docs/swagger.json b/docs/swagger.json index 7355be74..efdab2a3 100644 --- a/docs/swagger.json +++ b/docs/swagger.json @@ -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": { diff --git a/docs/swagger.yaml b/docs/swagger.yaml index 6bbf8815..a52e0600 100644 --- a/docs/swagger.yaml +++ b/docs/swagger.yaml @@ -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: 删除我的文件 diff --git a/internal/apps/agent/httpclient/client.go b/internal/apps/agent/httpclient/client.go index 34e1398a..4e4e2efb 100644 --- a/internal/apps/agent/httpclient/client.go +++ b/internal/apps/agent/httpclient/client.go @@ -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. diff --git a/internal/apps/agent/httpclient/client_test.go b/internal/apps/agent/httpclient/client_test.go new file mode 100644 index 00000000..613c2dda --- /dev/null +++ b/internal/apps/agent/httpclient/client_test.go @@ -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) + } +} diff --git a/internal/apps/agent/sync/pages.go b/internal/apps/agent/sync/pages.go index 67ab219a..5364ec4d 100644 --- a/internal/apps/agent/sync/pages.go +++ b/internal/apps/agent/sync/pages.go @@ -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[:]) -} diff --git a/internal/apps/agent/sync/pages_stream_test.go b/internal/apps/agent/sync/pages_stream_test.go new file mode 100644 index 00000000..8c16e50d --- /dev/null +++ b/internal/apps/agent/sync/pages_stream_test.go @@ -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) + } +} diff --git a/internal/apps/agent/sync/service.go b/internal/apps/agent/sync/service.go index 30562c91..6c684b02 100644 --- a/internal/apps/agent/sync/service.go +++ b/internal/apps/agent/sync/service.go @@ -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) } diff --git a/internal/apps/agent/sync/service_test.go b/internal/apps/agent/sync/service_test.go index ecec6338..01b0a2fd 100644 --- a/internal/apps/agent/sync/service_test.go +++ b/internal/apps/agent/sync/service_test.go @@ -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[:]) diff --git a/internal/apps/openflare/agent/routers.go b/internal/apps/openflare/agent/routers.go index e39b27be..dd840a63 100644 --- a/internal/apps/openflare/agent/routers.go +++ b/internal/apps/openflare/agent/routers.go @@ -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, })) } diff --git a/internal/apps/openflare/config_version/pages_snapshot.go b/internal/apps/openflare/config_version/pages_snapshot.go index d7012829..b7486f67 100644 --- a/internal/apps/openflare/config_version/pages_snapshot.go +++ b/internal/apps/openflare/config_version/pages_snapshot.go @@ -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 } diff --git a/internal/apps/openflare/config_version/pages_snapshot_test.go b/internal/apps/openflare/config_version/pages_snapshot_test.go index 9b6fb1a9..476a9ed5 100644 --- a/internal/apps/openflare/config_version/pages_snapshot_test.go +++ b/internal/apps/openflare/config_version/pages_snapshot_test.go @@ -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) diff --git a/internal/apps/openflare/pages/download_url.go b/internal/apps/openflare/pages/download_url.go index 9de2938e..5959d7e7 100644 --- a/internal/apps/openflare/pages/download_url.go +++ b/internal/apps/openflare/pages/download_url.go @@ -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" diff --git a/internal/apps/openflare/pages/errs.go b/internal/apps/openflare/pages/errs.go index b4591c78..602825a1 100644 --- a/internal/apps/openflare/pages/errs.go +++ b/internal/apps/openflare/pages/errs.go @@ -34,4 +34,6 @@ const ( errPagesPackageNotInActiveConfig = "pages 部署尚未进入激活配置" errPagesDeploymentHashMissing = "pages 部署包哈希缺失" errPagesInvalidSnapshotFormat = "配置快照格式无效" + errPagesActorMissing = "无法识别当前用户" + errPagesEntryFileMissing = "当前激活部署中不存在指定入口文件" ) diff --git a/internal/apps/openflare/pages/helpers.go b/internal/apps/openflare/pages/helpers.go index b1f2c959..90db9996 100644 --- a/internal/apps/openflare/pages/helpers.go +++ b/internal/apps/openflare/pages/helpers.go @@ -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, diff --git a/internal/apps/openflare/pages/logics.go b/internal/apps/openflare/pages/logics.go index f588c02d..622564bc 100644 --- a/internal/apps/openflare/pages/logics.go +++ b/internal/apps/openflare/pages/logics.go @@ -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 } diff --git a/internal/apps/openflare/pages/logics_test.go b/internal/apps/openflare/pages/logics_test.go index ba3d52ff..230f3352 100644 --- a/internal/apps/openflare/pages/logics_test.go +++ b/internal/apps/openflare/pages/logics_test.go @@ -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() diff --git a/internal/apps/openflare/pages/package_metadata.go b/internal/apps/openflare/pages/package_metadata.go new file mode 100644 index 00000000..0e99c358 --- /dev/null +++ b/internal/apps/openflare/pages/package_metadata.go @@ -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 +} diff --git a/internal/apps/openflare/pages/package_metadata_test.go b/internal/apps/openflare/pages/package_metadata_test.go new file mode 100644 index 00000000..d192a8c9 --- /dev/null +++ b/internal/apps/openflare/pages/package_metadata_test.go @@ -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) + } +} diff --git a/internal/apps/openflare/pages/rebind.go b/internal/apps/openflare/pages/rebind.go index d1950c9e..d26d34d8 100644 --- a/internal/apps/openflare/pages/rebind.go +++ b/internal/apps/openflare/pages/rebind.go @@ -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) { diff --git a/internal/apps/openflare/pages/rebind_test.go b/internal/apps/openflare/pages/rebind_test.go index c37602d1..94dfed4d 100644 --- a/internal/apps/openflare/pages/rebind_test.go +++ b/internal/apps/openflare/pages/rebind_test.go @@ -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"]) } diff --git a/internal/apps/openflare/pages/routers.go b/internal/apps/openflare/pages/routers.go index 5e117fb3..6767f478 100644 --- a/internal/apps/openflare/pages/routers.go +++ b/internal/apps/openflare/pages/routers.go @@ -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 } diff --git a/internal/apps/openflare/pages/routers_test.go b/internal/apps/openflare/pages/routers_test.go new file mode 100644 index 00000000..c73a55c5 --- /dev/null +++ b/internal/apps/openflare/pages/routers_test.go @@ -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()) +} diff --git a/internal/apps/openflare/proxy_route/logics.go b/internal/apps/openflare/proxy_route/logics.go index 1b781a84..1570be32 100644 --- a/internal/apps/openflare/proxy_route/logics.go +++ b/internal/apps/openflare/proxy_route/logics.go @@ -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 { diff --git a/internal/apps/openflare/proxy_route/logics_test.go b/internal/apps/openflare/proxy_route/logics_test.go index 1af61011..2267a089 100644 --- a/internal/apps/openflare/proxy_route/logics_test.go +++ b/internal/apps/openflare/proxy_route/logics_test.go @@ -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. diff --git a/internal/apps/upload/exports.go b/internal/apps/upload/exports.go index edbd7713..322acddf 100644 --- a/internal/apps/upload/exports.go +++ b/internal/apps/upload/exports.go @@ -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 diff --git a/internal/apps/upload/handler/file_management.go b/internal/apps/upload/handler/file_management.go index 075d72e8..6bb69d46 100644 --- a/internal/apps/upload/handler/file_management.go +++ b/internal/apps/upload/handler/file_management.go @@ -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 } diff --git a/internal/apps/upload/handler/routers.go b/internal/apps/upload/handler/routers.go index 025052c4..348f8f71 100644 --- a/internal/apps/upload/handler/routers.go +++ b/internal/apps/upload/handler/routers.go @@ -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 != "" { diff --git a/internal/apps/upload/handler/routers_test.go b/internal/apps/upload/handler/routers_test.go index e8ddce20..75459ff8 100644 --- a/internal/apps/upload/handler/routers_test.go +++ b/internal/apps/upload/handler/routers_test.go @@ -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 { diff --git a/internal/apps/upload/ingest/errors.go b/internal/apps/upload/ingest/errors.go index 6e501977..59ad4ed8 100644 --- a/internal/apps/upload/ingest/errors.go +++ b/internal/apps/upload/ingest/errors.go @@ -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) diff --git a/internal/apps/upload/ingest/helpers.go b/internal/apps/upload/ingest/helpers.go index 00580a90..b4b094ed 100644 --- a/internal/apps/upload/ingest/helpers.go +++ b/internal/apps/upload/ingest/helpers.go @@ -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 } diff --git a/internal/apps/upload/ingest/ingest_test.go b/internal/apps/upload/ingest/ingest_test.go index d70bedc7..4853d926 100644 --- a/internal/apps/upload/ingest/ingest_test.go +++ b/internal/apps/upload/ingest/ingest_test.go @@ -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 }, ) diff --git a/internal/apps/upload/ingest/remove.go b/internal/apps/upload/ingest/remove.go index 902ab5ef..05c6fb60 100644 --- a/internal/apps/upload/ingest/remove.go +++ b/internal/apps/upload/ingest/remove.go @@ -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) } diff --git a/internal/apps/upload/shared/constants.go b/internal/apps/upload/shared/constants.go index 14117af3..f3382112 100644 --- a/internal/apps/upload/shared/constants.go +++ b/internal/apps/upload/shared/constants.go @@ -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" ) diff --git a/internal/apps/upload/shared/errs.go b/internal/apps/upload/shared/errs.go index c37c864c..7f031341 100644 --- a/internal/apps/upload/shared/errs.go +++ b/internal/apps/upload/shared/errs.go @@ -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" diff --git a/internal/apps/upload/task/cleanup.go b/internal/apps/upload/task/cleanup.go index 0d70e35d..36b3f5b1 100644 --- a/internal/apps/upload/task/cleanup.go +++ b/internal/apps/upload/task/cleanup.go @@ -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 } } diff --git a/internal/apps/upload/task/tasks_test.go b/internal/apps/upload/task/tasks_test.go index 90572c8e..f75bec6b 100644 --- a/internal/apps/upload/task/tasks_test.go +++ b/internal/apps/upload/task/tasks_test.go @@ -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 diff --git a/internal/repository/upload.go b/internal/repository/upload.go index 34c8a96f..599f66c4 100644 --- a/internal/repository/upload.go +++ b/internal/repository/upload.go @@ -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. diff --git a/pkg/pagesarchive/entry.go b/pkg/pagesarchive/entry.go index 59b350aa..ce19734b 100644 --- a/pkg/pagesarchive/entry.go +++ b/pkg/pagesarchive/entry.go @@ -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 } diff --git a/pkg/pagesarchive/extract.go b/pkg/pagesarchive/extract.go index 25335264..c8eb1426 100644 --- a/pkg/pagesarchive/extract.go +++ b/pkg/pagesarchive/extract.go @@ -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 { diff --git a/pkg/pagesarchive/inspect.go b/pkg/pagesarchive/inspect.go index 6c3211d7..ae3a3125 100644 --- a/pkg/pagesarchive/inspect.go +++ b/pkg/pagesarchive/inspect.go @@ -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) } diff --git a/pkg/pagesarchive/list.go b/pkg/pagesarchive/list.go index 16f2fc8f..e5781f08 100644 --- a/pkg/pagesarchive/list.go +++ b/pkg/pagesarchive/list.go @@ -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 } diff --git a/pkg/pagesarchive/path.go b/pkg/pagesarchive/path.go index 4735ff86..94562c8b 100644 --- a/pkg/pagesarchive/path.go +++ b/pkg/pagesarchive/path.go @@ -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 } diff --git a/pkg/pagesarchive/security_test.go b/pkg/pagesarchive/security_test.go new file mode 100644 index 00000000..9194991e --- /dev/null +++ b/pkg/pagesarchive/security_test.go @@ -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() +} diff --git a/pkg/protocol/agent.go b/pkg/protocol/agent.go index 26572286..cd0ba6ee 100644 --- a/pkg/protocol/agent.go +++ b/pkg/protocol/agent.go @@ -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"` } diff --git a/pkg/protocol/agent_test.go b/pkg/protocol/agent_test.go index 4a0f71ab..0a1c99d4 100644 --- a/pkg/protocol/agent_test.go +++ b/pkg/protocol/agent_test.go @@ -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 {