merge: 合并 PR #22 Pages 部署源 V2 到 feat/pages-source-sync-v2

基于最新 main 合并 deqiying/feat/pages-source-sync-v2,
解决 docs/changelog/index.md 与限流相关条目的冲突。
This commit is contained in:
ryan
2026-07-19 19:36:48 +08:00
127 changed files with 24144 additions and 1933 deletions
+35 -5
View File
@@ -66,7 +66,7 @@ func DispatchTask(c *gin.Context) {
return
}
meta := task.GetTaskMeta(req.TaskType)
meta := getAdminTaskMeta(req.TaskType)
if meta == nil {
response.AbortBadRequest(c, InvalidTaskType)
return
@@ -202,7 +202,7 @@ func RetryTask(c *gin.Context) {
// ListSchedules 获取定时任务列表
// @Summary 获取定时任务列表
// @Description 返回系统所有的定时任务配置列表,包括名称、关联的异步任务类型、Cron 表达式和启用状态,需要管理员权限
// @Description 返回管理员可管理的定时任务配置列表,包括名称、关联的异步任务类型、Cron 表达式和启用状态;系统内部排程不会暴露,需要管理员权限
// @Tags admin
// @Produce json
// @Security SessionCookie
@@ -216,7 +216,15 @@ func ListSchedules(c *gin.Context) {
response.AbortInternal(c, err.Error())
return
}
c.JSON(http.StatusOK, response.OK(schedules))
visible := make([]model.Schedule, 0, len(schedules))
for _, schedule := range schedules {
meta := task.GetTaskMeta(schedule.TaskType)
if meta != nil && meta.InternalOnly {
continue
}
visible = append(visible, schedule)
}
c.JSON(http.StatusOK, response.OK(visible))
}
// CreateScheduleRequest 创建定时任务请求
@@ -256,7 +264,7 @@ func CreateSchedule(c *gin.Context) {
}
// 校验关联的异步任务类型
meta := task.GetTaskMeta(req.TaskType)
meta := getAdminTaskMeta(req.TaskType)
if meta == nil {
response.AbortBadRequest(c, InvalidTaskType)
return
@@ -338,6 +346,10 @@ func UpdateSchedule(c *gin.Context) {
response.AbortNotFound(c, ScheduleNotFound)
return
}
if existingMeta := task.GetTaskMeta(schedule.TaskType); existingMeta != nil && existingMeta.InternalOnly {
response.AbortBadRequest(c, InvalidTaskType)
return
}
// 校验 Cron 表达式
if _, err := cron.ParseStandard(req.Cron); err != nil {
@@ -346,7 +358,7 @@ func UpdateSchedule(c *gin.Context) {
}
// 校验关联的异步任务类型
meta := task.GetTaskMeta(req.TaskType)
meta := getAdminTaskMeta(req.TaskType)
if meta == nil {
response.AbortBadRequest(c, InvalidTaskType)
return
@@ -382,6 +394,14 @@ func UpdateSchedule(c *gin.Context) {
c.JSON(http.StatusOK, response.OK(schedule))
}
func getAdminTaskMeta(taskType string) *task.TaskMeta {
meta := task.GetTaskMeta(taskType)
if meta == nil || meta.InternalOnly {
return nil
}
return meta
}
// DeleteSchedule 删除定时任务
// @Summary 删除定时任务
// @Description 删除指定的定时任务配置,并触发调度器热加载,需要管理员权限
@@ -393,6 +413,7 @@ func UpdateSchedule(c *gin.Context) {
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 403 {object} response.Any "无管理员权限"
// @Failure 404 {object} response.Any "定时任务不存在"
// @Failure 500 {object} response.Any "删除定时任务失败"
// @Router /api/v1/admin/tasks/schedules/{id} [delete]
func DeleteSchedule(c *gin.Context) {
@@ -401,6 +422,15 @@ func DeleteSchedule(c *gin.Context) {
response.AbortBadRequest(c, "无效的定时任务ID")
return
}
schedule, err := model.GetScheduleByID(c.Request.Context(), id)
if err != nil {
response.AbortNotFound(c, ScheduleNotFound)
return
}
if meta := task.GetTaskMeta(schedule.TaskType); meta != nil && meta.InternalOnly {
response.AbortBadRequest(c, InvalidTaskType)
return
}
if err := model.DeleteSchedule(c.Request.Context(), id); err != nil {
response.AbortInternal(c, fmt.Sprintf("%s: %v", ScheduleDeleteFailed, err))
+236
View File
@@ -29,6 +29,20 @@ import (
"github.com/Rain-kl/Wavelet/internal/common/response"
)
const (
testInternalOnlyTaskType = "test_internal_only_admin"
testInternalOnlyAsynqTask = "test:internal_only_admin"
)
func registerInternalOnlyTaskMeta() {
task.RegisterTaskMeta(task.TaskMeta{
Type: testInternalOnlyTaskType,
AsynqTask: testInternalOnlyAsynqTask,
Name: "内部测试任务",
InternalOnly: true,
})
}
func setupTaskTestEnvironment(t *testing.T) func() {
_, mr, cleanup := testhelper.SetupTestEnvironment(t)
bootstrap.RegisterTasks()
@@ -61,12 +75,17 @@ func setupTestRouter(authUser *model.User) *gin.Engine {
adminGroup.GET("/tasks/executions", ListTaskExecutions)
adminGroup.GET("/tasks/executions/:id", GetTaskExecution)
adminGroup.POST("/tasks/executions/:id/retry", RetryTask)
adminGroup.GET("/tasks/schedules", ListSchedules)
adminGroup.POST("/tasks/schedules", CreateSchedule)
adminGroup.PUT("/tasks/schedules/:id", UpdateSchedule)
adminGroup.DELETE("/tasks/schedules/:id", DeleteSchedule)
return r
}
func TestListTaskTypes(t *testing.T) {
cleanup := setupTaskTestEnvironment(t)
defer cleanup()
registerInternalOnlyTaskMeta()
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true}
router := setupTestRouter(adminUser)
@@ -93,6 +112,9 @@ func TestListTaskTypes(t *testing.T) {
foundCleanup := false
foundWarmImageCache := false
for _, m := range taskMetas {
if m.Type == testInternalOnlyTaskType {
t.Errorf("internal-only task type %s must not be listed", testInternalOnlyTaskType)
}
if m.Type == uploadtask.TaskTypeSystemCleanup {
foundCleanup = true
}
@@ -108,6 +130,220 @@ func TestListTaskTypes(t *testing.T) {
}
}
func TestInternalOnlyTaskAdminBoundaries(t *testing.T) {
cleanup := setupTaskTestEnvironment(t)
defer cleanup()
registerInternalOnlyTaskMeta()
adminUser := &model.User{ID: 1001, Username: "admin", IsAdmin: true}
router := setupTestRouter(adminUser)
ctx := context.Background()
t.Run("list hides internal-only schedule", func(t *testing.T) {
internalSchedule := &model.Schedule{
Name: "隐藏的系统内部排程",
TaskType: testInternalOnlyTaskType,
Cron: "*/5 * * * *",
Payload: "{}",
IsActive: true,
}
publicSchedule := &model.Schedule{
Name: "可见的公开排程",
TaskType: uploadtask.TaskTypeSystemCleanup,
Cron: "0 * * * *",
Payload: "{}",
IsActive: true,
}
require.NoError(t, model.CreateSchedule(ctx, internalSchedule))
require.NoError(t, model.CreateSchedule(ctx, publicSchedule))
req := httptest.NewRequest(http.MethodGet, "/api/v1/admin/tasks/schedules", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
var resp response.Any
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
data, err := json.Marshal(resp.Data)
require.NoError(t, err)
var schedules []model.Schedule
require.NoError(t, json.Unmarshal(data, &schedules))
assert.NotContains(t, scheduleIDs(schedules), internalSchedule.ID)
assert.Contains(t, scheduleIDs(schedules), publicSchedule.ID)
})
t.Run("dispatch rejects internal-only task", func(t *testing.T) {
body, err := json.Marshal(DispatchTaskRequest{TaskType: testInternalOnlyTaskType})
require.NoError(t, err)
req := httptest.NewRequest(http.MethodPost, "/api/v1/admin/tasks/dispatch", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusBadRequest, w.Code)
var resp response.Any
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
assert.Equal(t, InvalidTaskType, resp.ErrorMsg)
})
t.Run("create schedule rejects internal-only task", func(t *testing.T) {
isActive := true
body, err := json.Marshal(CreateScheduleRequest{
Name: "内部任务排程",
TaskType: testInternalOnlyTaskType,
Cron: "0 * * * *",
IsActive: &isActive,
})
require.NoError(t, err)
req := httptest.NewRequest(http.MethodPost, "/api/v1/admin/tasks/schedules", bytes.NewReader(body))
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusBadRequest, w.Code)
var resp response.Any
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
assert.Equal(t, InvalidTaskType, resp.ErrorMsg)
})
t.Run("update cannot convert existing internal schedule to public task", func(t *testing.T) {
schedule := &model.Schedule{
Name: "系统内部排程",
TaskType: testInternalOnlyTaskType,
Cron: "0 * * * *",
IsActive: true,
}
require.NoError(t, model.CreateSchedule(ctx, schedule))
isActive := false
body, err := json.Marshal(UpdateScheduleRequest{
Name: "尝试修改内部排程",
TaskType: uploadtask.TaskTypeSystemCleanup,
Cron: "5 * * * *",
IsActive: &isActive,
})
require.NoError(t, err)
req := httptest.NewRequest(
http.MethodPut,
fmt.Sprintf("/api/v1/admin/tasks/schedules/%d", schedule.ID),
bytes.NewReader(body),
)
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusBadRequest, w.Code)
unchanged, err := model.GetScheduleByID(ctx, schedule.ID)
require.NoError(t, err)
assert.Equal(t, "系统内部排程", unchanged.Name)
assert.Equal(t, testInternalOnlyTaskType, unchanged.TaskType)
assert.True(t, unchanged.IsActive)
})
t.Run("update public schedule rejects internal-only target task", func(t *testing.T) {
schedule := &model.Schedule{
Name: "公开排程",
TaskType: uploadtask.TaskTypeSystemCleanup,
Cron: "0 * * * *",
IsActive: true,
}
require.NoError(t, model.CreateSchedule(ctx, schedule))
isActive := true
body, err := json.Marshal(UpdateScheduleRequest{
Name: "尝试切入内部任务",
TaskType: testInternalOnlyTaskType,
Cron: "10 * * * *",
IsActive: &isActive,
})
require.NoError(t, err)
req := httptest.NewRequest(
http.MethodPut,
fmt.Sprintf("/api/v1/admin/tasks/schedules/%d", schedule.ID),
bytes.NewReader(body),
)
req.Header.Set("Content-Type", "application/json")
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusBadRequest, w.Code)
unchanged, err := model.GetScheduleByID(ctx, schedule.ID)
require.NoError(t, err)
assert.Equal(t, "公开排程", unchanged.Name)
assert.Equal(t, uploadtask.TaskTypeSystemCleanup, unchanged.TaskType)
})
t.Run("delete rejects internal-only schedule", func(t *testing.T) {
schedule := &model.Schedule{
Name: "不可删除的系统内部排程",
TaskType: testInternalOnlyTaskType,
Cron: "*/5 * * * *",
IsActive: true,
}
require.NoError(t, model.CreateSchedule(ctx, schedule))
req := httptest.NewRequest(
http.MethodDelete,
fmt.Sprintf("/api/v1/admin/tasks/schedules/%d", schedule.ID),
nil,
)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusBadRequest, w.Code)
var resp response.Any
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
assert.Equal(t, InvalidTaskType, resp.ErrorMsg)
preserved, err := model.GetScheduleByID(ctx, schedule.ID)
require.NoError(t, err)
assert.Equal(t, testInternalOnlyTaskType, preserved.TaskType)
})
t.Run("delete missing schedule returns not found", func(t *testing.T) {
req := httptest.NewRequest(http.MethodDelete, "/api/v1/admin/tasks/schedules/999999", nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusNotFound, w.Code)
var resp response.Any
require.NoError(t, json.Unmarshal(w.Body.Bytes(), &resp))
assert.Equal(t, ScheduleNotFound, resp.ErrorMsg)
})
t.Run("delete public schedule remains allowed", func(t *testing.T) {
schedule := &model.Schedule{
Name: "可删除的公开排程",
TaskType: uploadtask.TaskTypeSystemCleanup,
Cron: "0 * * * *",
IsActive: false,
}
require.NoError(t, model.CreateSchedule(ctx, schedule))
req := httptest.NewRequest(
http.MethodDelete,
fmt.Sprintf("/api/v1/admin/tasks/schedules/%d", schedule.ID),
nil,
)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
assert.Equal(t, http.StatusOK, w.Code)
_, err := model.GetScheduleByID(ctx, schedule.ID)
assert.Error(t, err)
})
}
func scheduleIDs(schedules []model.Schedule) []uint64 {
ids := make([]uint64, 0, len(schedules))
for _, schedule := range schedules {
ids = append(ids, schedule.ID)
}
return ids
}
func TestDispatchTask(t *testing.T) {
cleanup := setupTaskTestEnvironment(t)
defer cleanup()
+113 -18
View File
@@ -1,9 +1,13 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package httpclient provides an authenticated HTTP client for the agent.
package httpclient
import (
"context"
"encoding/json"
"errors"
"fmt"
"io"
"net/http"
@@ -13,6 +17,8 @@ import (
edgehttp "github.com/Rain-kl/Wavelet/internal/apps/edge/httpclient"
)
const pagesControlResponseMaxBytes = int64(64 * 1024)
// Client is a HTTP client used by the agent to communicate with the control plane server.
type Client struct {
base *edgehttp.Client
@@ -98,23 +104,43 @@ func (c *Client) GetPagesDeploymentHash(ctx context.Context, deploymentID uint)
return resp.Data.Hash, nil
}
// DownloadPagesDeploymentPackage downloads the deployment package for the given Pages deployment ID.
func (c *Client) DownloadPagesDeploymentPackage(ctx context.Context, deploymentID uint) ([]byte, error) {
res, err := c.base.DoRaw(ctx, http.MethodGet, fmt.Sprintf("/api/v1/agent/pages/deployments/%d/package", deploymentID), nil)
if err != nil {
return nil, err
}
defer func() { _ = res.Body.Close() }()
if res.StatusCode != http.StatusOK {
return nil, edgehttp.ReadHTTPError(res)
}
return io.ReadAll(res.Body)
// DownloadPagesDeploymentPackage streams the deployment package into dst while
// enforcing maxBytes against both advertised and actual response sizes.
func (c *Client) DownloadPagesDeploymentPackage(
ctx context.Context,
deploymentID uint,
dst io.Writer,
maxBytes int64,
) (int64, error) {
return c.downloadPagesPackage(
ctx,
fmt.Sprintf("/api/v1/agent/pages/deployments/%d/package", deploymentID),
dst,
maxBytes,
)
}
// GetPagesProjectLatestHash returns the active deployment package hash for a Pages project.
func (c *Client) GetPagesProjectLatestHash(ctx context.Context, projectID uint) (*protocol.PagesProjectLatestHashResponse, error) {
res, err := c.base.DoRaw(
ctx,
http.MethodGet,
fmt.Sprintf("/api/v1/agent/pages/projects/%d/latest/hash", projectID),
nil,
)
if err != nil {
return nil, err
}
defer func() { _ = res.Body.Close() }()
body, err := readPagesControlResponse(res, pagesControlResponseMaxBytes)
if err != nil {
return nil, err
}
if res.StatusCode != http.StatusOK {
return nil, edgehttp.ReadBodyError(body, res.Status)
}
resp := protocol.APIResponse[protocol.PagesProjectLatestHashResponse]{}
if err := c.base.GetJSON(ctx, fmt.Sprintf("/api/v1/agent/pages/projects/%d/latest/hash", projectID), &resp); err != nil {
if err := json.Unmarshal(body, &resp); err != nil {
return nil, err
}
if err := edgehttp.APIError(resp.ErrorMsg); err != nil {
@@ -123,17 +149,86 @@ func (c *Client) GetPagesProjectLatestHash(ctx context.Context, projectID uint)
return &resp.Data, nil
}
// DownloadPagesProjectLatestPackage downloads the active deployment package for a Pages project.
func (c *Client) DownloadPagesProjectLatestPackage(ctx context.Context, projectID uint) ([]byte, error) {
res, err := c.base.DoRaw(ctx, http.MethodGet, fmt.Sprintf("/api/v1/agent/pages/projects/%d/latest/package", projectID), nil)
// DownloadPagesProjectLatestPackage streams the active deployment package into
// dst while enforcing maxBytes against both advertised and actual sizes.
func (c *Client) DownloadPagesProjectLatestPackage(
ctx context.Context,
projectID uint,
dst io.Writer,
maxBytes int64,
) (int64, error) {
return c.downloadPagesPackage(
ctx,
fmt.Sprintf("/api/v1/agent/pages/projects/%d/latest/package", projectID),
dst,
maxBytes,
)
}
func (c *Client) downloadPagesPackage(
ctx context.Context,
path string,
dst io.Writer,
maxBytes int64,
) (int64, error) {
if dst == nil {
return 0, errors.New("pages package destination is required")
}
if maxBytes <= 0 {
return 0, errors.New("pages package byte limit must be positive")
}
res, err := c.base.DoRaw(ctx, http.MethodGet, path, nil)
if err != nil {
return nil, err
return 0, err
}
defer func() { _ = res.Body.Close() }()
if res.StatusCode != http.StatusOK {
return nil, edgehttp.ReadHTTPError(res)
body, readErr := readPagesControlResponse(res, pagesControlResponseMaxBytes)
if readErr != nil {
return 0, readErr
}
return 0, edgehttp.ReadBodyError(body, res.Status)
}
return io.ReadAll(res.Body)
return copyPagesPackageResponse(dst, res, maxBytes)
}
func readPagesControlResponse(res *http.Response, maxBytes int64) ([]byte, error) {
if res.ContentLength > maxBytes {
return nil, fmt.Errorf(
"pages control response Content-Length %d exceeds limit %d",
res.ContentLength,
maxBytes,
)
}
limited := &io.LimitedReader{R: res.Body, N: maxBytes + 1}
body, err := io.ReadAll(limited)
if err != nil {
return nil, fmt.Errorf("read pages control response: %w", err)
}
if int64(len(body)) > maxBytes {
return nil, fmt.Errorf("pages control response body exceeds limit %d", maxBytes)
}
return body, nil
}
func copyPagesPackageResponse(dst io.Writer, res *http.Response, maxBytes int64) (int64, error) {
if res.ContentLength > maxBytes {
return 0, fmt.Errorf(
"pages package Content-Length %d exceeds limit %d",
res.ContentLength,
maxBytes,
)
}
limited := &io.LimitedReader{R: res.Body, N: maxBytes + 1}
written, err := io.Copy(dst, limited)
if err != nil {
return written, fmt.Errorf("stream pages package: %w", err)
}
if written > maxBytes {
return written, fmt.Errorf("pages package body exceeds limit %d", maxBytes)
}
return written, nil
}
// SetToken updates the authentication token used for API requests.
@@ -0,0 +1,109 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package httpclient
import (
"bytes"
"context"
"io"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
)
func TestDownloadPagesProjectLatestPackageRejectsChunkedBodyOverLimit(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusOK)
if flusher, ok := w.(http.Flusher); ok {
flusher.Flush()
}
_, _ = io.WriteString(w, "123456")
}))
defer server.Close()
client := New(server.URL, "test-token", time.Second)
var dst bytes.Buffer
written, err := client.DownloadPagesProjectLatestPackage(
context.Background(),
7,
&dst,
4,
)
if err == nil || !strings.Contains(err.Error(), "body exceeds limit") {
t.Fatalf("DownloadPagesProjectLatestPackage(chunked, limit=4) error = %v, want body limit error", err)
}
if written != 5 {
t.Errorf("DownloadPagesProjectLatestPackage(chunked, limit=4) written = %d, want 5", written)
}
}
func TestCopyPagesPackageResponseRejectsAdvertisedContentLengthBeforeWrite(t *testing.T) {
response := &http.Response{
Body: io.NopCloser(strings.NewReader("123456")),
ContentLength: 6,
}
var dst bytes.Buffer
written, err := copyPagesPackageResponse(&dst, response, 4)
if err == nil || !strings.Contains(err.Error(), "Content-Length") {
t.Fatalf("copyPagesPackageResponse(Content-Length=6, limit=4) error = %v, want Content-Length limit error", err)
}
if written != 0 || dst.Len() != 0 {
t.Errorf("copyPagesPackageResponse(Content-Length=6, limit=4) wrote (%d, %d buffered), want no writes", written, dst.Len())
}
}
func TestCopyPagesPackageResponseRejectsForgedSmallContentLength(t *testing.T) {
response := &http.Response{
Body: io.NopCloser(strings.NewReader("123456")),
ContentLength: 2,
}
var dst bytes.Buffer
written, err := copyPagesPackageResponse(&dst, response, 4)
if err == nil || !strings.Contains(err.Error(), "body exceeds limit") {
t.Fatalf("copyPagesPackageResponse(forged Content-Length=2, limit=4) error = %v, want body limit error", err)
}
if written != 5 {
t.Errorf("copyPagesPackageResponse(forged Content-Length=2, limit=4) written = %d, want 5", written)
}
}
func TestDownloadPagesProjectLatestPackageBoundsChunkedErrorResponse(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusBadRequest)
if flusher, ok := w.(http.Flusher); ok {
flusher.Flush()
}
_, _ = io.WriteString(w, strings.Repeat("x", int(pagesControlResponseMaxBytes+1)))
}))
defer server.Close()
client := New(server.URL, "test-token", time.Second)
var dst bytes.Buffer
_, err := client.DownloadPagesProjectLatestPackage(context.Background(), 7, &dst, 1024)
if err == nil || !strings.Contains(err.Error(), "control response body exceeds limit") {
t.Fatalf("DownloadPagesProjectLatestPackage(large chunked 400) error = %v, want bounded response error", err)
}
if dst.Len() != 0 {
t.Errorf("DownloadPagesProjectLatestPackage(large chunked 400) wrote %d package bytes, want 0", dst.Len())
}
}
func TestGetPagesProjectLatestHashBoundsChunkedMetadataResponse(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.WriteHeader(http.StatusOK)
if flusher, ok := w.(http.Flusher); ok {
flusher.Flush()
}
_, _ = io.WriteString(w, strings.Repeat("x", int(pagesControlResponseMaxBytes+1)))
}))
defer server.Close()
client := New(server.URL, "test-token", time.Second)
_, err := client.GetPagesProjectLatestHash(context.Background(), 7)
if err == nil || !strings.Contains(err.Error(), "control response body exceeds limit") {
t.Fatalf("GetPagesProjectLatestHash(large chunked metadata) error = %v, want bounded response error", err)
}
}
+565 -94
View File
@@ -1,3 +1,6 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package sync applies control-plane configuration to the local agent runtime.
package sync
@@ -20,9 +23,13 @@ import (
)
const (
pagesDirPerm = 0o755
pagesFilePerm = 0o644
pagesManifestFilePerm = 0o644
pagesDirPerm = 0o755
pagesFilePerm = 0o644
pagesManifestFilePerm = 0o644
agentPagesMaxPackageBytes = int64(2 * 1024 * 1024 * 1024)
agentPagesMaxFiles = 1000
agentPagesMaxFileBytes = int64(8 * 1024 * 1024 * 1024)
agentPagesMaxTotalBytes = int64(8 * 1024 * 1024 * 1024)
// pagesLatestPullAttempts covers a race where the active deployment changes
// between the hash probe and the package download.
pagesLatestPullAttempts = 2
@@ -45,6 +52,11 @@ type pagesProjectRef struct {
Checksum string
}
type pagesPackageLimits struct {
PackageBytes int64
Extraction pagesarchive.Limits
}
type pagesDeploymentMarker struct {
ProjectID uint `json:"project_id"`
DeploymentID uint `json:"deployment_id,omitempty"`
@@ -191,10 +203,11 @@ func (s *Service) ensurePagesProject(ctx context.Context, snapshot *state.Snapsh
if err != nil {
return fmt.Errorf("fetch Pages project %d latest hash: %w", projectID, err)
}
hash := strings.TrimSpace(latest.Hash)
if hash == "" {
return fmt.Errorf("pages project %d latest hash is empty", projectID)
limits, err := validatePagesPackageMetadata(projectID, latest)
if err != nil {
return err
}
hash := strings.TrimSpace(latest.Hash)
effective := pagesProjectRef{
ProjectID: projectID,
DeploymentID: latest.DeploymentID,
@@ -211,44 +224,65 @@ func (s *Service) ensurePagesProject(ctx context.Context, snapshot *state.Snapsh
return nil
}
packageBytes, err := s.client.DownloadPagesProjectLatestPackage(ctx, projectID)
packagePath, got, err := s.downloadPagesProjectPackage(ctx, projectID, latest, limits.PackageBytes)
if err != nil {
return fmt.Errorf("download Pages project %d latest package: %w", projectID, err)
}
got := checksumBytes(packageBytes)
// Re-probe latest after download to detect activation races.
// Accept the package only when its content hash still matches latest.
// A deployment-id-only change is still a latest-pointer race even when
// deduplication makes both deployments share the same package hash.
verify, err := s.client.GetPagesProjectLatestHash(ctx, projectID)
if err != nil {
_ = os.Remove(packagePath)
return fmt.Errorf("re-fetch Pages project %d latest hash: %w", projectID, err)
}
verifyHash := strings.TrimSpace(verify.Hash)
if verifyHash == "" {
return fmt.Errorf("pages project %d latest hash is empty", projectID)
if _, err := validatePagesPackageMetadata(projectID, verify); err != nil {
_ = os.Remove(packagePath)
return err
}
if got != verifyHash {
if !samePagesPackageMetadata(latest, verify) {
_ = os.Remove(packagePath)
lastErr = fmt.Errorf(
"pages project %d package/hash race: downloaded %s, latest now %s (attempt %d/%d)",
projectID, got, verifyHash, attempt+1, pagesLatestPullAttempts,
"pages project %d latest metadata changed during download: deployment %d/%s -> %d/%s (attempt %d/%d)",
projectID,
latest.DeploymentID,
strings.TrimSpace(latest.Hash),
verify.DeploymentID,
strings.TrimSpace(verify.Hash),
attempt+1,
pagesLatestPullAttempts,
)
slog.Warn("pages latest package race, retrying",
slog.Warn("pages latest metadata race, retrying",
"project_id", projectID,
"before_deployment_id", latest.DeploymentID,
"before_hash", strings.TrimSpace(latest.Hash),
"after_deployment_id", verify.DeploymentID,
"after_hash", strings.TrimSpace(verify.Hash),
"attempt", attempt+1,
)
continue
}
if got != hash {
_ = os.Remove(packagePath)
lastErr = fmt.Errorf(
"pages project %d package hash mismatch: downloaded %s, expected %s (attempt %d/%d)",
projectID, got, hash, attempt+1, pagesLatestPullAttempts,
)
slog.Warn("pages latest package hash mismatch, retrying",
"project_id", projectID,
"downloaded_hash", got,
"latest_hash", verifyHash,
"expected_hash", hash,
"attempt", attempt+1,
)
continue
}
effective = pagesProjectRef{
ProjectID: projectID,
DeploymentID: verify.DeploymentID,
Checksum: got,
}
releaseDir = pagesProjectReleaseDir(s.pagesDir, projectID, got)
if err := extractPagesPackage(packageBytes, releaseDir, effective); err != nil {
return err
extractErr := extractPagesPackageFile(packagePath, releaseDir, effective, limits.Extraction, latest)
_ = os.Remove(packagePath)
if extractErr != nil {
return extractErr
}
if err := switchPagesProjectCurrentDir(s.pagesDir, projectID, releaseDir); err != nil {
return err
@@ -264,6 +298,145 @@ func (s *Service) ensurePagesProject(ctx context.Context, snapshot *state.Snapsh
return fmt.Errorf("pages project %d latest pull failed", projectID)
}
func validatePagesPackageMetadata(
projectID uint,
metadata *protocol.PagesProjectLatestHashResponse,
) (pagesPackageLimits, error) {
if metadata == nil {
return pagesPackageLimits{}, fmt.Errorf("pages project %d latest metadata is missing", projectID)
}
if metadata.ProjectID != projectID {
return pagesPackageLimits{}, fmt.Errorf(
"pages project %d latest metadata has project id %d",
projectID,
metadata.ProjectID,
)
}
if metadata.DeploymentID == 0 {
return pagesPackageLimits{}, fmt.Errorf("pages project %d latest deployment id is missing", projectID)
}
if strings.TrimSpace(metadata.Hash) == "" {
return pagesPackageLimits{}, fmt.Errorf("pages project %d latest hash is empty", projectID)
}
if metadata.PackageSize <= 0 {
return pagesPackageLimits{}, fmt.Errorf("pages project %d package size must be positive", projectID)
}
if metadata.PackageSize > agentPagesMaxPackageBytes {
return pagesPackageLimits{}, fmt.Errorf(
"pages project %d package size %d exceeds agent limit %d",
projectID,
metadata.PackageSize,
agentPagesMaxPackageBytes,
)
}
if metadata.FileCount <= 0 {
return pagesPackageLimits{}, fmt.Errorf("pages project %d file count must be positive", projectID)
}
if metadata.FileCount > agentPagesMaxFiles {
return pagesPackageLimits{}, fmt.Errorf(
"pages project %d file count %d exceeds agent limit %d",
projectID,
metadata.FileCount,
agentPagesMaxFiles,
)
}
if metadata.TotalSize < 0 {
return pagesPackageLimits{}, fmt.Errorf("pages project %d total size cannot be negative", projectID)
}
if metadata.TotalSize > agentPagesMaxTotalBytes {
return pagesPackageLimits{}, fmt.Errorf(
"pages project %d total size %d exceeds agent limit %d",
projectID,
metadata.TotalSize,
agentPagesMaxTotalBytes,
)
}
// pagesarchive treats zero limits as defaults. A one-byte extraction guard
// plus the exact post-extraction manifest check below preserves the valid
// case of one or more zero-byte files while still enforcing total_size=0.
extractedBytes := metadata.TotalSize
if extractedBytes == 0 {
extractedBytes = 1
}
maxFileBytes := extractedBytes
if maxFileBytes > agentPagesMaxFileBytes {
maxFileBytes = agentPagesMaxFileBytes
}
return pagesPackageLimits{
PackageBytes: metadata.PackageSize,
Extraction: pagesarchive.Limits{
MaxFiles: metadata.FileCount,
MaxFileBytes: maxFileBytes,
MaxTotalBytes: extractedBytes,
},
}, nil
}
func samePagesPackageMetadata(
before *protocol.PagesProjectLatestHashResponse,
after *protocol.PagesProjectLatestHashResponse,
) bool {
if before == nil || after == nil {
return false
}
return before.ProjectID == after.ProjectID &&
before.DeploymentID == after.DeploymentID &&
strings.TrimSpace(before.Hash) == strings.TrimSpace(after.Hash) &&
before.PackageSize == after.PackageSize &&
before.FileCount == after.FileCount &&
before.TotalSize == after.TotalSize
}
func (s *Service) downloadPagesProjectPackage(
ctx context.Context,
projectID uint,
metadata *protocol.PagesProjectLatestHashResponse,
maxBytes int64,
) (packagePath string, hash string, err error) {
releasesRoot := filepath.Join(s.pagesDir, "projects", fmt.Sprintf("%d", projectID), "releases")
if err := os.MkdirAll(releasesRoot, pagesDirPerm); err != nil {
return "", "", err
}
packageFile, err := os.CreateTemp(releasesRoot, ".package-*.tmp")
if err != nil {
return "", "", err
}
packagePath = packageFile.Name()
keep := false
defer func() {
if closeErr := packageFile.Close(); err == nil && closeErr != nil {
err = closeErr
}
if !keep || err != nil {
_ = os.Remove(packagePath)
packagePath = ""
}
}()
hasher := sha256.New()
written, err := s.client.DownloadPagesProjectLatestPackage(
ctx,
projectID,
io.MultiWriter(packageFile, hasher),
maxBytes,
)
if err != nil {
return "", "", err
}
if written != metadata.PackageSize {
return "", "", fmt.Errorf(
"pages project %d package size %d does not match metadata %d",
projectID,
written,
metadata.PackageSize,
)
}
keep = true
return packagePath, hex.EncodeToString(hasher.Sum(nil)), nil
}
// cleanupPagesProjectStaleReleases keeps only keepHash under projects/{id}/releases.
// Must be called only after the keepHash release is ready and current points at it.
func cleanupPagesProjectStaleReleases(baseDir string, projectID uint, keepHash string) error {
@@ -372,40 +545,258 @@ type pagesDeploymentSource struct {
Checksum string `json:"checksum"`
}
func extractPagesPackage(packageBytes []byte, releaseDir string, project pagesProjectRef) error {
tmpDir := releaseDir + ".tmp"
_ = os.RemoveAll(tmpDir)
if err := os.MkdirAll(tmpDir, pagesDirPerm); err != nil {
func extractPagesPackageFile(
packagePath string,
releaseDir string,
project pagesProjectRef,
limits pagesarchive.Limits,
expected *protocol.PagesProjectLatestHashResponse,
) error {
if err := os.MkdirAll(filepath.Dir(releaseDir), pagesDirPerm); err != nil {
return err
}
format, err := pagesarchive.DetectFormat("", packageBytes)
stagingDir, err := os.MkdirTemp(
filepath.Dir(releaseDir),
"."+filepath.Base(releaseDir)+"-*.tmp",
)
if err != nil {
_ = os.RemoveAll(tmpDir)
return fmt.Errorf("detect Pages package format: %w", err)
return err
}
// Control plane already inspected and accepted this package.
if err := pagesarchive.ExtractBytes(packageBytes, format, tmpDir, pagesarchive.ExtractOptions{
cleanupStaging := true
defer func() {
if cleanupStaging {
removePagesStagingUnlessCurrent(stagingDir, pagesCurrentDirFromRelease(releaseDir))
}
}()
if err := pagesarchive.ExtractFile(packagePath, "", stagingDir, pagesarchive.ExtractOptions{
StripCommonRoot: true,
EnforceLimits: false,
EnforceLimits: true,
Limits: limits,
}); err != nil {
_ = os.RemoveAll(tmpDir)
return fmt.Errorf("extract Pages package: %w", err)
}
if err := writePagesMarker(tmpDir, project); err != nil {
_ = os.RemoveAll(tmpDir)
if err := validateExtractedPagesMetadata(stagingDir, expected); err != nil {
return err
}
_ = os.RemoveAll(releaseDir)
return os.Rename(tmpDir, releaseDir)
if err := writePagesMarker(stagingDir, project); err != nil {
return err
}
if err := promotePagesRelease(stagingDir, releaseDir, project); err != nil {
return err
}
cleanupStaging = false
return nil
}
func switchPagesProjectCurrentDir(baseDir string, projectID uint, releaseDir string) error {
currentDir := pagesProjectCurrentDir(baseDir, projectID)
previousDir := currentDir + ".previous"
_ = os.RemoveAll(previousDir)
func validateExtractedPagesMetadata(
dir string,
expected *protocol.PagesProjectLatestHashResponse,
) error {
if expected == nil {
return nil
}
fileCount := 0
totalSize := int64(0)
err := filepath.WalkDir(dir, func(path string, entry os.DirEntry, walkErr error) error {
if walkErr != nil {
return walkErr
}
if entry.IsDir() {
return nil
}
info, err := entry.Info()
if err != nil {
return err
}
if !info.Mode().IsRegular() {
return fmt.Errorf("pages extracted entry is not a regular file: %s", path)
}
fileCount++
if fileCount > agentPagesMaxFiles {
return fmt.Errorf("pages extracted file count exceeds agent limit %d", agentPagesMaxFiles)
}
if info.Size() < 0 || info.Size() > agentPagesMaxTotalBytes-totalSize {
return fmt.Errorf("pages extracted size exceeds agent limit %d", agentPagesMaxTotalBytes)
}
totalSize += info.Size()
return nil
})
if err != nil {
return fmt.Errorf("validate extracted Pages package: %w", err)
}
if fileCount != expected.FileCount || totalSize != expected.TotalSize {
return fmt.Errorf(
"pages extracted metadata mismatch: got %d files/%d bytes, expected %d files/%d bytes",
fileCount,
totalSize,
expected.FileCount,
expected.TotalSize,
)
}
return nil
}
func promotePagesRelease(stagingDir string, releaseDir string, project pagesProjectRef) error {
return promotePagesReleaseWithCopy(stagingDir, releaseDir, project, copyPagesDir)
}
func promotePagesReleaseWithCopy(
stagingDir string,
releaseDir string,
project pagesProjectRef,
copyDir func(string, string) error,
) error {
currentDir := pagesCurrentDirFromRelease(releaseDir)
defer removePagesStagingUnlessCurrent(stagingDir, currentDir)
currentUsesRelease, err := pagesCurrentTargetsRelease(currentDir, releaseDir)
if err != nil {
return err
}
if !currentUsesRelease {
if err := os.RemoveAll(releaseDir); err != nil {
return err
}
return os.Rename(stagingDir, releaseDir)
}
// A same-hash repair cannot remove releaseDir while current still resolves
// through it. Keep traffic on the fully validated staging tree, rebuild the
// canonical release, then atomically point current back to the canonical path.
if err := switchPagesCurrentDir(currentDir, stagingDir, os.Rename); err != nil {
return fmt.Errorf("switch Pages current to repair staging: %w", err)
}
backupDir := stagingDir + ".previous"
if err := os.Rename(releaseDir, backupDir); err != nil {
restoreErr := switchPagesCurrentDir(currentDir, releaseDir, os.Rename)
return errors.Join(
fmt.Errorf("move previous Pages release aside: %w", err),
restoreErr,
)
}
rollback := func(cause error) error {
var rollbackErrors []error
rollbackErrors = append(rollbackErrors, cause)
if err := os.RemoveAll(releaseDir); err != nil {
rollbackErrors = append(rollbackErrors, fmt.Errorf("remove failed Pages release repair: %w", err))
}
if err := os.Rename(backupDir, releaseDir); err != nil {
rollbackErrors = append(rollbackErrors, fmt.Errorf("restore previous Pages release: %w", err))
return errors.Join(rollbackErrors...)
}
if err := switchPagesCurrentDir(currentDir, releaseDir, os.Rename); err != nil {
rollbackErrors = append(rollbackErrors, fmt.Errorf("restore previous Pages current target: %w", err))
}
return errors.Join(rollbackErrors...)
}
if err := copyDir(stagingDir, releaseDir); err != nil {
return rollback(fmt.Errorf("copy repaired Pages release: %w", err))
}
if !pagesProjectReleaseReady(releaseDir, project) {
return rollback(errors.New("repaired Pages release is not ready"))
}
if err := switchPagesCurrentDir(currentDir, releaseDir, os.Rename); err != nil {
return rollback(fmt.Errorf("switch Pages current to repaired release: %w", err))
}
if err := os.RemoveAll(backupDir); err != nil {
slog.Warn("failed to remove previous Pages release", "path", backupDir, "error", err)
}
if err := os.RemoveAll(stagingDir); err != nil {
slog.Warn("failed to remove Pages repair staging", "path", stagingDir, "error", err)
}
return nil
}
func pagesCurrentDirFromRelease(releaseDir string) string {
return filepath.Join(filepath.Dir(filepath.Dir(releaseDir)), "current")
}
func pagesCurrentTargetsRelease(currentDir string, releaseDir string) (bool, error) {
if _, err := os.Lstat(currentDir); err != nil {
if os.IsNotExist(err) {
return false, nil
}
return false, err
}
currentInfo, err := os.Stat(currentDir)
if err != nil {
if os.IsNotExist(err) {
return false, nil
}
return false, fmt.Errorf("stat Pages current target: %w", err)
}
releaseInfo, err := os.Stat(releaseDir)
if err != nil {
if os.IsNotExist(err) {
return false, nil
}
return false, err
}
return os.SameFile(currentInfo, releaseInfo), nil
}
func removePagesStagingUnlessCurrent(stagingDir string, currentDir string) {
currentUsesStaging, err := pagesCurrentTargetsRelease(currentDir, stagingDir)
if err == nil && currentUsesStaging {
slog.Error("preserving Pages staging because current still references it", "path", stagingDir)
return
}
if removeErr := os.RemoveAll(stagingDir); removeErr != nil {
slog.Warn("failed to remove Pages staging", "path", stagingDir, "error", removeErr)
}
}
func verifyPagesCurrentTarget(currentDir string, releaseDir string) error {
currentInfo, err := os.Stat(currentDir)
if err != nil {
return fmt.Errorf("stat Pages current target: %w", err)
}
releaseInfo, err := os.Stat(releaseDir)
if err != nil {
return fmt.Errorf("stat Pages release target: %w", err)
}
if !os.SameFile(currentInfo, releaseInfo) {
return fmt.Errorf("pages current target does not resolve to release %s", releaseDir)
}
return nil
}
func switchPagesCurrentDir(
currentDir string,
releaseDir string,
rename func(string, string) error,
) error {
return switchPagesCurrentDirWithOps(currentDir, releaseDir, rename, os.Symlink)
}
func switchPagesCurrentDirWithOps(
currentDir string,
releaseDir string,
rename func(string, string) error,
symlink func(string, string) error,
) error {
if err := os.MkdirAll(filepath.Dir(currentDir), pagesDirPerm); err != nil {
return err
}
currentInfo, currentErr := os.Lstat(currentDir)
if currentErr != nil && !os.IsNotExist(currentErr) {
return currentErr
}
if currentErr == nil && currentInfo.Mode()&os.ModeSymlink == 0 {
return fallbackCopyPagesCurrentDir(currentDir, releaseDir, rename)
}
previousTarget := ""
hadPrevious := currentErr == nil
if hadPrevious {
var err error
previousTarget, err = os.Readlink(currentDir)
if err != nil {
return err
}
}
relTarget, err := filepath.Rel(filepath.Dir(currentDir), releaseDir)
if err != nil {
@@ -413,44 +804,123 @@ func switchPagesProjectCurrentDir(baseDir string, projectID uint, releaseDir str
}
tmpSymlink := currentDir + ".tmp"
_ = os.Remove(tmpSymlink)
symlinkErr := os.Symlink(relTarget, tmpSymlink)
if symlinkErr != nil {
return fallbackCopyPagesCurrentDir(currentDir, previousDir, releaseDir)
}
_ = os.Remove(tmpSymlink)
if _, err := os.Lstat(currentDir); err == nil {
if err := os.Rename(currentDir, previousDir); err != nil {
return err
}
}
if err := os.Symlink(relTarget, currentDir); err != nil {
if _, restoreErr := os.Lstat(previousDir); restoreErr == nil {
_ = os.Rename(previousDir, currentDir)
}
if err := os.Remove(tmpSymlink); err != nil && !os.IsNotExist(err) {
return err
}
_ = os.RemoveAll(previousDir)
if err := symlink(relTarget, tmpSymlink); err != nil {
_ = os.Remove(tmpSymlink)
return fallbackCopyPagesCurrentDir(currentDir, releaseDir, rename)
}
defer func() { _ = os.Remove(tmpSymlink) }()
if err := verifyPagesCurrentTarget(tmpSymlink, releaseDir); err != nil {
return err
}
if err := rename(tmpSymlink, currentDir); err != nil {
return err
}
if err := verifyPagesCurrentTarget(currentDir, releaseDir); err != nil {
rollbackErr := rollbackPagesCurrentSymlink(
currentDir,
previousTarget,
hadPrevious,
rename,
)
return errors.Join(err, rollbackErr)
}
return nil
}
func fallbackCopyPagesCurrentDir(currentDir, previousDir, releaseDir string) error {
if _, err := os.Lstat(currentDir); err == nil {
if err := os.Rename(currentDir, previousDir); err != nil {
return err
}
}
if err := copyPagesDir(releaseDir, currentDir); err != nil {
_ = os.RemoveAll(currentDir)
if _, restoreErr := os.Lstat(previousDir); restoreErr == nil {
_ = os.Rename(previousDir, currentDir)
}
func fallbackCopyPagesCurrentDir(
currentDir string,
releaseDir string,
rename func(string, string) error,
) error {
stagingDir := currentDir + ".copy.tmp"
previousDir := currentDir + ".previous"
if err := os.RemoveAll(stagingDir); err != nil {
return err
}
_ = os.RemoveAll(previousDir)
if err := copyPagesDir(releaseDir, stagingDir); err != nil {
_ = os.RemoveAll(stagingDir)
return err
}
if err := os.RemoveAll(previousDir); err != nil {
_ = os.RemoveAll(stagingDir)
return err
}
hadPrevious := false
if _, err := os.Lstat(currentDir); err == nil {
if err := rename(currentDir, previousDir); err != nil {
_ = os.RemoveAll(stagingDir)
return err
}
hadPrevious = true
} else if !os.IsNotExist(err) {
_ = os.RemoveAll(stagingDir)
return err
}
if err := rename(stagingDir, currentDir); err != nil {
var restoreErr error
if hadPrevious {
restoreErr = rename(previousDir, currentDir)
}
_ = os.RemoveAll(stagingDir)
return errors.Join(err, restoreErr)
}
if err := os.RemoveAll(previousDir); err != nil {
slog.Warn("failed to remove previous Pages current directory", "path", previousDir, "error", err)
}
return nil
}
func switchPagesProjectCurrentDir(baseDir string, projectID uint, releaseDir string) error {
return switchPagesProjectCurrentDirWithRename(baseDir, projectID, releaseDir, os.Rename)
}
func switchPagesProjectCurrentDirWithRename(
baseDir string,
projectID uint,
releaseDir string,
rename func(string, string) error,
) error {
return switchPagesCurrentDir(pagesProjectCurrentDir(baseDir, projectID), releaseDir, rename)
}
func rollbackPagesCurrentSymlink(
currentDir string,
previousTarget string,
hadPrevious bool,
rename func(string, string) error,
) error {
if !hadPrevious {
if err := os.Remove(currentDir); err != nil && !os.IsNotExist(err) {
return fmt.Errorf("remove unverified Pages current symlink: %w", err)
}
return nil
}
rollbackSymlink := currentDir + ".rollback.tmp"
if err := os.Remove(rollbackSymlink); err != nil && !os.IsNotExist(err) {
return err
}
if err := os.Symlink(previousTarget, rollbackSymlink); err != nil {
return err
}
defer func() { _ = os.Remove(rollbackSymlink) }()
if err := rename(rollbackSymlink, currentDir); err != nil {
return fmt.Errorf("restore previous Pages current symlink: %w", err)
}
gotTarget, err := os.Readlink(currentDir)
if err != nil {
return fmt.Errorf("verify restored Pages current symlink: %w", err)
}
if gotTarget != previousTarget {
return fmt.Errorf(
"restored Pages current symlink target %q does not match %q",
gotTarget,
previousTarget,
)
}
return nil
}
@@ -467,24 +937,30 @@ func copyPagesDir(sourceDir string, targetDir string) error {
if entry.IsDir() {
return os.MkdirAll(targetPath, pagesDirPerm)
}
input, err := os.Open(sourcePath) //nolint:gosec // sourcePath is under managed PagesDir walk root
if err != nil {
return err
}
defer func() { _ = input.Close() }()
if err := os.MkdirAll(filepath.Dir(targetPath), pagesDirPerm); err != nil {
return err
}
output, err := os.OpenFile(targetPath, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, pagesFilePerm) //nolint:gosec // targetPath is under managed PagesDir walk root
if err != nil {
return err
}
defer func() { _ = output.Close() }()
_, err = io.Copy(output, input)
return err
return copyPagesFile(sourcePath, targetPath)
})
}
func copyPagesFile(sourcePath string, targetPath string) error {
input, err := os.Open(sourcePath) //nolint:gosec // sourcePath is under managed PagesDir walk root
if err != nil {
return err
}
if err := os.MkdirAll(filepath.Dir(targetPath), pagesDirPerm); err != nil {
_ = input.Close()
return err
}
output, err := os.OpenFile(targetPath, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, pagesFilePerm) //nolint:gosec // targetPath is under managed PagesDir walk root
if err != nil {
_ = input.Close()
return err
}
_, copyErr := io.Copy(output, input)
outputCloseErr := output.Close()
inputCloseErr := input.Close()
return errors.Join(copyErr, outputCloseErr, inputCloseErr)
}
func markerMatches(dir string, project pagesProjectRef) bool {
data, err := os.ReadFile(filepath.Join(dir, ".openflare-pages.json")) //nolint:gosec // dir is managed PagesDir
if err != nil {
@@ -515,8 +991,3 @@ func pagesProjectCurrentDir(baseDir string, projectID uint) string {
func pagesProjectReleaseDir(baseDir string, projectID uint, checksum string) string {
return filepath.Join(baseDir, "projects", fmt.Sprintf("%d", projectID), "releases", checksum)
}
func checksumBytes(data []byte) string {
sum := sha256.Sum256(data)
return hex.EncodeToString(sum[:])
}
@@ -0,0 +1,388 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package sync
import (
"context"
"errors"
"os"
"path/filepath"
"strings"
"testing"
"github.com/Rain-kl/Wavelet/internal/apps/agent/protocol"
"github.com/Rain-kl/Wavelet/internal/apps/agent/state"
)
func TestEnsurePagesProjectRejectsMetadataBeyondAgentCapsBeforeDownload(t *testing.T) {
packageBytes := testPagesPackage(t, map[string]string{"index.html": "x"})
base := protocol.PagesProjectLatestHashResponse{
ProjectID: 1,
DeploymentID: 1,
Hash: testBytesChecksum(packageBytes),
PackageSize: int64(len(packageBytes)),
FileCount: 1,
TotalSize: 1,
}
tests := []struct {
name string
mutate func(*protocol.PagesProjectLatestHashResponse)
}{
{
name: "package size",
mutate: func(metadata *protocol.PagesProjectLatestHashResponse) {
metadata.PackageSize = agentPagesMaxPackageBytes + 1
},
},
{
name: "file count",
mutate: func(metadata *protocol.PagesProjectLatestHashResponse) {
metadata.FileCount = agentPagesMaxFiles + 1
},
},
{
name: "total size",
mutate: func(metadata *protocol.PagesProjectLatestHashResponse) {
metadata.TotalSize = agentPagesMaxTotalBytes + 1
},
},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
metadata := base
test.mutate(&metadata)
client := &fakeClient{
pagesPackages: map[uint][]byte{1: packageBytes},
pagesMetadata: map[uint]protocol.PagesProjectLatestHashResponse{1: metadata},
}
service := New(client, &fakeManager{}, nil)
service.SetPagesDir(t.TempDir())
err := service.ensurePagesProject(context.Background(), &state.Snapshot{}, 1)
if err == nil || !strings.Contains(err.Error(), "agent limit") {
t.Fatalf("ensurePagesProject(%s metadata) error = %v, want agent limit error", test.name, err)
}
if client.pagesPackageDownloads != 0 {
t.Errorf("ensurePagesProject(%s metadata) downloads = %d, want 0", test.name, client.pagesPackageDownloads)
}
})
}
}
func TestEnsurePagesProjectRetriesSameHashDifferentDeployment(t *testing.T) {
packageBytes := testPagesPackage(t, map[string]string{"index.html": "same"})
hash := testBytesChecksum(packageBytes)
client := &racingLatestClient{
pkgA: packageBytes,
pkgB: packageBytes,
hashA: hash,
hashB: hash,
}
service := New(client, &fakeManager{}, nil)
pagesDir := t.TempDir()
service.SetPagesDir(pagesDir)
snapshot := &state.Snapshot{PagesDeployments: []state.PagesDeployment{{ProjectID: 42}}}
if err := service.ensurePagesProject(context.Background(), snapshot, 42); err != nil {
t.Fatalf("ensurePagesProject(same hash deployment race) error = %v", err)
}
if client.downloadCalls != 2 {
t.Errorf("ensurePagesProject(same hash deployment race) downloads = %d, want 2", client.downloadCalls)
}
if snapshot.PagesDeployments[0].DeploymentID != 2 || snapshot.PagesDeployments[0].Hash != hash {
t.Errorf("snapshot Pages deployment = %+v, want deployment 2/hash %s", snapshot.PagesDeployments[0], hash)
}
}
func TestEnsurePagesProjectExtractionFailureCleansTempAndPreservesCurrent(t *testing.T) {
projectID := uint(9)
oldPackage := testPagesPackage(t, map[string]string{"index.html": "old"})
oldHash := testBytesChecksum(oldPackage)
newPackage := testPagesPackage(t, map[string]string{"index.html": "new"})
newHash := testBytesChecksum(newPackage)
pagesDir := t.TempDir()
oldRelease := pagesProjectReleaseDir(pagesDir, projectID, oldHash)
if err := extractTestPagesPackage(t, oldPackage, oldRelease, pagesProjectRef{
ProjectID: projectID,
DeploymentID: 1,
Checksum: oldHash,
}); err != nil {
t.Fatalf("extractTestPagesPackage(old) error = %v", err)
}
if err := switchPagesProjectCurrentDir(pagesDir, projectID, oldRelease); err != nil {
t.Fatalf("switchPagesProjectCurrentDir(old) error = %v", err)
}
client := &fakeClient{
pagesPackages: map[uint][]byte{projectID: newPackage},
pagesMetadata: map[uint]protocol.PagesProjectLatestHashResponse{
projectID: {
ProjectID: projectID,
DeploymentID: 2,
Hash: newHash,
PackageSize: int64(len(newPackage)),
FileCount: 1,
TotalSize: 2, // Smaller than the actual three-byte file.
},
},
}
service := New(client, &fakeManager{}, nil)
service.SetPagesDir(pagesDir)
err := service.ensurePagesProject(context.Background(), &state.Snapshot{}, projectID)
if err == nil {
t.Fatal("ensurePagesProject(metadata-tightened extraction) error = nil, want error")
}
current, readErr := os.ReadFile(pagesProjectCurrentDir(pagesDir, projectID) + "/index.html")
if readErr != nil {
t.Fatalf("read old current after failed extraction error = %v", readErr)
}
if string(current) != "old" {
t.Errorf("current content after failed extraction = %q, want %q", current, "old")
}
entries, readErr := os.ReadDir(filepath.Join(pagesDir, "projects", "9", "releases"))
if readErr != nil {
t.Fatalf("read releases after failed extraction error = %v", readErr)
}
if len(entries) != 1 || entries[0].Name() != oldHash {
t.Errorf("releases after failed extraction = %v, want only %s", entries, oldHash)
}
}
func TestEnsurePagesProjectAcceptsAllZeroByteFiles(t *testing.T) {
packageBytes := testPagesPackage(t, map[string]string{
"index.html": "",
".gitkeep": "",
})
client := &fakeClient{pagesPackages: map[uint][]byte{5: packageBytes}}
service := New(client, &fakeManager{}, nil)
pagesDir := t.TempDir()
service.SetPagesDir(pagesDir)
if err := service.ensurePagesProject(context.Background(), &state.Snapshot{}, 5); err != nil {
t.Fatalf("ensurePagesProject(all-zero files) error = %v", err)
}
for _, name := range []string{"index.html", ".gitkeep"} {
info, err := os.Stat(filepath.Join(pagesProjectCurrentDir(pagesDir, 5), name))
if err != nil {
t.Errorf("stat all-zero file %q error = %v", name, err)
continue
}
if info.Size() != 0 {
t.Errorf("all-zero file %q size = %d, want 0", name, info.Size())
}
}
}
func TestSwitchPagesProjectCurrentDirRenameFailureKeepsPreviousCurrent(t *testing.T) {
pagesDir := t.TempDir()
projectID := uint(21)
oldRelease := pagesProjectReleaseDir(pagesDir, projectID, "old")
newRelease := pagesProjectReleaseDir(pagesDir, projectID, "new")
for path, content := range map[string]string{
oldRelease: "old",
newRelease: "new",
} {
if err := os.MkdirAll(path, pagesDirPerm); err != nil {
t.Fatalf("mkdir release %q error = %v", path, err)
}
if err := os.WriteFile(filepath.Join(path, "index.html"), []byte(content), pagesFilePerm); err != nil {
t.Fatalf("write release %q error = %v", path, err)
}
}
if err := switchPagesProjectCurrentDir(pagesDir, projectID, oldRelease); err != nil {
t.Fatalf("seed previous current error = %v", err)
}
currentDir := pagesProjectCurrentDir(pagesDir, projectID)
renameErr := errors.New("injected current rename failure")
err := switchPagesProjectCurrentDirWithRename(
pagesDir,
projectID,
newRelease,
func(oldPath string, newPath string) error {
if oldPath == currentDir+".tmp" && newPath == currentDir {
return renameErr
}
return os.Rename(oldPath, newPath)
},
)
if !errors.Is(err, renameErr) {
t.Fatalf("switchPagesProjectCurrentDirWithRename() error = %v, want injected rename error", err)
}
current, err := os.ReadFile(filepath.Join(currentDir, "index.html"))
if err != nil {
t.Fatalf("read previous current after rename failure error = %v", err)
}
if string(current) != "old" {
t.Errorf("current after rename failure = %q, want %q", current, "old")
}
if _, err := os.Lstat(currentDir + ".tmp"); !os.IsNotExist(err) {
t.Errorf("temporary current symlink remains after rename failure: %v", err)
}
}
func TestPromoteSameHashReleaseFailureRestoresPreviousCurrent(t *testing.T) {
pagesDir := t.TempDir()
projectID := uint(22)
hash := "same-hash"
project := pagesProjectRef{ProjectID: projectID, DeploymentID: 2, Checksum: hash}
releaseDir := pagesProjectReleaseDir(pagesDir, projectID, hash)
if err := os.MkdirAll(releaseDir, pagesDirPerm); err != nil {
t.Fatalf("mkdir previous same-hash release error = %v", err)
}
if err := os.WriteFile(filepath.Join(releaseDir, "index.html"), []byte("old"), pagesFilePerm); err != nil {
t.Fatalf("write previous same-hash release error = %v", err)
}
if err := writePagesMarker(releaseDir, project); err != nil {
t.Fatalf("write previous same-hash marker error = %v", err)
}
if err := switchPagesProjectCurrentDir(pagesDir, projectID, releaseDir); err != nil {
t.Fatalf("seed same-hash current error = %v", err)
}
stagingDir, err := os.MkdirTemp(filepath.Dir(releaseDir), ".same-hash-*.tmp")
if err != nil {
t.Fatalf("create same-hash staging error = %v", err)
}
if err := os.WriteFile(filepath.Join(stagingDir, "index.html"), []byte("new"), pagesFilePerm); err != nil {
t.Fatalf("write repaired same-hash release error = %v", err)
}
if err := writePagesMarker(stagingDir, project); err != nil {
t.Fatalf("write repaired same-hash marker error = %v", err)
}
copyErr := errors.New("injected same-hash copy failure")
err = promotePagesReleaseWithCopy(
stagingDir,
releaseDir,
project,
func(_ string, targetDir string) error {
if err := os.MkdirAll(targetDir, pagesDirPerm); err != nil {
return err
}
if err := os.WriteFile(filepath.Join(targetDir, "index.html"), []byte("partial"), pagesFilePerm); err != nil {
return err
}
return copyErr
},
)
if !errors.Is(err, copyErr) {
t.Fatalf("promotePagesReleaseWithCopy() error = %v, want injected copy error", err)
}
for name, path := range map[string]string{
"current": filepath.Join(pagesProjectCurrentDir(pagesDir, projectID), "index.html"),
"release": filepath.Join(releaseDir, "index.html"),
} {
content, readErr := os.ReadFile(path)
if readErr != nil {
t.Fatalf("read restored %s after same-hash repair failure error = %v", name, readErr)
}
if string(content) != "old" {
t.Errorf("restored %s after same-hash repair failure = %q, want %q", name, content, "old")
}
}
if _, err := os.Stat(stagingDir); !os.IsNotExist(err) {
t.Errorf("same-hash staging remains after successful rollback: %v", err)
}
}
func TestPromotePagesReleaseRepairsDanglingCurrent(t *testing.T) {
pagesDir := t.TempDir()
projectID := uint(23)
releaseDir := pagesProjectReleaseDir(pagesDir, projectID, "new-hash")
currentDir := pagesProjectCurrentDir(pagesDir, projectID)
requireTestMkdirAll(t, filepath.Dir(currentDir))
relTarget, err := filepath.Rel(filepath.Dir(currentDir), releaseDir)
if err != nil {
t.Fatalf("relative release target error = %v", err)
}
if err := os.Symlink(relTarget, currentDir); err != nil {
t.Skipf("symlink unsupported: %v", err)
}
requireTestMkdirAll(t, filepath.Dir(releaseDir))
stagingDir, err := os.MkdirTemp(filepath.Dir(releaseDir), ".dangling-*.tmp")
if err != nil {
t.Fatalf("create dangling repair staging error = %v", err)
}
if err := os.WriteFile(filepath.Join(stagingDir, "index.html"), []byte("repaired"), pagesFilePerm); err != nil {
t.Fatalf("write dangling repair staging error = %v", err)
}
project := pagesProjectRef{ProjectID: projectID, DeploymentID: 1, Checksum: "new-hash"}
if err := writePagesMarker(stagingDir, project); err != nil {
t.Fatalf("write dangling repair marker error = %v", err)
}
if err := promotePagesRelease(stagingDir, releaseDir, project); err != nil {
t.Fatalf("promotePagesRelease(dangling current) error = %v", err)
}
content, err := os.ReadFile(filepath.Join(currentDir, "index.html"))
if err != nil {
t.Fatalf("read repaired dangling current error = %v", err)
}
if string(content) != "repaired" {
t.Errorf("repaired dangling current = %q, want %q", content, "repaired")
}
}
func TestSwitchPagesCurrentDirCopiesOverLegacyDirectory(t *testing.T) {
pagesDir := t.TempDir()
currentDir := filepath.Join(pagesDir, "current")
releaseDir := filepath.Join(pagesDir, "releases", "new")
requireTestMkdirAll(t, currentDir)
requireTestMkdirAll(t, releaseDir)
if err := os.WriteFile(filepath.Join(currentDir, "index.html"), []byte("old"), pagesFilePerm); err != nil {
t.Fatalf("write legacy current error = %v", err)
}
if err := os.WriteFile(filepath.Join(releaseDir, "index.html"), []byte("new"), pagesFilePerm); err != nil {
t.Fatalf("write new release error = %v", err)
}
if err := switchPagesCurrentDir(currentDir, releaseDir, os.Rename); err != nil {
t.Fatalf("switchPagesCurrentDir(legacy directory) error = %v", err)
}
content, err := os.ReadFile(filepath.Join(currentDir, "index.html"))
if err != nil {
t.Fatalf("read copied legacy current error = %v", err)
}
if string(content) != "new" {
t.Errorf("copied legacy current = %q, want %q", content, "new")
}
}
func TestSwitchPagesCurrentDirFallsBackWhenSymlinkUnavailable(t *testing.T) {
pagesDir := t.TempDir()
currentDir := filepath.Join(pagesDir, "current")
releaseDir := filepath.Join(pagesDir, "releases", "new")
requireTestMkdirAll(t, releaseDir)
if err := os.WriteFile(filepath.Join(releaseDir, "index.html"), []byte("new"), pagesFilePerm); err != nil {
t.Fatalf("write fallback release error = %v", err)
}
symlinkErr := errors.New("injected symlink unavailable")
if err := switchPagesCurrentDirWithOps(
currentDir,
releaseDir,
os.Rename,
func(string, string) error { return symlinkErr },
); err != nil {
t.Fatalf("switchPagesCurrentDirWithOps(symlink unavailable) error = %v", err)
}
info, err := os.Lstat(currentDir)
if err != nil {
t.Fatalf("lstat copied current error = %v", err)
}
if !info.IsDir() {
t.Errorf("copied current mode = %v, want directory", info.Mode())
}
content, err := os.ReadFile(filepath.Join(currentDir, "index.html"))
if err != nil {
t.Fatalf("read fallback current error = %v", err)
}
if string(content) != "new" {
t.Errorf("fallback current = %q, want %q", content, "new")
}
}
func requireTestMkdirAll(t *testing.T, dir string) {
t.Helper()
if err := os.MkdirAll(dir, pagesDirPerm); err != nil {
t.Fatalf("mkdir %q error = %v", dir, err)
}
}
+6 -2
View File
@@ -1,3 +1,6 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package sync
import (
@@ -7,6 +10,7 @@ import (
"encoding/json"
"errors"
"fmt"
"io"
"log/slog"
"sort"
"strconv"
@@ -31,9 +35,9 @@ const (
type ConfigClient interface {
GetActiveConfig(ctx context.Context) (*protocol.ActiveConfigResponse, error)
GetPagesDeploymentHash(ctx context.Context, deploymentID uint) (string, error)
DownloadPagesDeploymentPackage(ctx context.Context, deploymentID uint) ([]byte, error)
DownloadPagesDeploymentPackage(ctx context.Context, deploymentID uint, dst io.Writer, maxBytes int64) (int64, error)
GetPagesProjectLatestHash(ctx context.Context, projectID uint) (*protocol.PagesProjectLatestHashResponse, error)
DownloadPagesProjectLatestPackage(ctx context.Context, projectID uint) ([]byte, error)
DownloadPagesProjectLatestPackage(ctx context.Context, projectID uint, dst io.Writer, maxBytes int64) (int64, error)
ReportApplyLog(ctx context.Context, payload protocol.ApplyLogPayload) error
SyncWAFIPGroups(ctx context.Context, payload protocol.WAFIPGroupSyncRequest) (*protocol.WAFIPGroupSyncResponse, error)
}
+128 -18
View File
@@ -1,3 +1,6 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package sync
import (
@@ -7,6 +10,7 @@ import (
"crypto/sha256"
"encoding/hex"
"fmt"
"io"
"os"
"path/filepath"
"strings"
@@ -16,6 +20,7 @@ import (
"github.com/Rain-kl/Wavelet/internal/apps/agent/nginx"
"github.com/Rain-kl/Wavelet/internal/apps/agent/protocol"
"github.com/Rain-kl/Wavelet/internal/apps/agent/state"
"github.com/Rain-kl/Wavelet/pkg/pagesarchive"
)
type fakeExecutor struct {
@@ -28,15 +33,17 @@ func testPagesSourceConfigJSON(projectID, deploymentID uint, checksum string) st
}
type fakeClient struct {
config protocol.ActiveConfigResponse
reports []protocol.ApplyLogPayload
wafSyncCalls []protocol.WAFIPGroupSyncRequest
pagesPackages map[uint][]byte // key: project_id (latest package)
pagesHashes map[uint]string // key: project_id
pagesLatestDeployIDs map[uint]uint // key: project_id → deployment_id
wafSyncResult protocol.WAFIPGroupSyncResponse
fetchCalls int
hashCalls int
config protocol.ActiveConfigResponse
reports []protocol.ApplyLogPayload
wafSyncCalls []protocol.WAFIPGroupSyncRequest
pagesPackages map[uint][]byte // key: project_id (latest package)
pagesHashes map[uint]string // key: project_id
pagesLatestDeployIDs map[uint]uint // key: project_id → deployment_id
pagesMetadata map[uint]protocol.PagesProjectLatestHashResponse
pagesPackageDownloads int
wafSyncResult protocol.WAFIPGroupSyncResponse
fetchCalls int
hashCalls int
}
type fakeManager struct {
@@ -98,17 +105,28 @@ func (f *fakeClient) GetPagesDeploymentHash(ctx context.Context, deploymentID ui
return f.projectHash(deploymentID)
}
func (f *fakeClient) DownloadPagesDeploymentPackage(ctx context.Context, deploymentID uint) ([]byte, error) {
func (f *fakeClient) DownloadPagesDeploymentPackage(
ctx context.Context,
deploymentID uint,
dst io.Writer,
maxBytes int64,
) (int64, error) {
for projectID, depID := range f.pagesLatestDeployIDs {
if depID == deploymentID {
return f.projectPackage(projectID)
return f.writeProjectPackage(projectID, dst, maxBytes)
}
}
return f.projectPackage(deploymentID)
return f.writeProjectPackage(deploymentID, dst, maxBytes)
}
func (f *fakeClient) GetPagesProjectLatestHash(ctx context.Context, projectID uint) (*protocol.PagesProjectLatestHashResponse, error) {
f.hashCalls++
if f.pagesMetadata != nil {
if metadata, ok := f.pagesMetadata[projectID]; ok {
result := metadata
return &result, nil
}
}
hash, err := f.projectHash(projectID)
if err != nil {
return nil, err
@@ -119,15 +137,32 @@ func (f *fakeClient) GetPagesProjectLatestHash(ctx context.Context, projectID ui
deploymentID = id
}
}
packageBytes, err := f.projectPackage(projectID)
if err != nil {
return nil, err
}
fileCount, totalSize, err := testPagesPackageStats(packageBytes)
if err != nil {
return nil, err
}
return &protocol.PagesProjectLatestHashResponse{
ProjectID: projectID,
DeploymentID: deploymentID,
Hash: hash,
PackageSize: int64(len(packageBytes)),
FileCount: fileCount,
TotalSize: totalSize,
}, nil
}
func (f *fakeClient) DownloadPagesProjectLatestPackage(ctx context.Context, projectID uint) ([]byte, error) {
return f.projectPackage(projectID)
func (f *fakeClient) DownloadPagesProjectLatestPackage(
ctx context.Context,
projectID uint,
dst io.Writer,
maxBytes int64,
) (int64, error) {
f.pagesPackageDownloads++
return f.writeProjectPackage(projectID, dst, maxBytes)
}
func (f *fakeClient) projectHash(projectID uint) (string, error) {
@@ -155,6 +190,22 @@ func (f *fakeClient) projectPackage(projectID uint) ([]byte, error) {
return packageBytes, nil
}
func (f *fakeClient) writeProjectPackage(projectID uint, dst io.Writer, maxBytes int64) (int64, error) {
packageBytes, err := f.projectPackage(projectID)
if err != nil {
return 0, err
}
limited := &io.LimitedReader{R: bytes.NewReader(packageBytes), N: maxBytes + 1}
written, err := io.Copy(dst, limited)
if err != nil {
return written, err
}
if written > maxBytes {
return written, fmt.Errorf("test pages package exceeds limit %d", maxBytes)
}
return written, nil
}
func (f *fakeClient) ReportApplyLog(ctx context.Context, payload protocol.ApplyLogPayload) error {
f.reports = append(f.reports, payload)
return nil
@@ -572,7 +623,9 @@ func TestSyncOnceRejectsPagesZipSlipBeforeApply(t *testing.T) {
service.SetPagesDir(t.TempDir())
err := service.SyncOnce(context.Background(), &protocol.ActiveConfigMeta{Version: "20260309-102", Checksum: "pages-config-checksum"})
if err == nil || (!strings.Contains(err.Error(), "escapes deployment root") && !strings.Contains(err.Error(), "escapes directory")) {
if err == nil || (!strings.Contains(err.Error(), "escapes deployment root") &&
!strings.Contains(err.Error(), "escapes directory") &&
!strings.Contains(err.Error(), "dot segment")) {
t.Fatalf("expected zip-slip rejection, got %v", err)
}
if len(manager.applyRouteContents) != 0 {
@@ -1316,7 +1369,7 @@ func TestSyncOnceRedownloadsPagesDeploymentWhenServerHashChanges(t *testing.T) {
}
pagesDir := t.TempDir()
releaseDir := pagesProjectReleaseDir(pagesDir, projectID, initialHash)
if err = extractPagesPackage(initialPackage, releaseDir, pagesProjectRef{
if err = extractTestPagesPackage(t, initialPackage, releaseDir, pagesProjectRef{
ProjectID: projectID,
Checksum: initialHash,
}); err != nil {
@@ -1385,20 +1438,42 @@ func (r *racingLatestClient) GetPagesProjectLatestHash(ctx context.Context, proj
// 2: verify after downloading B → B (race)
// 3+: stable on B for retry
hash, dep := r.hashA, uint(1)
packageBytes := r.pkgA
if r.hashCall >= 2 {
hash, dep = r.hashB, 2
packageBytes = r.pkgB
}
fileCount, totalSize, err := testPagesPackageStats(packageBytes)
if err != nil {
return nil, err
}
return &protocol.PagesProjectLatestHashResponse{
ProjectID: projectID,
DeploymentID: dep,
Hash: hash,
PackageSize: int64(len(packageBytes)),
FileCount: fileCount,
TotalSize: totalSize,
}, nil
}
func (r *racingLatestClient) DownloadPagesProjectLatestPackage(ctx context.Context, projectID uint) ([]byte, error) {
func (r *racingLatestClient) DownloadPagesProjectLatestPackage(
ctx context.Context,
projectID uint,
dst io.Writer,
maxBytes int64,
) (int64, error) {
r.downloadCalls++
// Always return package B (what "latest download" would stream mid-race / after).
return r.pkgB, nil
limited := &io.LimitedReader{R: bytes.NewReader(r.pkgB), N: maxBytes + 1}
written, err := io.Copy(dst, limited)
if err != nil {
return written, err
}
if written > maxBytes {
return written, fmt.Errorf("test pages package exceeds limit %d", maxBytes)
}
return written, nil
}
func TestEnsurePagesProjectSurvivesHashPackageRace(t *testing.T) {
@@ -1641,6 +1716,41 @@ func testPagesPackage(t *testing.T, files map[string]string) []byte {
return buffer.Bytes()
}
func extractTestPagesPackage(
t *testing.T,
packageBytes []byte,
releaseDir string,
project pagesProjectRef,
) error {
t.Helper()
packagePath := filepath.Join(t.TempDir(), "pages-package.zip")
if err := os.WriteFile(packagePath, packageBytes, pagesFilePerm); err != nil {
t.Fatalf("write test Pages package error = %v", err)
}
return extractPagesPackageFile(packagePath, releaseDir, project, pagesarchive.Limits{
MaxFiles: agentPagesMaxFiles,
MaxFileBytes: agentPagesMaxFileBytes,
MaxTotalBytes: agentPagesMaxTotalBytes,
}, nil)
}
func testPagesPackageStats(packageBytes []byte) (int, int64, error) {
reader, err := zip.NewReader(bytes.NewReader(packageBytes), int64(len(packageBytes)))
if err != nil {
return 0, 0, err
}
fileCount := 0
totalSize := int64(0)
for _, file := range reader.File {
if file.FileInfo().IsDir() {
continue
}
fileCount++
totalSize += int64(file.UncompressedSize64) //nolint:gosec // test packages are memory-bounded
}
return fileCount, totalSize, nil
}
func testBytesChecksum(data []byte) string {
sum := sha256.Sum256(data)
return hex.EncodeToString(sum[:])
+6 -3
View File
@@ -224,14 +224,17 @@ func GetPagesProjectLatestHashHandler(c *gin.Context) {
if !ok {
return
}
deploymentID, hash, err := pages.GetProjectLatestPackageHash(c.Request.Context(), projectID)
metadata, err := pages.GetProjectLatestPackageMetadata(c.Request.Context(), projectID)
if apiutil.AbortBadRequestOnError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(protocol.PagesProjectLatestHashResponse{
ProjectID: projectID,
DeploymentID: deploymentID,
Hash: hash,
DeploymentID: metadata.DeploymentID,
Hash: metadata.Hash,
PackageSize: metadata.PackageSize,
FileCount: metadata.FileCount,
TotalSize: metadata.TotalSize,
}))
}
@@ -7,9 +7,11 @@ import (
"context"
"errors"
"fmt"
"path"
"strings"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/pkg/pagesarchive"
openrestyrender "github.com/Rain-kl/Wavelet/pkg/render/openresty"
"gorm.io/gorm"
)
@@ -58,23 +60,41 @@ func buildPagesRouteSnapshot(
}
pagesProjectID = route.PagesProjectID
deployment = buildSnapshotPagesDeployment(project, activeDeployment)
deployment, err = buildSnapshotPagesDeployment(project, activeDeployment)
if err != nil {
return "", nil, nil, nil, fmt.Errorf("路由 %s Pages 配置无效: %w", route.SiteName, err)
}
originURL = fmt.Sprintf("openflare-pages://project/%d", project.ID)
return originURL, []string{originURL}, pagesProjectID, deployment, nil
}
func buildSnapshotPagesDeployment(project *model.PagesProject, activeDeployment *model.PagesDeployment) *openrestyrender.PagesDeployment {
func buildSnapshotPagesDeployment(
project *model.PagesProject,
activeDeployment *model.PagesDeployment,
) (*openrestyrender.PagesDeployment, error) {
if project == nil || activeDeployment == nil {
return nil
return nil, errors.New("pages 项目或部署为空")
}
rootDir, err := pagesarchive.NormalizeLogicalPath(strings.TrimSpace(project.RootDir), true)
if err != nil {
return nil, fmt.Errorf("pages 根目录不合法: %w", err)
}
entryFile := strings.TrimSpace(project.EntryFile)
if entryFile == "" {
entryFile = defaultPagesSnapshotEntryFile
}
entryFile, err = pagesarchive.NormalizeLogicalPath(entryFile, false)
if err != nil {
return nil, fmt.Errorf("pages 入口文件不合法: %w", err)
}
fallbackPath := strings.TrimSpace(project.SPAFallbackPath)
if fallbackPath == "" {
fallbackPath = defaultPagesSnapshotFallbackPath
}
localRoot := openrestyrender.PagesProjectLocalRoot(project.ID)
if rootDir != "" {
localRoot = path.Join(localRoot, rootDir)
}
return &openrestyrender.PagesDeployment{
ProjectID: project.ID,
ProjectSlug: strings.TrimSpace(project.Slug),
@@ -90,6 +110,6 @@ func buildSnapshotPagesDeployment(project *model.PagesProject, activeDeployment
APIProxyRewrite: strings.TrimSpace(project.APIProxyRewrite),
// Root is project-scoped so Agents can swap active packages without
// re-publishing main config (nginx root stays stable).
LocalRoot: openrestyrender.PagesProjectLocalRoot(project.ID),
}
LocalRoot: localRoot,
}, nil
}
@@ -28,6 +28,7 @@ func TestBuildSnapshotRoutesPages(t *testing.T) {
Enabled: true,
SPAFallbackEnabled: true,
SPAFallbackPath: "/index.html",
RootDir: "public/site",
EntryFile: "index.html",
}
require.NoError(t, conn.Create(project).Error)
@@ -64,7 +65,7 @@ func TestBuildSnapshotRoutesPages(t *testing.T) {
require.NotNil(t, snapshotRoute.PagesDeployment)
assert.Equal(t, deployment.ID, snapshotRoute.PagesDeployment.DeploymentID)
assert.Equal(t, deployment.Checksum, snapshotRoute.PagesDeployment.Checksum)
assert.Equal(t, "__OPENFLARE_PAGES_DIR__/projects/1/current", snapshotRoute.PagesDeployment.LocalRoot)
assert.Equal(t, "__OPENFLARE_PAGES_DIR__/projects/1/current/public/site", snapshotRoute.PagesDeployment.LocalRoot)
_, err = renderSnapshotConfig(bundle.SnapshotJSON, nil)
require.NoError(t, err)
@@ -78,6 +79,17 @@ func TestBuildSnapshotRoutesPages(t *testing.T) {
require.NotNil(t, decoded.Routes[0].PagesDeployment)
}
func TestBuildSnapshotPagesDeploymentRejectsUnsafeStoredPaths(t *testing.T) {
deployment := &model.PagesDeployment{ID: 1, ProjectID: 1, Checksum: "checksum"}
for _, project := range []*model.PagesProject{
{ID: 1, RootDir: "../escape", EntryFile: "index.html"},
{ID: 1, RootDir: "public", EntryFile: "/index.html"},
} {
_, err := buildSnapshotPagesDeployment(project, deployment)
require.Error(t, err)
}
}
func requireDB(t *testing.T, ctx context.Context) *gorm.DB {
t.Helper()
conn := db.DB(ctx)
+24 -242
View File
@@ -5,211 +5,38 @@ package pages
import (
"context"
"crypto/sha256"
"crypto/tls"
"encoding/hex"
"errors"
"fmt"
"io"
"mime"
"net"
"net/http"
"net/url"
"os"
"path"
"path/filepath"
"strings"
"time"
"github.com/Rain-kl/Wavelet/pkg/pagesarchive"
)
const (
pagesURLDownloadTimeout = 10 * time.Minute
pagesURLMaxRedirects = 5
pagesMagicSniffBytes = 16
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"
pagesBrowserAccept = "text/html,application/xhtml+xml,application/xml;q=0.9,image/avif,image/webp,image/apng,*/*;q=0.8,application/signed-exchange;v=b3;q=0.7"
pagesBrowserAcceptLang = "zh-CN,zh;q=0.9,en-US;q=0.8,en;q=0.7"
pagesBrowserSecCHUA = `"Google Chrome";v="131", "Chromium";v="131", "Not_A Brand";v="24"`
pagesBrowserSecCHUAMobile = "?0"
pagesBrowserSecCHUAPlat = `"Windows"`
)
// downloadPagesPackageFromURL fetches a remote archive with browser-like headers
// and writes it to a temp file. Allows private/LAN hosts and insecure TLS certs
// (self-signed / internal CA) so operators can pull from internal artifact stores.
func downloadPagesPackageFromURL(ctx context.Context, rawURL string, maxPackageBytes int64) (tempPath string, checksum string, size int64, format pagesarchive.Format, fileName string, err error) {
parsed, err := parseAndValidatePagesDownloadURL(rawURL)
if err != nil {
// downloadPagesPackageFromURL is the deprecated one-shot URL adapter. It uses
// the same bounded downloader as persisted sources, with the legacy trusted
// network policy that permits operator-managed internal artifact services.
func downloadPagesPackageFromURL(
ctx context.Context,
rawURL string,
maxPackageBytes int64,
) (tempPath string, checksum string, size int64, format pagesarchive.Format, fileName string, err error) {
if _, err := parseAndValidatePagesDownloadURL(rawURL); err != nil {
return "", "", 0, "", "", err
}
resp, err := doBrowserDownload(ctx, parsed)
candidate, err := FetchRemoteSource(ctx, RemoteSourceRequest{
URL: strings.TrimSpace(rawURL),
NetworkPolicy: RemoteNetworkPolicyTrustedInternal,
MaxPackageBytes: maxPackageBytes,
})
if err != nil {
return "", "", 0, "", "", err
}
defer func() { _ = resp.Body.Close() }()
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
return "", "", 0, "", "", fmt.Errorf("%s: HTTP %d", errPagesPackageURLDownloadFailed, resp.StatusCode)
}
if resp.ContentLength > 0 && resp.ContentLength > maxPackageBytes {
return "", "", 0, "", "", errors.New(errPagesPackageURLTooLarge)
}
fileName = fileNameFromDownload(resp, parsed)
format, _ = pagesarchive.DetectFormatFromName(fileName)
tempPath, checksum, size, err = writeLimitedPackageTemp(resp.Body, format, maxPackageBytes)
if err != nil {
return "", "", 0, "", "", err
}
format, fileName, err = ensurePackageFormat(tempPath, format, fileName)
if err != nil {
_ = os.Remove(tempPath)
return "", "", 0, "", "", err
}
return tempPath, checksum, size, format, fileName, nil
}
func newPagesURLDownloadClient() *http.Client {
transport := &http.Transport{
Proxy: http.ProxyFromEnvironment,
DialContext: (&net.Dialer{
Timeout: pagesURLDialTimeout,
KeepAlive: pagesURLDialTimeout,
}).DialContext,
ForceAttemptHTTP2: true,
MaxIdleConns: 32,
IdleConnTimeout: 90 * time.Second,
TLSHandshakeTimeout: pagesURLTLSHandshake,
ExpectContinueTimeout: time.Second,
// Allow self-signed / internal certificates for artifact hosts.
TLSClientConfig: &tls.Config{InsecureSkipVerify: true}, //nolint:gosec // intentional for internal/dev artifact URLs
}
client := &http.Client{
Timeout: pagesURLDownloadTimeout,
Transport: transport,
}
client.CheckRedirect = func(req *http.Request, via []*http.Request) error {
if len(via) >= pagesURLMaxRedirects {
return errors.New(errPagesPackageURLDownloadFailed)
if strings.Contains(err.Error(), errPagesSourceRemoteURLInvalid) {
return "", "", 0, "", "", errors.New(errPagesPackageURLInvalid)
}
if err := validatePagesDownloadURLValue(req.URL); err != nil {
return err
}
applyBrowserDownloadHeaders(req, via[0].URL.String())
return nil
}
return client
}
func doBrowserDownload(ctx context.Context, parsed *url.URL) (*http.Response, error) {
client := newPagesURLDownloadClient()
req, err := http.NewRequestWithContext(ctx, http.MethodGet, parsed.String(), nil)
if err != nil {
return nil, errors.New(errPagesPackageURLInvalid)
}
applyBrowserDownloadHeaders(req, "")
resp, err := client.Do(req) //nolint:gosec // scheme validated; private hosts and insecure TLS intentionally allowed
if err != nil {
return nil, fmt.Errorf("%s: %w", errPagesPackageURLDownloadFailed, err)
}
return resp, nil
}
func writeLimitedPackageTemp(body io.Reader, format pagesarchive.Format, maxPackageBytes int64) (tempPath, checksum string, size int64, err error) {
temp, err := os.CreateTemp("", "openflare-pages-url-*."+safeTempSuffixOrBin(format))
if err != nil {
return "", "", 0, err
}
tempPath = temp.Name()
defer func() {
_ = temp.Close()
if err != nil {
_ = os.Remove(tempPath)
}
}()
hash := sha256.New()
written, copyErr := io.Copy(io.MultiWriter(temp, hash), io.LimitReader(body, maxPackageBytes+1))
if copyErr != nil {
err = fmt.Errorf("%s: %w", errPagesPackageURLDownloadFailed, copyErr)
return "", "", 0, err
}
if written > maxPackageBytes {
err = errors.New(errPagesPackageURLTooLarge)
return "", "", 0, err
}
if written == 0 {
err = errors.New(errPagesPackageEmpty)
return "", "", 0, err
}
return tempPath, hex.EncodeToString(hash.Sum(nil)), written, nil
}
func ensurePackageFormat(tempPath string, format pagesarchive.Format, fileName string) (pagesarchive.Format, string, error) {
if format != "" {
return format, fileName, nil
}
detected, ok := sniffPackageFormat(tempPath)
if !ok {
return "", fileName, errors.New(errPagesPackageUnsupported)
}
if !strings.Contains(strings.ToLower(fileName), ".") {
fileName = fileName + "." + pagesarchive.Extension(detected)
}
return detected, fileName, nil
}
func sniffPackageFormat(tempPath string) (pagesarchive.Format, bool) {
file, err := os.Open(tempPath) //nolint:gosec // temp path created by us
if err != nil {
return "", false
}
defer func() { _ = file.Close() }()
head := make([]byte, pagesMagicSniffBytes)
n, _ := io.ReadFull(file, head)
if n <= 0 {
return "", false
}
return pagesarchive.DetectFormatFromBytes(head[:n])
}
func safeTempSuffixOrBin(format pagesarchive.Format) string {
if format == "" {
return "bin"
}
return safeTempSuffix(format)
}
func applyBrowserDownloadHeaders(req *http.Request, referer string) {
if req == nil {
return
}
req.Header.Set("User-Agent", pagesBrowserUserAgent)
req.Header.Set("Accept", pagesBrowserAccept)
req.Header.Set("Accept-Language", pagesBrowserAcceptLang)
req.Header.Set("Cache-Control", "no-cache")
req.Header.Set("Pragma", "no-cache")
req.Header.Set("Upgrade-Insecure-Requests", "1")
req.Header.Set("Sec-Fetch-Dest", "document")
req.Header.Set("Sec-Fetch-Mode", "navigate")
req.Header.Set("Sec-Fetch-Site", "none")
req.Header.Set("Sec-Fetch-User", "?1")
req.Header.Set("Sec-Ch-Ua", pagesBrowserSecCHUA)
req.Header.Set("Sec-Ch-Ua-Mobile", pagesBrowserSecCHUAMobile)
req.Header.Set("Sec-Ch-Ua-Platform", pagesBrowserSecCHUAPlat)
if referer != "" {
req.Header.Set("Referer", referer)
req.Header.Set("Sec-Fetch-Site", "cross-site")
return
}
if req.URL != nil {
req.Header.Set("Referer", req.URL.Scheme+"://"+req.URL.Host+"/")
return "", "", 0, "", "", err
}
// Ownership transfers to the existing one-shot caller, which removes the
// temporary file after the candidate deployment has been created.
return candidate.TempPath, candidate.Checksum, candidate.PackageSize, candidate.Format, candidate.SafeLabel, nil
}
func parseAndValidatePagesDownloadURL(raw string) (*url.URL, error) {
@@ -218,57 +45,12 @@ func parseAndValidatePagesDownloadURL(raw string) (*url.URL, error) {
return nil, errors.New(errPagesPackageURLRequired)
}
parsed, err := url.Parse(value)
if err != nil {
if err != nil || parsed.User != nil || parsed.Fragment != "" || parsed.Opaque != "" {
return nil, errors.New(errPagesPackageURLInvalid)
}
if err := validatePagesDownloadURLValue(parsed); err != nil {
return nil, err
scheme := strings.ToLower(strings.TrimSpace(parsed.Scheme))
if (scheme != remoteSourceSchemeHTTP && scheme != remoteSourceSchemeHTTPS) || strings.TrimSpace(parsed.Hostname()) == "" {
return nil, errors.New(errPagesPackageURLInvalid)
}
return parsed, nil
}
func validatePagesDownloadURLValue(parsed *url.URL) error {
if parsed == nil {
return errors.New(errPagesPackageURLInvalid)
}
scheme := strings.ToLower(strings.TrimSpace(parsed.Scheme))
if scheme != "http" && scheme != "https" {
return errors.New(errPagesPackageURLInvalid)
}
if strings.TrimSpace(parsed.Hostname()) == "" {
return errors.New(errPagesPackageURLInvalid)
}
return nil
}
func fileNameFromDownload(resp *http.Response, parsed *url.URL) string {
if name := fileNameFromContentDisposition(resp); name != "" {
return name
}
if parsed != nil {
base := path.Base(parsed.Path)
if base != "" && base != "." && base != "/" {
return base
}
}
return "package.bin"
}
func fileNameFromContentDisposition(resp *http.Response) string {
if resp == nil {
return ""
}
cd := resp.Header.Get("Content-Disposition")
if cd == "" {
return ""
}
_, params, err := mime.ParseMediaType(cd)
if err != nil {
return ""
}
name := strings.TrimSpace(params["filename"])
if name == "" {
return ""
}
return path.Base(filepath.ToSlash(name))
}
@@ -10,7 +10,6 @@ import (
"net/http"
"net/http/httptest"
"os"
"strings"
"testing"
"github.com/stretchr/testify/assert"
@@ -48,10 +47,10 @@ func TestDownloadPagesPackageFromURLAllowsPrivateHost(t *testing.T) {
require.NoError(t, zw.Close())
zipBytes := body.Bytes()
var sawBrowserUA bool
var sawProviderUA bool
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if strings.Contains(r.Header.Get("User-Agent"), "Mozilla") {
sawBrowserUA = true
if r.Header.Get("User-Agent") == remoteSourceUserAgent {
sawProviderUA = true
}
w.Header().Set("Content-Disposition", `attachment; filename="remote-site.zip"`)
w.Header().Set("Content-Type", "application/zip")
@@ -59,12 +58,6 @@ func TestDownloadPagesPackageFromURLAllowsPrivateHost(t *testing.T) {
}))
t.Cleanup(server.Close)
req, err := http.NewRequest(http.MethodGet, server.URL+"/pkg.zip", nil)
require.NoError(t, err)
applyBrowserDownloadHeaders(req, "")
assert.Contains(t, req.Header.Get("User-Agent"), "Mozilla")
assert.Contains(t, req.Header.Get("Sec-Fetch-Mode"), "navigate")
tempPath, checksum, size, format, fileName, err := downloadPagesPackageFromURL(
context.Background(),
server.URL+"/pkg.zip",
@@ -72,11 +65,11 @@ func TestDownloadPagesPackageFromURLAllowsPrivateHost(t *testing.T) {
)
require.NoError(t, err)
t.Cleanup(func() { _ = os.Remove(tempPath) })
assert.True(t, sawBrowserUA)
assert.True(t, sawProviderUA)
assert.NotEmpty(t, checksum)
assert.Positive(t, size)
assert.Equal(t, "zip", string(format))
assert.Equal(t, "remote-site.zip", fileName)
assert.Equal(t, "pkg.zip", fileName)
}
func TestUploadDeploymentFromURLPrivateHost(t *testing.T) {
+60 -29
View File
@@ -5,33 +5,64 @@
package pages
const (
errPagesProjectNotFound = "pages 项目不存在"
errPagesSlugExists = "pages 项目标识已存在"
errPagesNameRequired = "pages 项目名称不能为空"
errPagesSlugInvalid = "pages 项目标识只能包含小写字母、数字和连字符"
errPagesDeleteReferenced = "pages 项目已被规则引用,不能删除"
errPagesDeploymentNotFound = "pages 部署不存在"
errPagesDeploymentMismatch = "pages 部署不属于该项目"
errPagesDeleteActiveDeploy = "不能删除当前激活的 Pages 部署"
errPagesPackageMissing = "缺少 Pages 部署包"
errPagesPackageURLRequired = "请填写部署包下载链接"
errPagesPackageURLInvalid = "部署包下载链接无效,仅支持 http/https"
errPagesPackageURLDownloadFailed = "从链接下载部署包失败"
errPagesPackageURLTooLarge = "链接指向的部署包超过大小限制"
errPagesPackageNotZip = "pages 部署包必须是 .zip 文件" // legacy alias kept for tests
errPagesPackageUnsupported = "pages 部署包仅支持 zip、tar.gz、tar.xz、tar.bz2、tar、7z 格式"
errPagesPackageInvalidZip = "pages 部署包不是有效 zip 文件" // legacy alias
errPagesPackageInvalid = "pages 部署包不是有效的压缩文件"
errPagesPackageEmpty = "pages 部署包不能为空"
errPagesPackageExtractedTooLarge = "pages 部署包展开后体积超过限制"
errPagesPackageFileTooLarge = "pages 部署包内文件过大"
errPagesAPIProxyPathRequired = "启用 API 反代时,匹配路径不能为空"
errPagesAPIProxyPathPrefix = "API 反代匹配路径必须以 '/' 开头"
errPagesAPIProxyPassRequired = "启用 API 反代时,后端服务地址不能为空" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errPagesAPIProxyPassInvalid = "API 反代后端服务地址必须是有效的 HTTP/HTTPS URL" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errPagesPackagePathEmpty = "pages 部署包路径为空"
errPagesPackageUploadMissing = "pages 部署包上传记录不存在"
errPagesPackageNotInActiveConfig = "pages 部署尚未进入激活配置"
errPagesDeploymentHashMissing = "pages 部署包哈希缺失"
errPagesInvalidSnapshotFormat = "配置快照格式无效"
errPagesProjectNotFound = "pages 项目不存在"
errPagesSlugExists = "pages 项目标识已存在"
errPagesNameRequired = "pages 项目名称不能为空"
errPagesSlugInvalid = "pages 项目标识只能包含小写字母、数字和连字符"
errPagesDeleteReferenced = "pages 项目已被规则引用,不能删除"
errPagesDeploymentNotFound = "pages 部署不存在"
errPagesDeploymentMismatch = "pages 部署不属于该项目"
errPagesDeleteActiveDeploy = "不能删除当前激活的 Pages 部署"
errPagesPackageMissing = "缺少 Pages 部署包"
errPagesPackageURLRequired = "请填写部署包下载链接"
errPagesPackageURLInvalid = "部署包下载链接无效,仅支持 http/https"
errPagesPackageURLDownloadFailed = "从链接下载部署包失败"
errPagesPackageURLTooLarge = "链接指向的部署包超过大小限制"
errPagesPackageNotZip = "pages 部署包必须是 .zip 文件" // legacy alias kept for tests
errPagesPackageUnsupported = "pages 部署包仅支持 zip、tar.gz、tar.xz、tar.bz2、tar、7z 格式"
errPagesPackageInvalidZip = "pages 部署包不是有效 zip 文件" // legacy alias
errPagesPackageInvalid = "pages 部署包不是有效的压缩文件"
errPagesPackageEmpty = "pages 部署包不能为空"
errPagesPackageExtractedTooLarge = "pages 部署包展开后体积超过限制"
errPagesPackageFileTooLarge = "pages 部署包内文件过大"
errPagesAPIProxyPathRequired = "启用 API 反代时,匹配路径不能为空"
errPagesAPIProxyPathPrefix = "API 反代匹配路径必须以 '/' 开头"
errPagesAPIProxyPassRequired = "启用 API 反代时,后端服务地址不能为空" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errPagesAPIProxyPassInvalid = "API 反代后端服务地址必须是有效的 HTTP/HTTPS URL" //nolint:gosec // false positive: this is an error message, not hardcoded credentials
errPagesPackagePathEmpty = "pages 部署包路径为空"
errPagesPackageUploadMissing = "pages 部署包上传记录不存在"
errPagesPackageNotInActiveConfig = "pages 部署尚未进入激活配置"
errPagesDeploymentHashMissing = "pages 部署包哈希缺失"
errPagesInvalidSnapshotFormat = "配置快照格式无效"
errPagesActorMissing = "无法识别当前用户"
errPagesEntryFileMissing = "当前激活部署中不存在指定入口文件"
errPagesSourceNotFound = "pages 部署源不存在"
errPagesSourceTypeRequired = "请选择 pages 部署源类型"
errPagesSourceTypeUnsupported = "pages 部署源类型不受支持"
errPagesSourceRemoteFields = "远程地址来源不能包含 GitHub 或自动更新配置"
errPagesSourceRemoteURLRequired = "请提供远程部署包地址"
errPagesSourceRemoteURLMode = "remote_url_set 与 remote_url 参数不匹配"
errPagesSourceRemoteURLInvalid = "远程部署包地址无效,仅支持不含用户信息和片段的 http/https 地址"
errPagesSourceNetworkPolicy = "远程地址网络策略仅支持 public 或 trusted_internal"
errPagesSourceGitHubFields = "GitHub Release 来源不能包含远程地址配置"
errPagesSourceRepositoryInvalid = "GitHub 仓库地址无效,仅支持 https://github.com/{owner}/{repo}"
errPagesSourceSelectorInvalid = "GitHub Release 选择方式无效"
errPagesSourceAssetNameInvalid = "GitHub Release 资源名称必须是安全的文件名"
errPagesSourceCheckInterval = "GitHub latest 检查间隔必须在 5 到 1440 分钟之间"
errPagesSourceAutoNotAvailable = "自动更新将在后续阶段开放,当前必须保持关闭"
errPagesSourceReleaseNotFound = "未找到符合配置的 GitHub Release 资源"
errPagesSourceDigestInvalid = "GitHub Release 资源摘要格式无效"
errPagesSourceDigestMismatch = "GitHub Release 资源摘要校验失败"
errPagesSourceConfirmationNeeded = "检测到同一 Release 的资源已被替换,请刷新并确认当前版本"
errPagesSourceConfirmationStale = "确认的版本已变化,请刷新后重新确认"
errPagesSourceInitialCheckWarning = "部署源已保存,但首次检查任务入队失败,请稍后手动检查"
errPagesSourceCheckUnsupported = "远程地址来源不支持检查更新,请使用立即同步"
errPagesSourceActionBusy = "pages 部署源任务正在执行"
errPagesSourceActionInvalid = "pages 部署源任务参数无效"
errPagesSourceActionStale = "pages 部署源配置已变化,本次任务已跳过"
errPagesSourceLeaseLost = "pages 部署源任务执行权已失效"
errPagesSourceLeaseExpired = "上次 pages 部署源任务租约已过期"
errPagesSourceSyncFailed = "pages 部署源同步失败"
errPagesSourceTaskDispatchFailed = "pages 部署源任务入队失败"
errPagesSourceInternal = "pages 部署源操作失败,请稍后重试"
)
@@ -0,0 +1,347 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pages
import (
"context"
"crypto/sha256"
"encoding/binary"
"encoding/hex"
"errors"
"net/url"
"path"
"regexp"
"strings"
"time"
"unicode"
"unicode/utf8"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/pkg/logger"
"gorm.io/gorm"
"gorm.io/gorm/clause"
)
const (
githubReleaseSelectorLatest = "latest"
githubReleaseSelectorTag = "tag"
githubSourceIdentityDomain = "openflare:pages:github-release:v2"
initialCheckRetryDelay = 5 * time.Minute
githubRepositoryPathParts = 2
githubCheckJitterRange = 301
githubCheckJitterCenter = 150
)
var (
githubOwnerPattern = regexp.MustCompile(`^[A-Za-z0-9](?:[A-Za-z0-9-]{0,37}[A-Za-z0-9])?$`)
githubRepoPattern = regexp.MustCompile(`^[A-Za-z0-9._-]+$`)
)
type githubSourceConfig struct {
Repository string
Selector string
Tag string
AssetName string
AutoUpdate bool
CheckInterval int
SourceIdentity string
}
func validateGitHubSourceInput(input SourceUpdateInput) error {
if strings.TrimSpace(input.SourceType) != PagesSourceTypeGitHubRelease {
return errors.New(errPagesSourceTypeUnsupported)
}
if input.RemoteURLSet || strings.TrimSpace(input.RemoteURL) != "" ||
strings.TrimSpace(input.RemoteNetworkPolicy) != "" {
return errors.New(errPagesSourceGitHubFields)
}
if _, err := normalizeGitHubRepositoryURL(input.RepositoryURL); err != nil {
return err
}
selector := strings.TrimSpace(input.ReleaseSelector)
if selector == "" {
selector = githubReleaseSelectorLatest
}
assetName := input.AssetName
if assetName == "" {
assetName = defaultGitHubAssetName
}
if !validGitHubAssetName(assetName) {
return errors.New(errPagesSourceAssetNameInvalid)
}
switch selector {
case githubReleaseSelectorLatest:
if input.ReleaseTag != "" {
return errors.New(errPagesSourceSelectorInvalid)
}
interval := input.CheckIntervalMinutes
if interval != 0 && (interval < minimumCheckInterval || interval > maximumCheckInterval) {
return errors.New(errPagesSourceCheckInterval)
}
case githubReleaseSelectorTag:
if !validGitHubReleaseTagConfig(input.ReleaseTag) || input.AutoUpdateEnabled || input.CheckIntervalMinutes != 0 {
return errors.New(errPagesSourceSelectorInvalid)
}
default:
return errors.New(errPagesSourceSelectorInvalid)
}
return nil
}
func buildGitHubSourceConfig(input SourceUpdateInput) (githubSourceConfig, error) {
repository, err := normalizeGitHubRepositoryURL(input.RepositoryURL)
if err != nil {
return githubSourceConfig{}, err
}
selector := strings.TrimSpace(input.ReleaseSelector)
if selector == "" {
selector = githubReleaseSelectorLatest
}
tag := input.ReleaseTag
assetName := input.AssetName
if assetName == "" {
assetName = defaultGitHubAssetName
}
interval := input.CheckIntervalMinutes
if selector == githubReleaseSelectorLatest && interval == 0 {
interval = defaultCheckInterval
}
autoUpdate := input.AutoUpdateEnabled
if selector == githubReleaseSelectorTag {
autoUpdate = false
interval = 0
}
return githubSourceConfig{
Repository: repository,
Selector: selector,
Tag: tag,
AssetName: assetName,
AutoUpdate: autoUpdate,
CheckInterval: interval,
SourceIdentity: buildGitHubSourceIdentity(repository, selector, tag, assetName),
}, nil
}
func buildGitHubSourceIdentity(repository, selector, tag, assetName string) string {
fields := [...]string{repository, selector, tag, assetName}
encoded := make([]byte, 0, len(githubSourceIdentityDomain)+len(fields)*8+
len(repository)+len(selector)+len(tag)+len(assetName))
encoded = append(encoded, githubSourceIdentityDomain...)
var fieldLength [8]byte
for _, field := range fields {
// Go strings hold the validated UTF-8 bytes used by GitHub. Prefixing each
// field with its byte length prevents delimiter characters from creating
// ambiguous identities across field boundaries.
binary.BigEndian.PutUint64(fieldLength[:], uint64(len(field)))
encoded = append(encoded, fieldLength[:]...)
encoded = append(encoded, field...)
}
identityHash := sha256.Sum256(encoded)
return hex.EncodeToString(identityHash[:])
}
func normalizeGitHubRepositoryURL(raw string) (string, error) {
parsed, err := url.Parse(raw)
if err != nil || parsed.Scheme != "https" || !strings.EqualFold(parsed.Host, "github.com") ||
parsed.User != nil || parsed.RawQuery != "" || parsed.ForceQuery || parsed.Fragment != "" ||
strings.Contains(raw, "#") ||
parsed.EscapedPath() != parsed.Path || !strings.HasPrefix(parsed.Path, "/") ||
strings.HasPrefix(parsed.Path, "//") || strings.HasSuffix(parsed.Path, "/") {
return "", errors.New(errPagesSourceRepositoryInvalid)
}
parts := strings.Split(strings.TrimPrefix(parsed.Path, "/"), "/")
if len(parts) != githubRepositoryPathParts {
return "", errors.New(errPagesSourceRepositoryInvalid)
}
owner := parts[0]
repository := parts[1]
repository = strings.TrimSuffix(repository, ".git")
if !githubOwnerPattern.MatchString(owner) || !githubRepoPattern.MatchString(repository) ||
len(repository) > 100 || repository == "." || repository == ".." {
return "", errors.New(errPagesSourceRepositoryInvalid)
}
return owner + "/" + repository, nil
}
func validGitHubReleaseTagConfig(value string) bool {
if !validGitHubReleaseDisplayTag(value) ||
strings.ContainsAny(value, " ~^:?*[\\") || strings.HasPrefix(value, "/") ||
strings.HasSuffix(value, "/") || strings.HasSuffix(value, ".") ||
strings.Contains(value, "//") || strings.Contains(value, "..") || strings.Contains(value, "@{") {
return false
}
for component := range strings.SplitSeq(value, "/") {
if strings.HasPrefix(component, ".") || strings.HasSuffix(component, ".lock") {
return false
}
}
return true
}
func validGitHubReleaseDisplayTag(value string) bool {
if value == "" || len(value) > 255 || !utf8.ValidString(value) {
return false
}
for _, character := range value {
if unsafeGitHubInputRune(character) {
return false
}
}
return true
}
func validGitHubAssetName(value string) bool {
if value == "" || len(value) > 255 || !utf8.ValidString(value) ||
path.Base(value) != value || strings.Contains(value, "\\") ||
value == "." || value == ".." {
return false
}
for _, character := range value {
if unsafeGitHubInputRune(character) {
return false
}
}
return true
}
func unsafeGitHubInputRune(character rune) bool {
return unicode.IsControl(character) || character == '\u2028' || character == '\u2029' ||
character == '\u061c' || character == '\u200e' || character == '\u200f' ||
(character >= '\u202a' && character <= '\u202e') ||
(character >= '\u2066' && character <= '\u2069')
}
func updateGitHubSourceTx(tx *gorm.DB, projectID uint, input SourceUpdateInput) (bool, error) {
var project model.PagesProject
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&project, projectID).Error; err != nil {
return false, err
}
existing, hasExisting, err := loadProjectSourceForUpdate(tx, projectID)
if err != nil {
return false, err
}
config, err := buildGitHubSourceConfig(input)
if err != nil {
return false, err
}
if !hasExisting {
return true, createGitHubSourceTx(tx, projectID, config)
}
if !githubSourceConfigChanged(existing, config) {
return false, nil
}
var runtime model.PagesProjectSourceRuntime
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).
Where("source_id = ?", existing.ID).First(&runtime).Error; err != nil {
return false, err
}
identityChanged := existing.SourceIdentity != config.SourceIdentity
if err := tx.Model(existing).Updates(githubSourceUpdates(config, existing.ConfigVersion+1)).Error; err != nil {
return false, err
}
if err := resetRuntimeAfterGitHubUpdate(tx, &runtime, config, identityChanged); err != nil {
return false, err
}
return true, nil
}
func createGitHubSourceTx(tx *gorm.DB, projectID uint, config githubSourceConfig) error {
source := &model.PagesProjectSource{
ProjectID: projectID,
SourceType: PagesSourceTypeGitHubRelease,
GitHubRepository: config.Repository,
ReleaseSelector: config.Selector,
ReleaseTag: config.Tag,
AssetName: config.AssetName,
AutoUpdateEnabled: config.AutoUpdate,
CheckIntervalMinutes: config.CheckInterval,
ConfigVersion: 1,
SourceIdentity: config.SourceIdentity,
}
if err := tx.Create(source).Error; err != nil {
return err
}
runtime := &model.PagesProjectSourceRuntime{SourceID: source.ID, SyncStatus: pagesSourceStatusIdle}
if config.Selector == githubReleaseSelectorLatest {
next := nextGitHubCheckAt(time.Now(), source.ID, config.CheckInterval)
runtime.NextCheckAt = &next
}
return tx.Create(runtime).Error
}
func githubSourceUpdates(config githubSourceConfig, version int) map[string]any {
return map[string]any{
"source_type": PagesSourceTypeGitHubRelease,
"remote_url": "",
"remote_network_policy": "",
"github_repository": config.Repository,
"release_selector": config.Selector,
"release_tag": config.Tag,
"asset_name": config.AssetName,
sourceColumnAutoUpdateEnabled: config.AutoUpdate,
"check_interval_minutes": config.CheckInterval,
sourceColumnConfigVersion: version,
"source_identity": config.SourceIdentity,
}
}
func githubSourceConfigChanged(existing *model.PagesProjectSource, config githubSourceConfig) bool {
return existing.SourceType != PagesSourceTypeGitHubRelease || existing.RemoteURL != "" ||
existing.RemoteNetworkPolicy != "" || existing.GitHubRepository != config.Repository ||
existing.ReleaseSelector != config.Selector || existing.ReleaseTag != config.Tag ||
existing.AssetName != config.AssetName || existing.AutoUpdateEnabled != config.AutoUpdate ||
existing.CheckIntervalMinutes != config.CheckInterval
}
func resetRuntimeAfterGitHubUpdate(
tx *gorm.DB,
runtime *model.PagesProjectSourceRuntime,
config githubSourceConfig,
identityChanged bool,
) error {
if err := resetRuntimeAfterSourceUpdate(tx, runtime, identityChanged); err != nil {
return err
}
var nextCheckAt any
if config.Selector == githubReleaseSelectorLatest {
next := nextGitHubCheckAt(time.Now(), runtime.SourceID, config.CheckInterval)
nextCheckAt = &next
}
return tx.Model(runtime).Update("next_check_at", nextCheckAt).Error
}
func nextGitHubCheckAt(now time.Time, sourceID uint, intervalMinutes int) time.Time {
// A stable, bounded offset avoids a thundering herd without persisting
// another scheduling field. Scanner Phase 3 reuses this calculation.
jitterSeconds := int64(sourceID%githubCheckJitterRange) - githubCheckJitterCenter
return now.Add(time.Duration(intervalMinutes)*time.Minute + time.Duration(jitterSeconds)*time.Second)
}
func markInitialCheckDispatchFailed(ctx context.Context, sourceID uint, configVersion int) {
updates := map[string]any{
sourceRuntimeColumnSyncStatus: pagesSourceStatusFailed,
sourceRuntimeColumnLastError: errPagesSourceInitialCheckWarning,
}
var source model.PagesProjectSource
if err := db.DB(ctx).Where("id = ? AND config_version = ?", sourceID, configVersion).First(&source).Error; err != nil {
if !errors.Is(err, gorm.ErrRecordNotFound) {
logger.ErrorF(ctx, "[PagesSource] load initial check source snapshot failed: source_id=%d error=%v", sourceID, err)
}
return
}
if source.ReleaseSelector == githubReleaseSelectorLatest {
next := time.Now().Add(initialCheckRetryDelay)
updates["next_check_at"] = &next
}
now := time.Now()
result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).
Where("source_id = ?", sourceID).
Where("lease_expires_at IS NULL OR lease_expires_at <= ?", now).
Where("EXISTS (SELECT 1 FROM of_pages_project_sources source WHERE source.id = ? AND source.config_version = ?)", sourceID, configVersion).
Updates(updates)
if result.Error != nil {
logger.ErrorF(ctx, "[PagesSource] mark initial check dispatch failure: source_id=%d error=%v", sourceID, result.Error)
}
}
@@ -0,0 +1,754 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pages
import (
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"regexp"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/integration/githubrelease"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/task"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/Rain-kl/Wavelet/pkg/pagesarchive"
"gorm.io/gorm"
"gorm.io/gorm/clause"
)
const githubSourceDetailProvider = "github"
var githubDigestPattern = regexp.MustCompile(`^sha256:[0-9a-f]{64}$`)
type githubSourceProviderDomainError struct {
message string
permanent bool
retryAt *time.Time
statusCode int
}
func (domainError *githubSourceProviderDomainError) Error() string {
return domainError.message
}
type githubReleaseAPI interface {
Resolve(context.Context, githubrelease.ResolveRequest) (githubrelease.ResolveResult, error)
Download(context.Context, githubrelease.DownloadRequest) (*githubrelease.DownloadResult, error)
}
var newGitHubReleaseClient = func() githubReleaseAPI {
return githubrelease.NewClient()
}
type githubSourceTarget struct {
Revision string
Detail sourceDetail
DetailJSON string
Release githubrelease.Release
Asset githubrelease.Asset
RetryAt *time.Time
}
type githubCheckTaskResult struct {
Message string
Detail string
Revision string
Status string
RetryAt *time.Time
Stale bool
}
type preparedGitHubSource struct {
target *githubSourceTarget
download *githubrelease.DownloadResult
format pagesarchive.Format
manifest *deploymentManifest
ingestState *sourceIngestState
limits pagesLimits
}
func checkGitHubSource(
ctx context.Context,
snapshot *sourceExecutionSnapshot,
) (*githubCheckTaskResult, error) {
if snapshot == nil || snapshot.SourceType != PagesSourceTypeGitHubRelease {
return nil, errors.New(errPagesSourceTypeUnsupported)
}
task.AppendLog(ctx, "[check] 正在检查 GitHub Release:repo=%s asset=%s", snapshot.GitHubRepository, snapshot.AssetName)
client := newGitHubReleaseClient()
result, err := client.Resolve(ctx, githubrelease.ResolveRequest{
Repository: snapshot.GitHubRepository,
Selector: githubrelease.Selector(snapshot.ReleaseSelector),
Tag: snapshot.ReleaseTag,
AssetName: snapshot.AssetName,
ETag: snapshot.ETag,
})
if err != nil {
logger.WarnF(ctx, "[PagesSource] GitHub resolve failed: source_id=%d repo=%s error=%v", snapshot.SourceID, snapshot.GitHubRepository, err)
retryAt, _ := githubrelease.RetryAt(err)
domainErr := githubSourceDomainError(err)
if failErr := failGitHubCheckLease(ctx, snapshot, domainErr.Error(), retryAt); failErr != nil {
if errors.Is(failErr, errSourceFinalFence) {
return &githubCheckTaskResult{Message: errPagesSourceActionStale, Stale: true}, nil
}
return nil, failErr
}
return nil, domainErr
}
if result.NotModified {
revision, status, err := finishGitHubCheckNotModified(ctx, snapshot, result)
if err != nil {
if errors.Is(err, errSourceFinalFence) {
return &githubCheckTaskResult{Message: errPagesSourceActionStale, Stale: true}, nil
}
return nil, err
}
detail, _ := json.Marshal(map[string]string{"revision": revision, pagesDeploymentColumnStatus: status})
return &githubCheckTaskResult{
Message: "GitHub Release 检查完成,内容未变化",
Detail: string(detail),
Revision: revision,
Status: status,
RetryAt: result.RetryAt,
}, nil
}
target, err := buildGitHubSourceTarget(result.Release, result.Asset, result.RetryAt)
if err != nil {
retryAt := time.Time{}
if result.RetryAt != nil {
retryAt = result.RetryAt.UTC()
}
if failErr := failGitHubCheckLease(ctx, snapshot, err.Error(), retryAt); failErr != nil {
if errors.Is(failErr, errSourceFinalFence) {
return &githubCheckTaskResult{Message: errPagesSourceActionStale, Stale: true}, nil
}
return nil, failErr
}
return nil, err
}
status, err := finishGitHubCheckTarget(ctx, snapshot, result, target)
if err != nil {
if errors.Is(err, errSourceFinalFence) {
return &githubCheckTaskResult{Message: errPagesSourceActionStale, Stale: true}, nil
}
return nil, err
}
detail, _ := json.Marshal(map[string]string{"revision": target.Revision, pagesDeploymentColumnStatus: status})
message := "GitHub Release 检查完成"
switch status {
case pagesSourceStatusUpdateAvailable:
message = "发现新的 GitHub Release 部署包"
case pagesSourceStatusAttention:
message = "检测到同一 Release 的资源被替换,需要确认"
}
return &githubCheckTaskResult{
Message: message, Detail: string(detail), Revision: target.Revision,
Status: status, RetryAt: result.RetryAt,
}, nil
}
func buildGitHubSourceTarget(
release githubrelease.Release,
asset githubrelease.Asset,
retryAt *time.Time,
) (*githubSourceTarget, error) {
digest := strings.ToLower(strings.TrimSpace(asset.Digest))
if digest != "" && !githubDigestPattern.MatchString(digest) {
return nil, errors.New(errPagesSourceDigestInvalid)
}
if strings.TrimSpace(release.ID) == "" || strings.TrimSpace(asset.ID) == "" ||
!validGitHubReleaseDisplayTag(release.Tag) || !validGitHubAssetName(asset.Name) ||
asset.State != "uploaded" || asset.UpdatedAt.IsZero() {
return nil, errors.New(errPagesSourceReleaseNotFound)
}
updatedAt := asset.UpdatedAt.UTC().Format(time.RFC3339Nano)
rawRevision := "github:" + release.ID + ":" + asset.ID + ":" + updatedAt + ":" + digest
sum := sha256.Sum256([]byte(rawRevision))
detail := sourceDetail{
Provider: githubSourceDetailProvider,
Tag: release.Tag,
AssetName: asset.Name,
ReleaseID: release.ID,
AssetID: asset.ID,
AssetUpdatedAt: updatedAt,
Digest: digest,
}
detailJSON, err := json.Marshal(detail)
if err != nil {
return nil, errors.New(errPagesSourceSyncFailed)
}
return &githubSourceTarget{
Revision: hex.EncodeToString(sum[:]),
Detail: detail,
DetailJSON: string(detailJSON),
Release: release,
Asset: asset,
RetryAt: retryAt,
}, nil
}
func finishGitHubCheckNotModified(
ctx context.Context,
snapshot *sourceExecutionSnapshot,
result githubrelease.ResolveResult,
) (string, string, error) {
var revision string
var status string
err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
runtime, now, err := lockOwnedSourceRuntime(tx, snapshot)
if err != nil {
return err
}
revision = runtime.LastSeenRevision
status = normalizedSourceRuntimeStatus(runtime)
updates := githubCheckTerminalUpdates(snapshot, now, result.RetryAt)
updates["etag"] = result.ETag
updates[sourceRuntimeColumnSyncStatus] = status
return tx.Model(runtime).Updates(updates).Error
})
return revision, status, err
}
func finishGitHubCheckTarget(
ctx context.Context,
snapshot *sourceExecutionSnapshot,
result githubrelease.ResolveResult,
target *githubSourceTarget,
) (string, error) {
status := pagesSourceStatusIdle
err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
runtime, now, err := lockOwnedSourceRuntime(tx, snapshot)
if err != nil {
return err
}
status = targetRuntimeStatus(target, runtime.LastAppliedRevision, runtime.LastAppliedDetail)
updates := githubCheckTerminalUpdates(snapshot, now, result.RetryAt)
updates["etag"] = result.ETag
updates["last_seen_revision"] = target.Revision
updates["last_seen_detail"] = target.DetailJSON
updates[sourceRuntimeColumnSyncStatus] = status
return tx.Model(runtime).Updates(updates).Error
})
return status, err
}
func githubCheckTerminalUpdates(
snapshot *sourceExecutionSnapshot,
now time.Time,
retryAt *time.Time,
) map[string]any {
updates := map[string]any{
sourceRuntimeColumnLastError: "",
sourceRuntimeColumnLastCheckedAt: &now,
sourceRuntimeColumnLeaseToken: "",
sourceRuntimeColumnLeaseExpiresAt: nil,
}
updates[sourceRuntimeColumnNextCheckAt] = nextCheckAfterGitHubResponse(snapshot, now, retryAt)
return updates
}
func nextCheckAfterGitHubResponse(
snapshot *sourceExecutionSnapshot,
now time.Time,
retryAt *time.Time,
) any {
if snapshot.ReleaseSelector != githubReleaseSelectorLatest {
return nil
}
next := nextGitHubCheckAt(now, snapshot.SourceID, snapshot.CheckIntervalMinutes)
if retryAt != nil && retryAt.After(next) {
next = retryAt.In(now.Location())
}
return &next
}
func lockOwnedSourceRuntime(
tx *gorm.DB,
snapshot *sourceExecutionSnapshot,
) (*model.PagesProjectSourceRuntime, time.Time, error) {
var runtime model.PagesProjectSourceRuntime
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).
Where("source_id = ?", snapshot.SourceID).First(&runtime).Error; err != nil {
return nil, time.Time{}, err
}
now := time.Now()
if runtime.LeaseToken != snapshot.LeaseToken || runtime.LeaseExpiresAt == nil ||
!runtime.LeaseExpiresAt.After(now) {
return nil, time.Time{}, errSourceFinalFence
}
return &runtime, now, nil
}
func failGitHubCheckLease(
ctx context.Context,
snapshot *sourceExecutionSnapshot,
message string,
retryAt time.Time,
) error {
now := time.Now()
next := now.Add(initialCheckRetryDelay)
if retryAt.After(next) {
next = retryAt.In(now.Location())
}
updates := map[string]any{
sourceRuntimeColumnSyncStatus: pagesSourceStatusFailed,
sourceRuntimeColumnLastError: safeSourceRuntimeError(message),
sourceRuntimeColumnLastCheckedAt: &now,
sourceRuntimeColumnLeaseToken: "",
sourceRuntimeColumnLeaseExpiresAt: nil,
}
if snapshot.ReleaseSelector == githubReleaseSelectorLatest {
updates[sourceRuntimeColumnNextCheckAt] = &next
} else {
updates[sourceRuntimeColumnNextCheckAt] = nil
}
result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).
Where("source_id = ? AND lease_token = ? AND lease_expires_at > ?", snapshot.SourceID, snapshot.LeaseToken, now).
Updates(updates)
if result.Error != nil {
return result.Error
}
if result.RowsAffected != 1 {
return errSourceFinalFence
}
return nil
}
func targetRuntimeStatus(
target *githubSourceTarget,
appliedRevision string,
appliedDetail string,
) string {
if target == nil || target.Revision == appliedRevision {
return pagesSourceStatusIdle
}
applied := sourceDetail{}
if unmarshalSourceDetail(appliedDetail, &applied) == nil && target.Detail.ReleaseID != "" &&
target.Detail.ReleaseID == applied.ReleaseID {
return pagesSourceStatusAttention
}
return pagesSourceStatusUpdateAvailable
}
func preflightGitHubSyncConfirmation(ctx context.Context, sourceID uint, confirmedRevision string) error {
var runtime model.PagesProjectSourceRuntime
if err := db.DB(ctx).Where("source_id = ?", sourceID).First(&runtime).Error; err != nil {
return err
}
replacement := sourceHasSameReleaseReplacement(&runtime)
if replacement && confirmedRevision == "" {
return errors.New(errPagesSourceConfirmationNeeded)
}
if confirmedRevision != "" && (!replacement || confirmedRevision != runtime.LastSeenRevision) {
return errors.New(errPagesSourceConfirmationStale)
}
return nil
}
func syncGitHubSourceWithTrigger(
ctx context.Context,
snapshot *sourceExecutionSnapshot,
actor string,
targetRevision string,
confirmedRevision string,
triggerType string,
) (outcome *sourceSyncOutcome, resultErr error) {
if snapshot == nil || snapshot.SourceType != PagesSourceTypeGitHubRelease || !validPagesSourceActor(actor) {
return nil, errors.New(errPagesSourceActionInvalid)
}
if !validSourceDeploymentTrigger(triggerType) {
return nil, errors.New(errPagesSourceActionInvalid)
}
defer func() {
resultErr = finalizeGitHubSyncFailure(ctx, snapshot, resultErr)
}()
workCtx, heartbeat, err := startSourceLeaseHeartbeat(
ctx, snapshot, pagesSourceSyncLeaseDuration, pagesSourceHeartbeatInterval,
)
if err != nil {
return sourceHeartbeatOutcome(err)
}
defer func() { _ = heartbeat.stop() }()
client := newGitHubReleaseClient()
target, guardedOutcome, err := resolveAndGuardGitHubSync(
workCtx, client, snapshot, targetRevision, confirmedRevision,
)
if err != nil {
return nil, err
}
if guardedOutcome != nil {
return guardedOutcome, nil
}
prepared, err := prepareGitHubSyncPackage(workCtx, client, snapshot, target)
if err != nil {
return nil, err
}
defer func() {
if cleanupErr := prepared.download.Cleanup(); cleanupErr != nil {
logger.WarnF(ctx, "[PagesSource] cleanup GitHub package failed: source_id=%d error=%v", snapshot.SourceID, cleanupErr)
}
}()
defer compensateSourceIngest(ctx, snapshot, prepared.ingestState)
if heartbeatErr := heartbeat.stop(); heartbeatErr != nil {
return sourceHeartbeatOutcome(heartbeatErr)
}
renewed, err := renewSourceLease(ctx, snapshot, pagesSourceSyncLeaseDuration)
if err != nil {
return nil, err
}
if !renewed {
return &sourceSyncOutcome{Stale: true}, nil
}
return activatePreparedGitHubSource(ctx, snapshot, actor, triggerType, prepared)
}
func finalizeGitHubSyncFailure(
ctx context.Context,
snapshot *sourceExecutionSnapshot,
resultErr error,
) error {
if resultErr == nil {
return nil
}
cleanupCtx, cancel := sourceCleanupContext(ctx)
defer cancel()
finalizerErr := persistGitHubSyncFailure(cleanupCtx, snapshot, resultErr)
if finalizerErr == nil {
return resultErr
}
logger.WarnF(
cleanupCtx,
"[PagesSource] finalize GitHub sync failure failed: source_id=%d source_error=%s error=%v",
snapshot.SourceID, safeGitHubSourceError(resultErr), finalizerErr,
)
// final fence 丢失表示已有新任务接管 runtime,不应覆盖;数据库
// finalizer 失败则保持可重试,避免继承永久错误或 provider deadline 分类。
if errors.Is(finalizerErr, errSourceFinalFence) {
return resultErr
}
return errors.New(errPagesSourceSyncFailed)
}
func persistGitHubSyncFailure(
ctx context.Context,
snapshot *sourceExecutionSnapshot,
resultErr error,
) error {
var domainError *githubSourceProviderDomainError
if errors.As(resultErr, &domainError) && domainError.retryAt != nil {
return failGitHubCheckLease(ctx, snapshot, domainError.message, *domainError.retryAt)
}
return failSourceLease(ctx, snapshot, safeGitHubSourceError(resultErr))
}
func activatePreparedGitHubSource(
ctx context.Context,
snapshot *sourceExecutionSnapshot,
actor string,
triggerType string,
prepared *preparedGitHubSource,
) (*sourceSyncOutcome, error) {
task.AppendLog(ctx, "[activate] 正在原子切换 GitHub Release 部署")
deployment, reused, referenced, err := commitSourceDeploymentWithTrigger(
ctx, snapshot, prepared.target.Revision, prepared.download.SHA256,
prepared.target.Detail, prepared.target.DetailJSON, actor, triggerType, prepared.manifest,
prepared.ingestState.Result, prepared.ingestState.HasIngest, prepared.target.RetryAt,
)
prepared.ingestState.Referenced = referenced
if errors.Is(err, errSourceFinalFence) {
return &sourceSyncOutcome{Stale: true}, nil
}
if err != nil {
return nil, err
}
prepared.ingestState.Referenced = prepared.ingestState.HasIngest && deployment.UploadID == prepared.ingestState.Result.Upload.ID
if pruneErr := pruneProjectDeploymentHistory(ctx, snapshot.ProjectID, prepared.limits.HistoryCount, 0); pruneErr != nil {
logger.ErrorF(ctx, "[PagesSource] strict prune failed after GitHub sync: project_id=%d source_id=%d error=%v", snapshot.ProjectID, snapshot.SourceID, pruneErr)
}
view := buildDeploymentView(deployment)
return &sourceSyncOutcome{Deployment: &view, Reused: reused}, nil
}
func resolveAndGuardGitHubSync(
ctx context.Context,
client githubReleaseAPI,
snapshot *sourceExecutionSnapshot,
targetRevision string,
confirmedRevision string,
) (*githubSourceTarget, *sourceSyncOutcome, error) {
task.AppendLog(ctx, "[resolve] 正在解析 GitHub Release:repo=%s asset=%s", snapshot.GitHubRepository, snapshot.AssetName)
resolved, err := client.Resolve(ctx, githubrelease.ResolveRequest{
Repository: snapshot.GitHubRepository,
Selector: githubrelease.Selector(snapshot.ReleaseSelector),
Tag: snapshot.ReleaseTag,
AssetName: snapshot.AssetName,
})
if err != nil {
logger.WarnF(ctx, "[PagesSource] GitHub resolve failed: source_id=%d repo=%s error=%v", snapshot.SourceID, snapshot.GitHubRepository, err)
return nil, nil, githubSourceDomainError(err)
}
if resolved.NotModified {
return nil, nil, errors.New(errPagesSourceReleaseNotFound)
}
target, err := buildGitHubSourceTarget(resolved.Release, resolved.Asset, resolved.RetryAt)
if err != nil {
return nil, nil, &githubSourceProviderDomainError{
message: safeGitHubSourceError(err),
permanent: isPermanentSourceSyncError(err),
retryAt: resolved.RetryAt,
}
}
guardedOutcome, err := guardGitHubSyncTarget(ctx, snapshot, target, targetRevision, confirmedRevision)
return target, guardedOutcome, err
}
func guardGitHubSyncTarget(
ctx context.Context,
snapshot *sourceExecutionSnapshot,
target *githubSourceTarget,
targetRevision string,
confirmedRevision string,
) (*sourceSyncOutcome, error) {
status := targetRuntimeStatus(target, snapshot.LastAppliedRevision, snapshot.LastAppliedDetail)
if targetRevision != "" && targetRevision != target.Revision {
return releaseGuardedGitHubTarget(ctx, snapshot, target, status, "", true, true)
}
if confirmedRevision != "" && (confirmedRevision != snapshot.LastSeenRevision || confirmedRevision != target.Revision) {
return releaseGuardedGitHubTarget(ctx, snapshot, target, status, errPagesSourceConfirmationStale, false, false)
}
if status == pagesSourceStatusAttention && confirmedRevision != target.Revision {
return releaseGuardedGitHubTarget(ctx, snapshot, target, status, errPagesSourceConfirmationNeeded, false, false)
}
if confirmedRevision != "" && status != pagesSourceStatusAttention {
return nil, errors.New(errPagesSourceConfirmationStale)
}
return nil, nil
}
func releaseGuardedGitHubTarget(
ctx context.Context,
snapshot *sourceExecutionSnapshot,
target *githubSourceTarget,
status string,
lastError string,
expedite bool,
staleSuccess bool,
) (*sourceSyncOutcome, error) {
err := releaseGitHubSyncWithoutActivation(ctx, snapshot, target, status, lastError, expedite, target.RetryAt)
if errors.Is(err, errSourceFinalFence) || (err == nil && staleSuccess) {
return &sourceSyncOutcome{Stale: true}, nil
}
if err != nil {
return nil, err
}
return nil, errors.New(lastError)
}
func prepareGitHubSyncPackage(
ctx context.Context,
client githubReleaseAPI,
snapshot *sourceExecutionSnapshot,
target *githubSourceTarget,
) (*preparedGitHubSource, error) {
limits := resolvePagesLimits(ctx)
task.AppendLog(ctx, "[download] 正在下载 GitHub Release asset:repo=%s asset=%s", snapshot.GitHubRepository, snapshot.AssetName)
download, err := client.Download(ctx, githubrelease.DownloadRequest{
Repository: snapshot.GitHubRepository,
Asset: target.Asset,
MaxBytes: limits.PackageBytes,
})
if err != nil {
logger.WarnF(ctx, "[PagesSource] GitHub download failed: source_id=%d repo=%s asset=%s error=%v", snapshot.SourceID, snapshot.GitHubRepository, snapshot.AssetName, err)
return nil, githubSourceDomainError(err)
}
prepared, err := inspectAndIngestGitHubPackage(ctx, snapshot, target, download, limits)
if err != nil {
if cleanupErr := download.Cleanup(); cleanupErr != nil {
logger.WarnF(ctx, "[PagesSource] cleanup GitHub package after preparation failure failed: source_id=%d error=%v", snapshot.SourceID, cleanupErr)
}
return nil, err
}
return prepared, nil
}
func inspectAndIngestGitHubPackage(
ctx context.Context,
snapshot *sourceExecutionSnapshot,
target *githubSourceTarget,
download *githubrelease.DownloadResult,
limits pagesLimits,
) (*preparedGitHubSource, error) {
if download.SHA256 == "" || download.Path == "" {
return nil, errors.New(errPagesSourceSyncFailed)
}
if target.Detail.Digest != "" && "sha256:"+download.SHA256 != target.Detail.Digest {
return nil, errors.New(errPagesSourceDigestMismatch)
}
format, ok := pagesarchive.DetectFormatFromName(target.Asset.Name)
var err error
if !ok {
format, _, err = detectRemoteSourceFormat(download.Path, target.Asset.Name, "")
if err != nil {
return nil, err
}
}
rootDir, err := validateAndNormalizePagesRootDir(snapshot.RootDir)
if err != nil {
return nil, err
}
entryFile, err := validateAndNormalizePagesEntryFile(snapshot.EntryFile)
if err != nil {
return nil, err
}
task.AppendLog(ctx, "[verify] 正在校验 GitHub Release 归档与入口")
manifest, err := inspectPagesPackage(download.Path, format, rootDir, entryFile, limits)
if err != nil {
return nil, err
}
ingestState, err := resolveGitHubSourceIngest(ctx, snapshot, target, download, format)
if err != nil {
return nil, err
}
return &preparedGitHubSource{
target: target, download: download, format: format,
manifest: manifest, ingestState: ingestState, limits: limits,
}, nil
}
func resolveGitHubSourceIngest(
ctx context.Context,
snapshot *sourceExecutionSnapshot,
target *githubSourceTarget,
download *githubrelease.DownloadResult,
format pagesarchive.Format,
) (*sourceIngestState, error) {
if _, err := findSourceDeployment(ctx, snapshot.ProjectID, snapshot.SourceIdentity, target.Revision); err == nil {
return &sourceIngestState{}, nil
} else if !errors.Is(err, gorm.ErrRecordNotFound) {
return nil, err
}
task.AppendLog(ctx, "[ingest] 正在保存 GitHub Release 部署包")
result, err := ingestPagesDeploymentPackageWithSource(
ctx, download.Path, download.SHA256, snapshot.ProjectID, snapshot.SourceID, target.Asset.Name, format,
)
if err != nil {
return nil, err
}
return &sourceIngestState{Result: result, HasIngest: true}, nil
}
func releaseGitHubSyncWithoutActivation(
ctx context.Context,
snapshot *sourceExecutionSnapshot,
target *githubSourceTarget,
status string,
lastError string,
expedite bool,
retryAt *time.Time,
) error {
now := time.Now()
nextCheckAt := nextCheckAfterGitHubResponse(snapshot, now, retryAt)
if expedite && snapshot.ReleaseSelector == githubReleaseSelectorLatest {
next := now.Add(initialCheckRetryDelay)
if retryAt != nil && retryAt.After(next) {
next = retryAt.In(now.Location())
}
nextCheckAt = &next
}
updates := map[string]any{
"last_seen_revision": target.Revision,
"last_seen_detail": target.DetailJSON,
sourceRuntimeColumnSyncStatus: status,
sourceRuntimeColumnLastError: lastError,
sourceRuntimeColumnLastCheckedAt: &now,
sourceRuntimeColumnNextCheckAt: nextCheckAt,
sourceRuntimeColumnLeaseToken: "",
sourceRuntimeColumnLeaseExpiresAt: nil,
}
result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).
Where("source_id = ? AND lease_token = ? AND lease_expires_at > ?", snapshot.SourceID, snapshot.LeaseToken, now).
Updates(updates)
if result.Error != nil {
return result.Error
}
if result.RowsAffected != 1 {
return errSourceFinalFence
}
return nil
}
func safeGitHubSourceError(err error) string {
if err == nil {
return errPagesSourceSyncFailed
}
message := strings.TrimSpace(err.Error())
for _, safeMessage := range []string{
errPagesSourceSyncFailed,
errPagesSourceReleaseNotFound,
errPagesSourceDigestInvalid,
errPagesSourceDigestMismatch,
errPagesSourceConfirmationNeeded,
errPagesSourceConfirmationStale,
errPagesPackageURLTooLarge,
errPagesPackageEmpty,
errPagesPackageUnsupported,
errPagesPackageInvalid,
errPagesPackageExtractedTooLarge,
errPagesPackageFileTooLarge,
errPagesEntryFileMissing,
} {
if message == safeMessage {
return safeMessage
}
}
return errPagesSourceSyncFailed
}
func githubSourceDomainError(err error) error {
message := errPagesSourceSyncFailed
statusCode := 0
var providerError *githubrelease.Error
if errors.As(err, &providerError) {
statusCode = providerError.StatusCode
}
retryAt, hasRetryAt := githubrelease.RetryAt(err)
var retryDeadline *time.Time
if hasRetryAt {
retryDeadline = &retryAt
}
if err == nil {
return &githubSourceProviderDomainError{message: message, permanent: false, statusCode: statusCode}
}
if githubrelease.IsDigestError(err) {
message = errPagesSourceDigestMismatch
return &githubSourceProviderDomainError{message: message, permanent: true, statusCode: statusCode}
}
if githubrelease.IsNotFound(err) {
message = errPagesSourceReleaseNotFound
return &githubSourceProviderDomainError{message: message, permanent: true, statusCode: statusCode}
}
if errors.Is(err, githubrelease.ErrAssetTooLarge) {
message = errPagesPackageURLTooLarge
return &githubSourceProviderDomainError{message: message, permanent: true, statusCode: statusCode}
}
if errors.Is(err, githubrelease.ErrEmptyAsset) {
message = errPagesPackageEmpty
return &githubSourceProviderDomainError{message: message, permanent: true, statusCode: statusCode}
}
return &githubSourceProviderDomainError{
message: message, permanent: !githubrelease.IsRetryable(err), retryAt: retryDeadline, statusCode: statusCode,
}
}
func shouldSkipGitHubActionRetry(err error) bool {
var domainError *githubSourceProviderDomainError
return errors.As(err, &domainError) && (domainError.permanent || domainError.retryAt != nil)
}
@@ -0,0 +1,111 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pages
import (
"strings"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
)
func TestGitHubSourceIdentityLengthPrefixesFieldsAndResetsRuntime(t *testing.T) {
firstInput := SourceUpdateInput{
SourceType: PagesSourceTypeGitHubRelease,
RepositoryURL: "https://github.com/OpenFlare/site",
ReleaseSelector: githubReleaseSelectorTag,
ReleaseTag: "release|foo",
AssetName: "bar.zip",
}
secondInput := SourceUpdateInput{
SourceType: PagesSourceTypeGitHubRelease,
RepositoryURL: "https://github.com/OpenFlare/site",
ReleaseSelector: githubReleaseSelectorTag,
ReleaseTag: "release",
AssetName: "foo|bar.zip",
}
firstConfig, err := buildGitHubSourceConfig(firstInput)
if err != nil {
t.Fatalf("buildGitHubSourceConfig(first) error = %v, want nil", err)
}
secondConfig, err := buildGitHubSourceConfig(secondInput)
if err != nil {
t.Fatalf("buildGitHubSourceConfig(second) error = %v, want nil", err)
}
legacyIdentityInput := func(config githubSourceConfig) string {
return "github|" + config.Repository + "|" + config.Selector + "|" +
config.Tag + "|" + config.AssetName
}
if firstLegacy, secondLegacy := legacyIdentityInput(firstConfig), legacyIdentityInput(secondConfig); firstLegacy != secondLegacy {
t.Fatalf("legacy identity inputs differ: %q != %q; collision fixture is invalid", firstLegacy, secondLegacy)
}
if firstConfig.SourceIdentity == secondConfig.SourceIdentity {
t.Fatalf("length-prefixed identities collide: %q", firstConfig.SourceIdentity)
}
ctx := setupPagesSourceTest(t)
project := mustCreatePagesSourceProject(t, ctx, "github-identity-collision")
firstSource, _ := mustConfigureGitHubSourceWithoutDispatch(t, ctx, project.ID, firstInput)
if got, want := firstSource.SourceIdentity, firstConfig.SourceIdentity; got != want {
t.Fatalf("first source identity = %q, want %q", got, want)
}
checkedAt := time.Now().Add(-time.Minute)
syncedAt := time.Now().Add(-30 * time.Second)
nextCheckAt := time.Now().Add(time.Hour)
leaseExpiresAt := time.Now().Add(time.Minute)
if err := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).
Where("source_id = ?", firstSource.ID).
Updates(map[string]any{
"etag": `"old-etag"`,
"last_seen_revision": strings.Repeat("a", 64),
"last_seen_detail": `{"provider":"github_release","tag":"release|foo"}`,
"last_applied_revision": strings.Repeat("b", 64),
"last_applied_detail": `{"provider":"github_release","tag":"older"}`,
"sync_status": pagesSourceStatusSyncing,
"last_error": "old error",
"last_checked_at": &checkedAt,
"last_synced_at": &syncedAt,
"next_check_at": &nextCheckAt,
"lease_expires_at": &leaseExpiresAt,
"lease_token": "old-lease",
}).Error; err != nil {
t.Fatalf("seed runtime cursors error = %v, want nil", err)
}
secondSource, runtime := mustConfigureGitHubSourceWithoutDispatch(t, ctx, project.ID, secondInput)
if secondSource.ID != firstSource.ID {
t.Errorf("updated source ID = %d, want unchanged %d", secondSource.ID, firstSource.ID)
}
if got, want := secondSource.SourceIdentity, secondConfig.SourceIdentity; got != want {
t.Errorf("updated source identity = %q, want %q", got, want)
}
if got, want := secondSource.ConfigVersion, firstSource.ConfigVersion+1; got != want {
t.Errorf("updated source config version = %d, want %d", got, want)
}
if runtime.ETag != "" || runtime.LastSeenRevision != "" || runtime.LastSeenDetail != "" ||
runtime.LastAppliedRevision != "" || runtime.LastAppliedDetail != "" {
t.Errorf("identity change retained runtime cursors: %+v", runtime)
}
if runtime.LastCheckedAt != nil || runtime.LastSyncedAt != nil || runtime.NextCheckAt != nil {
t.Errorf(
"identity change retained runtime timestamps: checked=%v synced=%v next=%v",
runtime.LastCheckedAt,
runtime.LastSyncedAt,
runtime.NextCheckAt,
)
}
if runtime.SyncStatus != pagesSourceStatusIdle || runtime.LastError != "" ||
runtime.LeaseToken != "" || runtime.LeaseExpiresAt != nil {
t.Errorf(
"identity change retained runtime state: status=%q error=%q lease=(%q, %v)",
runtime.SyncStatus,
runtime.LastError,
runtime.LeaseToken,
runtime.LeaseExpiresAt,
)
}
}
@@ -0,0 +1,889 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pages
import (
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"net/http"
"os"
"path/filepath"
"strings"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/integration/githubrelease"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/hibiken/asynq"
"gorm.io/gorm"
)
type fakeGitHubReleaseClient struct {
resolve func(context.Context, githubrelease.ResolveRequest) (githubrelease.ResolveResult, error)
download func(context.Context, githubrelease.DownloadRequest) (*githubrelease.DownloadResult, error)
}
func (client *fakeGitHubReleaseClient) Resolve(
ctx context.Context,
request githubrelease.ResolveRequest,
) (githubrelease.ResolveResult, error) {
return client.resolve(ctx, request)
}
func (client *fakeGitHubReleaseClient) Download(
ctx context.Context,
request githubrelease.DownloadRequest,
) (*githubrelease.DownloadResult, error) {
return client.download(ctx, request)
}
func useFakeGitHubReleaseClient(t *testing.T, client githubReleaseAPI) {
t.Helper()
previous := newGitHubReleaseClient
newGitHubReleaseClient = func() githubReleaseAPI { return client }
t.Cleanup(func() { newGitHubReleaseClient = previous })
}
func mustConfigureGitHubSourceWithoutDispatch(
t *testing.T,
ctx context.Context,
projectID uint,
input SourceUpdateInput,
) (*model.PagesProjectSource, *model.PagesProjectSourceRuntime) {
t.Helper()
if err := validateGitHubSourceInput(input); err != nil {
t.Fatalf("validateGitHubSourceInput(%+v) error = %v, want nil", input, err)
}
if err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
_, err := updateGitHubSourceTx(tx, projectID, input)
return err
}); err != nil {
t.Fatalf("updateGitHubSourceTx(project=%d) error = %v, want nil", projectID, err)
}
return mustLoadPagesSource(t, ctx, projectID)
}
func mustLoadPagesSource(
t *testing.T,
ctx context.Context,
projectID uint,
) (*model.PagesProjectSource, *model.PagesProjectSourceRuntime) {
t.Helper()
var source model.PagesProjectSource
if err := db.DB(ctx).Where("project_id = ?", projectID).First(&source).Error; err != nil {
t.Fatalf("load source for project %d error = %v, want nil", projectID, err)
}
var runtime model.PagesProjectSourceRuntime
if err := db.DB(ctx).Where("source_id = ?", source.ID).First(&runtime).Error; err != nil {
t.Fatalf("load runtime for source %d error = %v, want nil", source.ID, err)
}
return &source, &runtime
}
func TestGitHubSourceValidationNormalizationAndProviderSwitch(t *testing.T) {
ctx := setupPagesSourceTest(t)
setupPagesSourceDispatchTest(t)
project := mustCreatePagesSourceProject(t, ctx, "github-config")
input := SourceUpdateInput{
SourceType: PagesSourceTypeGitHubRelease,
RepositoryURL: "https://github.com/OpenFlare/site.git",
}
result, err := UpdateSourceAs(ctx, project.ID, input, "user:42")
if err != nil {
t.Fatalf("UpdateSourceAs(GitHub) error = %v, want nil", err)
}
if result.CheckTask == nil || result.CheckTask.Action != sourceActionCheck || result.Warning != "" {
t.Errorf("UpdateSourceAs(GitHub) result = %+v, want initial check receipt without warning", result)
}
execution, err := model.GetTaskExecutionByTaskID(ctx, result.CheckTask.TaskID)
if err != nil {
t.Fatalf("GetTaskExecutionByTaskID(%q) error = %v, want nil", result.CheckTask.TaskID, err)
}
var actionPayload SourceActionPayload
if err := json.Unmarshal([]byte(execution.Payload), &actionPayload); err != nil {
t.Fatalf("json.Unmarshal(initial check payload) error = %v, want nil", err)
}
if actionPayload.Actor != "user:42" || actionPayload.Action != sourceActionCheck ||
actionPayload.TargetRevision != "" || actionPayload.ConfirmedRevision != "" {
t.Errorf("initial check payload = %+v, want real actor and credential-free check", actionPayload)
}
source, runtime := mustLoadPagesSource(t, ctx, project.ID)
if got, want := source.GitHubRepository, "OpenFlare/site"; got != want {
t.Errorf("GitHubRepository = %q, want %q", got, want)
}
if got, want := source.ReleaseSelector, githubReleaseSelectorLatest; got != want {
t.Errorf("ReleaseSelector = %q, want %q", got, want)
}
if got, want := source.AssetName, defaultGitHubAssetName; got != want {
t.Errorf("AssetName = %q, want %q", got, want)
}
if got, want := source.CheckIntervalMinutes, defaultCheckInterval; got != want {
t.Errorf("CheckIntervalMinutes = %d, want %d", got, want)
}
if got, want := source.SourceIdentity, "dbbd25307aaa3b88bc25353476940a049428655bd8421ac63045fdcb5fb23c9d"; got != want {
t.Errorf("SourceIdentity = %q, want %q", got, want)
}
if runtime.NextCheckAt == nil {
t.Error("GitHub latest NextCheckAt = nil, want scheduled value")
}
var taskCount int64
if err := db.DB(ctx).Model(&model.TaskExecution{}).Count(&taskCount).Error; err != nil {
t.Fatalf("count initial checks error = %v, want nil", err)
}
if _, err := UpdateSourceAs(ctx, project.ID, input, "user:42"); err != nil {
t.Fatalf("UpdateSourceAs(GitHub no-op) error = %v, want nil", err)
}
var noOpTaskCount int64
if err := db.DB(ctx).Model(&model.TaskExecution{}).Count(&noOpTaskCount).Error; err != nil {
t.Fatalf("count no-op checks error = %v, want nil", err)
}
if noOpTaskCount != taskCount {
t.Errorf("no-op initial check count = %d, want unchanged %d", noOpTaskCount, taskCount)
}
secret := "provider-switch-secret"
if _, err := UpdateSource(ctx, project.ID, SourceUpdateInput{
SourceType: PagesSourceTypeRemoteURL,
RemoteURLSet: true,
RemoteURL: "https://artifacts.example.com/site.zip?token=" + secret,
RemoteNetworkPolicy: RemoteNetworkPolicyPublic,
}); err != nil {
t.Fatalf("UpdateSource(GitHub to Remote) error = %v, want nil", err)
}
remote, _ := mustLoadPagesSource(t, ctx, project.ID)
if remote.GitHubRepository != "" || remote.ReleaseSelector != "" || remote.AssetName != "" ||
remote.AutoUpdateEnabled || remote.CheckIntervalMinutes != 0 {
t.Errorf("Remote switched source retained GitHub fields: %+v", remote)
}
if _, err := UpdateSourceAs(ctx, project.ID, input, "user:42"); err != nil {
t.Fatalf("UpdateSourceAs(Remote to GitHub) error = %v, want nil", err)
}
github, _ := mustLoadPagesSource(t, ctx, project.ID)
if github.RemoteURL != "" || github.RemoteNetworkPolicy != "" {
t.Errorf("GitHub switched source retained Remote fields: URL=%q policy=%q", github.RemoteURL, github.RemoteNetworkPolicy)
}
}
func TestGitHubSourceSaveSurvivesInitialCheckDispatchFailure(t *testing.T) {
ctx := setupPagesSourceTest(t)
project := mustCreatePagesSourceProject(t, ctx, "github-dispatch-warning")
result, err := UpdateSourceAs(ctx, project.ID, SourceUpdateInput{
SourceType: PagesSourceTypeGitHubRelease,
RepositoryURL: "https://github.com/a/b",
}, "user:9")
if err != nil {
t.Fatalf("UpdateSourceAs(dispatch failure) error = %v, want saved source with warning", err)
}
if result.CheckTask != nil || result.Warning != errPagesSourceInitialCheckWarning {
t.Errorf("UpdateSourceAs(dispatch failure) result = %+v, want warning and nil check task", result)
}
source, runtime := mustLoadPagesSource(t, ctx, project.ID)
if source.GitHubRepository != "a/b" || runtime.SyncStatus != pagesSourceStatusFailed ||
runtime.LastError != errPagesSourceInitialCheckWarning {
t.Errorf("saved source/runtime = repo:%q status:%q error:%q", source.GitHubRepository, runtime.SyncStatus, runtime.LastError)
}
}
func TestGitHubSourceRejectsUnsafeOrModeIncompatibleFields(t *testing.T) {
tests := []SourceUpdateInput{
{SourceType: PagesSourceTypeGitHubRelease, RepositoryURL: "http://github.com/a/b"},
{SourceType: PagesSourceTypeGitHubRelease, RepositoryURL: "https://github.com/a%20b/repo"},
{SourceType: PagesSourceTypeGitHubRelease, RepositoryURL: "https://github.com/a/b/extra"},
{SourceType: PagesSourceTypeGitHubRelease, RepositoryURL: "https://github.com//a/b"},
{SourceType: PagesSourceTypeGitHubRelease, RepositoryURL: "https://github.com/a/b/"},
{SourceType: PagesSourceTypeGitHubRelease, RepositoryURL: "https://github.com/a/b?"},
{SourceType: PagesSourceTypeGitHubRelease, RepositoryURL: "https://github.com/a/b#"},
{SourceType: PagesSourceTypeGitHubRelease, RepositoryURL: "https://github.com/a/b", AssetName: "dist\n.zip"},
{SourceType: PagesSourceTypeGitHubRelease, RepositoryURL: "https://github.com/a/b", AssetName: "dist\u202e.zip"},
{SourceType: PagesSourceTypeGitHubRelease, RepositoryURL: "https://github.com/a/b", AssetName: "dir/dist.zip"},
{SourceType: PagesSourceTypeGitHubRelease, RepositoryURL: "https://github.com/a/b", ReleaseSelector: "tag", ReleaseTag: "v1", CheckIntervalMinutes: 60},
{SourceType: PagesSourceTypeGitHubRelease, RepositoryURL: "https://github.com/a/b", ReleaseSelector: "tag", ReleaseTag: "v1", AutoUpdateEnabled: true},
{SourceType: PagesSourceTypeGitHubRelease, RepositoryURL: "https://github.com/a/b", ReleaseSelector: "tag", ReleaseTag: " v1"},
{SourceType: PagesSourceTypeGitHubRelease, RepositoryURL: "https://github.com/a/b", ReleaseSelector: "tag", ReleaseTag: "v1\n"},
{SourceType: PagesSourceTypeGitHubRelease, RepositoryURL: "https://github.com/a/b", ReleaseSelector: "tag", ReleaseTag: "v1\u2028draft"},
{SourceType: PagesSourceTypeGitHubRelease, RepositoryURL: "https://github.com/a/b", ReleaseSelector: "tag", ReleaseTag: `v1\draft`},
{SourceType: PagesSourceTypeGitHubRelease, RepositoryURL: "https://github.com/a/b", ReleaseSelector: "tag", ReleaseTag: "release//v1"},
{SourceType: PagesSourceTypeGitHubRelease, RepositoryURL: "https://github.com/a/b", ReleaseSelector: "tag", ReleaseTag: "release/.draft"},
{SourceType: PagesSourceTypeGitHubRelease, RepositoryURL: "https://github.com/a/b", ReleaseSelector: "tag", ReleaseTag: "release/v1.lock"},
{SourceType: PagesSourceTypeGitHubRelease, RepositoryURL: "https://github.com/a/b", ReleaseSelector: "latest", ReleaseTag: "v1"},
}
for _, input := range tests {
if err := validateGitHubSourceInput(input); err == nil {
t.Errorf("validateGitHubSourceInput(%+v) error = nil, want non-nil", input)
}
}
}
func TestGitHubSourceAcceptsLegalAssetAndTagCharacters(t *testing.T) {
tests := []SourceUpdateInput{
{
SourceType: PagesSourceTypeGitHubRelease,
RepositoryURL: "https://github.com/a/b",
AssetName: "dist?channel=stable&part#1.zip",
},
{
SourceType: PagesSourceTypeGitHubRelease,
RepositoryURL: "https://github.com/a/b",
AssetName: " dist.zip ",
},
{
SourceType: PagesSourceTypeGitHubRelease,
RepositoryURL: "https://github.com/a/b",
ReleaseSelector: "tag",
ReleaseTag: "release/v1#stable&build=1",
AssetName: "dist.zip",
},
{
SourceType: PagesSourceTypeGitHubRelease,
RepositoryURL: "https://github.com/a/b.git",
ReleaseSelector: "tag",
ReleaseTag: "@",
AssetName: "dist.zip",
},
{
SourceType: PagesSourceTypeGitHubRelease,
RepositoryURL: "https://github.com/a/b",
ReleaseSelector: "tag",
ReleaseTag: "release/v1.LOCK",
AssetName: "dist.zip",
},
}
for _, input := range tests {
if err := validateGitHubSourceInput(input); err != nil {
t.Errorf("validateGitHubSourceInput(%+v) error = %v, want nil", input, err)
}
}
}
func TestInitialCheckFailureUsesExactConfigFence(t *testing.T) {
ctx := setupPagesSourceTest(t)
project := mustCreatePagesSourceProject(t, ctx, "github-initial-fence")
source, _ := mustConfigureGitHubSourceWithoutDispatch(t, ctx, project.ID, SourceUpdateInput{
SourceType: PagesSourceTypeGitHubRelease,
RepositoryURL: "https://github.com/a/b",
})
staleVersion := source.ConfigVersion
if err := db.DB(ctx).Model(source).Update("config_version", staleVersion+1).Error; err != nil {
t.Fatalf("increment source config version error = %v, want nil", err)
}
markInitialCheckDispatchFailed(ctx, source.ID, staleVersion)
_, runtime := mustLoadPagesSource(t, ctx, project.ID)
if runtime.SyncStatus != pagesSourceStatusIdle || runtime.LastError != "" {
t.Errorf("stale initial failure runtime = status:%q error:%q, want unchanged idle", runtime.SyncStatus, runtime.LastError)
}
}
func TestGitHubCheckUsesETagAndDetectsSameReleaseReplacement(t *testing.T) {
ctx := setupPagesSourceTest(t)
project := mustCreatePagesSourceProject(t, ctx, "github-check")
source, _ := mustConfigureGitHubSourceWithoutDispatch(t, ctx, project.ID, SourceUpdateInput{
SourceType: PagesSourceTypeGitHubRelease,
RepositoryURL: "https://github.com/a/b",
})
appliedRevision := strings.Repeat("a", 64)
appliedDetail := `{"provider":"github","release_id":"100","asset_id":"1","tag":"release/v1","asset_name":"dist.zip"}`
if err := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).Where("source_id = ?", source.ID).Updates(map[string]any{
"etag": `"old-etag"`,
"last_applied_revision": appliedRevision,
"last_applied_detail": appliedDetail,
}).Error; err != nil {
t.Fatalf("seed GitHub runtime error = %v, want nil", err)
}
updatedAt := time.Date(2026, 7, 19, 10, 0, 0, 0, time.UTC)
var gotETag string
useFakeGitHubReleaseClient(t, &fakeGitHubReleaseClient{
resolve: func(_ context.Context, request githubrelease.ResolveRequest) (githubrelease.ResolveResult, error) {
gotETag = request.ETag
return githubrelease.ResolveResult{
ETag: `"new-etag"`,
Release: githubrelease.Release{ID: "100", Tag: "release/v1"},
Asset: githubrelease.Asset{ID: "2", Name: "dist.zip", State: "uploaded", UpdatedAt: updatedAt},
}, nil
},
download: func(context.Context, githubrelease.DownloadRequest) (*githubrelease.DownloadResult, error) {
t.Fatal("Download called during check, want resolve only")
return nil, nil
},
})
snapshot, outcome, err := acquireSourceLease(ctx, source.ID, source.ConfigVersion, sourceActionCheck)
if err != nil || outcome != sourceLeaseAcquired {
t.Fatalf("acquire check lease = (%+v, %q, %v), want acquired", snapshot, outcome, err)
}
result, err := checkGitHubSource(ctx, snapshot)
if err != nil {
t.Fatalf("checkGitHubSource() error = %v, want nil", err)
}
if result.Stale {
t.Error("checkGitHubSource() stale = true, want false")
}
if got, want := gotETag, `"old-etag"`; got != want {
t.Errorf("Resolve ETag = %q, want %q", got, want)
}
_, runtime := mustLoadPagesSource(t, ctx, project.ID)
if runtime.SyncStatus != pagesSourceStatusAttention || runtime.LastSeenRevision == "" {
t.Errorf("replacement runtime = status:%q seen:%q, want attention with revision", runtime.SyncStatus, runtime.LastSeenRevision)
}
view, err := GetSource(ctx, project.ID)
if err != nil {
t.Fatalf("GetSource() error = %v, want nil", err)
}
if view.LastSeen == nil || view.LastSeen.Label != "release/v1" {
t.Errorf("LastSeen = %+v, want full tag with slash", view.LastSeen)
}
if err := preflightGitHubSyncConfirmation(ctx, source.ID, ""); err == nil || err.Error() != errPagesSourceConfirmationNeeded {
t.Errorf("preflight without confirmation error = %v, want %q", err, errPagesSourceConfirmationNeeded)
}
if err := preflightGitHubSyncConfirmation(ctx, source.ID, runtime.LastSeenRevision); err != nil {
t.Errorf("preflight exact confirmation error = %v, want nil", err)
}
}
func TestGitHubCheckNotModifiedRefreshesRuntimeWithoutDeployment(t *testing.T) {
ctx := setupPagesSourceTest(t)
project := mustCreatePagesSourceProject(t, ctx, "github-304")
source, _ := mustConfigureGitHubSourceWithoutDispatch(t, ctx, project.ID, SourceUpdateInput{
SourceType: PagesSourceTypeGitHubRelease,
RepositoryURL: "https://github.com/a/b",
})
useFakeGitHubReleaseClient(t, &fakeGitHubReleaseClient{
resolve: func(_ context.Context, request githubrelease.ResolveRequest) (githubrelease.ResolveResult, error) {
return githubrelease.ResolveResult{NotModified: true, ETag: `"same"`}, nil
},
download: func(context.Context, githubrelease.DownloadRequest) (*githubrelease.DownloadResult, error) {
t.Fatal("Download called for 304 check")
return nil, nil
},
})
snapshot, _, _ := acquireSourceLease(ctx, source.ID, source.ConfigVersion, sourceActionCheck)
if _, err := checkGitHubSource(ctx, snapshot); err != nil {
t.Fatalf("checkGitHubSource(304) error = %v, want nil", err)
}
_, runtime := mustLoadPagesSource(t, ctx, project.ID)
if runtime.ETag != `"same"` || runtime.LastCheckedAt == nil || runtime.NextCheckAt == nil || runtime.LeaseToken != "" {
t.Errorf("304 runtime = %+v, want refreshed timestamps/etag and released lease", runtime)
}
var deployments int64
if err := db.DB(ctx).Model(&model.PagesDeployment{}).Where("project_id = ?", project.ID).Count(&deployments).Error; err != nil {
t.Fatalf("count deployments error = %v, want nil", err)
}
if deployments != 0 {
t.Errorf("deployments after check = %d, want 0", deployments)
}
}
func TestGitHubTargetMismatchPreservesAttentionAndExpeditesRecheck(t *testing.T) {
ctx := setupPagesSourceTest(t)
project := mustCreatePagesSourceProject(t, ctx, "github-target-mismatch")
source, _ := mustConfigureGitHubSourceWithoutDispatch(t, ctx, project.ID, SourceUpdateInput{
SourceType: PagesSourceTypeGitHubRelease,
RepositoryURL: "https://github.com/a/b",
})
appliedRevision := strings.Repeat("a", 64)
if err := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).Where("source_id = ?", source.ID).Updates(map[string]any{
"last_applied_revision": appliedRevision,
"last_applied_detail": `{"provider":"github","release_id":"100","asset_id":"1","tag":"v1","asset_name":"dist.zip"}`,
}).Error; err != nil {
t.Fatalf("seed applied runtime error = %v, want nil", err)
}
retryAt := time.Now().Add(2 * time.Hour).UTC()
useFakeGitHubReleaseClient(t, &fakeGitHubReleaseClient{
resolve: func(context.Context, githubrelease.ResolveRequest) (githubrelease.ResolveResult, error) {
return githubrelease.ResolveResult{
Release: githubrelease.Release{ID: "100", Tag: "v1"},
Asset: githubrelease.Asset{
ID: "2", Name: "dist.zip", State: "uploaded",
UpdatedAt: time.Date(2026, 7, 19, 12, 0, 0, 0, time.UTC),
},
RetryAt: &retryAt,
}, nil
},
download: func(context.Context, githubrelease.DownloadRequest) (*githubrelease.DownloadResult, error) {
t.Fatal("Download called after target mismatch")
return nil, nil
},
})
snapshot, _, _ := acquireSourceLease(ctx, source.ID, source.ConfigVersion, sourceActionSync)
outcome, err := syncGitHubSource(ctx, snapshot, pagesSourceCreatedBySystem, strings.Repeat("b", 64), "")
if err != nil {
t.Fatalf("syncGitHubSource(target mismatch) error = %v, want nil stale outcome", err)
}
if outcome == nil || !outcome.Stale {
t.Errorf("syncGitHubSource(target mismatch) = %+v, want stale", outcome)
}
_, runtime := mustLoadPagesSource(t, ctx, project.ID)
if runtime.SyncStatus != pagesSourceStatusAttention {
t.Errorf("target mismatch SyncStatus = %q, want %q", runtime.SyncStatus, pagesSourceStatusAttention)
}
if runtime.NextCheckAt == nil || runtime.NextCheckAt.Before(retryAt) {
t.Errorf("target mismatch NextCheckAt = %v, want server deadline >= %v", runtime.NextCheckAt, retryAt)
}
var deployments int64
if err := db.DB(ctx).Model(&model.PagesDeployment{}).Where("project_id = ?", project.ID).Count(&deployments).Error; err != nil {
t.Fatalf("count mismatch deployments error = %v, want nil", err)
}
if deployments != 0 {
t.Errorf("target mismatch deployments = %d, want 0", deployments)
}
}
func TestGitHubCheckLostLeaseReturnsStaleWithoutOverwritingRuntime(t *testing.T) {
ctx := setupPagesSourceTest(t)
project := mustCreatePagesSourceProject(t, ctx, "github-check-fence")
source, _ := mustConfigureGitHubSourceWithoutDispatch(t, ctx, project.ID, SourceUpdateInput{
SourceType: PagesSourceTypeGitHubRelease,
RepositoryURL: "https://github.com/a/b",
})
useFakeGitHubReleaseClient(t, &fakeGitHubReleaseClient{
resolve: func(context.Context, githubrelease.ResolveRequest) (githubrelease.ResolveResult, error) {
return githubrelease.ResolveResult{}, errors.New("transient provider failure")
},
download: func(context.Context, githubrelease.DownloadRequest) (*githubrelease.DownloadResult, error) {
return nil, errors.New("unexpected")
},
})
snapshot, _, _ := acquireSourceLease(ctx, source.ID, source.ConfigVersion, sourceActionCheck)
if err := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).Where("source_id = ?", source.ID).Updates(map[string]any{
"lease_token": "new-owner",
"lease_expires_at": time.Now().Add(time.Minute),
"sync_status": pagesSourceStatusSyncing,
"last_error": "new-owner-state",
}).Error; err != nil {
t.Fatalf("replace lease owner error = %v, want nil", err)
}
result, err := checkGitHubSource(ctx, snapshot)
if err != nil {
t.Fatalf("checkGitHubSource(lost lease) error = %v, want stale no-op", err)
}
if result == nil || !result.Stale {
t.Errorf("checkGitHubSource(lost lease) = %+v, want stale", result)
}
_, runtime := mustLoadPagesSource(t, ctx, project.ID)
if runtime.LeaseToken != "new-owner" || runtime.LastError != "new-owner-state" || runtime.SyncStatus != pagesSourceStatusSyncing {
t.Errorf("lost lease runtime = token:%q error:%q status:%q, want new owner state", runtime.LeaseToken, runtime.LastError, runtime.SyncStatus)
}
}
func TestGitHubCheckRateLimitUsesServerDeadlineAndSuppressesFastRetry(t *testing.T) {
ctx := setupPagesSourceTest(t)
project := mustCreatePagesSourceProject(t, ctx, "github-rate-limit")
source, _ := mustConfigureGitHubSourceWithoutDispatch(t, ctx, project.ID, SourceUpdateInput{
SourceType: PagesSourceTypeGitHubRelease,
RepositoryURL: "https://github.com/a/b",
})
retryAt := time.Now().Add(2 * time.Hour).UTC()
useFakeGitHubReleaseClient(t, &fakeGitHubReleaseClient{
resolve: func(context.Context, githubrelease.ResolveRequest) (githubrelease.ResolveResult, error) {
return githubrelease.ResolveResult{}, &githubrelease.Error{
Kind: githubrelease.ErrMetadata, StatusCode: 429, RetryAt: &retryAt,
}
},
download: func(context.Context, githubrelease.DownloadRequest) (*githubrelease.DownloadResult, error) {
return nil, errors.New("unexpected")
},
})
snapshot, _, _ := acquireSourceLease(ctx, source.ID, source.ConfigVersion, sourceActionCheck)
_, err := checkGitHubSource(ctx, snapshot)
if err == nil || err.Error() != errPagesSourceSyncFailed {
t.Fatalf("checkGitHubSource(rate limit) error = %v, want safe sync failure", err)
}
if !shouldSkipGitHubActionRetry(err) {
t.Error("shouldSkipGitHubActionRetry(rate limit) = false, want true")
}
_, runtime := mustLoadPagesSource(t, ctx, project.ID)
if runtime.NextCheckAt == nil || runtime.NextCheckAt.Before(retryAt) || runtime.SyncStatus != pagesSourceStatusFailed {
t.Errorf("rate limit runtime = next:%v status:%q, want deadline >= %v and failed", runtime.NextCheckAt, runtime.SyncStatus, retryAt)
}
}
func TestGitHubCheckInvalidResolvedTargetUsesServerDeadline(t *testing.T) {
ctx := setupPagesSourceTest(t)
project := mustCreatePagesSourceProject(t, ctx, "github-invalid-check-target")
source, _ := mustConfigureGitHubSourceWithoutDispatch(t, ctx, project.ID, SourceUpdateInput{
SourceType: PagesSourceTypeGitHubRelease,
RepositoryURL: "https://github.com/a/b",
})
retryAt := time.Now().Add(2 * time.Hour).UTC()
useFakeGitHubReleaseClient(t, &fakeGitHubReleaseClient{
resolve: func(context.Context, githubrelease.ResolveRequest) (githubrelease.ResolveResult, error) {
return githubrelease.ResolveResult{
Release: githubrelease.Release{ID: "1", Tag: "v1"},
Asset: githubrelease.Asset{
ID: "2", Name: "dist.zip", State: "uploaded",
UpdatedAt: time.Now().UTC(), Digest: "sha256:invalid",
},
RetryAt: &retryAt,
}, nil
},
download: func(context.Context, githubrelease.DownloadRequest) (*githubrelease.DownloadResult, error) {
t.Fatal("Download called after invalid check target")
return nil, nil
},
})
snapshot, _, _ := acquireSourceLease(ctx, source.ID, source.ConfigVersion, sourceActionCheck)
_, err := checkGitHubSource(ctx, snapshot)
if err == nil || err.Error() != errPagesSourceDigestInvalid {
t.Fatalf("checkGitHubSource(invalid target) error = %v, want %q", err, errPagesSourceDigestInvalid)
}
if !isPermanentSourceSyncError(err) {
t.Error("invalid check target classification = retryable, want permanent")
}
_, runtime := mustLoadPagesSource(t, ctx, project.ID)
if runtime.SyncStatus != pagesSourceStatusFailed || runtime.LastError != errPagesSourceDigestInvalid ||
runtime.NextCheckAt == nil || runtime.NextCheckAt.Before(retryAt) || runtime.LeaseToken != "" {
t.Errorf(
"invalid check target runtime = status:%q error:%q next:%v lease:%q, want failed/%q/deadline >= %v/cleared",
runtime.SyncStatus, runtime.LastError, runtime.NextCheckAt, runtime.LeaseToken,
errPagesSourceDigestInvalid, retryAt,
)
}
}
func TestGitHubSyncInvalidResolvedTargetUsesServerDeadline(t *testing.T) {
ctx := setupPagesSourceTest(t)
project := mustCreatePagesSourceProject(t, ctx, "github-invalid-sync-target")
source, _ := mustConfigureGitHubSourceWithoutDispatch(t, ctx, project.ID, SourceUpdateInput{
SourceType: PagesSourceTypeGitHubRelease,
RepositoryURL: "https://github.com/a/b",
})
retryAt := time.Now().Add(2 * time.Hour).UTC()
useFakeGitHubReleaseClient(t, &fakeGitHubReleaseClient{
resolve: func(context.Context, githubrelease.ResolveRequest) (githubrelease.ResolveResult, error) {
return githubrelease.ResolveResult{
Release: githubrelease.Release{ID: "1", Tag: "v1"},
Asset: githubrelease.Asset{
ID: "2", Name: "dist.zip", State: "uploaded",
UpdatedAt: time.Now().UTC(), Digest: "sha256:invalid",
},
RetryAt: &retryAt,
}, nil
},
download: func(context.Context, githubrelease.DownloadRequest) (*githubrelease.DownloadResult, error) {
t.Fatal("Download called after invalid sync target")
return nil, nil
},
})
snapshot, _, _ := acquireSourceLease(ctx, source.ID, source.ConfigVersion, sourceActionSync)
_, err := syncGitHubSource(ctx, snapshot, "user:7", "", "")
if err == nil || err.Error() != errPagesSourceDigestInvalid {
t.Fatalf("syncGitHubSource(invalid target) error = %v, want %q", err, errPagesSourceDigestInvalid)
}
if !isPermanentSourceSyncError(err) {
t.Error("invalid sync target classification = retryable, want permanent")
}
_, runtime := mustLoadPagesSource(t, ctx, project.ID)
if runtime.SyncStatus != pagesSourceStatusFailed || runtime.LastError != errPagesSourceDigestInvalid ||
runtime.NextCheckAt == nil || runtime.NextCheckAt.Before(retryAt) || runtime.LeaseToken != "" {
t.Errorf(
"invalid sync target runtime = status:%q error:%q next:%v lease:%q, want failed/%q/deadline >= %v/cleared",
runtime.SyncStatus, runtime.LastError, runtime.NextCheckAt, runtime.LeaseToken,
errPagesSourceDigestInvalid, retryAt,
)
}
}
func TestGitHubCheckHandlerSkipsProviderFastRetry(t *testing.T) {
tests := []struct {
name string
status int
retryDate bool
}{
{name: "bad request", status: http.StatusBadRequest},
{name: "rate limited forbidden", status: http.StatusForbidden, retryDate: true},
{name: "too many requests", status: http.StatusTooManyRequests, retryDate: true},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
ctx := setupPagesSourceTest(t)
project := mustCreatePagesSourceProject(t, ctx, "github-handler-"+strings.ReplaceAll(test.name, " ", "-"))
source, _ := mustConfigureGitHubSourceWithoutDispatch(t, ctx, project.ID, SourceUpdateInput{
SourceType: PagesSourceTypeGitHubRelease,
RepositoryURL: "https://github.com/a/b",
})
var retryAt *time.Time
if test.retryDate {
deadline := time.Now().Add(2 * time.Hour).UTC()
retryAt = &deadline
}
useFakeGitHubReleaseClient(t, &fakeGitHubReleaseClient{
resolve: func(context.Context, githubrelease.ResolveRequest) (githubrelease.ResolveResult, error) {
return githubrelease.ResolveResult{}, &githubrelease.Error{
Kind: githubrelease.ErrMetadata, StatusCode: test.status, RetryAt: retryAt,
}
},
download: func(context.Context, githubrelease.DownloadRequest) (*githubrelease.DownloadResult, error) {
t.Fatal("Download called after provider check failure")
return nil, nil
},
})
raw, err := json.Marshal(SourceActionPayload{
SourceID: source.ID, ConfigVersion: source.ConfigVersion,
Action: sourceActionCheck, Actor: "user:7",
})
if err != nil {
t.Fatalf("json.Marshal(check payload) error = %v, want nil", err)
}
result, err := (&SourceActionHandler{}).Execute(ctx, raw)
if result != nil || err == nil || !errors.Is(err, asynq.SkipRetry) {
t.Fatalf("SourceActionHandler.Execute(status %d) = result:%+v error:%v, want SkipRetry", test.status, result, err)
}
_, runtime := mustLoadPagesSource(t, ctx, project.ID)
if runtime.SyncStatus != pagesSourceStatusFailed || runtime.NextCheckAt == nil || runtime.LeaseToken != "" {
t.Errorf("provider failure runtime = status:%q next:%v lease:%q", runtime.SyncStatus, runtime.NextCheckAt, runtime.LeaseToken)
}
if retryAt != nil && (runtime.NextCheckAt == nil || runtime.NextCheckAt.Before(*retryAt)) {
t.Errorf("provider failure NextCheckAt = %v, want deadline >= %v", runtime.NextCheckAt, *retryAt)
}
})
}
}
func TestGitHubSyncActivatesWithMetadataRevisionAndPackageChecksum(t *testing.T) {
ctx := setupPagesSourceSyncTest(t)
project := mustCreatePagesSourceProject(t, ctx, "github-sync")
source, _ := mustConfigureGitHubSourceWithoutDispatch(t, ctx, project.ID, SourceUpdateInput{
SourceType: PagesSourceTypeGitHubRelease,
RepositoryURL: "https://github.com/a/b",
AssetName: "dist.zip",
})
packageBytes := testPagesZip(t, map[string]string{"index.html": "github-v1"})
packageHash := sha256.Sum256(packageBytes)
updatedAt := time.Date(2026, 7, 19, 11, 0, 0, 0, time.UTC)
retryAt := time.Now().Add(2 * time.Hour).UTC()
release := githubrelease.Release{ID: "200", Tag: "release/v2"}
asset := githubrelease.Asset{ID: "10", Name: "dist.zip", State: "uploaded", UpdatedAt: updatedAt}
client := &fakeGitHubReleaseClient{
resolve: func(context.Context, githubrelease.ResolveRequest) (githubrelease.ResolveResult, error) {
return githubrelease.ResolveResult{Release: release, Asset: asset, RetryAt: &retryAt}, nil
},
download: func(_ context.Context, request githubrelease.DownloadRequest) (*githubrelease.DownloadResult, error) {
path := filepath.Join(t.TempDir(), "download")
if err := os.WriteFile(path, packageBytes, 0o600); err != nil {
t.Fatalf("os.WriteFile(download) error = %v, want nil", err)
}
return &githubrelease.DownloadResult{
Path: path, Size: int64(len(packageBytes)), SHA256: hex.EncodeToString(packageHash[:]),
}, nil
},
}
useFakeGitHubReleaseClient(t, client)
snapshot, _, _ := acquireSourceLease(ctx, source.ID, source.ConfigVersion, sourceActionSync)
outcome, err := syncGitHubSource(ctx, snapshot, "user:7", "", "")
if err != nil {
t.Fatalf("syncGitHubSource() error = %v, want nil", err)
}
if outcome == nil || outcome.Deployment == nil || outcome.Stale {
t.Fatalf("syncGitHubSource() = %+v, want active deployment", outcome)
}
deployment, err := model.GetPagesDeploymentByID(ctx, outcome.Deployment.ID)
if err != nil {
t.Fatalf("GetPagesDeploymentByID(%d) error = %v, want nil", outcome.Deployment.ID, err)
}
if got, want := deployment.Checksum, hex.EncodeToString(packageHash[:]); got != want {
t.Errorf("deployment Checksum = %q, want package hash %q", got, want)
}
if deployment.SourceRevision == nil || *deployment.SourceRevision == deployment.Checksum {
t.Errorf("deployment SourceRevision = %v, want metadata revision distinct from package checksum", deployment.SourceRevision)
}
if got, want := deployment.SourceLabel, "release/v2"; got != want {
t.Errorf("deployment SourceLabel = %q, want %q", got, want)
}
if deployment.SourceType != PagesSourceTypeGitHubRelease || deployment.TriggerType != pagesSourceTriggerManualSync ||
deployment.CreatedBy != "user:7" {
t.Errorf("deployment provenance = type:%q trigger:%q actor:%q", deployment.SourceType, deployment.TriggerType, deployment.CreatedBy)
}
if strings.Contains(deployment.SourceMeta, "http") || strings.Contains(deployment.SourceMeta, "token") {
t.Errorf("deployment SourceMeta = %q, want no URL or token", deployment.SourceMeta)
}
if !strings.Contains(deployment.SourceMeta, `"tag":"release/v2"`) || strings.Contains(deployment.SourceMeta, `"label"`) {
t.Errorf("deployment SourceMeta = %q, want provider-specific tag field", deployment.SourceMeta)
}
_, runtime := mustLoadPagesSource(t, ctx, project.ID)
if runtime.NextCheckAt == nil || runtime.NextCheckAt.Before(retryAt) {
t.Errorf("sync runtime NextCheckAt = %v, want server deadline >= %v", runtime.NextCheckAt, retryAt)
}
secondSnapshot, _, _ := acquireSourceLease(ctx, source.ID, source.ConfigVersion, sourceActionSync)
second, err := syncGitHubSource(ctx, secondSnapshot, "user:7", "", "")
if err != nil {
t.Fatalf("syncGitHubSource(idempotent) error = %v, want nil", err)
}
if second == nil || !second.Reused || second.Deployment == nil || second.Deployment.ID != outcome.Deployment.ID {
t.Errorf("syncGitHubSource(idempotent) = %+v, want reused deployment %d", second, outcome.Deployment.ID)
}
}
func TestGitHubSyncActivatesExactConfirmedReplacement(t *testing.T) {
ctx := setupPagesSourceSyncTest(t)
project := mustCreatePagesSourceProject(t, ctx, "github-confirm-replacement")
source, _ := mustConfigureGitHubSourceWithoutDispatch(t, ctx, project.ID, SourceUpdateInput{
SourceType: PagesSourceTypeGitHubRelease,
RepositoryURL: "https://github.com/a/b",
AssetName: "dist.zip",
})
release := githubrelease.Release{ID: "300", Tag: "v3"}
asset := githubrelease.Asset{
ID: "12", Name: "dist.zip", State: "uploaded",
UpdatedAt: time.Date(2026, 7, 19, 13, 0, 0, 0, time.UTC),
}
target, err := buildGitHubSourceTarget(release, asset, nil)
if err != nil {
t.Fatalf("buildGitHubSourceTarget() error = %v, want nil", err)
}
if err := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).Where("source_id = ?", source.ID).Updates(map[string]any{
"last_seen_revision": target.Revision,
"last_seen_detail": target.DetailJSON,
"last_applied_revision": strings.Repeat("a", 64),
"last_applied_detail": `{"provider":"github","release_id":"300","asset_id":"11","tag":"v3","asset_name":"dist.zip"}`,
"sync_status": pagesSourceStatusAttention,
}).Error; err != nil {
t.Fatalf("seed replacement cursor error = %v, want nil", err)
}
packageBytes := testPagesZip(t, map[string]string{"index.html": "confirmed-v3"})
packageHash := sha256.Sum256(packageBytes)
useFakeGitHubReleaseClient(t, &fakeGitHubReleaseClient{
resolve: func(context.Context, githubrelease.ResolveRequest) (githubrelease.ResolveResult, error) {
return githubrelease.ResolveResult{Release: release, Asset: asset}, nil
},
download: func(context.Context, githubrelease.DownloadRequest) (*githubrelease.DownloadResult, error) {
path := filepath.Join(t.TempDir(), "confirmed.zip")
if err := os.WriteFile(path, packageBytes, 0o600); err != nil {
t.Fatalf("os.WriteFile(confirmed package) error = %v, want nil", err)
}
return &githubrelease.DownloadResult{
Path: path, Size: int64(len(packageBytes)), SHA256: hex.EncodeToString(packageHash[:]),
}, nil
},
})
snapshot, _, _ := acquireSourceLease(ctx, source.ID, source.ConfigVersion, sourceActionSync)
outcome, err := syncGitHubSource(ctx, snapshot, "user:9", "", target.Revision)
if err != nil {
t.Fatalf("syncGitHubSource(confirmed replacement) error = %v, want nil", err)
}
if outcome == nil || outcome.Deployment == nil || outcome.Stale {
t.Fatalf("syncGitHubSource(confirmed replacement) = %+v, want active deployment", outcome)
}
_, runtime := mustLoadPagesSource(t, ctx, project.ID)
if runtime.SyncStatus != pagesSourceStatusIdle || runtime.LastAppliedRevision != target.Revision {
t.Errorf("confirmed replacement runtime = status:%q applied:%q, want idle/%q", runtime.SyncStatus, runtime.LastAppliedRevision, target.Revision)
}
}
func TestSourceActionPayloadSeparatesSystemTargetAndUserConfirmation(t *testing.T) {
handler := &SourceActionHandler{}
revision := strings.Repeat("a", 64)
invalid := []SourceActionPayload{
{SourceID: 1, ConfigVersion: 1, Action: sourceActionSync, Actor: "user:1", TargetRevision: revision},
{SourceID: 1, ConfigVersion: 1, Action: sourceActionSync, Actor: pagesSourceCreatedBySystem, TriggerType: pagesSourceTriggerManualSync},
{SourceID: 1, ConfigVersion: 1, Action: sourceActionSync, Actor: pagesSourceCreatedBySystem, ConfirmedRevision: revision},
{SourceID: 1, ConfigVersion: 1, Action: sourceActionSync, Actor: pagesSourceCreatedBySystem, TargetRevision: revision, ConfirmedRevision: revision},
}
for _, payload := range invalid {
raw, _ := json.Marshal(payload)
if normalized, err := handler.ValidatePayload(raw); err == nil {
t.Errorf("ValidatePayload(%+v) = %s, nil; want error", payload, normalized)
}
}
valid := []SourceActionPayload{
{SourceID: 1, ConfigVersion: 1, Action: sourceActionSync, Actor: pagesSourceCreatedBySystem, TriggerType: pagesSourceTriggerScheduledAutoUpdate, TargetRevision: revision},
{SourceID: 1, ConfigVersion: 1, Action: sourceActionSync, Actor: "user:1", TriggerType: pagesSourceTriggerManualSync, ConfirmedRevision: revision},
}
for _, payload := range valid {
raw, _ := json.Marshal(payload)
if _, err := handler.ValidatePayload(raw); err != nil {
t.Errorf("ValidatePayload(%+v) error = %v, want nil", payload, err)
}
}
}
func TestGitHubProviderErrorsMapToSafeRetryClassification(t *testing.T) {
tests := []struct {
name string
provider error
want string
permanent bool
skipRetry bool
}{
{
name: "asset missing", provider: &githubrelease.Error{Kind: githubrelease.ErrAssetNotFound, StatusCode: 200},
want: errPagesSourceReleaseNotFound, permanent: true, skipRetry: true,
},
{
name: "digest mismatch", provider: &githubrelease.Error{Kind: githubrelease.ErrDigestMismatch, StatusCode: 200},
want: errPagesSourceDigestMismatch, permanent: true, skipRetry: true,
},
{
name: "rate limit", provider: &githubrelease.Error{
Kind: githubrelease.ErrMetadata, StatusCode: 429,
RetryAt: func() *time.Time { value := time.Now().Add(time.Hour); return &value }(),
},
want: errPagesSourceSyncFailed, permanent: false, skipRetry: true,
},
{
name: "network", provider: &githubrelease.Error{Kind: githubrelease.ErrDownload},
want: errPagesSourceSyncFailed, permanent: false, skipRetry: false,
},
{
name: "forbidden without retry", provider: &githubrelease.Error{Kind: githubrelease.ErrMetadata, StatusCode: 403},
want: errPagesSourceSyncFailed, permanent: true, skipRetry: true,
},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
domainErr := githubSourceDomainError(test.provider)
if got := domainErr.Error(); got != test.want {
t.Errorf("githubSourceDomainError() = %q, want %q", got, test.want)
}
var typedDomainError *githubSourceProviderDomainError
if !errors.As(domainErr, &typedDomainError) {
t.Fatalf("githubSourceDomainError() type = %T, want *githubSourceProviderDomainError", domainErr)
}
if got := typedDomainError.permanent; got != test.permanent {
t.Errorf("githubSourceProviderDomainError.permanent = %t, want %t", got, test.permanent)
}
if got := shouldSkipGitHubActionRetry(domainErr); got != test.skipRetry {
t.Errorf("shouldSkipGitHubActionRetry() = %t, want %t", got, test.skipRetry)
}
if strings.Contains(domainErr.Error(), "status=") || strings.Contains(domainErr.Error(), "repo=") {
t.Errorf("githubSourceDomainError() = %q, want stable Pages message", domainErr)
}
})
}
}
func TestGitHubSyncRejectsStaleConfirmationWithoutChangingActive(t *testing.T) {
ctx := setupPagesSourceSyncTest(t)
project := mustCreatePagesSourceProject(t, ctx, "github-confirm-stale")
oldActive := mustCreateActiveManualDeployment(t, ctx, project.ID, "old")
source, _ := mustConfigureGitHubSourceWithoutDispatch(t, ctx, project.ID, SourceUpdateInput{
SourceType: PagesSourceTypeGitHubRelease,
RepositoryURL: "https://github.com/a/b",
})
asset := githubrelease.Asset{ID: "2", Name: "dist.zip", State: "uploaded", UpdatedAt: time.Now().UTC()}
useFakeGitHubReleaseClient(t, &fakeGitHubReleaseClient{
resolve: func(context.Context, githubrelease.ResolveRequest) (githubrelease.ResolveResult, error) {
return githubrelease.ResolveResult{Release: githubrelease.Release{ID: "1", Tag: "v1"}, Asset: asset}, nil
},
download: func(context.Context, githubrelease.DownloadRequest) (*githubrelease.DownloadResult, error) {
t.Fatal("Download called for stale confirmation")
return nil, nil
},
})
snapshot, _, _ := acquireSourceLease(ctx, source.ID, source.ConfigVersion, sourceActionSync)
_, err := syncGitHubSource(ctx, snapshot, "user:1", "", strings.Repeat("f", 64))
if err == nil || err.Error() != errPagesSourceConfirmationStale {
t.Errorf("syncGitHubSource(stale confirmation) error = %v, want %q", err, errPagesSourceConfirmationStale)
}
storedProject, loadErr := model.GetPagesProjectByID(ctx, project.ID)
if loadErr != nil {
t.Fatalf("GetPagesProjectByID() error = %v, want nil", loadErr)
}
if storedProject.ActiveDeploymentID == nil || *storedProject.ActiveDeploymentID != oldActive.ID {
t.Errorf("ActiveDeploymentID = %v, want old active %d", storedProject.ActiveDeploymentID, oldActive.ID)
}
}
+99 -37
View File
@@ -15,13 +15,17 @@ import (
"path"
"path/filepath"
"regexp"
"strconv"
"strings"
"github.com/Rain-kl/Wavelet/internal/apps/upload"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/Rain-kl/Wavelet/pkg/pagesarchive"
"gorm.io/gorm"
"gorm.io/gorm/clause"
)
const (
@@ -31,11 +35,15 @@ 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"
pagesSourceIDMetadataKey = "pages_source_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 +123,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 +178,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,41 +261,56 @@ func ingestPagesDeploymentPackage(
ctx context.Context,
localPath string,
checksum string,
projectSlug string,
projectID uint,
fileName string,
format pagesarchive.Format,
) (upload.IngestResult, error) {
return ingestPagesDeploymentPackageWithSource(ctx, localPath, checksum, projectID, 0, fileName, format)
}
func ingestPagesDeploymentPackageWithSource(
ctx context.Context,
localPath string,
checksum string,
projectID uint,
sourceID uint,
fileName string,
format pagesarchive.Format,
) (upload.IngestResult, error) {
systemUser := repository.GetSystemUser(ctx)
accessMode := 0
extension := pagesarchive.NormalizeNameExtension(fileName, format)
extra := map[string]any{
pagesIngestMarkerKey: pagesIngestMarkerV2,
pagesProjectIDMetadataKey: strconv.FormatUint(uint64(projectID), 10),
}
if sourceID != 0 {
extra[pagesSourceIDMetadataKey] = strconv.FormatUint(uint64(sourceID), 10)
}
return upload.IngestFromLocalPath(ctx, localPath, upload.IngestRequest{
UserID: systemUser.ID,
FileName: fileName,
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),
},
Extra: extra,
},
})
}
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 +318,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,
+544 -137
View File
@@ -12,6 +12,7 @@ import (
"mime/multipart"
"net/url"
"os"
"path"
"strings"
"time"
@@ -21,6 +22,7 @@ import (
"github.com/Rain-kl/Wavelet/pkg/logger"
"github.com/Rain-kl/Wavelet/pkg/pagesarchive"
"gorm.io/gorm"
"gorm.io/gorm/clause"
)
// DeploymentPackage is a streamable Pages deployment artifact for agent download.
@@ -58,6 +60,9 @@ type DeploymentView struct {
FileCount int `json:"file_count"`
TotalSize int64 `json:"total_size"`
CreatedBy string `json:"created_by"`
SourceType string `json:"source_type"`
SourceLabel string `json:"source_label"`
TriggerType string `json:"trigger_type"`
CreatedAt time.Time `json:"created_at"`
ActivatedAt *time.Time `json:"activated_at"`
}
@@ -137,28 +142,57 @@ 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
}
contentConfigChanged := existing.RootDir != project.RootDir || existing.EntryFile != project.EntryFile
if contentConfigChanged &&
existing.ActiveDeploymentID != nil && *existing.ActiveDeploymentID != 0 {
if err := ensureDeploymentEntry(tx, *existing.ActiveDeploymentID, project.RootDir, project.EntryFile); err != nil {
return err
}
}
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,
}
if contentConfigChanged {
updates["content_config_version"] = existing.ContentConfigVersion + 1
var source model.PagesProjectSource
sourceErr := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).
Where("project_id = ?", existing.ID).
First(&source).Error
if sourceErr != nil && !errors.Is(sourceErr, gorm.ErrRecordNotFound) {
return sourceErr
}
if sourceErr == nil {
if err := fenceAndNormalizeRuntime(tx, source.ID); err != nil {
return err
}
}
}
return tx.Model(&existing).Updates(updates).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 +201,71 @@ 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)
}
}
var source model.PagesProjectSource
sourceErr := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).
Where("project_id = ?", project.ID).
First(&source).Error
if sourceErr != nil && !errors.Is(sourceErr, gorm.ErrRecordNotFound) {
return sourceErr
}
if sourceErr == nil {
var runtime model.PagesProjectSourceRuntime
runtimeErr := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).
Where("source_id = ?", source.ID).
First(&runtime).Error
if runtimeErr != nil && !errors.Is(runtimeErr, gorm.ErrRecordNotFound) {
return runtimeErr
}
if err := tx.Where("source_id = ?", source.ID).Delete(&model.PagesProjectSourceRuntime{}).Error; err != nil {
return err
}
if err := tx.Delete(&source).Error; err != nil {
return err
}
}
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),
@@ -194,14 +275,18 @@ func DeleteProject(ctx context.Context, id uint) error {
if err := tx.Where("project_id = ?", project.ID).Delete(&model.PagesDeployment{}).Error; err != nil {
return err
}
if err := tx.Delete(project).Error; err != nil {
if err := tx.Delete(&lockedProject).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 列出项目的全部部署。
@@ -258,7 +343,18 @@ func UploadDeployment(ctx context.Context, projectID uint, fileHeader *multipart
return nil, err
}
defer func() { _ = os.Remove(tempPath) }()
return createDeploymentFromTempPackage(ctx, project, tempPath, checksum, format, fileHeader.Filename, createdBy, limits)
return createDeploymentFromTempPackage(
ctx,
project,
tempPath,
checksum,
format,
fileHeader.Filename,
createdBy,
"manual_upload",
"manual_upload",
limits,
)
}
// UploadFromURLInput is the request body for downloading a deployment package from a remote URL.
@@ -278,7 +374,18 @@ func UploadDeploymentFromURL(ctx context.Context, projectID uint, rawURL string,
return nil, err
}
defer func() { _ = os.Remove(tempPath) }()
return createDeploymentFromTempPackage(ctx, project, tempPath, checksum, format, fileName, createdBy, limits)
return createDeploymentFromTempPackage(
ctx,
project,
tempPath,
checksum,
format,
fileName,
createdBy,
"manual_url",
"manual_url",
limits,
)
}
func createDeploymentFromTempPackage(
@@ -289,6 +396,8 @@ func createDeploymentFromTempPackage(
format pagesarchive.Format,
fileName string,
createdBy string,
sourceType string,
triggerType string,
limits pagesLimits,
) (*DeploymentView, error) {
if project == nil {
@@ -298,7 +407,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 +419,7 @@ func createDeploymentFromTempPackage(
ctx,
tempPath,
checksum,
project.Slug,
project.ID,
fileName,
format,
)
@@ -317,11 +429,29 @@ 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 uploadRecord model.Upload
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).
Where("id = ?", ingestResult.Upload.ID).
First(&uploadRecord).Error; err != nil {
return errors.New(errPagesPackageUploadMissing)
}
if uploadRecord.Status != model.UploadStatusUsed || uploadRecord.Type != upload.ReservedPagesDeploymentType {
return errors.New(errPagesPackageUploadMissing)
}
var maxNumber int
if err := tx.Model(&model.PagesDeployment{}).
Where("project_id = ?", project.ID).
@@ -338,6 +468,9 @@ func createDeploymentFromTempPackage(
FileCount: manifest.FileCount,
TotalSize: manifest.TotalSize,
CreatedBy: strings.TrimSpace(createdBy),
SourceType: sourceType,
SourceLabel: safeRemoteSourceLabel(fileName),
TriggerType: triggerType,
}
if err := tx.Create(deployment).Error; err != nil {
return err
@@ -357,7 +490,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 +514,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 +523,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 +537,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 +625,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 {
@@ -510,37 +660,208 @@ func selectDeploymentsToPrune(deployments []model.PagesDeployment, activeID uint
// ActivateDeployment 激活 Pages 部署。
func ActivateDeployment(ctx context.Context, projectID uint, deploymentID uint) (*View, error) {
project, err := model.GetPagesProjectByID(ctx, projectID)
return ActivateDeploymentAs(ctx, projectID, deploymentID, "system:pages-manual-activation")
}
// ActivateDeploymentAs activates a historical deployment and fences any
// configured source when the active deployment actually changes.
func ActivateDeploymentAs(ctx context.Context, projectID uint, deploymentID uint, actor string) (*View, error) {
if err := ensureActivationDeploymentUpload(ctx, projectID, deploymentID); err != nil {
return nil, err
}
audit, err := activateDeploymentTransaction(ctx, projectID, deploymentID, time.Now())
if err != nil {
return nil, err
}
if audit.Noop {
return GetProject(ctx, projectID)
}
logger.InfoF(ctx,
"[Pages] manual activation: actor=%s project_id=%d old_deployment_id=%d new_deployment_id=%d source_type=%s source_identity=%s auto_disabled=%t",
strings.TrimSpace(actor), projectID, audit.OldDeploymentID, deploymentID,
audit.SourceType, audit.SourceIdentity, audit.AutoDisabled,
)
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)
}
type deploymentActivationAudit struct {
OldDeploymentID uint
SourceType string
SourceIdentity string
AutoDisabled bool
Noop bool
}
type deploymentActivationSource struct {
Source *model.PagesProjectSource
Runtime *model.PagesProjectSourceRuntime
}
func ensureActivationDeploymentUpload(ctx context.Context, projectID uint, deploymentID uint) error {
deployment, err := model.GetPagesDeploymentByID(ctx, deploymentID)
if err != nil {
return err
}
if deployment.ProjectID != projectID {
return errors.New(errPagesDeploymentMismatch)
}
if deployment.UploadID != 0 {
return nil
}
return ensureDeploymentUploadRecord(ctx, deployment)
}
func activateDeploymentTransaction(
ctx context.Context,
projectID uint,
deploymentID uint,
now time.Time,
) (deploymentActivationAudit, error) {
audit := deploymentActivationAudit{}
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 project.ActiveDeploymentID != nil {
audit.OldDeploymentID = *project.ActiveDeploymentID
}
if audit.OldDeploymentID == deploymentID {
audit.Noop = true
return nil
}
sourceState, err := lockDeploymentActivationSource(tx, project.ID)
if err != nil {
return err
}
deployment, err := loadDeploymentActivationTarget(tx, &project, deploymentID)
if err != nil {
return err
}
if err := fenceDeploymentActivationSource(tx, sourceState, deployment, &audit); err != nil {
return err
}
return switchActiveDeploymentTx(tx, &project, deployment, now)
})
return audit, err
}
func lockDeploymentActivationSource(tx *gorm.DB, projectID uint) (*deploymentActivationSource, error) {
var source model.PagesProjectSource
err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).
Where("project_id = ?", projectID).
First(&source).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, nil
}
if err != nil {
return nil, err
}
var runtime model.PagesProjectSourceRuntime
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).
Where("source_id = ?", source.ID).
First(&runtime).Error; err != nil {
return nil, err
}
return &deploymentActivationSource{Source: &source, Runtime: &runtime}, nil
}
func loadDeploymentActivationTarget(
tx *gorm.DB,
project *model.PagesProject,
deploymentID uint,
) (*model.PagesDeployment, error) {
var deployment model.PagesDeployment
if err := tx.First(&deployment, deploymentID).Error; err != nil {
return nil, err
}
if deployment.ProjectID != project.ID {
return nil, errors.New(errPagesDeploymentMismatch)
}
now := time.Now()
if err = db.DB(ctx).Transaction(func(tx *gorm.DB) error {
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{
"status": model.PagesDeploymentStatusActive,
"activated_at": &now,
}).Error; err != nil {
return err
}
return tx.Model(project).Updates(map[string]any{
"active_deployment_id": deployment.ID,
}).Error
}); err != nil {
rootDir, err := validateAndNormalizePagesRootDir(project.RootDir)
if err != nil {
return nil, err
}
return GetProject(ctx, project.ID)
entryFile, err := validateAndNormalizePagesEntryFile(project.EntryFile)
if err != nil {
return nil, err
}
if err := ensureDeploymentEntry(tx, deployment.ID, rootDir, entryFile); err != nil {
return nil, err
}
var uploadRecord model.Upload
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).
Where("id = ?", deployment.UploadID).
First(&uploadRecord).Error; err != nil {
return nil, errors.New(errPagesPackageUploadMissing)
}
if uploadRecord.Status != model.UploadStatusUsed || uploadRecord.Type != upload.ReservedPagesDeploymentType {
return nil, errors.New(errPagesPackageUploadMissing)
}
return &deployment, nil
}
func fenceDeploymentActivationSource(
tx *gorm.DB,
state *deploymentActivationSource,
deployment *model.PagesDeployment,
audit *deploymentActivationAudit,
) error {
if state == nil {
return nil
}
audit.SourceType = state.Source.SourceType
audit.SourceIdentity = state.Source.SourceIdentity
audit.AutoDisabled = state.Source.AutoUpdateEnabled
if err := tx.Model(state.Source).Updates(map[string]any{
sourceColumnConfigVersion: state.Source.ConfigVersion + 1,
sourceColumnAutoUpdateEnabled: false,
}).Error; err != nil {
return err
}
if deployment.SourceIdentity != nil && *deployment.SourceIdentity == state.Source.SourceIdentity &&
deployment.SourceRevision != nil {
state.Runtime.LastAppliedRevision = *deployment.SourceRevision
state.Runtime.LastAppliedDetail = deployment.SourceMeta
} else {
state.Runtime.LastAppliedRevision = ""
state.Runtime.LastAppliedDetail = ""
}
return tx.Model(state.Runtime).Updates(map[string]any{
"last_applied_revision": state.Runtime.LastAppliedRevision,
"last_applied_detail": state.Runtime.LastAppliedDetail,
sourceRuntimeColumnSyncStatus: normalizedSourceRuntimeStatus(state.Runtime),
sourceRuntimeColumnLeaseToken: "",
sourceRuntimeColumnLeaseExpiresAt: nil,
}).Error
}
func switchActiveDeploymentTx(
tx *gorm.DB,
project *model.PagesProject,
deployment *model.PagesDeployment,
now time.Time,
) error {
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{
pagesDeploymentColumnStatus: model.PagesDeploymentStatusActive,
"activated_at": &now,
}).Error; err != nil {
return err
}
return tx.Model(project).Update("active_deployment_id", deployment.ID).Error
}
// GetDeploymentPackageHash returns the upload SHA-256 hash of the deployment package.
@@ -728,22 +1049,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 +1238,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) {
@@ -910,7 +1310,7 @@ func buildProject(existing *model.PagesProject, input Input) (*model.PagesProjec
return nil, errors.New(errPagesAPIProxyPassRequired)
}
parsedURL, err := url.Parse(apiProxyPass)
if err != nil || (parsedURL.Scheme != "http" && parsedURL.Scheme != "https") || parsedURL.Host == "" {
if err != nil || (parsedURL.Scheme != remoteSourceSchemeHTTP && parsedURL.Scheme != remoteSourceSchemeHTTPS) || parsedURL.Host == "" {
return nil, errors.New(errPagesAPIProxyPassInvalid)
}
}
@@ -923,7 +1323,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
}
@@ -979,6 +1383,9 @@ func buildDeploymentView(deployment *model.PagesDeployment) DeploymentView {
FileCount: deployment.FileCount,
TotalSize: deployment.TotalSize,
CreatedBy: deployment.CreatedBy,
SourceType: deployment.SourceType,
SourceLabel: deployment.SourceLabel,
TriggerType: deployment.TriggerType,
CreatedAt: deployment.CreatedAt,
ActivatedAt: deployment.ActivatedAt,
}
+232 -6
View File
@@ -17,6 +17,7 @@ import (
"path/filepath"
"testing"
"github.com/Rain-kl/Wavelet/internal/apps/upload"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
@@ -42,6 +43,8 @@ func setupPagesTestDB(t *testing.T) func() {
&model.PagesProject{},
&model.PagesDeployment{},
&model.PagesDeploymentFile{},
&model.PagesProjectSource{},
&model.PagesProjectSourceRuntime{},
&model.ConfigVersion{},
&model.SystemConfig{},
))
@@ -163,6 +166,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 +299,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 +541,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 +641,131 @@ func TestPruneProjectDeploymentHistory(t *testing.T) {
assert.True(t, hasLatest, "newest deployment must fill remaining slot")
}
func TestHistoryCountOnePreservesFreshCandidateUntilActivation(t *testing.T) {
cleanup := setupPagesTestDB(t)
defer cleanup()
_, disableStorage := setupPagesStorageMock(t)
defer disableStorage()
ctx := context.Background()
require.NoError(t, db.DB(ctx).Model(&model.SystemConfig{}).
Where("key = ?", model.ConfigKeyPagesMaxHistoryCount).
Update("value", "1").Error)
require.NoError(t, repository.InvalidateSystemConfigCache(ctx, model.ConfigKeyPagesMaxHistoryCount))
project, err := CreateProject(ctx, Input{Name: "Single History", Slug: "single-history", Enabled: true})
require.NoError(t, err)
active, err := UploadDeployment(ctx, project.ID, testPagesMultipartFile(t, "v1.zip", testPagesZip(t, map[string]string{
"index.html": "v1",
})), "user:1")
require.NoError(t, err)
_, err = ActivateDeployment(ctx, project.ID, active.ID)
require.NoError(t, err)
oldCandidate, err := UploadDeployment(ctx, project.ID, testPagesMultipartFile(t, "v2.zip", testPagesZip(t, map[string]string{
"index.html": "v2",
})), "user:1")
require.NoError(t, err)
deployments, err := model.ListPagesDeployments(ctx, project.ID)
require.NoError(t, err)
require.Len(t, deployments, 2)
newCandidate, err := UploadDeployment(ctx, project.ID, testPagesMultipartFile(t, "v3.zip", testPagesZip(t, map[string]string{
"index.html": "v3",
})), "user:1")
require.NoError(t, err)
deployments, err = model.ListPagesDeployments(ctx, project.ID)
require.NoError(t, err)
require.Len(t, deployments, 2)
kept := map[uint]bool{}
for _, deployment := range deployments {
kept[deployment.ID] = true
}
assert.True(t, kept[active.ID])
assert.True(t, kept[newCandidate.ID])
assert.False(t, kept[oldCandidate.ID])
var removedUpload model.Upload
require.NoError(t, db.DB(ctx).First(&removedUpload, oldCandidate.UploadID).Error)
assert.Equal(t, model.UploadStatusDeleted, removedUpload.Status)
_, err = ActivateDeployment(ctx, project.ID, newCandidate.ID)
require.NoError(t, err)
deployments, err = model.ListPagesDeployments(ctx, project.ID)
require.NoError(t, err)
require.Len(t, deployments, 1)
assert.Equal(t, newCandidate.ID, deployments[0].ID)
}
func TestPruneUsesLockTimeNewestCandidateInsteadOfStaleCaller(t *testing.T) {
cleanup := setupPagesTestDB(t)
defer cleanup()
_, disableStorage := setupPagesStorageMock(t)
defer disableStorage()
ctx := context.Background()
project, err := CreateProject(ctx, Input{Name: "Concurrent Candidate", Slug: "concurrent-candidate", Enabled: true})
require.NoError(t, err)
active, err := UploadDeployment(ctx, project.ID, testPagesMultipartFile(t, "v1.zip", testPagesZip(t, map[string]string{
"index.html": "v1",
})), "user:1")
require.NoError(t, err)
_, err = ActivateDeployment(ctx, project.ID, active.ID)
require.NoError(t, err)
staleCandidate, err := UploadDeployment(ctx, project.ID, testPagesMultipartFile(t, "v2.zip", testPagesZip(t, map[string]string{
"index.html": "v2",
})), "user:1")
require.NoError(t, err)
newCandidate, err := UploadDeployment(ctx, project.ID, testPagesMultipartFile(t, "v3.zip", testPagesZip(t, map[string]string{
"index.html": "v3",
})), "user:1")
require.NoError(t, err)
deleted, err := pruneProjectDeploymentHistoryOnce(ctx, project.ID, 1, staleCandidate.ID)
require.NoError(t, err)
assert.Equal(t, 1, deleted)
deployments, err := model.ListPagesDeployments(ctx, project.ID)
require.NoError(t, err)
require.Len(t, deployments, 2)
kept := map[uint]bool{}
for _, deployment := range deployments {
kept[deployment.ID] = true
}
assert.True(t, kept[active.ID])
assert.True(t, kept[newCandidate.ID])
assert.False(t, kept[staleCandidate.ID])
}
func TestDeleteDeploymentAndProjectSoftDeleteUnreferencedArtifacts(t *testing.T) {
cleanup := setupPagesTestDB(t)
defer cleanup()
_, disableStorage := setupPagesStorageMock(t)
defer disableStorage()
ctx := context.Background()
project, err := CreateProject(ctx, Input{Name: "Delete Artifacts", Slug: "delete-artifacts", Enabled: true})
require.NoError(t, err)
first, err := UploadDeployment(ctx, project.ID, testPagesMultipartFile(t, "first.zip", testPagesZip(t, map[string]string{
"index.html": "first",
})), "user:1")
require.NoError(t, err)
second, err := UploadDeployment(ctx, project.ID, testPagesMultipartFile(t, "second.zip", testPagesZip(t, map[string]string{
"index.html": "second",
})), "user:1")
require.NoError(t, err)
require.NoError(t, DeleteDeployment(ctx, project.ID, second.ID))
var secondUpload model.Upload
require.NoError(t, db.DB(ctx).First(&secondUpload, second.UploadID).Error)
assert.Equal(t, model.UploadStatusDeleted, secondUpload.Status)
require.NoError(t, DeleteProject(ctx, project.ID))
var firstUpload model.Upload
require.NoError(t, db.DB(ctx).First(&firstUpload, first.UploadID).Error)
assert.Equal(t, model.UploadStatusDeleted, firstUpload.Status)
_, err = model.GetPagesProjectByID(ctx, project.ID)
assert.Error(t, err)
}
func testPagesTarGz(t *testing.T, files map[string]string) []byte {
t.Helper()
@@ -0,0 +1,64 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pages
import (
"context"
"errors"
"fmt"
"strings"
"github.com/Rain-kl/Wavelet/internal/apps/upload"
"github.com/Rain-kl/Wavelet/internal/model"
)
// ProjectLatestPackageMetadata describes the active package limits published to Agents.
type ProjectLatestPackageMetadata struct {
DeploymentID uint
Hash string
PackageSize int64
FileCount int
TotalSize int64
}
// GetProjectLatestPackageMetadata returns one coherent metadata snapshot for a
// project's currently active deployment.
func GetProjectLatestPackageMetadata(ctx context.Context, projectID uint) (*ProjectLatestPackageMetadata, error) {
deployment, err := resolveProjectActiveDeploymentForAgent(ctx, projectID)
if err != nil {
return nil, err
}
if deployment.UploadID == 0 {
if err := ensureDeploymentUploadRecord(ctx, deployment); err != nil {
return nil, err
}
deployment, err = model.GetPagesDeploymentByID(ctx, deployment.ID)
if err != nil {
return nil, err
}
}
if deployment.UploadID == 0 {
return nil, errors.New(errPagesDeploymentNotFound)
}
uploadRecord, err := upload.GetActiveUpload(ctx, deployment.UploadID)
if err != nil {
return nil, fmt.Errorf("pages 部署包不存在: %w", err)
}
hash := strings.TrimSpace(uploadRecord.Hash)
if hash == "" {
hash = strings.TrimSpace(deployment.Checksum)
}
if hash == "" {
return nil, errors.New(errPagesDeploymentHashMissing)
}
return &ProjectLatestPackageMetadata{
DeploymentID: deployment.ID,
Hash: hash,
PackageSize: uploadRecord.FileSize,
FileCount: deployment.FileCount,
TotalSize: deployment.TotalSize,
}, nil
}
@@ -0,0 +1,81 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pages
import (
"context"
"crypto/sha256"
"encoding/hex"
"fmt"
"testing"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
)
func TestGetProjectLatestPackageMetadataPreservesZeroTotalSize(t *testing.T) {
cleanup := setupPagesTestDB(t)
defer cleanup()
_, disableStorage := setupPagesStorageMock(t)
defer disableStorage()
ctx := context.Background()
project, err := CreateProject(ctx, Input{
Name: "Empty Files",
Slug: "empty-files",
Enabled: true,
EntryFile: "index.html",
})
if err != nil {
t.Fatalf("CreateProject() error = %v", err)
}
packageBytes := testPagesZip(t, map[string]string{
"index.html": "",
".gitkeep": "",
})
deployment, err := UploadDeployment(
ctx,
project.ID,
testPagesMultipartFile(t, "empty-files.zip", packageBytes),
"test",
)
if err != nil {
t.Fatalf("UploadDeployment() error = %v", err)
}
if _, err := ActivateDeployment(ctx, project.ID, deployment.ID); err != nil {
t.Fatalf("ActivateDeployment() error = %v", err)
}
if err := db.DB(ctx).Create(&model.ConfigVersion{
Version: "v-package-metadata",
SnapshotJSON: fmt.Sprintf(
`{"routes":[{"upstream_type":"pages","pages_project_id":%d}]}`,
project.ID,
),
SupportFilesJSON: "[]",
Checksum: "package-metadata-config",
IsActive: true,
CreatedBy: "test",
}).Error; err != nil {
t.Fatalf("create active ConfigVersion error = %v", err)
}
got, err := GetProjectLatestPackageMetadata(ctx, project.ID)
if err != nil {
t.Fatalf("GetProjectLatestPackageMetadata(%d) error = %v", project.ID, err)
}
wantHashBytes := sha256.Sum256(packageBytes)
wantHash := hex.EncodeToString(wantHashBytes[:])
if got.DeploymentID != deployment.ID || got.Hash != wantHash {
t.Errorf("GetProjectLatestPackageMetadata(%d) identity = (%d, %q), want (%d, %q)",
project.ID, got.DeploymentID, got.Hash, deployment.ID, wantHash)
}
if got.PackageSize != int64(len(packageBytes)) {
t.Errorf("GetProjectLatestPackageMetadata(%d).PackageSize = %d, want %d",
project.ID, got.PackageSize, len(packageBytes))
}
if got.FileCount != 2 || got.TotalSize != 0 {
t.Errorf("GetProjectLatestPackageMetadata(%d) content = (%d files, %d bytes), want (2 files, 0 bytes)",
project.ID, got.FileCount, got.TotalSize)
}
}
+22 -7
View File
@@ -7,6 +7,7 @@ import (
"context"
"encoding/json"
"fmt"
"path"
"strings"
"github.com/Rain-kl/Wavelet/internal/model"
@@ -102,7 +103,10 @@ func rebindPagesRouteMaps(ctx context.Context, routes []map[string]json.RawMessa
if err != nil {
return false, err
}
deployment := buildLivePagesDeployment(project, activeDeployment)
deployment, err := buildLivePagesDeployment(project, activeDeployment)
if err != nil {
return false, err
}
projectIDCopy := project.ID
originURL := fmt.Sprintf("openflare-pages://project/%d", project.ID)
@@ -181,15 +185,26 @@ func loadActivePagesProject(ctx context.Context, projectID uint, siteName string
return project, activeDeployment, nil
}
func buildLivePagesDeployment(project *model.PagesProject, active *model.PagesDeployment) *openrestyrender.PagesDeployment {
entryFile := strings.TrimSpace(project.EntryFile)
if entryFile == "" {
entryFile = defaultPagesEntryFile
func buildLivePagesDeployment(
project *model.PagesProject,
active *model.PagesDeployment,
) (*openrestyrender.PagesDeployment, error) {
rootDir, err := validateAndNormalizePagesRootDir(project.RootDir)
if err != nil {
return nil, err
}
entryFile, err := validateAndNormalizePagesEntryFile(project.EntryFile)
if err != nil {
return nil, err
}
fallbackPath := strings.TrimSpace(project.SPAFallbackPath)
if fallbackPath == "" {
fallbackPath = defaultPagesFallbackPath
}
localRoot := openrestyrender.PagesProjectLocalRoot(project.ID)
if rootDir != "" {
localRoot = path.Join(localRoot, rootDir)
}
return &openrestyrender.PagesDeployment{
ProjectID: project.ID,
ProjectSlug: strings.TrimSpace(project.Slug),
@@ -203,8 +218,8 @@ func buildLivePagesDeployment(project *model.PagesProject, active *model.PagesDe
APIProxyPath: strings.TrimSpace(project.APIProxyPath),
APIProxyPass: strings.TrimSpace(project.APIProxyPass),
APIProxyRewrite: strings.TrimSpace(project.APIProxyRewrite),
LocalRoot: openrestyrender.PagesProjectLocalRoot(project.ID),
}
LocalRoot: localRoot,
}, nil
}
func rawJSONString(raw json.RawMessage) (string, bool) {
+6 -3
View File
@@ -20,9 +20,11 @@ func TestRebindSnapshotPagesToCurrentActive(t *testing.T) {
ctx := context.Background()
project, err := CreateProject(ctx, Input{
Name: "Rebind Site",
Slug: "rebind-site",
Enabled: true,
Name: "Rebind Site",
Slug: "rebind-site",
Enabled: true,
RootDir: "public/site",
EntryFile: "index.html",
})
require.NoError(t, err)
@@ -84,4 +86,5 @@ func TestRebindSnapshotPagesToCurrentActive(t *testing.T) {
deployment := route["pages_deployment"].(map[string]any)
assert.EqualValues(t, active.ID, deployment["deployment_id"])
assert.Equal(t, "new-checksum", deployment["checksum"])
assert.Equal(t, "__OPENFLARE_PAGES_DIR__/projects/1/current/public/site", deployment["local_root"])
}
+254 -4
View File
@@ -4,12 +4,20 @@
package pages
import (
"encoding/json"
"errors"
"fmt"
"io"
"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/Rain-kl/Wavelet/pkg/logger"
"github.com/gin-gonic/gin"
"gorm.io/gorm"
)
func handleLogicError(c *gin.Context, err error) bool {
@@ -19,6 +27,64 @@ func handleLogicError(c *gin.Context, err error) bool {
return apiutil.AbortNotFoundIfMissing(c, err, errPagesProjectNotFound)
}
func handleSourceLogicError(c *gin.Context, err error) bool {
if err == nil {
return false
}
if errors.Is(err, gorm.ErrRecordNotFound) || err.Error() == errPagesSourceNotFound {
response.AbortNotFound(c, errPagesSourceNotFound)
return true
}
switch err.Error() {
case errPagesSourceActionBusy:
response.AbortConflict(c, errPagesSourceActionBusy)
case errPagesSourceTypeRequired,
errPagesSourceTypeUnsupported,
errPagesSourceRemoteFields,
errPagesSourceRemoteURLRequired,
errPagesSourceRemoteURLMode,
errPagesSourceRemoteURLInvalid,
errPagesSourceNetworkPolicy,
errPagesSourceGitHubFields,
errPagesSourceRepositoryInvalid,
errPagesSourceSelectorInvalid,
errPagesSourceAssetNameInvalid,
errPagesSourceCheckInterval,
errPagesSourceAutoNotAvailable,
errPagesSourceReleaseNotFound,
errPagesSourceDigestInvalid,
errPagesSourceDigestMismatch,
errPagesSourceConfirmationNeeded,
errPagesSourceConfirmationStale,
errPagesSourceCheckUnsupported,
errPagesSourceActionInvalid:
response.AbortBadRequest(c, err.Error())
case errPagesSourceTaskDispatchFailed:
response.AbortInternal(c, errPagesSourceInternal)
default:
logger.ErrorF(c.Request.Context(), "[PagesSource] API operation failed: error=%v", err)
response.AbortInternal(c, errPagesSourceInternal)
}
return true
}
func decodeStrictJSON(c *gin.Context, target any, allowEmpty bool) bool {
decoder := json.NewDecoder(c.Request.Body)
decoder.DisallowUnknownFields()
if err := decoder.Decode(target); err != nil {
if allowEmpty && errors.Is(err, io.EOF) {
return true
}
response.AbortBadRequest(c, errPagesSourceActionInvalid)
return false
}
if err := ensureJSONEOF(decoder); err != nil {
response.AbortBadRequest(c, errPagesSourceActionInvalid)
return false
}
return true
}
func deploymentIDParam(c *gin.Context) (uint, bool) {
raw := c.Param("deployment_id")
if raw == "" {
@@ -33,6 +99,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 项目,需要管理员权限
@@ -162,6 +237,168 @@ func DeleteProjectHandler(c *gin.Context) {
c.JSON(http.StatusOK, response.OKNil())
}
// GetSourceHandler 获取 Pages 项目的部署源。
// @Summary 获取 Pages 部署源
// @Description 返回脱敏后的项目部署源配置与运行状态,需要管理员权限
// @Tags openflare-pages
// @Produce json
// @Security SessionCookie
// @Param id path int true "项目 ID"
// @Success 200 {object} response.Any{data=pages.SourceView} "部署源"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "项目或部署源不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/pages/{id}/source [get]
func GetSourceHandler(c *gin.Context) {
projectID, ok := apiutil.IDParam(c)
if !ok {
return
}
source, err := GetSource(c.Request.Context(), projectID)
if handleSourceLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(source))
}
// UpdateSourceHandler 创建或更新 Pages 项目部署源。
// @Summary 更新 Pages 部署源
// @Description 支持 Remote URL 与公开 GitHub Release 来源;敏感地址仅写入,不会在响应中返回
// @Tags openflare-pages
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param id path int true "项目 ID"
// @Param request body pages.SourceUpdateInput true "部署源配置"
// @Success 200 {object} response.Any{data=pages.SourceUpdateResult} "更新结果"
// @Failure 400 {object} response.Any "配置无效"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "项目不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/pages/{id}/source/update [post]
func UpdateSourceHandler(c *gin.Context) {
projectID, ok := apiutil.IDParam(c)
if !ok {
return
}
var input SourceUpdateInput
if !decodeStrictJSON(c, &input, false) {
return
}
actor, ok := currentPagesActor(c)
if !ok {
return
}
result, err := UpdateSourceAs(c.Request.Context(), projectID, input, actor)
if handleSourceLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(result))
}
// DeleteSourceHandler 将 Pages 项目切换回手动部署模式。
// @Summary 删除 Pages 部署源
// @Description 幂等删除持久部署源;已有部署历史与当前生产部署保持不变
// @Tags openflare-pages
// @Produce json
// @Security SessionCookie
// @Param id path int true "项目 ID"
// @Success 200 {object} response.Any{data=pages.SourceView} "手动来源视图"
// @Failure 400 {object} response.Any "参数错误"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "项目不存在"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/pages/{id}/source/delete [post]
func DeleteSourceHandler(c *gin.Context) {
projectID, ok := apiutil.IDParam(c)
if !ok {
return
}
source, err := DeleteSource(c.Request.Context(), projectID)
if handleSourceLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(source))
}
// CheckSourceHandler 请求检查 Pages 部署源。
// @Summary 检查 Pages 部署源
// @Description 异步检查 GitHub Release 来源;Remote URL 来源不支持检查更新
// @Tags openflare-pages
// @Produce json
// @Security SessionCookie
// @Param id path int true "项目 ID"
// @Success 200 {object} response.Any{data=pages.SourceActionReceipt} "任务回执"
// @Failure 400 {object} response.Any "当前来源不支持检查"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "部署源不存在"
// @Failure 409 {object} response.Any "来源任务正在执行"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/pages/{id}/source/check [post]
func CheckSourceHandler(c *gin.Context) {
projectID, ok := apiutil.IDParam(c)
if !ok {
return
}
actor, ok := currentPagesActor(c)
if !ok {
return
}
receipt, err := DispatchSourceAction(c.Request.Context(), projectID, sourceActionCheck, actor, "")
if handleSourceLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(receipt))
}
// SourceSyncInput is the optional source sync action payload.
type SourceSyncInput struct {
ConfirmedRevision string `json:"confirmed_revision"`
}
// SyncSourceHandler 请求同步并发布 Pages 部署源。
// @Summary 同步并发布 Pages 部署源
// @Description 异步下载、校验并原子激活来源部署包;空请求体与空 JSON 对象均有效
// @Tags openflare-pages
// @Accept json
// @Produce json
// @Security SessionCookie
// @Param id path int true "项目 ID"
// @Param request body pages.SourceSyncInput false "同步参数"
// @Success 200 {object} response.Any{data=pages.SourceActionReceipt} "任务回执"
// @Failure 400 {object} response.Any "参数或来源类型无效"
// @Failure 401 {object} response.Any "未登录"
// @Failure 404 {object} response.Any "部署源不存在"
// @Failure 409 {object} response.Any "来源任务正在执行"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/d/pages/{id}/source/sync [post]
func SyncSourceHandler(c *gin.Context) {
projectID, ok := apiutil.IDParam(c)
if !ok {
return
}
var input SourceSyncInput
if !decodeStrictJSON(c, &input, true) {
return
}
actor, ok := currentPagesActor(c)
if !ok {
return
}
receipt, err := DispatchSourceAction(
c.Request.Context(),
projectID,
sourceActionSync,
actor,
input.ConfirmedRevision,
)
if handleSourceLogicError(c, err) {
return
}
c.JSON(http.StatusOK, response.OK(receipt))
}
// ListDeploymentsHandler 列出项目的全部部署。
// @Summary 列出 Pages 部署
// @Description 返回指定项目的全部部署记录,需要管理员权限
@@ -214,7 +451,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
}
@@ -223,7 +464,8 @@ func UploadDeploymentHandler(c *gin.Context) {
// UploadDeploymentFromURLHandler 从 URL 下载并创建 Pages 部署。
// @Summary 从 URL 导入 Pages 部署包
// @Description 从用户提供的 HTTP(S) 链接下载部署包并创建部署记录;服务端使用浏览器伪装请求头拉取,允许内网地址与不安全 TLS 证书,需要管理员权限
// @Description 已弃用的一次性 URL 导入;使用 trusted_internal 策略兼容内网与自签名证书,不创建持久部署源
// @Deprecated
// @Tags openflare-pages
// @Accept json
// @Produce json
@@ -247,7 +489,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
}
@@ -278,7 +524,11 @@ func ActivateDeploymentHandler(c *gin.Context) {
if !ok {
return
}
project, err := ActivateDeployment(c.Request.Context(), projectID, deploymentID)
actor, ok := currentPagesActor(c)
if !ok {
return
}
project, err := ActivateDeploymentAs(c.Request.Context(), projectID, deploymentID, actor)
if handleLogicError(c, err) {
return
}
@@ -0,0 +1,224 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pages
import (
"bytes"
"encoding/json"
"fmt"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/oauth"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/task"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/alicebob/miniredis/v2"
"github.com/gin-gonic/gin"
"github.com/hibiken/asynq"
"github.com/redis/go-redis/v9"
)
type sourceHandlerEnvelope struct {
ErrorMsg string `json:"error_msg"`
Data json.RawMessage `json:"data"`
}
func newPagesSourceTestRouter(userID uint64) *gin.Engine {
router := testhelper.NewTestGinEngine(func(ctx *gin.Context) {
oauth.SetToContext(ctx, oauth.UserObjKey, &model.User{ID: userID})
ctx.Next()
})
router.GET("/api/v1/d/pages/:id/source", GetSourceHandler)
router.POST("/api/v1/d/pages/:id/source/update", UpdateSourceHandler)
router.POST("/api/v1/d/pages/:id/source/delete", DeleteSourceHandler)
router.POST("/api/v1/d/pages/:id/source/check", CheckSourceHandler)
router.POST("/api/v1/d/pages/:id/source/sync", SyncSourceHandler)
return router
}
func performPagesSourceRequest(
t *testing.T,
router http.Handler,
method string,
path string,
body []byte,
) (int, sourceHandlerEnvelope) {
t.Helper()
request := httptest.NewRequest(method, path, bytes.NewReader(body))
if body != nil {
request.Header.Set("Content-Type", "application/json")
}
recorder := httptest.NewRecorder()
router.ServeHTTP(recorder, request)
var envelope sourceHandlerEnvelope
if err := json.Unmarshal(recorder.Body.Bytes(), &envelope); err != nil {
t.Fatalf("json.Unmarshal(%s %s response %q) error = %v, want nil", method, path, recorder.Body.String(), err)
}
return recorder.Code, envelope
}
func setupPagesSourceDispatchTest(t *testing.T) {
t.Helper()
miniRedis, err := miniredis.Run()
if err != nil {
t.Fatalf("miniredis.Run() error = %v, want nil", err)
}
redisClient := redis.NewClient(&redis.Options{Addr: miniRedis.Addr()})
asynqClient := asynq.NewClient(asynq.RedisClientOpt{Addr: miniRedis.Addr()})
previousRedis := db.Redis
previousAsynqClient := task.AsynqClient
db.Redis = redisClient
task.AsynqClient = asynqClient
task.RegisterTaskMeta(PagesSourceActionMeta)
t.Cleanup(func() {
_ = asynqClient.Close()
_ = redisClient.Close()
miniRedis.Close()
task.AsynqClient = previousAsynqClient
db.Redis = previousRedis
})
}
func TestPagesSourceHandlersReturnStableActionErrors(t *testing.T) {
ctx := setupPagesSourceTest(t)
router := newPagesSourceTestRouter(42)
manualProject := mustCreatePagesSourceProject(t, ctx, "handler-no-source")
code, envelope := performPagesSourceRequest(
t,
router,
http.MethodPost,
fmt.Sprintf("/api/v1/d/pages/%d/source/sync", manualProject.ID),
nil,
)
if got, want := code, http.StatusNotFound; got != want {
t.Errorf("POST source/sync without source status = %d, want %d", got, want)
}
if got, want := envelope.ErrorMsg, errPagesSourceNotFound; got != want {
t.Errorf("POST source/sync without source error = %q, want %q", got, want)
}
remoteProject := mustCreatePagesSourceProject(t, ctx, "handler-check")
_, _ = mustConfigureRemoteSource(
t,
ctx,
remoteProject.ID,
"https://example.com/site.zip?token=handler-secret",
RemoteNetworkPolicyPublic,
)
code, envelope = performPagesSourceRequest(
t,
router,
http.MethodPost,
fmt.Sprintf("/api/v1/d/pages/%d/source/check", remoteProject.ID),
nil,
)
if got, want := code, http.StatusBadRequest; got != want {
t.Errorf("POST remote source/check status = %d, want %d", got, want)
}
if got, want := envelope.ErrorMsg, errPagesSourceCheckUnsupported; got != want {
t.Errorf("POST remote source/check error = %q, want %q", got, want)
}
if strings.Contains(string(envelope.Data), "handler-secret") || strings.Contains(envelope.ErrorMsg, "handler-secret") {
t.Errorf("POST remote source/check response = %+v, want no URL secret", envelope)
}
busyProject := mustCreatePagesSourceProject(t, ctx, "handler-busy")
busySource, _ := mustConfigureRemoteSource(
t,
ctx,
busyProject.ID,
"https://example.com/site.zip",
RemoteNetworkPolicyPublic,
)
future := time.Now().Add(time.Minute)
if err := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).
Where("source_id = ?", busySource.ID).
Updates(map[string]any{
"sync_status": pagesSourceStatusSyncing,
"lease_token": "busy-owner",
"lease_expires_at": &future,
}).Error; err != nil {
t.Fatalf("seed busy source runtime error = %v, want nil", err)
}
code, envelope = performPagesSourceRequest(
t,
router,
http.MethodPost,
fmt.Sprintf("/api/v1/d/pages/%d/source/sync", busyProject.ID),
[]byte(`{}`),
)
if got, want := code, http.StatusConflict; got != want {
t.Errorf("POST busy source/sync status = %d, want %d", got, want)
}
if got, want := envelope.ErrorMsg, errPagesSourceActionBusy; got != want {
t.Errorf("POST busy source/sync error = %q, want %q", got, want)
}
}
func TestSyncSourceHandlerAcceptsEmptyBodyAndEmptyObject(t *testing.T) {
ctx := setupPagesSourceTest(t)
setupPagesSourceDispatchTest(t)
router := newPagesSourceTestRouter(77)
project := mustCreatePagesSourceProject(t, ctx, "handler-empty-sync")
source, _ := mustConfigureRemoteSource(
t,
ctx,
project.ID,
"https://example.com/site.zip?token=dispatch-secret",
RemoteNetworkPolicyPublic,
)
path := fmt.Sprintf("/api/v1/d/pages/%d/source/sync", project.ID)
for _, test := range []struct {
name string
body []byte
}{
{name: "empty body", body: nil},
{name: "empty object", body: []byte(`{}`)},
} {
t.Run(test.name, func(t *testing.T) {
code, envelope := performPagesSourceRequest(t, router, http.MethodPost, path, test.body)
if got, want := code, http.StatusOK; got != want {
t.Fatalf("POST source/sync (%s) status = %d, want %d; error=%q", test.name, got, want, envelope.ErrorMsg)
}
if envelope.ErrorMsg != "" {
t.Errorf("POST source/sync (%s) error = %q, want empty", test.name, envelope.ErrorMsg)
}
var receipt SourceActionReceipt
if err := json.Unmarshal(envelope.Data, &receipt); err != nil {
t.Fatalf("json.Unmarshal(source/sync %s receipt) error = %v, want nil", test.name, err)
}
if receipt.TaskID == "" || receipt.ExecutionID == "" || receipt.Action != sourceActionSync {
t.Errorf("POST source/sync (%s) receipt = %+v, want task/execution IDs and action %q", test.name, receipt, sourceActionSync)
}
})
}
var executions []model.TaskExecution
if err := db.DB(ctx).Where("task_type = ?", PagesSourceActionTask).Order("id asc").Find(&executions).Error; err != nil {
t.Fatalf("list Pages source task executions error = %v, want nil", err)
}
if got, want := len(executions), 2; got != want {
t.Fatalf("Pages source task execution count = %d, want %d", got, want)
}
for _, execution := range executions {
if strings.Contains(execution.Payload, "dispatch-secret") || strings.Contains(execution.Payload, "http") {
t.Errorf("task execution %q payload = %s, want no Remote URL secret", execution.TaskID, execution.Payload)
}
var payload SourceActionPayload
if err := json.Unmarshal([]byte(execution.Payload), &payload); err != nil {
t.Errorf("json.Unmarshal(task execution %q payload) error = %v, want nil", execution.TaskID, err)
continue
}
if payload.SourceID != source.ID || payload.ConfigVersion != source.ConfigVersion || payload.Actor != "user:77" {
t.Errorf("task execution %q payload = %+v, want source=%d config=%d actor=user:77", execution.TaskID, payload, source.ID, source.ConfigVersion)
}
}
}
@@ -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())
}
+570
View File
@@ -0,0 +1,570 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pages
import (
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"net"
"net/url"
"path"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"gorm.io/gorm"
"gorm.io/gorm/clause"
)
const (
// PagesSourceTypeManual represents projects without a persisted source row.
PagesSourceTypeManual = "manual"
// PagesSourceTypeRemoteURL represents a persisted artifact URL.
PagesSourceTypeRemoteURL = "remote_url"
// PagesSourceTypeGitHubRelease represents a public GitHub Release asset.
PagesSourceTypeGitHubRelease = "github_release"
pagesSourceStatusIdle = "idle"
pagesSourceStatusChecking = "checking"
pagesSourceStatusUpdateAvailable = "update_available"
pagesSourceStatusSyncing = "syncing"
pagesSourceStatusFailed = "failed"
pagesSourceStatusAttention = "attention"
defaultRemoteAssetLabel = "pages-package"
defaultGitHubAssetName = "dist.zip"
defaultCheckInterval = 60
minimumCheckInterval = 5
maximumCheckInterval = 1440
)
// SourceUpdateInput is the discriminated source configuration payload.
// GitHub fields are accepted by the decoder so mode-incompatible values can be
// rejected deterministically; GitHub itself is enabled in Phase 2.
type SourceUpdateInput struct {
SourceType string `json:"source_type"`
RemoteURLSet bool `json:"remote_url_set"`
RemoteURL string `json:"remote_url"`
RemoteNetworkPolicy string `json:"remote_network_policy"`
RepositoryURL string `json:"repository_url"`
ReleaseSelector string `json:"release_selector"`
ReleaseTag string `json:"release_tag"`
AssetName string `json:"asset_name"`
AutoUpdateEnabled bool `json:"auto_update_enabled"`
CheckIntervalMinutes int `json:"check_interval_minutes"`
}
// SourceRevisionView is a credential-free source cursor.
type SourceRevisionView struct {
Revision string `json:"revision"`
Label string `json:"label"`
AssetName string `json:"asset_name,omitempty"`
}
// SourceView is the safe discriminated source view returned to the console.
type SourceView struct {
SourceType string `json:"source_type"`
HasRemoteURL bool `json:"has_remote_url,omitempty"`
DisplayURL string `json:"display_url,omitempty"`
RemoteNetworkPolicy string `json:"remote_network_policy,omitempty"`
GitHubRepository string `json:"github_repository,omitempty"`
ReleaseSelector string `json:"release_selector,omitempty"`
ReleaseTag string `json:"release_tag,omitempty"`
AssetName string `json:"asset_name,omitempty"`
AutoUpdateEnabled *bool `json:"auto_update_enabled,omitempty"`
CheckIntervalMinutes int `json:"check_interval_minutes,omitempty"`
SyncStatus string `json:"sync_status,omitempty"`
UpdateAvailable bool `json:"update_available,omitempty"`
LastSeen *SourceRevisionView `json:"last_seen,omitempty"`
LastApplied *SourceRevisionView `json:"last_applied,omitempty"`
LastCheckedAt *time.Time `json:"last_checked_at,omitempty"`
LastSyncedAt *time.Time `json:"last_synced_at,omitempty"`
NextCheckAt *time.Time `json:"next_check_at,omitempty"`
LastError string `json:"last_error,omitempty"`
}
// SourceActionReceipt identifies the internal task execution created by an action API.
type SourceActionReceipt struct {
TaskID string `json:"task_id"`
ExecutionID string `json:"execution_id"`
Action string `json:"action"`
}
// SourceUpdateResult is returned after persisting source configuration.
type SourceUpdateResult struct {
Source *SourceView `json:"source"`
CheckTask *SourceActionReceipt `json:"check_task"`
Warning string `json:"warning"`
}
type sourceDetail struct {
Provider string `json:"provider"`
DisplayName string `json:"display_name,omitempty"`
Tag string `json:"tag,omitempty"`
LegacyLabel string `json:"label,omitempty"`
AssetName string `json:"asset_name,omitempty"`
ReleaseID string `json:"release_id,omitempty"`
AssetID string `json:"asset_id,omitempty"`
AssetUpdatedAt string `json:"asset_updated_at,omitempty"`
Digest string `json:"digest,omitempty"`
}
type remoteSourceConfig struct {
URL string
Policy string
Identity string
}
// GetSource returns the current persisted source or a manual discriminator.
func GetSource(ctx context.Context, projectID uint) (*SourceView, error) {
if _, err := model.GetPagesProjectByID(ctx, projectID); err != nil {
return nil, err
}
source, runtime, err := loadSourceByProject(ctx, projectID)
if errors.Is(err, gorm.ErrRecordNotFound) {
return &SourceView{SourceType: PagesSourceTypeManual}, nil
}
if err != nil {
return nil, err
}
return buildSourceView(source, runtime)
}
// UpdateSource creates or updates a source. Direct callers use the system actor;
// HTTP handlers should call UpdateSourceAs so the initial check is auditable.
func UpdateSource(ctx context.Context, projectID uint, input SourceUpdateInput) (*SourceUpdateResult, error) {
return UpdateSourceAs(ctx, projectID, input, pagesSourceCreatedBySystem)
}
// UpdateSourceAs persists source configuration and queues the first GitHub check
// after commit when the GitHub configuration was materially changed.
func UpdateSourceAs(
ctx context.Context,
projectID uint,
input SourceUpdateInput,
actor string,
) (*SourceUpdateResult, error) {
if !validPagesSourceActor(actor) {
return nil, errors.New(errPagesSourceActionInvalid)
}
if err := validateSourceUpdateInput(input); err != nil {
return nil, err
}
changed := false
var persistedSource model.PagesProjectSource
err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
var err error
switch strings.TrimSpace(input.SourceType) {
case PagesSourceTypeRemoteURL:
changed, err = updateRemoteSourceTx(tx, projectID, input)
case PagesSourceTypeGitHubRelease:
changed, err = updateGitHubSourceTx(tx, projectID, input)
default:
err = errors.New(errPagesSourceTypeUnsupported)
}
if err != nil || !changed || strings.TrimSpace(input.SourceType) != PagesSourceTypeGitHubRelease {
return err
}
return tx.Where("project_id = ?", projectID).First(&persistedSource).Error
})
if err != nil {
return nil, err
}
view, err := GetSource(ctx, projectID)
if err != nil {
return nil, err
}
result := &SourceUpdateResult{Source: view, Warning: ""}
if changed && strings.TrimSpace(input.SourceType) == PagesSourceTypeGitHubRelease {
receipt, dispatchErr := dispatchSourceActionSnapshot(ctx, persistedSource, sourceActionCheck, actor, "", "", "manual")
if dispatchErr != nil {
result.Warning = errPagesSourceInitialCheckWarning
markInitialCheckDispatchFailed(ctx, persistedSource.ID, persistedSource.ConfigVersion)
} else {
result.CheckTask = receipt
}
}
return result, nil
}
func updateRemoteSourceTx(tx *gorm.DB, projectID uint, input SourceUpdateInput) (bool, error) {
var project model.PagesProject
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).First(&project, projectID).Error; err != nil {
return false, err
}
existing, hasExisting, err := loadProjectSourceForUpdate(tx, projectID)
if err != nil {
return false, err
}
config, err := buildRemoteSourceConfig(existing, hasExisting, input)
if err != nil {
return false, err
}
if !hasExisting {
return true, createRemoteSourceTx(tx, projectID, config)
}
return updateExistingRemoteSourceTx(tx, existing, config)
}
func loadProjectSourceForUpdate(tx *gorm.DB, projectID uint) (*model.PagesProjectSource, bool, error) {
var source model.PagesProjectSource
err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).
Where("project_id = ?", projectID).
First(&source).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return &source, false, nil
}
if err != nil {
return nil, false, err
}
return &source, true, nil
}
func buildRemoteSourceConfig(
existing *model.PagesProjectSource,
hasExisting bool,
input SourceUpdateInput,
) (remoteSourceConfig, error) {
remoteURL, err := resolveUpdatedRemoteURL(existing, hasExisting, input)
if err != nil {
return remoteSourceConfig{}, err
}
parsedURL, err := parseRemoteSourceURL(remoteURL)
if err != nil {
return remoteSourceConfig{}, err
}
policy := strings.TrimSpace(input.RemoteNetworkPolicy)
if policy == "" {
policy = RemoteNetworkPolicyPublic
}
return remoteSourceConfig{
URL: remoteURL,
Policy: policy,
Identity: remoteSourceIdentity(parsedURL),
}, nil
}
func createRemoteSourceTx(tx *gorm.DB, projectID uint, config remoteSourceConfig) error {
source := &model.PagesProjectSource{
ProjectID: projectID,
SourceType: PagesSourceTypeRemoteURL,
RemoteURL: config.URL,
RemoteNetworkPolicy: config.Policy,
AutoUpdateEnabled: false,
CheckIntervalMinutes: 0,
ConfigVersion: 1,
SourceIdentity: config.Identity,
}
if err := tx.Create(source).Error; err != nil {
return err
}
return tx.Create(&model.PagesProjectSourceRuntime{
SourceID: source.ID,
SyncStatus: pagesSourceStatusIdle,
}).Error
}
func updateExistingRemoteSourceTx(
tx *gorm.DB,
existing *model.PagesProjectSource,
config remoteSourceConfig,
) (bool, error) {
if !remoteSourceConfigChanged(existing, config) {
return false, nil
}
var runtime model.PagesProjectSourceRuntime
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).
Where("source_id = ?", existing.ID).
First(&runtime).Error; err != nil {
return false, err
}
identityChanged := existing.SourceIdentity != config.Identity
if err := tx.Model(existing).Updates(map[string]any{
"source_type": PagesSourceTypeRemoteURL,
"remote_url": config.URL,
"remote_network_policy": config.Policy,
"github_repository": "",
"release_selector": "",
"release_tag": "",
"asset_name": "",
sourceColumnAutoUpdateEnabled: false,
"check_interval_minutes": 0,
sourceColumnConfigVersion: existing.ConfigVersion + 1,
"source_identity": config.Identity,
}).Error; err != nil {
return false, err
}
return true, resetRuntimeAfterSourceUpdate(tx, &runtime, identityChanged)
}
func remoteSourceConfigChanged(existing *model.PagesProjectSource, config remoteSourceConfig) bool {
return existing.SourceType != PagesSourceTypeRemoteURL ||
existing.RemoteURL != config.URL ||
existing.RemoteNetworkPolicy != config.Policy ||
existing.GitHubRepository != "" ||
existing.ReleaseSelector != "" ||
existing.ReleaseTag != "" ||
existing.AssetName != "" ||
existing.AutoUpdateEnabled ||
existing.CheckIntervalMinutes != 0
}
// DeleteSource idempotently switches a project back to manual mode.
func DeleteSource(ctx context.Context, projectID uint) (*SourceView, error) {
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 source model.PagesProjectSource
err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).
Where("project_id = ?", projectID).
First(&source).Error
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil
}
if err != nil {
return err
}
var runtime model.PagesProjectSourceRuntime
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).
Where("source_id = ?", source.ID).
First(&runtime).Error; err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
return err
}
if err := tx.Where("source_id = ?", source.ID).Delete(&model.PagesProjectSourceRuntime{}).Error; err != nil {
return err
}
return tx.Delete(&source).Error
})
if err != nil {
return nil, err
}
return &SourceView{SourceType: PagesSourceTypeManual}, nil
}
func validateRemoteSourceInput(input SourceUpdateInput) error {
sourceType := strings.TrimSpace(input.SourceType)
if sourceType == "" {
return errors.New(errPagesSourceTypeRequired)
}
if sourceType != PagesSourceTypeRemoteURL {
return errors.New(errPagesSourceTypeUnsupported)
}
if strings.TrimSpace(input.RepositoryURL) != "" || strings.TrimSpace(input.ReleaseSelector) != "" ||
strings.TrimSpace(input.ReleaseTag) != "" || strings.TrimSpace(input.AssetName) != "" ||
input.AutoUpdateEnabled || input.CheckIntervalMinutes != 0 {
return errors.New(errPagesSourceRemoteFields)
}
policy := strings.TrimSpace(input.RemoteNetworkPolicy)
if policy != "" && policy != RemoteNetworkPolicyPublic && policy != RemoteNetworkPolicyTrustedInternal {
return errors.New(errPagesSourceNetworkPolicy)
}
if !input.RemoteURLSet && strings.TrimSpace(input.RemoteURL) != "" {
return errors.New(errPagesSourceRemoteURLMode)
}
if input.RemoteURLSet && strings.TrimSpace(input.RemoteURL) == "" {
return errors.New(errPagesSourceRemoteURLRequired)
}
return nil
}
func validateSourceUpdateInput(input SourceUpdateInput) error {
switch strings.TrimSpace(input.SourceType) {
case PagesSourceTypeRemoteURL:
return validateRemoteSourceInput(input)
case PagesSourceTypeGitHubRelease:
return validateGitHubSourceInput(input)
case "":
return errors.New(errPagesSourceTypeRequired)
default:
return errors.New(errPagesSourceTypeUnsupported)
}
}
func resolveUpdatedRemoteURL(existing *model.PagesProjectSource, hasExisting bool, input SourceUpdateInput) (string, error) {
if input.RemoteURLSet {
return strings.TrimSpace(input.RemoteURL), nil
}
if !hasExisting || existing.SourceType != PagesSourceTypeRemoteURL || strings.TrimSpace(existing.RemoteURL) == "" {
return "", errors.New(errPagesSourceRemoteURLRequired)
}
return existing.RemoteURL, nil
}
func parseRemoteSourceURL(raw string) (*url.URL, error) {
parsed, err := url.Parse(strings.TrimSpace(raw))
if err != nil || parsed.Host == "" || parsed.User != nil || parsed.Fragment != "" {
return nil, errors.New(errPagesSourceRemoteURLInvalid)
}
parsed.Scheme = strings.ToLower(parsed.Scheme)
if parsed.Scheme != remoteSourceSchemeHTTP && parsed.Scheme != remoteSourceSchemeHTTPS {
return nil, errors.New(errPagesSourceRemoteURLInvalid)
}
if strings.TrimSpace(parsed.Hostname()) == "" {
return nil, errors.New(errPagesSourceRemoteURLInvalid)
}
return parsed, nil
}
func remoteSourceIdentity(parsed *url.URL) string {
hostname := strings.ToLower(parsed.Hostname())
port := parsed.Port()
if (parsed.Scheme == "https" && port == "443") || (parsed.Scheme == "http" && port == "80") {
port = ""
}
host := hostname
if port != "" {
host = net.JoinHostPort(hostname, port)
} else if strings.Contains(hostname, ":") {
host = "[" + hostname + "]"
}
canonicalPath := parsed.EscapedPath()
if canonicalPath == "" {
canonicalPath = "/"
}
canonicalPath = path.Clean("/" + strings.TrimPrefix(canonicalPath, "/"))
canonical := parsed.Scheme + "://" + host + canonicalPath
sum := sha256.Sum256([]byte("remote_url|" + canonical))
return hex.EncodeToString(sum[:])
}
func displayRemoteSourceURL(raw string) string {
parsed, err := parseRemoteSourceURL(raw)
if err != nil {
return ""
}
hadQuery := parsed.RawQuery != ""
parsed.RawQuery = ""
display := parsed.String()
if hadQuery {
display += "?***"
}
return display
}
func loadSourceByProject(ctx context.Context, projectID uint) (*model.PagesProjectSource, *model.PagesProjectSourceRuntime, error) {
var source model.PagesProjectSource
if err := db.DB(ctx).Where("project_id = ?", projectID).First(&source).Error; err != nil {
return nil, nil, err
}
var runtime model.PagesProjectSourceRuntime
if err := db.DB(ctx).Where("source_id = ?", source.ID).First(&runtime).Error; err != nil {
return nil, nil, err
}
return &source, &runtime, nil
}
func buildSourceView(source *model.PagesProjectSource, runtime *model.PagesProjectSourceRuntime) (*SourceView, error) {
if source == nil || runtime == nil {
return nil, errors.New(errPagesSourceNotFound)
}
view := &SourceView{
SourceType: source.SourceType,
SyncStatus: runtime.SyncStatus,
LastSyncedAt: runtime.LastSyncedAt,
LastError: runtime.LastError,
}
if runtime.LastAppliedRevision != "" {
view.LastApplied = revisionView(runtime.LastAppliedRevision, runtime.LastAppliedDetail)
}
switch source.SourceType {
case PagesSourceTypeRemoteURL:
view.HasRemoteURL = strings.TrimSpace(source.RemoteURL) != ""
view.DisplayURL = displayRemoteSourceURL(source.RemoteURL)
view.RemoteNetworkPolicy = source.RemoteNetworkPolicy
case PagesSourceTypeGitHubRelease:
view.LastCheckedAt = runtime.LastCheckedAt
view.NextCheckAt = runtime.NextCheckAt
view.UpdateAvailable = runtime.LastSeenRevision != "" && runtime.LastSeenRevision != runtime.LastAppliedRevision
if runtime.LastSeenRevision != "" {
view.LastSeen = revisionView(runtime.LastSeenRevision, runtime.LastSeenDetail)
}
view.GitHubRepository = source.GitHubRepository
view.ReleaseSelector = source.ReleaseSelector
view.ReleaseTag = source.ReleaseTag
view.AssetName = source.AssetName
autoUpdateEnabled := source.AutoUpdateEnabled
view.AutoUpdateEnabled = &autoUpdateEnabled
view.CheckIntervalMinutes = source.CheckIntervalMinutes
default:
return nil, errors.New(errPagesSourceTypeUnsupported)
}
return view, nil
}
func revisionView(revision string, detailJSON string) *SourceRevisionView {
detail := sourceDetail{}
_ = unmarshalSourceDetail(detailJSON, &detail)
label := sourceDetailLabel(detail)
if label == "" {
label = defaultRemoteAssetLabel
}
return &SourceRevisionView{
Revision: revision,
Label: label,
AssetName: detail.AssetName,
}
}
func sourceDetailLabel(detail sourceDetail) string {
if detail.Provider == githubSourceDetailProvider {
if label := strings.TrimSpace(detail.Tag); label != "" {
return label
}
return strings.TrimSpace(detail.LegacyLabel)
}
if label := strings.TrimSpace(detail.DisplayName); label != "" {
return label
}
return strings.TrimSpace(detail.LegacyLabel)
}
func unmarshalSourceDetail(raw string, detail *sourceDetail) error {
if detail == nil || strings.TrimSpace(raw) == "" {
return nil
}
return json.Unmarshal([]byte(raw), detail)
}
func resetRuntimeAfterSourceUpdate(tx *gorm.DB, runtime *model.PagesProjectSourceRuntime, identityChanged bool) error {
updates := map[string]any{
sourceRuntimeColumnLeaseToken: "",
sourceRuntimeColumnLeaseExpiresAt: nil,
sourceRuntimeColumnLastError: "",
}
if identityChanged {
updates["etag"] = ""
updates["last_seen_revision"] = ""
updates["last_seen_detail"] = ""
updates["last_applied_revision"] = ""
updates["last_applied_detail"] = ""
updates["last_checked_at"] = nil
updates["last_synced_at"] = nil
updates["next_check_at"] = nil
updates[sourceRuntimeColumnSyncStatus] = pagesSourceStatusIdle
} else {
updates[sourceRuntimeColumnSyncStatus] = normalizedSourceRuntimeStatus(runtime)
}
return tx.Model(runtime).Updates(updates).Error
}
func normalizedSourceRuntimeStatus(runtime *model.PagesProjectSourceRuntime) string {
if runtime == nil {
return pagesSourceStatusIdle
}
if sourceHasSameReleaseReplacement(runtime) {
return pagesSourceStatusAttention
}
if runtime.LastSeenRevision != "" && runtime.LastSeenRevision != runtime.LastAppliedRevision {
return pagesSourceStatusUpdateAvailable
}
return pagesSourceStatusIdle
}
@@ -0,0 +1,51 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pages
import (
"context"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/upload"
"github.com/Rain-kl/Wavelet/internal/model"
)
func syncRemoteSource(
ctx context.Context,
snapshot *sourceExecutionSnapshot,
actor string,
) (*sourceSyncOutcome, error) {
return syncRemoteSourceWithTrigger(ctx, snapshot, actor, pagesSourceTriggerManualSync)
}
func syncGitHubSource(
ctx context.Context,
snapshot *sourceExecutionSnapshot,
actor string,
targetRevision string,
confirmedRevision string,
) (*sourceSyncOutcome, error) {
return syncGitHubSourceWithTrigger(
ctx, snapshot, actor, targetRevision, confirmedRevision, pagesSourceTriggerManualSync,
)
}
func commitSourceDeployment(
ctx context.Context,
snapshot *sourceExecutionSnapshot,
revision string,
packageChecksum string,
detail sourceDetail,
detailJSON string,
actor string,
manifest *deploymentManifest,
ingestResult upload.IngestResult,
hasIngest bool,
nextCheckNotBefore *time.Time,
) (*model.PagesDeployment, bool, bool, error) {
return commitSourceDeploymentWithTrigger(
ctx, snapshot, revision, packageChecksum, detail, detailJSON, actor,
pagesSourceTriggerManualSync, manifest, ingestResult, hasIngest, nextCheckNotBefore,
)
}
@@ -0,0 +1,328 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pages
import (
"context"
"errors"
"fmt"
"strconv"
"time"
"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"
"gorm.io/gorm"
"gorm.io/gorm/clause"
)
const pagesOrphanUploadIsolation = 2 * time.Hour
// PagesOrphanCleanupSummary describes one bounded delayed compensation pass.
// Every candidate is counted exactly once in one outcome field.
//
//nolint:revive // Keep the domain-qualified exported name for scanner/task result clarity.
type PagesOrphanCleanupSummary struct {
Candidates int `json:"candidates"`
Reconciled int `json:"reconciled"`
Referenced int `json:"referenced"`
LeaseBusy int `json:"lease_busy"`
InvalidMarker int `json:"invalid_marker"`
Skipped int `json:"skipped"`
Failed int `json:"failed"`
}
type pagesOrphanMarker struct {
ProjectID uint
SourceID *uint
}
type pagesOrphanCleanupOutcome uint8
const (
pagesOrphanCleanupSkipped pagesOrphanCleanupOutcome = iota
pagesOrphanCleanupReconciled
pagesOrphanCleanupReferenced
pagesOrphanCleanupLeaseBusy
pagesOrphanCleanupInvalidMarker
)
// ReconcilePagesOrphanUploads performs one bounded delayed compensation pass.
// Individual candidate failures are counted and logged so they do not prevent
// the scanner from continuing with source checks.
func ReconcilePagesOrphanUploads(
ctx context.Context,
now time.Time,
) (PagesOrphanCleanupSummary, error) {
if now.IsZero() {
now = time.Now()
}
cutoff := now.UTC().Add(-pagesOrphanUploadIsolation)
systemUser := repository.GetSystemUser(ctx)
candidates, err := model.ListPagesOrphanUploadCandidates(ctx, model.PagesOrphanUploadCandidateQuery{
SystemUserID: systemUser.ID,
UploadType: upload.ReservedPagesDeploymentType,
Marker: pagesIngestMarkerV2,
CreatedBefore: cutoff,
})
if err != nil {
return PagesOrphanCleanupSummary{}, err
}
summary := PagesOrphanCleanupSummary{Candidates: len(candidates)}
for index := range candidates {
if err := ctx.Err(); err != nil {
return summary, err
}
candidate := &candidates[index]
marker, err := parsePagesOrphanMarker(candidate.Metadata)
if err != nil {
summary.InvalidMarker++
logger.WarnF(ctx, "[PagesSource] orphan upload marker invalid: upload_id=%d error=%v", candidate.ID, err)
continue
}
outcome, err := reconcilePagesOrphanUploadCandidate(
ctx,
candidate,
marker,
systemUser.ID,
cutoff,
)
if err != nil {
summary.Failed++
logger.WarnF(ctx, "[PagesSource] orphan upload reconciliation failed: upload_id=%d error=%v", candidate.ID, err)
continue
}
summary.add(outcome)
}
return summary, nil
}
func (summary *PagesOrphanCleanupSummary) add(outcome pagesOrphanCleanupOutcome) {
switch outcome {
case pagesOrphanCleanupReconciled:
summary.Reconciled++
case pagesOrphanCleanupReferenced:
summary.Referenced++
case pagesOrphanCleanupLeaseBusy:
summary.LeaseBusy++
case pagesOrphanCleanupInvalidMarker:
summary.InvalidMarker++
default:
summary.Skipped++
}
}
func reconcilePagesOrphanUploadCandidate(
ctx context.Context,
candidate *model.Upload,
marker pagesOrphanMarker,
systemUserID uint64,
cutoff time.Time,
) (pagesOrphanCleanupOutcome, error) {
if candidate == nil || candidate.ID == 0 {
return pagesOrphanCleanupSkipped, nil
}
outcome := pagesOrphanCleanupSkipped
uploadLocked := false
err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
scopeOutcome, proceed, err := lockPagesOrphanCleanupScope(ctx, tx, candidate.ID, marker)
if err != nil {
return err
}
if !proceed {
outcome = scopeOutcome
return nil
}
lockedOutcome, locked, err := reconcileLockedPagesOrphanUpload(
ctx,
tx,
candidate.ID,
marker,
systemUserID,
cutoff,
)
if err != nil {
return err
}
outcome = lockedOutcome
uploadLocked = locked
return nil
})
if err != nil {
return pagesOrphanCleanupSkipped, err
}
if uploadLocked {
// Also heal a prior post-commit cache invalidation interruption when the
// status transition was an idempotent no-op.
upload.InvalidateUploadMetaCache(ctx, candidate.ID)
}
return outcome, nil
}
func lockPagesOrphanCleanupScope(
ctx context.Context,
tx *gorm.DB,
uploadID uint64,
marker pagesOrphanMarker,
) (pagesOrphanCleanupOutcome, bool, error) {
var project model.PagesProject
if _, err := lockOptionalPagesCleanupRecord(tx, &project, "id = ?", marker.ProjectID); err != nil {
return pagesOrphanCleanupSkipped, false, err
}
if marker.SourceID == nil {
return pagesOrphanCleanupSkipped, true, nil
}
var source model.PagesProjectSource
sourceExists, err := lockOptionalPagesCleanupRecord(tx, &source, "id = ?", *marker.SourceID)
if err != nil {
return pagesOrphanCleanupSkipped, false, err
}
if !sourceExists {
return pagesOrphanCleanupSkipped, true, nil
}
if source.ProjectID != marker.ProjectID {
logger.WarnF(ctx,
"[PagesSource] orphan upload source ownership mismatch: upload_id=%d project_id=%d source_id=%d source_project_id=%d",
uploadID,
marker.ProjectID,
*marker.SourceID,
source.ProjectID,
)
return pagesOrphanCleanupInvalidMarker, false, nil
}
var runtime model.PagesProjectSourceRuntime
runtimeExists, err := lockOptionalPagesCleanupRecord(tx, &runtime, "source_id = ?", source.ID)
if err != nil {
return pagesOrphanCleanupSkipped, false, err
}
// Read the real clock only after obtaining the runtime row lock. The scanner
// snapshot time is only an isolation cutoff and may be stale after lock wait.
leaseCheckedAt := time.Now()
if runtimeExists && runtime.LeaseExpiresAt != nil && runtime.LeaseExpiresAt.After(leaseCheckedAt) {
return pagesOrphanCleanupLeaseBusy, false, nil
}
return pagesOrphanCleanupSkipped, true, nil
}
func reconcileLockedPagesOrphanUpload(
ctx context.Context,
tx *gorm.DB,
uploadID uint64,
marker pagesOrphanMarker,
systemUserID uint64,
cutoff time.Time,
) (pagesOrphanCleanupOutcome, bool, error) {
var lockedUpload model.Upload
found, err := lockOptionalPagesCleanupRecord(tx, &lockedUpload, "id = ?", uploadID)
if err != nil || !found {
return pagesOrphanCleanupSkipped, false, err
}
lockedMarker, err := parsePagesOrphanMarker(lockedUpload.Metadata)
if err != nil {
logger.WarnF(ctx, "[PagesSource] orphan upload marker changed or invalid: upload_id=%d error=%v", uploadID, err)
return pagesOrphanCleanupInvalidMarker, true, nil
}
if lockedUpload.Status != model.UploadStatusUsed ||
lockedUpload.UserID != systemUserID ||
lockedUpload.Type != upload.ReservedPagesDeploymentType ||
!lockedUpload.CreatedAt.Before(cutoff) {
return pagesOrphanCleanupSkipped, true, nil
}
if lockedMarker.ProjectID != marker.ProjectID || !sameOptionalPagesSourceID(lockedMarker.SourceID, marker.SourceID) {
logger.WarnF(ctx, "[PagesSource] orphan upload marker changed during reconciliation: upload_id=%d", uploadID)
return pagesOrphanCleanupInvalidMarker, true, nil
}
var references int64
if err := tx.Model(&model.PagesDeployment{}).
Where("upload_id = ?", lockedUpload.ID).
Count(&references).Error; err != nil {
return pagesOrphanCleanupSkipped, true, err
}
if references > 0 {
return pagesOrphanCleanupReferenced, true, nil
}
transitioned, err := upload.RemoveLockedTx(tx, &lockedUpload)
if err != nil {
return pagesOrphanCleanupSkipped, true, err
}
if transitioned {
return pagesOrphanCleanupReconciled, true, nil
}
return pagesOrphanCleanupSkipped, true, nil
}
func lockOptionalPagesCleanupRecord(
tx *gorm.DB,
value any,
query string,
args ...any,
) (bool, error) {
err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).
Where(query, args...).
First(value).Error
if err == nil {
return true, nil
}
if errors.Is(err, gorm.ErrRecordNotFound) {
return false, nil
}
return false, err
}
func parsePagesOrphanMarker(metadata model.UploadMetadata) (pagesOrphanMarker, error) {
if metadata.Extra == nil {
return pagesOrphanMarker{}, errors.New("pages marker metadata missing")
}
marker, ok := metadata.Extra[pagesIngestMarkerKey].(string)
if !ok || marker != pagesIngestMarkerV2 {
return pagesOrphanMarker{}, errors.New("pages marker version invalid")
}
projectID, err := parsePagesOrphanMetadataID(metadata.Extra, pagesProjectIDMetadataKey)
if err != nil {
return pagesOrphanMarker{}, err
}
result := pagesOrphanMarker{ProjectID: projectID}
if _, exists := metadata.Extra[pagesSourceIDMetadataKey]; exists {
sourceID, err := parsePagesOrphanMetadataID(metadata.Extra, pagesSourceIDMetadataKey)
if err != nil {
return pagesOrphanMarker{}, err
}
result.SourceID = &sourceID
}
return result, nil
}
func parsePagesOrphanMetadataID(extra map[string]any, key string) (uint, error) {
raw, exists := extra[key]
if !exists {
return 0, fmt.Errorf("pages marker %s missing", key)
}
value, ok := raw.(string)
if !ok || value == "" {
return 0, fmt.Errorf("pages marker %s must be a decimal string", key)
}
parsed, err := strconv.ParseUint(value, 10, 64)
maxModelID := uint64(^uint(0) >> 1)
if err != nil || parsed == 0 || parsed > maxModelID || strconv.FormatUint(parsed, 10) != value {
return 0, fmt.Errorf("pages marker %s is not a canonical non-zero decimal ID", key)
}
return uint(parsed), nil
}
func sameOptionalPagesSourceID(left, right *uint) bool {
if left == nil || right == nil {
return left == nil && right == nil
}
return *left == *right
}
@@ -0,0 +1,381 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pages
import (
"context"
"errors"
"strconv"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/apps/upload"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"gorm.io/gorm"
"gorm.io/gorm/clause"
)
func TestParsePagesOrphanMarker(t *testing.T) {
tests := []struct {
name string
extra map[string]any
wantProject uint
wantSource uint
wantSourceOK bool
wantErr bool
}{
{
name: "manual source marker",
extra: map[string]any{
pagesIngestMarkerKey: pagesIngestMarkerV2,
pagesProjectIDMetadataKey: "12",
},
wantProject: 12,
},
{
name: "persistent source marker",
extra: map[string]any{
pagesIngestMarkerKey: pagesIngestMarkerV2,
pagesProjectIDMetadataKey: "12",
pagesSourceIDMetadataKey: "34",
},
wantProject: 12,
wantSource: 34,
wantSourceOK: true,
},
{
name: "project ID must be canonical decimal",
extra: map[string]any{
pagesIngestMarkerKey: pagesIngestMarkerV2,
pagesProjectIDMetadataKey: "012",
},
wantErr: true,
},
{
name: "source ID must be a string",
extra: map[string]any{
pagesIngestMarkerKey: pagesIngestMarkerV2,
pagesProjectIDMetadataKey: "12",
pagesSourceIDMetadataKey: float64(34),
},
wantErr: true,
},
{
name: "zero ID rejected",
extra: map[string]any{
pagesIngestMarkerKey: pagesIngestMarkerV2,
pagesProjectIDMetadataKey: "0",
},
wantErr: true,
},
{
name: "wrong marker version rejected",
extra: map[string]any{
pagesIngestMarkerKey: "pages_deployment_v1",
pagesProjectIDMetadataKey: "12",
},
wantErr: true,
},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
got, err := parsePagesOrphanMarker(model.UploadMetadata{Extra: test.extra})
if gotErr := err != nil; gotErr != test.wantErr {
t.Fatalf("parsePagesOrphanMarker(%v) error = %v, want error presence = %t", test.extra, err, test.wantErr)
}
if test.wantErr {
return
}
if got.ProjectID != test.wantProject {
t.Errorf("parsePagesOrphanMarker(%v).ProjectID = %d, want %d", test.extra, got.ProjectID, test.wantProject)
}
if gotSourceOK := got.SourceID != nil; gotSourceOK != test.wantSourceOK {
t.Fatalf("parsePagesOrphanMarker(%v).SourceID presence = %t, want %t", test.extra, gotSourceOK, test.wantSourceOK)
}
if got.SourceID != nil && *got.SourceID != test.wantSource {
t.Errorf("parsePagesOrphanMarker(%v).SourceID = %d, want %d", test.extra, *got.SourceID, test.wantSource)
}
})
}
}
func TestReconcilePagesOrphanUploadsDeletesEligibleUploadOnce(t *testing.T) {
cleanup := setupPagesTestDB(t)
defer cleanup()
ctx := context.Background()
now := time.Now().UTC()
project := createPagesOrphanProject(t, ctx, "eligible-orphan")
candidate := createPagesOrphanUpload(t, ctx, now.Add(-3*time.Hour), project.ID, nil)
if err := upload.RebuildUploadStats(ctx); err != nil {
t.Fatalf("RebuildUploadStats() error = %v, want nil", err)
}
summary, err := ReconcilePagesOrphanUploads(ctx, now)
if err != nil {
t.Fatalf("ReconcilePagesOrphanUploads() error = %v, want nil", err)
}
if summary.Candidates != 1 || summary.Reconciled != 1 || cleanupOutcomeTotal(summary) != 1 {
t.Errorf("ReconcilePagesOrphanUploads() summary = %+v, want one reconciled candidate", summary)
}
assertPagesCleanupUploadStatus(t, ctx, candidate.ID, model.UploadStatusDeleted)
assertPagesCleanupTotalStat(t, ctx, 0)
second, err := ReconcilePagesOrphanUploads(ctx, now.Add(time.Minute))
if err != nil {
t.Fatalf("second ReconcilePagesOrphanUploads() error = %v, want nil", err)
}
if second.Candidates != 0 || cleanupOutcomeTotal(second) != 0 {
t.Errorf("second ReconcilePagesOrphanUploads() summary = %+v, want empty", second)
}
assertPagesCleanupTotalStat(t, ctx, 0)
}
func TestReconcilePagesOrphanUploadsAllowsDeletedProjectAndSource(t *testing.T) {
cleanup := setupPagesTestDB(t)
defer cleanup()
ctx := context.Background()
now := time.Now().UTC()
missingSourceID := uint(9876)
candidate := createPagesOrphanUpload(t, ctx, now.Add(-3*time.Hour), 8765, &missingSourceID)
summary, err := ReconcilePagesOrphanUploads(ctx, now)
if err != nil {
t.Fatalf("ReconcilePagesOrphanUploads() error = %v, want nil", err)
}
if summary.Candidates != 1 || summary.Reconciled != 1 || cleanupOutcomeTotal(summary) != 1 {
t.Errorf("ReconcilePagesOrphanUploads() summary = %+v, want deleted project/source treated as one orphan", summary)
}
assertPagesCleanupUploadStatus(t, ctx, candidate.ID, model.UploadStatusDeleted)
}
func TestReconcilePagesOrphanUploadsSkipsBusyLeaseAndSourceMismatch(t *testing.T) {
t.Run("unexpired source lease", func(t *testing.T) {
cleanup := setupPagesTestDB(t)
defer cleanup()
ctx := context.Background()
realNow := time.Now().UTC()
// A deliberately future scanner snapshot proves lease freshness uses the
// real clock after the runtime lock, not this isolation-cutoff input.
scannerNow := realNow.Add(24 * time.Hour)
project := createPagesOrphanProject(t, ctx, "busy-orphan")
source := createPagesOrphanSource(t, ctx, project.ID)
future := realNow.Add(time.Hour)
if err := db.DB(ctx).Create(&model.PagesProjectSourceRuntime{
SourceID: source.ID,
LeaseToken: "busy-worker",
LeaseExpiresAt: &future,
}).Error; err != nil {
t.Fatalf("create busy source runtime error = %v, want nil", err)
}
candidate := createPagesOrphanUpload(t, ctx, realNow.Add(-3*time.Hour), project.ID, &source.ID)
summary, err := ReconcilePagesOrphanUploads(ctx, scannerNow)
if err != nil {
t.Fatalf("ReconcilePagesOrphanUploads() error = %v, want nil", err)
}
if summary.LeaseBusy != 1 || cleanupOutcomeTotal(summary) != 1 {
t.Errorf("ReconcilePagesOrphanUploads() summary = %+v, want one lease-busy candidate", summary)
}
assertPagesCleanupUploadStatus(t, ctx, candidate.ID, model.UploadStatusUsed)
})
t.Run("source belongs to another project", func(t *testing.T) {
cleanup := setupPagesTestDB(t)
defer cleanup()
ctx := context.Background()
now := time.Now().UTC()
markerProject := createPagesOrphanProject(t, ctx, "marker-project")
actualProject := createPagesOrphanProject(t, ctx, "actual-project")
source := createPagesOrphanSource(t, ctx, actualProject.ID)
candidate := createPagesOrphanUpload(t, ctx, now.Add(-3*time.Hour), markerProject.ID, &source.ID)
summary, err := ReconcilePagesOrphanUploads(ctx, now)
if err != nil {
t.Fatalf("ReconcilePagesOrphanUploads() error = %v, want nil", err)
}
if summary.InvalidMarker != 1 || cleanupOutcomeTotal(summary) != 1 {
t.Errorf("ReconcilePagesOrphanUploads() summary = %+v, want one ownership mismatch", summary)
}
assertPagesCleanupUploadStatus(t, ctx, candidate.ID, model.UploadStatusUsed)
})
}
func TestReconcilePagesOrphanUploadsRejectsMalformedMarker(t *testing.T) {
cleanup := setupPagesTestDB(t)
defer cleanup()
ctx := context.Background()
now := time.Now().UTC()
candidate := createPagesOrphanUpload(t, ctx, now.Add(-3*time.Hour), 1, nil)
metadata := candidate.Metadata
metadata.Extra[pagesProjectIDMetadataKey] = "01"
candidate.Metadata = metadata
if err := db.DB(ctx).Save(candidate).Error; err != nil {
t.Fatalf("seed malformed candidate marker error = %v, want nil", err)
}
summary, err := ReconcilePagesOrphanUploads(ctx, now)
if err != nil {
t.Fatalf("ReconcilePagesOrphanUploads() error = %v, want nil", err)
}
if summary.InvalidMarker != 1 || cleanupOutcomeTotal(summary) != 1 {
t.Errorf("ReconcilePagesOrphanUploads() summary = %+v, want one invalid marker", summary)
}
assertPagesCleanupUploadStatus(t, ctx, candidate.ID, model.UploadStatusUsed)
}
func TestPagesOrphanCleanupAndDeploymentCommitInterleavings(t *testing.T) {
t.Run("deployment reference commits first", func(t *testing.T) {
cleanup := setupPagesTestDB(t)
defer cleanup()
ctx := context.Background()
now := time.Now().UTC()
project := createPagesOrphanProject(t, ctx, "deployment-first")
candidate := createPagesOrphanUpload(t, ctx, now.Add(-3*time.Hour), project.ID, nil)
marker, err := parsePagesOrphanMarker(candidate.Metadata)
if err != nil {
t.Fatalf("parsePagesOrphanMarker() error = %v, want nil", err)
}
if err := db.DB(ctx).Create(&model.PagesDeployment{
ProjectID: project.ID,
DeploymentNumber: 1,
Checksum: "deployment-first",
Status: model.PagesDeploymentStatusUploaded,
UploadID: candidate.ID,
}).Error; err != nil {
t.Fatalf("create deployment reference error = %v, want nil", err)
}
outcome, err := reconcilePagesOrphanUploadCandidate(ctx, candidate, marker, 999, now.Add(-2*time.Hour))
if err != nil {
t.Fatalf("reconcilePagesOrphanUploadCandidate() error = %v, want nil", err)
}
if outcome != pagesOrphanCleanupReferenced {
t.Errorf("reconcilePagesOrphanUploadCandidate() outcome = %d, want %d", outcome, pagesOrphanCleanupReferenced)
}
assertPagesCleanupUploadStatus(t, ctx, candidate.ID, model.UploadStatusUsed)
})
t.Run("cleanup commits first", func(t *testing.T) {
cleanup := setupPagesTestDB(t)
defer cleanup()
ctx := context.Background()
now := time.Now().UTC()
project := createPagesOrphanProject(t, ctx, "cleanup-first")
candidate := createPagesOrphanUpload(t, ctx, now.Add(-3*time.Hour), project.ID, nil)
summary, err := ReconcilePagesOrphanUploads(ctx, now)
if err != nil {
t.Fatalf("ReconcilePagesOrphanUploads() error = %v, want nil", err)
}
if summary.Reconciled != 1 {
t.Fatalf("ReconcilePagesOrphanUploads() summary = %+v, want one reconciled candidate", summary)
}
target := &model.PagesDeployment{ProjectID: project.ID, UploadID: candidate.ID}
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
}
return lockSourceDeploymentUploadsTx(tx, target, upload.IngestResult{}, false)
})
if !errors.Is(err, errSourceFinalFence) {
t.Errorf("final deployment upload lock after cleanup error = %v, want %v", err, errSourceFinalFence)
}
var references int64
if err := db.DB(ctx).Model(&model.PagesDeployment{}).Where("upload_id = ?", candidate.ID).Count(&references).Error; err != nil {
t.Fatalf("count deployment references error = %v, want nil", err)
}
if references != 0 {
t.Errorf("deployment references after cleanup-first interleaving = %d, want 0", references)
}
})
}
func cleanupOutcomeTotal(summary PagesOrphanCleanupSummary) int {
return summary.Reconciled + summary.Referenced + summary.LeaseBusy + summary.InvalidMarker + summary.Skipped + summary.Failed
}
func createPagesOrphanProject(t *testing.T, ctx context.Context, slug string) *model.PagesProject {
t.Helper()
project := &model.PagesProject{Name: slug, Slug: slug, Enabled: true}
if err := db.DB(ctx).Create(project).Error; err != nil {
t.Fatalf("create Pages orphan project %q error = %v, want nil", slug, err)
}
return project
}
func createPagesOrphanSource(t *testing.T, ctx context.Context, projectID uint) *model.PagesProjectSource {
t.Helper()
source := &model.PagesProjectSource{
ProjectID: projectID,
SourceType: PagesSourceTypeRemoteURL,
ConfigVersion: 1,
SourceIdentity: "orphan-source-identity",
}
if err := db.DB(ctx).Create(source).Error; err != nil {
t.Fatalf("create Pages orphan source for project %d error = %v, want nil", projectID, err)
}
return source
}
func createPagesOrphanUpload(
t *testing.T,
ctx context.Context,
createdAt time.Time,
projectID uint,
sourceID *uint,
) *model.Upload {
t.Helper()
extra := map[string]any{
pagesIngestMarkerKey: pagesIngestMarkerV2,
pagesProjectIDMetadataKey: strconv.FormatUint(uint64(projectID), 10),
}
if sourceID != nil {
extra[pagesSourceIDMetadataKey] = strconv.FormatUint(uint64(*sourceID), 10)
}
candidate := &model.Upload{
UserID: 999,
FileName: "site.zip",
FilePath: "pages/orphan-site.zip",
FileSize: 64,
MimeType: "application/zip",
Extension: "zip",
Hash: "orphan-checksum",
Type: upload.ReservedPagesDeploymentType,
Status: model.UploadStatusUsed,
AccessMode: 0,
Metadata: model.UploadMetadata{Extra: extra},
CreatedAt: createdAt,
UpdatedAt: createdAt,
}
if err := db.DB(ctx).Create(candidate).Error; err != nil {
t.Fatalf("create Pages orphan upload error = %v, want nil", err)
}
return candidate
}
func assertPagesCleanupUploadStatus(t *testing.T, ctx context.Context, uploadID uint64, want model.UploadStatus) {
t.Helper()
var got model.Upload
if err := db.DB(ctx).First(&got, uploadID).Error; err != nil {
t.Fatalf("load upload %d error = %v, want nil", uploadID, err)
}
if got.Status != want {
t.Errorf("upload %d status = %q, want %q", uploadID, got.Status, want)
}
}
func assertPagesCleanupTotalStat(t *testing.T, ctx context.Context, want int64) {
t.Helper()
var stat model.UploadStat
if err := db.DB(ctx).Where("dimension = ? AND stat_key = ?", model.UploadStatDimensionTotal, "").First(&stat).Error; err != nil {
t.Fatalf("load total upload stat error = %v, want nil", err)
}
if stat.FileCount != want {
t.Errorf("total upload stat FileCount = %d, want %d", stat.FileCount, want)
}
}
@@ -0,0 +1,549 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pages
import (
"context"
"crypto/sha256"
"crypto/tls"
"encoding/hex"
"errors"
"fmt"
"io"
"math"
"net"
"net/http"
"net/netip"
"net/url"
"os"
"path"
"strings"
"time"
"github.com/Rain-kl/Wavelet/pkg/httppool"
"github.com/Rain-kl/Wavelet/pkg/pagesarchive"
)
const (
// RemoteNetworkPolicyPublic only permits publicly routable targets and
// performs DNS validation again for every connection.
RemoteNetworkPolicyPublic = "public"
// RemoteNetworkPolicyTrustedInternal permits private targets and self-signed
// TLS certificates. It is an explicit administrator trust boundary.
RemoteNetworkPolicyTrustedInternal = "trusted_internal"
remoteSourceDownloadTimeout = 10 * time.Minute
remoteSourceResponseHeaderTimeout = 30 * time.Second
remoteSourceDialTimeout = 30 * time.Second
remoteSourceDialKeepAlive = 30 * time.Second
remoteSourceMaxRedirects = 5
remoteSourceMagicSniffBytes = 512
remoteSourceMaxSafeLabelBytes = 255
remoteSourceFallbackLabel = "package"
remoteSourceUserAgent = "OpenFlare Pages Source/2"
remoteSourceSchemeHTTP = "http"
remoteSourceSchemeHTTPS = "https"
)
type remoteProviderError string
func (providerError remoteProviderError) Error() string {
return string(providerError)
}
const (
errRemoteProviderInvalidPolicy remoteProviderError = "远程来源网络策略无效"
errRemoteProviderInvalidLimit remoteProviderError = "远程来源部署包大小限制无效"
errRemoteProviderBlockedAddress remoteProviderError = "远程来源 public 策略禁止访问非公网地址"
errRemoteProviderResolveFailed remoteProviderError = "远程来源地址解析失败"
errRemoteProviderRedirectLimit remoteProviderError = "远程来源重定向次数超过限制"
errRemoteProviderDownloadFailed remoteProviderError = errPagesPackageURLDownloadFailed
errRemoteProviderTooLarge remoteProviderError = errPagesPackageURLTooLarge
errRemoteProviderEmpty remoteProviderError = errPagesPackageEmpty
errRemoteProviderUnsupported remoteProviderError = errPagesPackageUnsupported
errRemoteProviderCleanupFailed remoteProviderError = "清理远程来源临时文件失败"
)
var remoteSourceNonPublicPrefixes = []netip.Prefix{
// IPv4 special-use, private, link-local, documentation, multicast and
// reserved ranges. A conservative deny list is intentional for SSRF safety.
netip.MustParsePrefix("0.0.0.0/8"),
netip.MustParsePrefix("10.0.0.0/8"),
netip.MustParsePrefix("100.64.0.0/10"),
netip.MustParsePrefix("127.0.0.0/8"),
netip.MustParsePrefix("169.254.0.0/16"),
netip.MustParsePrefix("172.16.0.0/12"),
netip.MustParsePrefix("192.0.0.0/24"),
netip.MustParsePrefix("192.0.2.0/24"),
netip.MustParsePrefix("192.88.99.0/24"),
netip.MustParsePrefix("192.168.0.0/16"),
netip.MustParsePrefix("198.18.0.0/15"),
netip.MustParsePrefix("198.51.100.0/24"),
netip.MustParsePrefix("203.0.113.0/24"),
netip.MustParsePrefix("224.0.0.0/4"),
netip.MustParsePrefix("240.0.0.0/4"),
// IPv6 protocol-assignment, documentation and transition ranges that are
// not acceptable as direct public artifact origins.
netip.MustParsePrefix("2001::/23"),
netip.MustParsePrefix("2001:db8::/32"),
netip.MustParsePrefix("2002::/16"),
netip.MustParsePrefix("3fff::/20"),
}
var remoteSourcePublicIPv6Prefix = netip.MustParsePrefix("2000::/3")
// RemoteSourceRequest describes one immutable Remote URL package fetch.
type RemoteSourceRequest struct {
URL string
NetworkPolicy string
MaxPackageBytes int64
}
// SourceCandidate is a constrained, immutable archive downloaded to a
// provider-owned temporary file. The caller owns the file after a successful
// fetch and must call Cleanup when processing finishes.
type SourceCandidate struct {
TempPath string
Checksum string
PackageSize int64
Format pagesarchive.Format
SafeLabel string
}
// Cleanup removes the candidate temporary file. It is safe to call repeatedly.
func (candidate *SourceCandidate) Cleanup() error {
if candidate == nil || candidate.TempPath == "" {
return nil
}
tempPath := candidate.TempPath
err := os.Remove(tempPath)
if err == nil || errors.Is(err, os.ErrNotExist) {
candidate.TempPath = ""
return nil
}
return errRemoteProviderCleanupFailed
}
type remoteSourceResolver interface {
LookupNetIP(context.Context, string, string) ([]netip.Addr, error)
}
type remoteSourceDependencies struct {
resolver remoteSourceResolver
dialContext func(context.Context, string, string) (net.Conn, error)
createTemp func(string, string) (*os.File, error)
}
// FetchRemoteSource downloads a Remote URL package without writing deployment
// state. Errors are reduced to safe domain messages and never contain the raw
// URL, query, response headers or response body.
func FetchRemoteSource(ctx context.Context, request RemoteSourceRequest) (*SourceCandidate, error) {
dialer := &net.Dialer{
Timeout: remoteSourceDialTimeout,
KeepAlive: remoteSourceDialKeepAlive,
}
dependencies := remoteSourceDependencies{
resolver: net.DefaultResolver,
dialContext: dialer.DialContext,
createTemp: os.CreateTemp,
}
return fetchRemoteSource(ctx, request, dependencies)
}
func fetchRemoteSource(ctx context.Context, request RemoteSourceRequest, dependencies remoteSourceDependencies) (*SourceCandidate, error) {
if request.MaxPackageBytes <= 0 {
return nil, errRemoteProviderInvalidLimit
}
if dependencies.dialContext == nil || dependencies.createTemp == nil {
return nil, errRemoteProviderDownloadFailed
}
policy, err := normalizeRemoteNetworkPolicy(request.NetworkPolicy)
if err != nil {
return nil, err
}
parsed, err := parseRemoteSourceURL(request.URL)
if err != nil {
return nil, err
}
if err := validateRemoteSourceTarget(ctx, parsed, policy, dependencies.resolver); err != nil {
return nil, sanitizeRemoteProviderError(ctx, err)
}
safeLabel, namedFormat := remoteSourceLabel(parsed)
client := newRemoteSourceClient(policy, dependencies)
defer client.CloseIdleConnections()
response, err := requestRemoteSource(ctx, client, parsed)
if err != nil {
return nil, err
}
defer func() { _ = response.Body.Close() }()
if response.StatusCode < http.StatusOK || response.StatusCode >= http.StatusMultipleChoices {
return nil, fmt.Errorf("%w: HTTP %d", errRemoteProviderDownloadFailed, response.StatusCode)
}
if response.ContentLength > request.MaxPackageBytes {
return nil, errRemoteProviderTooLarge
}
tempPath, checksum, packageSize, err := streamRemoteSourcePackage(
response.Body,
request.MaxPackageBytes,
dependencies.createTemp,
)
if err != nil {
return nil, sanitizeRemoteProviderError(ctx, err)
}
format, safeLabel, err := detectRemoteSourceFormat(tempPath, safeLabel, namedFormat)
if err != nil {
if removeErr := os.Remove(tempPath); removeErr != nil && !errors.Is(removeErr, os.ErrNotExist) {
return nil, errRemoteProviderCleanupFailed
}
return nil, err
}
return &SourceCandidate{
TempPath: tempPath,
Checksum: checksum,
PackageSize: packageSize,
Format: format,
SafeLabel: safeLabel,
}, nil
}
func normalizeRemoteNetworkPolicy(policy string) (string, error) {
switch strings.TrimSpace(policy) {
case "", RemoteNetworkPolicyPublic:
return RemoteNetworkPolicyPublic, nil
case RemoteNetworkPolicyTrustedInternal:
return RemoteNetworkPolicyTrustedInternal, nil
default:
return "", errRemoteProviderInvalidPolicy
}
}
func newRemoteSourceClient(policy string, dependencies remoteSourceDependencies) *http.Client {
tlsConfig := &tls.Config{MinVersion: tls.VersionTLS12}
dialContext := dependencies.dialContext
if policy == RemoteNetworkPolicyPublic {
dialContext = newPublicRemoteSourceDialer(dependencies.resolver, dependencies.dialContext)
} else {
// trusted_internal is an explicit administrator-selected boundary for
// private artifact services using an internal CA or self-signed cert.
tlsConfig.InsecureSkipVerify = true //nolint:gosec // required trusted_internal semantics
}
client := &http.Client{
Timeout: remoteSourceDownloadTimeout,
Transport: httppool.NewTransport(httppool.TransportOptions{
Proxy: nil,
DialContext: dialContext,
TLSClientConfig: tlsConfig,
ResponseHeaderTimeout: remoteSourceResponseHeaderTimeout,
TraceFilter: remoteSourceTraceFilter,
}),
}
client.CheckRedirect = func(next *http.Request, previous []*http.Request) error {
if len(previous) > remoteSourceMaxRedirects {
return errRemoteProviderRedirectLimit
}
stripRemoteSourceRedirectHeaders(next)
if err := validateRemoteSourceTarget(next.Context(), next.URL, policy, dependencies.resolver); err != nil {
return err
}
applyRemoteSourceHeaders(next)
return nil
}
return client
}
func requestRemoteSource(ctx context.Context, client *http.Client, parsed *url.URL) (*http.Response, error) {
request, err := http.NewRequestWithContext(ctx, http.MethodGet, parsed.String(), nil)
if err != nil {
return nil, errors.New(errPagesSourceRemoteURLInvalid)
}
applyRemoteSourceHeaders(request)
response, err := client.Do(request) //nolint:gosec // scheme and every dial target are validated above
if err == nil {
return response, nil
}
if response != nil && response.Body != nil {
_ = response.Body.Close()
}
return nil, sanitizeRemoteProviderError(ctx, err)
}
func applyRemoteSourceHeaders(request *http.Request) {
request.Header.Set("User-Agent", remoteSourceUserAgent)
request.Header.Set("Accept", "application/octet-stream,application/zip,application/x-tar,application/gzip,*/*;q=0.1")
// Preserve the artifact bytes exactly as stored. Automatic HTTP gzip
// decompression would change the checksum, size and archive format.
request.Header.Set("Accept-Encoding", "identity")
}
func stripRemoteSourceRedirectHeaders(request *http.Request) {
request.Header.Del("Authorization")
request.Header.Del("Cookie")
request.Header.Del("Proxy-Authorization")
request.Header.Del("Referer")
}
func remoteSourceTraceFilter(request *http.Request) bool {
// otelhttp records url.full. Signed query strings must never enter traces.
return request.URL == nil || request.URL.RawQuery == ""
}
func validateRemoteSourceTarget(ctx context.Context, target *url.URL, policy string, resolver remoteSourceResolver) error {
if target == nil || target.User != nil || target.Fragment != "" || target.Opaque != "" {
return errors.New(errPagesSourceRemoteURLInvalid)
}
scheme := strings.ToLower(strings.TrimSpace(target.Scheme))
if (scheme != remoteSourceSchemeHTTP && scheme != remoteSourceSchemeHTTPS) || strings.TrimSpace(target.Hostname()) == "" {
return errors.New(errPagesSourceRemoteURLInvalid)
}
if policy != RemoteNetworkPolicyPublic {
return nil
}
_, err := resolvePublicRemoteSourceIPs(ctx, resolver, target.Hostname())
return err
}
func newPublicRemoteSourceDialer(
resolver remoteSourceResolver,
directDial func(context.Context, string, string) (net.Conn, error),
) func(context.Context, string, string) (net.Conn, error) {
return func(ctx context.Context, network string, address string) (net.Conn, error) {
host, port, err := net.SplitHostPort(address)
if err != nil {
return nil, errRemoteProviderDownloadFailed
}
addresses, err := resolvePublicRemoteSourceIPs(ctx, resolver, host)
if err != nil {
return nil, err
}
for _, address := range addresses {
if !remoteSourceIPMatchesNetwork(address, network) {
continue
}
connection, dialErr := directDial(ctx, network, net.JoinHostPort(address.String(), port))
if dialErr == nil {
return connection, nil
}
}
return nil, errRemoteProviderDownloadFailed
}
}
func resolvePublicRemoteSourceIPs(ctx context.Context, resolver remoteSourceResolver, host string) ([]netip.Addr, error) {
if strings.Contains(host, "%") {
return nil, errRemoteProviderBlockedAddress
}
if literal, parseErr := netip.ParseAddr(host); parseErr == nil {
if !isPublicRemoteSourceIP(literal) {
return nil, errRemoteProviderBlockedAddress
}
return []netip.Addr{literal}, nil
}
if resolver == nil {
return nil, errRemoteProviderResolveFailed
}
addresses, err := resolver.LookupNetIP(ctx, "ip", host)
if err != nil || len(addresses) == 0 {
return nil, errRemoteProviderResolveFailed
}
for _, address := range addresses {
if !isPublicRemoteSourceIP(address) {
return nil, errRemoteProviderBlockedAddress
}
}
return addresses, nil
}
func isPublicRemoteSourceIP(address netip.Addr) bool {
if !address.IsValid() || address.Zone() != "" {
return false
}
address = address.Unmap()
if !address.IsGlobalUnicast() {
return false
}
if address.Is6() && !remoteSourcePublicIPv6Prefix.Contains(address) {
return false
}
for _, prefix := range remoteSourceNonPublicPrefixes {
if prefix.Contains(address) {
return false
}
}
return true
}
func remoteSourceIPMatchesNetwork(address netip.Addr, network string) bool {
switch network {
case "tcp4":
return address.Unmap().Is4()
case "tcp6":
return address.Unmap().Is6()
default:
return true
}
}
func streamRemoteSourcePackage(
body io.Reader,
maxPackageBytes int64,
createTemp func(string, string) (*os.File, error),
) (tempPath string, checksum string, packageSize int64, err error) {
if createTemp == nil {
return "", "", 0, errRemoteProviderDownloadFailed
}
tempFile, err := createTemp("", "openflare-pages-source-*")
if err != nil {
return "", "", 0, errRemoteProviderDownloadFailed
}
createdTempPath := tempFile.Name()
tempPath = createdTempPath
defer func() {
closeErr := tempFile.Close()
if err == nil && closeErr != nil {
err = errRemoteProviderDownloadFailed
}
if err != nil {
if removeErr := os.Remove(createdTempPath); removeErr != nil && !errors.Is(removeErr, os.ErrNotExist) {
err = errRemoteProviderCleanupFailed
}
}
}()
hasher := sha256.New()
readLimit := maxPackageBytes
if readLimit < math.MaxInt64 {
readLimit++
}
packageSize, err = io.Copy(io.MultiWriter(tempFile, hasher), io.LimitReader(body, readLimit))
if err != nil {
return "", "", 0, errRemoteProviderDownloadFailed
}
if packageSize > maxPackageBytes {
return "", "", 0, errRemoteProviderTooLarge
}
if packageSize == 0 {
return "", "", 0, errRemoteProviderEmpty
}
checksum = hex.EncodeToString(hasher.Sum(nil))
return tempPath, checksum, packageSize, nil
}
func detectRemoteSourceFormat(
tempPath string,
safeLabel string,
namedFormat pagesarchive.Format,
) (pagesarchive.Format, string, error) {
if namedFormat != "" {
return namedFormat, safeLabel, nil
}
tempFile, err := os.Open(tempPath) //nolint:gosec // path is a provider-created temporary file
if err != nil {
return "", safeLabel, errRemoteProviderDownloadFailed
}
defer func() { _ = tempFile.Close() }()
head := make([]byte, remoteSourceMagicSniffBytes)
readBytes, readErr := io.ReadFull(tempFile, head)
if readErr != nil && !errors.Is(readErr, io.EOF) && !errors.Is(readErr, io.ErrUnexpectedEOF) {
return "", safeLabel, errRemoteProviderDownloadFailed
}
format, ok := pagesarchive.DetectFormatFromBytes(head[:readBytes])
if !ok {
return "", safeLabel, errRemoteProviderUnsupported
}
return format, appendRemoteSourceLabelExtension(safeLabel, format), nil
}
func remoteSourceLabel(parsed *url.URL) (string, pagesarchive.Format) {
baseName := path.Base(parsed.Path)
if baseName == "" || baseName == "." || baseName == "/" {
baseName = remoteSourceFallbackLabel
}
safeLabel := sanitizeRemoteSourceLabel(baseName)
format, _ := pagesarchive.DetectFormatFromName(safeLabel)
return limitRemoteSourceLabel(safeLabel, format), format
}
func sanitizeRemoteSourceLabel(label string) string {
var builder strings.Builder
lastReplacement := false
for _, character := range label {
if isRemoteSourceLabelCharacter(character) {
builder.WriteRune(character)
lastReplacement = false
continue
}
if !lastReplacement {
builder.WriteByte('-')
lastReplacement = true
}
}
safeLabel := strings.TrimSpace(builder.String())
if safeLabel == "" || strings.Trim(safeLabel, "._-") == "" {
return remoteSourceFallbackLabel
}
return safeLabel
}
func isRemoteSourceLabelCharacter(character rune) bool {
return character >= 'a' && character <= 'z' ||
character >= 'A' && character <= 'Z' ||
character >= '0' && character <= '9' ||
character == '.' || character == '-' || character == '_'
}
func limitRemoteSourceLabel(label string, format pagesarchive.Format) string {
if len(label) <= remoteSourceMaxSafeLabelBytes {
return label
}
if format == "" {
return strings.TrimRight(label[:remoteSourceMaxSafeLabelBytes], ".-_")
}
extension := "." + pagesarchive.Extension(format)
prefixLength := remoteSourceMaxSafeLabelBytes - len(extension)
prefix := strings.TrimRight(label[:prefixLength], ".-_")
if prefix == "" {
prefix = remoteSourceFallbackLabel
}
return prefix + extension
}
func appendRemoteSourceLabelExtension(label string, format pagesarchive.Format) string {
extension := "." + pagesarchive.Extension(format)
maxPrefixLength := remoteSourceMaxSafeLabelBytes - len(extension)
if len(label) > maxPrefixLength {
label = strings.TrimRight(label[:maxPrefixLength], ".-_")
}
if label == "" {
label = remoteSourceFallbackLabel
}
return label + extension
}
func sanitizeRemoteProviderError(ctx context.Context, err error) error {
if ctxErr := ctx.Err(); ctxErr != nil {
return fmt.Errorf("%w: %w", errRemoteProviderDownloadFailed, ctxErr)
}
for _, safeError := range []error{
errRemoteProviderInvalidPolicy,
errRemoteProviderInvalidLimit,
errRemoteProviderBlockedAddress,
errRemoteProviderResolveFailed,
errRemoteProviderRedirectLimit,
errRemoteProviderTooLarge,
errRemoteProviderEmpty,
errRemoteProviderUnsupported,
errRemoteProviderCleanupFailed,
} {
if errors.Is(err, safeError) {
return safeError
}
}
return errRemoteProviderDownloadFailed
}
@@ -0,0 +1,459 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pages
import (
"archive/zip"
"bytes"
"context"
"crypto/sha256"
"encoding/hex"
"errors"
"io"
"log"
"net"
"net/http"
"net/http/httptest"
"net/netip"
"os"
"strconv"
"strings"
"sync/atomic"
"testing"
)
type remoteSourceResolverFunc func(context.Context, string, string) ([]netip.Addr, error)
func (function remoteSourceResolverFunc) LookupNetIP(
ctx context.Context,
network string,
host string,
) ([]netip.Addr, error) {
return function(ctx, network, host)
}
func TestFetchRemoteSourceTrustedInternalSelfSignedAndSafeLabel(t *testing.T) {
packageBytes := makeRemoteSourceZIP(t)
server := httptest.NewTLSServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
if request.URL.Query().Get("token") != "source-secret" {
t.Error("signed query did not reach the artifact server")
}
if request.Header.Get("Accept-Encoding") != "identity" {
t.Error("artifact request must disable automatic HTTP decompression")
}
writer.Header().Set("Content-Disposition", `attachment; filename="redirected.tar.gz"`)
_, _ = writer.Write(packageBytes)
}))
server.Config.ErrorLog = log.New(io.Discard, "", 0)
t.Cleanup(server.Close)
candidate, err := FetchRemoteSource(t.Context(), RemoteSourceRequest{
URL: server.URL + "/original/site.zip?token=source-secret",
NetworkPolicy: RemoteNetworkPolicyTrustedInternal,
MaxPackageBytes: int64(len(packageBytes) + 1),
})
if err != nil {
t.Fatalf("FetchRemoteSource() error = %v", err)
}
if candidate.Format != "zip" {
t.Fatalf("Format = %q, want zip", candidate.Format)
}
if candidate.SafeLabel != "site.zip" {
t.Fatalf("SafeLabel = %q, want original path basename", candidate.SafeLabel)
}
if candidate.PackageSize != int64(len(packageBytes)) {
t.Fatalf("PackageSize = %d, want %d", candidate.PackageSize, len(packageBytes))
}
wantChecksum := sha256.Sum256(packageBytes)
if candidate.Checksum != hex.EncodeToString(wantChecksum[:]) {
t.Fatalf("Checksum = %q, want SHA-256", candidate.Checksum)
}
downloaded, err := os.ReadFile(candidate.TempPath) //nolint:gosec // provider-owned test temp file
if err != nil {
t.Fatalf("ReadFile() error = %v", err)
}
if !bytes.Equal(downloaded, packageBytes) {
t.Fatal("downloaded package differs from response body")
}
tempPath := candidate.TempPath
if err := candidate.Cleanup(); err != nil {
t.Fatalf("Cleanup() error = %v", err)
}
if err := candidate.Cleanup(); err != nil {
t.Fatalf("second Cleanup() error = %v", err)
}
if _, err := os.Stat(tempPath); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("temporary file still exists: %v", err)
}
}
func TestFetchRemoteSourceKeepsOriginalLabelAcrossRedirect(t *testing.T) {
packageBytes := makeRemoteSourceZIP(t)
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
if request.URL.Path == "/original/site.zip" {
writer.Header().Set("Location", "/delivery/final.tar.gz?token=redirect-secret")
writer.WriteHeader(http.StatusFound)
return
}
if request.Header.Get("Referer") != "" {
t.Error("redirect must not forward a signed source URL as Referer")
}
writer.Header().Set("Content-Disposition", `attachment; filename="response.7z"`)
_, _ = writer.Write(packageBytes)
}))
t.Cleanup(server.Close)
candidate, err := FetchRemoteSource(t.Context(), RemoteSourceRequest{
URL: server.URL + "/original/site.zip?token=initial-secret",
NetworkPolicy: RemoteNetworkPolicyTrustedInternal,
MaxPackageBytes: int64(len(packageBytes) + 1),
})
if err != nil {
t.Fatalf("FetchRemoteSource() error = %v", err)
}
defer func() { _ = candidate.Cleanup() }()
if candidate.SafeLabel != "site.zip" || candidate.Format != "zip" {
t.Fatalf("candidate = label %q format %q, want original site.zip", candidate.SafeLabel, candidate.Format)
}
}
func TestFetchRemoteSourcePublicRejectsNonPublicAddresses(t *testing.T) {
tests := []string{
"http://127.0.0.1/site.zip?token=loopback-secret",
"http://[::1]/site.zip?token=ipv6-secret",
"http://100.64.0.1/site.zip?token=cgnat-secret",
"http://192.0.2.1/site.zip?token=documentation-secret",
}
for _, rawURL := range tests {
t.Run(rawURL, func(t *testing.T) {
_, err := FetchRemoteSource(t.Context(), RemoteSourceRequest{
URL: rawURL,
NetworkPolicy: RemoteNetworkPolicyPublic,
MaxPackageBytes: 1024,
})
if !errors.Is(err, errRemoteProviderBlockedAddress) {
t.Fatalf("FetchRemoteSource() error = %v, want blocked address", err)
}
assertRemoteSourceErrorRedacted(t, err, rawURL, "secret", "token=")
})
}
}
func TestFetchRemoteSourcePublicDialsValidatedIP(t *testing.T) {
packageBytes := makeRemoteSourceZIP(t)
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
_, _ = writer.Write(packageBytes)
}))
t.Cleanup(server.Close)
var dialedAddress string
dependencies := mappedRemoteSourceDependencies(server.Listener.Addr().String(), staticPublicRemoteSourceResolver())
dependencies.dialContext = func(ctx context.Context, network string, address string) (net.Conn, error) {
dialedAddress = address
dialer := &net.Dialer{}
return dialer.DialContext(ctx, network, server.Listener.Addr().String())
}
candidate, err := fetchRemoteSource(t.Context(), RemoteSourceRequest{
URL: "http://artifact.example/site.zip",
NetworkPolicy: RemoteNetworkPolicyPublic,
MaxPackageBytes: int64(len(packageBytes) + 1),
}, dependencies)
if err != nil {
t.Fatalf("fetchRemoteSource() error = %v", err)
}
defer func() { _ = candidate.Cleanup() }()
if dialedAddress != "93.184.216.34:80" {
t.Fatalf("direct dial address = %q, want validated IP", dialedAddress)
}
}
func TestFetchRemoteSourcePublicRejectsPrivateRedirect(t *testing.T) {
var requestCount atomic.Int32
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
requestCount.Add(1)
writer.Header().Set("Location", "http://127.0.0.1/private.zip?token=redirect-secret")
writer.WriteHeader(http.StatusFound)
}))
t.Cleanup(server.Close)
dependencies := mappedRemoteSourceDependencies(server.Listener.Addr().String(), staticPublicRemoteSourceResolver())
rawURL := "http://artifact.example/start.zip?token=initial-secret"
_, err := fetchRemoteSource(t.Context(), RemoteSourceRequest{
URL: rawURL,
NetworkPolicy: RemoteNetworkPolicyPublic,
MaxPackageBytes: 1024,
}, dependencies)
if !errors.Is(err, errRemoteProviderBlockedAddress) {
t.Fatalf("fetchRemoteSource() error = %v, want blocked redirect", err)
}
if requestCount.Load() != 1 {
t.Fatalf("request count = %d, private redirect must not be requested", requestCount.Load())
}
assertRemoteSourceErrorRedacted(t, err, rawURL, "initial-secret", "redirect-secret", "token=")
}
func TestFetchRemoteSourcePublicRejectsDNSRebinding(t *testing.T) {
var lookupCount atomic.Int32
resolver := remoteSourceResolverFunc(func(context.Context, string, string) ([]netip.Addr, error) {
if lookupCount.Add(1) == 1 {
return []netip.Addr{netip.MustParseAddr("93.184.216.34")}, nil
}
return []netip.Addr{netip.MustParseAddr("127.0.0.1")}, nil
})
var dialCount atomic.Int32
dependencies := remoteSourceDependencies{
resolver: resolver,
dialContext: func(context.Context, string, string) (net.Conn, error) {
dialCount.Add(1)
return nil, errors.New("unexpected dial")
},
createTemp: os.CreateTemp,
}
rawURL := "http://rebind.example/site.zip?signature=dns-secret"
_, err := fetchRemoteSource(t.Context(), RemoteSourceRequest{
URL: rawURL,
NetworkPolicy: RemoteNetworkPolicyPublic,
MaxPackageBytes: 1024,
}, dependencies)
if !errors.Is(err, errRemoteProviderBlockedAddress) {
t.Fatalf("fetchRemoteSource() error = %v, want DNS rebinding rejection", err)
}
if lookupCount.Load() != 2 {
t.Fatalf("DNS lookup count = %d, want preflight plus dial validation", lookupCount.Load())
}
if dialCount.Load() != 0 {
t.Fatalf("direct dial count = %d, rebound address must not be dialed", dialCount.Load())
}
assertRemoteSourceErrorRedacted(t, err, rawURL, "dns-secret", "signature=")
}
func TestFetchRemoteSourcePublicRejectsSelfSignedTLS(t *testing.T) {
packageBytes := makeRemoteSourceZIP(t)
server := httptest.NewTLSServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
_, _ = writer.Write(packageBytes)
}))
server.Config.ErrorLog = log.New(io.Discard, "", 0)
t.Cleanup(server.Close)
dependencies := mappedRemoteSourceDependencies(server.Listener.Addr().String(), staticPublicRemoteSourceResolver())
rawURL := "https://artifact.example/site.zip?signature=tls-secret"
_, err := fetchRemoteSource(t.Context(), RemoteSourceRequest{
URL: rawURL,
NetworkPolicy: RemoteNetworkPolicyPublic,
MaxPackageBytes: int64(len(packageBytes) + 1),
}, dependencies)
if !errors.Is(err, errRemoteProviderDownloadFailed) {
t.Fatalf("fetchRemoteSource() error = %v, want strict TLS failure", err)
}
assertRemoteSourceErrorRedacted(t, err, rawURL, "tls-secret", "signature=")
}
func TestFetchRemoteSourceRejectsChunkedBodyOverLimitAndCleansTemp(t *testing.T) {
const maxPackageBytes = int64(64)
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
_, _ = writer.Write(bytes.Repeat([]byte{'x'}, int(maxPackageBytes)))
if flusher, ok := writer.(http.Flusher); ok {
flusher.Flush()
}
_, _ = writer.Write([]byte("overflow"))
}))
t.Cleanup(server.Close)
tempDir := t.TempDir()
dependencies := defaultRemoteSourceDependenciesForTest()
dependencies.createTemp = func(_ string, pattern string) (*os.File, error) {
return os.CreateTemp(tempDir, pattern)
}
rawURL := server.URL + "/site.zip?token=chunk-secret"
_, err := fetchRemoteSource(t.Context(), RemoteSourceRequest{
URL: rawURL,
NetworkPolicy: RemoteNetworkPolicyTrustedInternal,
MaxPackageBytes: maxPackageBytes,
}, dependencies)
if !errors.Is(err, errRemoteProviderTooLarge) {
t.Fatalf("fetchRemoteSource() error = %v, want actual stream limit", err)
}
assertRemoteSourceTempDirEmpty(t, tempDir)
assertRemoteSourceErrorRedacted(t, err, rawURL, "chunk-secret", "token=")
}
func TestFetchRemoteSourceRejectsContentLengthBeforeCreatingTemp(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
writer.Header().Set("Content-Length", "4096")
writer.WriteHeader(http.StatusOK)
}))
t.Cleanup(server.Close)
var createCount atomic.Int32
dependencies := defaultRemoteSourceDependenciesForTest()
dependencies.createTemp = func(directory string, pattern string) (*os.File, error) {
createCount.Add(1)
return os.CreateTemp(directory, pattern)
}
_, err := fetchRemoteSource(t.Context(), RemoteSourceRequest{
URL: server.URL + "/site.zip",
NetworkPolicy: RemoteNetworkPolicyTrustedInternal,
MaxPackageBytes: 1024,
}, dependencies)
if !errors.Is(err, errRemoteProviderTooLarge) {
t.Fatalf("fetchRemoteSource() error = %v, want Content-Length rejection", err)
}
if createCount.Load() != 0 {
t.Fatalf("CreateTemp called %d times before Content-Length rejection", createCount.Load())
}
}
func TestFetchRemoteSourceSniffsAtLeast512BytesForTar(t *testing.T) {
packageBytes := make([]byte, remoteSourceMagicSniffBytes)
copy(packageBytes[257:], []byte("ustar"))
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
_, _ = writer.Write(packageBytes)
}))
t.Cleanup(server.Close)
candidate, err := FetchRemoteSource(t.Context(), RemoteSourceRequest{
URL: server.URL + "/download",
NetworkPolicy: RemoteNetworkPolicyTrustedInternal,
MaxPackageBytes: int64(len(packageBytes) + 1),
})
if err != nil {
t.Fatalf("FetchRemoteSource() error = %v", err)
}
defer func() { _ = candidate.Cleanup() }()
if candidate.Format != "tar" {
t.Fatalf("Format = %q, want tar detected at byte 257", candidate.Format)
}
if candidate.SafeLabel != "download.tar" {
t.Fatalf("SafeLabel = %q, want download.tar", candidate.SafeLabel)
}
}
func TestFetchRemoteSourceRedactsURLHeadersAndBodyFromErrors(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
writer.Header().Set("X-Artifact-Secret", "header-secret")
_, _ = writer.Write([]byte("response-body-secret"))
}))
t.Cleanup(server.Close)
tempDir := t.TempDir()
dependencies := defaultRemoteSourceDependenciesForTest()
dependencies.createTemp = func(_ string, pattern string) (*os.File, error) {
return os.CreateTemp(tempDir, pattern)
}
rawURL := server.URL + "/download?token=query-secret"
_, err := fetchRemoteSource(t.Context(), RemoteSourceRequest{
URL: rawURL,
NetworkPolicy: RemoteNetworkPolicyTrustedInternal,
MaxPackageBytes: 1024,
}, dependencies)
if !errors.Is(err, errRemoteProviderUnsupported) {
t.Fatalf("fetchRemoteSource() error = %v, want unsupported archive", err)
}
assertRemoteSourceTempDirEmpty(t, tempDir)
assertRemoteSourceErrorRedacted(
t,
err,
rawURL,
"query-secret",
"header-secret",
"response-body-secret",
"token=",
)
}
func TestFetchRemoteSourceAllowsFiveRedirectsOnly(t *testing.T) {
packageBytes := makeRemoteSourceZIP(t)
var requestCount atomic.Int32
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
requestCount.Add(1)
redirectNumber, _ := strconv.Atoi(strings.TrimPrefix(request.URL.Path, "/"))
if redirectNumber < remoteSourceMaxRedirects+1 {
writer.Header().Set("Location", "/"+strconv.Itoa(redirectNumber+1))
writer.WriteHeader(http.StatusFound)
return
}
_, _ = writer.Write(packageBytes)
}))
t.Cleanup(server.Close)
dependencies := mappedRemoteSourceDependencies(server.Listener.Addr().String(), staticPublicRemoteSourceResolver())
_, err := fetchRemoteSource(t.Context(), RemoteSourceRequest{
URL: "http://artifact.example/0",
NetworkPolicy: RemoteNetworkPolicyPublic,
MaxPackageBytes: int64(len(packageBytes) + 1),
}, dependencies)
if !errors.Is(err, errRemoteProviderRedirectLimit) {
t.Fatalf("fetchRemoteSource() error = %v, want redirect limit", err)
}
if requestCount.Load() != remoteSourceMaxRedirects+1 {
t.Fatalf("request count = %d, want initial plus five redirects", requestCount.Load())
}
}
func staticPublicRemoteSourceResolver() remoteSourceResolver {
return remoteSourceResolverFunc(func(context.Context, string, string) ([]netip.Addr, error) {
return []netip.Addr{netip.MustParseAddr("93.184.216.34")}, nil
})
}
func defaultRemoteSourceDependenciesForTest() remoteSourceDependencies {
dialer := &net.Dialer{}
return remoteSourceDependencies{
resolver: net.DefaultResolver,
dialContext: dialer.DialContext,
createTemp: os.CreateTemp,
}
}
func mappedRemoteSourceDependencies(targetAddress string, resolver remoteSourceResolver) remoteSourceDependencies {
dialer := &net.Dialer{}
return remoteSourceDependencies{
resolver: resolver,
dialContext: func(ctx context.Context, network string, _ string) (net.Conn, error) {
return dialer.DialContext(ctx, network, targetAddress)
},
createTemp: os.CreateTemp,
}
}
func makeRemoteSourceZIP(t *testing.T) []byte {
t.Helper()
var buffer bytes.Buffer
archive := zip.NewWriter(&buffer)
file, err := archive.Create("index.html")
if err != nil {
t.Fatalf("zip.Create() error = %v", err)
}
if _, err := file.Write([]byte("<h1>OpenFlare</h1>")); err != nil {
t.Fatalf("zip entry Write() error = %v", err)
}
if err := archive.Close(); err != nil {
t.Fatalf("zip.Close() error = %v", err)
}
return buffer.Bytes()
}
func assertRemoteSourceTempDirEmpty(t *testing.T, directory string) {
t.Helper()
entries, err := os.ReadDir(directory)
if err != nil {
t.Fatalf("ReadDir() error = %v", err)
}
if len(entries) != 0 {
t.Fatalf("temporary directory contains %d leaked files", len(entries))
}
}
func assertRemoteSourceErrorRedacted(t *testing.T, err error, sensitiveValues ...string) {
t.Helper()
if err == nil {
t.Fatal("expected an error")
}
message := err.Error()
for _, sensitiveValue := range sensitiveValues {
if sensitiveValue != "" && strings.Contains(message, sensitiveValue) {
t.Fatalf("error %q contains sensitive value %q", message, sensitiveValue)
}
}
}
@@ -0,0 +1,371 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pages
import (
"context"
"crypto/rand"
"encoding/hex"
"errors"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"gorm.io/gorm"
"gorm.io/gorm/clause"
)
const (
pagesSourceCheckLeaseDuration = 2 * time.Minute
pagesSourceSyncLeaseDuration = 15 * time.Minute
sourceLeaseTokenBytes = 32
sourceRuntimeErrorMaxBytes = 512
sourceRevisionHexLength = 64
sourceColumnAutoUpdateEnabled = "auto_update_enabled"
sourceColumnConfigVersion = "config_version"
sourceRuntimeColumnSyncStatus = "sync_status"
sourceRuntimeColumnLastError = "last_error"
sourceRuntimeColumnLastCheckedAt = "last_checked_at"
sourceRuntimeColumnNextCheckAt = "next_check_at"
sourceRuntimeColumnLeaseToken = "lease_token"
sourceRuntimeColumnLeaseExpiresAt = "lease_expires_at"
pagesDeploymentColumnStatus = "status"
)
type sourceLeaseOutcome string
const (
sourceLeaseAcquired sourceLeaseOutcome = "acquired"
sourceLeaseBusy sourceLeaseOutcome = "busy"
sourceLeaseStale sourceLeaseOutcome = "stale"
)
// sourceExecutionSnapshot captures every mutable value that can affect archive
// validation or the atomic activation decision. The queued payload deliberately
// does not carry project content configuration.
type sourceExecutionSnapshot struct {
ProjectID uint
SourceID uint
SourceConfigVersion int
ContentConfigVersion int
SourceType string
SourceIdentity string
RemoteURL string
RemoteNetworkPolicy string
GitHubRepository string
ReleaseSelector string
ReleaseTag string
AssetName string
AutoUpdateEnabled bool
CheckIntervalMinutes int
ETag string
LastSeenRevision string
LastSeenDetail string
LastAppliedRevision string
LastAppliedDetail string
RootDir string
EntryFile string
LeaseToken string
LeaseExpiresAt time.Time
}
func acquireSourceLease(
ctx context.Context,
sourceID uint,
expectedConfigVersion int,
action string,
) (*sourceExecutionSnapshot, sourceLeaseOutcome, error) {
leaseDuration, status, err := sourceLeaseParameters(action)
if err != nil {
return nil, sourceLeaseStale, err
}
token, err := newSourceLeaseToken()
if err != nil {
return nil, sourceLeaseStale, err
}
now := time.Now()
expiresAt := now.Add(leaseDuration)
result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).
Where("source_id = ?", sourceID).
Where("lease_expires_at IS NULL OR lease_expires_at <= ?", now).
Where(
"EXISTS (SELECT 1 FROM of_pages_project_sources source WHERE source.id = ? AND source.config_version = ?)",
sourceID,
expectedConfigVersion,
).
Updates(map[string]any{
sourceRuntimeColumnLeaseToken: token,
sourceRuntimeColumnLeaseExpiresAt: expiresAt,
sourceRuntimeColumnSyncStatus: status,
sourceRuntimeColumnLastError: "",
})
if result.Error != nil {
return nil, sourceLeaseStale, result.Error
}
if result.RowsAffected == 0 {
outcome, inspectErr := inspectSourceLeaseMiss(ctx, sourceID, expectedConfigVersion, now)
return nil, outcome, inspectErr
}
snapshot, err := loadSourceExecutionSnapshot(ctx, sourceID, token)
if err != nil {
if errors.Is(err, errSourceLeaseSnapshotStale) {
return nil, sourceLeaseStale, nil
}
return nil, sourceLeaseStale, err
}
return snapshot, sourceLeaseAcquired, nil
}
var errSourceLeaseSnapshotStale = errors.New("source lease snapshot stale")
func loadSourceExecutionSnapshot(
ctx context.Context,
sourceID uint,
token string,
) (*sourceExecutionSnapshot, error) {
var snapshot sourceExecutionSnapshot
err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
var source model.PagesProjectSource
if err := tx.Where("id = ?", sourceID).First(&source).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return errSourceLeaseSnapshotStale
}
return err
}
var project model.PagesProject
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).
First(&project, source.ProjectID).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return errSourceLeaseSnapshotStale
}
return err
}
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).
Where("id = ?", sourceID).
First(&source).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return errSourceLeaseSnapshotStale
}
return err
}
var runtime model.PagesProjectSourceRuntime
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).
Where("source_id = ?", source.ID).
First(&runtime).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return errSourceLeaseSnapshotStale
}
return err
}
// 必须在 runtime 行锁拿到后重新取时间,避免锁等待跨过
// lease expiry 时仍使用事务开始前的旧时间继续执行。
now := time.Now()
if runtime.LeaseToken != token || runtime.LeaseExpiresAt == nil || !runtime.LeaseExpiresAt.After(now) {
return errSourceLeaseSnapshotStale
}
snapshot = sourceExecutionSnapshot{
ProjectID: project.ID,
SourceID: source.ID,
SourceConfigVersion: source.ConfigVersion,
ContentConfigVersion: project.ContentConfigVersion,
SourceType: source.SourceType,
SourceIdentity: source.SourceIdentity,
RemoteURL: source.RemoteURL,
RemoteNetworkPolicy: source.RemoteNetworkPolicy,
GitHubRepository: source.GitHubRepository,
ReleaseSelector: source.ReleaseSelector,
ReleaseTag: source.ReleaseTag,
AssetName: source.AssetName,
AutoUpdateEnabled: source.AutoUpdateEnabled,
CheckIntervalMinutes: source.CheckIntervalMinutes,
ETag: runtime.ETag,
LastSeenRevision: runtime.LastSeenRevision,
LastSeenDetail: runtime.LastSeenDetail,
LastAppliedRevision: runtime.LastAppliedRevision,
LastAppliedDetail: runtime.LastAppliedDetail,
RootDir: project.RootDir,
EntryFile: project.EntryFile,
LeaseToken: token,
LeaseExpiresAt: *runtime.LeaseExpiresAt,
}
return nil
})
if err != nil {
return nil, err
}
return &snapshot, nil
}
func inspectSourceLeaseMiss(
ctx context.Context,
sourceID uint,
expectedConfigVersion int,
now time.Time,
) (sourceLeaseOutcome, error) {
var source model.PagesProjectSource
if err := db.DB(ctx).Where("id = ?", sourceID).First(&source).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return sourceLeaseStale, nil
}
return sourceLeaseStale, err
}
if source.ConfigVersion != expectedConfigVersion {
return sourceLeaseStale, nil
}
var runtime model.PagesProjectSourceRuntime
if err := db.DB(ctx).Where("source_id = ?", sourceID).First(&runtime).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return sourceLeaseStale, nil
}
return sourceLeaseStale, err
}
if runtime.LeaseExpiresAt != nil && runtime.LeaseExpiresAt.After(now) {
return sourceLeaseBusy, nil
}
return sourceLeaseStale, nil
}
func sourceLeaseParameters(action string) (time.Duration, string, error) {
switch action {
case sourceActionCheck:
return pagesSourceCheckLeaseDuration, pagesSourceStatusChecking, nil
case sourceActionSync:
return pagesSourceSyncLeaseDuration, pagesSourceStatusSyncing, nil
default:
return 0, "", errors.New(errPagesSourceActionInvalid)
}
}
func newSourceLeaseToken() (string, error) {
value := make([]byte, sourceLeaseTokenBytes)
if _, err := rand.Read(value); err != nil {
return "", err
}
return hex.EncodeToString(value), nil
}
func renewSourceLease(ctx context.Context, snapshot *sourceExecutionSnapshot, duration time.Duration) (bool, error) {
if snapshot == nil || snapshot.SourceID == 0 || snapshot.LeaseToken == "" || duration <= 0 {
return false, errors.New(errPagesSourceLeaseLost)
}
now := time.Now()
expiresAt := now.Add(duration)
result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).
Where("source_id = ? AND lease_token = ? AND lease_expires_at > ?", snapshot.SourceID, snapshot.LeaseToken, now).
Updates(map[string]any{sourceRuntimeColumnLeaseExpiresAt: expiresAt})
if result.Error != nil {
return false, result.Error
}
if result.RowsAffected == 0 {
return false, nil
}
snapshot.LeaseExpiresAt = expiresAt
return true, nil
}
func failSourceLease(ctx context.Context, snapshot *sourceExecutionSnapshot, message string) error {
if snapshot == nil || snapshot.SourceID == 0 || snapshot.LeaseToken == "" {
return nil
}
message = safeSourceRuntimeError(message)
now := time.Now()
return db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).
Where("source_id = ? AND lease_token = ? AND lease_expires_at > ?", snapshot.SourceID, snapshot.LeaseToken, now).
Updates(map[string]any{
sourceRuntimeColumnSyncStatus: pagesSourceStatusFailed,
sourceRuntimeColumnLastError: message,
sourceRuntimeColumnLeaseToken: "",
sourceRuntimeColumnLeaseExpiresAt: nil,
}).Error
}
func safeSourceRuntimeError(message string) string {
message = strings.TrimSpace(message)
if message == "" {
return errPagesSourceSyncFailed
}
if len(message) > sourceRuntimeErrorMaxBytes {
message = message[:sourceRuntimeErrorMaxBytes]
}
return message
}
func sourceLeaseIsBusy(ctx context.Context, sourceID uint) (bool, error) {
var runtime model.PagesProjectSourceRuntime
if err := db.DB(ctx).Where("source_id = ?", sourceID).First(&runtime).Error; err != nil {
return false, err
}
return runtime.LeaseExpiresAt != nil && runtime.LeaseExpiresAt.After(time.Now()), nil
}
// recoverExpiredSourceLease clears one exact expired lease owner. Matching the
// token, observed expiry and status prevents a scanner from overwriting a
// worker that renewed or was replaced after the candidate query.
func recoverExpiredSourceLease(
ctx context.Context,
sourceID uint,
token string,
expiresAt time.Time,
status string,
now time.Time,
nextCheckAt *time.Time,
) (bool, error) {
if sourceID == 0 || token == "" ||
(status != pagesSourceStatusChecking && status != pagesSourceStatusSyncing) {
return false, nil
}
updates := map[string]any{
sourceRuntimeColumnSyncStatus: pagesSourceStatusFailed,
sourceRuntimeColumnLastError: errPagesSourceLeaseExpired,
sourceRuntimeColumnLeaseToken: "",
sourceRuntimeColumnLeaseExpiresAt: nil,
sourceRuntimeColumnNextCheckAt: nextCheckAt,
}
result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).
Where("source_id = ?", sourceID).
Where("lease_token = ?", token).
Where("lease_expires_at = ? AND lease_expires_at <= ?", expiresAt, now).
Where("sync_status = ?", status).
Updates(updates)
if result.Error != nil {
return false, result.Error
}
return result.RowsAffected == 1, nil
}
// fenceAndNormalizeRuntime invalidates in-flight work while preserving safe
// seen/applied cursors. The caller must already hold the source row lock.
func fenceAndNormalizeRuntime(tx *gorm.DB, sourceID uint) error {
var runtime model.PagesProjectSourceRuntime
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).
Where("source_id = ?", sourceID).
First(&runtime).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil
}
return err
}
return tx.Model(&runtime).Updates(map[string]any{
sourceRuntimeColumnLeaseToken: "",
sourceRuntimeColumnLeaseExpiresAt: nil,
sourceRuntimeColumnSyncStatus: normalizedSourceRuntimeStatus(&runtime),
}).Error
}
func sourceHasSameReleaseReplacement(runtime *model.PagesProjectSourceRuntime) bool {
if runtime == nil || runtime.LastSeenRevision == "" || runtime.LastSeenRevision == runtime.LastAppliedRevision {
return false
}
seen := sourceDetail{}
applied := sourceDetail{}
if unmarshalSourceDetail(runtime.LastSeenDetail, &seen) != nil ||
unmarshalSourceDetail(runtime.LastAppliedDetail, &applied) != nil {
return false
}
return seen.ReleaseID != "" && seen.ReleaseID == applied.ReleaseID
}
@@ -0,0 +1,350 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pages
import (
"errors"
"strings"
"sync"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
)
func TestSourceLeaseHeartbeatRenewsAndCancelsOnOwnershipLoss(t *testing.T) {
ctx := setupPagesSourceTest(t)
project := mustCreatePagesSourceProject(t, ctx, "lease-heartbeat")
source, _ := mustConfigureRemoteSource(
t,
ctx,
project.ID,
"https://example.com/site.zip",
RemoteNetworkPolicyPublic,
)
snapshot, outcome, err := acquireSourceLease(ctx, source.ID, source.ConfigVersion, sourceActionSync)
if err != nil || outcome != sourceLeaseAcquired || snapshot == nil {
t.Fatalf("acquireSourceLease(heartbeat) = (%+v, %q, %v), want acquired", snapshot, outcome, err)
}
workCtx, heartbeat, err := startSourceLeaseHeartbeat(ctx, snapshot, 500*time.Millisecond, 20*time.Millisecond)
if err != nil {
t.Fatalf("startSourceLeaseHeartbeat() error = %v, want nil", err)
}
t.Cleanup(func() { _ = heartbeat.stop() })
var initial model.PagesProjectSourceRuntime
if err := db.DB(ctx).Where("source_id = ?", source.ID).First(&initial).Error; err != nil {
t.Fatalf("load initial heartbeat runtime error = %v, want nil", err)
}
if initial.LeaseExpiresAt == nil {
t.Fatal("initial heartbeat expiry = nil, want non-nil")
}
deadline := time.Now().Add(2 * time.Second)
for {
var renewedRuntime model.PagesProjectSourceRuntime
if err := db.DB(ctx).Where("source_id = ?", source.ID).First(&renewedRuntime).Error; err != nil {
t.Fatalf("load renewed heartbeat runtime error = %v, want nil", err)
}
if renewedRuntime.LeaseExpiresAt != nil && renewedRuntime.LeaseExpiresAt.After(*initial.LeaseExpiresAt) {
break
}
if time.Now().After(deadline) {
t.Fatal("heartbeat did not extend lease before deadline")
}
time.Sleep(10 * time.Millisecond)
}
if err := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).
Where("source_id = ?", source.ID).
Update("lease_token", "replacement-owner").Error; err != nil {
t.Fatalf("replace heartbeat lease owner error = %v, want nil", err)
}
select {
case <-workCtx.Done():
case <-time.After(2 * time.Second):
t.Fatal("heartbeat work context was not canceled after ownership loss")
}
if err := heartbeat.stop(); !errors.Is(err, errSourceLeaseHeartbeatLost) {
t.Fatalf("heartbeat.stop() error = %v, want %v", err, errSourceLeaseHeartbeatLost)
}
if err := heartbeat.stop(); !errors.Is(err, errSourceLeaseHeartbeatLost) {
t.Fatalf("heartbeat.stop() second error = %v, want stable %v", err, errSourceLeaseHeartbeatLost)
}
}
func TestAcquireSourceLeaseConcurrentOnlyOneOwner(t *testing.T) {
ctx := setupPagesSourceTest(t)
project := mustCreatePagesSourceProject(t, ctx, "lease-concurrent")
source, _ := mustConfigureRemoteSource(
t,
ctx,
project.ID,
"https://example.com/site.zip",
RemoteNetworkPolicyPublic,
)
type leaseResult struct {
snapshot *sourceExecutionSnapshot
outcome sourceLeaseOutcome
err error
}
results := make(chan leaseResult, 2)
var workers sync.WaitGroup
workers.Add(2)
for range 2 {
go func() {
defer workers.Done()
snapshot, outcome, err := acquireSourceLease(ctx, source.ID, source.ConfigVersion, sourceActionSync)
results <- leaseResult{snapshot: snapshot, outcome: outcome, err: err}
}()
}
workers.Wait()
close(results)
acquired := 0
busy := 0
for result := range results {
if result.err != nil {
t.Errorf("acquireSourceLease(concurrent) error = %v, want nil", result.err)
continue
}
switch result.outcome {
case sourceLeaseAcquired:
acquired++
if result.snapshot == nil || result.snapshot.LeaseToken == "" {
t.Errorf("acquireSourceLease(concurrent acquired) snapshot = %+v, want token-bearing snapshot", result.snapshot)
}
case sourceLeaseBusy:
busy++
if result.snapshot != nil {
t.Errorf("acquireSourceLease(concurrent busy) snapshot = %+v, want nil", result.snapshot)
}
default:
t.Errorf("acquireSourceLease(concurrent) outcome = %q, want acquired or busy", result.outcome)
}
}
if acquired != 1 || busy != 1 {
t.Errorf("concurrent lease outcomes = acquired:%d busy:%d, want 1 and 1", acquired, busy)
}
}
func TestAcquireSourceLeaseMutualExclusionExpiryAndTerminalOwnership(t *testing.T) {
ctx := setupPagesSourceTest(t)
project := mustCreatePagesSourceProject(t, ctx, "lease-cas")
source, _ := mustConfigureRemoteSource(
t,
ctx,
project.ID,
"https://example.com/site.zip",
RemoteNetworkPolicyPublic,
)
first, outcome, err := acquireSourceLease(ctx, source.ID, source.ConfigVersion, sourceActionSync)
if err != nil {
t.Fatalf("acquireSourceLease(first) error = %v, want nil", err)
}
if got, want := outcome, sourceLeaseAcquired; got != want {
t.Fatalf("acquireSourceLease(first) outcome = %q, want %q", got, want)
}
if first == nil || first.LeaseToken == "" {
t.Fatalf("acquireSourceLease(first) snapshot = %+v, want token-bearing snapshot", first)
}
second, outcome, err := acquireSourceLease(ctx, source.ID, source.ConfigVersion, sourceActionSync)
if err != nil {
t.Fatalf("acquireSourceLease(duplicate) error = %v, want nil", err)
}
if got, want := outcome, sourceLeaseBusy; got != want {
t.Errorf("acquireSourceLease(duplicate) outcome = %q, want %q", got, want)
}
if second != nil {
t.Errorf("acquireSourceLease(duplicate) snapshot = %+v, want nil", second)
}
past := time.Now().Add(-time.Second)
if err := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).
Where("source_id = ?", source.ID).
Update("lease_expires_at", &past).Error; err != nil {
t.Fatalf("expire first lease error = %v, want nil", err)
}
takeover, outcome, err := acquireSourceLease(ctx, source.ID, source.ConfigVersion, sourceActionSync)
if err != nil {
t.Fatalf("acquireSourceLease(takeover) error = %v, want nil", err)
}
if got, want := outcome, sourceLeaseAcquired; got != want {
t.Fatalf("acquireSourceLease(takeover) outcome = %q, want %q", got, want)
}
if takeover == nil {
t.Fatal("acquireSourceLease(takeover) snapshot = nil, want non-nil")
}
if takeover.LeaseToken == "" || takeover.LeaseToken == first.LeaseToken {
t.Fatalf("takeover LeaseToken = %q, want non-empty token distinct from %q", takeover.LeaseToken, first.LeaseToken)
}
renewed, err := renewSourceLease(ctx, first, pagesSourceSyncLeaseDuration)
if err != nil {
t.Fatalf("renewSourceLease(expired owner) error = %v, want nil", err)
}
if renewed {
t.Error("renewSourceLease(expired owner) = true, want false")
}
if err := failSourceLease(ctx, first, "stale worker must not win"); err != nil {
t.Fatalf("failSourceLease(expired owner) error = %v, want nil", err)
}
var runtime model.PagesProjectSourceRuntime
if err := db.DB(ctx).Where("source_id = ?", source.ID).First(&runtime).Error; err != nil {
t.Fatalf("load runtime after takeover error = %v, want nil", err)
}
if got, want := runtime.LeaseToken, takeover.LeaseToken; got != want {
t.Errorf("runtime LeaseToken after stale terminal write = %q, want %q", got, want)
}
if got, want := runtime.SyncStatus, pagesSourceStatusSyncing; got != want {
t.Errorf("runtime SyncStatus after stale terminal write = %q, want %q", got, want)
}
renewed, err = renewSourceLease(ctx, takeover, pagesSourceSyncLeaseDuration)
if err != nil {
t.Fatalf("renewSourceLease(current owner) error = %v, want nil", err)
}
if !renewed {
t.Error("renewSourceLease(current owner) = false, want true")
}
if err := failSourceLease(ctx, takeover, errPagesSourceSyncFailed); err != nil {
t.Fatalf("failSourceLease(current owner) error = %v, want nil", err)
}
var failedRuntime model.PagesProjectSourceRuntime
if err := db.DB(ctx).Where("source_id = ?", source.ID).First(&failedRuntime).Error; err != nil {
t.Fatalf("load failed runtime error = %v, want nil", err)
}
if got, want := failedRuntime.SyncStatus, pagesSourceStatusFailed; got != want {
t.Errorf("failed runtime SyncStatus = %q, want %q", got, want)
}
if failedRuntime.LeaseToken != "" || failedRuntime.LeaseExpiresAt != nil {
t.Errorf("failed runtime lease = (%q, %v), want cleared", failedRuntime.LeaseToken, failedRuntime.LeaseExpiresAt)
}
}
func TestSourceConfigAndProjectContentChangesFenceLease(t *testing.T) {
ctx := setupPagesSourceTest(t)
project := mustCreatePagesSourceProject(t, ctx, "lease-fence")
source, _ := mustConfigureRemoteSource(
t,
ctx,
project.ID,
"https://example.com/site.zip?token=first",
RemoteNetworkPolicyPublic,
)
configSnapshot, outcome, err := acquireSourceLease(ctx, source.ID, source.ConfigVersion, sourceActionSync)
if err != nil || outcome != sourceLeaseAcquired {
t.Fatalf("acquireSourceLease(config fence) = (%+v, %q, %v), want acquired", configSnapshot, outcome, err)
}
if _, err := UpdateSource(ctx, project.ID, SourceUpdateInput{
SourceType: PagesSourceTypeRemoteURL,
RemoteURLSet: true,
RemoteURL: "https://example.com/site.zip?token=second",
RemoteNetworkPolicy: RemoteNetworkPolicyPublic,
}); err != nil {
t.Fatalf("UpdateSource(config fence) error = %v, want nil", err)
}
renewed, err := renewSourceLease(ctx, configSnapshot, pagesSourceSyncLeaseDuration)
if err != nil {
t.Fatalf("renewSourceLease(after source update) error = %v, want nil", err)
}
if renewed {
t.Error("renewSourceLease(after source update) = true, want false")
}
var updatedSource model.PagesProjectSource
if err := db.DB(ctx).Where("id = ?", source.ID).First(&updatedSource).Error; err != nil {
t.Fatalf("load updated source error = %v, want nil", err)
}
if got, want := updatedSource.ConfigVersion, source.ConfigVersion+1; got != want {
t.Errorf("updated source ConfigVersion = %d, want %d", got, want)
}
if snapshot, staleOutcome, staleErr := acquireSourceLease(ctx, source.ID, source.ConfigVersion, sourceActionSync); staleErr != nil || staleOutcome != sourceLeaseStale || snapshot != nil {
t.Errorf("acquireSourceLease(old config) = (%+v, %q, %v), want (nil, %q, nil)", snapshot, staleOutcome, staleErr, sourceLeaseStale)
}
contentSnapshot, outcome, err := acquireSourceLease(ctx, source.ID, updatedSource.ConfigVersion, sourceActionSync)
if err != nil || outcome != sourceLeaseAcquired {
t.Fatalf("acquireSourceLease(content fence) = (%+v, %q, %v), want acquired", contentSnapshot, outcome, err)
}
if _, err := UpdateProject(ctx, project.ID, Input{
Name: project.Name,
Slug: project.Slug,
Enabled: true,
RootDir: "dist",
EntryFile: "index.html",
}); err != nil {
t.Fatalf("UpdateProject(content fence) error = %v, want nil", err)
}
renewed, err = renewSourceLease(ctx, contentSnapshot, pagesSourceSyncLeaseDuration)
if err != nil {
t.Fatalf("renewSourceLease(after content update) error = %v, want nil", err)
}
if renewed {
t.Error("renewSourceLease(after content update) = true, want false")
}
storedProject, err := model.GetPagesProjectByID(ctx, project.ID)
if err != nil {
t.Fatalf("GetPagesProjectByID(%d) error = %v, want nil", project.ID, err)
}
if got, want := storedProject.ContentConfigVersion, project.ContentConfigVersion+1; got != want {
t.Errorf("ContentConfigVersion after RootDir update = %d, want %d", got, want)
}
var runtime model.PagesProjectSourceRuntime
if err := db.DB(ctx).Where("source_id = ?", source.ID).First(&runtime).Error; err != nil {
t.Fatalf("load fenced runtime error = %v, want nil", err)
}
if runtime.LeaseToken != "" || runtime.LeaseExpiresAt != nil {
t.Errorf("content-fenced runtime lease = (%q, %v), want cleared", runtime.LeaseToken, runtime.LeaseExpiresAt)
}
}
func TestSourceRuntimeUsesOnlySixDocumentedStates(t *testing.T) {
states := []string{
pagesSourceStatusIdle,
pagesSourceStatusChecking,
pagesSourceStatusUpdateAvailable,
pagesSourceStatusSyncing,
pagesSourceStatusFailed,
pagesSourceStatusAttention,
}
seen := make(map[string]struct{}, len(states))
for _, state := range states {
if strings.TrimSpace(state) == "" {
t.Errorf("documented source state = %q, want non-empty", state)
}
if _, exists := seen[state]; exists {
t.Errorf("documented source state %q is duplicated", state)
}
seen[state] = struct{}{}
}
if got, want := len(seen), 6; got != want {
t.Errorf("unique source states = %d, want %d", got, want)
}
updateRuntime := &model.PagesProjectSourceRuntime{
LastSeenRevision: strings.Repeat("a", 64),
LastAppliedRevision: strings.Repeat("b", 64),
LastSeenDetail: `{"release_id":"new"}`,
LastAppliedDetail: `{"release_id":"old"}`,
}
if got, want := normalizedSourceRuntimeStatus(updateRuntime), pagesSourceStatusUpdateAvailable; got != want {
t.Errorf("normalizedSourceRuntimeStatus(update) = %q, want %q", got, want)
}
attentionRuntime := &model.PagesProjectSourceRuntime{
LastSeenRevision: strings.Repeat("a", 64),
LastAppliedRevision: strings.Repeat("b", 64),
LastSeenDetail: `{"release_id":"same"}`,
LastAppliedDetail: `{"release_id":"same"}`,
}
if got, want := normalizedSourceRuntimeStatus(attentionRuntime), pagesSourceStatusAttention; got != want {
t.Errorf("normalizedSourceRuntimeStatus(attention) = %q, want %q", got, want)
}
if got, want := normalizedSourceRuntimeStatus(&model.PagesProjectSourceRuntime{}), pagesSourceStatusIdle; got != want {
t.Errorf("normalizedSourceRuntimeStatus(idle) = %q, want %q", got, want)
}
}
@@ -0,0 +1,478 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pages
import (
"bytes"
"context"
"encoding/json"
"errors"
"fmt"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/task"
"github.com/Rain-kl/Wavelet/pkg/logger"
"gorm.io/gorm"
)
const (
// PagesSourceScanTask is the private Asynq task type for the periodic scanner.
PagesSourceScanTask = "openflare:pages_source_scan"
// TaskTypePagesSourceScan is the internal task meta type seeded in w_schedules.
TaskTypePagesSourceScan = "of_pages_source_scan"
pagesSourceScanBatchSize = 20
)
// PagesSourceScanMeta is available to the scheduler registry but hidden from
// generic Admin task dispatch and schedule mutation APIs.
var PagesSourceScanMeta = task.TaskMeta{
Type: TaskTypePagesSourceScan,
AsynqTask: PagesSourceScanTask,
Name: "OpenFlare Pages 部署源扫描",
Description: "补偿孤儿部署包、恢复过期执行权并串行检查到期的 GitHub latest 部署源",
SupportsTime: false,
MaxRetry: 0,
Queue: task.QueueDefault,
Retryable: false,
InternalOnly: true,
}
type pagesSourceScanPayload struct{}
type pagesSourceScanSummary struct {
ExpiredCandidates int `json:"expired_candidates"`
RecoveredLeases int `json:"recovered_leases"`
OrphanCleanup PagesOrphanCleanupSummary `json:"orphan_cleanup"`
DueSources int `json:"due_sources"`
SelectedSources int `json:"selected_sources"`
CheckedSources int `json:"checked_sources"`
UpdatesFound int `json:"updates_found"`
AttentionSources int `json:"attention_sources"`
DispatchedSyncs int `json:"dispatched_syncs"`
FailedDispatches int `json:"failed_dispatches"`
BusySources int `json:"busy_sources"`
StaleSources int `json:"stale_sources"`
FailedSources int `json:"failed_sources"`
Backlog int `json:"backlog"`
ProviderBackoffs []pagesSourceProviderBackoff `json:"provider_backoffs,omitempty"`
}
type pagesSourceProviderBackoff struct {
SourceID uint `json:"source_id"`
StatusCode int `json:"status_code"`
RetryAt string `json:"retry_at"`
}
type expiredSourceLeaseCandidate struct {
SourceID uint
LeaseToken string
LeaseExpiresAt time.Time
SyncStatus string
SourceType string
ReleaseSelector string
}
type dueGitHubSourceCandidate struct {
SourceID uint
ConfigVersion int
}
var (
pagesSourceScanNow = time.Now
reconcilePagesSourceOrphans = ReconcilePagesOrphanUploads
dispatchPagesSourceAutoSync = func(
ctx context.Context,
source model.PagesProjectSource,
targetRevision string,
) (*SourceActionReceipt, error) {
return dispatchSourceActionSnapshotWithTrigger(
ctx,
source,
sourceActionSync,
pagesSourceCreatedBySystem,
pagesSourceTriggerScheduledAutoUpdate,
targetRevision,
"",
"system",
)
}
)
// SourceScanHandler serializes provider checks inside one scheduled task. A
// source-level lease still permits overlapping scanner executions safely.
type SourceScanHandler struct{}
// ValidatePayload accepts only an empty object; the scanner has no user input.
func (handler *SourceScanHandler) ValidatePayload(payload []byte) ([]byte, error) {
if len(bytes.TrimSpace(payload)) == 0 {
payload = []byte("{}")
}
var input pagesSourceScanPayload
decoder := json.NewDecoder(bytes.NewReader(payload))
decoder.DisallowUnknownFields()
if err := decoder.Decode(&input); err != nil {
return nil, errors.New(errPagesSourceActionInvalid)
}
if err := ensureJSONEOF(decoder); err != nil {
return nil, errors.New(errPagesSourceActionInvalid)
}
return []byte("{}"), nil
}
// Execute recovers expired leases and checks at most 20 due latest sources in
// stable order. Provider and dispatch failures are isolated per source.
func (handler *SourceScanHandler) Execute(ctx context.Context, payload []byte) (*task.TaskResult, error) {
if _, err := handler.ValidatePayload(payload); err != nil {
return nil, task.PermanentError(errPagesSourceActionInvalid)
}
now := pagesSourceScanNow()
summary := pagesSourceScanSummary{}
if err := recoverExpiredPagesSourceLeases(ctx, now, &summary); err != nil {
return nil, err
}
orphanSummary, err := reconcilePagesSourceOrphans(ctx, now)
if err != nil {
return nil, err
}
summary.OrphanCleanup = orphanSummary
task.AppendLog(
ctx,
"[cleanup] orphan 候选=%d,已补偿=%d,仍被引用=%d,lease busy=%d,非法 marker=%d,跳过=%d,失败=%d",
orphanSummary.Candidates,
orphanSummary.Reconciled,
orphanSummary.Referenced,
orphanSummary.LeaseBusy,
orphanSummary.InvalidMarker,
orphanSummary.Skipped,
orphanSummary.Failed,
)
if err := scanDueGitHubSources(ctx, now, &summary); err != nil {
return nil, err
}
detail, err := json.Marshal(summary)
if err != nil {
return nil, err
}
message := fmt.Sprintf(
"Pages 部署源扫描完成:恢复 %d 个租约,补偿 %d 个孤儿记录,检查 %d 个来源,投递 %d 个自动更新,积压 %d 个",
summary.RecoveredLeases,
summary.OrphanCleanup.Reconciled,
summary.CheckedSources,
summary.DispatchedSyncs,
summary.Backlog,
)
return &task.TaskResult{Message: message, Detail: string(detail)}, nil
}
func recoverExpiredPagesSourceLeases(
ctx context.Context,
now time.Time,
summary *pagesSourceScanSummary,
) error {
var candidates []expiredSourceLeaseCandidate
err := db.DB(ctx).
Table("of_pages_project_source_runtime AS runtime").
Select(`runtime.source_id, runtime.lease_token, runtime.lease_expires_at,
runtime.sync_status, source.source_type, source.release_selector`).
Joins("JOIN of_pages_project_sources AS source ON source.id = runtime.source_id").
Where("runtime.lease_token <> ''").
Where("runtime.lease_expires_at IS NOT NULL AND runtime.lease_expires_at <= ?", now).
Where("runtime.sync_status IN ?", []string{pagesSourceStatusChecking, pagesSourceStatusSyncing}).
Order("runtime.source_id ASC").
Scan(&candidates).Error
if err != nil {
return err
}
summary.ExpiredCandidates = len(candidates)
for _, candidate := range candidates {
var nextCheckAt *time.Time
if candidate.SourceType == PagesSourceTypeGitHubRelease &&
candidate.ReleaseSelector == githubReleaseSelectorLatest {
next := nextGitHubCheckAt(now, candidate.SourceID, minimumCheckInterval)
nextCheckAt = &next
}
recovered, recoverErr := recoverExpiredSourceLease(
ctx,
candidate.SourceID,
candidate.LeaseToken,
candidate.LeaseExpiresAt,
candidate.SyncStatus,
now,
nextCheckAt,
)
if recoverErr != nil {
summary.FailedSources++
logger.WarnF(
ctx,
"[PagesSourceScan] recover expired lease failed: source_id=%d error=%v",
candidate.SourceID,
recoverErr,
)
continue
}
if recovered {
summary.RecoveredLeases++
task.AppendLog(ctx, "[recover] 已恢复过期来源租约:source_id=%d", candidate.SourceID)
}
}
return nil
}
func scanDueGitHubSources(
ctx context.Context,
now time.Time,
summary *pagesSourceScanSummary,
) error {
dueQuery := func() *gorm.DB {
return db.DB(ctx).
Table("of_pages_project_source_runtime AS runtime").
Joins("JOIN of_pages_project_sources AS source ON source.id = runtime.source_id").
Where("source.source_type = ?", PagesSourceTypeGitHubRelease).
Where("source.release_selector = ?", githubReleaseSelectorLatest).
Where("runtime.next_check_at IS NOT NULL AND runtime.next_check_at <= ?", now)
}
var dueCount int64
if err := dueQuery().Count(&dueCount).Error; err != nil {
return err
}
summary.DueSources = int(dueCount)
var candidates []dueGitHubSourceCandidate
if err := dueQuery().
Select("source.id AS source_id, source.config_version").
Order("runtime.next_check_at ASC").
Order("source.id ASC").
Limit(pagesSourceScanBatchSize).
Scan(&candidates).Error; err != nil {
return err
}
summary.SelectedSources = len(candidates)
task.AppendLog(
ctx,
"[scan] 到期来源=%d,本批=%d",
summary.DueSources,
summary.SelectedSources,
)
for _, candidate := range candidates {
scanOneDueGitHubSource(ctx, candidate, summary)
}
var remainingDue int64
if err := dueQuery().Count(&remainingDue).Error; err != nil {
return err
}
summary.Backlog = int(remainingDue)
task.AppendLog(ctx, "[scan] 本批处理后仍到期来源=%d", summary.Backlog)
return nil
}
func scanOneDueGitHubSource(
ctx context.Context,
candidate dueGitHubSourceCandidate,
summary *pagesSourceScanSummary,
) {
snapshot, outcome, err := acquireSourceLease(
ctx,
candidate.SourceID,
candidate.ConfigVersion,
sourceActionCheck,
)
if err != nil {
summary.FailedSources++
logger.WarnF(ctx, "[PagesSourceScan] acquire check lease failed: source_id=%d error=%v", candidate.SourceID, err)
return
}
switch outcome {
case sourceLeaseBusy:
summary.BusySources++
task.AppendLog(ctx, "[check] 来源正在执行其它任务,跳过:source_id=%d", candidate.SourceID)
return
case sourceLeaseStale:
summary.StaleSources++
return
}
if snapshot == nil || snapshot.SourceType != PagesSourceTypeGitHubRelease ||
snapshot.ReleaseSelector != githubReleaseSelectorLatest {
summary.StaleSources++
if snapshot != nil {
if finalizeErr := failSourceLease(ctx, snapshot, errPagesSourceActionStale); finalizeErr != nil {
logger.WarnF(
ctx,
"[PagesSourceScan] finalize stale source failed: source_id=%d error=%v",
snapshot.SourceID,
finalizeErr,
)
}
}
return
}
checkResult, checkErr := checkGitHubSource(ctx, snapshot)
if checkErr != nil {
summary.FailedSources++
recordPagesSourceProviderBackoff(ctx, candidate.SourceID, checkErr, summary)
logger.WarnF(
ctx,
"[PagesSourceScan] source check failed: source_id=%d error=%s",
candidate.SourceID,
safeGitHubSourceError(checkErr),
)
return
}
if checkResult == nil || checkResult.Stale {
summary.StaleSources++
return
}
handleCheckedGitHubSource(ctx, snapshot, checkResult, summary)
}
func recordPagesSourceProviderBackoff(
ctx context.Context,
sourceID uint,
checkErr error,
summary *pagesSourceScanSummary,
) {
var domainError *githubSourceProviderDomainError
if !errors.As(checkErr, &domainError) ||
(domainError.statusCode != 403 && domainError.statusCode != 429) {
return
}
retryAt := domainError.retryAt
var runtime model.PagesProjectSourceRuntime
if err := db.DB(ctx).
Select("next_check_at").
Where("source_id = ?", sourceID).
First(&runtime).Error; err != nil {
logger.WarnF(ctx, "[PagesSourceScan] load provider backoff deadline failed: source_id=%d error=%v", sourceID, err)
} else if runtime.NextCheckAt != nil {
retryAt = runtime.NextCheckAt
}
retryAtText := "unknown"
if retryAt != nil {
retryAtText = retryAt.UTC().Format(time.RFC3339)
}
summary.ProviderBackoffs = append(summary.ProviderBackoffs, pagesSourceProviderBackoff{
SourceID: sourceID, StatusCode: domainError.statusCode, RetryAt: retryAtText,
})
task.AppendLog(
ctx,
"[check] GitHub provider 退避:source_id=%d status=%d retry_at=%s",
sourceID,
domainError.statusCode,
retryAtText,
)
}
func handleCheckedGitHubSource(
ctx context.Context,
snapshot *sourceExecutionSnapshot,
checkResult *githubCheckTaskResult,
summary *pagesSourceScanSummary,
) {
summary.CheckedSources++
switch checkResult.Status {
case pagesSourceStatusUpdateAvailable:
summary.UpdatesFound++
case pagesSourceStatusAttention:
summary.AttentionSources++
}
if !snapshot.AutoUpdateEnabled || checkResult.Status != pagesSourceStatusUpdateAvailable ||
!validOptionalSourceRevision(checkResult.Revision) || checkResult.Revision == "" {
return
}
source := model.PagesProjectSource{
ID: snapshot.SourceID,
ProjectID: snapshot.ProjectID,
ConfigVersion: snapshot.SourceConfigVersion,
}
receipt, dispatchErr := dispatchPagesSourceAutoSync(ctx, source, checkResult.Revision)
if dispatchErr == nil {
summary.DispatchedSyncs++
if receipt != nil {
task.AppendLog(
ctx,
"[dispatch] 已投递自动更新:source_id=%d execution_id=%s revision=%s",
snapshot.SourceID,
receipt.ExecutionID,
checkResult.Revision,
)
}
return
}
summary.FailedSources++
summary.FailedDispatches++
logger.WarnF(
ctx,
"[PagesSourceScan] dispatch auto sync failed: source_id=%d revision=%s error=%v",
snapshot.SourceID,
checkResult.Revision,
dispatchErr,
)
updated, recordErr := recordPagesSourceAutoDispatchFailure(
ctx,
snapshot,
checkResult.Revision,
checkResult.RetryAt,
)
if recordErr != nil {
logger.WarnF(
ctx,
"[PagesSourceScan] record auto sync dispatch failure failed: source_id=%d error=%v",
snapshot.SourceID,
recordErr,
)
} else if !updated {
summary.StaleSources++
}
}
func recordPagesSourceAutoDispatchFailure(
ctx context.Context,
snapshot *sourceExecutionSnapshot,
revision string,
retryAt *time.Time,
) (bool, error) {
if snapshot == nil || revision == "" {
return false, nil
}
now := pagesSourceScanNow()
next := nextGitHubCheckAt(now, snapshot.SourceID, minimumCheckInterval)
if retryAt != nil && retryAt.After(next) {
next = retryAt.In(now.Location())
}
result := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).
Where("source_id = ?", snapshot.SourceID).
Where("sync_status = ? AND last_seen_revision = ?", pagesSourceStatusUpdateAvailable, revision).
Where("lease_expires_at IS NULL OR lease_expires_at <= ?", now).
Where(`EXISTS (
SELECT 1 FROM of_pages_project_sources AS source
WHERE source.id = ? AND source.config_version = ?
AND source.source_type = ? AND source.release_selector = ?
AND source.auto_update_enabled = ?
)`,
snapshot.SourceID,
snapshot.SourceConfigVersion,
PagesSourceTypeGitHubRelease,
githubReleaseSelectorLatest,
true,
).
Updates(map[string]any{
sourceRuntimeColumnSyncStatus: pagesSourceStatusUpdateAvailable,
sourceRuntimeColumnLastError: errPagesSourceTaskDispatchFailed,
sourceRuntimeColumnNextCheckAt: &next,
})
if result.Error != nil {
return false, result.Error
}
return result.RowsAffected == 1, nil
}
@@ -0,0 +1,511 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pages
import (
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"os"
"path/filepath"
"strings"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/integration/githubrelease"
"github.com/Rain-kl/Wavelet/internal/model"
"gorm.io/gorm"
)
type scannerDispatchedSync struct {
SourceID uint
Revision string
}
func TestGitHubLatestAutoConfigPreservesIdentityAndRuntimeCursor(t *testing.T) {
ctx := setupPagesSourceTest(t)
project := mustCreatePagesSourceProject(t, ctx, "scanner-auto-config")
source, _ := mustConfigureGitHubSourceWithoutDispatch(t, ctx, project.ID, SourceUpdateInput{
SourceType: PagesSourceTypeGitHubRelease,
RepositoryURL: "https://github.com/scanner/auto-config",
AutoUpdateEnabled: false,
CheckIntervalMinutes: 60,
})
identity := source.SourceIdentity
seenRevision := strings.Repeat("a", sourceRevisionHexLength)
appliedRevision := strings.Repeat("b", sourceRevisionHexLength)
future := time.Now().Add(time.Hour)
if err := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).
Where("source_id = ?", source.ID).
Updates(map[string]any{
"etag": `"cursor-etag"`,
"last_seen_revision": seenRevision,
"last_seen_detail": `{"provider":"github","release_id":"2","asset_id":"2","tag":"v2","asset_name":"dist.zip"}`,
"last_applied_revision": appliedRevision,
"last_applied_detail": `{"provider":"github","release_id":"1","asset_id":"1","tag":"v1","asset_name":"dist.zip"}`,
"sync_status": pagesSourceStatusSyncing,
"last_error": "old error",
"lease_token": "in-flight",
"lease_expires_at": &future,
}).Error; err != nil {
t.Fatalf("seed runtime error = %v, want nil", err)
}
input := SourceUpdateInput{
SourceType: PagesSourceTypeGitHubRelease,
RepositoryURL: "https://github.com/scanner/auto-config",
AutoUpdateEnabled: true,
CheckIntervalMinutes: 15,
}
if err := validateGitHubSourceInput(input); err != nil {
t.Fatalf("validateGitHubSourceInput(auto latest) error = %v, want nil", err)
}
if err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
changed, err := updateGitHubSourceTx(tx, project.ID, input)
if err == nil && !changed {
return errors.New("auto config update was treated as no-op")
}
return err
}); err != nil {
t.Fatalf("updateGitHubSourceTx(auto latest) error = %v, want nil", err)
}
updated, runtime := mustLoadPagesSource(t, ctx, project.ID)
if updated.SourceIdentity != identity || updated.ConfigVersion != source.ConfigVersion+1 {
t.Errorf(
"updated source = identity:%q version:%d, want identity:%q version:%d",
updated.SourceIdentity,
updated.ConfigVersion,
identity,
source.ConfigVersion+1,
)
}
if !updated.AutoUpdateEnabled || updated.CheckIntervalMinutes != 15 {
t.Errorf("updated auto config = enabled:%t interval:%d, want true/15", updated.AutoUpdateEnabled, updated.CheckIntervalMinutes)
}
if runtime.ETag != `"cursor-etag"` || runtime.LastSeenRevision != seenRevision ||
runtime.LastAppliedRevision != appliedRevision {
t.Errorf("runtime cursor changed after auto-only update: %+v", runtime)
}
if runtime.LeaseToken != "" || runtime.LeaseExpiresAt != nil || runtime.LastError != "" ||
runtime.SyncStatus != pagesSourceStatusUpdateAvailable || runtime.NextCheckAt == nil {
t.Errorf(
"runtime fence = token:%q expiry:%v error:%q status:%q next:%v",
runtime.LeaseToken,
runtime.LeaseExpiresAt,
runtime.LastError,
runtime.SyncStatus,
runtime.NextCheckAt,
)
}
tagConfig, err := buildGitHubSourceConfig(SourceUpdateInput{
SourceType: PagesSourceTypeGitHubRelease,
RepositoryURL: "https://github.com/scanner/auto-config",
ReleaseSelector: githubReleaseSelectorTag,
ReleaseTag: "v1",
AutoUpdateEnabled: true,
CheckIntervalMinutes: 60,
})
if err != nil {
t.Fatalf("buildGitHubSourceConfig(tag) error = %v, want nil", err)
}
if tagConfig.AutoUpdate || tagConfig.CheckInterval != 0 {
t.Errorf("tag config auto/interval = %t/%d, want false/0", tagConfig.AutoUpdate, tagConfig.CheckInterval)
}
}
func TestRecoverExpiredPagesSourceLeaseUsesExactCASAndStableJitter(t *testing.T) {
ctx := setupPagesSourceTest(t)
project := mustCreatePagesSourceProject(t, ctx, "scanner-expired-lease")
source, _ := mustConfigureGitHubSourceWithoutDispatch(t, ctx, project.ID, SourceUpdateInput{
SourceType: PagesSourceTypeGitHubRelease,
RepositoryURL: "https://github.com/scanner/expired-lease",
})
now := time.Date(2026, 7, 19, 12, 0, 0, 0, time.UTC)
usePagesSourceScannerClock(t, now)
expiredAt := now.Add(-time.Minute)
if err := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).
Where("source_id = ?", source.ID).
Updates(map[string]any{
"sync_status": pagesSourceStatusChecking,
"lease_token": "expired-owner",
"lease_expires_at": &expiredAt,
"next_check_at": &expiredAt,
}).Error; err != nil {
t.Fatalf("seed expired lease error = %v, want nil", err)
}
summary := pagesSourceScanSummary{}
if err := recoverExpiredPagesSourceLeases(ctx, now, &summary); err != nil {
t.Fatalf("recoverExpiredPagesSourceLeases() error = %v, want nil", err)
}
if summary.ExpiredCandidates != 1 || summary.RecoveredLeases != 1 || summary.FailedSources != 0 {
t.Errorf("recovery summary = %+v, want one recovered lease", summary)
}
_, runtime := mustLoadPagesSource(t, ctx, project.ID)
wantNext := nextGitHubCheckAt(now, source.ID, minimumCheckInterval)
if runtime.SyncStatus != pagesSourceStatusFailed || runtime.LastError != errPagesSourceLeaseExpired ||
runtime.LeaseToken != "" || runtime.LeaseExpiresAt != nil || runtime.NextCheckAt == nil ||
runtime.NextCheckAt.Sub(wantNext) != 0 {
t.Errorf(
"recovered runtime = status:%q error:%q token:%q expiry:%v next:%v, want next %v",
runtime.SyncStatus,
runtime.LastError,
runtime.LeaseToken,
runtime.LeaseExpiresAt,
runtime.NextCheckAt,
wantNext,
)
}
renewedExpiry := now.Add(time.Minute)
if err := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).
Where("source_id = ?", source.ID).
Updates(map[string]any{
"sync_status": pagesSourceStatusSyncing,
"lease_token": "renewed-owner",
"lease_expires_at": &renewedExpiry,
}).Error; err != nil {
t.Fatalf("seed renewed lease error = %v, want nil", err)
}
recovered, err := recoverExpiredSourceLease(
ctx,
source.ID,
"renewed-owner",
expiredAt,
pagesSourceStatusSyncing,
now,
&wantNext,
)
if err != nil || recovered {
t.Fatalf("recoverExpiredSourceLease(stale expiry) = %t, %v; want false, nil", recovered, err)
}
_, runtime = mustLoadPagesSource(t, ctx, project.ID)
if runtime.LeaseToken != "renewed-owner" || runtime.LeaseExpiresAt == nil ||
runtime.LeaseExpiresAt.Sub(renewedExpiry) != 0 || runtime.SyncStatus != pagesSourceStatusSyncing {
t.Errorf("stale recovery overwrote renewed lease: %+v", runtime)
}
}
func TestScheduledAutoSyncPersistsExplicitDeploymentTrigger(t *testing.T) {
ctx := setupPagesSourceSyncTest(t)
project := mustCreatePagesSourceProject(t, ctx, "scanner-scheduled-trigger")
source, _ := mustConfigureGitHubSourceWithoutDispatch(t, ctx, project.ID, SourceUpdateInput{
SourceType: PagesSourceTypeGitHubRelease,
RepositoryURL: "https://github.com/scanner/scheduled-trigger",
AutoUpdateEnabled: true,
})
packageBytes := testPagesZip(t, map[string]string{"index.html": "scheduled-v1"})
packageHash := sha256.Sum256(packageBytes)
release := githubrelease.Release{ID: "scheduled-release", Tag: "v1"}
asset := githubrelease.Asset{
ID: "scheduled-asset", Name: defaultGitHubAssetName, State: "uploaded",
UpdatedAt: time.Date(2026, 7, 19, 12, 30, 0, 0, time.UTC),
}
target, err := buildGitHubSourceTarget(release, asset, nil)
if err != nil {
t.Fatalf("buildGitHubSourceTarget() error = %v, want nil", err)
}
useFakeGitHubReleaseClient(t, &fakeGitHubReleaseClient{
resolve: func(context.Context, githubrelease.ResolveRequest) (githubrelease.ResolveResult, error) {
return githubrelease.ResolveResult{Release: release, Asset: asset}, nil
},
download: func(context.Context, githubrelease.DownloadRequest) (*githubrelease.DownloadResult, error) {
path := filepath.Join(t.TempDir(), "scheduled.zip")
if err := os.WriteFile(path, packageBytes, 0o600); err != nil {
t.Fatalf("os.WriteFile(scheduled package) error = %v", err)
}
return &githubrelease.DownloadResult{
Path: path, Size: int64(len(packageBytes)), SHA256: hex.EncodeToString(packageHash[:]),
}, nil
},
})
snapshot, outcome, err := acquireSourceLease(ctx, source.ID, source.ConfigVersion, sourceActionSync)
if err != nil || outcome != sourceLeaseAcquired {
t.Fatalf("acquireSourceLease() = %+v, %q, %v; want acquired", snapshot, outcome, err)
}
synced, err := syncGitHubSourceWithTrigger(
ctx,
snapshot,
pagesSourceCreatedBySystem,
target.Revision,
"",
pagesSourceTriggerScheduledAutoUpdate,
)
if err != nil || synced == nil || synced.Deployment == nil || synced.Stale {
t.Fatalf("syncGitHubSourceWithTrigger() = %+v, %v; want active deployment", synced, err)
}
deployment, err := model.GetPagesDeploymentByID(ctx, synced.Deployment.ID)
if err != nil {
t.Fatalf("GetPagesDeploymentByID(%d) error = %v", synced.Deployment.ID, err)
}
if deployment.TriggerType != pagesSourceTriggerScheduledAutoUpdate ||
deployment.CreatedBy != pagesSourceCreatedBySystem {
t.Errorf(
"scheduled provenance = trigger:%q actor:%q, want %q/%q",
deployment.TriggerType,
deployment.CreatedBy,
pagesSourceTriggerScheduledAutoUpdate,
pagesSourceCreatedBySystem,
)
}
}
func TestPagesSourceScannerSerialBatchIsolation304AndActualBacklog(t *testing.T) {
ctx := setupPagesSourceTest(t)
now := time.Now().Truncate(time.Second)
usePagesSourceScannerClock(t, now)
dueAt := now.Add(-time.Hour)
type fixture struct {
source *model.PagesProjectSource
runtime *model.PagesProjectSourceRuntime
repository string
}
fixtures := make([]fixture, 0, 22)
byRepository := make(map[string]int, 22)
for index := 1; index <= 22; index++ {
project := mustCreatePagesSourceProject(t, ctx, fmt.Sprintf("scanner-batch-%02d", index))
repository := fmt.Sprintf("scanner/source-%02d", index)
source, runtime := mustConfigureGitHubSourceWithoutDispatch(t, ctx, project.ID, SourceUpdateInput{
SourceType: PagesSourceTypeGitHubRelease,
RepositoryURL: "https://github.com/" + repository,
AutoUpdateEnabled: index != 4,
CheckIntervalMinutes: 60,
})
if err := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).
Where("source_id = ?", source.ID).
Update("next_check_at", &dueAt).Error; err != nil {
t.Fatalf("mark source %d due error = %v, want nil", source.ID, err)
}
fixtures = append(fixtures, fixture{source: source, runtime: runtime, repository: repository})
byRepository[repository] = index
}
busyUntil := now.Add(time.Hour)
if err := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).
Where("source_id = ?", fixtures[0].source.ID).
Updates(map[string]any{
"sync_status": pagesSourceStatusChecking,
"lease_token": "busy-owner",
"lease_expires_at": &busyUntil,
}).Error; err != nil {
t.Fatalf("seed busy source error = %v, want nil", err)
}
stored304Revision := strings.Repeat("3", sourceRevisionHexLength)
if err := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).
Where("source_id = ?", fixtures[2].source.ID).
Updates(map[string]any{
"etag": `"stored-etag"`,
"last_seen_revision": stored304Revision,
"last_seen_detail": `{"provider":"github","release_id":"release-3","asset_id":"3","tag":"v3","asset_name":"dist.zip"}`,
"sync_status": pagesSourceStatusUpdateAvailable,
}).Error; err != nil {
t.Fatalf("seed 304 cursor error = %v, want nil", err)
}
if err := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).
Where("source_id = ?", fixtures[4].source.ID).
Updates(map[string]any{
"last_applied_revision": strings.Repeat("a", sourceRevisionHexLength),
"last_applied_detail": `{"provider":"github","release_id":"shared-release","asset_id":"old","tag":"v5","asset_name":"dist.zip"}`,
}).Error; err != nil {
t.Fatalf("seed replacement cursor error = %v, want nil", err)
}
retryAt := now.Add(2 * time.Hour)
calledRepositories := make([]string, 0, pagesSourceScanBatchSize)
useFakeGitHubReleaseClient(t, &fakeGitHubReleaseClient{
resolve: func(_ context.Context, request githubrelease.ResolveRequest) (githubrelease.ResolveResult, error) {
index := byRepository[request.Repository]
calledRepositories = append(calledRepositories, request.Repository)
switch index {
case 2:
return githubrelease.ResolveResult{}, &githubrelease.Error{
Kind: githubrelease.ErrMetadata, StatusCode: 429, RetryAt: &retryAt,
}
case 3:
if request.ETag != `"stored-etag"` {
t.Errorf("304 source ETag = %q, want stored ETag", request.ETag)
}
return githubrelease.ResolveResult{NotModified: true, ETag: request.ETag}, nil
default:
releaseID := fmt.Sprintf("release-%d", index)
if index == 5 {
releaseID = "shared-release"
}
result := githubrelease.ResolveResult{
ETag: fmt.Sprintf(`"etag-%d"`, index),
Release: githubrelease.Release{ID: releaseID, Tag: fmt.Sprintf("v%d", index)},
Asset: githubrelease.Asset{
ID: fmt.Sprintf("asset-%d", index),
Name: defaultGitHubAssetName,
State: "uploaded",
UpdatedAt: now.Add(time.Duration(index) * time.Minute),
},
}
if index == 6 {
result.RetryAt = &retryAt
}
return result, nil
}
},
download: func(context.Context, githubrelease.DownloadRequest) (*githubrelease.DownloadResult, error) {
t.Fatal("scanner downloaded an asset; want check-only behavior")
return nil, nil
},
})
dispatched := make([]scannerDispatchedSync, 0, pagesSourceScanBatchSize)
previousDispatch := dispatchPagesSourceAutoSync
dispatchPagesSourceAutoSync = func(
_ context.Context,
source model.PagesProjectSource,
revision string,
) (*SourceActionReceipt, error) {
if source.ID == fixtures[5].source.ID {
return nil, errors.New("injected dispatch failure")
}
dispatched = append(dispatched, scannerDispatchedSync{SourceID: source.ID, Revision: revision})
return &SourceActionReceipt{ExecutionID: fmt.Sprintf("%d", source.ID), Action: sourceActionSync}, nil
}
t.Cleanup(func() { dispatchPagesSourceAutoSync = previousDispatch })
result, err := (&SourceScanHandler{}).Execute(ctx, []byte("{}"))
if err != nil {
t.Fatalf("SourceScanHandler.Execute() error = %v, want nil", err)
}
var summary pagesSourceScanSummary
if err := json.Unmarshal([]byte(result.Detail), &summary); err != nil {
t.Fatalf("json.Unmarshal(scan detail) error = %v, want nil", err)
}
if summary.DueSources != 22 || summary.SelectedSources != pagesSourceScanBatchSize ||
summary.CheckedSources != 18 || summary.UpdatesFound != 17 || summary.AttentionSources != 1 ||
summary.DispatchedSyncs != 15 || summary.FailedDispatches != 1 ||
summary.BusySources != 1 || summary.FailedSources != 2 ||
summary.Backlog != 3 {
t.Errorf("scan summary = %+v, want due=22 selected=20 checked=18 updates=17 attention=1 dispatched=15 dispatch_failed=1 busy=1 failed=2 backlog=3", summary)
}
if len(summary.ProviderBackoffs) != 1 ||
summary.ProviderBackoffs[0].SourceID != fixtures[1].source.ID ||
summary.ProviderBackoffs[0].StatusCode != 429 ||
summary.ProviderBackoffs[0].RetryAt != retryAt.UTC().Format(time.RFC3339) {
t.Errorf("scan provider backoffs = %+v, want source=%d status=429 retry_at=%s", summary.ProviderBackoffs, fixtures[1].source.ID, retryAt.UTC().Format(time.RFC3339))
}
if len(calledRepositories) != 19 {
t.Fatalf("Resolve calls = %d, want 19 (one busy source in selected batch)", len(calledRepositories))
}
for index, repository := range calledRepositories {
want := fixtures[index+1].repository
if repository != want {
t.Fatalf("Resolve order[%d] = %q, want %q", index, repository, want)
}
}
if !containsDispatchedSource(dispatched, fixtures[2].source.ID, stored304Revision) {
t.Errorf("304 stored revision was not dispatched: %+v", dispatched)
}
if containsDispatchedSource(dispatched, fixtures[3].source.ID, "") {
t.Errorf("auto=false source was dispatched: %+v", dispatched)
}
if containsDispatchedSource(dispatched, fixtures[4].source.ID, "") {
t.Errorf("attention source was dispatched: %+v", dispatched)
}
_, dispatchFailedRuntime := mustLoadPagesSource(t, ctx, fixtures[5].source.ProjectID)
if dispatchFailedRuntime.SyncStatus != pagesSourceStatusUpdateAvailable ||
dispatchFailedRuntime.LastError != errPagesSourceTaskDispatchFailed ||
dispatchFailedRuntime.NextCheckAt == nil || dispatchFailedRuntime.NextCheckAt.Before(retryAt) ||
dispatchFailedRuntime.LastSeenRevision == "" {
t.Errorf(
"dispatch failure runtime = status:%q error:%q next:%v seen:%q, want preserved update and provider deadline >= %v",
dispatchFailedRuntime.SyncStatus,
dispatchFailedRuntime.LastError,
dispatchFailedRuntime.NextCheckAt,
dispatchFailedRuntime.LastSeenRevision,
retryAt,
)
}
}
func TestPagesSourceScannerIncludesOrphanCleanupSummary(t *testing.T) {
ctx := setupPagesSourceTest(t)
now := time.Now().Truncate(time.Second)
usePagesSourceScannerClock(t, now)
previousReconcile := reconcilePagesSourceOrphans
reconcilePagesSourceOrphans = func(
_ context.Context,
gotNow time.Time,
) (PagesOrphanCleanupSummary, error) {
if !gotNow.Equal(now) {
t.Errorf("orphan cleanup now = %v, want %v", gotNow, now)
}
return PagesOrphanCleanupSummary{
Candidates: 7,
Reconciled: 1,
Referenced: 2,
LeaseBusy: 1,
InvalidMarker: 1,
Skipped: 1,
Failed: 1,
}, nil
}
t.Cleanup(func() { reconcilePagesSourceOrphans = previousReconcile })
result, err := (&SourceScanHandler{}).Execute(ctx, []byte("{}"))
if err != nil {
t.Fatalf("SourceScanHandler.Execute() error = %v, want nil", err)
}
var summary pagesSourceScanSummary
if err := json.Unmarshal([]byte(result.Detail), &summary); err != nil {
t.Fatalf("json.Unmarshal(scan detail) error = %v, want nil", err)
}
if summary.OrphanCleanup.Candidates != 7 || summary.OrphanCleanup.Reconciled != 1 ||
summary.OrphanCleanup.Referenced != 2 || summary.OrphanCleanup.LeaseBusy != 1 ||
summary.OrphanCleanup.InvalidMarker != 1 || summary.OrphanCleanup.Skipped != 1 ||
summary.OrphanCleanup.Failed != 1 {
t.Errorf("orphan cleanup summary = %+v, want injected result", summary.OrphanCleanup)
}
}
func TestPagesSourceScanPayloadAndMetaAreInternalOnly(t *testing.T) {
handler := &SourceScanHandler{}
if normalized, err := handler.ValidatePayload(nil); err != nil || string(normalized) != "{}" {
t.Errorf("ValidatePayload(nil) = %s, %v; want {}, nil", normalized, err)
}
if _, err := handler.ValidatePayload([]byte(`{"unexpected":true}`)); err == nil {
t.Error("ValidatePayload(unknown field) error = nil, want non-nil")
}
if !PagesSourceScanMeta.InternalOnly || PagesSourceScanMeta.Type != TaskTypePagesSourceScan ||
PagesSourceScanMeta.AsynqTask != PagesSourceScanTask || PagesSourceScanMeta.MaxRetry != 0 {
t.Errorf("PagesSourceScanMeta = %+v, want internal bounded scheduled scanner", PagesSourceScanMeta)
}
if PagesSourceScanMeta.SupportsTime {
t.Error("PagesSourceScanMeta.SupportsTime = true, want empty scanner payload")
}
}
func usePagesSourceScannerClock(t *testing.T, now time.Time) {
t.Helper()
previous := pagesSourceScanNow
pagesSourceScanNow = func() time.Time { return now }
t.Cleanup(func() { pagesSourceScanNow = previous })
}
func containsDispatchedSource(dispatched []scannerDispatchedSync, sourceID uint, revision string) bool {
for _, item := range dispatched {
if item.SourceID == sourceID && (revision == "" || item.Revision == revision) {
return true
}
}
return false
}
@@ -0,0 +1,745 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pages
import (
"context"
"encoding/json"
"errors"
"fmt"
"path"
"sort"
"strings"
"sync"
"time"
"unicode/utf8"
"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/task"
"github.com/Rain-kl/Wavelet/pkg/logger"
"gorm.io/gorm"
"gorm.io/gorm/clause"
)
const (
pagesSourceTriggerManualSync = "manual_sync"
pagesSourceTriggerScheduledAutoUpdate = "scheduled_auto_update"
pagesSourceCreatedBySystem = "system:pages-source-sync"
pagesSourceHeartbeatInterval = pagesSourceSyncLeaseDuration / 3
pagesSourceCleanupTimeout = 15 * time.Second
)
var (
errSourceFinalFence = errors.New("pages source final fence rejected")
errSourceLeaseHeartbeatLost = errors.New("pages source lease heartbeat lost")
sourceCommitNow = time.Now
)
type sourceSyncOutcome struct {
Deployment *DeploymentView
Reused bool
Stale bool
}
type preparedRemoteSource struct {
Candidate *SourceCandidate
Manifest *deploymentManifest
Detail sourceDetail
DetailJSON string
}
type sourceIngestState struct {
Result upload.IngestResult
HasIngest bool
Referenced bool
}
type sourceCommitState struct {
Project *model.PagesProject
Source *model.PagesProjectSource
Runtime *model.PagesProjectSourceRuntime
Now time.Time
}
type sourceLeaseHeartbeat struct {
cancel context.CancelFunc
done <-chan error
stopOnce sync.Once
stopErr error
}
func startSourceLeaseHeartbeat(
ctx context.Context,
snapshot *sourceExecutionSnapshot,
leaseDuration time.Duration,
interval time.Duration,
) (context.Context, *sourceLeaseHeartbeat, error) {
if snapshot == nil || leaseDuration <= 0 || interval <= 0 || interval >= leaseDuration {
return nil, nil, errors.New(errPagesSourceLeaseLost)
}
renewed, err := renewSourceLease(ctx, snapshot, leaseDuration)
if err != nil {
return nil, nil, err
}
if !renewed {
return nil, nil, errSourceLeaseHeartbeatLost
}
workCtx, cancel := context.WithCancel(ctx)
done := make(chan error, 1)
heartbeat := &sourceLeaseHeartbeat{cancel: cancel, done: done}
go func() {
ticker := time.NewTicker(interval)
defer ticker.Stop()
for {
select {
case <-workCtx.Done():
done <- nil
return
case <-ticker.C:
renewed, renewErr := renewSourceLease(workCtx, snapshot, leaseDuration)
if renewErr != nil {
if workCtx.Err() != nil {
done <- nil
return
}
done <- renewErr
cancel()
return
}
if !renewed {
done <- errSourceLeaseHeartbeatLost
cancel()
return
}
}
}
}()
return workCtx, heartbeat, nil
}
func (heartbeat *sourceLeaseHeartbeat) stop() error {
if heartbeat == nil {
return nil
}
heartbeat.stopOnce.Do(func() {
heartbeat.cancel()
heartbeat.stopErr = <-heartbeat.done
})
return heartbeat.stopErr
}
func sourceHeartbeatOutcome(err error) (*sourceSyncOutcome, error) {
if errors.Is(err, errSourceLeaseHeartbeatLost) {
return &sourceSyncOutcome{Stale: true}, nil
}
return nil, err
}
func recordSourceLeaseFailure(ctx context.Context, snapshot *sourceExecutionSnapshot) {
cleanupCtx, cancel := sourceCleanupContext(ctx)
defer cancel()
if err := failSourceLease(cleanupCtx, snapshot, errPagesSourceSyncFailed); err != nil {
var sourceID uint
if snapshot != nil {
sourceID = snapshot.SourceID
}
logger.WarnF(cleanupCtx, "[PagesSource] record failed runtime state failed: source_id=%d error=%v", sourceID, err)
}
}
func sourceCleanupContext(ctx context.Context) (context.Context, context.CancelFunc) {
return context.WithTimeout(context.WithoutCancel(ctx), pagesSourceCleanupTimeout)
}
func syncRemoteSourceWithTrigger(
ctx context.Context,
snapshot *sourceExecutionSnapshot,
actor string,
triggerType string,
) (outcome *sourceSyncOutcome, resultErr error) {
if snapshot == nil || snapshot.SourceType != PagesSourceTypeRemoteURL {
return nil, errors.New(errPagesSourceTypeUnsupported)
}
actor = strings.TrimSpace(actor)
if actor == "" || !validSourceDeploymentTrigger(triggerType) {
return nil, errors.New(errPagesSourceActionInvalid)
}
defer func() {
if resultErr != nil {
recordSourceLeaseFailure(ctx, snapshot)
}
}()
workCtx, heartbeat, err := startSourceLeaseHeartbeat(
ctx,
snapshot,
pagesSourceSyncLeaseDuration,
pagesSourceHeartbeatInterval,
)
if err != nil {
return sourceHeartbeatOutcome(err)
}
defer func() {
_ = heartbeat.stop()
}()
limits := resolvePagesLimits(workCtx)
prepared, err := prepareRemoteSource(workCtx, snapshot, limits)
if err != nil {
if heartbeatErr := heartbeat.stop(); heartbeatErr != nil {
return sourceHeartbeatOutcome(heartbeatErr)
}
return nil, err
}
defer func() {
if cleanupErr := prepared.Candidate.Cleanup(); cleanupErr != nil {
logger.WarnF(ctx, "[PagesSource] cleanup temporary package failed: source_id=%d error=%v", snapshot.SourceID, cleanupErr)
}
}()
ingestState, err := resolveSourceIngest(workCtx, snapshot, prepared)
if err != nil {
if heartbeatErr := heartbeat.stop(); heartbeatErr != nil {
return sourceHeartbeatOutcome(heartbeatErr)
}
return nil, err
}
defer func() {
compensateSourceIngest(ctx, snapshot, ingestState)
}()
if heartbeatErr := heartbeat.stop(); heartbeatErr != nil {
return sourceHeartbeatOutcome(heartbeatErr)
}
renewed, err := renewSourceLease(ctx, snapshot, pagesSourceSyncLeaseDuration)
if err != nil {
return nil, err
}
if !renewed {
return &sourceSyncOutcome{Stale: true}, nil
}
task.AppendLog(ctx, "[activate] 正在原子切换生产部署")
deployment, reused, referenced, err := commitSourceDeploymentWithTrigger(
ctx,
snapshot,
prepared.Candidate.Checksum,
prepared.Candidate.Checksum,
prepared.Detail,
prepared.DetailJSON,
actor,
triggerType,
prepared.Manifest,
ingestState.Result,
ingestState.HasIngest,
nil,
)
ingestState.Referenced = referenced
if errors.Is(err, errSourceFinalFence) {
return &sourceSyncOutcome{Stale: true}, nil
}
if err != nil {
return nil, err
}
ingestState.Referenced = ingestState.HasIngest && deployment.UploadID == ingestState.Result.Upload.ID
if pruneErr := pruneProjectDeploymentHistory(ctx, snapshot.ProjectID, limits.HistoryCount, 0); pruneErr != nil {
logger.ErrorF(ctx,
"[PagesSource] strict prune failed: project_id=%d source_id=%d keep=%d error=%v",
snapshot.ProjectID,
snapshot.SourceID,
limits.HistoryCount,
pruneErr,
)
}
view := buildDeploymentView(deployment)
return &sourceSyncOutcome{Deployment: &view, Reused: reused}, nil
}
func prepareRemoteSource(
ctx context.Context,
snapshot *sourceExecutionSnapshot,
limits pagesLimits,
) (*preparedRemoteSource, error) {
task.AppendLog(ctx, "[download] 正在获取远程部署包")
candidate, err := FetchRemoteSource(ctx, RemoteSourceRequest{
URL: snapshot.RemoteURL,
NetworkPolicy: snapshot.RemoteNetworkPolicy,
MaxPackageBytes: limits.PackageBytes,
})
if err != nil {
return nil, err
}
if candidate == nil || candidate.TempPath == "" || candidate.Checksum == "" || candidate.Format == "" {
return nil, errors.New(errPagesSourceSyncFailed)
}
rootDir, err := validateAndNormalizePagesRootDir(snapshot.RootDir)
if err != nil {
cleanupFailedRemoteCandidate(ctx, snapshot, candidate)
return nil, err
}
entryFile, err := validateAndNormalizePagesEntryFile(snapshot.EntryFile)
if err != nil {
cleanupFailedRemoteCandidate(ctx, snapshot, candidate)
return nil, err
}
task.AppendLog(ctx, "[verify] 正在校验归档结构与入口文件")
manifest, err := inspectPagesPackage(candidate.TempPath, candidate.Format, rootDir, entryFile, limits)
if err != nil {
cleanupFailedRemoteCandidate(ctx, snapshot, candidate)
return nil, err
}
detail := sourceDetail{Provider: PagesSourceTypeRemoteURL, DisplayName: safeRemoteSourceLabel(candidate.SafeLabel)}
detailJSON, err := json.Marshal(detail)
if err != nil {
cleanupFailedRemoteCandidate(ctx, snapshot, candidate)
return nil, err
}
return &preparedRemoteSource{
Candidate: candidate,
Manifest: manifest,
Detail: detail,
DetailJSON: string(detailJSON),
}, nil
}
func cleanupFailedRemoteCandidate(
ctx context.Context,
snapshot *sourceExecutionSnapshot,
candidate *SourceCandidate,
) {
if err := candidate.Cleanup(); err != nil {
logger.WarnF(ctx, "[PagesSource] cleanup failed preparation package: source_id=%d error=%v", snapshot.SourceID, err)
}
}
func resolveSourceIngest(
ctx context.Context,
snapshot *sourceExecutionSnapshot,
prepared *preparedRemoteSource,
) (*sourceIngestState, error) {
_, err := findSourceDeployment(
ctx,
snapshot.ProjectID,
snapshot.SourceIdentity,
prepared.Candidate.Checksum,
)
if err == nil {
return &sourceIngestState{}, nil
}
if !errors.Is(err, gorm.ErrRecordNotFound) {
return nil, err
}
task.AppendLog(ctx, "[ingest] 正在保存受管部署包")
result, err := ingestPagesDeploymentPackageWithSource(
ctx,
prepared.Candidate.TempPath,
prepared.Candidate.Checksum,
snapshot.ProjectID,
snapshot.SourceID,
sourceDetailLabel(prepared.Detail),
prepared.Candidate.Format,
)
if err != nil {
return nil, err
}
return &sourceIngestState{Result: result, HasIngest: true}, nil
}
func compensateSourceIngest(
ctx context.Context,
snapshot *sourceExecutionSnapshot,
state *sourceIngestState,
) {
if state == nil || !state.HasIngest || !state.Result.Created || state.Referenced {
return
}
cleanupCtx, cancel := sourceCleanupContext(ctx)
defer cancel()
task.AppendLog(cleanupCtx, "[cleanup] 正在补偿未引用的部署包记录")
if err := removePagesUploadIfUnreferenced(cleanupCtx, snapshot.ProjectID, state.Result.Upload.ID); err != nil {
logger.ErrorF(cleanupCtx,
"[PagesSource] compensate upload failed: project_id=%d source_id=%d upload_id=%d error=%v",
snapshot.ProjectID, snapshot.SourceID, state.Result.Upload.ID, err,
)
}
}
func findSourceDeployment(
ctx context.Context,
projectID uint,
sourceIdentity string,
revision string,
) (*model.PagesDeployment, error) {
var deployment model.PagesDeployment
err := db.DB(ctx).
Where("project_id = ? AND source_identity = ? AND source_revision = ?", projectID, sourceIdentity, revision).
First(&deployment).Error
if err != nil {
return nil, err
}
return &deployment, nil
}
func commitSourceDeploymentWithTrigger(
ctx context.Context,
snapshot *sourceExecutionSnapshot,
revision string,
packageChecksum string,
detail sourceDetail,
detailJSON string,
actor string,
triggerType string,
manifest *deploymentManifest,
ingestResult upload.IngestResult,
hasIngest bool,
nextCheckNotBefore *time.Time,
) (*model.PagesDeployment, bool, bool, error) {
if snapshot == nil || manifest == nil {
return nil, false, false, errors.New(errPagesSourceSyncFailed)
}
if !validSourceDeploymentTrigger(triggerType) {
return nil, false, false, errors.New(errPagesSourceActionInvalid)
}
var committed model.PagesDeployment
reused := false
ingestReferenced := false
err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
state, err := lockSourceCommitState(tx, snapshot)
if err != nil {
return err
}
target, targetReused, err := resolveSourceDeploymentTx(
tx, state, revision, packageChecksum, detail, detailJSON, actor, triggerType,
manifest, ingestResult, hasIngest,
)
if err != nil {
return err
}
if err := lockSourceDeploymentUploadsTx(tx, target, ingestResult, hasIngest); err != nil {
return err
}
if err := ensureDeploymentEntry(tx, target.ID, state.Project.RootDir, state.Project.EntryFile); err != nil {
return err
}
if err := refreshSourceCommitLease(state, snapshot); err != nil {
return err
}
if err := activateSourceDeploymentTx(tx, state, target, revision, detailJSON, nextCheckNotBefore); err != nil {
return err
}
committed = *target
reused = targetReused
ingestReferenced = hasIngest && target.UploadID == ingestResult.Upload.ID
return nil
})
if err != nil {
return nil, false, false, err
}
return &committed, reused, ingestReferenced, nil
}
func lockSourceCommitState(tx *gorm.DB, snapshot *sourceExecutionSnapshot) (*sourceCommitState, error) {
state := &sourceCommitState{}
var project model.PagesProject
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).
First(&project, snapshot.ProjectID).Error; err != nil {
return nil, sourceFenceRecordError(err)
}
if project.ContentConfigVersion != snapshot.ContentConfigVersion {
return nil, errSourceFinalFence
}
var source model.PagesProjectSource
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).
Where("id = ? AND project_id = ?", snapshot.SourceID, snapshot.ProjectID).
First(&source).Error; err != nil {
return nil, sourceFenceRecordError(err)
}
if source.ConfigVersion != snapshot.SourceConfigVersion ||
source.SourceIdentity != snapshot.SourceIdentity ||
source.SourceType != snapshot.SourceType {
return nil, errSourceFinalFence
}
var runtime model.PagesProjectSourceRuntime
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).
Where("source_id = ?", source.ID).
First(&runtime).Error; err != nil {
return nil, sourceFenceRecordError(err)
}
state.Project = &project
state.Source = &source
state.Runtime = &runtime
if err := refreshSourceCommitLease(state, snapshot); err != nil {
return nil, err
}
return state, nil
}
func refreshSourceCommitLease(state *sourceCommitState, snapshot *sourceExecutionSnapshot) error {
if state == nil || state.Runtime == nil || snapshot == nil {
return errSourceFinalFence
}
now := sourceCommitNow()
if state.Runtime.LeaseToken != snapshot.LeaseToken || state.Runtime.LeaseExpiresAt == nil ||
!state.Runtime.LeaseExpiresAt.After(now) {
return errSourceFinalFence
}
state.Now = now
return nil
}
func sourceFenceRecordError(err error) error {
if errors.Is(err, gorm.ErrRecordNotFound) {
return errSourceFinalFence
}
return err
}
func resolveSourceDeploymentTx(
tx *gorm.DB,
state *sourceCommitState,
revision string,
packageChecksum string,
detail sourceDetail,
detailJSON string,
actor string,
triggerType string,
manifest *deploymentManifest,
ingestResult upload.IngestResult,
hasIngest bool,
) (*model.PagesDeployment, bool, error) {
var target model.PagesDeployment
err := tx.Where(
"project_id = ? AND source_identity = ? AND source_revision = ?",
state.Project.ID,
state.Source.SourceIdentity,
revision,
).First(&target).Error
if err == nil {
return &target, true, nil
}
if !errors.Is(err, gorm.ErrRecordNotFound) {
return nil, false, err
}
if !hasIngest {
return nil, false, errSourceFinalFence
}
return createSourceDeploymentTx(
tx, state, revision, packageChecksum, detail, detailJSON, actor, triggerType, manifest, ingestResult,
)
}
func createSourceDeploymentTx(
tx *gorm.DB,
state *sourceCommitState,
revision string,
packageChecksum string,
detail sourceDetail,
detailJSON string,
actor string,
triggerType string,
manifest *deploymentManifest,
ingestResult upload.IngestResult,
) (*model.PagesDeployment, bool, error) {
var maxNumber int
if err := tx.Model(&model.PagesDeployment{}).
Where("project_id = ?", state.Project.ID).
Select("COALESCE(MAX(deployment_number), 0)").
Scan(&maxNumber).Error; err != nil {
return nil, false, err
}
identity := state.Source.SourceIdentity
revisionValue := revision
target := &model.PagesDeployment{
ProjectID: state.Project.ID,
DeploymentNumber: maxNumber + 1,
Checksum: packageChecksum,
Status: model.PagesDeploymentStatusUploaded,
UploadID: ingestResult.Upload.ID,
FileCount: manifest.FileCount,
TotalSize: manifest.TotalSize,
CreatedBy: actor,
SourceType: state.Source.SourceType,
SourceIdentity: &identity,
SourceRevision: &revisionValue,
SourceLabel: sourceDetailLabel(detail),
SourceMeta: detailJSON,
TriggerType: triggerType,
}
result := tx.Clauses(clause.OnConflict{DoNothing: true}).Create(target)
if result.Error != nil {
return nil, false, result.Error
}
if result.RowsAffected == 0 {
return reloadSourceDeploymentTx(tx, state.Project.ID, identity, revision)
}
if err := createSourceDeploymentFilesTx(tx, target.ID, manifest.Files); err != nil {
return nil, false, err
}
return target, false, nil
}
func validSourceDeploymentTrigger(triggerType string) bool {
return triggerType == pagesSourceTriggerManualSync || triggerType == pagesSourceTriggerScheduledAutoUpdate
}
func reloadSourceDeploymentTx(
tx *gorm.DB,
projectID uint,
identity string,
revision string,
) (*model.PagesDeployment, bool, error) {
var target model.PagesDeployment
err := tx.Where(
"project_id = ? AND source_identity = ? AND source_revision = ?",
projectID,
identity,
revision,
).First(&target).Error
return &target, true, err
}
func createSourceDeploymentFilesTx(tx *gorm.DB, deploymentID uint, files []model.PagesDeploymentFile) error {
if len(files) == 0 {
return nil
}
for index := range files {
files[index].DeploymentID = deploymentID
}
return tx.Create(&files).Error
}
func lockSourceDeploymentUploadsTx(
tx *gorm.DB,
target *model.PagesDeployment,
ingestResult upload.IngestResult,
hasIngest bool,
) error {
uploadIDs := []uint64{target.UploadID}
if hasIngest && ingestResult.Upload.ID != 0 && ingestResult.Upload.ID != target.UploadID {
uploadIDs = append(uploadIDs, ingestResult.Upload.ID)
}
sort.Slice(uploadIDs, func(i, j int) bool { return uploadIDs[i] < uploadIDs[j] })
var records []model.Upload
if err := tx.Clauses(clause.Locking{Strength: pagesRowLockStrength}).
Where("id IN ?", uploadIDs).
Order("id asc").
Find(&records).Error; err != nil {
return err
}
for index := range records {
if records[index].ID != target.UploadID {
continue
}
if records[index].Status == model.UploadStatusUsed && records[index].Type == upload.ReservedPagesDeploymentType {
return nil
}
break
}
return errSourceFinalFence
}
func activateSourceDeploymentTx(
tx *gorm.DB,
state *sourceCommitState,
target *model.PagesDeployment,
revision string,
detailJSON string,
nextCheckNotBefore *time.Time,
) error {
if err := tx.Model(&model.PagesDeployment{}).
Where("project_id = ?", state.Project.ID).
Update("status", model.PagesDeploymentStatusUploaded).Error; err != nil {
return err
}
if err := tx.Model(target).Updates(map[string]any{
pagesDeploymentColumnStatus: model.PagesDeploymentStatusActive,
"activated_at": &state.Now,
}).Error; err != nil {
return err
}
if err := tx.Model(state.Project).Update("active_deployment_id", target.ID).Error; err != nil {
return err
}
finishedAt := sourceCommitNow()
var nextCheckAt any
if state.Source.SourceType == PagesSourceTypeGitHubRelease &&
state.Source.ReleaseSelector == githubReleaseSelectorLatest {
next := nextGitHubCheckAt(finishedAt, state.Source.ID, state.Source.CheckIntervalMinutes)
if nextCheckNotBefore != nil && nextCheckNotBefore.After(next) {
next = nextCheckNotBefore.In(finishedAt.Location())
}
nextCheckAt = &next
}
result := tx.Model(&model.PagesProjectSourceRuntime{}).
Where("source_id = ? AND lease_token = ? AND lease_expires_at > ?",
state.Runtime.SourceID,
state.Runtime.LeaseToken,
finishedAt,
).
Updates(map[string]any{
"last_seen_revision": revision,
"last_seen_detail": detailJSON,
"last_applied_revision": revision,
"last_applied_detail": detailJSON,
sourceRuntimeColumnSyncStatus: pagesSourceStatusIdle,
sourceRuntimeColumnLastError: "",
sourceRuntimeColumnLastCheckedAt: &finishedAt,
"last_synced_at": &finishedAt,
"next_check_at": nextCheckAt,
sourceRuntimeColumnLeaseToken: "",
sourceRuntimeColumnLeaseExpiresAt: nil,
})
if result.Error != nil {
return result.Error
}
if result.RowsAffected != 1 {
return errSourceFinalFence
}
return nil
}
func safeRemoteSourceLabel(raw string) string {
value := strings.ReplaceAll(strings.ToValidUTF8(raw, ""), "\\", "/")
value = path.Base(strings.TrimSpace(value))
if value == "." || value == "/" {
value = ""
}
var builder strings.Builder
for _, character := range value {
if character >= 0x20 && character != 0x7f {
builder.WriteRune(character)
}
}
value = strings.TrimSpace(builder.String())
if value == "" {
value = defaultRemoteAssetLabel
}
if len(value) > remoteSourceMaxSafeLabelBytes {
value = value[:remoteSourceMaxSafeLabelBytes]
for !utf8.ValidString(value) {
_, size := utf8.DecodeLastRuneInString(value)
value = value[:len(value)-size]
}
}
return value
}
func sourceSyncResultDetail(outcome *sourceSyncOutcome) string {
if outcome == nil || outcome.Deployment == nil {
return ""
}
payload := map[string]any{
"deployment_id": outcome.Deployment.ID,
"reused": outcome.Reused,
}
encoded, err := json.Marshal(payload)
if err != nil {
return fmt.Sprintf(`{"deployment_id":%d}`, outcome.Deployment.ID)
}
return string(encoded)
}
@@ -0,0 +1,608 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pages
import (
"context"
"crypto/sha256"
"encoding/hex"
"errors"
"fmt"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"strings"
"testing"
"time"
"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/pkg/pagesarchive"
"gorm.io/gorm"
)
func setupPagesSourceSyncTest(t *testing.T) context.Context {
t.Helper()
ctx := setupPagesSourceTest(t)
_, disableStorage := setupPagesStorageMock(t)
t.Cleanup(disableStorage)
return ctx
}
func mustAcquireRemoteSyncLease(
t *testing.T,
ctx context.Context,
source *model.PagesProjectSource,
) *sourceExecutionSnapshot {
t.Helper()
snapshot, outcome, err := acquireSourceLease(ctx, source.ID, source.ConfigVersion, sourceActionSync)
if err != nil {
t.Fatalf("acquireSourceLease(source=%d) error = %v, want nil", source.ID, err)
}
if got, want := outcome, sourceLeaseAcquired; got != want {
t.Fatalf("acquireSourceLease(source=%d) outcome = %q, want %q", source.ID, got, want)
}
if snapshot == nil {
t.Fatalf("acquireSourceLease(source=%d) snapshot = nil, want non-nil", source.ID)
}
return snapshot
}
func mustCreateActiveManualDeployment(
t *testing.T,
ctx context.Context,
projectID uint,
content string,
) *model.PagesDeployment {
t.Helper()
view, err := UploadDeployment(
ctx,
projectID,
testPagesMultipartFile(t, "manual.zip", testPagesZip(t, map[string]string{"index.html": content})),
"user:1",
)
if err != nil {
t.Fatalf("UploadDeployment(project=%d) error = %v, want nil", projectID, err)
}
if _, err := ActivateDeploymentAs(ctx, projectID, view.ID, "user:1"); err != nil {
t.Fatalf("ActivateDeploymentAs(project=%d, deployment=%d) error = %v, want nil", projectID, view.ID, err)
}
deployment, err := model.GetPagesDeploymentByID(ctx, view.ID)
if err != nil {
t.Fatalf("GetPagesDeploymentByID(%d) error = %v, want nil", view.ID, err)
}
return deployment
}
func newPagesArchiveServer(t *testing.T, status int, body []byte, beforeWrite func() error) *httptest.Server {
t.Helper()
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
if beforeWrite != nil {
if err := beforeWrite(); err != nil {
writer.WriteHeader(http.StatusInternalServerError)
return
}
}
writer.Header().Set("Content-Type", "application/zip")
writer.Header().Set("Content-Disposition", `attachment; filename="site.zip"`)
writer.WriteHeader(status)
_, _ = writer.Write(body)
}))
t.Cleanup(server.Close)
return server
}
func TestSyncRemoteSourceAtomicallyActivatesAndReusesChecksum(t *testing.T) {
ctx := setupPagesSourceSyncTest(t)
project := mustCreatePagesSourceProject(t, ctx, "sync-success")
packageBytes := testPagesZip(t, map[string]string{
"index.html": "remote-v1",
"assets/app.js": "console.log('ok')",
})
server := newPagesArchiveServer(t, http.StatusOK, packageBytes, nil)
secret := "sync-query-secret"
source, _ := mustConfigureRemoteSource(
t,
ctx,
project.ID,
server.URL+"/site.zip?token="+secret,
RemoteNetworkPolicyTrustedInternal,
)
firstSnapshot := mustAcquireRemoteSyncLease(t, ctx, source)
first, err := syncRemoteSource(ctx, firstSnapshot, "user:42")
if err != nil {
t.Fatalf("syncRemoteSource(first) error = %v, want nil", err)
}
if first == nil || first.Stale || first.Reused || first.Deployment == nil {
t.Fatalf("syncRemoteSource(first) = %+v, want new active deployment", first)
}
storedProject, err := model.GetPagesProjectByID(ctx, project.ID)
if err != nil {
t.Fatalf("GetPagesProjectByID(%d) error = %v, want nil", project.ID, err)
}
if storedProject.ActiveDeploymentID == nil || *storedProject.ActiveDeploymentID != first.Deployment.ID {
t.Fatalf("project ActiveDeploymentID = %v, want %d", storedProject.ActiveDeploymentID, first.Deployment.ID)
}
deployment, err := model.GetPagesDeploymentByID(ctx, first.Deployment.ID)
if err != nil {
t.Fatalf("GetPagesDeploymentByID(%d) error = %v, want nil", first.Deployment.ID, err)
}
expectedHash := sha256.Sum256(packageBytes)
if got, want := deployment.Checksum, hex.EncodeToString(expectedHash[:]); got != want {
t.Errorf("deployment Checksum = %q, want %q", got, want)
}
if got, want := deployment.Status, model.PagesDeploymentStatusActive; got != want {
t.Errorf("deployment Status = %q, want %q", got, want)
}
if got, want := deployment.SourceType, PagesSourceTypeRemoteURL; got != want {
t.Errorf("deployment SourceType = %q, want %q", got, want)
}
if deployment.SourceIdentity == nil || *deployment.SourceIdentity != source.SourceIdentity {
t.Errorf("deployment SourceIdentity = %v, want %q", deployment.SourceIdentity, source.SourceIdentity)
}
if deployment.SourceRevision == nil || *deployment.SourceRevision != deployment.Checksum {
t.Errorf("deployment SourceRevision = %v, want %q", deployment.SourceRevision, deployment.Checksum)
}
if got, want := deployment.CreatedBy, "user:42"; got != want {
t.Errorf("deployment CreatedBy = %q, want %q", got, want)
}
if got, want := deployment.TriggerType, pagesSourceTriggerManualSync; got != want {
t.Errorf("deployment TriggerType = %q, want %q", got, want)
}
if strings.Contains(deployment.SourceMeta, secret) || strings.Contains(deployment.SourceLabel, secret) {
t.Errorf("deployment provenance = label:%q meta:%q, want no query secret", deployment.SourceLabel, deployment.SourceMeta)
}
var runtime model.PagesProjectSourceRuntime
if err := db.DB(ctx).Where("source_id = ?", source.ID).First(&runtime).Error; err != nil {
t.Fatalf("load source runtime error = %v, want nil", err)
}
if got, want := runtime.SyncStatus, pagesSourceStatusIdle; got != want {
t.Errorf("runtime SyncStatus = %q, want %q", got, want)
}
if got, want := runtime.LastAppliedRevision, deployment.Checksum; got != want {
t.Errorf("runtime LastAppliedRevision = %q, want %q", got, want)
}
if runtime.LeaseToken != "" || runtime.LeaseExpiresAt != nil {
t.Errorf("runtime lease = (%q, %v), want cleared", runtime.LeaseToken, runtime.LeaseExpiresAt)
}
var uploadRecord model.Upload
if err := db.DB(ctx).First(&uploadRecord, deployment.UploadID).Error; err != nil {
t.Fatalf("load deployment upload %d error = %v, want nil", deployment.UploadID, err)
}
if got, want := uploadRecord.Status, model.UploadStatusUsed; got != want {
t.Errorf("deployment upload Status = %q, want %q", got, want)
}
if got, want := uploadRecord.Type, upload.ReservedPagesDeploymentType; got != want {
t.Errorf("deployment upload Type = %q, want %q", got, want)
}
if got, want := fmt.Sprint(uploadRecord.Metadata.Extra[pagesSourceIDMetadataKey]), fmt.Sprint(source.ID); got != want {
t.Errorf("deployment upload pages_source_id = %q, want %q", got, want)
}
secondSnapshot := mustAcquireRemoteSyncLease(t, ctx, source)
second, err := syncRemoteSource(ctx, secondSnapshot, "user:42")
if err != nil {
t.Fatalf("syncRemoteSource(second) error = %v, want nil", err)
}
if second == nil || second.Stale || !second.Reused || second.Deployment == nil {
t.Fatalf("syncRemoteSource(second) = %+v, want reused active deployment", second)
}
if got, want := second.Deployment.ID, first.Deployment.ID; got != want {
t.Errorf("reused deployment ID = %d, want %d", got, want)
}
var deploymentCount, uploadCount int64
if err := db.DB(ctx).Model(&model.PagesDeployment{}).Where("project_id = ?", project.ID).Count(&deploymentCount).Error; err != nil {
t.Fatalf("count source deployments error = %v, want nil", err)
}
if err := db.DB(ctx).Model(&model.Upload{}).Count(&uploadCount).Error; err != nil {
t.Fatalf("count source uploads error = %v, want nil", err)
}
if got, want := deploymentCount, int64(1); got != want {
t.Errorf("deployment count after identical sync = %d, want %d", got, want)
}
if got, want := uploadCount, int64(1); got != want {
t.Errorf("upload count after identical sync = %d, want %d", got, want)
}
}
func TestSyncRemoteSourceDownloadFailureKeepsOldActive(t *testing.T) {
ctx := setupPagesSourceSyncTest(t)
project := mustCreatePagesSourceProject(t, ctx, "sync-download-fail")
oldActive := mustCreateActiveManualDeployment(t, ctx, project.ID, "old-active")
server := newPagesArchiveServer(t, http.StatusBadGateway, []byte("upstream failed"), nil)
secret := "download-failure-secret"
source, _ := mustConfigureRemoteSource(
t,
ctx,
project.ID,
server.URL+"/site.zip?token="+secret,
RemoteNetworkPolicyTrustedInternal,
)
_, err := syncRemoteSource(ctx, mustAcquireRemoteSyncLease(t, ctx, source), "user:2")
if err == nil {
t.Fatal("syncRemoteSource(download failure) error = nil, want non-nil")
}
if strings.Contains(err.Error(), secret) {
t.Errorf("syncRemoteSource(download failure) error = %q, want no query secret", err)
}
assertPagesSyncFailureState(t, ctx, project.ID, source.ID, oldActive.ID, 1)
}
func TestSyncRemoteSourceArchiveFailureKeepsOldActive(t *testing.T) {
ctx := setupPagesSourceSyncTest(t)
project := mustCreatePagesSourceProject(t, ctx, "sync-archive-fail")
oldActive := mustCreateActiveManualDeployment(t, ctx, project.ID, "old-active")
server := newPagesArchiveServer(t, http.StatusOK, []byte("not-a-valid-zip"), nil)
secret := "archive-failure-secret"
source, _ := mustConfigureRemoteSource(
t,
ctx,
project.ID,
server.URL+"/site.zip?token="+secret,
RemoteNetworkPolicyTrustedInternal,
)
_, err := syncRemoteSource(ctx, mustAcquireRemoteSyncLease(t, ctx, source), "user:3")
if err == nil {
t.Fatal("syncRemoteSource(archive failure) error = nil, want non-nil")
}
if strings.Contains(err.Error(), secret) {
t.Errorf("syncRemoteSource(archive failure) error = %q, want no query secret", err)
}
assertPagesSyncFailureState(t, ctx, project.ID, source.ID, oldActive.ID, 1)
}
func TestSyncRemoteSourceFinalFenceCompensatesIngest(t *testing.T) {
ctx := setupPagesSourceSyncTest(t)
project := mustCreatePagesSourceProject(t, ctx, "sync-final-fence")
oldActive := mustCreateActiveManualDeployment(t, ctx, project.ID, "old-active")
packageBytes := testPagesZip(t, map[string]string{"index.html": "never-activate"})
mutationResult := make(chan error, 1)
server := newPagesArchiveServer(t, http.StatusOK, packageBytes, func() error {
err := db.DB(context.Background()).Model(&model.PagesProject{}).
Where("id = ?", project.ID).
Update("content_config_version", gorm.Expr("content_config_version + 1")).Error
mutationResult <- err
return err
})
source, _ := mustConfigureRemoteSource(
t,
ctx,
project.ID,
server.URL+"/site.zip?token=final-fence-secret",
RemoteNetworkPolicyTrustedInternal,
)
outcome, err := syncRemoteSource(ctx, mustAcquireRemoteSyncLease(t, ctx, source), "user:4")
if err != nil {
t.Fatalf("syncRemoteSource(final fence) error = %v, want nil stale outcome", err)
}
select {
case mutationErr := <-mutationResult:
if mutationErr != nil {
t.Fatalf("content version mutation error = %v, want nil", mutationErr)
}
case <-time.After(2 * time.Second):
t.Fatal("content version mutation was not observed")
}
if outcome == nil || !outcome.Stale || outcome.Deployment != nil {
t.Fatalf("syncRemoteSource(final fence) = %+v, want stale outcome without deployment", outcome)
}
storedProject, err := model.GetPagesProjectByID(ctx, project.ID)
if err != nil {
t.Fatalf("GetPagesProjectByID(%d) error = %v, want nil", project.ID, err)
}
if storedProject.ActiveDeploymentID == nil || *storedProject.ActiveDeploymentID != oldActive.ID {
t.Errorf("ActiveDeploymentID after final fence = %v, want %d", storedProject.ActiveDeploymentID, oldActive.ID)
}
var deployments []model.PagesDeployment
if err := db.DB(ctx).Where("project_id = ?", project.ID).Find(&deployments).Error; err != nil {
t.Fatalf("list deployments after final fence error = %v, want nil", err)
}
if got, want := len(deployments), 1; got != want {
t.Errorf("deployment count after final fence = %d, want %d", got, want)
}
var uploads []model.Upload
if err := db.DB(ctx).Order("id asc").Find(&uploads).Error; err != nil {
t.Fatalf("list uploads after final fence error = %v, want nil", err)
}
var compensated *model.Upload
for index := range uploads {
if fmt.Sprint(uploads[index].Metadata.Extra[pagesSourceIDMetadataKey]) == fmt.Sprint(source.ID) {
compensated = &uploads[index]
break
}
}
if compensated == nil {
t.Fatalf("source upload after final fence = nil, want compensated upload record")
}
if got, want := compensated.Status, model.UploadStatusDeleted; got != want {
t.Errorf("compensated upload Status = %q, want %q", got, want)
}
var danglingCount int64
if err := db.DB(ctx).Model(&model.PagesDeployment{}).
Where("upload_id = ?", compensated.ID).
Count(&danglingCount).Error; err != nil {
t.Fatalf("count compensated upload references error = %v, want nil", err)
}
if got, want := danglingCount, int64(0); got != want {
t.Errorf("deployments referencing compensated upload = %d, want %d", got, want)
}
}
func TestCommitSourceDeploymentRechecksLeaseAfterUploadLocks(t *testing.T) {
ctx := setupPagesSourceSyncTest(t)
project := mustCreatePagesSourceProject(t, ctx, "sync-expiry-recheck")
packageBytes := testPagesZip(t, map[string]string{"index.html": "expiry-recheck"})
server := newPagesArchiveServer(t, http.StatusOK, packageBytes, nil)
source, _ := mustConfigureRemoteSource(
t,
ctx,
project.ID,
server.URL+"/site.zip",
RemoteNetworkPolicyTrustedInternal,
)
first, err := syncRemoteSource(ctx, mustAcquireRemoteSyncLease(t, ctx, source), "user:5")
if err != nil || first == nil || first.Deployment == nil {
t.Fatalf("syncRemoteSource(seed) = (%+v, %v), want deployment", first, err)
}
deployment, err := model.GetPagesDeploymentByID(ctx, first.Deployment.ID)
if err != nil {
t.Fatalf("GetPagesDeploymentByID(%d) error = %v, want nil", first.Deployment.ID, err)
}
if err := db.DB(ctx).Model(&model.PagesProject{}).
Where("id = ?", project.ID).
Update("active_deployment_id", nil).Error; err != nil {
t.Fatalf("clear active deployment error = %v, want nil", err)
}
if err := db.DB(ctx).Model(&model.PagesDeployment{}).
Where("id = ?", deployment.ID).
Update("status", model.PagesDeploymentStatusUploaded).Error; err != nil {
t.Fatalf("reset deployment status error = %v, want nil", err)
}
snapshot := mustAcquireRemoteSyncLease(t, ctx, source)
expiresAt := time.Now().Add(time.Hour)
if err := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).
Where("source_id = ? AND lease_token = ?", source.ID, snapshot.LeaseToken).
Update("lease_expires_at", &expiresAt).Error; err != nil {
t.Fatalf("set deterministic lease expiry error = %v, want nil", err)
}
originalNow := sourceCommitNow
nowCalls := 0
sourceCommitNow = func() time.Time {
nowCalls++
if nowCalls == 1 {
return expiresAt.Add(-time.Second)
}
return expiresAt.Add(time.Second)
}
t.Cleanup(func() { sourceCommitNow = originalNow })
_, _, _, err = commitSourceDeployment(
ctx,
snapshot,
deployment.Checksum,
deployment.Checksum,
sourceDetail{Provider: PagesSourceTypeRemoteURL, DisplayName: deployment.SourceLabel},
deployment.SourceMeta,
"user:5",
&deploymentManifest{},
upload.IngestResult{},
false,
nil,
)
if !errors.Is(err, errSourceFinalFence) {
t.Fatalf("commitSourceDeployment(expired after upload lock) error = %v, want %v", err, errSourceFinalFence)
}
if nowCalls != 2 {
t.Fatalf("sourceCommitNow calls = %d, want runtime-lock and post-upload checks", nowCalls)
}
storedProject, err := model.GetPagesProjectByID(ctx, project.ID)
if err != nil {
t.Fatalf("GetPagesProjectByID(%d) error = %v, want nil", project.ID, err)
}
if storedProject.ActiveDeploymentID != nil {
t.Errorf("ActiveDeploymentID after expiry recheck = %v, want nil", storedProject.ActiveDeploymentID)
}
storedDeployment, err := model.GetPagesDeploymentByID(ctx, deployment.ID)
if err != nil {
t.Fatalf("GetPagesDeploymentByID(%d) after expiry error = %v, want nil", deployment.ID, err)
}
if got, want := storedDeployment.Status, model.PagesDeploymentStatusUploaded; got != want {
t.Errorf("deployment status after expiry recheck = %q, want %q", got, want)
}
}
func TestCompensateSourceIngestSurvivesCanceledParentContext(t *testing.T) {
ctx := setupPagesSourceSyncTest(t)
project := mustCreatePagesSourceProject(t, ctx, "sync-canceled-compensation")
source, _ := mustConfigureRemoteSource(
t,
ctx,
project.ID,
"https://example.com/site.zip",
RemoteNetworkPolicyPublic,
)
packageBytes := testPagesZip(t, map[string]string{"index.html": "cancel-compensation"})
packagePath := filepath.Join(t.TempDir(), "site.zip")
if err := os.WriteFile(packagePath, packageBytes, 0o600); err != nil {
t.Fatalf("write test package error = %v, want nil", err)
}
digest := sha256.Sum256(packageBytes)
result, err := ingestPagesDeploymentPackageWithSource(
ctx,
packagePath,
hex.EncodeToString(digest[:]),
project.ID,
source.ID,
"site.zip",
pagesarchive.FormatZip,
)
if err != nil {
t.Fatalf("ingestPagesDeploymentPackageWithSource() error = %v, want nil", err)
}
if !result.Created {
t.Fatal("ingest result Created = false, want a compensatable record")
}
canceledCtx, cancel := context.WithCancel(ctx)
cancel()
compensateSourceIngest(canceledCtx, &sourceExecutionSnapshot{
ProjectID: project.ID,
SourceID: source.ID,
}, &sourceIngestState{
Result: result,
HasIngest: true,
})
var uploadRecord model.Upload
if err := db.DB(ctx).Where("id = ?", result.Upload.ID).First(&uploadRecord).Error; err != nil {
t.Fatalf("load compensated upload error = %v, want nil", err)
}
if got, want := uploadRecord.Status, model.UploadStatusDeleted; got != want {
t.Errorf("compensated upload status = %q, want %q", got, want)
}
}
func assertPagesSyncFailureState(
t *testing.T,
ctx context.Context,
projectID uint,
sourceID uint,
oldActiveID uint,
wantDeploymentCount int64,
) {
t.Helper()
project, err := model.GetPagesProjectByID(ctx, projectID)
if err != nil {
t.Fatalf("GetPagesProjectByID(%d) error = %v, want nil", projectID, err)
}
if project.ActiveDeploymentID == nil || *project.ActiveDeploymentID != oldActiveID {
t.Errorf("project %d ActiveDeploymentID = %v, want %d", projectID, project.ActiveDeploymentID, oldActiveID)
}
var deploymentCount int64
if err := db.DB(ctx).Model(&model.PagesDeployment{}).
Where("project_id = ?", projectID).
Count(&deploymentCount).Error; err != nil {
t.Fatalf("count project %d deployments error = %v, want nil", projectID, err)
}
if got, want := deploymentCount, wantDeploymentCount; got != want {
t.Errorf("project %d deployment count = %d, want %d", projectID, got, want)
}
var runtime model.PagesProjectSourceRuntime
if err := db.DB(ctx).Where("source_id = ?", sourceID).First(&runtime).Error; err != nil {
t.Fatalf("load source %d runtime error = %v, want nil", sourceID, err)
}
if got, want := runtime.SyncStatus, pagesSourceStatusFailed; got != want {
t.Errorf("source %d runtime SyncStatus = %q, want %q", sourceID, got, want)
}
if runtime.LeaseToken != "" || runtime.LeaseExpiresAt != nil {
t.Errorf("source %d runtime lease = (%q, %v), want cleared", sourceID, runtime.LeaseToken, runtime.LeaseExpiresAt)
}
}
func TestCommitSourceDeploymentRejectsDeletedTargetUpload(t *testing.T) {
ctx := setupPagesSourceSyncTest(t)
project := mustCreatePagesSourceProject(t, ctx, "sync-deleted-upload")
packageBytes := testPagesZip(t, map[string]string{"index.html": "content"})
server := newPagesArchiveServer(t, http.StatusOK, packageBytes, nil)
source, _ := mustConfigureRemoteSource(
t,
ctx,
project.ID,
server.URL+"/site.zip",
RemoteNetworkPolicyTrustedInternal,
)
snapshot := mustAcquireRemoteSyncLease(t, ctx, source)
// A pre-existing source deployment whose upload was removed must never be
// reactivated into a dangling active pointer.
identity := source.SourceIdentity
revision := strings.Repeat("d", 64)
uploadRecord := &model.Upload{
ID: 987654321,
UserID: 999,
FileName: "deleted.zip",
FilePath: "deleted.zip",
FileSize: 1,
MimeType: "application/zip",
Extension: "zip",
Hash: revision,
Type: upload.ReservedPagesDeploymentType,
Status: model.UploadStatusDeleted,
AccessMode: 0,
}
if err := db.DB(ctx).Create(uploadRecord).Error; err != nil {
t.Fatalf("create deleted upload error = %v, want nil", err)
}
deployment := &model.PagesDeployment{
ProjectID: project.ID,
DeploymentNumber: 1,
Checksum: revision,
Status: model.PagesDeploymentStatusUploaded,
UploadID: uploadRecord.ID,
FileCount: 1,
TotalSize: 1,
CreatedBy: "user:1",
SourceType: PagesSourceTypeRemoteURL,
SourceIdentity: &identity,
SourceRevision: &revision,
SourceLabel: "deleted.zip",
SourceMeta: `{"provider":"remote_url","display_name":"deleted.zip"}`,
TriggerType: pagesSourceTriggerManualSync,
}
if err := db.DB(ctx).Create(deployment).Error; err != nil {
t.Fatalf("create source deployment error = %v, want nil", err)
}
if err := db.DB(ctx).Create(&model.PagesDeploymentFile{
DeploymentID: deployment.ID,
Path: "index.html",
Size: 1,
Checksum: revision,
}).Error; err != nil {
t.Fatalf("create source deployment file error = %v, want nil", err)
}
manifest := &deploymentManifest{
FileCount: 1,
TotalSize: 1,
EntryFile: "index.html",
}
_, _, _, err := commitSourceDeployment(
ctx,
snapshot,
revision,
revision,
sourceDetail{Provider: PagesSourceTypeRemoteURL, DisplayName: "deleted.zip"},
`{"provider":"remote_url","display_name":"deleted.zip"}`,
"user:1",
manifest,
upload.IngestResult{},
false,
nil,
)
if !errors.Is(err, errSourceFinalFence) {
t.Errorf("commitSourceDeployment(deleted upload) error = %v, want %v", err, errSourceFinalFence)
}
storedProject, err := model.GetPagesProjectByID(ctx, project.ID)
if err != nil {
t.Fatalf("GetPagesProjectByID(%d) error = %v, want nil", project.ID, err)
}
if storedProject.ActiveDeploymentID != nil {
t.Errorf("ActiveDeploymentID after deleted upload rejection = %v, want nil", storedProject.ActiveDeploymentID)
}
var activeCount int64
if err := db.DB(ctx).Model(&model.PagesDeployment{}).
Where("project_id = ? AND status = ?", project.ID, model.PagesDeploymentStatusActive).
Count(&activeCount).Error; err != nil && !errors.Is(err, gorm.ErrRecordNotFound) {
t.Fatalf("count active deployments error = %v, want nil", err)
}
if got, want := activeCount, int64(0); got != want {
t.Errorf("active deployment count after deleted upload rejection = %d, want %d", got, want)
}
}
@@ -0,0 +1,431 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pages
import (
"bytes"
"context"
"encoding/hex"
"encoding/json"
"errors"
"io"
"strconv"
"strings"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/task"
"github.com/Rain-kl/Wavelet/pkg/logger"
"gorm.io/gorm"
)
const (
// PagesSourceActionTask is the private Asynq task type for source actions.
PagesSourceActionTask = "openflare:pages_source_action"
// TaskTypePagesSourceAction is the internal task meta type.
TaskTypePagesSourceAction = "of_pages_source_action"
sourceActionCheck = "check"
sourceActionSync = "sync"
)
var errUnexpectedJSONTrailingValue = errors.New("unexpected trailing JSON value")
// PagesSourceActionMeta is intentionally hidden from generic Admin task APIs.
var PagesSourceActionMeta = task.TaskMeta{
Type: TaskTypePagesSourceAction,
AsynqTask: PagesSourceActionTask,
Name: "OpenFlare Pages 部署源操作",
Description: "检查或同步 Pages 项目部署源",
SupportsTime: false,
MaxRetry: 2,
Queue: task.QueueDefault,
Retryable: false,
InternalOnly: true,
}
// SourceActionPayload is the credential-free internal queue contract.
type SourceActionPayload struct {
SourceID uint `json:"source_id"`
ConfigVersion int `json:"config_version"`
Action string `json:"action"`
Actor string `json:"actor"`
TriggerType string `json:"trigger_type"`
TargetRevision string `json:"target_revision"`
ConfirmedRevision string `json:"confirmed_revision"`
}
// SourceActionHandler executes a validated source action.
type SourceActionHandler struct{}
// ValidatePayload rejects unknown keys and normalizes the internal contract.
func (h *SourceActionHandler) ValidatePayload(payload []byte) ([]byte, error) {
var input SourceActionPayload
decoder := json.NewDecoder(bytes.NewReader(payload))
decoder.DisallowUnknownFields()
if err := decoder.Decode(&input); err != nil {
return nil, errors.New(errPagesSourceActionInvalid)
}
if err := ensureJSONEOF(decoder); err != nil {
return nil, errors.New(errPagesSourceActionInvalid)
}
input.Action = strings.TrimSpace(input.Action)
input.Actor = strings.TrimSpace(input.Actor)
input.TriggerType = strings.TrimSpace(input.TriggerType)
input.TargetRevision = strings.TrimSpace(input.TargetRevision)
input.ConfirmedRevision = strings.TrimSpace(input.ConfirmedRevision)
if input.Action == sourceActionSync && input.TriggerType == "" {
// Keep already queued Phase 2 payloads valid while making every new
// dispatch carry an explicit deployment trigger.
input.TriggerType = pagesSourceTriggerManualSync
}
if !validSourceActionPayload(input) {
return nil, errors.New(errPagesSourceActionInvalid)
}
return json.Marshal(input)
}
func validSourceActionPayload(input SourceActionPayload) bool {
if input.SourceID == 0 || input.ConfigVersion <= 0 {
return false
}
if input.Action != sourceActionCheck && input.Action != sourceActionSync {
return false
}
if !validPagesSourceActor(input.Actor) {
return false
}
if !validOptionalSourceRevision(input.TargetRevision) ||
!validOptionalSourceRevision(input.ConfirmedRevision) {
return false
}
if input.Action == sourceActionCheck {
return input.TriggerType == "" && input.TargetRevision == "" && input.ConfirmedRevision == ""
}
return validSourceSyncPayload(input)
}
func validSourceSyncPayload(input SourceActionPayload) bool {
if !validSourceDeploymentTrigger(input.TriggerType) ||
(input.TargetRevision != "" && input.ConfirmedRevision != "") {
return false
}
switch input.TriggerType {
case pagesSourceTriggerScheduledAutoUpdate:
return input.Actor == pagesSourceCreatedBySystem &&
input.TargetRevision != "" && input.ConfirmedRevision == ""
case pagesSourceTriggerManualSync:
if input.TargetRevision != "" || !strings.HasPrefix(input.Actor, "user:") {
return false
}
return true
default:
return false
}
}
// Execute validates again inside the worker and performs the source action.
func (h *SourceActionHandler) Execute(ctx context.Context, payload []byte) (*task.TaskResult, error) {
normalized, err := h.ValidatePayload(payload)
if err != nil {
return nil, task.PermanentError(errPagesSourceActionInvalid)
}
var input SourceActionPayload
if err := json.Unmarshal(normalized, &input); err != nil {
return nil, task.PermanentError(errPagesSourceActionInvalid)
}
var source model.PagesProjectSource
if err := db.DB(ctx).Where("id = ?", input.SourceID).First(&source).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
task.AppendLog(ctx, "[resolve] 来源已不存在,本次任务跳过")
return &task.TaskResult{Message: errPagesSourceActionStale}, nil
}
logger.ErrorF(ctx, "[PagesSource] load source failed: source_id=%d error=%v", input.SourceID, err)
return nil, errors.New(errPagesSourceSyncFailed)
}
if source.ConfigVersion != input.ConfigVersion {
task.AppendLog(ctx, "[resolve] 来源配置已变化,本次任务跳过")
return &task.TaskResult{Message: errPagesSourceActionStale}, nil
}
if input.Action == sourceActionCheck && source.SourceType == PagesSourceTypeRemoteURL {
return nil, task.PermanentError(errPagesSourceCheckUnsupported)
}
if source.SourceType != PagesSourceTypeRemoteURL && source.SourceType != PagesSourceTypeGitHubRelease {
return nil, task.PermanentError(errPagesSourceTypeUnsupported)
}
if source.SourceType == PagesSourceTypeRemoteURL && (input.TargetRevision != "" || input.ConfirmedRevision != "") {
return nil, task.PermanentError(errPagesSourceActionInvalid)
}
if input.Action == sourceActionCheck && (input.TargetRevision != "" || input.ConfirmedRevision != "") {
return nil, task.PermanentError(errPagesSourceActionInvalid)
}
task.AppendLog(ctx, "[resolve] 正在获取来源执行权")
snapshot, outcome, err := acquireSourceLease(ctx, input.SourceID, input.ConfigVersion, input.Action)
if err != nil {
logger.ErrorF(ctx, "[PagesSource] acquire lease failed: source_id=%d error=%v", input.SourceID, err)
return nil, errors.New(errPagesSourceSyncFailed)
}
switch outcome {
case sourceLeaseBusy:
task.AppendLog(ctx, "[resolve] 已有来源任务正在执行,本次任务跳过")
return &task.TaskResult{Message: errPagesSourceActionBusy}, nil
case sourceLeaseStale:
task.AppendLog(ctx, "[resolve] 来源配置或执行权已变化,本次任务跳过")
return &task.TaskResult{Message: errPagesSourceActionStale}, nil
}
if input.Action == sourceActionCheck {
return executeGitHubCheckAction(ctx, snapshot)
}
return executeSourceSyncAction(ctx, &source, snapshot, input)
}
func executeGitHubCheckAction(ctx context.Context, snapshot *sourceExecutionSnapshot) (*task.TaskResult, error) {
checkResult, checkErr := checkGitHubSource(ctx, snapshot)
if checkErr != nil {
logger.ErrorF(ctx, "[PagesSource] check failed: project_id=%d source_id=%d error=%v", snapshot.ProjectID, snapshot.SourceID, checkErr)
if isPermanentSourceSyncError(checkErr) || shouldSkipGitHubActionRetry(checkErr) {
return nil, task.PermanentError(checkErr.Error())
}
return nil, errors.New(errPagesSourceSyncFailed)
}
if checkResult == nil || checkResult.Stale {
return &task.TaskResult{Message: errPagesSourceActionStale}, nil
}
return &task.TaskResult{Message: checkResult.Message, Detail: checkResult.Detail}, nil
}
func executeSourceSyncAction(
ctx context.Context,
source *model.PagesProjectSource,
snapshot *sourceExecutionSnapshot,
input SourceActionPayload,
) (*task.TaskResult, error) {
var result *sourceSyncOutcome
var err error
if source.SourceType == PagesSourceTypeGitHubRelease {
result, err = syncGitHubSourceWithTrigger(
ctx, snapshot, input.Actor, input.TargetRevision, input.ConfirmedRevision, input.TriggerType,
)
} else {
result, err = syncRemoteSourceWithTrigger(ctx, snapshot, input.Actor, input.TriggerType)
}
if err != nil {
logger.ErrorF(ctx, "[PagesSource] sync failed: project_id=%d source_id=%d error=%v", snapshot.ProjectID, snapshot.SourceID, err)
if isPermanentSourceSyncError(err) || shouldSkipGitHubActionRetry(err) {
return nil, task.PermanentError(errPagesSourceSyncFailed)
}
return nil, errors.New(errPagesSourceSyncFailed)
}
if result == nil || result.Stale {
task.AppendLog(ctx, "[activate] 来源配置或执行权已变化,本次任务未切换部署")
return &task.TaskResult{Message: errPagesSourceActionStale}, nil
}
message := "Pages 部署源同步并发布成功"
if result.Reused {
message = "Pages 部署源内容未变化,已重新激活现有部署"
}
return &task.TaskResult{Message: message, Detail: sourceSyncResultDetail(result)}, nil
}
func ensureJSONEOF(decoder *json.Decoder) error {
var trailing any
err := decoder.Decode(&trailing)
if errors.Is(err, io.EOF) {
return nil
}
if err == nil {
return errUnexpectedJSONTrailingValue
}
return err
}
func validPagesSourceActor(actor string) bool {
if actor == pagesSourceCreatedBySystem {
return true
}
if !strings.HasPrefix(actor, "user:") {
return false
}
id, err := strconv.ParseUint(strings.TrimPrefix(actor, "user:"), 10, 64)
return err == nil && id > 0
}
func validOptionalSourceRevision(value string) bool {
if value == "" {
return true
}
if len(value) != sourceRevisionHexLength {
return false
}
decoded, err := hex.DecodeString(value)
return err == nil && len(decoded) == 32
}
func isPermanentSourceSyncError(err error) bool {
if err == nil {
return false
}
message := err.Error()
return strings.Contains(message, errPagesPackageUnsupported) ||
strings.Contains(message, errPagesPackageURLTooLarge) ||
strings.Contains(message, errPagesPackageInvalid) ||
strings.Contains(message, errPagesPackageEmpty) ||
strings.Contains(message, errPagesPackageExtractedTooLarge) ||
strings.Contains(message, errPagesPackageFileTooLarge) ||
strings.Contains(message, errPagesEntryFileMissing) ||
strings.Contains(message, errPagesSourceRemoteURLInvalid) ||
strings.Contains(message, errPagesSourceNetworkPolicy) ||
strings.Contains(message, errPagesSourceReleaseNotFound) ||
strings.Contains(message, errPagesSourceDigestInvalid) ||
strings.Contains(message, errPagesSourceDigestMismatch) ||
strings.Contains(message, errPagesSourceConfirmationNeeded) ||
strings.Contains(message, errPagesSourceConfirmationStale)
}
// DispatchSourceAction performs API preflight and enqueues a credential-free action.
func DispatchSourceAction(
ctx context.Context,
projectID uint,
action string,
actor string,
confirmedRevision string,
) (*SourceActionReceipt, error) {
return dispatchSourceActionByProject(ctx, projectID, action, actor, "", confirmedRevision)
}
func dispatchSourceActionByProject(
ctx context.Context,
projectID uint,
action string,
actor string,
targetRevision string,
confirmedRevision string,
) (*SourceActionReceipt, error) {
action = strings.TrimSpace(action)
targetRevision = strings.TrimSpace(targetRevision)
confirmedRevision = strings.TrimSpace(confirmedRevision)
if action != sourceActionCheck && action != sourceActionSync {
return nil, errors.New(errPagesSourceActionInvalid)
}
if !validPagesSourceActor(actor) {
return nil, errors.New(errPagesSourceActionInvalid)
}
var source model.PagesProjectSource
if err := db.DB(ctx).Where("project_id = ?", projectID).First(&source).Error; err != nil {
if errors.Is(err, gorm.ErrRecordNotFound) {
return nil, errors.New(errPagesSourceNotFound)
}
return nil, err
}
if err := validateSourceActionPreflight(ctx, &source, action, targetRevision, confirmedRevision); err != nil {
return nil, err
}
busy, err := sourceLeaseIsBusy(ctx, source.ID)
if err != nil {
return nil, err
}
if busy {
return nil, errors.New(errPagesSourceActionBusy)
}
return dispatchSourceActionSnapshot(ctx, source, action, actor, targetRevision, confirmedRevision, "manual")
}
func validateSourceActionPreflight(
ctx context.Context,
source *model.PagesProjectSource,
action string,
targetRevision string,
confirmedRevision string,
) error {
if source == nil {
return errors.New(errPagesSourceNotFound)
}
if source.SourceType != PagesSourceTypeRemoteURL && source.SourceType != PagesSourceTypeGitHubRelease {
return errors.New(errPagesSourceTypeUnsupported)
}
if action == sourceActionCheck && source.SourceType == PagesSourceTypeRemoteURL {
return errors.New(errPagesSourceCheckUnsupported)
}
if source.SourceType == PagesSourceTypeRemoteURL && (targetRevision != "" || confirmedRevision != "") {
return errors.New(errPagesSourceActionInvalid)
}
if action == sourceActionCheck && (targetRevision != "" || confirmedRevision != "") {
return errors.New(errPagesSourceActionInvalid)
}
if source.SourceType == PagesSourceTypeGitHubRelease && action == sourceActionSync {
if err := preflightGitHubSyncConfirmation(ctx, source.ID, confirmedRevision); err != nil {
return err
}
}
return nil
}
func dispatchSourceActionSnapshot(
ctx context.Context,
source model.PagesProjectSource,
action string,
actor string,
targetRevision string,
confirmedRevision string,
triggeredBy string,
) (*SourceActionReceipt, error) {
triggerType := ""
if action == sourceActionSync {
triggerType = pagesSourceTriggerManualSync
}
return dispatchSourceActionSnapshotWithTrigger(
ctx, source, action, actor, triggerType, targetRevision, confirmedRevision, triggeredBy,
)
}
func dispatchSourceActionSnapshotWithTrigger(
ctx context.Context,
source model.PagesProjectSource,
action string,
actor string,
triggerType string,
targetRevision string,
confirmedRevision string,
triggeredBy string,
) (*SourceActionReceipt, error) {
if task.AsynqClient == nil {
return nil, errors.New(errPagesSourceTaskDispatchFailed)
}
handler := &SourceActionHandler{}
rawPayload, err := json.Marshal(SourceActionPayload{
SourceID: source.ID,
ConfigVersion: source.ConfigVersion,
Action: action,
Actor: actor,
TriggerType: triggerType,
TargetRevision: targetRevision,
ConfirmedRevision: confirmedRevision,
})
if err != nil {
return nil, errors.New(errPagesSourceActionInvalid)
}
payload, err := handler.ValidatePayload(rawPayload)
if err != nil {
return nil, err
}
taskID, err := task.DispatchTask(ctx, TaskTypePagesSourceAction, payload, triggeredBy)
if err != nil {
logger.ErrorF(ctx, "[PagesSource] dispatch action failed: project_id=%d source_id=%d action=%s error=%v", source.ProjectID, source.ID, action, err)
return nil, errors.New(errPagesSourceTaskDispatchFailed)
}
execution, err := model.GetTaskExecutionByTaskID(ctx, taskID)
if err != nil {
logger.ErrorF(ctx, "[PagesSource] load dispatched execution failed: source_id=%d task_id=%s error=%v", source.ID, taskID, err)
return nil, errors.New(errPagesSourceTaskDispatchFailed)
}
return &SourceActionReceipt{
TaskID: taskID,
ExecutionID: strconv.FormatUint(execution.ID, 10),
Action: action,
}, nil
}
@@ -0,0 +1,184 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pages
import (
"encoding/json"
"errors"
"strings"
"testing"
"github.com/hibiken/asynq"
)
func TestSourceActionPayloadValidationIsStrictAndCredentialFree(t *testing.T) {
handler := &SourceActionHandler{}
valid := SourceActionPayload{
SourceID: 7,
ConfigVersion: 3,
Action: sourceActionSync,
Actor: "user:42",
TriggerType: pagesSourceTriggerManualSync,
}
raw, err := json.Marshal(valid)
if err != nil {
t.Fatalf("json.Marshal(valid payload) error = %v, want nil", err)
}
normalized, err := handler.ValidatePayload(raw)
if err != nil {
t.Fatalf("ValidatePayload(valid) error = %v, want nil", err)
}
var got SourceActionPayload
if err := json.Unmarshal(normalized, &got); err != nil {
t.Fatalf("json.Unmarshal(normalized payload) error = %v, want nil", err)
}
if got != valid {
t.Errorf("ValidatePayload(valid) = %+v, want %+v", got, valid)
}
legacy := valid
legacy.TriggerType = ""
legacyRaw, err := json.Marshal(legacy)
if err != nil {
t.Fatalf("json.Marshal(legacy payload) error = %v, want nil", err)
}
legacyNormalized, err := handler.ValidatePayload(legacyRaw)
if err != nil {
t.Fatalf("ValidatePayload(legacy payload) error = %v, want nil", err)
}
var legacyGot SourceActionPayload
if err := json.Unmarshal(legacyNormalized, &legacyGot); err != nil {
t.Fatalf("json.Unmarshal(legacy normalized payload) error = %v, want nil", err)
}
if legacyGot.TriggerType != pagesSourceTriggerManualSync {
t.Errorf("legacy payload trigger_type = %q, want %q", legacyGot.TriggerType, pagesSourceTriggerManualSync)
}
for _, forbidden := range []string{"remote_url", "content_config_version", "expected_revision", "lease_token", "etag"} {
if strings.Contains(string(normalized), forbidden) {
t.Errorf("normalized payload = %s, want no forbidden field %q", normalized, forbidden)
}
}
invalidPayloads := []struct {
name string
raw string
}{
{
name: "unknown remote URL field",
raw: `{"source_id":7,"config_version":3,"action":"sync","actor":"user:42","remote_url":"https://example.com/site.zip?token=secret"}`,
},
{
name: "unknown content version field",
raw: `{"source_id":7,"config_version":3,"action":"sync","actor":"user:42","content_config_version":9}`,
},
{
name: "empty actor",
raw: `{"source_id":7,"config_version":3,"action":"sync","actor":""}`,
},
{
name: "untrusted system actor",
raw: `{"source_id":7,"config_version":3,"action":"sync","actor":"system"}`,
},
{
name: "zero user actor",
raw: `{"source_id":7,"config_version":3,"action":"sync","actor":"user:0"}`,
},
{
name: "invalid action",
raw: `{"source_id":7,"config_version":3,"action":"activate","actor":"user:42"}`,
},
{
name: "multiple JSON values",
raw: `{"source_id":7,"config_version":3,"action":"sync","actor":"user:42"} {}`,
},
}
for _, test := range invalidPayloads {
t.Run(test.name, func(t *testing.T) {
normalized, err := handler.ValidatePayload([]byte(test.raw))
if err == nil {
t.Errorf("ValidatePayload(%s) = %s, nil; want non-nil error", test.raw, normalized)
}
if err != nil && strings.Contains(err.Error(), "secret") {
t.Errorf("ValidatePayload(%s) error = %q, want credential-free error", test.name, err)
}
})
}
}
func TestSourceActionPayloadAcceptsOnlyRealActors(t *testing.T) {
tests := []struct {
actor string
want bool
}{
{actor: "user:1", want: true},
{actor: "user:18446744073709551615", want: true},
{actor: pagesSourceCreatedBySystem, want: true},
{actor: "", want: false},
{actor: "user:0", want: false},
{actor: "user:-1", want: false},
{actor: "user:not-a-number", want: false},
{actor: "system", want: false},
}
for _, test := range tests {
if got := validPagesSourceActor(test.actor); got != test.want {
t.Errorf("validPagesSourceActor(%q) = %t, want %t", test.actor, got, test.want)
}
}
}
func TestRemoteCheckActionIsPermanentWithoutExposingURL(t *testing.T) {
ctx := setupPagesSourceTest(t)
project := mustCreatePagesSourceProject(t, ctx, "task-remote-check")
secret := "task-query-secret"
source, _ := mustConfigureRemoteSource(
t,
ctx,
project.ID,
"https://example.com/site.zip?token="+secret,
RemoteNetworkPolicyPublic,
)
raw, err := json.Marshal(SourceActionPayload{
SourceID: source.ID,
ConfigVersion: source.ConfigVersion,
Action: sourceActionCheck,
Actor: "user:9",
})
if err != nil {
t.Fatalf("json.Marshal(check payload) error = %v, want nil", err)
}
result, err := (&SourceActionHandler{}).Execute(ctx, raw)
if result != nil {
t.Errorf("SourceActionHandler.Execute(remote check) result = %+v, want nil", result)
}
if err == nil {
t.Fatal("SourceActionHandler.Execute(remote check) error = nil, want permanent error")
}
if !errors.Is(err, asynq.SkipRetry) {
t.Errorf("SourceActionHandler.Execute(remote check) error = %v, want errors.Is(asynq.SkipRetry)", err)
}
if got, want := err.Error(), errPagesSourceCheckUnsupported; got != want {
t.Errorf("SourceActionHandler.Execute(remote check) error = %q, want %q", got, want)
}
if strings.Contains(err.Error(), secret) || strings.Contains(string(raw), secret) {
t.Errorf("remote check result error/payload = %q / %s, want no URL secret", err, raw)
}
}
func TestPagesSourceActionMetaIsInternalOnly(t *testing.T) {
if !PagesSourceActionMeta.InternalOnly {
t.Error("PagesSourceActionMeta.InternalOnly = false, want true")
}
if PagesSourceActionMeta.Type != TaskTypePagesSourceAction {
t.Errorf("PagesSourceActionMeta.Type = %q, want %q", PagesSourceActionMeta.Type, TaskTypePagesSourceAction)
}
if PagesSourceActionMeta.AsynqTask != PagesSourceActionTask {
t.Errorf("PagesSourceActionMeta.AsynqTask = %q, want %q", PagesSourceActionMeta.AsynqTask, PagesSourceActionTask)
}
if PagesSourceActionMeta.Retryable {
t.Error("PagesSourceActionMeta.Retryable = true, want false for manual retry API")
}
if PagesSourceActionMeta.MaxRetry <= 0 {
t.Errorf("PagesSourceActionMeta.MaxRetry = %d, want bounded transient retries", PagesSourceActionMeta.MaxRetry)
}
}
@@ -0,0 +1,413 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package pages
import (
"context"
"encoding/json"
"fmt"
"strings"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
)
func setupPagesSourceTest(t *testing.T) context.Context {
t.Helper()
cleanup := setupPagesTestDB(t)
t.Cleanup(cleanup)
sqlDB, err := db.DB(t.Context()).DB()
if err != nil {
t.Fatalf("db.DB().DB() error = %v, want nil", err)
}
// SQLite :memory: is scoped to one connection. Keeping one connection also
// makes lease tests exercise the production CAS without creating empty
// per-connection databases.
sqlDB.SetMaxOpenConns(1)
return t.Context()
}
func TestRevisionViewReadsLegacySourceDetailLabel(t *testing.T) {
tests := []struct {
name string
detail string
want string
}{
{
name: "remote",
detail: `{"provider":"remote_url","label":"legacy.zip"}`,
want: "legacy.zip",
},
{
name: "github",
detail: `{"provider":"github","label":"v1.2.3","asset_name":"dist.zip"}`,
want: "v1.2.3",
},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
view := revisionView(strings.Repeat("a", 64), test.detail)
if view.Label != test.want {
t.Errorf("revisionView(%s).Label = %q, want %q", test.name, view.Label, test.want)
}
})
}
}
func mustCreatePagesSourceProject(t *testing.T, ctx context.Context, slug string) *model.PagesProject {
t.Helper()
view, err := CreateProject(ctx, Input{
Name: "Source " + slug,
Slug: slug,
Enabled: true,
EntryFile: "index.html",
})
if err != nil {
t.Fatalf("CreateProject(%q) error = %v, want nil", slug, err)
}
project, err := model.GetPagesProjectByID(ctx, view.ID)
if err != nil {
t.Fatalf("GetPagesProjectByID(%d) error = %v, want nil", view.ID, err)
}
return project
}
func mustConfigureRemoteSource(
t *testing.T,
ctx context.Context,
projectID uint,
remoteURL string,
policy string,
) (*model.PagesProjectSource, *model.PagesProjectSourceRuntime) {
t.Helper()
_, err := UpdateSource(ctx, projectID, SourceUpdateInput{
SourceType: PagesSourceTypeRemoteURL,
RemoteURLSet: true,
RemoteURL: remoteURL,
RemoteNetworkPolicy: policy,
})
if err != nil {
t.Fatalf("UpdateSource(%d, %q) error = %v, want nil", projectID, remoteURL, err)
}
var source model.PagesProjectSource
if err := db.DB(ctx).Where("project_id = ?", projectID).First(&source).Error; err != nil {
t.Fatalf("load source for project %d error = %v, want nil", projectID, err)
}
var runtime model.PagesProjectSourceRuntime
if err := db.DB(ctx).Where("source_id = ?", source.ID).First(&runtime).Error; err != nil {
t.Fatalf("load runtime for source %d error = %v, want nil", source.ID, err)
}
return &source, &runtime
}
func TestValidateRemoteSourceInputRejectsModeIncompatibleFields(t *testing.T) {
tests := []struct {
name string
input SourceUpdateInput
}{
{
name: "missing source type",
input: SourceUpdateInput{
RemoteURLSet: true,
RemoteURL: "https://example.com/site.zip",
},
},
{
name: "github type reserved for phase two",
input: SourceUpdateInput{
SourceType: PagesSourceTypeGitHubRelease,
RepositoryURL: "https://github.com/example/site",
},
},
{
name: "remote rejects repository field",
input: SourceUpdateInput{
SourceType: PagesSourceTypeRemoteURL,
RemoteURLSet: true,
RemoteURL: "https://example.com/site.zip",
RepositoryURL: "https://github.com/example/site",
},
},
{
name: "remote rejects automatic updates",
input: SourceUpdateInput{
SourceType: PagesSourceTypeRemoteURL,
RemoteURLSet: true,
RemoteURL: "https://example.com/site.zip",
AutoUpdateEnabled: true,
},
},
{
name: "url value requires replacement flag",
input: SourceUpdateInput{
SourceType: PagesSourceTypeRemoteURL,
RemoteURL: "https://example.com/site.zip",
},
},
{
name: "invalid network policy",
input: SourceUpdateInput{
SourceType: PagesSourceTypeRemoteURL,
RemoteURLSet: true,
RemoteURL: "https://example.com/site.zip",
RemoteNetworkPolicy: "private",
},
},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
if err := validateRemoteSourceInput(test.input); err == nil {
t.Errorf("validateRemoteSourceInput(%+v) error = nil, want non-nil", test.input)
}
})
}
}
func TestUpdateSourceNewRemoteRequiresExplicitURL(t *testing.T) {
ctx := setupPagesSourceTest(t)
project := mustCreatePagesSourceProject(t, ctx, "remote-requires-url")
_, err := UpdateSource(ctx, project.ID, SourceUpdateInput{
SourceType: PagesSourceTypeRemoteURL,
})
if err == nil {
t.Fatal("UpdateSource(new remote without URL) error = nil, want non-nil")
}
if got, want := err.Error(), errPagesSourceRemoteURLRequired; got != want {
t.Errorf("UpdateSource(new remote without URL) error = %q, want %q", got, want)
}
}
func TestRemoteSourceCRUDPreservesSecretAndResetsRuntimeByIdentity(t *testing.T) {
ctx := setupPagesSourceTest(t)
project := mustCreatePagesSourceProject(t, ctx, "remote-crud")
firstURL := "https://Artifacts.Example.com:443/dist/site.zip?token=first-secret&expires=1"
source, runtime := mustConfigureRemoteSource(t, ctx, project.ID, firstURL, RemoteNetworkPolicyPublic)
if got, want := source.ConfigVersion, 1; got != want {
t.Errorf("new source ConfigVersion = %d, want %d", got, want)
}
if got, want := runtime.SyncStatus, pagesSourceStatusIdle; got != want {
t.Errorf("new runtime SyncStatus = %q, want %q", got, want)
}
view, err := GetSource(ctx, project.ID)
if err != nil {
t.Fatalf("GetSource(%d) error = %v, want nil", project.ID, err)
}
if got, want := view.DisplayURL, "https://Artifacts.Example.com:443/dist/site.zip?***"; got != want {
t.Errorf("GetSource(%d).DisplayURL = %q, want %q", project.ID, got, want)
}
encodedView, err := json.Marshal(view)
if err != nil {
t.Fatalf("json.Marshal(GetSource(%d)) error = %v, want nil", project.ID, err)
}
if strings.Contains(string(encodedView), "first-secret") || strings.Contains(string(encodedView), "expires=1") {
t.Errorf("GetSource(%d) JSON = %s, want credential-free view", project.ID, encodedView)
}
if _, err := UpdateSource(ctx, project.ID, SourceUpdateInput{
SourceType: PagesSourceTypeRemoteURL,
RemoteURLSet: true,
RemoteURL: firstURL,
RemoteNetworkPolicy: RemoteNetworkPolicyPublic,
}); err != nil {
t.Fatalf("UpdateSource(%d, no-op) error = %v, want nil", project.ID, err)
}
var unchangedSource model.PagesProjectSource
if err := db.DB(ctx).Where("id = ?", source.ID).First(&unchangedSource).Error; err != nil {
t.Fatalf("load no-op source error = %v, want nil", err)
}
if got, want := unchangedSource.ConfigVersion, source.ConfigVersion; got != want {
t.Errorf("no-op source ConfigVersion = %d, want unchanged %d", got, want)
}
seenRevision := strings.Repeat("a", 64)
appliedRevision := strings.Repeat("b", 64)
future := time.Now().Add(time.Hour)
if err := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).
Where("source_id = ?", source.ID).
Updates(map[string]any{
"last_seen_revision": seenRevision,
"last_seen_detail": `{"provider":"remote_url","display_name":"new.zip"}`,
"last_applied_revision": appliedRevision,
"last_applied_detail": `{"provider":"remote_url","display_name":"old.zip"}`,
"sync_status": pagesSourceStatusSyncing,
"lease_token": "in-flight",
"lease_expires_at": &future,
}).Error; err != nil {
t.Fatalf("seed source runtime error = %v, want nil", err)
}
// Omit the secret URL while changing policy. The stored URL and cursor must
// survive, while the in-flight lease is fenced.
if _, err := UpdateSource(ctx, project.ID, SourceUpdateInput{
SourceType: PagesSourceTypeRemoteURL,
RemoteURLSet: false,
RemoteNetworkPolicy: RemoteNetworkPolicyTrustedInternal,
}); err != nil {
t.Fatalf("UpdateSource(%d, preserve URL) error = %v, want nil", project.ID, err)
}
preservedSource, preservedRuntime, err := loadSourceByProject(ctx, project.ID)
if err != nil {
t.Fatalf("loadSourceByProject(%d) error = %v, want nil", project.ID, err)
}
if got, want := preservedSource.RemoteURL, firstURL; got != want {
t.Errorf("preserved RemoteURL = %q, want %q", got, want)
}
if got, want := preservedSource.ConfigVersion, 2; got != want {
t.Errorf("preserved source ConfigVersion = %d, want %d", got, want)
}
if got, want := preservedSource.SourceIdentity, source.SourceIdentity; got != want {
t.Errorf("preserved source identity = %q, want %q", got, want)
}
if got, want := preservedRuntime.LastSeenRevision, seenRevision; got != want {
t.Errorf("preserved LastSeenRevision = %q, want %q", got, want)
}
if got, want := preservedRuntime.SyncStatus, pagesSourceStatusUpdateAvailable; got != want {
t.Errorf("preserved runtime SyncStatus = %q, want %q", got, want)
}
if preservedRuntime.LeaseToken != "" || preservedRuntime.LeaseExpiresAt != nil {
t.Errorf("preserved runtime lease = (%q, %v), want cleared", preservedRuntime.LeaseToken, preservedRuntime.LeaseExpiresAt)
}
// Replacing only the query secret keeps the canonical identity and cursors.
queryReplacementURL := "https://artifacts.example.com/dist/site.zip?token=second-secret"
if _, err := UpdateSource(ctx, project.ID, SourceUpdateInput{
SourceType: PagesSourceTypeRemoteURL,
RemoteURLSet: true,
RemoteURL: queryReplacementURL,
RemoteNetworkPolicy: RemoteNetworkPolicyTrustedInternal,
}); err != nil {
t.Fatalf("UpdateSource(%d, query replacement) error = %v, want nil", project.ID, err)
}
querySource, queryRuntime, err := loadSourceByProject(ctx, project.ID)
if err != nil {
t.Fatalf("loadSourceByProject(%d) after query replacement error = %v, want nil", project.ID, err)
}
if got, want := querySource.SourceIdentity, source.SourceIdentity; got != want {
t.Errorf("query replacement identity = %q, want %q", got, want)
}
if got, want := queryRuntime.LastSeenRevision, seenRevision; got != want {
t.Errorf("query replacement LastSeenRevision = %q, want %q", got, want)
}
// Replacing the path changes identity and clears all remote cursors.
pathReplacementURL := "https://artifacts.example.com/dist/other.zip?token=third-secret"
if _, err := UpdateSource(ctx, project.ID, SourceUpdateInput{
SourceType: PagesSourceTypeRemoteURL,
RemoteURLSet: true,
RemoteURL: pathReplacementURL,
RemoteNetworkPolicy: RemoteNetworkPolicyTrustedInternal,
}); err != nil {
t.Fatalf("UpdateSource(%d, path replacement) error = %v, want nil", project.ID, err)
}
pathSource, pathRuntime, err := loadSourceByProject(ctx, project.ID)
if err != nil {
t.Fatalf("loadSourceByProject(%d) after path replacement error = %v, want nil", project.ID, err)
}
if pathSource.SourceIdentity == source.SourceIdentity {
t.Errorf("path replacement identity = %q, want a new identity", pathSource.SourceIdentity)
}
if pathRuntime.LastSeenRevision != "" || pathRuntime.LastAppliedRevision != "" {
t.Errorf("path replacement cursors = (%q, %q), want empty", pathRuntime.LastSeenRevision, pathRuntime.LastAppliedRevision)
}
if got, want := pathRuntime.SyncStatus, pagesSourceStatusIdle; got != want {
t.Errorf("path replacement SyncStatus = %q, want %q", got, want)
}
pathView, err := GetSource(ctx, project.ID)
if err != nil {
t.Fatalf("GetSource(%d) after path replacement error = %v, want nil", project.ID, err)
}
pathJSON, err := json.Marshal(pathView)
if err != nil {
t.Fatalf("json.Marshal(path view) error = %v, want nil", err)
}
for _, secret := range []string{"first-secret", "second-secret", "third-secret"} {
if strings.Contains(string(pathJSON), secret) {
t.Errorf("path view JSON = %s, want no secret %q", pathJSON, secret)
}
}
}
func TestRemoteSourceIdentityIgnoresQueryAndNormalizesDefaultPort(t *testing.T) {
first, err := parseRemoteSourceURL("HTTPS://Artifacts.Example.com:443/dist/../dist/site.zip?token=one")
if err != nil {
t.Fatalf("parseRemoteSourceURL(first) error = %v, want nil", err)
}
second, err := parseRemoteSourceURL("https://artifacts.example.com/dist/site.zip?token=two")
if err != nil {
t.Fatalf("parseRemoteSourceURL(second) error = %v, want nil", err)
}
if got, want := remoteSourceIdentity(first), remoteSourceIdentity(second); got != want {
t.Errorf("remoteSourceIdentity(first) = %q, want %q", got, want)
}
}
func TestDeleteSourceIsIdempotentAndKeepsDeploymentState(t *testing.T) {
ctx := setupPagesSourceTest(t)
project := mustCreatePagesSourceProject(t, ctx, "source-delete")
source, _ := mustConfigureRemoteSource(
t,
ctx,
project.ID,
"https://example.com/site.zip?token=delete-secret",
RemoteNetworkPolicyPublic,
)
deployment := &model.PagesDeployment{
ProjectID: project.ID,
DeploymentNumber: 1,
Checksum: strings.Repeat("c", 64),
Status: model.PagesDeploymentStatusActive,
CreatedBy: "user:1",
SourceType: "manual_upload",
TriggerType: "manual_upload",
}
if err := db.DB(ctx).Create(deployment).Error; err != nil {
t.Fatalf("create deployment error = %v, want nil", err)
}
if err := db.DB(ctx).Model(&model.PagesProject{}).
Where("id = ?", project.ID).
Update("active_deployment_id", deployment.ID).Error; err != nil {
t.Fatalf("set active deployment error = %v, want nil", err)
}
for attempt := 1; attempt <= 2; attempt++ {
view, err := DeleteSource(ctx, project.ID)
if err != nil {
t.Fatalf("DeleteSource(%d), attempt %d error = %v, want nil", project.ID, attempt, err)
}
if got, want := view.SourceType, PagesSourceTypeManual; got != want {
t.Errorf("DeleteSource(%d), attempt %d SourceType = %q, want %q", project.ID, attempt, got, want)
}
}
var sourceCount, runtimeCount, deploymentCount int64
if err := db.DB(ctx).Model(&model.PagesProjectSource{}).Where("id = ?", source.ID).Count(&sourceCount).Error; err != nil {
t.Fatalf("count source error = %v, want nil", err)
}
if err := db.DB(ctx).Model(&model.PagesProjectSourceRuntime{}).Where("source_id = ?", source.ID).Count(&runtimeCount).Error; err != nil {
t.Fatalf("count runtime error = %v, want nil", err)
}
if err := db.DB(ctx).Model(&model.PagesDeployment{}).Where("id = ?", deployment.ID).Count(&deploymentCount).Error; err != nil {
t.Fatalf("count deployment error = %v, want nil", err)
}
if sourceCount != 0 || runtimeCount != 0 || deploymentCount != 1 {
t.Errorf("DeleteSource counts = source:%d runtime:%d deployment:%d, want 0, 0, 1", sourceCount, runtimeCount, deploymentCount)
}
storedProject, err := model.GetPagesProjectByID(ctx, project.ID)
if err != nil {
t.Fatalf("GetPagesProjectByID(%d) error = %v, want nil", project.ID, err)
}
if storedProject.ActiveDeploymentID == nil || *storedProject.ActiveDeploymentID != deployment.ID {
t.Errorf("active deployment = %v, want %d", storedProject.ActiveDeploymentID, deployment.ID)
}
manual, err := GetSource(ctx, project.ID)
if err != nil {
t.Fatalf("GetSource(%d) after delete error = %v, want nil", project.ID, err)
}
if got, want := fmt.Sprint(manual.SourceType), PagesSourceTypeManual; got != want {
t.Errorf("GetSource(%d).SourceType = %q, want %q", project.ID, got, want)
}
}
@@ -6,12 +6,14 @@ package proxy_route
import (
"context"
"errors"
"sort"
"strings"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"gorm.io/gorm"
"gorm.io/gorm/clause"
)
// CustomHeaderInput 自定义响应头。
@@ -122,6 +124,9 @@ func CreateProxyRoute(ctx context.Context, input Input) (*View, error) {
return nil, err
}
if err = db.DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := lockPagesProjectsForRouteMutation(tx, 0, route); err != nil {
return err
}
if err := tx.Create(route).Error; err != nil {
return err
}
@@ -141,11 +146,15 @@ func UpdateProxyRoute(ctx context.Context, id uint, input Input) (*View, error)
if err != nil {
return nil, err
}
previousPagesProjectID := pagesProjectIDForRoute(route)
route, _, err = buildProxyRoute(ctx, route, input)
if err != nil {
return nil, err
}
if err = db.DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := lockPagesProjectsForRouteMutation(tx, previousPagesProjectID, route); err != nil {
return err
}
if err := updateProxyRouteRecord(tx, route); err != nil {
return err
}
@@ -159,6 +168,58 @@ func UpdateProxyRoute(ctx context.Context, id uint, input Input) (*View, error)
return buildProxyRouteView(ctx, route)
}
func pagesProjectIDForRoute(route *model.ProxyRoute) uint {
if route == nil || route.UpstreamType != proxyRouteUpstreamTypePages || route.PagesProjectID == nil {
return 0
}
return *route.PagesProjectID
}
func lockPagesProjectsForRouteMutation(tx *gorm.DB, previousProjectID uint, route *model.ProxyRoute) error {
nextProjectID := pagesProjectIDForRoute(route)
var projectIDs []uint
if previousProjectID != 0 {
projectIDs = append(projectIDs, previousProjectID)
}
if nextProjectID != 0 && nextProjectID != previousProjectID {
projectIDs = append(projectIDs, nextProjectID)
}
sort.Slice(projectIDs, func(i int, j int) bool { return projectIDs[i] < projectIDs[j] })
for _, projectID := range projectIDs {
var project model.PagesProject
err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).First(&project, projectID).Error
if errors.Is(err, gorm.ErrRecordNotFound) && projectID != nextProjectID {
continue
}
if errors.Is(err, gorm.ErrRecordNotFound) {
return errors.New(errProxyRoutePagesNotFound)
}
if err != nil {
return err
}
if projectID == nextProjectID {
if err := validateLockedPagesRouteProject(&project); err != nil {
return err
}
}
}
return nil
}
func validateLockedPagesRouteProject(project *model.PagesProject) error {
if project == nil {
return errors.New(errProxyRoutePagesNotFound)
}
if !project.Enabled {
return errors.New(errProxyRoutePagesDisabled)
}
if project.ActiveDeploymentID == nil || *project.ActiveDeploymentID == 0 {
return errors.New(errProxyRoutePagesNoDeploy)
}
return nil
}
// DeleteProxyRoute 删除代理规则。
func DeleteProxyRoute(ctx context.Context, id uint) error {
if _, err := model.GetProxyRouteByID(ctx, id); err != nil {
@@ -19,7 +19,14 @@ func setupProxyRouteTestDB(t *testing.T) func() {
t.Helper()
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{DisableForeignKeyConstraintWhenMigrating: true})
require.NoError(t, err)
require.NoError(t, sqliteDB.AutoMigrate(&model.ProxyRoute{}, &model.Origin{}, &model.Zone{}, &model.ZoneDomain{}, &model.TLSCertificate{}))
require.NoError(t, sqliteDB.AutoMigrate(
&model.ProxyRoute{},
&model.Origin{},
&model.Zone{},
&model.ZoneDomain{},
&model.TLSCertificate{},
&model.PagesProject{},
))
db.SetDB(sqliteDB)
return func() { db.SetDB(nil) }
}
@@ -81,6 +88,52 @@ func TestCreateProxyRouteHTTPSRequiresCoveringCertificate(t *testing.T) {
require.EqualError(t, err, errProxyRouteCertRequired)
}
func TestPagesRouteLocksAndRevalidatesTargetProject(t *testing.T) {
cleanup := setupProxyRouteTestDB(t)
defer cleanup()
ctx := context.Background()
domain := createZoneDomain(t, ctx, "pages.example.com", nil)
activeDeploymentID := uint(99)
project := &model.PagesProject{
Name: "Pages Site",
Slug: "pages-site",
Enabled: true,
ActiveDeploymentID: &activeDeploymentID,
}
require.NoError(t, db.DB(ctx).Create(project).Error)
view, err := CreateProxyRoute(ctx, Input{
SiteName: "pages",
ZoneDomainIDs: []uint{domain.ID},
UpstreamType: proxyRouteUpstreamTypePages,
PagesProjectID: &project.ID,
Enabled: true,
})
require.NoError(t, err)
require.NotNil(t, view.PagesProjectID)
assert.Equal(t, project.ID, *view.PagesProjectID)
require.NoError(t, db.DB(ctx).Delete(&model.PagesProject{}, project.ID).Error)
route := &model.ProxyRoute{UpstreamType: proxyRouteUpstreamTypePages, PagesProjectID: &project.ID}
err = db.DB(ctx).Transaction(func(tx *gorm.DB) error {
return lockPagesProjectsForRouteMutation(tx, 0, route)
})
require.EqualError(t, err, errProxyRoutePagesNotFound)
}
func TestRouteCanMoveAwayFromAlreadyMissingPagesProject(t *testing.T) {
cleanup := setupProxyRouteTestDB(t)
defer cleanup()
ctx := context.Background()
missingProjectID := uint(404)
route := &model.ProxyRoute{UpstreamType: "direct"}
err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
return lockPagesProjectsForRouteMutation(tx, missingProjectID, route)
})
require.NoError(t, err)
}
func TestNormalizeCachePolicyDefaultsAndLegacy(t *testing.T) {
assert.Equal(t, "", normalizeCachePolicy(false, "static"))
// Empty/url on write = legacy all (compat); UI sends static explicitly for new default.
+15 -9
View File
@@ -8,6 +8,7 @@ import (
"github.com/Rain-kl/Wavelet/internal/apps/upload/filesrv"
"github.com/Rain-kl/Wavelet/internal/apps/upload/handler"
"github.com/Rain-kl/Wavelet/internal/apps/upload/ingest"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
uploadtask "github.com/Rain-kl/Wavelet/internal/apps/upload/task"
"github.com/Rain-kl/Wavelet/internal/apps/upload/util"
@@ -31,15 +32,17 @@ var (
// Programmatic ingest API
var (
Ingest = ingest.Ingest
Remove = ingest.Remove
RemoveOwned = ingest.RemoveOwned
FindByHash = ingest.FindByHash
GetActiveUpload = ingest.GetActive
OpenStoredUpload = ingest.OpenActiveObject
ActiveUploadHash = ingest.ActiveHash
ResolveLocalFile = ingest.ResolveLocalFile
IngestFromLocalPath = ingest.FromLocalPath
Ingest = ingest.Ingest
Remove = ingest.Remove
RemoveOwned = ingest.RemoveOwned
RemoveLockedTx = ingest.RemoveLockedTx
InvalidateUploadMetaCache = ingest.InvalidateUploadMetaCache
FindByHash = ingest.FindByHash
GetActiveUpload = ingest.GetActive
OpenStoredUpload = ingest.OpenActiveObject
ActiveUploadHash = ingest.ActiveHash
ResolveLocalFile = ingest.ResolveLocalFile
IngestFromLocalPath = ingest.FromLocalPath
)
type (
@@ -54,6 +57,8 @@ const (
PolicyCreate = ingest.PolicyCreate
PolicyDedupNewRecord = ingest.PolicyDedupNewRecord
PolicyResolveExisting = ingest.PolicyResolveExisting
// ReservedPagesDeploymentType is managed exclusively by the Pages domain.
ReservedPagesDeploymentType = shared.ReservedPagesDeploymentType
)
type (
@@ -69,6 +74,7 @@ type (
var (
ErrIngestForbidden = ingest.ErrForbidden
ErrIngestStorageReadOnly = ingest.ErrStorageReadOnly
ErrReservedUploadType = ingest.ErrReservedUploadType
)
// Cache management
@@ -4,6 +4,7 @@
package handler
import (
"errors"
"net/http"
"strconv"
@@ -96,6 +97,7 @@ func ListFiles(c *gin.Context) {
// @Success 200 {object} response.Any "删除成功"
// @Failure 403 {object} response.Any "无权操作"
// @Failure 404 {object} response.Any "文件不存在"
// @Failure 409 {object} response.Any "系统保留类型或存储只读"
// @Router /api/v1/admin/uploads/{id} [delete]
func DeleteFile(c *gin.Context) {
ctx := c.Request.Context()
@@ -111,6 +113,10 @@ func DeleteFile(c *gin.Context) {
}
if _, err := softDeleteUpload(ctx, uploadID); err != nil {
if errors.Is(err, ingest.ErrReservedUploadType) {
response.AbortConflict(c, shared.ErrReservedUploadType)
return
}
if isRecordNotFound(err) {
response.AbortNotFound(c, "文件记录未找到")
return
@@ -216,6 +222,7 @@ func ListMyFiles(c *gin.Context) {
// @Success 200 {object} response.Any "删除成功"
// @Failure 403 {object} response.Any "无权操作"
// @Failure 404 {object} response.Any "文件不存在"
// @Failure 409 {object} response.Any "系统保留类型或存储只读"
// @Router /api/v1/upload/{id} [delete]
func DeleteMyFile(c *gin.Context) {
currUser, _ := oauth.GetFromContext[*model.User](c, oauth.UserObjKey)
@@ -232,11 +239,15 @@ func DeleteMyFile(c *gin.Context) {
}
if _, err := softDeleteOwnedUpload(ctx, currUser.ID, uploadID); err != nil {
if errors.Is(err, ingest.ErrReservedUploadType) {
response.AbortConflict(c, shared.ErrReservedUploadType)
return
}
if isRecordNotFound(err) {
response.AbortNotFound(c, "文件记录未找到")
return
}
if err == ingest.ErrForbidden {
if errors.Is(err, ingest.ErrForbidden) {
response.AbortForbidden(c, "无权操作")
return
}
+5
View File
@@ -53,6 +53,7 @@ type batchDownloadRequest struct {
// @Success 200 {object} response.Any{data=model.Upload} "上传成功"
// @Failure 400 {object} response.Any "请求参数错误或文件受限"
// @Failure 401 {object} response.Any "未登录"
// @Failure 409 {object} response.Any "系统保留类型或存储只读"
// @Failure 500 {object} response.Any "内部错误"
// @Router /api/v1/upload [post]
//
@@ -107,6 +108,10 @@ func UploadFile(c *gin.Context) {
}
uploadType := c.DefaultPostForm("type", "generic")
if uploadType == shared.ReservedPagesDeploymentType {
response.AbortConflict(c, shared.ErrReservedUploadType)
return
}
accessMode, errMsg := resolveUploadAccessMode(c, uploadType)
if errMsg != "" {
@@ -226,6 +226,42 @@ func TestUploadFile(t *testing.T) {
}
})
t.Run("upload rejects Pages reserved type", func(t *testing.T) {
putCountBefore := putCount
imgContent := []byte("\x89PNG\r\n\x1a\n\x00\x00\x00\rIHDR\x00\x00\x00\x01\x00\x00\x00\x01")
contentType, body := createMultipartRequest(t, "file", "pages.png", imgContent, map[string]string{
"type": shared.ReservedPagesDeploymentType,
})
req, _ := http.NewRequest("POST", "/api/v1/upload", body)
req.Header.Set("Content-Type", contentType)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusConflict {
t.Fatalf("expected status 409, got %d. Body: %s", w.Code, w.Body.String())
}
var resp testResponse
if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
t.Fatalf("unmarshal reserved type response: %v", err)
}
if resp.ErrorMsg != shared.ErrReservedUploadType {
t.Fatalf("reserved type error = %q, want %q", resp.ErrorMsg, shared.ErrReservedUploadType)
}
if putCount != putCountBefore {
t.Fatalf("reserved upload wrote storage object: put count %d -> %d", putCountBefore, putCount)
}
var count int64
if err := dbConn.Model(&model.Upload{}).
Where("type = ?", shared.ReservedPagesDeploymentType).
Count(&count).Error; err != nil {
t.Fatalf("count reserved uploads: %v", err)
}
if count != 0 {
t.Fatalf("reserved upload record count = %d, want 0", count)
}
})
t.Run("instant upload deduplication (秒传)", func(t *testing.T) {
putCount = 0
imgContent := []byte("\x89PNG\r\n\x1a\n\x00\x00\x00\rIHDR\x00\x00\x00\x01\x00\x00\x00\x01")
@@ -993,6 +1029,62 @@ func TestUserUploadManagement(t *testing.T) {
})
}
func TestDeleteReservedUploadType(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
authUser := &model.User{ID: 1001, Username: "reserved_owner"}
router := setupTestRouter(authUser)
reserved := model.Upload{
ID: 4101,
UserID: authUser.ID,
FileName: "pages.zip",
FilePath: "uploads/pages.zip",
FileSize: 128,
MimeType: "application/zip",
Extension: "zip",
Hash: "pages-reserved-hash",
Type: shared.ReservedPagesDeploymentType,
Status: model.UploadStatusUsed,
CreatedAt: time.Now(),
}
if err := dbConn.Create(&reserved).Error; err != nil {
t.Fatalf("seed reserved upload: %v", err)
}
for _, tc := range []struct {
name string
path string
}{
{name: "admin delete", path: "/api/v1/admin/uploads/4101"},
{name: "owner delete", path: "/api/v1/upload/4101"},
} {
t.Run(tc.name, func(t *testing.T) {
req, _ := http.NewRequest(http.MethodDelete, tc.path, nil)
w := httptest.NewRecorder()
router.ServeHTTP(w, req)
if w.Code != http.StatusConflict {
t.Fatalf("expected status 409, got %d. Body: %s", w.Code, w.Body.String())
}
var resp testResponse
if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
t.Fatalf("unmarshal delete response: %v", err)
}
if resp.ErrorMsg != shared.ErrReservedUploadType {
t.Fatalf("reserved delete error = %q, want %q", resp.ErrorMsg, shared.ErrReservedUploadType)
}
})
}
var persisted model.Upload
if err := dbConn.First(&persisted, reserved.ID).Error; err != nil {
t.Fatalf("reload reserved upload: %v", err)
}
if persisted.Status != model.UploadStatusUsed {
t.Fatalf("reserved upload status = %s, want used", persisted.Status)
}
}
func configureLocalStorageRoot(t *testing.T, dbConn *gorm.DB, tempDir string) {
var sc model.SystemConfig
if err := dbConn.Where("key = ?", model.ConfigKeyStorageConfig).First(&sc).Error; err != nil {
+3
View File
@@ -12,5 +12,8 @@ import (
// ErrForbidden indicates the caller is not allowed to mutate the upload record.
var ErrForbidden = errors.New("upload forbidden")
// ErrReservedUploadType indicates that a generic mutation targeted a domain-reserved upload type.
var ErrReservedUploadType = errors.New(shared.ErrReservedUploadType)
// ErrStorageReadOnly indicates the storage backend is in migration read-only mode.
var ErrStorageReadOnly = errors.New(shared.ErrStorageReadOnly)
+18 -9
View File
@@ -100,13 +100,10 @@ func storeObject(ctx context.Context, objectKey string, reader io.Reader, size i
return result.Key, nil
}
func persistUploadRecord(ctx context.Context, upload *model.Upload, objectKey string) error {
func persistUploadRecord(ctx context.Context, upload *model.Upload, objectKey string, storedByRequest bool) error {
if err := createUploadWithStats(ctx, upload); err != nil {
_, backend, backendErr := storage.Active(ctx)
if backendErr == nil {
if deleteErr := backend.Delete(ctx, objectKey); deleteErr != nil {
logger.WarnF(ctx, "清理未写入数据库的上传对象失败: %v", deleteErr)
}
if storedByRequest {
cleanupUnpersistedObject(ctx, objectKey)
}
return err
}
@@ -114,6 +111,16 @@ func persistUploadRecord(ctx context.Context, upload *model.Upload, objectKey st
return nil
}
func cleanupUnpersistedObject(ctx context.Context, objectKey string) {
_, backend, err := storage.Active(ctx)
if err != nil {
return
}
if err := backend.Delete(ctx, objectKey); err != nil {
logger.WarnF(ctx, "清理未写入数据库的上传对象失败: %v", err)
}
}
func createUploadWithStats(ctx context.Context, upload *model.Upload) error {
return db.DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := repository.CreateUploadTx(tx, upload); err != nil {
@@ -125,6 +132,8 @@ func createUploadWithStats(ctx context.Context, upload *model.Upload) error {
func createDedupRecord(ctx context.Context, existing model.Upload, req Request) (Result, error) {
accessMode := resolveAccessMode(req.Type, req.AccessMode)
metadata := req.Metadata
metadata.Bucket = existing.Metadata.Bucket
newUpload := model.Upload{
ID: idgen.NextUint64ID(),
UserID: req.UserID,
@@ -137,9 +146,9 @@ func createDedupRecord(ctx context.Context, existing model.Upload, req Request)
Type: req.Type,
Status: req.Status,
AccessMode: accessMode,
Metadata: existing.Metadata,
Metadata: metadata,
}
if err := persistUploadRecord(ctx, &newUpload, existing.FilePath); err != nil {
if err := persistUploadRecord(ctx, &newUpload, existing.FilePath, false); err != nil {
return Result{}, err
}
logger.InfoF(ctx, "文件触发秒传成功! ID: %d, Path: %s", newUpload.ID, existing.FilePath)
@@ -186,7 +195,7 @@ func createNewUpload(ctx context.Context, req Request) (Result, error) {
AccessMode: accessMode,
Metadata: req.Metadata,
}
if err := persistUploadRecord(ctx, &upload, storedKey); err != nil {
if err := persistUploadRecord(ctx, &upload, storedKey, true); err != nil {
return Result{}, err
}
+233 -2
View File
@@ -8,15 +8,20 @@ import (
"context"
"crypto/sha256"
"encoding/hex"
"errors"
"io"
"os"
"sync"
"testing"
"time"
uploadcache "github.com/Rain-kl/Wavelet/internal/apps/upload/cache"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/storage"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"gorm.io/gorm"
)
func TestIngestPolicyCreateIncrementsStats(t *testing.T) {
@@ -141,7 +146,11 @@ func TestIngestPolicyDedupNewRecordCreatesSecondRecord(t *testing.T) {
Extension: "png",
Hash: hashStr,
Type: "avatar",
Policy: PolicyDedupNewRecord,
Metadata: model.UploadMetadata{
UserAgent: "first-agent",
Extra: map[string]any{"record": "first"},
},
Policy: PolicyDedupNewRecord,
})
if err != nil {
t.Fatalf("first Ingest returned error: %v", err)
@@ -149,6 +158,10 @@ func TestIngestPolicyDedupNewRecordCreatesSecondRecord(t *testing.T) {
if putCount != 1 {
t.Fatalf("putCount after first ingest = %d, want 1", putCount)
}
first.Upload.Metadata.Bucket = "shared-bucket"
if err := dbConn.Save(&first.Upload).Error; err != nil {
t.Fatalf("update first upload metadata failed: %v", err)
}
second, err := Ingest(ctx, Request{
UserID: 1002,
@@ -159,7 +172,12 @@ func TestIngestPolicyDedupNewRecordCreatesSecondRecord(t *testing.T) {
Extension: "png",
Hash: hashStr,
Type: "avatar",
Policy: PolicyDedupNewRecord,
Metadata: model.UploadMetadata{
UserAgent: "second-agent",
Bucket: "caller-bucket-must-not-survive",
Extra: map[string]any{"record": "second"},
},
Policy: PolicyDedupNewRecord,
})
if err != nil {
t.Fatalf("second Ingest returned error: %v", err)
@@ -173,6 +191,15 @@ func TestIngestPolicyDedupNewRecordCreatesSecondRecord(t *testing.T) {
if first.Upload.ID == second.Upload.ID {
t.Fatal("dedup records should have unique IDs")
}
if second.Upload.Metadata.Bucket != "shared-bucket" {
t.Fatalf("dedup bucket = %q, want inherited shared-bucket", second.Upload.Metadata.Bucket)
}
if second.Upload.Metadata.UserAgent != "second-agent" {
t.Fatalf("dedup user agent = %q, want caller metadata", second.Upload.Metadata.UserAgent)
}
if second.Upload.Metadata.Extra["record"] != "second" {
t.Fatalf("dedup extra metadata = %#v, want caller metadata", second.Upload.Metadata.Extra)
}
var count int64
if err := dbConn.Model(&model.Upload{}).Where("hash = ?", hashStr).Count(&count).Error; err != nil {
@@ -183,6 +210,77 @@ func TestIngestPolicyDedupNewRecordCreatesSecondRecord(t *testing.T) {
}
}
func TestDedupRecordFailureDoesNotDeleteSharedObject(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
ctx := context.Background()
content := []byte("\x89PNG\r\n\x1a\nshared-object")
hash := sha256.Sum256(content)
hashStr := hex.EncodeToString(hash[:])
deleteCount := 0
restoreStorage, disableStorage := setupMockStorageWithDeleteCount(t, nil, &deleteCount)
defer restoreStorage()
defer disableStorage()
first, err := Ingest(ctx, Request{
UserID: 1001,
Reader: bytes.NewReader(content),
Size: int64(len(content)),
FileName: "shared.png",
MimeType: "image/png",
Extension: "png",
Hash: hashStr,
Type: "avatar",
Policy: PolicyDedupNewRecord,
})
if err != nil {
t.Fatalf("first Ingest returned error: %v", err)
}
const callbackName = "test:reject_dedup_upload_record"
if err := dbConn.Callback().Create().Before("gorm:create").Register(callbackName, func(tx *gorm.DB) {
upload, ok := tx.Statement.Dest.(*model.Upload)
if ok && upload.FileName == "dedup-fail.png" {
tx.AddError(errors.New("injected upload create failure"))
}
}); err != nil {
t.Fatalf("register create failure callback: %v", err)
}
defer func() { _ = dbConn.Callback().Create().Remove(callbackName) }()
_, err = Ingest(ctx, Request{
UserID: 1002,
Reader: bytes.NewReader(content),
Size: int64(len(content)),
FileName: "dedup-fail.png",
MimeType: "image/png",
Extension: "png",
Hash: hashStr,
Type: "avatar",
Metadata: model.UploadMetadata{
Extra: map[string]any{"record": "dedup-failure"},
},
Policy: PolicyDedupNewRecord,
})
if err == nil {
t.Fatal("dedup Ingest expected injected persistence error")
}
if deleteCount != 0 {
t.Fatalf("shared object delete count = %d, want 0", deleteCount)
}
_, backend, err := storage.Active(ctx)
if err != nil {
t.Fatalf("load active storage: %v", err)
}
obj, err := backend.Get(ctx, first.Upload.FilePath)
if err != nil {
t.Fatalf("shared object became unreadable after dedup failure: %v", err)
}
_ = obj.Body.Close()
}
func TestCreateUploadWithStatsRollsBackOnCreateFailure(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
@@ -259,6 +357,18 @@ func TestRemoveDecrementsStats(t *testing.T) {
if _, err := Remove(ctx, result.Upload.ID); err != nil {
t.Fatalf("Remove(%d) returned error: %v", result.Upload.ID, err)
}
stale := result.Upload
uploadcache.SetUploadMetaCache(ctx, &stale)
removedAgain, err := Remove(ctx, result.Upload.ID)
if err != nil {
t.Fatalf("second Remove(%d) returned error: %v", result.Upload.ID, err)
}
if removedAgain.Status != model.UploadStatusDeleted {
t.Fatalf("second Remove status = %s, want deleted", removedAgain.Status)
}
if _, err := uploadcache.GetUploadByID(ctx, result.Upload.ID); !errors.Is(err, gorm.ErrRecordNotFound) {
t.Fatalf("cache lookup after idempotent Remove error = %v, want record not found", err)
}
stats, err := loadTotalStats(ctx)
if err != nil {
@@ -269,6 +379,120 @@ func TestRemoveDecrementsStats(t *testing.T) {
}
}
func TestConcurrentRemoveDecrementsStatsOnce(t *testing.T) {
_, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
ctx := context.Background()
content := []byte("\x89PNG\r\n\x1a\nconcurrent-remove")
hash := sha256.Sum256(content)
restoreStorage, disableStorage := setupMockStorage(t, nil)
defer restoreStorage()
defer disableStorage()
result, err := Ingest(ctx, Request{
UserID: 1001,
Reader: bytes.NewReader(content),
Size: int64(len(content)),
FileName: "concurrent.png",
MimeType: "image/png",
Extension: "png",
Hash: hex.EncodeToString(hash[:]),
Type: "generic",
Policy: PolicyCreate,
})
if err != nil {
t.Fatalf("Ingest returned error: %v", err)
}
const workers = 8
start := make(chan struct{})
errs := make(chan error, workers)
var wg sync.WaitGroup
for i := 0; i < workers; i++ {
wg.Add(1)
go func() {
defer wg.Done()
<-start
_, removeErr := Remove(ctx, result.Upload.ID)
errs <- removeErr
}()
}
close(start)
wg.Wait()
close(errs)
for removeErr := range errs {
if removeErr != nil {
t.Fatalf("concurrent Remove returned error: %v", removeErr)
}
}
stats, err := loadTotalStats(ctx)
if err != nil {
t.Fatalf("loadTotalStats returned error: %v", err)
}
if stats.TotalCount != 0 || stats.TotalSize != 0 {
t.Fatalf("stats after concurrent remove = count %d size %d, want zero", stats.TotalCount, stats.TotalSize)
}
}
func TestRemoveOwnedAndReservedTypeBoundaries(t *testing.T) {
dbConn, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
ctx := context.Background()
ordinary := model.Upload{
ID: 99101,
UserID: 1001,
FileName: "owned.txt",
FilePath: "uploads/owned.txt",
FileSize: 16,
MimeType: "text/plain",
Extension: "txt",
Hash: "owned-hash",
Type: "generic",
Status: model.UploadStatusUsed,
CreatedAt: time.Now(),
}
reserved := model.Upload{
ID: 99102,
UserID: 1001,
FileName: "pages.zip",
FilePath: "uploads/pages.zip",
FileSize: 32,
MimeType: "application/zip",
Extension: "zip",
Hash: "reserved-hash",
Type: shared.ReservedPagesDeploymentType,
Status: model.UploadStatusUsed,
CreatedAt: time.Now(),
}
if err := dbConn.Create(&ordinary).Error; err != nil {
t.Fatalf("seed ordinary upload: %v", err)
}
if err := dbConn.Create(&reserved).Error; err != nil {
t.Fatalf("seed reserved upload: %v", err)
}
if _, err := RemoveOwned(ctx, 2002, ordinary.ID); !errors.Is(err, ErrForbidden) {
t.Fatalf("RemoveOwned non-owner error = %v, want ErrForbidden", err)
}
if _, err := Remove(ctx, reserved.ID); !errors.Is(err, ErrReservedUploadType) {
t.Fatalf("Remove reserved error = %v, want ErrReservedUploadType", err)
}
if _, err := RemoveOwned(ctx, reserved.UserID, reserved.ID); !errors.Is(err, ErrReservedUploadType) {
t.Fatalf("RemoveOwned reserved error = %v, want ErrReservedUploadType", err)
}
var persisted model.Upload
if err := dbConn.First(&persisted, reserved.ID).Error; err != nil {
t.Fatalf("reload reserved upload: %v", err)
}
if persisted.Status != model.UploadStatusUsed {
t.Fatalf("reserved upload status = %s, want used", persisted.Status)
}
}
type totalStatsSnapshot struct {
TotalCount int64
TotalSize int64
@@ -289,6 +513,10 @@ func loadTotalStats(ctx context.Context) (totalStatsSnapshot, error) {
}
func setupMockStorage(t *testing.T, putCount *int) (restore func(), disable func()) {
return setupMockStorageWithDeleteCount(t, putCount, nil)
}
func setupMockStorageWithDeleteCount(t *testing.T, putCount, deleteCount *int) (restore func(), disable func()) {
t.Helper()
mockFiles := make(map[string][]byte)
restore = storage.MockStorage(
@@ -316,6 +544,9 @@ func setupMockStorage(t *testing.T, putCount *int) (restore func(), disable func
},
func(ctx context.Context, key string) error {
delete(mockFiles, key)
if deleteCount != nil {
*deleteCount++
}
return nil
},
)
+46 -22
View File
@@ -7,52 +7,76 @@ import (
"context"
uploadcache "github.com/Rain-kl/Wavelet/internal/apps/upload/cache"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/repository"
"gorm.io/gorm"
"gorm.io/gorm/clause"
)
// Remove soft-deletes an upload and decrements incremental stats.
// Remove soft-deletes an ordinary upload and decrements incremental stats once.
func Remove(ctx context.Context, uploadID uint64) (model.Upload, error) {
upload, err := repository.GetActiveUploadByID(ctx, uploadID)
upload, err := remove(ctx, 0, uploadID, false)
if err != nil {
return model.Upload{}, err
}
if err := softDeleteUploadWithStats(ctx, &upload); err != nil {
return model.Upload{}, err
}
upload.Status = model.UploadStatusDeleted
return upload, nil
}
// RemoveOwned soft-deletes an upload owned by userID and decrements incremental stats.
// RemoveOwned soft-deletes an ordinary upload owned by userID and decrements incremental stats once.
func RemoveOwned(ctx context.Context, userID, uploadID uint64) (model.Upload, error) {
upload, err := repository.GetActiveUploadByID(ctx, uploadID)
upload, err := remove(ctx, userID, uploadID, true)
if err != nil {
return model.Upload{}, err
}
if upload.UserID != userID {
return model.Upload{}, ErrForbidden
}
if err := softDeleteUploadWithStats(ctx, &upload); err != nil {
return model.Upload{}, err
}
upload.Status = model.UploadStatusDeleted
return upload, nil
}
func softDeleteUploadWithStats(ctx context.Context, upload *model.Upload) error {
statsSnapshot := *upload
func remove(ctx context.Context, userID, uploadID uint64, owned bool) (model.Upload, error) {
var upload model.Upload
if err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := repository.SoftDeleteUploadTx(tx, upload); err != nil {
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).
Where("id = ?", uploadID).
First(&upload).Error; err != nil {
return err
}
return uploadstats.ApplyUploadStatsDeltaTx(tx, &statsSnapshot, -1)
}); err != nil {
if owned && upload.UserID != userID {
return ErrForbidden
}
if upload.Type == shared.ReservedPagesDeploymentType {
return ErrReservedUploadType
}
_, err := RemoveLockedTx(tx, &upload)
return err
}); err != nil {
return model.Upload{}, err
}
uploadcache.InvalidateUploadMetaCache(ctx, upload.ID)
return nil
InvalidateUploadMetaCache(ctx, uploadID)
upload.Status = model.UploadStatusDeleted
return upload, nil
}
// RemoveLockedTx performs the idempotent active-to-deleted transition for a row
// that the caller has already locked in its surrounding transaction.
func RemoveLockedTx(tx *gorm.DB, upload *model.Upload) (bool, error) {
rowsAffected, err := repository.SoftDeleteUploadTx(tx, upload)
if err != nil {
return false, err
}
if rowsAffected == 0 {
return false, nil
}
if err := uploadstats.ApplyUploadStatsDeltaTx(tx, upload, -1); err != nil {
return false, err
}
upload.Status = model.UploadStatusDeleted
return true, nil
}
// InvalidateUploadMetaCache invalidates upload metadata after the caller commits its transaction.
func InvalidateUploadMetaCache(ctx context.Context, uploadID uint64) {
uploadcache.InvalidateUploadMetaCache(ctx, uploadID)
}
+2
View File
@@ -17,4 +17,6 @@ const (
FileStatsTrendDays = 7
MaxS3KeyLength = 1024
AccessCacheTTL = 5 // seconds; multiplied by time.Second at use site
// ReservedPagesDeploymentType is managed exclusively by the Pages domain.
ReservedPagesDeploymentType = "openflare_pages_deployment"
)
+1
View File
@@ -27,6 +27,7 @@ const (
ErrQueryFileCountFailed = "查询文件数量失败"
ErrQueryFileListFailed = "查询文件列表失败"
ErrDeleteFileFailed = "删除文件失败"
ErrReservedUploadType = "系统保留的文件类型不能通过通用文件接口操作"
ErrStorageReadOnly = "存储迁移维护中,当前仅允许读取文件"
ErrS3KeyRequired = "s3 key must not be empty"
ErrS3KeyTooLongFormat = "s3 key exceeds maximum length of %d"
+16 -18
View File
@@ -10,16 +10,15 @@ import (
"fmt"
"time"
uploadcache "github.com/Rain-kl/Wavelet/internal/apps/upload/cache"
"github.com/Rain-kl/Wavelet/internal/apps/upload/ingest"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
uploadstorage "github.com/Rain-kl/Wavelet/internal/apps/upload/storage"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/storage"
"github.com/Rain-kl/Wavelet/internal/task"
"github.com/Rain-kl/Wavelet/pkg/logger"
"gorm.io/gorm"
"gorm.io/gorm/clause"
)
const (
@@ -77,32 +76,31 @@ func (h *SystemCleanupHandler) Execute(ctx context.Context, _ []byte) (*task.Tas
for _, u := range unusedUploads {
totalProcessed++
transitioned := false
if err := db.DB(ctx).Transaction(func(tx *gorm.DB) error {
if err := tx.Model(&model.Upload{}).
Where("id = ? AND status = ?", u.ID, model.UploadStatusPending).
Update("status", model.UploadStatusDeleted).Error; err != nil {
var locked model.Upload
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).
Where("id = ?", u.ID).
First(&locked).Error; err != nil {
return err
}
_, backend, err := storage.Active(ctx)
if err != nil {
return err
if locked.Status != model.UploadStatusPending || !locked.CreatedAt.Before(oneHourAgo) {
return nil
}
if err := backend.Delete(ctx, u.FilePath); err != nil {
return err
}
return nil
var err error
transitioned, err = ingest.RemoveLockedTx(tx, &locked)
return err
}); err != nil {
task.AppendLog(ctx, "清理上传文件失败 [ID:%d]: %v", u.ID, err)
lastID = u.ID
continue
}
uploadstats.RecordUploadStatsRemove(ctx, &u)
uploadcache.InvalidateUploadMetaCache(ctx, u.ID)
totalDeleted++
ingest.InvalidateUploadMetaCache(ctx, u.ID)
if transitioned {
totalDeleted++
}
lastID = u.ID
}
}
+28 -2
View File
@@ -19,6 +19,7 @@ import (
"github.com/Rain-kl/Wavelet/internal/apps/upload/filesrv"
"github.com/Rain-kl/Wavelet/internal/apps/upload/shared"
uploadstats "github.com/Rain-kl/Wavelet/internal/apps/upload/stats"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/diskcache"
"github.com/Rain-kl/Wavelet/internal/model"
@@ -33,13 +34,17 @@ func TestSystemCleanupHandler_Execute(t *testing.T) {
_, _, cleanup := testhelper.SetupTestEnvironment(t)
defer cleanup()
// Mock S3 存储(让 DeleteObject 总是成功)
deleteCount := 0
// Mock S3 存储并记录 Delete,cleanup 不应物理删除共享对象。
storageMock := storage.MockStorage(
func(ctx context.Context, key string, body io.Reader, size int64, contentType string) error {
return nil
},
func(ctx context.Context, key string) (*storage.Object, error) { return nil, nil },
func(ctx context.Context, key string) error { return nil },
func(ctx context.Context, key string) error {
deleteCount++
return nil
},
)
defer storageMock()
storage.IsEnabledFunc = func() bool { return true }
@@ -86,6 +91,7 @@ func TestSystemCleanupHandler_Execute(t *testing.T) {
for _, r := range records {
err := db.DB(ctx).Create(r).Error
require.NoError(t, err)
require.NoError(t, uploadstats.ApplyUploadStatsAdd(ctx, r))
}
// 准备推送历史测试数据:1个旧的(应删除),1个新的(应保留)
@@ -147,6 +153,26 @@ func TestSystemCleanupHandler_Execute(t *testing.T) {
var usedCount int64
db.DB(ctx).Model(&model.Upload{}).Where("status = ?", model.UploadStatusUsed).Count(&usedCount)
assert.Equal(t, int64(1), usedCount, "used 状态的文件不应受影响")
assert.Equal(t, 0, deleteCount, "记录级 cleanup 不应调用 storage backend Delete")
var totalStats model.UploadStat
err = db.DB(ctx).
Where("dimension = ? AND stat_key = ?", model.UploadStatDimensionTotal, "").
First(&totalStats).Error
require.NoError(t, err)
assert.Equal(t, int64(2), totalStats.FileCount, "cleanup 后统计只应保留 used 与最近 pending 记录")
assert.Equal(t, int64(768), totalStats.FileSize, "cleanup 后统计大小应只扣减一次")
_, err = handler.Execute(ctx, nil)
require.NoError(t, err)
var statsAfterSecondRun model.UploadStat
err = db.DB(ctx).
Where("dimension = ? AND stat_key = ?", model.UploadStatDimensionTotal, "").
First(&statsAfterSecondRun).Error
require.NoError(t, err)
assert.Equal(t, totalStats.FileCount, statsAfterSecondRun.FileCount, "重复 cleanup 不应再次扣减统计")
assert.Equal(t, totalStats.FileSize, statsAfterSecondRun.FileSize, "重复 cleanup 不应再次扣减统计大小")
assert.Equal(t, 0, deleteCount, "重复 cleanup 仍不应调用 storage backend Delete")
// 验证推送历史数据状态:10天前的应被删除,今天的应保留
var pushCount int64
@@ -0,0 +1,82 @@
-- +goose Up
ALTER TABLE of_pages_projects
ADD COLUMN IF NOT EXISTS content_config_version INTEGER NOT NULL DEFAULT 0;
ALTER TABLE of_pages_deployments
ADD COLUMN IF NOT EXISTS source_type VARCHAR(32) NOT NULL DEFAULT '',
ADD COLUMN IF NOT EXISTS source_identity CHAR(64),
ADD COLUMN IF NOT EXISTS source_revision CHAR(64),
ADD COLUMN IF NOT EXISTS source_label VARCHAR(255) NOT NULL DEFAULT '',
ADD COLUMN IF NOT EXISTS source_meta TEXT NOT NULL DEFAULT '',
ADD COLUMN IF NOT EXISTS trigger_type VARCHAR(32) NOT NULL DEFAULT '';
UPDATE of_pages_deployments
SET source_type = 'manual_upload',
trigger_type = 'manual_upload';
CREATE TABLE IF NOT EXISTS of_pages_project_sources (
id BIGSERIAL PRIMARY KEY,
project_id BIGINT NOT NULL,
source_type VARCHAR(32) NOT NULL DEFAULT '',
remote_url TEXT NOT NULL DEFAULT '',
remote_network_policy VARCHAR(32) NOT NULL DEFAULT '',
github_repository VARCHAR(255) NOT NULL DEFAULT '',
release_selector VARCHAR(16) NOT NULL DEFAULT '',
release_tag VARCHAR(255) NOT NULL DEFAULT '',
asset_name VARCHAR(255) NOT NULL DEFAULT '',
auto_update_enabled BOOLEAN NOT NULL DEFAULT FALSE,
check_interval_minutes INTEGER NOT NULL DEFAULT 0,
config_version INTEGER NOT NULL DEFAULT 0,
source_identity CHAR(64) NOT NULL DEFAULT '',
created_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP,
updated_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE UNIQUE INDEX IF NOT EXISTS idx_of_pages_project_sources_project_id
ON of_pages_project_sources (project_id);
CREATE TABLE IF NOT EXISTS of_pages_project_source_runtime (
source_id BIGINT PRIMARY KEY,
etag VARCHAR(512) NOT NULL DEFAULT '',
last_seen_revision CHAR(64) NOT NULL DEFAULT '',
last_seen_detail TEXT NOT NULL DEFAULT '',
last_applied_revision CHAR(64) NOT NULL DEFAULT '',
last_applied_detail TEXT NOT NULL DEFAULT '',
sync_status VARCHAR(32) NOT NULL DEFAULT '',
last_error TEXT NOT NULL DEFAULT '',
last_checked_at TIMESTAMPTZ,
last_synced_at TIMESTAMPTZ,
next_check_at TIMESTAMPTZ,
lease_expires_at TIMESTAMPTZ,
lease_token VARCHAR(64) NOT NULL DEFAULT '',
updated_at TIMESTAMPTZ NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_of_pages_project_source_runtime_next_check_at
ON of_pages_project_source_runtime (next_check_at);
CREATE UNIQUE INDEX IF NOT EXISTS idx_of_pages_deployments_project_number
ON of_pages_deployments (project_id, deployment_number);
CREATE UNIQUE INDEX IF NOT EXISTS idx_of_pages_deployments_source_revision
ON of_pages_deployments (project_id, source_identity, source_revision)
WHERE source_identity IS NOT NULL AND source_revision IS NOT NULL;
-- +goose Down
DROP INDEX IF EXISTS idx_of_pages_deployments_source_revision;
DROP INDEX IF EXISTS idx_of_pages_deployments_project_number;
DROP INDEX IF EXISTS idx_of_pages_project_source_runtime_next_check_at;
DROP TABLE IF EXISTS of_pages_project_source_runtime;
DROP INDEX IF EXISTS idx_of_pages_project_sources_project_id;
DROP TABLE IF EXISTS of_pages_project_sources;
ALTER TABLE of_pages_deployments
DROP COLUMN IF EXISTS trigger_type,
DROP COLUMN IF EXISTS source_meta,
DROP COLUMN IF EXISTS source_label,
DROP COLUMN IF EXISTS source_revision,
DROP COLUMN IF EXISTS source_identity,
DROP COLUMN IF EXISTS source_type;
ALTER TABLE of_pages_projects
DROP COLUMN IF EXISTS content_config_version;
@@ -0,0 +1,38 @@
-- +goose Up
-- Earlier built-in schedules used explicit IDs, so advance the identity only
-- when it trails either existing rows or an already-higher sequence value.
SELECT setval(
pg_get_serial_sequence('w_schedules', 'id'),
GREATEST(
1,
COALESCE((SELECT MAX(id) FROM w_schedules), 0),
COALESCE((
SELECT sequences.last_value
FROM pg_sequences AS sequences
WHERE format('%I.%I', sequences.schemaname, sequences.sequencename)::regclass =
pg_get_serial_sequence('w_schedules', 'id')::regclass
), 0)
),
TRUE
);
INSERT INTO w_schedules (name, task_type, cron, payload, is_active, created_at, updated_at)
SELECT
'OpenFlare Pages 部署源扫描',
'of_pages_source_scan',
'*/5 * * * *',
'{}',
TRUE,
CURRENT_TIMESTAMP,
CURRENT_TIMESTAMP
WHERE NOT EXISTS (
SELECT 1 FROM w_schedules WHERE task_type = 'of_pages_source_scan'
);
-- +goose Down
DELETE FROM w_schedules
WHERE task_type = 'of_pages_source_scan'
AND name = 'OpenFlare Pages 部署源扫描'
AND cron = '*/5 * * * *'
AND payload = '{}'
AND is_active = TRUE;
@@ -0,0 +1,151 @@
-- +goose Up
ALTER TABLE of_pages_projects
ADD COLUMN content_config_version INTEGER NOT NULL DEFAULT 0;
ALTER TABLE of_pages_deployments
ADD COLUMN source_type TEXT NOT NULL DEFAULT '';
ALTER TABLE of_pages_deployments
ADD COLUMN source_identity TEXT;
ALTER TABLE of_pages_deployments
ADD COLUMN source_revision TEXT;
ALTER TABLE of_pages_deployments
ADD COLUMN source_label TEXT NOT NULL DEFAULT '';
ALTER TABLE of_pages_deployments
ADD COLUMN source_meta TEXT NOT NULL DEFAULT '';
ALTER TABLE of_pages_deployments
ADD COLUMN trigger_type TEXT NOT NULL DEFAULT '';
UPDATE of_pages_deployments
SET source_type = 'manual_upload',
trigger_type = 'manual_upload';
CREATE TABLE IF NOT EXISTS of_pages_project_sources (
id INTEGER PRIMARY KEY AUTOINCREMENT,
project_id INTEGER NOT NULL,
source_type TEXT NOT NULL DEFAULT '',
remote_url TEXT NOT NULL DEFAULT '',
remote_network_policy TEXT NOT NULL DEFAULT '',
github_repository TEXT NOT NULL DEFAULT '',
release_selector TEXT NOT NULL DEFAULT '',
release_tag TEXT NOT NULL DEFAULT '',
asset_name TEXT NOT NULL DEFAULT '',
auto_update_enabled INTEGER NOT NULL DEFAULT 0,
check_interval_minutes INTEGER NOT NULL DEFAULT 0,
config_version INTEGER NOT NULL DEFAULT 0,
source_identity TEXT NOT NULL DEFAULT '',
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP,
updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE UNIQUE INDEX IF NOT EXISTS idx_of_pages_project_sources_project_id
ON of_pages_project_sources (project_id);
CREATE TABLE IF NOT EXISTS of_pages_project_source_runtime (
source_id INTEGER PRIMARY KEY,
etag TEXT NOT NULL DEFAULT '',
last_seen_revision TEXT NOT NULL DEFAULT '',
last_seen_detail TEXT NOT NULL DEFAULT '',
last_applied_revision TEXT NOT NULL DEFAULT '',
last_applied_detail TEXT NOT NULL DEFAULT '',
sync_status TEXT NOT NULL DEFAULT '',
last_error TEXT NOT NULL DEFAULT '',
last_checked_at DATETIME,
last_synced_at DATETIME,
next_check_at DATETIME,
lease_expires_at DATETIME,
lease_token TEXT NOT NULL DEFAULT '',
updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
);
CREATE INDEX IF NOT EXISTS idx_of_pages_project_source_runtime_next_check_at
ON of_pages_project_source_runtime (next_check_at);
CREATE UNIQUE INDEX IF NOT EXISTS idx_of_pages_deployments_project_number
ON of_pages_deployments (project_id, deployment_number);
CREATE UNIQUE INDEX IF NOT EXISTS idx_of_pages_deployments_source_revision
ON of_pages_deployments (project_id, source_identity, source_revision)
WHERE source_identity IS NOT NULL AND source_revision IS NOT NULL;
-- +goose Down
DROP INDEX IF EXISTS idx_of_pages_deployments_source_revision;
DROP INDEX IF EXISTS idx_of_pages_deployments_project_number;
DROP INDEX IF EXISTS idx_of_pages_project_source_runtime_next_check_at;
DROP TABLE IF EXISTS of_pages_project_source_runtime;
DROP INDEX IF EXISTS idx_of_pages_project_sources_project_id;
DROP TABLE IF EXISTS of_pages_project_sources;
-- SQLite 的 Down 必须重建受影响表,完整移除新增列并保留原有数据与索引。
CREATE TABLE of_pages_deployments_before_source (
id INTEGER PRIMARY KEY AUTOINCREMENT,
project_id INTEGER NOT NULL,
deployment_number INTEGER NOT NULL,
checksum TEXT NOT NULL,
status TEXT NOT NULL DEFAULT 'uploaded',
upload_id INTEGER NOT NULL DEFAULT 0,
artifact_path TEXT NOT NULL,
file_count INTEGER NOT NULL DEFAULT 0,
total_size INTEGER NOT NULL DEFAULT 0,
created_by TEXT NOT NULL DEFAULT '',
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP,
activated_at DATETIME
);
INSERT INTO of_pages_deployments_before_source (
id, project_id, deployment_number, checksum, status, upload_id, artifact_path,
file_count, total_size, created_by, created_at, activated_at
)
SELECT
id, project_id, deployment_number, checksum, status, upload_id, artifact_path,
file_count, total_size, created_by, created_at, activated_at
FROM of_pages_deployments;
DROP TABLE of_pages_deployments;
ALTER TABLE of_pages_deployments_before_source RENAME TO of_pages_deployments;
CREATE INDEX IF NOT EXISTS idx_of_pages_deployments_project_id
ON of_pages_deployments (project_id);
CREATE INDEX IF NOT EXISTS idx_of_pages_deployments_checksum
ON of_pages_deployments (checksum);
CREATE INDEX IF NOT EXISTS idx_of_pages_deployments_status
ON of_pages_deployments (status);
CREATE INDEX IF NOT EXISTS idx_of_pages_deployments_upload_id
ON of_pages_deployments (upload_id);
CREATE TABLE of_pages_projects_before_source (
id INTEGER PRIMARY KEY AUTOINCREMENT,
name TEXT NOT NULL,
slug TEXT NOT NULL,
description TEXT NOT NULL DEFAULT '',
enabled INTEGER NOT NULL DEFAULT 1,
spa_fallback_enabled INTEGER NOT NULL DEFAULT 0,
spa_fallback_path TEXT NOT NULL DEFAULT '/index.html',
api_proxy_enabled INTEGER NOT NULL DEFAULT 0,
api_proxy_path TEXT NOT NULL DEFAULT '',
api_proxy_pass TEXT NOT NULL DEFAULT '',
api_proxy_rewrite TEXT NOT NULL DEFAULT '',
active_deployment_id INTEGER,
root_dir TEXT NOT NULL DEFAULT '',
entry_file TEXT NOT NULL DEFAULT 'index.html',
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP,
updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
);
INSERT INTO of_pages_projects_before_source (
id, name, slug, description, enabled, spa_fallback_enabled, spa_fallback_path,
api_proxy_enabled, api_proxy_path, api_proxy_pass, api_proxy_rewrite,
active_deployment_id, root_dir, entry_file, created_at, updated_at
)
SELECT
id, name, slug, description, enabled, spa_fallback_enabled, spa_fallback_path,
api_proxy_enabled, api_proxy_path, api_proxy_pass, api_proxy_rewrite,
active_deployment_id, root_dir, entry_file, created_at, updated_at
FROM of_pages_projects;
DROP TABLE of_pages_projects;
ALTER TABLE of_pages_projects_before_source RENAME TO of_pages_projects;
CREATE UNIQUE INDEX IF NOT EXISTS idx_of_pages_projects_slug
ON of_pages_projects (slug);
CREATE INDEX IF NOT EXISTS idx_of_pages_projects_active_deployment_id
ON of_pages_projects (active_deployment_id);
@@ -0,0 +1,21 @@
-- +goose Up
INSERT INTO w_schedules (name, task_type, cron, payload, is_active, created_at, updated_at)
SELECT
'OpenFlare Pages 部署源扫描',
'of_pages_source_scan',
'*/5 * * * *',
'{}',
1,
CURRENT_TIMESTAMP,
CURRENT_TIMESTAMP
WHERE NOT EXISTS (
SELECT 1 FROM w_schedules WHERE task_type = 'of_pages_source_scan'
);
-- +goose Down
DELETE FROM w_schedules
WHERE task_type = 'of_pages_source_scan'
AND name = 'OpenFlare Pages 部署源扫描'
AND cron = '*/5 * * * *'
AND payload = '{}'
AND is_active = 1;
+3 -3
View File
@@ -18,9 +18,9 @@ import (
"gorm.io/gorm"
)
// expectedMigratedSystemConfigCount 包含初始 32 项系统配置,以及 202606220004
// 从 of_options 迁移过来的 48 项业务配置(OpenFlare/UptimeKuma/OpenResty)。
const expectedMigratedSystemConfigCount = 80
// expectedMigratedSystemConfigCount 包含初始 32 项系统配置、202606220004
// 从 of_options 迁移过来的 48 项业务配置,以及 Pages 的 2 项业务配置。
const expectedMigratedSystemConfigCount = 82
func TestMigrateInitializesSQLiteDatabase(t *testing.T) {
sqliteDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
@@ -0,0 +1,322 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package migrator
import (
"database/sql"
"fmt"
"os"
"strings"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/glebarez/sqlite"
"github.com/pressly/goose/v3"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/driver/postgres"
"gorm.io/gorm"
)
const (
pagesSourcePreviousMigration = int64(202607180001)
pagesSourceMigration = int64(202607190001)
pagesMigrationProjectID = uint(900001)
pagesMigrationDeploymentID = uint(900001)
)
func TestPagesSourceMigrationSQLiteUpDownUp(t *testing.T) {
dbPath := t.TempDir() + "/pages-source-migration.db"
gormDB, err := gorm.Open(sqlite.Open(dbPath), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
require.NoError(t, err)
sqlDB, err := gormDB.DB()
require.NoError(t, err)
sqlDB.SetMaxOpenConns(1)
t.Cleanup(func() { require.NoError(t, sqlDB.Close()) })
runPagesSourceMigrationUpDownUp(t, gormDB, sqlDB, dialectSqlite, "goose/sqlite")
var indexSQL string
require.NoError(t, gormDB.Raw(
"SELECT sql FROM sqlite_master WHERE type = 'index' AND name = ?",
"idx_of_pages_deployments_source_revision",
).Scan(&indexSQL).Error)
assert.Contains(t, strings.ToUpper(indexSQL), "WHERE SOURCE_IDENTITY IS NOT NULL AND SOURCE_REVISION IS NOT NULL")
}
func TestPagesSourceMigrationPostgresUpDownUp(t *testing.T) {
dsn := strings.TrimSpace(os.Getenv("OPENFLARE_TEST_POSTGRES_DSN"))
if dsn == "" {
t.Skip("OPENFLARE_TEST_POSTGRES_DSN is not set")
}
gormDB, err := gorm.Open(postgres.Open(dsn), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
require.NoError(t, err)
sqlDB, err := gormDB.DB()
require.NoError(t, err)
sqlDB.SetMaxOpenConns(1)
schema := fmt.Sprintf("pages_source_migration_%d", time.Now().UnixNano())
require.Regexp(t, `^[a-z0-9_]+$`, schema)
require.NoError(t, gormDB.Exec(`CREATE SCHEMA "`+schema+`"`).Error)
require.NoError(t, gormDB.Exec(`SET search_path TO "`+schema+`"`).Error)
t.Cleanup(func() {
assert.NoError(t, gormDB.Exec("SET search_path TO public").Error)
assert.NoError(t, gormDB.Exec(`DROP SCHEMA IF EXISTS "`+schema+`" CASCADE`).Error)
assert.NoError(t, sqlDB.Close())
})
runPagesSourceMigrationUpDownUp(t, gormDB, sqlDB, dialectPostgres, "goose/postgres")
}
func runPagesSourceMigrationUpDownUp(
t *testing.T,
gormDB *gorm.DB,
sqlDB *sql.DB,
dialect string,
dir string,
) {
t.Helper()
goose.SetBaseFS(migrationFS)
require.NoError(t, goose.SetDialect(dialect))
require.NoError(t, goose.UpTo(sqlDB, dir, pagesSourcePreviousMigration))
seedPrePagesSourceMigrationData(t, gormDB)
require.NoError(t, goose.UpTo(sqlDB, dir, pagesSourceMigration))
assertPagesSourceMigrationUp(t, gormDB)
require.NoError(t, goose.DownTo(sqlDB, dir, pagesSourcePreviousMigration))
assertPagesSourceMigrationDown(t, gormDB)
require.NoError(t, goose.UpTo(sqlDB, dir, pagesSourceMigration))
assertPagesSourceMigrationUpAgain(t, gormDB)
}
func seedPrePagesSourceMigrationData(t *testing.T, gormDB *gorm.DB) {
t.Helper()
require.NoError(t, gormDB.Exec(`
INSERT INTO of_pages_projects (
id, name, slug, description, enabled, active_deployment_id, root_dir, entry_file
) VALUES (?, ?, ?, ?, ?, ?, ?, ?)
`,
pagesMigrationProjectID,
"Migration Site",
"migration-site",
"keep-project-data",
true,
pagesMigrationDeploymentID,
"public",
"home.html",
).Error)
require.NoError(t, gormDB.Exec(`
INSERT INTO of_pages_deployments (
id, project_id, deployment_number, checksum, status, upload_id, artifact_path,
file_count, total_size, created_by
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
`,
pagesMigrationDeploymentID,
pagesMigrationProjectID,
1,
strings.Repeat("a", 64),
model.PagesDeploymentStatusActive,
uint64(700001),
"legacy/package.zip",
2,
int64(128),
"user:1",
).Error)
}
func assertPagesSourceMigrationUp(t *testing.T, gormDB *gorm.DB) {
t.Helper()
migrator := gormDB.Migrator()
assert.True(t, migrator.HasTable(&model.PagesProjectSource{}))
assert.True(t, migrator.HasTable(&model.PagesProjectSourceRuntime{}))
assert.True(t, migrator.HasColumn(&model.PagesProject{}, "ContentConfigVersion"))
assert.True(t, migrator.HasColumn(&model.PagesDeployment{}, "SourceType"))
assert.True(t, migrator.HasColumn(&model.PagesDeployment{}, "SourceIdentity"))
assert.True(t, migrator.HasColumn(&model.PagesDeployment{}, "SourceRevision"))
assert.True(t, migrator.HasIndex(&model.PagesProjectSource{}, "idx_of_pages_project_sources_project_id"))
assert.True(t, migrator.HasIndex(&model.PagesProjectSourceRuntime{}, "idx_of_pages_project_source_runtime_next_check_at"))
assert.True(t, migrator.HasIndex(&model.PagesDeployment{}, "idx_of_pages_deployments_project_number"))
assert.True(t, migrator.HasIndex(&model.PagesDeployment{}, "idx_of_pages_deployments_source_revision"))
var project model.PagesProject
require.NoError(t, gormDB.First(&project, pagesMigrationProjectID).Error)
assert.Equal(t, 0, project.ContentConfigVersion)
assert.Equal(t, "keep-project-data", project.Description)
assert.Equal(t, "public", project.RootDir)
assert.Equal(t, "home.html", project.EntryFile)
var deployment model.PagesDeployment
require.NoError(t, gormDB.First(&deployment, pagesMigrationDeploymentID).Error)
assert.Equal(t, "manual_upload", deployment.SourceType)
assert.Equal(t, "manual_upload", deployment.TriggerType)
assert.Nil(t, deployment.SourceIdentity)
assert.Nil(t, deployment.SourceRevision)
assert.Equal(t, uint64(700001), deployment.UploadID)
sourceID := createMigrationSourceRuntime(t, gormDB)
assertPagesSourceConstraints(t, gormDB, sourceID)
}
func createMigrationSourceRuntime(t *testing.T, gormDB *gorm.DB) uint {
t.Helper()
source := model.PagesProjectSource{
ProjectID: pagesMigrationProjectID,
SourceType: "remote_url",
RemoteURL: "https://example.com/site.zip?token=secret",
RemoteNetworkPolicy: "public",
CheckIntervalMinutes: 0,
ConfigVersion: 1,
SourceIdentity: strings.Repeat("b", 64),
}
require.NoError(t, gormDB.Create(&source).Error)
require.NotZero(t, source.ID)
require.NoError(t, gormDB.Create(&model.PagesProjectSourceRuntime{
SourceID: source.ID,
SyncStatus: "idle",
}).Error)
return source.ID
}
func assertPagesSourceConstraints(t *testing.T, gormDB *gorm.DB, sourceID uint) {
t.Helper()
duplicateSource := model.PagesProjectSource{
ProjectID: pagesMigrationProjectID,
SourceType: "remote_url",
ConfigVersion: 1,
SourceIdentity: strings.Repeat("c", 64),
}
assert.Error(t, gormDB.Create(&duplicateSource).Error)
for number := 2; number <= 3; number++ {
require.NoError(t, createMigrationDeployment(
gormDB,
number,
strings.Repeat(string(rune('a'+number)), 64),
nil,
nil,
))
}
identity := strings.Repeat("d", 64)
revision := strings.Repeat("e", 64)
require.NoError(t, createMigrationDeployment(
gormDB,
4,
strings.Repeat("f", 64),
&identity,
&revision,
))
assert.Error(t, createMigrationDeployment(
gormDB,
5,
strings.Repeat("0", 64),
&identity,
&revision,
))
assert.Error(t, createMigrationDeployment(
gormDB,
1,
strings.Repeat("1", 64),
nil,
nil,
))
var runtime model.PagesProjectSourceRuntime
require.NoError(t, gormDB.First(&runtime, sourceID).Error)
assert.Equal(t, "idle", runtime.SyncStatus)
}
func createMigrationDeployment(
gormDB *gorm.DB,
deploymentNumber int,
checksum string,
identity *string,
revision *string,
) error {
return gormDB.Create(&model.PagesDeployment{
ProjectID: pagesMigrationProjectID,
DeploymentNumber: deploymentNumber,
Checksum: checksum,
Status: model.PagesDeploymentStatusUploaded,
UploadID: uint64(710000 + deploymentNumber),
ArtifactPath: fmt.Sprintf("legacy/%d.zip", deploymentNumber),
SourceType: "manual_upload",
SourceIdentity: identity,
SourceRevision: revision,
TriggerType: "manual_upload",
}).Error
}
func assertPagesSourceMigrationDown(t *testing.T, gormDB *gorm.DB) {
t.Helper()
migrator := gormDB.Migrator()
assert.False(t, migrator.HasTable(&model.PagesProjectSource{}))
assert.False(t, migrator.HasTable(&model.PagesProjectSourceRuntime{}))
assert.False(t, migrator.HasColumn(&model.PagesProject{}, "ContentConfigVersion"))
assert.False(t, migrator.HasColumn(&model.PagesDeployment{}, "SourceType"))
assert.False(t, migrator.HasColumn(&model.PagesDeployment{}, "SourceIdentity"))
assert.False(t, migrator.HasColumn(&model.PagesDeployment{}, "SourceRevision"))
assert.False(t, migrator.HasIndex(&model.PagesDeployment{}, "idx_of_pages_deployments_project_number"))
assert.False(t, migrator.HasIndex(&model.PagesDeployment{}, "idx_of_pages_deployments_source_revision"))
assert.True(t, migrator.HasIndex(&model.PagesDeployment{}, "idx_of_pages_deployments_project_id"))
assert.True(t, migrator.HasIndex(&model.PagesDeployment{}, "idx_of_pages_deployments_upload_id"))
var project struct {
Description string
RootDir string
EntryFile string
ActiveDeploymentID *uint
}
require.NoError(t, gormDB.Table("of_pages_projects").Where("id = ?", pagesMigrationProjectID).Take(&project).Error)
assert.Equal(t, "keep-project-data", project.Description)
assert.Equal(t, "public", project.RootDir)
assert.Equal(t, "home.html", project.EntryFile)
require.NotNil(t, project.ActiveDeploymentID)
assert.Equal(t, pagesMigrationDeploymentID, *project.ActiveDeploymentID)
var deployment struct {
UploadID uint64
ArtifactPath string
FileCount int
TotalSize int64
}
require.NoError(t, gormDB.Table("of_pages_deployments").Where("id = ?", pagesMigrationDeploymentID).Take(&deployment).Error)
assert.Equal(t, uint64(700001), deployment.UploadID)
assert.Equal(t, "legacy/package.zip", deployment.ArtifactPath)
assert.Equal(t, 2, deployment.FileCount)
assert.Equal(t, int64(128), deployment.TotalSize)
var count int64
require.NoError(t, gormDB.Table("of_pages_deployments").Where("project_id = ?", pagesMigrationProjectID).Count(&count).Error)
assert.Equal(t, int64(4), count)
}
func assertPagesSourceMigrationUpAgain(t *testing.T, gormDB *gorm.DB) {
t.Helper()
assert.True(t, gormDB.Migrator().HasTable(&model.PagesProjectSource{}))
assert.True(t, gormDB.Migrator().HasTable(&model.PagesProjectSourceRuntime{}))
assert.True(t, gormDB.Migrator().HasColumn(&model.PagesProject{}, "ContentConfigVersion"))
assert.True(t, gormDB.Migrator().HasColumn(&model.PagesDeployment{}, "SourceRevision"))
var count int64
require.NoError(t, gormDB.Table("of_pages_deployments").
Where("project_id = ? AND source_type = ? AND trigger_type = ?", pagesMigrationProjectID, "manual_upload", "manual_upload").
Count(&count).Error)
assert.Equal(t, int64(4), count)
require.NoError(t, gormDB.Table("of_pages_project_sources").Count(&count).Error)
assert.Zero(t, count, "source config is intentionally removed by Down and is not reconstructable")
}
@@ -0,0 +1,160 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package migrator
import (
"database/sql"
"fmt"
"os"
"strings"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/glebarez/sqlite"
"github.com/pressly/goose/v3"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/driver/postgres"
"gorm.io/gorm"
)
const (
pagesSourceScanPreviousMigration = int64(202607190001)
pagesSourceScanMigration = int64(202607190002)
pagesSourceScanTaskType = "of_pages_source_scan"
)
func TestPagesSourceScanScheduleMigrationSQLite(t *testing.T) {
dbPath := t.TempDir() + "/pages-source-scan-migration.db"
gormDB, err := gorm.Open(sqlite.Open(dbPath), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
require.NoError(t, err)
sqlDB, err := gormDB.DB()
require.NoError(t, err)
sqlDB.SetMaxOpenConns(1)
t.Cleanup(func() { require.NoError(t, sqlDB.Close()) })
runPagesSourceScanScheduleMigration(t, gormDB, sqlDB, dialectSqlite, "goose/sqlite")
}
func TestPagesSourceScanScheduleMigrationPostgres(t *testing.T) {
dsn := strings.TrimSpace(os.Getenv("OPENFLARE_TEST_POSTGRES_DSN"))
if dsn == "" {
t.Skip("OPENFLARE_TEST_POSTGRES_DSN is not set")
}
gormDB, err := gorm.Open(postgres.Open(dsn), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
require.NoError(t, err)
sqlDB, err := gormDB.DB()
require.NoError(t, err)
sqlDB.SetMaxOpenConns(1)
schema := fmt.Sprintf("pages_source_scan_migration_%d", time.Now().UnixNano())
require.Regexp(t, `^[a-z0-9_]+$`, schema)
require.NoError(t, gormDB.Exec(`CREATE SCHEMA "`+schema+`"`).Error)
require.NoError(t, gormDB.Exec(`SET search_path TO "`+schema+`"`).Error)
t.Cleanup(func() {
assert.NoError(t, gormDB.Exec("SET search_path TO public").Error)
assert.NoError(t, gormDB.Exec(`DROP SCHEMA IF EXISTS "`+schema+`" CASCADE`).Error)
assert.NoError(t, sqlDB.Close())
})
runPagesSourceScanScheduleMigration(t, gormDB, sqlDB, dialectPostgres, "goose/postgres")
}
func runPagesSourceScanScheduleMigration(
t *testing.T,
gormDB *gorm.DB,
sqlDB *sql.DB,
dialect string,
dir string,
) {
t.Helper()
goose.SetBaseFS(migrationFS)
require.NoError(t, goose.SetDialect(dialect))
require.NoError(t, goose.UpTo(sqlDB, dir, pagesSourceScanPreviousMigration))
var previousMaxID uint64
require.NoError(t, gormDB.Table("w_schedules").Select("COALESCE(MAX(id), 0)").Scan(&previousMaxID).Error)
require.NoError(t, goose.UpTo(sqlDB, dir, pagesSourceScanMigration))
seeded := assertPagesSourceScanSchedule(t, gormDB)
assert.NotZero(t, seeded.ID)
if dialect == dialectPostgres {
assert.Greater(t, seeded.ID, previousMaxID)
}
require.NoError(t, goose.DownTo(sqlDB, dir, pagesSourceScanPreviousMigration))
assertPagesSourceScanScheduleMissing(t, gormDB)
custom := model.Schedule{
ID: 900001,
Name: "用户保留的 Pages 扫描任务",
TaskType: pagesSourceScanTaskType,
Cron: "0 * * * *",
Payload: `{"custom":true}`,
IsActive: false,
}
require.NoError(t, gormDB.Create(&custom).Error)
require.NoError(t, goose.UpTo(sqlDB, dir, pagesSourceScanMigration))
var schedules []model.Schedule
require.NoError(t, gormDB.Where("task_type = ?", pagesSourceScanTaskType).Find(&schedules).Error)
require.Len(t, schedules, 1)
assert.Equal(t, custom.ID, schedules[0].ID)
assert.Equal(t, custom.Name, schedules[0].Name)
require.NoError(t, goose.DownTo(sqlDB, dir, pagesSourceScanPreviousMigration))
var retained model.Schedule
require.NoError(t, gormDB.First(&retained, custom.ID).Error)
assert.Equal(t, custom.TaskType, retained.TaskType)
}
func TestPagesSourceScanScheduleMigrationsUseDatabaseGeneratedIDs(t *testing.T) {
for _, name := range []string{
"goose/postgres/202607190002_seed_pages_source_scan.sql",
"goose/sqlite/202607190002_seed_pages_source_scan.sql",
} {
t.Run(name, func(t *testing.T) {
content, err := migrationFS.ReadFile(name)
require.NoError(t, err)
normalized := strings.ToLower(string(content))
assert.NotContains(t, normalized, "insert into w_schedules (id,")
assert.NotContains(t, normalized, "coalesce(max(id)")
})
}
postgresContent, err := migrationFS.ReadFile("goose/postgres/202607190002_seed_pages_source_scan.sql")
require.NoError(t, err)
compactPostgres := strings.Join(strings.Fields(strings.ToLower(string(postgresContent))), " ")
assert.Contains(
t,
compactPostgres,
"select setval( pg_get_serial_sequence('w_schedules', 'id'), greatest( 1,",
"sequence synchronization must retain a valid lower bound for an empty table",
)
}
func assertPagesSourceScanSchedule(t *testing.T, gormDB *gorm.DB) model.Schedule {
t.Helper()
var schedules []model.Schedule
require.NoError(t, gormDB.Where("task_type = ?", pagesSourceScanTaskType).Find(&schedules).Error)
require.Len(t, schedules, 1)
schedule := schedules[0]
assert.Equal(t, "OpenFlare Pages 部署源扫描", schedule.Name)
assert.Equal(t, "*/5 * * * *", schedule.Cron)
assert.Equal(t, "{}", schedule.Payload)
assert.True(t, schedule.IsActive)
return schedule
}
func assertPagesSourceScanScheduleMissing(t *testing.T, gormDB *gorm.DB) {
t.Helper()
var count int64
require.NoError(t, gormDB.Model(&model.Schedule{}).
Where("task_type = ?", pagesSourceScanTaskType).
Count(&count).Error)
assert.Zero(t, count)
}
+21 -16
View File
@@ -57,13 +57,10 @@ func initSQLite() {
// Trace 注入
if err = db.Use(
tracing.NewPlugin(
tracing.WithoutMetrics(),
tracing.WithAttributes(
attribute.String("db.instance", sqlitePath),
attribute.String("db.system", "SQLite"),
),
),
newGORMTracingPlugin([]attribute.KeyValue{
attribute.String("db.instance", sqlitePath),
attribute.String("db.system", "SQLite"),
}),
); err != nil {
log.Fatalf("[SQLite] init trace failed: %v\n", err)
}
@@ -98,15 +95,12 @@ func initPostgres() {
// Trace 注入
if err = db.Use(
tracing.NewPlugin(
tracing.WithoutMetrics(),
tracing.WithAttributes(
attribute.String("db.instance", dbConfig.Database),
attribute.String("db.ip", dbConfig.Host),
attribute.String("server.address", net.JoinHostPort(dbConfig.Host, strconv.Itoa(dbConfig.Port))),
attribute.String("db.system", "PostgreSQL"),
),
),
newGORMTracingPlugin([]attribute.KeyValue{
attribute.String("db.instance", dbConfig.Database),
attribute.String("db.ip", dbConfig.Host),
attribute.String("server.address", net.JoinHostPort(dbConfig.Host, strconv.Itoa(dbConfig.Port))),
attribute.String("db.system", "PostgreSQL"),
}),
); err != nil {
log.Fatalf("[PostgreSQL] init trace failed: %v\n", err)
}
@@ -160,6 +154,17 @@ func initPostgres() {
}
// newGORMTracingPlugin 构造数据库链路追踪插件。查询参数只保留占位符,避免凭据等绑定值进入 Span。
func newGORMTracingPlugin(attrs []attribute.KeyValue, extraOptions ...tracing.Option) gorm.Plugin {
options := []tracing.Option{
tracing.WithoutMetrics(),
tracing.WithoutQueryVariables(),
tracing.WithAttributes(attrs...),
}
options = append(options, extraOptions...)
return tracing.NewPlugin(options...)
}
// buildDSN 构建 PostgreSQL DSN
func buildDSN(host string, port int, username, password string) string {
cfg := config.Config.Database
+5
View File
@@ -49,6 +49,11 @@ func (l *gormZapLogger) Error(ctx context.Context, fmt string, args ...interface
}
}
// ParamsFilter 让 GORM 的 Trace 回调只接收参数化 SQL,避免绑定值被 Dialector.Explain 展开到日志。
func (l *gormZapLogger) ParamsFilter(_ context.Context, sql string, _ ...interface{}) (string, []interface{}) {
return sql, nil
}
func (l *gormZapLogger) Trace(ctx context.Context, begin time.Time, fc func() (sql string, rowsAffected int64), err error) {
elapsed := time.Since(begin)
switch {
+76
View File
@@ -4,11 +4,40 @@
package db
import (
"context"
"strings"
"testing"
"time"
"github.com/glebarez/sqlite"
"gorm.io/gorm"
gormLogger "gorm.io/gorm/logger"
)
type paramsFilterCaptureLogger struct {
filter *gormZapLogger
traces []string
}
func (l *paramsFilterCaptureLogger) LogMode(gormLogger.LogLevel) gormLogger.Interface {
return l
}
func (l *paramsFilterCaptureLogger) Info(context.Context, string, ...interface{}) {}
func (l *paramsFilterCaptureLogger) Warn(context.Context, string, ...interface{}) {}
func (l *paramsFilterCaptureLogger) Error(context.Context, string, ...interface{}) {}
func (l *paramsFilterCaptureLogger) ParamsFilter(ctx context.Context, sql string, params ...interface{}) (string, []interface{}) {
return l.filter.ParamsFilter(ctx, sql, params...)
}
func (l *paramsFilterCaptureLogger) Trace(_ context.Context, _ time.Time, fc func() (string, int64), _ error) {
sql, _ := fc()
l.traces = append(l.traces, sql)
}
func TestParseLogLevel(t *testing.T) {
t.Parallel()
@@ -37,3 +66,50 @@ func TestParseLogLevel(t *testing.T) {
})
}
}
func TestGormZapLoggerParamsFilterDropsBoundValues(t *testing.T) {
t.Parallel()
const (
query = "UPDATE openflare_pages_sources SET remote_url = ? WHERE id = ?"
secret = "https://example.test/release.zip?token=super-secret"
)
filteredSQL, filteredParams := (&gormZapLogger{}).ParamsFilter(t.Context(), query, secret, int64(42))
if filteredSQL != query {
t.Fatalf("ParamsFilter() sql = %q, want %q", filteredSQL, query)
}
if filteredParams != nil {
t.Fatalf("ParamsFilter() params = %#v, want nil", filteredParams)
}
}
func TestGormZapLoggerKeepsParameterizedSQLInTrace(t *testing.T) {
t.Parallel()
const secret = "https://example.test/release.zip?token=trace-secret"
capture := &paramsFilterCaptureLogger{filter: &gormZapLogger{}}
testDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{Logger: capture})
if err != nil {
t.Fatalf("open sqlite: %v", err)
}
if err := testDB.Exec("CREATE TABLE source_secrets (remote_url TEXT NOT NULL)").Error; err != nil {
t.Fatalf("create table: %v", err)
}
capture.traces = nil
if err := testDB.Exec("INSERT INTO source_secrets (remote_url) VALUES (?)", secret).Error; err != nil {
t.Fatalf("insert source secret: %v", err)
}
if len(capture.traces) != 1 {
t.Fatalf("trace count = %d, want 1", len(capture.traces))
}
traceSQL := capture.traces[0]
if strings.Contains(traceSQL, secret) || strings.Contains(traceSQL, "trace-secret") {
t.Fatalf("trace SQL leaked bound value: %q", traceSQL)
}
if !strings.Contains(traceSQL, "VALUES (?)") {
t.Fatalf("trace SQL = %q, want parameter placeholder", traceSQL)
}
}
+69
View File
@@ -0,0 +1,69 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package db
import (
"context"
"strings"
"testing"
"github.com/glebarez/sqlite"
"go.opentelemetry.io/otel/attribute"
sdktrace "go.opentelemetry.io/otel/sdk/trace"
"go.opentelemetry.io/otel/sdk/trace/tracetest"
semconv "go.opentelemetry.io/otel/semconv/v1.30.0"
"gorm.io/gorm"
gormLogger "gorm.io/gorm/logger"
"gorm.io/plugin/opentelemetry/tracing"
)
func TestGORMTracingPluginDoesNotRecordQueryVariables(t *testing.T) {
t.Parallel()
const secret = "https://example.test/release.zip?token=otel-secret"
spanRecorder := tracetest.NewSpanRecorder()
tracerProvider := sdktrace.NewTracerProvider(sdktrace.WithSpanProcessor(spanRecorder))
t.Cleanup(func() {
if err := tracerProvider.Shutdown(context.Background()); err != nil {
t.Errorf("shutdown tracer provider: %v", err)
}
})
testDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
Logger: gormLogger.Default.LogMode(gormLogger.Silent),
})
if err != nil {
t.Fatalf("open sqlite: %v", err)
}
if err := testDB.Use(newGORMTracingPlugin(
[]attribute.KeyValue{attribute.String("db.instance", "trace-test")},
tracing.WithTracerProvider(tracerProvider),
)); err != nil {
t.Fatalf("register tracing plugin: %v", err)
}
if err := testDB.Exec("CREATE TABLE source_secrets (remote_url TEXT NOT NULL)").Error; err != nil {
t.Fatalf("create table: %v", err)
}
if err := testDB.Exec("INSERT INTO source_secrets (remote_url) VALUES (?)", secret).Error; err != nil {
t.Fatalf("insert source secret: %v", err)
}
var queryText string
for _, span := range spanRecorder.Ended() {
for _, attr := range span.Attributes() {
if attr.Key == semconv.DBQueryTextKey && strings.Contains(attr.Value.AsString(), "INSERT INTO source_secrets") {
queryText = attr.Value.AsString()
}
}
}
if queryText == "" {
t.Fatal("database query text attribute not found")
}
if strings.Contains(queryText, secret) || strings.Contains(queryText, "otel-secret") {
t.Fatalf("db.query.text leaked bound value: %q", queryText)
}
if !strings.Contains(queryText, "VALUES (?)") {
t.Fatalf("db.query.text = %q, want parameter placeholder", queryText)
}
}
@@ -0,0 +1,725 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
// Package githubrelease resolves and downloads public GitHub Release assets.
// It deliberately does not know about Pages projects, deployments or runtime
// state so other callers can reuse the same constrained HTTP contract.
package githubrelease
import (
"bytes"
"context"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"io"
"math"
"net/http"
"net/url"
"os"
"path"
"regexp"
"strconv"
"strings"
"time"
"unicode"
"unicode/utf8"
)
const (
// APIVersion is the GitHub REST API contract used by this package.
APIVersion = "2026-03-10"
// SelectorLatest uses GitHub's repository latest-release endpoint.
SelectorLatest Selector = "latest"
// SelectorTag resolves one exact GitHub release tag.
SelectorTag Selector = "tag"
defaultAPIBaseURL = "https://api.github.com"
defaultUserAgent = "OpenFlare-GitHubRelease/1.0"
metadataAccept = "application/vnd.github+json"
assetAccept = "application/octet-stream"
maxMetadataBytes = 4 << 20
maxAssetErrorNames = 10
maxSafeTextBytes = 255
maxSafeAssetNameLen = 96
maxDigestBytes = 96
maxETagBytes = 512
safePartsCapacity = 6
)
var (
errInvalidRequest = errors.New("GitHub Release 请求参数无效")
errMetadata = errors.New("GitHub Release 元数据响应无效")
errAssetMissing = errors.New("GitHub Release 中未找到指定的已上传 asset")
errDownload = errors.New("GitHub Release asset 下载失败")
errTooLarge = errors.New("GitHub Release asset 超过大小限制")
errEmptyAsset = errors.New("GitHub Release asset 内容为空")
errDigest = errors.New("GitHub Release asset digest 无效或校验失败")
errCleanup = errors.New("GitHub Release 临时文件清理失败")
ownerPattern = regexp.MustCompile(`^[A-Za-z0-9](?:[A-Za-z0-9-]{0,37}[A-Za-z0-9])?$`)
repoPattern = regexp.MustCompile(`^[A-Za-z0-9._-]+$`)
hexPattern = regexp.MustCompile(`^[0-9a-fA-F]{64}$`)
)
var (
// ErrInvalidRequest identifies caller configuration errors.
ErrInvalidRequest = errInvalidRequest
// ErrMetadata identifies malformed, unavailable or failed Release metadata requests.
ErrMetadata = errMetadata
// ErrAssetNotFound identifies an otherwise valid Release without the exact uploaded asset.
ErrAssetNotFound = errAssetMissing
// ErrDownload identifies network or HTTP failures while downloading an asset.
ErrDownload = errDownload
// ErrAssetTooLarge identifies assets that exceed the caller's hard byte limit.
ErrAssetTooLarge = errTooLarge
// ErrEmptyAsset identifies an empty downloaded asset.
ErrEmptyAsset = errEmptyAsset
// ErrDigestMismatch identifies malformed or mismatched declared SHA-256 digests.
ErrDigestMismatch = errDigest
)
// Selector identifies GitHub's own latest endpoint or one exact tag.
type Selector string
// ResolveRequest describes one public repository release asset lookup.
type ResolveRequest struct {
Repository string
Selector Selector
Tag string
AssetName string
ETag string
}
// Release contains only metadata safe and necessary for source resolution.
type Release struct {
ID string `json:"release_id"`
Tag string `json:"tag"`
Name string `json:"name,omitempty"`
Draft bool `json:"draft"`
Prerelease bool `json:"prerelease"`
PublishedAt time.Time `json:"published_at,omitempty"`
}
// Asset contains the immutable target metadata returned by a resolve call.
type Asset struct {
ID string `json:"asset_id"`
Name string `json:"asset_name"`
State string `json:"state"`
Size int64 `json:"size"`
UpdatedAt time.Time `json:"updated_at,omitempty"`
Digest string `json:"digest,omitempty"`
}
// ResolveResult is either a selected uploaded asset or a not-modified marker.
type ResolveResult struct {
NotModified bool `json:"not_modified"`
ETag string `json:"etag,omitempty"`
Release Release `json:"release,omitempty"`
Asset Asset `json:"asset,omitempty"`
RetryAt *time.Time `json:"retry_at,omitempty"`
}
// DownloadRequest identifies an already resolved asset. Asset IDs never come
// from an untrusted URL and the download endpoint is built locally.
type DownloadRequest struct {
Repository string
Asset Asset
MaxBytes int64
}
// DownloadResult owns a temporary file. Call Cleanup after ingestion.
type DownloadResult struct {
Path string
Size int64
SHA256 string
DeclaredDigest string
}
// Cleanup removes the temporary file and is safe to call more than once.
func (result *DownloadResult) Cleanup() error {
if result == nil || result.Path == "" {
return nil
}
name := result.Path
err := os.Remove(name)
if err == nil || errors.Is(err, os.ErrNotExist) {
result.Path = ""
return nil
}
return errCleanup
}
// Error is a safe provider error. It never retains a response body, request
// URL, redirect location or request headers.
type Error struct {
Kind error
StatusCode int
RequestID string
Repository string
Tag string
AssetName string
AvailableAssets []string
RetryAt *time.Time
}
func (providerError *Error) Error() string {
if providerError == nil {
return "GitHub Release 请求失败"
}
message := "GitHub Release 请求失败"
if providerError.Kind != nil {
message = providerError.Kind.Error()
}
parts := make([]string, 0, safePartsCapacity)
if providerError.StatusCode != 0 {
parts = append(parts, "status="+strconv.Itoa(providerError.StatusCode))
}
if providerError.RequestID != "" {
parts = append(parts, "request_id="+providerError.RequestID)
}
if providerError.Repository != "" {
parts = append(parts, "repo="+providerError.Repository)
}
if providerError.Tag != "" {
parts = append(parts, "tag="+providerError.Tag)
}
if providerError.AssetName != "" {
parts = append(parts, "asset="+providerError.AssetName)
}
if len(providerError.AvailableAssets) > 0 {
parts = append(parts, "available="+strings.Join(providerError.AvailableAssets, ","))
}
if len(parts) == 0 {
return message
}
return message + " (" + strings.Join(parts, " ") + ")"
}
func (providerError *Error) Unwrap() error {
if providerError == nil {
return nil
}
return providerError.Kind
}
// RetryAt extracts the server-directed retry deadline from an error.
func RetryAt(err error) (time.Time, bool) {
var providerError *Error
if !errors.As(err, &providerError) || providerError.RetryAt == nil {
return time.Time{}, false
}
return *providerError.RetryAt, true
}
// RetryTime is retained as a compatibility alias for early callers.
//
// Deprecated: use RetryAt.
func RetryTime(err error) (time.Time, bool) {
return RetryAt(err)
}
// IsNotFound reports both a missing Release endpoint and a Release that lacks
// the exact uploaded asset requested by the caller.
func IsNotFound(err error) bool {
if errors.Is(err, ErrAssetNotFound) {
return true
}
var providerError *Error
return errors.As(err, &providerError) && providerError.StatusCode == http.StatusNotFound
}
// IsDigestError reports malformed or mismatched declared asset digests.
func IsDigestError(err error) bool {
return errors.Is(err, ErrDigestMismatch)
}
// IsRetryable classifies provider failures without relying on localized error
// strings. Configuration, not-found, size, empty-content and digest failures
// are permanent. Network failures, 408/425/429 and 5xx responses are retryable.
func IsRetryable(err error) bool {
if err == nil || errors.Is(err, ErrInvalidRequest) || IsNotFound(err) ||
errors.Is(err, ErrAssetTooLarge) || errors.Is(err, ErrEmptyAsset) || IsDigestError(err) {
return false
}
var providerError *Error
if !errors.As(err, &providerError) {
return false
}
if providerError.StatusCode == 0 {
return errors.Is(err, ErrMetadata) || errors.Is(err, ErrDownload) || errors.Is(err, errCleanup)
}
if providerError.StatusCode < http.StatusBadRequest {
return errors.Is(err, ErrMetadata) || errors.Is(err, ErrDownload) || errors.Is(err, errCleanup)
}
if providerError.RetryAt != nil {
return true
}
return providerError.StatusCode == http.StatusRequestTimeout ||
providerError.StatusCode == http.StatusTooEarly ||
providerError.StatusCode == http.StatusTooManyRequests ||
providerError.StatusCode >= http.StatusInternalServerError
}
// Client accesses public GitHub Releases using a fixed, constrained transport.
type Client struct {
httpClient *http.Client
baseURL string
createTemp func(string, string) (*os.File, error)
now func() time.Time
}
// NewClient constructs a production client for api.github.com. Public
// repositories do not require or send a token.
func NewClient() *Client {
return newClient(defaultClientOptions())
}
// Resolve calls GitHub's latest or exact-tag endpoint and selects one exact,
// case-sensitive uploaded asset. It never falls back to source archives.
func (client *Client) Resolve(ctx context.Context, request ResolveRequest) (ResolveResult, error) {
repository, tag, endpoint, err := normalizeResolveRequest(client.baseURL, request)
if err != nil {
return ResolveResult{}, safeError(errInvalidRequest, 0, "", repository, tag, validErrorAssetName(request.AssetName), nil, nil)
}
httpRequest, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil)
if err != nil {
return ResolveResult{}, safeError(errInvalidRequest, 0, "", repository, tag, validErrorAssetName(request.AssetName), nil, nil)
}
applyMetadataHeaders(httpRequest, request.ETag)
response, err := client.httpClient.Do(httpRequest) //nolint:gosec // endpoint and every dial target are constrained
if err != nil {
return ResolveResult{}, safeError(errMetadata, 0, "", repository, tag, request.AssetName, nil, nil)
}
defer func() { _ = response.Body.Close() }()
retryAt := responseRetryAt(response, client.now())
etag := safeETag(response.Header.Get("ETag"))
if response.StatusCode == http.StatusNotModified {
if etag == "" {
etag = safeETag(request.ETag)
}
return ResolveResult{NotModified: true, ETag: etag, RetryAt: retryAt}, nil
}
if response.StatusCode != http.StatusOK {
return ResolveResult{}, safeHTTPError(errMetadata, response, repository, tag, request.AssetName, retryAt)
}
body, readErr := io.ReadAll(io.LimitReader(response.Body, maxMetadataBytes+1))
if readErr != nil || len(body) > maxMetadataBytes || !utf8.Valid(body) {
return ResolveResult{}, safeHTTPError(errMetadata, response, repository, tag, request.AssetName, retryAt)
}
var payload releasePayload
decoder := json.NewDecoder(bytes.NewReader(body))
decoder.UseNumber()
if err := decoder.Decode(&payload); err != nil {
return ResolveResult{}, safeHTTPError(errMetadata, response, repository, tag, request.AssetName, retryAt)
}
if err := ensureJSONEOF(decoder); err != nil {
return ResolveResult{}, safeHTTPError(errMetadata, response, repository, tag, request.AssetName, retryAt)
}
release, assets, err := convertRelease(payload)
if err != nil {
return ResolveResult{}, safeHTTPError(errMetadata, response, repository, tag, request.AssetName, retryAt)
}
for _, asset := range assets {
if asset.State == "uploaded" && asset.Name == request.AssetName {
return ResolveResult{
ETag: etag,
Release: release,
Asset: asset,
RetryAt: retryAt,
}, nil
}
}
available := safeAssetNames(assets)
return ResolveResult{}, safeError(
errAssetMissing,
response.StatusCode,
response.Header.Get("X-GitHub-Request-Id"),
repository,
release.Tag,
request.AssetName,
available,
retryAt,
)
}
// Download streams an asset into a package-owned temporary file while
// enforcing a hard byte limit and verifying GitHub's declared sha256 digest.
func (client *Client) Download(ctx context.Context, request DownloadRequest) (*DownloadResult, error) {
repository, err := normalizeRepository(request.Repository)
if err != nil || request.MaxBytes <= 0 || !validPositiveID(request.Asset.ID) ||
!validAssetName(request.Asset.Name) || request.Asset.Size < 0 {
return nil, safeError(errInvalidRequest, 0, "", repository, "", validErrorAssetName(request.Asset.Name), nil, nil)
}
if request.Asset.Size > request.MaxBytes {
return nil, safeError(errTooLarge, 0, "", repository, "", validErrorAssetName(request.Asset.Name), nil, nil)
}
endpoint := strings.TrimRight(client.baseURL, "/") + "/repos/" + repository + "/releases/assets/" + request.Asset.ID
httpRequest, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil)
if err != nil {
return nil, safeError(errInvalidRequest, 0, "", repository, "", request.Asset.Name, nil, nil)
}
applyAssetHeaders(httpRequest)
response, err := client.httpClient.Do(httpRequest) //nolint:gosec // endpoint and every dial target are constrained
if err != nil {
return nil, safeError(errDownload, 0, "", repository, "", request.Asset.Name, nil, nil)
}
defer func() { _ = response.Body.Close() }()
retryAt := responseRetryAt(response, client.now())
if response.StatusCode != http.StatusOK {
return nil, safeHTTPError(errDownload, response, repository, "", request.Asset.Name, retryAt)
}
if response.ContentLength > request.MaxBytes {
return nil, safeHTTPError(errTooLarge, response, repository, "", request.Asset.Name, retryAt)
}
result, err := client.streamAsset(response.Body, request.MaxBytes, request.Asset.Digest)
if err != nil {
return nil, safeError(err, response.StatusCode, response.Header.Get("X-GitHub-Request-Id"), repository, "", request.Asset.Name, nil, retryAt)
}
return result, nil
}
func (client *Client) streamAsset(body io.Reader, maxBytes int64, declaredDigest string) (result *DownloadResult, resultErr error) {
tempFile, err := client.createTemp("", "openflare-github-release-*")
if err != nil {
return nil, errDownload
}
tempPath := tempFile.Name()
defer func() {
closeErr := tempFile.Close()
if resultErr == nil && closeErr != nil {
resultErr = errDownload
}
if resultErr != nil {
if removeErr := os.Remove(tempPath); removeErr != nil && !errors.Is(removeErr, os.ErrNotExist) {
resultErr = errCleanup
}
}
}()
hasher := sha256.New()
readLimit := maxBytes
if readLimit < math.MaxInt64 {
readLimit++
}
size, err := io.Copy(io.MultiWriter(tempFile, hasher), io.LimitReader(body, readLimit))
if err != nil {
return nil, errDownload
}
if size > maxBytes {
return nil, errTooLarge
}
if size == 0 {
return nil, errEmptyAsset
}
checksum := hex.EncodeToString(hasher.Sum(nil))
if err := verifyDeclaredDigest(declaredDigest, checksum); err != nil {
return nil, err
}
return &DownloadResult{
Path: tempPath,
Size: size,
SHA256: checksum,
DeclaredDigest: strings.ToLower(strings.TrimSpace(declaredDigest)),
}, nil
}
type releasePayload struct {
ID json.Number `json:"id"`
Tag string `json:"tag_name"`
Name string `json:"name"`
Draft bool `json:"draft"`
Prerelease bool `json:"prerelease"`
PublishedAt string `json:"published_at"`
Assets []assetPayload `json:"assets"`
}
type assetPayload struct {
ID json.Number `json:"id"`
Name string `json:"name"`
State string `json:"state"`
Size int64 `json:"size"`
UpdatedAt string `json:"updated_at"`
Digest string `json:"digest"`
}
func convertRelease(payload releasePayload) (Release, []Asset, error) {
releaseID, err := positiveJSONID(payload.ID)
if err != nil {
return Release{}, nil, err
}
if !validReleaseDisplayTag(payload.Tag) {
return Release{}, nil, errMetadata
}
publishedAt, err := parseOptionalTime(payload.PublishedAt)
if err != nil {
return Release{}, nil, err
}
release := Release{
ID: releaseID,
Tag: payload.Tag,
Name: safeText(payload.Name, maxSafeTextBytes),
Draft: payload.Draft,
Prerelease: payload.Prerelease,
PublishedAt: publishedAt,
}
assets := make([]Asset, 0, len(payload.Assets))
for _, rawAsset := range payload.Assets {
assetID, assetErr := positiveJSONID(rawAsset.ID)
if assetErr != nil || rawAsset.Size < 0 {
return Release{}, nil, errMetadata
}
updatedAt, assetErr := parseOptionalTime(rawAsset.UpdatedAt)
if assetErr != nil {
return Release{}, nil, errMetadata
}
assets = append(assets, Asset{
ID: assetID,
Name: rawAsset.Name,
State: rawAsset.State,
Size: rawAsset.Size,
UpdatedAt: updatedAt,
Digest: safeText(rawAsset.Digest, maxDigestBytes),
})
}
return release, assets, nil
}
func normalizeResolveRequest(baseURL string, request ResolveRequest) (string, string, string, error) {
repository, err := normalizeRepository(request.Repository)
if err != nil || !validAssetName(request.AssetName) {
return repository, validErrorTag(request.Tag), "", errInvalidRequest
}
baseURL = strings.TrimRight(baseURL, "/")
switch request.Selector {
case SelectorLatest:
if strings.TrimSpace(request.Tag) != "" {
return repository, "", "", errInvalidRequest
}
return repository, "latest", baseURL + "/repos/" + repository + "/releases/latest", nil
case SelectorTag:
if !validTag(request.Tag) {
return repository, validErrorTag(request.Tag), "", errInvalidRequest
}
return repository, request.Tag, baseURL + "/repos/" + repository + "/releases/tags/" + url.PathEscape(request.Tag), nil
default:
return repository, validErrorTag(request.Tag), "", errInvalidRequest
}
}
func normalizeRepository(repository string) (string, error) {
repository = strings.TrimSpace(repository)
parts := strings.Split(repository, "/")
if len(parts) != 2 || !ownerPattern.MatchString(parts[0]) || !repoPattern.MatchString(parts[1]) ||
len(parts[1]) > 100 || parts[1] == "." || parts[1] == ".." {
return "", errInvalidRequest
}
return parts[0] + "/" + parts[1], nil
}
func validAssetName(assetName string) bool {
return validLogText(assetName, maxSafeTextBytes, false) && path.Base(assetName) == assetName &&
assetName != "." && assetName != ".." && !strings.ContainsAny(assetName, `/\`)
}
func validTag(tag string) bool {
if !validLogText(tag, maxSafeTextBytes, false) || strings.ContainsAny(tag, " ~^:?*[\\") ||
strings.Contains(tag, "..") || strings.Contains(tag, "@{") || strings.Contains(tag, "//") ||
strings.HasPrefix(tag, "/") || strings.HasSuffix(tag, "/") || strings.HasSuffix(tag, ".") {
return false
}
for _, component := range strings.Split(tag, "/") {
if component == "" || strings.HasPrefix(component, ".") || strings.HasSuffix(component, ".lock") {
return false
}
}
return true
}
func validReleaseDisplayTag(tag string) bool {
return validLogText(tag, maxSafeTextBytes, false)
}
func validLogText(value string, maxBytes int, allowEmpty bool) bool {
if (!allowEmpty && value == "") || len(value) > maxBytes || !utf8.ValidString(value) {
return false
}
for _, character := range value {
if isLogControl(character) {
return false
}
}
return true
}
func isLogControl(character rune) bool {
if unicode.IsControl(character) || character == '\u2028' || character == '\u2029' {
return true
}
switch character {
case '\u061c', '\u200e', '\u200f',
'\u202a', '\u202b', '\u202c', '\u202d', '\u202e',
'\u2066', '\u2067', '\u2068', '\u2069':
return true
default:
return false
}
}
func validErrorTag(tag string) string {
if !validTag(tag) || containsSecretDelimiter(tag) {
return ""
}
return tag
}
func validErrorAssetName(assetName string) string {
if !validAssetName(assetName) || containsSecretDelimiter(assetName) {
return ""
}
return assetName
}
func containsSecretDelimiter(value string) bool {
return strings.ContainsAny(value, "?&=#") || strings.Contains(value, "://")
}
func validPositiveID(id string) bool {
parsed, err := strconv.ParseInt(id, 10, 64)
return err == nil && parsed > 0 && strconv.FormatInt(parsed, 10) == id
}
func positiveJSONID(id json.Number) (string, error) {
parsed, err := strconv.ParseInt(id.String(), 10, 64)
if err != nil || parsed <= 0 {
return "", errMetadata
}
return strconv.FormatInt(parsed, 10), nil
}
func parseOptionalTime(value string) (time.Time, error) {
if value == "" {
return time.Time{}, nil
}
parsed, err := time.Parse(time.RFC3339, value)
if err != nil {
return time.Time{}, errMetadata
}
return parsed, nil
}
func ensureJSONEOF(decoder *json.Decoder) error {
var trailing any
if err := decoder.Decode(&trailing); errors.Is(err, io.EOF) {
return nil
}
return errMetadata
}
func verifyDeclaredDigest(declaredDigest string, checksum string) error {
declaredDigest = strings.TrimSpace(declaredDigest)
if declaredDigest == "" {
return nil
}
algorithm, digest, ok := strings.Cut(declaredDigest, ":")
if !ok || !strings.EqualFold(algorithm, "sha256") || !hexPattern.MatchString(digest) ||
!strings.EqualFold(digest, checksum) {
return errDigest
}
return nil
}
func safeAssetNames(assets []Asset) []string {
count := len(assets)
if count > maxAssetErrorNames {
count = maxAssetErrorNames
}
names := make([]string, 0, count)
for _, asset := range assets[:count] {
name := safeText(asset.Name, maxSafeAssetNameLen)
if containsSecretDelimiter(name) {
name = "<redacted>"
}
names = append(names, name)
}
return names
}
func safeText(value string, maxBytes int) string {
var builder strings.Builder
for _, character := range value {
if isLogControl(character) {
builder.WriteByte('?')
continue
}
builder.WriteRune(character)
if builder.Len() >= maxBytes {
break
}
}
result := builder.String()
for len(result) > maxBytes {
_, size := utf8.DecodeLastRuneInString(result)
result = result[:len(result)-size]
}
return result
}
func safeETag(value string) string {
value = strings.TrimSpace(value)
if len(value) > maxETagBytes || safeText(value, maxETagBytes) != value {
return ""
}
return value
}
func safeHTTPError(kind error, response *http.Response, repository string, tag string, assetName string, retryAt *time.Time) error {
return safeError(
kind,
response.StatusCode,
response.Header.Get("X-GitHub-Request-Id"),
repository,
tag,
assetName,
nil,
retryAt,
)
}
func safeError(
kind error,
statusCode int,
requestID string,
repository string,
tag string,
assetName string,
availableAssets []string,
retryAt *time.Time,
) error {
return &Error{
Kind: kind,
StatusCode: statusCode,
RequestID: safeErrorToken(requestID, maxSafeTextBytes),
Repository: safeErrorToken(repository, maxSafeTextBytes),
Tag: safeErrorToken(tag, maxSafeTextBytes),
AssetName: safeErrorToken(assetName, maxSafeAssetNameLen),
AvailableAssets: availableAssets,
RetryAt: retryAt,
}
}
func safeErrorToken(value string, maxBytes int) string {
if !validLogText(value, maxBytes, true) {
return ""
}
value = safeText(value, maxBytes)
if containsSecretDelimiter(value) {
return ""
}
return value
}
@@ -0,0 +1,748 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package githubrelease
import (
"context"
"crypto/sha256"
"crypto/tls"
"encoding/hex"
"errors"
"fmt"
"net"
"net/http"
"net/http/httptest"
"net/netip"
"net/url"
"os"
"path/filepath"
"strconv"
"strings"
"sync"
"sync/atomic"
"testing"
"time"
)
type resolverFunc func(context.Context, string, string) ([]netip.Addr, error)
func (resolve resolverFunc) LookupNetIP(ctx context.Context, network string, host string) ([]netip.Addr, error) {
return resolve(ctx, network, host)
}
func TestResolveLatestUsesGitHubContractAndSelectsUploadedAsset(t *testing.T) {
t.Parallel()
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
if request.URL.Path != "/repos/acme/site/releases/latest" {
t.Errorf("path = %q", request.URL.Path)
}
assertHeader(t, request, "Accept", metadataAccept)
assertHeader(t, request, "User-Agent", defaultUserAgent)
assertHeader(t, request, "X-GitHub-Api-Version", APIVersion)
assertHeader(t, request, "If-None-Match", `W/"old"`)
writer.Header().Set("ETag", `W/"new"`)
writer.Header().Set("X-RateLimit-Remaining", "0")
writer.Header().Set("X-RateLimit-Reset", "1800000000")
_, _ = writer.Write([]byte(`{
"id": 9007199254740991,
"tag_name": "v1.2.3",
"name": "Stable",
"published_at": "2026-07-18T12:00:00Z",
"assets": [
{"id": 11, "name": "dist.zip", "state": "new", "size": 1},
{"id": 9007199254740990, "name": "dist.zip", "state": "uploaded", "size": 42,
"updated_at": "2026-07-18T12:10:00Z", "digest": "sha256:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa"}
]
}`))
}))
defer server.Close()
client := newTestClient(t, server.URL, nil)
result, err := client.Resolve(context.Background(), ResolveRequest{
Repository: "acme/site",
Selector: SelectorLatest,
AssetName: "dist.zip",
ETag: `W/"old"`,
})
if err != nil {
t.Fatalf("Resolve() error = %v", err)
}
if result.Release.ID != "9007199254740991" || result.Asset.ID != "9007199254740990" {
t.Fatalf("IDs lost precision: release=%q asset=%q", result.Release.ID, result.Asset.ID)
}
if result.ETag != `W/"new"` || result.Asset.Name != "dist.zip" || result.Asset.State != "uploaded" {
t.Fatalf("Resolve() = %+v", result)
}
if result.RetryAt == nil || result.RetryAt.Unix() != 1800000000 {
t.Fatalf("RetryAt = %v", result.RetryAt)
}
}
func TestResolveTagEscapesPathAndHandlesNotModified(t *testing.T) {
t.Parallel()
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
if request.RequestURI != "/repos/acme/site/releases/tags/release%2Fcandidate" {
t.Errorf("RequestURI = %q", request.RequestURI)
}
writer.WriteHeader(http.StatusNotModified)
}))
defer server.Close()
client := newTestClient(t, server.URL, nil)
result, err := client.Resolve(context.Background(), ResolveRequest{
Repository: "acme/site",
Selector: SelectorTag,
Tag: "release/candidate",
AssetName: "dist.zip",
ETag: `"cached"`,
})
if err != nil {
t.Fatalf("Resolve() error = %v", err)
}
if !result.NotModified || result.ETag != `"cached"` {
t.Fatalf("Resolve() = %+v", result)
}
}
func TestResolveAssetMissingTruncatesSafeNamesAndNeverIncludesBody(t *testing.T) {
t.Parallel()
assets := make([]string, 0, 12)
for index := 0; index < 12; index++ {
assets = append(assets, fmt.Sprintf(`{"id":%d,"name":"asset-%02d.zip","state":"uploaded","size":1}`, index+1, index))
}
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
_, _ = writer.Write([]byte(`{"id":1,"tag_name":"v1","message":"body-token","assets":[` + strings.Join(assets, ",") + `]}`))
}))
defer server.Close()
client := newTestClient(t, server.URL, nil)
_, err := client.Resolve(context.Background(), ResolveRequest{
Repository: "acme/site", Selector: SelectorLatest, AssetName: "dist.zip",
})
if !errors.Is(err, errAssetMissing) {
t.Fatalf("Resolve() error = %v", err)
}
message := err.Error()
if !strings.Contains(message, "asset-00.zip") || !strings.Contains(message, "asset-09.zip") {
t.Fatalf("error misses safe truncated names: %s", message)
}
if strings.Contains(message, "asset-10.zip") || strings.Contains(message, "asset-11.zip") || strings.Contains(message, "body-token") {
t.Fatalf("error leaked/truncation failed: %s", message)
}
}
func TestResolveHTTPErrorParsesRateLimitWithoutBodyLeak(t *testing.T) {
t.Parallel()
now := time.Date(2026, time.July, 19, 10, 0, 0, 0, time.UTC)
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
writer.Header().Set("Retry-After", "90")
writer.Header().Set("X-GitHub-Request-Id", "request-123")
writer.WriteHeader(http.StatusTooManyRequests)
_, _ = writer.Write([]byte(`{"message":"signed_url=https://secret.example/a?token=hidden"}`))
}))
defer server.Close()
client := newTestClient(t, server.URL, func(options *clientOptions) { options.now = func() time.Time { return now } })
_, err := client.Resolve(context.Background(), ResolveRequest{
Repository: "acme/site", Selector: SelectorLatest, AssetName: "dist.zip",
})
if err == nil || !strings.Contains(err.Error(), "status=429") || !strings.Contains(err.Error(), "request_id=request-123") {
t.Fatalf("Resolve() error = %v", err)
}
if strings.Contains(err.Error(), "secret.example") || strings.Contains(err.Error(), "hidden") {
t.Fatalf("error leaked body: %s", err)
}
retryAt, ok := RetryTime(err)
if !ok || !retryAt.Equal(now.Add(90*time.Second)) {
t.Fatalf("RetryTime() = %v, %v", retryAt, ok)
}
}
func TestDownloadStreamsVerifiesDigestAndCleansUp(t *testing.T) {
t.Parallel()
payload := []byte("package bytes")
digest := sha256.Sum256(payload)
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
if request.URL.Path != "/repos/acme/site/releases/assets/42" {
t.Errorf("path = %q", request.URL.Path)
}
assertHeader(t, request, "Accept", assetAccept)
assertHeader(t, request, "Accept-Encoding", "identity")
_, _ = writer.Write(payload)
}))
defer server.Close()
client := newTestClient(t, server.URL, nil)
result, err := client.Download(context.Background(), DownloadRequest{
Repository: "acme/site",
Asset: Asset{
ID: "42", Name: "dist.zip", Digest: "sha256:" + hex.EncodeToString(digest[:]),
},
MaxBytes: 1024,
})
if err != nil {
t.Fatalf("Download() error = %v", err)
}
if result.Size != int64(len(payload)) || result.SHA256 != hex.EncodeToString(digest[:]) {
t.Fatalf("Download() = %+v", result)
}
if _, err := os.Stat(result.Path); err != nil {
t.Fatalf("temp file stat: %v", err)
}
if err := result.Cleanup(); err != nil {
t.Fatalf("Cleanup() error = %v", err)
}
if err := result.Cleanup(); err != nil {
t.Fatalf("second Cleanup() error = %v", err)
}
}
func TestDownloadFollows302AndStripsCrossHostSensitiveHeaders(t *testing.T) {
t.Parallel()
var targetHost string
target := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
for _, header := range []string{"Authorization", "Cookie", "Proxy-Authorization", "Referer", "If-None-Match", "If-Modified-Since", "X-GitHub-Api-Version"} {
if value := request.Header.Get(header); value != "" {
t.Errorf("redirect leaked %s=%q", header, value)
}
}
_, _ = writer.Write([]byte("redirected package"))
}))
defer target.Close()
targetURL, _ := url.Parse(target.URL)
targetHost = "asset.example.test:" + targetURL.Port()
api := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
writer.Header().Set("Location", "http://"+targetHost+"/signed/package.zip?token=must-not-leak")
writer.WriteHeader(http.StatusFound)
}))
defer api.Close()
apiURL, _ := url.Parse(api.URL)
baseURL := "http://api.example.test:" + apiURL.Port()
client := newMappedTestClient(t, baseURL, nil)
result, err := client.Download(context.Background(), DownloadRequest{
Repository: "acme/site", Asset: Asset{ID: "42", Name: "dist.zip"}, MaxBytes: 1024,
})
if err != nil {
t.Fatalf("Download() redirect error = %v", err)
}
if cleanupErr := result.Cleanup(); cleanupErr != nil {
t.Fatalf("Cleanup() error = %v", cleanupErr)
}
request, _ := http.NewRequestWithContext(context.Background(), http.MethodGet, baseURL+"/repos/acme/site/releases/assets/42", nil)
applyAssetHeaders(request)
request.Header.Set("Authorization", "Bearer secret")
request.Header.Set("Cookie", "session=secret")
request.Header.Set("Proxy-Authorization", "proxy-secret")
request.Header.Set("Referer", "https://secret.example/path?token=x")
request.Header.Set("If-None-Match", `"secret-etag"`)
request.Header.Set("If-Modified-Since", time.Now().Format(http.TimeFormat))
response, err := client.httpClient.Do(request)
if err != nil {
t.Fatalf("Do() error = %v", err)
}
_ = response.Body.Close()
}
func TestRedirectSSRFAndDNSRebindingAreRejectedWithoutURLLeak(t *testing.T) {
t.Parallel()
t.Run("literal private redirect", func(t *testing.T) {
api := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
writer.Header().Set("Location", "http://127.0.0.1/private?token=secret-query")
writer.WriteHeader(http.StatusFound)
}))
defer api.Close()
client := newTestClient(t, api.URL, nil)
_, err := client.Download(context.Background(), DownloadRequest{
Repository: "acme/site", Asset: Asset{ID: "1", Name: "dist.zip"}, MaxBytes: 100,
})
if err == nil || strings.Contains(err.Error(), "secret-query") || strings.Contains(err.Error(), "127.0.0.1") {
t.Fatalf("Download() error = %v", err)
}
})
t.Run("DNS rebind between redirect and dial", func(t *testing.T) {
api := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
writer.Header().Set("Location", "http://rebind.example.test/package.zip")
writer.WriteHeader(http.StatusFound)
}))
defer api.Close()
var lock sync.Mutex
calls := map[string]int{}
resolve := resolverFunc(func(_ context.Context, _ string, host string) ([]netip.Addr, error) {
lock.Lock()
defer lock.Unlock()
calls[host]++
if host == "rebind.example.test" && calls[host] > 1 {
return []netip.Addr{netip.MustParseAddr("127.0.0.1")}, nil
}
return []netip.Addr{netip.MustParseAddr("8.8.8.8")}, nil
})
client := newTestClient(t, api.URL, func(options *clientOptions) { options.resolver = resolve })
_, err := client.Download(context.Background(), DownloadRequest{
Repository: "acme/site", Asset: Asset{ID: "1", Name: "dist.zip"}, MaxBytes: 100,
})
if err == nil {
t.Fatal("Download() error = nil")
}
lock.Lock()
defer lock.Unlock()
if calls["rebind.example.test"] != 2 {
t.Fatalf("rebind lookup calls = %d", calls["rebind.example.test"])
}
})
}
func TestDownloadFailureRemovesTemporaryFile(t *testing.T) {
t.Parallel()
payload := []byte("package bytes")
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
_, _ = writer.Write(payload)
}))
defer server.Close()
tempDir := t.TempDir()
client := newTestClient(t, server.URL, func(options *clientOptions) {
options.createTemp = func(_ string, pattern string) (*os.File, error) {
return os.CreateTemp(tempDir, pattern)
}
})
_, err := client.Download(context.Background(), DownloadRequest{
Repository: "acme/site",
Asset: Asset{
ID: "42", Name: "dist.zip", Digest: "sha256:" + strings.Repeat("0", 64),
},
MaxBytes: 1024,
})
if !errors.Is(err, errDigest) {
t.Fatalf("Download() error = %v", err)
}
files, readErr := filepath.Glob(filepath.Join(tempDir, "*"))
if readErr != nil || len(files) != 0 {
t.Fatalf("temporary files after failure = %v, err=%v", files, readErr)
}
}
func TestResolveRejectsInvalidRepositoryAndAssetWithoutRequest(t *testing.T) {
t.Parallel()
client := NewClient()
for _, request := range []ResolveRequest{
{Repository: "https://github.com/acme/site", Selector: SelectorLatest, AssetName: "dist.zip"},
{Repository: "acme/site/extra", Selector: SelectorLatest, AssetName: "dist.zip"},
{Repository: "acme/site", Selector: SelectorLatest, AssetName: "../dist.zip"},
{Repository: "acme/site", Selector: SelectorLatest, AssetName: `dir\dist.zip`},
{Repository: "acme/site", Selector: SelectorLatest, AssetName: string([]byte{'d', 'i', 's', 't', 0xff})},
{Repository: "acme/site", Selector: SelectorLatest, AssetName: "dist\nsecret.zip"},
{Repository: "acme/site", Selector: SelectorLatest, AssetName: "dist\u2028secret.zip"},
{Repository: "acme/site", Selector: SelectorLatest, AssetName: "dist\u202esecret.zip"},
{Repository: "acme/site", Selector: SelectorTag, AssetName: "dist.zip"},
} {
_, err := client.Resolve(context.Background(), request)
if !errors.Is(err, errInvalidRequest) {
t.Errorf("Resolve(%+v) error = %v", request, err)
}
}
}
func TestResolveAndDownloadAssetNameWithDelimiters(t *testing.T) {
t.Parallel()
assetName := "dist?channel=stable#1&x.zip"
payload := []byte("package with delimiter name")
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
switch request.URL.Path {
case "/repos/acme/site/releases/latest":
if request.Header.Get("If-None-Match") == "missing" {
_, _ = writer.Write([]byte(`{"id":1,"tag_name":"v1","assets":[]}`))
return
}
_, _ = fmt.Fprintf(writer, `{"id":1,"tag_name":"release/v1","assets":[{"id":42,"name":%q,"state":"uploaded","size":%d}]}`,
assetName, len(payload))
case "/repos/acme/site/releases/assets/42":
_, _ = writer.Write(payload)
default:
http.NotFound(writer, request)
}
}))
defer server.Close()
client := newTestClient(t, server.URL, nil)
resolved, err := client.Resolve(context.Background(), ResolveRequest{
Repository: "acme/site", Selector: SelectorLatest, AssetName: assetName,
})
if err != nil {
t.Fatalf("Resolve() error = %v", err)
}
if resolved.Asset.Name != assetName || resolved.Release.Tag != "release/v1" {
t.Fatalf("Resolve() = %+v", resolved)
}
download, err := client.Download(context.Background(), DownloadRequest{
Repository: "acme/site", Asset: resolved.Asset, MaxBytes: 1024,
})
if err != nil {
t.Fatalf("Download() error = %v", err)
}
if cleanupErr := download.Cleanup(); cleanupErr != nil {
t.Fatalf("Cleanup() error = %v", cleanupErr)
}
_, err = client.Resolve(context.Background(), ResolveRequest{
Repository: "acme/site", Selector: SelectorLatest, AssetName: assetName, ETag: "missing",
})
if !errors.Is(err, ErrAssetNotFound) {
t.Fatalf("missing Resolve() error = %v", err)
}
if strings.Contains(err.Error(), assetName) || strings.Contains(err.Error(), "channel=stable") {
t.Fatalf("missing error leaked delimiter-bearing name: %v", err)
}
}
func TestFixedTagGitRefRulesAndEscaping(t *testing.T) {
t.Parallel()
valid := []string{"@", "release/v1#stable&channel=prod", "foo.LOCK", "中文/发布=稳定"}
for _, tag := range valid {
if !validTag(tag) {
t.Errorf("validTag(%q) = false", tag)
}
}
invalid := []string{
"", "release v1", "release~v1", "release^v1", "release:v1", "release?v1", "release*v1",
"release[v1", `release\v1`, "release..v1", "release@{v1", "release//v1", "/release", "release/",
"release.", ".release", "release/.candidate", "release.lock", "release/v1.lock", "release\nsecret",
"release\u2028secret", "release\u202esecret", string([]byte{'v', '1', 0xff}), strings.Repeat("a", 256),
}
for _, tag := range invalid {
if validTag(tag) {
t.Errorf("validTag(%q) = true", tag)
}
}
tag := "release/v1#stable&channel=prod"
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
wantURI := "/repos/acme/site/releases/tags/" + url.PathEscape(tag)
if request.RequestURI != wantURI || request.URL.RawQuery != "" || request.URL.Fragment != "" {
t.Errorf("tag request = %q query=%q fragment=%q, want %q", request.RequestURI, request.URL.RawQuery, request.URL.Fragment, wantURI)
}
_, _ = fmt.Fprintf(writer, `{"id":1,"tag_name":%q,"assets":[{"id":2,"name":"dist.zip","state":"uploaded","size":1}]}`, tag)
}))
defer server.Close()
client := newTestClient(t, server.URL, nil)
result, err := client.Resolve(context.Background(), ResolveRequest{
Repository: "acme/site", Selector: SelectorTag, Tag: tag, AssetName: "dist.zip",
})
if err != nil {
t.Fatalf("Resolve() error = %v", err)
}
if result.Release.Tag != tag {
t.Fatalf("Release.Tag = %q", result.Release.Tag)
}
}
func TestResolveDoesNotMatchSanitizedRemoteAssetName(t *testing.T) {
t.Parallel()
tests := []struct {
name string
remote string
requested string
}{
{name: "unicode line separator", remote: "dist\u2028.zip", requested: "dist?.zip"},
{name: "overlong", remote: strings.Repeat("a", 256), requested: strings.Repeat("a", 255)},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
t.Parallel()
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
_, _ = fmt.Fprintf(writer, `{"id":1,"tag_name":"release/v1","assets":[{"id":2,"name":%q,"state":"uploaded","size":1}]}`, test.remote)
}))
defer server.Close()
client := newTestClient(t, server.URL, nil)
_, err := client.Resolve(context.Background(), ResolveRequest{
Repository: "acme/site", Selector: SelectorLatest, AssetName: test.requested,
})
if !errors.Is(err, ErrAssetNotFound) {
t.Fatalf("Resolve() error = %v", err)
}
})
}
}
func TestReleaseDisplayTagValidation(t *testing.T) {
t.Parallel()
for _, tag := range []string{"release/v1", "release v1", "release/v1#stable&channel=prod"} {
if !validReleaseDisplayTag(tag) {
t.Errorf("validReleaseDisplayTag(%q) = false", tag)
}
}
for _, tag := range []string{"", strings.Repeat("a", 256), "release\nsecret", "release\u2028secret", "release\u202esecret"} {
if validReleaseDisplayTag(tag) {
t.Errorf("validReleaseDisplayTag(%q) = true", tag)
}
}
}
func TestResolveRejectsInvalidUTF8Metadata(t *testing.T) {
t.Parallel()
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
_, _ = writer.Write(append([]byte(`{"id":1,"tag_name":"v1","assets":[{"id":2,"name":"dist`),
append([]byte{0xff}, []byte(`.zip","state":"uploaded","size":1}]}`)...)...))
}))
defer server.Close()
client := newTestClient(t, server.URL, nil)
_, err := client.Resolve(context.Background(), ResolveRequest{
Repository: "acme/site", Selector: SelectorLatest, AssetName: "dist�.zip",
})
if !errors.Is(err, ErrMetadata) {
t.Fatalf("Resolve() error = %v", err)
}
}
func TestDownloadRejectsImpossibleMetadataBeforeNetwork(t *testing.T) {
t.Parallel()
var requests atomic.Int64
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
requests.Add(1)
writer.WriteHeader(http.StatusInternalServerError)
}))
defer server.Close()
client := newTestClient(t, server.URL, nil)
tests := []struct {
name string
size int64
kind error
limit int64
}{
{name: "negative", size: -1, kind: ErrInvalidRequest, limit: 100},
{name: "declared too large", size: 101, kind: ErrAssetTooLarge, limit: 100},
}
for _, test := range tests {
_, err := client.Download(context.Background(), DownloadRequest{
Repository: "acme/site",
Asset: Asset{ID: "1", Name: "dist?token=hidden#asset.zip", Size: test.size},
MaxBytes: test.limit,
})
if !errors.Is(err, test.kind) {
t.Errorf("%s Download() error = %v", test.name, err)
}
if strings.Contains(err.Error(), "token=hidden") {
t.Errorf("%s error leaked asset name: %v", test.name, err)
}
}
if got := requests.Load(); got != 0 {
t.Fatalf("HTTP requests = %d, want 0", got)
}
}
func TestLogControlCharactersNeverEnterSafeErrors(t *testing.T) {
t.Parallel()
controls := []string{"\u2028", "\u2029", "\u061c", "\u200e", "\u200f", "\u202e", "\u2066", "\u2069"}
for _, control := range controls {
secret := "before" + control + "after"
err := safeError(errInvalidRequest, 0, secret, secret, secret, secret, nil, nil)
message := err.Error()
if strings.Contains(message, secret) || strings.Contains(message, control) || strings.Contains(message, "before") {
t.Errorf("safe error retained control %U: %q", []rune(control)[0], message)
}
}
}
func TestResolveRejectsMetadataOverHardLimit(t *testing.T) {
t.Parallel()
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
_, _ = writer.Write([]byte(`{"id":1,"tag_name":"v1","assets":[]}` + strings.Repeat(" ", maxMetadataBytes)))
}))
defer server.Close()
client := newTestClient(t, server.URL, nil)
_, err := client.Resolve(context.Background(), ResolveRequest{
Repository: "acme/site", Selector: SelectorLatest, AssetName: "dist.zip",
})
if !errors.Is(err, ErrMetadata) {
t.Fatalf("Resolve() error = %v", err)
}
}
func TestProductionTransportRejectsSelfSignedTLS(t *testing.T) {
t.Parallel()
server := httptest.NewTLSServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
_, _ = writer.Write([]byte("package"))
}))
defer server.Close()
parsed, _ := url.Parse(server.URL)
dialer := &net.Dialer{Timeout: time.Second}
client := newClient(clientOptions{
baseURL: "https://api.example.test:" + parsed.Port(),
resolver: resolverFunc(func(_ context.Context, _ string, _ string) ([]netip.Addr, error) {
return []netip.Addr{netip.MustParseAddr("8.8.8.8")}, nil
}),
dialContext: func(ctx context.Context, network string, _ string) (net.Conn, error) {
return dialer.DialContext(ctx, network, server.Listener.Addr().String())
},
tlsConfig: &tls.Config{MinVersion: tls.VersionTLS12},
createTemp: os.CreateTemp,
now: time.Now,
clientTimeout: 5 * time.Second,
})
_, err := client.Download(context.Background(), DownloadRequest{
Repository: "acme/site", Asset: Asset{ID: "1", Name: "dist.zip"}, MaxBytes: 100,
})
if !errors.Is(err, ErrDownload) {
t.Fatalf("Download() error = %v", err)
}
if strings.Contains(err.Error(), "api.example.test") || strings.Contains(err.Error(), server.URL) {
t.Fatalf("TLS error leaked URL: %v", err)
}
}
func TestStableErrorClassification(t *testing.T) {
t.Parallel()
now := time.Now()
assetMissing := safeError(errAssetMissing, http.StatusOK, "", "acme/site", "v1", "dist.zip", nil, nil)
if !IsNotFound(assetMissing) || IsRetryable(assetMissing) {
t.Fatalf("asset missing classification failed: %v", assetMissing)
}
metadata404 := safeError(errMetadata, http.StatusNotFound, "", "acme/site", "v1", "dist.zip", nil, nil)
if !IsNotFound(metadata404) || IsRetryable(metadata404) {
t.Fatalf("metadata 404 classification failed: %v", metadata404)
}
digest := safeError(errDigest, http.StatusOK, "", "acme/site", "", "dist.zip", nil, nil)
if !IsDigestError(digest) || IsRetryable(digest) {
t.Fatalf("digest classification failed: %v", digest)
}
for _, retryable := range []error{
safeError(errMetadata, 0, "", "acme/site", "", "dist.zip", nil, nil),
safeError(errMetadata, http.StatusOK, "", "acme/site", "", "dist.zip", nil, nil),
safeError(errDownload, http.StatusOK, "", "acme/site", "", "dist.zip", nil, nil),
safeError(errMetadata, http.StatusInternalServerError, "", "acme/site", "", "dist.zip", nil, nil),
safeError(errMetadata, http.StatusForbidden, "", "acme/site", "", "dist.zip", nil, &now),
} {
if !IsRetryable(retryable) {
t.Errorf("IsRetryable(%v) = false", retryable)
}
}
}
func TestDownloadRedirectLimitIsSafe(t *testing.T) {
t.Parallel()
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
step, _ := strconv.Atoi(request.URL.Query().Get("step"))
writer.Header().Set("Location", fmt.Sprintf("/repos/acme/site/releases/assets/1?step=%d&token=redirect-secret", step+1))
writer.WriteHeader(http.StatusFound)
}))
defer server.Close()
client := newTestClient(t, server.URL, nil)
_, err := client.Download(context.Background(), DownloadRequest{
Repository: "acme/site", Asset: Asset{ID: "1", Name: "dist.zip"}, MaxBytes: 100,
})
if err == nil {
t.Fatal("Download() error = nil")
}
if strings.Contains(err.Error(), "redirect-secret") || strings.Contains(err.Error(), "step=") {
t.Fatalf("redirect error leaked Location: %v", err)
}
}
func TestInvalidRequestDoesNotEchoURLQueryTagOrAsset(t *testing.T) {
t.Parallel()
client := NewClient()
requests := []ResolveRequest{
{Repository: "https://github.com/acme/site?token=repo-secret", Selector: SelectorLatest, AssetName: "dist.zip"},
{Repository: "acme/site", Selector: SelectorTag, Tag: "?token=tag-secret", AssetName: "dist.zip"},
{Repository: "acme/site", Selector: SelectorLatest, AssetName: "../dist.zip?token=asset-secret"},
}
for _, request := range requests {
_, err := client.Resolve(context.Background(), request)
if err == nil {
t.Fatalf("Resolve(%+v) error = nil", request)
}
for _, secret := range []string{"repo-secret", "tag-secret", "asset-secret", "https://github.com"} {
if strings.Contains(err.Error(), secret) {
t.Fatalf("Resolve(%+v) leaked %q: %v", request, secret, err)
}
}
}
}
func newTestClient(t *testing.T, rawBaseURL string, customize func(*clientOptions)) *Client {
t.Helper()
parsed, err := url.Parse(rawBaseURL)
if err != nil {
t.Fatal(err)
}
baseURL := "http://api.example.test:" + parsed.Port()
return newMappedTestClient(t, baseURL, customize)
}
func newMappedTestClient(t *testing.T, baseURL string, customize func(*clientOptions)) *Client {
t.Helper()
resolve := resolverFunc(func(_ context.Context, _ string, _ string) ([]netip.Addr, error) {
return []netip.Addr{netip.MustParseAddr("8.8.8.8")}, nil
})
dialer := &net.Dialer{Timeout: time.Second}
options := clientOptions{
baseURL: baseURL,
resolver: resolve,
allowHTTP: true,
tlsConfig: &tls.Config{MinVersion: tls.VersionTLS12},
dialContext: func(ctx context.Context, network string, address string) (net.Conn, error) {
_, port, splitErr := net.SplitHostPort(address)
if splitErr != nil {
return nil, splitErr
}
return dialer.DialContext(ctx, network, net.JoinHostPort("127.0.0.1", port))
},
createTemp: os.CreateTemp,
now: time.Now,
clientTimeout: 5 * time.Second,
}
if customize != nil {
customize(&options)
}
return newClient(options)
}
func assertHeader(t *testing.T, request *http.Request, name string, expected string) {
t.Helper()
if actual := request.Header.Get(name); actual != expected {
t.Errorf("%s = %q, want %q", name, actual, expected)
}
}
func TestResponseRetryAtHTTPDate(t *testing.T) {
t.Parallel()
want := time.Date(2026, time.July, 19, 12, 30, 0, 0, time.UTC)
response := &http.Response{Header: make(http.Header)}
response.Header.Set("Retry-After", want.Format(http.TimeFormat))
if got := responseRetryAt(response, time.Time{}); got == nil || !got.Equal(want) {
t.Fatalf("responseRetryAt() = %v", got)
}
}
func TestResponseRetryAtRejectsDurationOverflow(t *testing.T) {
t.Parallel()
response := &http.Response{Header: make(http.Header)}
response.Header.Set("Retry-After", strconv.FormatInt(maxRetryAfterSeconds+1, 10))
if got := responseRetryAt(response, time.Now()); got != nil {
t.Fatalf("responseRetryAt(overflow) = %v", got)
}
}
func TestSafeETagDropsOversizedOrControlValue(t *testing.T) {
t.Parallel()
if got := safeETag(strings.Repeat("x", 513)); got != "" {
t.Fatalf("safeETag(overlong) = %q", got)
}
if got := safeETag("ok\nsecret"); got != "" {
t.Fatalf("safeETag(control) = %q", got)
}
}
func TestDownloadSizeLimit(t *testing.T) {
t.Parallel()
server := httptest.NewServer(http.HandlerFunc(func(writer http.ResponseWriter, _ *http.Request) {
writer.Header().Set("Content-Length", strconv.Itoa(20))
_, _ = writer.Write([]byte(strings.Repeat("x", 20)))
}))
defer server.Close()
client := newTestClient(t, server.URL, nil)
_, err := client.Download(context.Background(), DownloadRequest{
Repository: "acme/site", Asset: Asset{ID: "1", Name: "dist.zip"}, MaxBytes: 10,
})
if !errors.Is(err, errTooLarge) {
t.Fatalf("Download() error = %v", err)
}
}
@@ -0,0 +1,312 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package githubrelease
import (
"context"
"crypto/tls"
"errors"
"math"
"net"
"net/http"
"net/netip"
"net/url"
"os"
"strconv"
"strings"
"time"
"github.com/Rain-kl/Wavelet/pkg/httppool"
)
const (
clientTimeout = 10 * time.Minute
dialTimeout = 30 * time.Second
dialKeepAlive = 30 * time.Second
responseHeaderTimeout = 30 * time.Second
maxRedirects = 5
maxRetryAfterSeconds = math.MaxInt64 / int64(time.Second)
)
var (
errBlockedTarget = errors.New("GitHub Release 请求目标不是公网地址")
errResolveTarget = errors.New("GitHub Release 请求目标解析失败")
errRedirectLimit = errors.New("GitHub Release asset 重定向次数过多")
publicIPv6Prefix = netip.MustParsePrefix("2000::/3")
nonPublicPrefixes = []netip.Prefix{
netip.MustParsePrefix("0.0.0.0/8"),
netip.MustParsePrefix("10.0.0.0/8"),
netip.MustParsePrefix("100.64.0.0/10"),
netip.MustParsePrefix("127.0.0.0/8"),
netip.MustParsePrefix("169.254.0.0/16"),
netip.MustParsePrefix("172.16.0.0/12"),
netip.MustParsePrefix("192.0.0.0/24"),
netip.MustParsePrefix("192.0.2.0/24"),
netip.MustParsePrefix("192.168.0.0/16"),
netip.MustParsePrefix("198.18.0.0/15"),
netip.MustParsePrefix("198.51.100.0/24"),
netip.MustParsePrefix("203.0.113.0/24"),
netip.MustParsePrefix("224.0.0.0/4"),
netip.MustParsePrefix("240.0.0.0/4"),
netip.MustParsePrefix("::/128"),
netip.MustParsePrefix("::1/128"),
netip.MustParsePrefix("::ffff:0:0/96"),
netip.MustParsePrefix("64:ff9b::/96"),
netip.MustParsePrefix("100::/64"),
netip.MustParsePrefix("2001:db8::/32"),
netip.MustParsePrefix("fc00::/7"),
netip.MustParsePrefix("fe80::/10"),
netip.MustParsePrefix("ff00::/8"),
}
)
type resolver interface {
LookupNetIP(context.Context, string, string) ([]netip.Addr, error)
}
type clientOptions struct {
baseURL string
resolver resolver
dialContext func(context.Context, string, string) (net.Conn, error)
tlsConfig *tls.Config
allowHTTP bool
createTemp func(string, string) (*os.File, error)
now func() time.Time
clientTimeout time.Duration
}
func defaultClientOptions() clientOptions {
dialer := &net.Dialer{Timeout: dialTimeout, KeepAlive: dialKeepAlive}
return clientOptions{
baseURL: defaultAPIBaseURL,
resolver: net.DefaultResolver,
dialContext: dialer.DialContext,
tlsConfig: &tls.Config{MinVersion: tls.VersionTLS12},
createTemp: os.CreateTemp,
now: time.Now,
clientTimeout: clientTimeout,
}
}
func newClient(options clientOptions) *Client {
if options.baseURL == "" {
options.baseURL = defaultAPIBaseURL
}
if options.resolver == nil {
options.resolver = net.DefaultResolver
}
if options.dialContext == nil {
dialer := &net.Dialer{Timeout: dialTimeout, KeepAlive: dialKeepAlive}
options.dialContext = dialer.DialContext
}
if options.tlsConfig == nil {
options.tlsConfig = &tls.Config{MinVersion: tls.VersionTLS12}
}
if options.createTemp == nil {
options.createTemp = os.CreateTemp
}
if options.now == nil {
options.now = time.Now
}
if options.clientTimeout <= 0 {
options.clientTimeout = clientTimeout
}
secureDial := publicDialer(options.resolver, options.dialContext)
transport := httppool.NewTransport(httppool.TransportOptions{
Proxy: nil,
DialContext: secureDial,
TLSClientConfig: options.tlsConfig,
ResponseHeaderTimeout: responseHeaderTimeout,
TraceFilter: func(request *http.Request) bool {
return request.URL == nil || request.URL.RawQuery == ""
},
})
httpClient := &http.Client{Timeout: options.clientTimeout, Transport: transport}
httpClient.CheckRedirect = func(next *http.Request, previous []*http.Request) error {
if len(previous) > maxRedirects {
return errRedirectLimit
}
if err := validateTarget(next.Context(), next.URL, options.resolver, options.allowHTTP); err != nil {
return err
}
if len(previous) > 0 && !sameHost(previous[len(previous)-1].URL, next.URL) {
stripCrossHostHeaders(next)
}
return nil
}
return &Client{
httpClient: httpClient,
baseURL: strings.TrimRight(options.baseURL, "/"),
createTemp: options.createTemp,
now: options.now,
}
}
func applyMetadataHeaders(request *http.Request, etag string) {
request.Header.Set("Accept", metadataAccept)
request.Header.Set("User-Agent", defaultUserAgent)
request.Header.Set("X-GitHub-Api-Version", APIVersion)
if etag = safeETag(etag); etag != "" {
request.Header.Set("If-None-Match", etag)
}
}
func applyAssetHeaders(request *http.Request) {
request.Header.Set("Accept", assetAccept)
request.Header.Set("Accept-Encoding", "identity")
request.Header.Set("User-Agent", defaultUserAgent)
request.Header.Set("X-GitHub-Api-Version", APIVersion)
}
func stripCrossHostHeaders(request *http.Request) {
for _, header := range []string{
"Authorization",
"Cookie",
"Proxy-Authorization",
"Referer",
"If-None-Match",
"If-Modified-Since",
"X-GitHub-Api-Version",
} {
request.Header.Del(header)
}
}
func sameHost(left *url.URL, right *url.URL) bool {
if left == nil || right == nil {
return false
}
return strings.EqualFold(left.Hostname(), right.Hostname()) && effectivePort(left) == effectivePort(right)
}
func effectivePort(target *url.URL) string {
if port := target.Port(); port != "" {
return port
}
if strings.EqualFold(target.Scheme, "https") {
return "443"
}
return "80"
}
func validateTarget(ctx context.Context, target *url.URL, targetResolver resolver, allowHTTP bool) error {
if target == nil || target.User != nil || target.Fragment != "" || target.Opaque != "" || target.Hostname() == "" {
return errBlockedTarget
}
isHTTPS := strings.EqualFold(target.Scheme, "https")
isAllowedHTTP := allowHTTP && strings.EqualFold(target.Scheme, "http")
if !isHTTPS && !isAllowedHTTP {
return errBlockedTarget
}
_, err := resolvePublicIPs(ctx, targetResolver, target.Hostname())
return err
}
func publicDialer(
targetResolver resolver,
directDial func(context.Context, string, string) (net.Conn, error),
) func(context.Context, string, string) (net.Conn, error) {
return func(ctx context.Context, network string, address string) (net.Conn, error) {
host, port, err := net.SplitHostPort(address)
if err != nil {
return nil, errResolveTarget
}
addresses, err := resolvePublicIPs(ctx, targetResolver, host)
if err != nil {
return nil, err
}
for _, resolved := range addresses {
if !ipMatchesNetwork(resolved, network) {
continue
}
connection, dialErr := directDial(ctx, network, net.JoinHostPort(resolved.String(), port))
if dialErr == nil {
return connection, nil
}
}
return nil, errDownload
}
}
func resolvePublicIPs(ctx context.Context, targetResolver resolver, host string) ([]netip.Addr, error) {
if strings.Contains(host, "%") {
return nil, errBlockedTarget
}
if literal, err := netip.ParseAddr(host); err == nil {
if !isPublicIP(literal) {
return nil, errBlockedTarget
}
return []netip.Addr{literal}, nil
}
if targetResolver == nil {
return nil, errResolveTarget
}
addresses, err := targetResolver.LookupNetIP(ctx, "ip", host)
if err != nil || len(addresses) == 0 {
return nil, errResolveTarget
}
for _, address := range addresses {
if !isPublicIP(address) {
return nil, errBlockedTarget
}
}
return addresses, nil
}
func isPublicIP(address netip.Addr) bool {
if !address.IsValid() || address.Zone() != "" {
return false
}
address = address.Unmap()
if !address.IsGlobalUnicast() {
return false
}
if address.Is6() && !publicIPv6Prefix.Contains(address) {
return false
}
for _, prefix := range nonPublicPrefixes {
if prefix.Contains(address) {
return false
}
}
return true
}
func ipMatchesNetwork(address netip.Addr, network string) bool {
switch network {
case "tcp4":
return address.Unmap().Is4()
case "tcp6":
return address.Unmap().Is6()
default:
return true
}
}
func responseRetryAt(response *http.Response, now time.Time) *time.Time {
if response == nil {
return nil
}
if retryAfter := strings.TrimSpace(response.Header.Get("Retry-After")); retryAfter != "" {
if seconds, err := strconv.ParseInt(retryAfter, 10, 64); err == nil && seconds >= 0 && seconds <= maxRetryAfterSeconds {
retryAt := now.Add(time.Duration(seconds) * time.Second)
return &retryAt
}
if retryAt, err := http.ParseTime(retryAfter); err == nil {
retryAt = retryAt.UTC()
return &retryAt
}
}
if strings.TrimSpace(response.Header.Get("X-RateLimit-Remaining")) != "0" {
return nil
}
reset, err := strconv.ParseInt(strings.TrimSpace(response.Header.Get("X-RateLimit-Reset")), 10, 64)
if err != nil || reset <= 0 {
return nil
}
retryAt := time.Unix(reset, 0).UTC()
return &retryAt
}
+25 -18
View File
@@ -18,22 +18,23 @@ const (
// PagesProject OpenFlare Pages 静态托管项目。
type PagesProject struct {
ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
Name string `json:"name" gorm:"size:255;not null"`
Slug string `json:"slug" gorm:"uniqueIndex;size:128;not null"`
Description string `json:"description" gorm:"type:text;not null;default:''"`
Enabled bool `json:"enabled" gorm:"not null;default:true"`
SPAFallbackEnabled bool `json:"spa_fallback_enabled" gorm:"not null;default:false"`
SPAFallbackPath string `json:"spa_fallback_path" gorm:"size:512;not null;default:'/index.html'"`
APIProxyEnabled bool `json:"api_proxy_enabled" gorm:"not null;default:false"`
APIProxyPath string `json:"api_proxy_path" gorm:"size:255;not null;default:''"`
APIProxyPass string `json:"api_proxy_pass" gorm:"size:2048;not null;default:''"`
APIProxyRewrite string `json:"api_proxy_rewrite" gorm:"size:255;not null;default:''"`
ActiveDeploymentID *uint `json:"active_deployment_id" gorm:"index"`
RootDir string `json:"root_dir" gorm:"size:512;not null;default:''"`
EntryFile string `json:"entry_file" gorm:"size:512;not null;default:'index.html'"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
Name string `json:"name" gorm:"size:255;not null"`
Slug string `json:"slug" gorm:"uniqueIndex;size:128;not null"`
Description string `json:"description" gorm:"type:text;not null;default:''"`
Enabled bool `json:"enabled" gorm:"not null;default:true"`
SPAFallbackEnabled bool `json:"spa_fallback_enabled" gorm:"not null;default:false"`
SPAFallbackPath string `json:"spa_fallback_path" gorm:"size:512;not null;default:'/index.html'"`
APIProxyEnabled bool `json:"api_proxy_enabled" gorm:"not null;default:false"`
APIProxyPath string `json:"api_proxy_path" gorm:"size:255;not null;default:''"`
APIProxyPass string `json:"api_proxy_pass" gorm:"size:2048;not null;default:''"`
APIProxyRewrite string `json:"api_proxy_rewrite" gorm:"size:255;not null;default:''"`
ActiveDeploymentID *uint `json:"active_deployment_id" gorm:"index"`
RootDir string `json:"root_dir" gorm:"size:512;not null;default:''"`
EntryFile string `json:"entry_file" gorm:"size:512;not null;default:'index.html'"`
ContentConfigVersion int `json:"-" gorm:"not null;default:0"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
UpdatedAt time.Time `json:"updated_at" gorm:"autoUpdateTime"`
}
// TableName 表名。
@@ -44,8 +45,8 @@ func (PagesProject) TableName() string {
// PagesDeployment OpenFlare Pages 不可变部署记录。
type PagesDeployment struct {
ID uint `json:"id" gorm:"primaryKey;autoIncrement"`
ProjectID uint `json:"project_id" gorm:"not null;index"`
DeploymentNumber int `json:"deployment_number" gorm:"not null"`
ProjectID uint `json:"project_id" gorm:"not null;index;uniqueIndex:idx_of_pages_deployments_project_number,priority:1;uniqueIndex:idx_of_pages_deployments_source_revision,priority:1,where:source_identity IS NOT NULL AND source_revision IS NOT NULL"`
DeploymentNumber int `json:"deployment_number" gorm:"not null;uniqueIndex:idx_of_pages_deployments_project_number,priority:2"`
Checksum string `json:"checksum" gorm:"size:64;not null;index"`
Status string `json:"status" gorm:"size:32;not null;default:'uploaded';index"`
UploadID uint64 `json:"upload_id,string" gorm:"not null;default:0;index"`
@@ -53,6 +54,12 @@ type PagesDeployment struct {
FileCount int `json:"file_count" gorm:"not null;default:0"`
TotalSize int64 `json:"total_size" gorm:"not null;default:0"`
CreatedBy string `json:"created_by" gorm:"size:64;not null;default:''"`
SourceType string `json:"source_type" gorm:"size:32;not null;default:''"`
SourceIdentity *string `json:"-" gorm:"type:char(64);uniqueIndex:idx_of_pages_deployments_source_revision,priority:2,where:source_identity IS NOT NULL AND source_revision IS NOT NULL"`
SourceRevision *string `json:"-" gorm:"type:char(64);uniqueIndex:idx_of_pages_deployments_source_revision,priority:3,where:source_identity IS NOT NULL AND source_revision IS NOT NULL"`
SourceLabel string `json:"source_label" gorm:"size:255;not null;default:''"`
SourceMeta string `json:"-" gorm:"type:text;not null;default:''"`
TriggerType string `json:"trigger_type" gorm:"size:32;not null;default:''"`
CreatedAt time.Time `json:"created_at" gorm:"autoCreateTime"`
ActivatedAt *time.Time `json:"activated_at"`
}
+75
View File
@@ -0,0 +1,75 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"context"
"errors"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
)
const (
// PagesOrphanUploadCandidateLimit bounds one delayed Pages upload cleanup pass.
PagesOrphanUploadCandidateLimit = 100
pagesOrphanMarkerPredicatePostgres = "w_uploads.metadata #>> '{extra,pages_ingest_marker}' = ?"
pagesOrphanMarkerPredicateSQLite = "CASE WHEN json_valid(w_uploads.metadata) THEN json_extract(w_uploads.metadata, '$.extra.pages_ingest_marker') ELSE NULL END = ?"
)
// PagesOrphanUploadCandidateQuery describes the fail-closed SQL candidate set
// for delayed Pages upload compensation.
type PagesOrphanUploadCandidateQuery struct {
SystemUserID uint64
UploadType string
Marker string
CreatedBefore time.Time
}
// ListPagesOrphanUploadCandidates returns at most 100 unreferenced, isolated
// Pages V2 upload records. Callers must still lock and recheck every condition
// before deleting a candidate.
func ListPagesOrphanUploadCandidates(
ctx context.Context,
input PagesOrphanUploadCandidateQuery,
) ([]Upload, error) {
if input.SystemUserID == 0 || input.UploadType == "" || input.Marker == "" || input.CreatedBefore.IsZero() {
return nil, errors.New("invalid pages orphan upload candidate query")
}
markerPredicate, err := pagesOrphanMarkerPredicate(db.DB(ctx).Name())
if err != nil {
return nil, err
}
deploymentTable := (PagesDeployment{}).TableName()
uploadTable := (Upload{}).TableName()
var candidates []Upload
err = db.DB(ctx).
Model(&Upload{}).
Where(uploadTable+".status = ?", UploadStatusUsed).
Where(uploadTable+".user_id = ?", input.SystemUserID).
Where(uploadTable+".type = ?", input.UploadType).
Where(uploadTable+".created_at < ?", input.CreatedBefore).
Where(markerPredicate, input.Marker).
Where("NOT EXISTS (SELECT 1 FROM " + deploymentTable + " WHERE " + deploymentTable + ".upload_id = " + uploadTable + ".id)").
Order(uploadTable + ".id ASC").
Limit(PagesOrphanUploadCandidateLimit).
Find(&candidates).Error
if err != nil {
return nil, err
}
return candidates, nil
}
func pagesOrphanMarkerPredicate(dialect string) (string, error) {
switch dialect {
case "postgres":
return pagesOrphanMarkerPredicatePostgres, nil
case "sqlite":
return pagesOrphanMarkerPredicateSQLite, nil
default:
return "", errors.New("unsupported database dialect for Pages orphan cleanup")
}
}
@@ -0,0 +1,193 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"context"
"strings"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/glebarez/sqlite"
"gorm.io/gorm"
)
func TestPagesOrphanMarkerPredicate(t *testing.T) {
tests := []struct {
name string
dialect string
want string
wantErr bool
}{
{
name: "postgres jsonb path",
dialect: "postgres",
want: "metadata #>> '{extra,pages_ingest_marker}'",
},
{
name: "sqlite guarded json extract",
dialect: "sqlite",
want: "CASE WHEN json_valid(w_uploads.metadata) THEN json_extract",
},
{
name: "unknown dialect rejected",
dialect: "mysql",
wantErr: true,
},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
got, err := pagesOrphanMarkerPredicate(test.dialect)
if gotErr := err != nil; gotErr != test.wantErr {
t.Fatalf("pagesOrphanMarkerPredicate(%q) error = %v, want error presence = %t", test.dialect, err, test.wantErr)
}
if test.want != "" && !strings.Contains(got, test.want) {
t.Errorf("pagesOrphanMarkerPredicate(%q) = %q, want substring %q", test.dialect, got, test.want)
}
})
}
}
func TestListPagesOrphanUploadCandidatesFiltersAndLimits(t *testing.T) {
ctx := context.Background()
gormDB := setupPagesCleanupModelTestDB(t)
cutoff := time.Now().UTC().Add(-2 * time.Hour)
old := cutoff.Add(-time.Minute)
marker := UploadMetadata{Extra: map[string]any{
"pages_ingest_marker": "pages_deployment_v2",
"pages_project_id": "1",
}}
valid := make([]Upload, 0, PagesOrphanUploadCandidateLimit+1)
for index := 0; index < PagesOrphanUploadCandidateLimit+1; index++ {
valid = append(valid, pagesCleanupModelUpload(uint64(index+100), 999, "openflare_pages_deployment", UploadStatusUsed, old, marker))
}
if err := gormDB.Create(&valid).Error; err != nil {
t.Fatalf("create valid candidates error = %v, want nil", err)
}
referenced := pagesCleanupModelUpload(1, 999, "openflare_pages_deployment", UploadStatusUsed, old, marker)
wrongOwner := pagesCleanupModelUpload(2, 1000, "openflare_pages_deployment", UploadStatusUsed, old, marker)
wrongType := pagesCleanupModelUpload(3, 999, "generic", UploadStatusUsed, old, marker)
wrongStatus := pagesCleanupModelUpload(4, 999, "openflare_pages_deployment", UploadStatusPending, old, marker)
fresh := pagesCleanupModelUpload(5, 999, "openflare_pages_deployment", UploadStatusUsed, cutoff, marker)
wrongMarker := pagesCleanupModelUpload(6, 999, "openflare_pages_deployment", UploadStatusUsed, old, UploadMetadata{Extra: map[string]any{
"pages_ingest_marker": "pages_deployment_v1",
"pages_project_id": "1",
}})
for _, upload := range []Upload{referenced, wrongOwner, wrongType, wrongStatus, fresh, wrongMarker} {
if err := gormDB.Create(&upload).Error; err != nil {
t.Fatalf("create filtered upload %d error = %v, want nil", upload.ID, err)
}
}
if err := gormDB.Create(&PagesDeployment{
ProjectID: 1,
DeploymentNumber: 1,
Checksum: "referenced",
Status: PagesDeploymentStatusUploaded,
UploadID: referenced.ID,
}).Error; err != nil {
t.Fatalf("create referenced deployment error = %v, want nil", err)
}
invalidJSON := pagesCleanupModelUpload(7, 999, "openflare_pages_deployment", UploadStatusUsed, old, marker)
if err := gormDB.Create(&invalidJSON).Error; err != nil {
t.Fatalf("create invalid JSON upload error = %v, want nil", err)
}
if err := gormDB.Table((Upload{}).TableName()).Where("id = ?", invalidJSON.ID).
UpdateColumn("metadata", "{invalid").Error; err != nil {
t.Fatalf("corrupt upload metadata error = %v, want nil", err)
}
got, err := ListPagesOrphanUploadCandidates(ctx, PagesOrphanUploadCandidateQuery{
SystemUserID: 999,
UploadType: "openflare_pages_deployment",
Marker: "pages_deployment_v2",
CreatedBefore: cutoff,
})
if err != nil {
t.Fatalf("ListPagesOrphanUploadCandidates() error = %v, want nil", err)
}
if len(got) != PagesOrphanUploadCandidateLimit {
t.Fatalf("ListPagesOrphanUploadCandidates() count = %d, want %d", len(got), PagesOrphanUploadCandidateLimit)
}
for index, candidate := range got {
wantID := uint64(index + 100)
if candidate.ID != wantID {
t.Errorf("ListPagesOrphanUploadCandidates()[%d].ID = %d, want %d", index, candidate.ID, wantID)
}
}
}
func TestListPagesOrphanUploadCandidatesSkipsInvalidSQLiteJSON(t *testing.T) {
ctx := context.Background()
gormDB := setupPagesCleanupModelTestDB(t)
cutoff := time.Now().UTC().Add(-2 * time.Hour)
upload := pagesCleanupModelUpload(1, 999, "openflare_pages_deployment", UploadStatusUsed, cutoff.Add(-time.Minute), UploadMetadata{})
if err := gormDB.Create(&upload).Error; err != nil {
t.Fatalf("create invalid JSON candidate error = %v, want nil", err)
}
if err := gormDB.Table((Upload{}).TableName()).Where("id = ?", upload.ID).
UpdateColumn("metadata", "{invalid").Error; err != nil {
t.Fatalf("corrupt upload metadata error = %v, want nil", err)
}
got, err := ListPagesOrphanUploadCandidates(ctx, PagesOrphanUploadCandidateQuery{
SystemUserID: 999,
UploadType: "openflare_pages_deployment",
Marker: "pages_deployment_v2",
CreatedBefore: cutoff,
})
if err != nil {
t.Fatalf("ListPagesOrphanUploadCandidates(invalid JSON) error = %v, want nil", err)
}
if len(got) != 0 {
t.Errorf("ListPagesOrphanUploadCandidates(invalid JSON) count = %d, want 0", len(got))
}
}
func setupPagesCleanupModelTestDB(t *testing.T) *gorm.DB {
t.Helper()
gormDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
if err != nil {
t.Fatalf("open Pages cleanup model test database error = %v, want nil", err)
}
if err := gormDB.AutoMigrate(&Upload{}, &PagesDeployment{}); err != nil {
t.Fatalf("migrate Pages cleanup model test database error = %v, want nil", err)
}
db.SetDB(gormDB)
t.Cleanup(func() { db.SetDB(nil) })
return gormDB
}
func pagesCleanupModelUpload(
id uint64,
userID uint64,
uploadType string,
status UploadStatus,
createdAt time.Time,
metadata UploadMetadata,
) Upload {
return Upload{
ID: id,
UserID: userID,
FileName: "site.zip",
FilePath: "pages/site.zip",
FileSize: 10,
MimeType: "application/zip",
Extension: "zip",
Hash: "checksum",
Type: uploadType,
Status: status,
AccessMode: 0,
Metadata: metadata,
CreatedAt: createdAt,
UpdatedAt: createdAt,
}
}
+59
View File
@@ -0,0 +1,59 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import "time"
// PagesProjectSource 保存 Pages 项目的持久部署源配置。
//
// RemoteURL 可能包含签名参数,禁止直接序列化 model;对外接口必须映射到
// pages 包内的脱敏 source view。
type PagesProjectSource struct {
ID uint `json:"-" gorm:"primaryKey;autoIncrement"`
ProjectID uint `json:"-" gorm:"not null;uniqueIndex:idx_of_pages_project_sources_project_id"`
SourceType string `json:"-" gorm:"size:32;not null;default:''"`
RemoteURL string `json:"-" gorm:"type:text;not null;default:''"`
RemoteNetworkPolicy string `json:"-" gorm:"size:32;not null;default:''"`
GitHubRepository string `json:"-" gorm:"column:github_repository;size:255;not null;default:''"`
ReleaseSelector string `json:"-" gorm:"size:16;not null;default:''"`
ReleaseTag string `json:"-" gorm:"size:255;not null;default:''"`
AssetName string `json:"-" gorm:"size:255;not null;default:''"`
AutoUpdateEnabled bool `json:"-" gorm:"not null;default:false"`
CheckIntervalMinutes int `json:"-" gorm:"not null;default:0"`
ConfigVersion int `json:"-" gorm:"not null;default:0"`
SourceIdentity string `json:"-" gorm:"type:char(64);not null;default:''"`
CreatedAt time.Time `json:"-" gorm:"autoCreateTime"`
UpdatedAt time.Time `json:"-" gorm:"autoUpdateTime"`
}
// TableName 返回 Pages 项目部署源配置表名。
func (PagesProjectSource) TableName() string {
return "of_pages_project_sources"
}
// PagesProjectSourceRuntime 保存 Pages 项目部署源的可变运行态。
//
// Runtime 不冗余 project_id;调用方通过 SourceID 关联配置,并在最终提交时
// 同时校验 source config version 与 project content config version。
type PagesProjectSourceRuntime struct {
SourceID uint `json:"-" gorm:"primaryKey;autoIncrement:false"`
ETag string `json:"-" gorm:"column:etag;size:512;not null;default:''"`
LastSeenRevision string `json:"-" gorm:"type:char(64);not null;default:''"`
LastSeenDetail string `json:"-" gorm:"type:text;not null;default:''"`
LastAppliedRevision string `json:"-" gorm:"type:char(64);not null;default:''"`
LastAppliedDetail string `json:"-" gorm:"type:text;not null;default:''"`
SyncStatus string `json:"-" gorm:"size:32;not null;default:''"`
LastError string `json:"-" gorm:"type:text;not null;default:''"`
LastCheckedAt *time.Time `json:"-"`
LastSyncedAt *time.Time `json:"-"`
NextCheckAt *time.Time `json:"-" gorm:"index:idx_of_pages_project_source_runtime_next_check_at"`
LeaseExpiresAt *time.Time `json:"-"`
LeaseToken string `json:"-" gorm:"size:64;not null;default:''"`
UpdatedAt time.Time `json:"-" gorm:"autoUpdateTime"`
}
// TableName 返回 Pages 项目部署源运行态表名。
func (PagesProjectSourceRuntime) TableName() string {
return "of_pages_project_source_runtime"
}
@@ -0,0 +1,80 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package model
import (
"encoding/json"
"strings"
"testing"
"github.com/glebarez/sqlite"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/gorm"
)
func TestPagesSourceModelsMatchMigrationSchema(t *testing.T) {
gormDB, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
DisableForeignKeyConstraintWhenMigrating: true,
})
require.NoError(t, err)
require.NoError(t, gormDB.AutoMigrate(
&PagesProject{},
&PagesDeployment{},
&PagesProjectSource{},
&PagesProjectSourceRuntime{},
))
assert.Equal(t, "of_pages_project_sources", (PagesProjectSource{}).TableName())
assert.Equal(t, "of_pages_project_source_runtime", (PagesProjectSourceRuntime{}).TableName())
assert.True(t, gormDB.Migrator().HasColumn(&PagesProjectSource{}, "github_repository"))
assert.False(t, gormDB.Migrator().HasColumn(&PagesProjectSource{}, "git_hub_repository"))
assert.True(t, gormDB.Migrator().HasColumn(&PagesProjectSourceRuntime{}, "etag"))
assert.False(t, gormDB.Migrator().HasColumn(&PagesProjectSourceRuntime{}, "e_tag"))
var indexSQL string
require.NoError(t, gormDB.Raw(
"SELECT sql FROM sqlite_master WHERE type = 'index' AND name = ?",
"idx_of_pages_deployments_source_revision",
).Scan(&indexSQL).Error)
assert.Contains(t, strings.ToUpper(indexSQL), "WHERE SOURCE_IDENTITY IS NOT NULL AND SOURCE_REVISION IS NOT NULL")
}
func TestPagesSourceModelsDoNotSerializeSecretsOrFencingState(t *testing.T) {
sourceJSON, err := json.Marshal(PagesProjectSource{
ID: 1,
ProjectID: 2,
RemoteURL: "https://example.com/site.zip?token=secret",
ConfigVersion: 3,
SourceIdentity: strings.Repeat("a", 64),
})
require.NoError(t, err)
assert.JSONEq(t, `{}`, string(sourceJSON))
runtimeJSON, err := json.Marshal(PagesProjectSourceRuntime{
SourceID: 1,
ETag: `"secret-etag"`,
LeaseToken: "secret-lease",
})
require.NoError(t, err)
assert.JSONEq(t, `{}`, string(runtimeJSON))
identity := strings.Repeat("b", 64)
revision := strings.Repeat("c", 64)
deploymentJSON, err := json.Marshal(PagesDeployment{
SourceType: "remote_url",
SourceIdentity: &identity,
SourceRevision: &revision,
SourceLabel: "site.zip",
SourceMeta: `{"provider":"remote_url","private":"secret"}`,
TriggerType: "manual_sync",
})
require.NoError(t, err)
assert.NotContains(t, string(deploymentJSON), identity)
assert.NotContains(t, string(deploymentJSON), revision)
assert.NotContains(t, string(deploymentJSON), "private")
assert.Contains(t, string(deploymentJSON), `"source_type":"remote_url"`)
assert.Contains(t, string(deploymentJSON), `"source_label":"site.zip"`)
assert.Contains(t, string(deploymentJSON), `"trigger_type":"manual_sync"`)
}
+12 -5
View File
@@ -62,15 +62,22 @@ func GetActiveUploadByID(ctx context.Context, id uint64) (model.Upload, error) {
return upload, nil
}
// SoftDeleteUpload marks an upload as deleted.
// SoftDeleteUpload marks an active upload as deleted and reports whether the row transitioned.
// External modules must use upload.Remove or upload.RemoveOwned; only internal/apps/upload may call this.
func SoftDeleteUpload(ctx context.Context, upload *model.Upload) error {
func SoftDeleteUpload(ctx context.Context, upload *model.Upload) (int64, error) {
return SoftDeleteUploadTx(db.DB(ctx), upload)
}
// SoftDeleteUploadTx marks an upload as deleted within an existing transaction.
func SoftDeleteUploadTx(tx *gorm.DB, upload *model.Upload) error {
return tx.Model(upload).Update("status", model.UploadStatusDeleted).Error
// SoftDeleteUploadTx marks an active upload as deleted within an existing transaction.
// RowsAffected is one only for the single successful active-to-deleted transition.
func SoftDeleteUploadTx(tx *gorm.DB, upload *model.Upload) (int64, error) {
result := tx.Model(&model.Upload{}).
Where("id = ? AND status IN ?", upload.ID, []model.UploadStatus{
model.UploadStatusPending,
model.UploadStatusUsed,
}).
Update("status", model.UploadStatusDeleted)
return result.RowsAffected, result.Error
}
// UpdateUpload applies partial field updates to an upload record.
@@ -18,6 +18,11 @@ func registerPagesRoutes(apiGroup *gin.RouterGroup) {
apiutil.RegisterCollection(pagesRoute, "POST", pages.CreateProjectHandler)
pagesRoute.POST("/:id/update", pages.UpdateProjectHandler)
pagesRoute.POST("/:id/delete", pages.DeleteProjectHandler)
pagesRoute.GET("/:id/source", pages.GetSourceHandler)
pagesRoute.POST("/:id/source/update", pages.UpdateSourceHandler)
pagesRoute.POST("/:id/source/delete", pages.DeleteSourceHandler)
pagesRoute.POST("/:id/source/check", pages.CheckSourceHandler)
pagesRoute.POST("/:id/source/sync", pages.SyncSourceHandler)
pagesRoute.GET("/:id/deployments", pages.ListDeploymentsHandler)
pagesRoute.POST("/:id/deployments/upload", pages.UploadDeploymentHandler)
pagesRoute.POST("/:id/deployments/upload-from-url", pages.UploadDeploymentFromURLHandler)
+5 -1
View File
@@ -428,7 +428,7 @@ func notifyTaskCompleted(ctx context.Context, execution *model.TaskExecution, re
}
func shouldFlushTaskExecutionLog(ctx context.Context, execErr error) bool {
if execErr == nil {
if isTerminalTaskExecutionError(execErr) {
return true
}
@@ -440,6 +440,10 @@ func shouldFlushTaskExecutionLog(ctx context.Context, execErr error) bool {
return retryCount >= maxRetry
}
func isTerminalTaskExecutionError(execErr error) bool {
return execErr == nil || errors.Is(execErr, asynq.SkipRetry)
}
func handleFailedTask(ctx context.Context, execution *model.TaskExecution, t *asynq.Task, duration time.Duration, execErr error, span trace.Span) {
execution.Status = model.TaskExecutionStatusFailed
execution.ErrorMessage = execErr.Error()
+51
View File
@@ -6,10 +6,12 @@ package task
import (
"context"
"errors"
"fmt"
"testing"
"time"
"github.com/Rain-kl/Wavelet/internal/db"
"github.com/Rain-kl/Wavelet/internal/model"
"github.com/Rain-kl/Wavelet/internal/testhelper"
"github.com/hibiken/asynq"
@@ -284,6 +286,55 @@ func TestCompleteTaskExecutionFlushesLog(t *testing.T) {
assert.Contains(t, found.Log, "任务执行成功")
}
func TestCompleteTaskExecutionFlushesPermanentFailureLog(t *testing.T) {
cleanup := setupTest(t)
defer cleanup()
ctx := context.Background()
execution := &model.TaskExecution{
TaskID: "complete_permanent_flush_001",
TaskType: testTaskType,
TaskName: "测试任务",
Status: model.TaskExecutionStatusRunning,
Retryable: true,
MaxRetry: 3,
TriggeredBy: "manual",
}
err := model.CreateTaskExecution(ctx, execution)
require.NoError(t, err)
ctx = withTaskID(ctx, execution.TaskID)
AppendLog(ctx, "永久失败前的日志")
execErr := PermanentError("来源配置无效")
finishTime := time.Now()
completeTaskExecution(
ctx,
execution,
asynq.NewTask(testTaskType, nil),
100*time.Millisecond,
finishTime,
nil,
execErr,
trace.SpanFromContext(ctx),
)
found, err := model.GetTaskExecutionByTaskID(ctx, execution.TaskID)
require.NoError(t, err)
assert.Equal(t, model.TaskExecutionStatusFailed, found.Status)
assert.Equal(t, "来源配置无效", found.ErrorMessage)
assert.Contains(t, found.Log, "永久失败前的日志")
assert.Contains(t, found.Log, "任务执行失败")
keys, err := db.Redis.Keys(ctx, "*"+execution.TaskID+"*").Result()
require.NoError(t, err)
assert.Empty(t, keys)
}
func TestPermanentErrorIsTerminalForLogFlush(t *testing.T) {
assert.True(t, isTerminalTaskExecutionError(PermanentError("配置无效")))
assert.False(t, isTerminalTaskExecutionError(errors.New("temporary failure")))
}
func TestRetryTask(t *testing.T) {
cleanup := setupTest(t)
defer cleanup()
+8
View File
@@ -8,6 +8,7 @@ package handlers
import (
"github.com/Rain-kl/Wavelet/internal/apps/admin/push"
"github.com/Rain-kl/Wavelet/internal/apps/openflare"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/pages"
"github.com/Rain-kl/Wavelet/internal/apps/openflare/tls"
"github.com/Rain-kl/Wavelet/internal/apps/upload"
"github.com/Rain-kl/Wavelet/internal/apps/user"
@@ -51,6 +52,13 @@ func Register() {
task.RegisterHandler(openflare.UptimeKumaSyncTask, &openflare.UptimeKumaSyncHandler{})
task.RegisterTaskMeta(openflare.UptimeKumaSyncMeta)
// pages source actions are only dispatched by the Pages domain API/scanner.
task.RegisterHandler(pages.PagesSourceScanTask, &pages.SourceScanHandler{})
task.RegisterTaskMeta(pages.PagesSourceScanMeta)
task.RegisterHandler(pages.PagesSourceActionTask, &pages.SourceActionHandler{})
task.RegisterTaskMeta(pages.PagesSourceActionMeta)
// tls single renew
task.RegisterHandler(tls.SSLSingleRenewTask, &tls.SSLSingleRenewHandler{})
task.RegisterTaskMeta(tls.SSLSingleRenewMeta)
+9 -3
View File
@@ -31,6 +31,7 @@ type TaskMeta struct {
MaxRetry int `json:"max_retry"`
Queue string `json:"queue"`
Retryable bool `json:"retryable"` // 是否支持手动重试
InternalOnly bool `json:"-"` // 是否仅允许内部业务入口调度
Params []TaskParam `json:"params,omitempty"`
}
@@ -51,13 +52,18 @@ func RegisterTaskMeta(meta TaskMeta) {
dispatchableTasks = append(dispatchableTasks, meta)
}
// GetDispatchableTasks 获取所有已注册的元数据列表(返回副本以避免并发并发读写冲突)
// GetDispatchableTasks 获取允许通过通用 Admin 入口调度的元数据列表。
func GetDispatchableTasks() []TaskMeta {
dispatchableTasksMutex.RLock()
defer dispatchableTasksMutex.RUnlock()
metas := make([]TaskMeta, len(dispatchableTasks))
copy(metas, dispatchableTasks)
metas := make([]TaskMeta, 0, len(dispatchableTasks))
for _, meta := range dispatchableTasks {
if meta.InternalOnly {
continue
}
metas = append(metas, meta)
}
return metas
}
+25
View File
@@ -29,3 +29,28 @@ func TestDuplicateTaskMeta(t *testing.T) {
}
}
}
func TestInternalOnlyTaskMetaIsHiddenFromDispatchableTasks(t *testing.T) {
const taskType = "test_internal_only_meta"
meta := task.TaskMeta{
Type: taskType,
AsynqTask: "test:internal_only_meta",
Name: "内部测试任务",
InternalOnly: true,
}
task.RegisterTaskMeta(meta)
registered := task.GetTaskMeta(taskType)
if registered == nil {
t.Fatal("GetTaskMeta() did not return internal-only metadata")
}
if !registered.InternalOnly {
t.Fatal("GetTaskMeta() lost InternalOnly flag")
}
for _, dispatchable := range task.GetDispatchableTasks() {
if dispatchable.Type == taskType {
t.Fatalf("GetDispatchableTasks() exposed internal-only task %q", taskType)
}
}
}
+35
View File
@@ -0,0 +1,35 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package task
import (
"strings"
"github.com/hibiken/asynq"
)
const defaultPermanentErrorMessage = "任务无法继续执行"
type permanentTaskError struct {
message string
}
// PermanentError marks a safe domain message as a non-retryable task failure.
// It intentionally accepts no underlying error so Error never exposes provider,
// URL, header, response-body, or other sensitive implementation details.
func PermanentError(message string) error {
message = strings.TrimSpace(message)
if message == "" {
message = defaultPermanentErrorMessage
}
return &permanentTaskError{message: message}
}
func (e *permanentTaskError) Error() string {
return e.message
}
func (e *permanentTaskError) Unwrap() error {
return asynq.SkipRetry
}
+27
View File
@@ -0,0 +1,27 @@
// Copyright 2026 Arctel.net
// SPDX-License-Identifier: Apache-2.0
package task
import (
"errors"
"testing"
"github.com/hibiken/asynq"
"github.com/stretchr/testify/assert"
)
func TestPermanentErrorSkipsRetryWithoutExposingAsynqMessage(t *testing.T) {
err := PermanentError(" 来源配置无效 ")
assert.True(t, errors.Is(err, asynq.SkipRetry))
assert.Equal(t, "来源配置无效", err.Error())
assert.NotContains(t, err.Error(), asynq.SkipRetry.Error())
}
func TestPermanentErrorUsesSafeFallbackForBlankMessage(t *testing.T) {
err := PermanentError(" ")
assert.True(t, errors.Is(err, asynq.SkipRetry))
assert.Equal(t, defaultPermanentErrorMessage, err.Error())
}