From 4d1dd37438b7c7001f572695feb5e1b56b9fc78d Mon Sep 17 00:00:00 2001 From: Rogee Date: Sun, 13 Sep 2026 21:52:52 +0800 Subject: [PATCH] feat: complete CreatorHub gateway review fixes --- AGENTS.md | 12 +- Dockerfile | 16 +- cmd/__init__.py | 0 cmd/control-plane/creator.go | 416 ++- cmd/control-plane/creator_events.go | 464 ++++ cmd/control-plane/creator_events_test.go | 99 + cmd/control-plane/creator_material.go | 210 ++ cmd/control-plane/creator_material_test.go | 53 + cmd/control-plane/hub.go | 16 +- cmd/control-plane/main.go | 37 +- cmd/docker-gateway/douyin.go | 509 ---- .../douyin_cdp_integration_test.go | 205 -- cmd/docker-gateway/douyin_test.go | 410 --- cmd/docker-gateway/main.go | 1547 ----------- cmd/docker-gateway/main_test.go | 2459 ----------------- cmd/docker-gateway/proxy.go | 435 --- cmd/docker-gateway/proxy_test.go | 252 -- cmd/docker_gateway/__init__.py | 0 cmd/docker_gateway/docker_client.py | 975 +++++++ cmd/docker_gateway/douyin.py | 1333 +++++++++ cmd/docker_gateway/gateway.py | 1608 +++++++++++ cmd/docker_gateway/proxy.py | 656 +++++ cmd/docker_gateway/test_gateway.py | 2076 ++++++++++++++ compose.yaml | 17 +- docker/browser-wrapper/Dockerfile | 5 + docker/browser-wrapper/README.md | 12 + docker/browser-wrapper/docker-entrypoint.sh | 32 + docs/architecture/container-control.md | 11 +- docs/deployment.md | 23 +- docs/plan01.md | 30 +- docs/python-gateway-branch-review.md | 477 ++++ internal/creator/accounts.go | 125 +- internal/creator/actions.go | 282 +- internal/creator/bailian.go | 22 +- internal/creator/bailian_test.go | 14 + internal/creator/collection.go | 97 +- internal/creator/content.go | 60 +- internal/creator/integration_test.go | 24 +- internal/creator/logic.go | 4 + internal/creator/metrics.go | 32 + .../migrations/023_password_references.sql | 8 + .../migrations/024_event_processing_times.sql | 2 + .../migrations/025_competitor_sync_tokens.sql | 2 + .../migrations/026_event_message_text.sql | 2 + internal/creator/models.go | 37 +- internal/creator/recovery_integration_test.go | 41 + internal/creator/rules.go | 13 +- internal/creator/scheduler.go | 4 +- internal/creator/scheduler_test.go | 19 +- internal/creator/settings.go | 66 +- internal/creator/store.go | 25 + internal/douyin/connector.go | 44 +- internal/douyin/connector_test.go | 40 +- internal/douyin/creator_collector.go | 32 +- requirements-gateway-dev.lock | 2 + requirements-gateway.lock | 1 + web/src/CreatorAccountsPage.jsx | 258 +- web/src/CreatorCompetitorsPage.jsx | 148 +- web/src/CreatorPages.test.jsx | 109 +- web/src/CreatorSettingsPage.jsx | 229 +- web/src/CreatorWorkbenchPage.jsx | 369 ++- 61 files changed, 10284 insertions(+), 6222 deletions(-) create mode 100644 cmd/__init__.py create mode 100644 cmd/control-plane/creator_events.go create mode 100644 cmd/control-plane/creator_events_test.go create mode 100644 cmd/control-plane/creator_material.go create mode 100644 cmd/control-plane/creator_material_test.go delete mode 100644 cmd/docker-gateway/douyin.go delete mode 100644 cmd/docker-gateway/douyin_cdp_integration_test.go delete mode 100644 cmd/docker-gateway/douyin_test.go delete mode 100644 cmd/docker-gateway/main.go delete mode 100644 cmd/docker-gateway/main_test.go delete mode 100644 cmd/docker-gateway/proxy.go delete mode 100644 cmd/docker-gateway/proxy_test.go create mode 100644 cmd/docker_gateway/__init__.py create mode 100644 cmd/docker_gateway/docker_client.py create mode 100644 cmd/docker_gateway/douyin.py create mode 100644 cmd/docker_gateway/gateway.py create mode 100644 cmd/docker_gateway/proxy.py create mode 100644 cmd/docker_gateway/test_gateway.py create mode 100644 docker/browser-wrapper/Dockerfile create mode 100644 docker/browser-wrapper/README.md create mode 100644 docker/browser-wrapper/docker-entrypoint.sh create mode 100644 docs/python-gateway-branch-review.md create mode 100644 internal/creator/migrations/023_password_references.sql create mode 100644 internal/creator/migrations/024_event_processing_times.sql create mode 100644 internal/creator/migrations/025_competitor_sync_tokens.sql create mode 100644 internal/creator/migrations/026_event_message_text.sql create mode 100644 internal/creator/recovery_integration_test.go create mode 100644 requirements-gateway-dev.lock create mode 100644 requirements-gateway.lock diff --git a/AGENTS.md b/AGENTS.md index b52a4dc..3e6b2ab 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -65,11 +65,11 @@ - 当前业务范围与验收以 [docs/plan01.md](docs/plan01.md) 为准:竞品分析、账号与大小号响应、环境与代理、评论线索和私信;先完成抖音完整流程,再完成小红书。 - 与旧探索规划或现有功能冲突时,以上述需求及使用者最新确认为准。自动响应按策略和 UID 冷却执行;人工发送逐次确认,两者不得混淆。自有账号互动与私信采用事件监听,不用轮询或 Mock 冒充实际能力。 -- 上述内容是目标范围,不代表已经实现。平台能力缺口、新的业务歧义必须向使用者确认;禁止自行删减需求、隐藏失败或以未要求的通用框架扩大实现范围。技术栈保持不变。 +- 上述内容是目标范围,不代表已经实现。平台能力缺口、新的业务歧义必须向使用者确认;禁止自行删减需求、隐藏失败或以未要求的通用框架扩大实现范围。控制面继续使用 Go;经使用者确认,Docker/浏览器 gateway 改用 Python。 ## 已批准的技术栈 -### Go 后端 +### Go 控制面 - 使用 Go 1.26、Fiber v3(HTTP 路由与服务生命周期)、Viper(配置)、Logrus(应用日志)、Cobra(可执行入口)。依赖版本由 `go.mod` 和 `go.sum` 精确锁定。 - 保留成熟的标准库集成,如反向代理和 Docker HTTP 客户端,不重复造轮子。`net/http` handler 跨越 Fiber 边界时,使用 Fiber 官方适配器。 @@ -78,6 +78,12 @@ - 每个服务只保留一个最小化的 Cobra 根命令。仅当存在真实的运维工作流需求时,才添加子命令、持久化 flag、代码生成器或补全。 - 未经评审的需求批准,不添加 ORM、Redis、任务框架或另一套 HTTP/配置/日志/CLI 技术栈。 +### Python Docker/浏览器 gateway + +- Docker/浏览器 gateway 使用 Python 3.12+;优先使用标准库 HTTP、Docker Engine Unix socket、socket/ssl/asyncio 与显式输入校验。只有真实浏览器会话需要时才使用已锁定的 Patchright/WebSocket 依赖。 +- gateway 继续是唯一挂载 Docker socket 的服务,只提供领域路由;Docker 生命周期、网络代际、代理、浏览器 CDP 和抖音页面内动作由 Python 实现,账号身份必须在每次写操作前核对。 +- Python 依赖必须写入锁定文件;不允许自动登录、任意 CDP、Cookie/验证码/密码回显或把不确定写结果转换为成功。 + ### React 前端 - 使用 React 19、Vite 8、Refine 与 shadcn/ui、Tailwind CSS。在 `web/package-lock.json` 中锁定精确的已安装版本。 @@ -93,6 +99,6 @@ - 开始任务前先定义完成标准。交付前依此验证,发现问题就修好再测,不把未完成的工作交回给使用者。只有确认完成,或遇到真正需要使用者介入的障碍时,才回报。 - 每个非平凡行为变更附带最小的回归测试,且该测试在无此变更时会失败。在信任与集成边界覆盖成功、校验、失败和兼容路径;单元测试覆盖率保证 65% 以上。 -- 后端变更必须通过 `go test ./...`、`go vet ./...`,并完成 `./cmd/control-plane` 和 `./cmd/docker-gateway` 双端构建;涉及并发、生命周期或共享状态的变更须运行 `go test -race ./...`。Docker 或 Compose 变更还须通过 `docker compose config --quiet`。 +- 控制面变更必须通过 `go test ./...`、`go vet ./...`,并构建 `./cmd/control-plane`;涉及并发、生命周期或共享状态的变更须运行 `go test -race ./...`。Python gateway 必须通过其非交互式单元测试与覆盖率检查,并完成 Compose 构建/健康检查。Docker 或 Compose 变更还须通过 `docker compose config --quiet`。 - 前端变更必须从 lockfile 安装、通过仓库的非交互式测试命令,并通过 `npm --prefix web run build`。主题、Layout、导航、资源动作或 data provider 的变更需要聚焦的交互覆盖,包括适用的错误与禁用状态。 - 除非 issue 明确批准契约变更,保持既有 API 与 Docker 生命周期行为不变。在 PR 中文档化任何状态码、载荷、配置、迁移、安全或重试方面的影响。 diff --git a/Dockerfile b/Dockerfile index c540fef..feeccc3 100644 --- a/Dockerfile +++ b/Dockerfile @@ -11,14 +11,18 @@ WORKDIR /src COPY go.mod go.sum ./ COPY cmd/ ./cmd/ COPY internal/ ./internal/ -RUN CGO_ENABLED=0 go build -buildvcs=false -trimpath -ldflags='-s -w' -o /out/control-plane ./cmd/control-plane \ - && CGO_ENABLED=0 go build -buildvcs=false -trimpath -ldflags='-s -w' -o /out/docker-gateway ./cmd/docker-gateway +RUN CGO_ENABLED=0 go build -buildvcs=false -trimpath -ldflags='-s -w' -o /out/control-plane ./cmd/control-plane -FROM alpine:3.22@sha256:14358309a308569c32bdc37e2e0e9694be33a9d99e68afb0f5ff33cc1f695dce -RUN addgroup -g 65532 app && adduser -D -u 65532 -G app app \ - && install -d -o app -g app -m 0700 /var/lib/creatorhub/credentials +FROM python:3.13-alpine@sha256:7415fbc3c9e4979cc717d92377ab2bc7b2b4a2af1ac03cc52b5f3f88efedaf3a +COPY requirements-gateway.lock /tmp/requirements-gateway.lock +RUN apk add --no-cache ffmpeg \ + && python -m pip install --no-cache-dir -r /tmp/requirements-gateway.lock \ + && addgroup -g 65532 app && adduser -D -u 65532 -G app app \ + && install -d -o app -g app -m 0700 /var/lib/creatorhub/credentials /var/lib/creatorhub/materials WORKDIR /app -COPY --from=go /out/control-plane /out/docker-gateway /app/ +COPY --from=go /out/control-plane /app/ +COPY cmd/__init__.py /app/cmd/__init__.py +COPY cmd/docker_gateway /app/cmd/docker_gateway COPY --from=web /src/web/dist /app/web USER 65532:65532 ENV WEB_DIR=/app/web diff --git a/cmd/__init__.py b/cmd/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/cmd/control-plane/creator.go b/cmd/control-plane/creator.go index d38a7f0..a924704 100644 --- a/cmd/control-plane/creator.go +++ b/cmd/control-plane/creator.go @@ -10,6 +10,7 @@ import ( "net/http" "net/url" "strconv" + "strings" "time" "git.ipao.vip/rogee/creator-hub/internal/creator" @@ -21,7 +22,11 @@ import ( ) func registerCreator(app *fiber.App, store *creator.Store, phaseAStore *phasea.Store, hubStore *hub.Store, credentials phasea.CredentialBridge) { - registerCreatorWithServices(app, store, phaseAStore, hubStore, credentials, nil, nil, nil) + var executor creator.ActionExecutor + if store != nil && phaseAStore != nil && hubStore != nil { + executor = creatorGatewayActionExecutor{store: store, phaseAStore: phaseAStore, hubStore: hubStore, credentials: credentials} + } + registerCreatorWithServices(app, store, phaseAStore, hubStore, credentials, executor, nil, nil) } func registerCreatorWithServices(app *fiber.App, store *creator.Store, phaseAStore *phasea.Store, hubStore *hub.Store, credentials phasea.CredentialBridge, executor creator.ActionExecutor, generator creator.TextGenerator, analyzer creator.ThemeAnalyzer) { @@ -289,17 +294,8 @@ func registerCreatorWithServices(app *fiber.App, store *creator.Store, phaseASto } return c.Status(status).JSON(item) }) - app.Post("/api/creator/works/:id/material/step", func(c fiber.Ctx) error { - var input struct { - Step string `json:"step"` - Status string `json:"status"` - Reference string `json:"reference"` - Reason string `json:"reason"` - } - if err := decodeCreator(c, &input); err != nil { - return creatorError(c, err) - } - item, err := store.SetMaterialStep(c.Context(), c.Params("id"), input.Step, input.Status, input.Reference, input.Reason) + app.Post("/api/creator/works/:id/material/process", func(c fiber.Ctx) error { + item, err := processCreatorMaterial(c.Context(), store, c.Params("id")) if err != nil { return creatorError(c, err) } @@ -632,6 +628,160 @@ type creatorGatewayBrowser struct { environment hub.EnvironmentContext } +type creatorGatewayActionExecutor struct { + store *creator.Store + phaseAStore *phasea.Store + hubStore *hub.Store + credentials phasea.CredentialBridge +} + +func (executor creatorGatewayActionExecutor) Execute(ctx context.Context, request creator.ActionRequest) (creator.ActionResult, error) { + if request.Platform != creator.PlatformDouyin || executor.store == nil || executor.phaseAStore == nil || executor.hubStore == nil { + return creator.ActionResult{}, creator.ErrUnavailable + } + profile, err := executor.store.GetAccountProfile(ctx, request.AccountID) + if err != nil { + return creator.ActionResult{}, err + } + account, err := executor.phaseAStore.GetAccount(ctx, request.AccountID) + if err != nil { + return creator.ActionResult{}, err + } + if profile.Platform != creator.PlatformDouyin || account.Platform != creator.PlatformDouyin || profile.PlatformAccountKey == "" || profile.PlatformAccountKey != account.PlatformAccountKey { + return creator.ActionResult{}, creator.ErrConflict + } + environment, err := executor.hubStore.GetEnvironmentContextForAccount(ctx, request.AccountID) + if err != nil { + return creator.ActionResult{}, err + } + gateway, err := executor.hubStore.GetGateway(ctx, environment.Gateway) + if err != nil { + return creator.ActionResult{}, err + } + browser := creatorGatewayBrowser{gateway: gateway, environment: environment} + uid, identityErr := browser.Identity(ctx, profile.PlatformAccountKey) + if identityErr != nil { + resolver, ok := executor.credentials.(phasea.CredentialResolver) + if !ok { + return creator.ActionResult{}, creator.ErrUnavailable + } + rawCredential, resolveErr := executor.phaseAStore.ResolveAccountCredential(ctx, request.AccountID, resolver) + if resolveErr != nil { + return creator.ActionResult{}, resolveErr + } + cookies, parseErr := douyin.ParseCredential(rawCredential) + if parseErr != nil { + return creator.ActionResult{}, fmt.Errorf("parse account credential: %w", parseErr) + } + if setErr := browser.SetCookies(ctx, cookies); setErr != nil { + return creator.ActionResult{}, setErr + } + uid, identityErr = browser.Identity(ctx, profile.PlatformAccountKey) + if identityErr != nil { + return creator.ActionResult{}, fmt.Errorf("verify account identity after credential injection: %w", identityErr) + } + } + if _, verifyErr := executor.store.RecordVerifiedLoginResult(ctx, request.AccountID, uid); verifyErr != nil { + return creator.ActionResult{}, fmt.Errorf("persist verified account identity: %w", verifyErr) + } + payload := gatewayGenerationPayload(environment) + payload["expected_uid"] = uid + payload["action"] = request.Action + payload["target_uid"] = request.TargetUID + // UI and persistence use opaque internal IDs; the platform gateway receives only + // the verified platform keys and the comment's owning work key. + if request.TargetCommentID != "" { + comment, targetErr := executor.store.GetComment(ctx, request.TargetCommentID) + if targetErr != nil { + return creator.ActionResult{}, targetErr + } + payload["target_comment_id"] = comment.CommentKey + if request.TargetWorkID == "" { + request.TargetWorkID = comment.WorkID + } + } + if request.TargetWorkID != "" { + work, targetErr := executor.store.GetWork(ctx, request.TargetWorkID) + if targetErr != nil { + return creator.ActionResult{}, targetErr + } + payload["target_work_id"] = work.WorkKey + } + payload["text"] = request.Text + payload["confirm"] = true + status, body, err := gatewayCall(ctx, gateway, http.MethodPost, "/v1/browsers/"+url.PathEscape(environment.Alias)+"/douyin/action", payload, 30*time.Second) + if err != nil { + return creator.ActionResult{}, err + } + if status != http.StatusOK { + return creator.ActionResult{State: "uncertain", Reason: fmt.Sprintf("gateway returned HTTP %d", status)}, nil + } + var response struct { + Status string `json:"status"` + Code string `json:"code"` + Action string `json:"action"` + Evidence any `json:"evidence"` + } + if err := json.Unmarshal(body, &response); err != nil || (response.Status != "succeeded" && response.Status != "failed" && response.Status != "unknown") { + return creator.ActionResult{}, errors.New("gateway returned an invalid action result") + } + evidence := map[string]string{"gateway_status": response.Status} + if response.Action != "" { + evidence["action"] = response.Action + } + if response.Code != "" { + evidence["code"] = response.Code + } + evidenceCount := flattenActionEvidence(evidence, "evidence", response.Evidence) + state := response.Status + reason := response.Code + if state == "unknown" || state == "succeeded" && evidenceCount == 0 || state == "failed" && strings.TrimSpace(response.Code) == "" { + state = "uncertain" + if reason == "" { + reason = "写后确认证据不足" + } + } + return creator.ActionResult{State: state, Evidence: evidence, Reason: reason}, nil +} + +func flattenActionEvidence(destination map[string]string, prefix string, value any) int { + switch typed := value.(type) { + case string: + if typed == "" { + return 0 + } + destination[prefix] = typed + return 1 + case map[string]any: + count := 0 + for key, item := range typed { + count += flattenActionEvidence(destination, prefix+"."+key, item) + } + return count + default: + return 0 + } +} + +func (browser creatorGatewayBrowser) Identity(ctx context.Context, expectedKey string) (string, error) { + payload := gatewayGenerationPayload(browser.environment) + payload["expected_account_key"] = expectedKey + status, body, err := gatewayCall(ctx, browser.gateway, http.MethodPost, "/v1/browsers/"+url.PathEscape(browser.environment.Alias)+"/douyin/identity", payload, 30*time.Second) + if err != nil { + return "", err + } + if status != http.StatusOK { + return "", fmt.Errorf("douyin identity verification rejected with HTTP %d: %s", status, string(body)) + } + var identity struct { + UID string `json:"uid"` + } + if err := json.Unmarshal(body, &identity); err != nil || identity.UID == "" { + return "", errors.New("douyin identity response omitted uid") + } + return identity.UID, nil +} + func (browser creatorGatewayBrowser) SetCookies(ctx context.Context, cookies []douyin.Cookie) error { payload := gatewayGenerationPayload(browser.environment) payload["cookies"] = cookies @@ -675,7 +825,7 @@ func syncCreatorCompetitorDue(ctx context.Context, store *creator.Store, phaseAS } func syncCreatorCompetitorWithClaim(ctx context.Context, store *creator.Store, phaseAStore *phasea.Store, hubStore *hub.Store, credentials phasea.CredentialBridge, competitorID, accountID string, force bool) (creator.CollectionReport, error) { - if store == nil || phaseAStore == nil || hubStore == nil || credentials == nil || accountID == "" { + if store == nil || phaseAStore == nil || hubStore == nil || accountID == "" { return creator.CollectionReport{}, creator.ErrUnavailable } competitor, err := store.GetCompetitor(ctx, competitorID) @@ -687,7 +837,7 @@ func syncCreatorCompetitorWithClaim(ctx context.Context, store *creator.Store, p return creator.CollectionReport{}, err } now := time.Now().UTC() - claimed, err := store.ClaimCompetitorSync(ctx, competitorID, force, now) + leaseToken, claimed, err := store.ClaimCompetitorSync(ctx, competitorID, force, now) if err != nil { return creator.CollectionReport{}, err } @@ -695,7 +845,7 @@ func syncCreatorCompetitorWithClaim(ctx context.Context, store *creator.Store, p return creator.CollectionReport{}, creator.ErrConflict } blocked := func(blockErr error) (creator.CollectionReport, error) { - markErr := store.MarkCompetitorSync(ctx, competitorID, "blocked", "", blockErr.Error(), nil) + markErr := store.MarkCompetitorSync(ctx, competitorID, leaseToken, "blocked", "", blockErr.Error(), nil) return creator.CollectionReport{}, errors.Join(blockErr, markErr) } if competitor.Platform != creator.PlatformDouyin { @@ -708,23 +858,12 @@ func syncCreatorCompetitorWithClaim(ctx context.Context, store *creator.Store, p if account.Platform != creator.PlatformDouyin || account.AuthorizationStatus != "authorized" { return blocked(creator.ErrConflict) } - if _, err := store.GetAccountProfile(ctx, accountID); err != nil { + profile, err := store.GetAccountProfile(ctx, accountID) + if err != nil { return blocked(err) } - resolver, ok := credentials.(phasea.CredentialResolver) - if !ok { - return blocked(creator.ErrUnavailable) - } - rawCredential, err := phaseAStore.ResolveAccountCredential(ctx, accountID, resolver) - if err != nil { - return blocked(fmt.Errorf("%w: resolve account credential: %v", creator.ErrUnavailable, err)) - } - cookies, err := douyin.ParseCookieHeader(rawCredential) - if err != nil { - cookies, err = douyin.ParseCookieBundle(rawCredential) - } - if err != nil { - return blocked(fmt.Errorf("%w: invalid account cookie bundle", creator.ErrConflict)) + if profile.BusinessStatus != "normal" || profile.LoginStatus != "logged_in" { + return blocked(creator.ErrConflict) } environment, err := hubStore.GetEnvironmentContextForAccount(ctx, accountID) if err != nil { @@ -738,16 +877,44 @@ func syncCreatorCompetitorWithClaim(ctx context.Context, store *creator.Store, p return blocked(fmt.Errorf("%w: gateway unavailable: %v", creator.ErrUnavailable, err)) } browser := creatorGatewayBrowser{gateway: gateway, environment: environment} - if err := browser.SetCookies(ctx, cookies); err != nil { - return blocked(fmt.Errorf("%w: set account cookies: %v", creator.ErrUnavailable, err)) + if _, identityErr := browser.Identity(ctx, account.PlatformAccountKey); identityErr != nil { + resolver, ok := credentials.(phasea.CredentialResolver) + if !ok { + return blocked(creator.ErrUnavailable) + } + rawCredential, resolveErr := phaseAStore.ResolveAccountCredential(ctx, accountID, resolver) + if resolveErr != nil { + return blocked(fmt.Errorf("%w: resolve account credential: %v", creator.ErrUnavailable, resolveErr)) + } + cookies, parseErr := douyin.ParseCredential(rawCredential) + if parseErr != nil { + return blocked(fmt.Errorf("%w: invalid account credential", creator.ErrConflict)) + } + if setErr := browser.SetCookies(ctx, cookies); setErr != nil { + return blocked(fmt.Errorf("%w: set account cookies: %v", creator.ErrUnavailable, setErr)) + } + if _, verifyErr := browser.Identity(ctx, account.PlatformAccountKey); verifyErr != nil { + return blocked(fmt.Errorf("%w: account identity verification failed: %v", creator.ErrConflict, verifyErr)) + } } collector := douyin.CreatorCollector{Browser: browser, AccountKey: competitor.PlatformAccountKey, SourceType: creator.SourceCompetitor, SourceID: competitor.ID} - if err := collector.VerifyIdentity(ctx, account.PlatformAccountKey); err != nil { + canonicalSecUID, err := collector.CanonicalSecUID(ctx, account.PlatformAccountKey) + if err != nil { return blocked(fmt.Errorf("%w: account identity verification failed: %v", creator.ErrConflict, err)) } - report, collectErr := store.CollectSource(ctx, competitor.Platform, creator.SourceCompetitor, competitor.ID, collector, now) + collector.AccountKey = canonicalSecUID + collectionNow := now + if competitor.NextSyncAt != nil && !competitor.NextSyncAt.After(now) { + collectionNow = competitor.NextSyncAt.UTC() + } + report, collectErr := store.CollectSource(ctx, competitor.Platform, creator.SourceCompetitor, competitor.ID, collector, collectionNow) if collectErr != nil { - status, next := "failed", now.Add(time.Duration(settings.NewWorkIntervalSeconds)*time.Second) + nextBase := now + if competitor.NextSyncAt != nil { + nextBase = competitor.NextSyncAt.UTC() + } + next := creator.NextFixedRun(nextBase, time.Now().UTC(), time.Duration(settings.NewWorkIntervalSeconds)*time.Second) + status := "failed" if errors.Is(collectErr, creator.ErrUnavailable) || errors.Is(collectErr, creator.ErrConflict) { status, next = "blocked", time.Time{} } @@ -755,11 +922,15 @@ func syncCreatorCompetitorWithClaim(ctx context.Context, store *creator.Store, p if !next.IsZero() { nextAt = &next } - markErr := store.MarkCompetitorSync(ctx, competitorID, status, "", collectErr.Error(), nextAt) + markErr := store.MarkCompetitorSync(ctx, competitorID, leaseToken, status, "", collectErr.Error(), nextAt) return report, errors.Join(collectErr, markErr) } - next := time.Now().UTC().Add(time.Duration(settings.NewWorkIntervalSeconds) * time.Second) - if err := store.MarkCompetitorSync(ctx, competitorID, "idle", "", "", &next); err != nil { + nextBase := now + if competitor.NextSyncAt != nil { + nextBase = competitor.NextSyncAt.UTC() + } + next := creator.NextFixedRun(nextBase, time.Now().UTC(), time.Duration(settings.NewWorkIntervalSeconds)*time.Second) + if err := store.MarkCompetitorSync(ctx, competitorID, leaseToken, "idle", "", "", &next); err != nil { return report, err } return report, nil @@ -781,7 +952,16 @@ func runCreatorScheduleOnce(ctx context.Context, store *creator.Store, phaseASto for _, competitor := range competitors { accountID, err := creatorCollectionAccount(ctx, store, phaseAStore, hubStore, competitor.Platform) if err != nil { - _ = store.MarkCompetitorSync(ctx, competitor.ID, "blocked", "", err.Error(), nil) + leaseToken, claimed, claimErr := store.ClaimCompetitorSync(ctx, competitor.ID, false, now) + if claimErr != nil { + logrus.WithError(claimErr).WithField("competitor_id", competitor.ID).Warn("creator competitor sync claim failed") + continue + } + if claimed { + if markErr := store.MarkCompetitorSync(ctx, competitor.ID, leaseToken, "blocked", "", err.Error(), nil); markErr != nil { + logrus.WithError(markErr).WithField("competitor_id", competitor.ID).Warn("creator competitor sync block update failed") + } + } logrus.WithError(err).WithField("competitor_id", competitor.ID).Warn("creator competitor sync blocked") continue } @@ -798,11 +978,113 @@ func runCreatorScheduleOnce(ctx context.Context, store *creator.Store, phaseASto logrus.WithError(err).WithField("account_id", accountID).Warn("creator owned scheduled sync failed") } } + if err := runCreatorMetricScheduleOnce(ctx, store, phaseAStore, hubStore, credentials, now); err != nil { + return err + } return nil } +func runCreatorMetricScheduleOnce(ctx context.Context, store *creator.Store, phaseAStore *phasea.Store, hubStore *hub.Store, credentials phasea.CredentialBridge, now time.Time) error { + settings, err := store.GetSettings(ctx) + if err != nil { + return err + } + works, err := store.ListDueMetricWorks(ctx, now) + if err != nil { + return err + } + for _, work := range works { + accountID := work.SourceID + if work.SourceType == creator.SourceCompetitor { + accountID, err = creatorCollectionAccount(ctx, store, phaseAStore, hubStore, work.Platform) + if err != nil { + logrus.WithError(err).WithField("work_id", work.ID).Warn("creator metric refresh account unavailable") + continue + } + } + if err := refreshCreatorMetricWork(ctx, store, phaseAStore, hubStore, credentials, work, accountID, settings, now); err != nil { + logrus.WithError(err).WithField("work_id", work.ID).Warn("creator metric refresh failed") + } + } + return nil +} + +func refreshCreatorMetricWork(ctx context.Context, store *creator.Store, phaseAStore *phasea.Store, hubStore *hub.Store, credentials phasea.CredentialBridge, work creator.Work, accountID string, settings creator.Settings, now time.Time) error { + account, err := phaseAStore.GetAccount(ctx, accountID) + if err != nil { + return err + } + profile, err := store.GetAccountProfile(ctx, accountID) + if err != nil { + return err + } + if account.Platform != creator.PlatformDouyin || account.AuthorizationStatus != "authorized" || profile.BusinessStatus != "normal" || profile.LoginStatus != "logged_in" { + return creator.ErrConflict + } + environment, err := hubStore.GetEnvironmentContextForAccount(ctx, accountID) + if err != nil { + return fmt.Errorf("%w: account environment unavailable: %v", creator.ErrUnavailable, err) + } + gateway, err := hubStore.GetGateway(ctx, environment.Gateway) + if err != nil { + return fmt.Errorf("%w: gateway unavailable: %v", creator.ErrUnavailable, err) + } + browser := creatorGatewayBrowser{gateway: gateway, environment: environment} + if _, identityErr := browser.Identity(ctx, account.PlatformAccountKey); identityErr != nil { + resolver, ok := credentials.(phasea.CredentialResolver) + if !ok { + return creator.ErrUnavailable + } + raw, resolveErr := phaseAStore.ResolveAccountCredential(ctx, accountID, resolver) + if resolveErr != nil { + return fmt.Errorf("%w: resolve account credential: %v", creator.ErrUnavailable, resolveErr) + } + cookies, parseErr := douyin.ParseCredential(raw) + if parseErr != nil { + return fmt.Errorf("%w: invalid account credential", creator.ErrConflict) + } + if setErr := browser.SetCookies(ctx, cookies); setErr != nil { + return fmt.Errorf("%w: set account cookies: %v", creator.ErrUnavailable, setErr) + } + if _, verifyErr := browser.Identity(ctx, account.PlatformAccountKey); verifyErr != nil { + return fmt.Errorf("%w: account identity verification failed: %v", creator.ErrConflict, verifyErr) + } + } + collector := douyin.CreatorCollector{Browser: browser, AccountKey: account.PlatformAccountKey, SourceType: work.SourceType, SourceID: work.SourceID} + canonical, err := collector.CanonicalSecUID(ctx, account.PlatformAccountKey) + if err != nil { + return fmt.Errorf("%w: account identity verification failed: %v", creator.ErrConflict, err) + } + collector.AccountKey = canonical + cursor := "" + for page := 0; page < 100; page++ { + result, pageErr := collector.ListWorks(ctx, canonical, cursor) + if pageErr != nil { + return pageErr + } + for _, item := range result.Items { + if item.WorkKey != work.WorkKey { + continue + } + if item.Likes == nil && item.CommentsCount == nil && item.Shares == nil { + return creator.ErrUnavailable + } + _, metricErr := store.RecordMetric(ctx, creator.MetricInput{WorkID: work.ID, CollectedAt: now, Likes: item.Likes, CommentsCount: item.CommentsCount, Shares: item.Shares}, settings, now) + return metricErr + } + if !result.HasMore { + return creator.ErrNotFound + } + if result.NextCursor == "" || result.NextCursor == cursor { + return creator.ErrInvalid + } + cursor = result.NextCursor + } + return creator.ErrInvalid +} + func syncCreatorOwned(ctx context.Context, store *creator.Store, phaseAStore *phasea.Store, hubStore *hub.Store, credentials phasea.CredentialBridge, accountID string, now time.Time) error { - if store == nil || phaseAStore == nil || hubStore == nil || credentials == nil || accountID == "" { + if store == nil || phaseAStore == nil || hubStore == nil || accountID == "" { return creator.ErrUnavailable } account, err := phaseAStore.GetAccount(ctx, accountID) @@ -812,23 +1094,19 @@ func syncCreatorOwned(ctx context.Context, store *creator.Store, phaseAStore *ph if account.Platform != creator.PlatformDouyin || account.AuthorizationStatus != "authorized" { return creator.ErrConflict } - if _, err := store.GetAccountProfile(ctx, accountID); err != nil { + profile, err := store.GetAccountProfile(ctx, accountID) + if err != nil { return err } - resolver, ok := credentials.(phasea.CredentialResolver) - if !ok { - return creator.ErrUnavailable + if profile.BusinessStatus != "normal" || profile.LoginStatus != "logged_in" { + return creator.ErrConflict } - rawCredential, err := phaseAStore.ResolveAccountCredential(ctx, accountID, resolver) + settings, err := store.GetSettings(ctx) if err != nil { - return fmt.Errorf("%w: resolve account credential: %v", creator.ErrUnavailable, err) + return err } - cookies, err := douyin.ParseCookieHeader(rawCredential) - if err != nil { - cookies, err = douyin.ParseCookieBundle(rawCredential) - } - if err != nil { - return fmt.Errorf("%w: invalid account cookie bundle", creator.ErrConflict) + blockOwned := func(blockErr error) error { + return errors.Join(blockErr, store.MarkCollectionBlocked(ctx, creator.SourceOwned, accountID, blockErr.Error(), now, settings.LookbackDays)) } environment, err := hubStore.GetEnvironmentContextForAccount(ctx, accountID) if err != nil { @@ -842,14 +1120,38 @@ func syncCreatorOwned(ctx context.Context, store *creator.Store, phaseAStore *ph return fmt.Errorf("%w: gateway unavailable: %v", creator.ErrUnavailable, err) } browser := creatorGatewayBrowser{gateway: gateway, environment: environment} - if err := browser.SetCookies(ctx, cookies); err != nil { - return fmt.Errorf("%w: set account cookies: %v", creator.ErrUnavailable, err) + if _, identityErr := browser.Identity(ctx, account.PlatformAccountKey); identityErr != nil { + resolver, ok := credentials.(phasea.CredentialResolver) + if !ok { + return blockOwned(creator.ErrUnavailable) + } + rawCredential, resolveErr := phaseAStore.ResolveAccountCredential(ctx, accountID, resolver) + if resolveErr != nil { + return blockOwned(fmt.Errorf("%w: resolve account credential: %v", creator.ErrUnavailable, resolveErr)) + } + cookies, parseErr := douyin.ParseCredential(rawCredential) + if parseErr != nil { + return blockOwned(fmt.Errorf("%w: invalid account credential", creator.ErrConflict)) + } + if setErr := browser.SetCookies(ctx, cookies); setErr != nil { + return blockOwned(fmt.Errorf("%w: set account cookies: %v", creator.ErrUnavailable, setErr)) + } + if _, verifyErr := browser.Identity(ctx, account.PlatformAccountKey); verifyErr != nil { + return blockOwned(fmt.Errorf("%w: account identity verification failed: %v", creator.ErrConflict, verifyErr)) + } } collector := douyin.CreatorCollector{Browser: browser, AccountKey: account.PlatformAccountKey, SourceType: creator.SourceOwned, SourceID: account.ID} - if err := collector.VerifyIdentity(ctx, account.PlatformAccountKey); err != nil { - return fmt.Errorf("%w: account identity verification failed: %v", creator.ErrConflict, err) + canonicalSecUID, err := collector.CanonicalSecUID(ctx, account.PlatformAccountKey) + if err != nil { + blockErr := store.MarkCollectionBlocked(ctx, creator.SourceOwned, account.ID, err.Error(), now, settings.LookbackDays) + return errors.Join(fmt.Errorf("%w: account identity verification failed: %v", creator.ErrConflict, err), blockErr) } - _, err = store.CollectSource(ctx, account.Platform, creator.SourceOwned, account.ID, collector, now) + collector.AccountKey = canonicalSecUID + _, collectionNow, windowErr := store.NextCollectionWindow(ctx, creator.SourceOwned, account.ID, now, time.Duration(settings.NewWorkIntervalSeconds)*time.Second, settings.LookbackDays) + if windowErr != nil { + return windowErr + } + _, err = store.CollectSource(ctx, account.Platform, creator.SourceOwned, account.ID, collector, collectionNow) return err } diff --git a/cmd/control-plane/creator_events.go b/cmd/control-plane/creator_events.go new file mode 100644 index 0000000..664de75 --- /dev/null +++ b/cmd/control-plane/creator_events.go @@ -0,0 +1,464 @@ +package main + +import ( + "context" + "encoding/json" + "fmt" + "net/http" + "net/url" + "strings" + "sync" + "time" + + "git.ipao.vip/rogee/creator-hub/internal/creator" + "git.ipao.vip/rogee/creator-hub/internal/hub" + "git.ipao.vip/rogee/creator-hub/internal/phasea" + "github.com/sirupsen/logrus" +) + +const creatorEventReconcileInterval = 10 * time.Second + +type creatorGatewayEvent struct { + DeliveryID string `json:"delivery_id,omitempty"` + Kind string `json:"kind"` + Reason string `json:"reason,omitempty"` + Continuity string `json:"continuity,omitempty"` + BoundaryAt string `json:"boundary_at,omitempty"` + Baseline bool `json:"baseline,omitempty"` + Notice *creatorGatewayEventNotice `json:"notice,omitempty"` +} + +type creatorGatewayEventNotice struct { + EventKey string `json:"event_key"` + EventType string `json:"event_type"` + InteractorUID string `json:"interactor_uid"` + CommentID string `json:"comment_id"` + WorkID string `json:"work_id"` + MessageText string `json:"message_text,omitempty"` + PlatformEventAt string `json:"platform_event_at,omitempty"` + GatewayReceivedAt string `json:"gateway_received_at,omitempty"` +} + +type creatorEventBinding struct { + accountID string + uid string + env hub.EnvironmentContext + gateway hub.Gateway +} + +func (binding creatorEventBinding) key() string { + return fmt.Sprintf("%s\x00%s\x00%s\x00%d\x00%s\x00%s\x00%s\x00%s", binding.gateway.Name, binding.gateway.Endpoint, binding.gateway.Token, binding.env.BindingVersion, binding.env.RuntimeID, binding.env.RuntimeNetworkID, binding.env.Exit.ID, binding.uid) +} + +type creatorEventListenerHandle struct { + cancel context.CancelFunc + done chan struct{} + key string +} + +type creatorEventListenerManager struct { + mu sync.Mutex + items map[string]creatorEventListenerHandle +} + +func runCreatorEventListeners(ctx context.Context, store *creator.Store, phaseAStore *phasea.Store, hubStore *hub.Store, executor creator.ActionExecutor, generator creator.TextGenerator) { + manager := &creatorEventListenerManager{items: map[string]creatorEventListenerHandle{}} + ticker := time.NewTicker(creatorEventReconcileInterval) + defer ticker.Stop() + defer manager.close() + for { + if err := manager.reconcile(ctx, store, phaseAStore, hubStore, executor, generator); err != nil && ctx.Err() == nil { + logrus.WithField("service", "control-plane").WithError(err).Warn("creator event listener reconciliation failed") + } + select { + case <-ctx.Done(): + return + case <-ticker.C: + } + } +} + +func (manager *creatorEventListenerManager) reconcile(ctx context.Context, store *creator.Store, phaseAStore *phasea.Store, hubStore *hub.Store, executor creator.ActionExecutor, generator creator.TextGenerator) error { + if store == nil || phaseAStore == nil || hubStore == nil { + return creator.ErrUnavailable + } + if recovered, err := store.RecoverStaleProcessing(ctx, time.Now().UTC()); err != nil { + return err + } else if recovered > 0 { + logrus.WithField("count", recovered).Warn("recovered stale creator operations as uncertain") + } + accounts, err := phaseAStore.ListAccounts(ctx) + if err != nil { + return err + } + desired := make(map[string]creatorEventBinding) + for _, account := range accounts { + if account.Platform != creator.PlatformDouyin || account.AuthorizationStatus != "authorized" { + continue + } + profile, profileErr := store.GetAccountProfile(ctx, account.ID) + if profileErr != nil { + logrus.WithError(profileErr).WithField("account_id", account.ID).Warn("creator event listener account profile unavailable") + continue + } + if profile.LoginStatus != "logged_in" { + continue + } + if !creatorEventUID(profile.PlatformAccountKey) { + logrus.WithField("account_id", account.ID).Warn("creator event listener account UID is unavailable") + continue + } + environment, environmentErr := hubStore.GetEnvironmentContextForAccount(ctx, account.ID) + if environmentErr != nil { + logrus.WithError(environmentErr).WithField("account_id", account.ID).Warn("creator event listener environment unavailable") + continue + } + if environment.RuntimeID == "" || environment.RuntimeNetworkID == "" || environment.BindingVersion <= 0 { + continue + } + gateway, gatewayErr := hubStore.GetGateway(ctx, environment.Gateway) + if gatewayErr != nil { + logrus.WithError(gatewayErr).WithField("account_id", account.ID).Warn("creator event listener gateway unavailable") + continue + } + desired[account.ID] = creatorEventBinding{accountID: account.ID, uid: profile.PlatformAccountKey, env: environment, gateway: gateway} + } + + var stopping []creatorEventListenerHandle + manager.mu.Lock() + for accountID, current := range manager.items { + binding, ok := desired[accountID] + if ok && current.key == binding.key() { + continue + } + current.cancel() + delete(manager.items, accountID) + stopping = append(stopping, current) + } + manager.mu.Unlock() + for _, current := range stopping { + <-current.done + } + + manager.mu.Lock() + defer manager.mu.Unlock() + for accountID, binding := range desired { + if _, ok := manager.items[accountID]; ok { + continue + } + listenerContext, cancel := context.WithCancel(ctx) + done := make(chan struct{}) + listenerBinding := binding + manager.items[accountID] = creatorEventListenerHandle{cancel: cancel, done: done, key: binding.key()} + go func() { + defer close(done) + runCreatorEventListener(listenerContext, store, listenerBinding, executor, generator) + }() + } + return nil +} + +func (manager *creatorEventListenerManager) close() { + manager.mu.Lock() + handles := make([]creatorEventListenerHandle, 0, len(manager.items)) + for accountID, current := range manager.items { + current.cancel() + handles = append(handles, current) + delete(manager.items, accountID) + } + manager.mu.Unlock() + for _, current := range handles { + <-current.done + } +} + +func runCreatorEventListener(ctx context.Context, store *creator.Store, binding creatorEventBinding, executor creator.ActionExecutor, generator creator.TextGenerator) { + path := "/v1/browsers/" + url.PathEscape(binding.env.Alias) + "/douyin/events" + generation := gatewayGenerationPayload(binding.env) + startPayload := make(map[string]any, len(generation)+1) + for key, value := range generation { + startPayload[key] = value + } + startPayload["expected_uid"] = binding.uid + defer stopCreatorEventListener(binding.accountID, binding.gateway, path, generation) + + backoff := time.Second + for ctx.Err() == nil { + status, _, err := gatewayCall(ctx, binding.gateway, http.MethodPost, path, startPayload, 30*time.Second) + if err != nil || status != http.StatusOK { + if err == nil { + err = fmt.Errorf("gateway returned HTTP %d", status) + } + logrus.WithError(err).WithField("account_id", binding.accountID).Warn("creator event listener start failed") + if !waitCreatorEventBackoff(ctx, backoff) { + return + } + backoff *= 2 + if backoff > 30*time.Second { + backoff = 30 * time.Second + } + continue + } + backoff = time.Second + // Every successful start creates a new boundary. Events observed before + // its marker remain historical even when the previous poll loop was ready. + ready := false + var boundaryAt time.Time + for ctx.Err() == nil { + status, body, err := gatewayCall(ctx, binding.gateway, http.MethodGet, path+"?limit=100&wait=25", generation, 35*time.Second) + if err != nil || status != http.StatusOK { + if err == nil { + err = fmt.Errorf("gateway returned HTTP %d", status) + } + logrus.WithError(err).WithField("account_id", binding.accountID).Warn("creator event listener poll failed") + break + } + var events []creatorGatewayEvent + if err := json.Unmarshal(body, &events); err != nil { + logrus.WithError(err).WithField("account_id", binding.accountID).Warn("creator event listener response is invalid") + break + } + for _, event := range events { + if event.Kind == "baseline" { + if event.BoundaryAt != "" { + parsedBoundary, boundaryErr := time.Parse(time.RFC3339Nano, event.BoundaryAt) + if boundaryErr != nil { + logrus.WithError(boundaryErr).WithField("account_id", binding.accountID).Warn("creator event baseline timestamp is invalid") + ready = false + boundaryAt = time.Time{} + } else { + boundaryAt = parsedBoundary.UTC() + ready = true + } + } else { + // A marker without a platform-verified boundary cannot prove + // continuity. Keep all timestamped notices non-actionable. + ready = false + boundaryAt = time.Time{} + } + } else if event.Kind == "open" || event.Kind == "error" || event.Kind == "close" || event.Kind == "reconnected" { + // A transport event never proves continuity. Only the explicit + // boundary marker permits automatic writes again. + ready = false + } + if event.Kind == "notice" { + if needsBaseline, reason := creatorGatewayEventNeedsBaseline(event, ready); needsBaseline { + event.Baseline = true + event.Reason = reason + } else if creatorEventBeforeBoundary(event, boundaryAt) { + event.Baseline = true + event.Reason = "平台事件早于监听边界" + } + } + handleCreatorGatewayEvent(ctx, store, binding, event, executor, generator) + } + } + if !waitCreatorEventBackoff(ctx, backoff) { + return + } + backoff *= 2 + if backoff > 30*time.Second { + backoff = 30 * time.Second + } + } +} + +func stopCreatorEventListener(accountID string, gateway hub.Gateway, path string, generation map[string]any) { + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + if status, _, err := gatewayCall(ctx, gateway, http.MethodDelete, path, generation, 5*time.Second); err != nil || status != http.StatusNoContent { + if err == nil { + err = fmt.Errorf("gateway returned HTTP %d", status) + } + logrus.WithError(err).WithField("account_id", accountID).Warn("creator event listener stop failed") + } +} + +func waitCreatorEventBackoff(ctx context.Context, delay time.Duration) bool { + timer := time.NewTimer(delay) + defer timer.Stop() + select { + case <-ctx.Done(): + return false + case <-timer.C: + return true + } +} + +func pathForCreatorEvent(environment hub.EnvironmentContext) string { + return "/v1/browsers/" + url.PathEscape(environment.Alias) + "/douyin/events" +} + +func handleCreatorGatewayEvent(ctx context.Context, store *creator.Store, binding creatorEventBinding, event creatorGatewayEvent, executor creator.ActionExecutor, generator creator.TextGenerator) { + ack := func() { + if event.DeliveryID == "" { + return + } + // Acknowledgement is deliberately detached from the listener poll loop: + // receipt and event classification must not be serialized behind a slow + // gateway request. The gateway keeps the delivery until this succeeds. + go func(deliveryID string) { + ackCtx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 5*time.Second) + defer cancel() + ackPath := pathForCreatorEvent(binding.env) + status, _, ackErr := gatewayCall(ackCtx, binding.gateway, http.MethodGet, ackPath+"?ack="+url.QueryEscape(deliveryID)+"&limit=1&wait=0", gatewayGenerationPayload(binding.env), 5*time.Second) + if ackErr != nil || status != http.StatusOK { + if ackErr == nil { + ackErr = fmt.Errorf("gateway returned HTTP %d", status) + } + logrus.WithError(ackErr).WithField("account_id", binding.accountID).Warn("creator event acknowledgement failed") + } + }(event.DeliveryID) + } + switch event.Kind { + case "error": + logrus.WithFields(logrus.Fields{"account_id": binding.accountID, "reason": event.Reason, "continuity": event.Continuity}).Warn("creator event listener reported an error") + ack() + return + case "reconnected": + logrus.WithField("account_id", binding.accountID).Warn("creator event listener reconnected") + ack() + return + case "open", "baseline": + ack() + return + case "close": + logrus.WithField("account_id", binding.accountID).Warn("creator event listener connection closed") + ack() + return + case "notice": + default: + logrus.WithFields(logrus.Fields{"account_id": binding.accountID, "kind": event.Kind}).Warn("creator event listener returned an unknown event") + ack() + return + } + if event.Notice == nil { + logrus.WithField("account_id", binding.accountID).Warn("creator event listener notice is missing") + ack() + return + } + input, err := creatorEventFromGatewayNotice(binding.accountID, *event.Notice) + if err != nil { + logrus.WithError(err).WithField("account_id", binding.accountID).Warn("creator event listener notice was rejected") + ack() + return + } + input.Baseline = event.Baseline + input.BaselineReason = event.Reason + if input.PlatformEventAt == nil { + input.Baseline = true + input.BaselineReason = "缺少平台事件时间" + } + if store == nil { + return + } + // Receipt is durable before the potentially slow action. This prevents an + // executor outage from erasing the platform notification and lets polling + // continue while a prior action is still in flight. + received, err := store.RecordEvent(ctx, input) + if err != nil { + logrus.WithError(err).WithFields(logrus.Fields{"account_id": binding.accountID, "event_key": input.EventKey}).Warn("creator event receipt failed") + return + } + if input.EventType == "dm" && input.InteractorUID != "" { + messageAt := input.PlatformEventAt + if messageAt == nil { + receivedAt := input.ReceivedAt + if receivedAt.IsZero() { + receivedAt = time.Now().UTC() + } + messageAt = &receivedAt + } + if _, _, messageErr := store.SaveMessage(ctx, creator.MessageInput{Platform: input.Platform, AccountID: input.ReceivingAccountID, PeerUID: input.InteractorUID, PlatformMessageKey: "event:" + input.EventKey, Direction: "inbound", MessageType: "text", Text: input.MessageText, SentState: "received", MessageAt: messageAt}); messageErr != nil { + logrus.WithError(messageErr).WithFields(logrus.Fields{"account_id": binding.accountID, "event_key": input.EventKey}).Warn("creator direct message persistence failed") + return + } + } + // The event is now durable. An action may be slow or unavailable, but that + // must not hold ingestion or cause the same receipt to be fetched forever. + ack() + if input.Baseline || received.Event.State != "received" { + return + } + go func() { + result, processErr := store.ProcessAutomaticEvent(ctx, input, executor, generator) + if processErr != nil { + logrus.WithError(processErr).WithFields(logrus.Fields{"account_id": binding.accountID, "event_key": input.EventKey}).Warn("creator event processing failed") + return + } + logrus.WithFields(logrus.Fields{"account_id": binding.accountID, "event_key": input.EventKey, "event_type": input.EventType, "state": result.Event.State}).Info("creator event processed") + }() +} + +func creatorEventBeforeBoundary(event creatorGatewayEvent, boundaryAt time.Time) bool { + if boundaryAt.IsZero() || event.Notice == nil { + return false + } + platformAt, err := time.Parse(time.RFC3339Nano, event.Notice.PlatformEventAt) + return err != nil || !platformAt.After(boundaryAt) +} + +func creatorGatewayEventNeedsBaseline(event creatorGatewayEvent, ready bool) (bool, string) { + if event.Baseline { + if event.Reason != "" { + return true, event.Reason + } + return true, "监听基线" + } + if !ready { + return true, "监听边界未确认" + } + if event.Notice == nil || strings.TrimSpace(event.Notice.PlatformEventAt) == "" { + return true, "缺少平台事件时间" + } + return false, "" +} + +func creatorEventFromGatewayNotice(accountID string, notice creatorGatewayEventNotice) (creator.InteractionEvent, error) { + if accountID == "" || !creatorEventID(notice.EventKey) || !creator.ValidEventType(notice.EventType) { + return creator.InteractionEvent{}, creator.ErrInvalid + } + if notice.InteractorUID != "" && !creatorEventUID(notice.InteractorUID) { + return creator.InteractionEvent{}, creator.ErrInvalid + } + if (notice.CommentID != "" && !creatorEventID(notice.CommentID)) || (notice.WorkID != "" && !creatorEventID(notice.WorkID)) { + return creator.InteractionEvent{}, creator.ErrInvalid + } + result := creator.InteractionEvent{Platform: creator.PlatformDouyin, ReceivingAccountID: accountID, EventKey: notice.EventKey, EventType: notice.EventType, InteractorUID: notice.InteractorUID, CommentID: notice.CommentID, WorkID: notice.WorkID, MessageText: strings.TrimSpace(notice.MessageText)} + if strings.TrimSpace(notice.PlatformEventAt) != "" { + at, err := time.Parse(time.RFC3339Nano, notice.PlatformEventAt) + if err != nil { + return creator.InteractionEvent{}, fmt.Errorf("invalid platform event time: %w", err) + } + at = at.UTC() + result.PlatformEventAt = &at + } + if strings.TrimSpace(notice.GatewayReceivedAt) != "" { + receivedAt, err := time.Parse(time.RFC3339Nano, notice.GatewayReceivedAt) + if err != nil { + return creator.InteractionEvent{}, fmt.Errorf("invalid gateway receipt time: %w", err) + } + result.ReceivedAt = receivedAt.UTC() + } + return result, nil +} + +func creatorEventUID(value string) bool { + return creatorEventDigits(value, 20) +} + +func creatorEventID(value string) bool { + return creatorEventDigits(value, 64) +} + +func creatorEventDigits(value string, max int) bool { + if value == "" || len(value) > max || value[0] < '1' || value[0] > '9' { + return false + } + for _, char := range value[1:] { + if char < '0' || char > '9' { + return false + } + } + return true +} diff --git a/cmd/control-plane/creator_events_test.go b/cmd/control-plane/creator_events_test.go new file mode 100644 index 0000000..9a9116a --- /dev/null +++ b/cmd/control-plane/creator_events_test.go @@ -0,0 +1,99 @@ +package main + +import ( + "testing" + "time" + + "git.ipao.vip/rogee/creator-hub/internal/hub" +) + +func TestGatewayGenerationPayloadIncludesCurrentProxyExit(t *testing.T) { + payload := gatewayGenerationPayload(hub.EnvironmentContext{ + BindingVersion: 3, + RuntimeID: "runtime", + RuntimeNetworkID: "network", + Exit: hub.NetworkExit{ID: "exit-current"}, + }) + if payload["binding_version"] != int64(3) || payload["runtime_id"] != "runtime" || payload["network_id"] != "network" || payload["network_exit_id"] != "exit-current" { + t.Fatalf("unexpected generation payload: %#v", payload) + } +} + +func TestCreatorGatewayEventNeedsBaseline(t *testing.T) { + if ok, reason := creatorGatewayEventNeedsBaseline(creatorGatewayEvent{Kind: "notice"}, true); !ok || reason != "缺少平台事件时间" { + t.Fatalf("missing platform time must be held at baseline: %v %q", ok, reason) + } + if ok, reason := creatorGatewayEventNeedsBaseline(creatorGatewayEvent{Kind: "notice", Notice: &creatorGatewayEventNotice{PlatformEventAt: "2024-01-01T00:00:00Z"}}, false); !ok || reason != "监听边界未确认" { + t.Fatalf("unconfirmed listener boundary must be held: %v %q", ok, reason) + } + if ok, reason := creatorGatewayEventNeedsBaseline(creatorGatewayEvent{Kind: "notice", Notice: &creatorGatewayEventNotice{PlatformEventAt: "2024-01-01T00:00:00Z"}}, true); ok || reason != "" { + t.Fatalf("confirmed event boundary should be actionable: %v %q", ok, reason) + } +} + +func TestCreatorEventBeforeBoundary(t *testing.T) { + boundary := time.Date(2024, 1, 1, 0, 0, 0, 0, time.UTC) + before := creatorGatewayEvent{Notice: &creatorGatewayEventNotice{PlatformEventAt: "2023-12-31T23:59:59Z"}} + after := creatorGatewayEvent{Notice: &creatorGatewayEventNotice{PlatformEventAt: "2024-01-01T00:00:01Z"}} + invalid := creatorGatewayEvent{Notice: &creatorGatewayEventNotice{PlatformEventAt: "not-a-time"}} + if !creatorEventBeforeBoundary(before, boundary) || !creatorEventBeforeBoundary(invalid, boundary) || creatorEventBeforeBoundary(after, boundary) { + t.Fatalf("unexpected boundary classification") + } +} + +func TestCreatorEventFromGatewayNotice(t *testing.T) { + event, err := creatorEventFromGatewayNotice("account-1", creatorGatewayEventNotice{ + EventKey: "9007199254740993", + EventType: "comment", + InteractorUID: "7654321", + CommentID: "987654", + WorkID: "123456", + PlatformEventAt: "2023-11-14T22:13:20+00:00", + GatewayReceivedAt: "2023-11-14T22:13:21+00:00", + }) + if err != nil { + t.Fatal(err) + } + if event.Platform != "douyin" || event.ReceivingAccountID != "account-1" || event.EventKey != "9007199254740993" || event.InteractorUID != "7654321" || event.CommentID != "987654" || event.WorkID != "123456" { + t.Fatalf("unexpected event: %+v", event) + } + if event.PlatformEventAt == nil || event.PlatformEventAt.UTC().Format("2006-01-02T15:04:05Z07:00") != "2023-11-14T22:13:20Z" { + t.Fatalf("unexpected event time: %+v", event.PlatformEventAt) + } + if event.ReceivedAt.UTC().Format("2006-01-02T15:04:05Z07:00") != "2023-11-14T22:13:21Z" { + t.Fatalf("unexpected gateway receipt time: %v", event.ReceivedAt) + } +} + +func TestCreatorEventFromGatewayNoticeRejectsInvalidIdentity(t *testing.T) { + _, err := creatorEventFromGatewayNotice("account-1", creatorGatewayEventNotice{ + EventKey: "1", + EventType: "follow", + InteractorUID: "not-a-uid", + }) + if err == nil { + t.Fatal("expected invalid interactor UID") + } +} + +func TestCreatorEventFromGatewayNoticeRejectsInvalidGatewayReceiptTime(t *testing.T) { + _, err := creatorEventFromGatewayNotice("account-1", creatorGatewayEventNotice{ + EventKey: "1", + EventType: "like", + GatewayReceivedAt: "not-a-time", + }) + if err == nil { + t.Fatal("expected invalid gateway receipt time") + } +} + +func TestCreatorEventFromGatewayNoticeRejectsInvalidTime(t *testing.T) { + _, err := creatorEventFromGatewayNotice("account-1", creatorGatewayEventNotice{ + EventKey: "1", + EventType: "like", + PlatformEventAt: "not-a-time", + }) + if err == nil { + t.Fatal("expected invalid platform event time") + } +} diff --git a/cmd/control-plane/creator_material.go b/cmd/control-plane/creator_material.go new file mode 100644 index 0000000..9b4df92 --- /dev/null +++ b/cmd/control-plane/creator_material.go @@ -0,0 +1,210 @@ +package main + +import ( + "context" + "fmt" + "io" + "net/http" + "os" + "os/exec" + "path/filepath" + "strings" + "time" + + "git.ipao.vip/rogee/creator-hub/internal/creator" +) + +const maxCreatorMaterialBytes int64 = 512 << 20 + +func processCreatorMaterial(ctx context.Context, store *creator.Store, workID string) (creator.MaterialJob, error) { + if store == nil || workID == "" || filepath.Base(workID) != workID { + return creator.MaterialJob{}, creator.ErrInvalid + } + work, err := store.GetWork(ctx, workID) + if err != nil { + return creator.MaterialJob{}, err + } + job, err := store.GetMaterial(ctx, workID) + if err != nil { + return creator.MaterialJob{}, err + } + if !job.Selected { + return creator.MaterialJob{}, creator.ErrConflict + } + root := os.Getenv("CREATOR_MEDIA_DIR") + if root == "" { + root = "/var/lib/creatorhub/materials" + } + root = filepath.Clean(root) + dir := filepath.Join(root, workID) + if err := os.MkdirAll(dir, 0o700); err != nil { + return creator.MaterialJob{}, fmt.Errorf("create material directory: %w", err) + } + videoPath := filepath.Join(dir, "source") + videoReference := filepath.Join(workID, "source") + if job.DownloadStatus != "succeeded" || !fileExists(videoPath) { + if _, err := store.SetMaterialStep(ctx, workID, "download", "running", "", ""); err != nil { + return creator.MaterialJob{}, err + } + if err := downloadCreatorMaterial(ctx, work.OriginalURL, videoPath); err != nil { + return setMaterialFailure(ctx, store, workID, "download", err) + } + job, err = store.SetMaterialStep(ctx, workID, "download", "succeeded", videoReference, "") + if err != nil { + return creator.MaterialJob{}, err + } + } + + audioPath := filepath.Join(dir, "audio.wav") + audioReference := filepath.Join(workID, "audio.wav") + if job.AudioStatus != "succeeded" && job.AudioStatus != "no_audio" { + if _, err := store.SetMaterialStep(ctx, workID, "audio", "running", "", ""); err != nil { + return creator.MaterialJob{}, err + } + hasAudio, err := creatorMaterialHasAudio(ctx, videoPath) + if err != nil { + return setMaterialFailure(ctx, store, workID, "audio", err) + } + if !hasAudio { + job, err = store.SetMaterialStep(ctx, workID, "audio", "no_audio", "", "视频没有音轨") + if err != nil { + return creator.MaterialJob{}, err + } + } else if err := extractCreatorAudio(ctx, videoPath, audioPath); err != nil { + return setMaterialFailure(ctx, store, workID, "audio", err) + } else { + job, err = store.SetMaterialStep(ctx, workID, "audio", "succeeded", audioReference, "") + if err != nil { + return creator.MaterialJob{}, err + } + } + } + + if job.TranscriptionStatus != "succeeded" && job.TranscriptionStatus != "no_speech" { + if job.AudioStatus == "no_audio" { + job, err = store.SetMaterialStep(ctx, workID, "transcription", "no_speech", "", "没有可转写的音轨") + } else { + if _, err = store.SetMaterialStep(ctx, workID, "transcription", "running", "", ""); err == nil { + transcript, transcribeErr := transcribeCreatorAudio(ctx, audioPath) + if transcribeErr != nil { + job, err = setMaterialFailure(ctx, store, workID, "transcription", transcribeErr) + } else if strings.TrimSpace(transcript) == "" { + job, err = store.SetMaterialStep(ctx, workID, "transcription", "no_speech", "", "转写未检测到语音") + } else { + transcriptPath := filepath.Join(dir, "transcript.txt") + if writeErr := os.WriteFile(transcriptPath, []byte(transcript), 0o600); writeErr != nil { + job, err = setMaterialFailure(ctx, store, workID, "transcription", writeErr) + } else { + job, err = store.SetMaterialStep(ctx, workID, "transcription", "succeeded", filepath.Join(workID, "transcript.txt"), "") + } + } + } + } + if err != nil { + return creator.MaterialJob{}, err + } + } + return job, nil +} + +func downloadCreatorMaterial(ctx context.Context, rawURL, destination string) error { + parsed, err := http.NewRequestWithContext(ctx, http.MethodGet, rawURL, nil) + if err != nil || (parsed.URL.Scheme != "http" && parsed.URL.Scheme != "https") || parsed.URL.Host == "" { + return creator.ErrInvalid + } + client := &http.Client{Timeout: 2 * time.Minute} + response, err := client.Do(parsed) + if err != nil { + return fmt.Errorf("download material: %w", err) + } + defer response.Body.Close() + if response.StatusCode < 200 || response.StatusCode >= 300 { + return fmt.Errorf("download material: HTTP %s", response.Status) + } + if response.ContentLength > maxCreatorMaterialBytes { + return fmt.Errorf("download material exceeds size limit") + } + temporary := destination + ".tmp" + defer os.Remove(temporary) + file, err := os.OpenFile(temporary, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0o600) + if err != nil { + return fmt.Errorf("create material file: %w", err) + } + written, copyErr := io.Copy(file, io.LimitReader(response.Body, maxCreatorMaterialBytes+1)) + closeErr := file.Close() + if copyErr != nil { + return fmt.Errorf("write material: %w", copyErr) + } + if closeErr != nil { + return fmt.Errorf("close material file: %w", closeErr) + } + if written > maxCreatorMaterialBytes { + return fmt.Errorf("download material exceeds size limit") + } + if err := os.Rename(temporary, destination); err != nil { + return fmt.Errorf("publish material: %w", err) + } + return nil +} + +func creatorMaterialHasAudio(ctx context.Context, videoPath string) (bool, error) { + if _, err := exec.LookPath("ffprobe"); err != nil { + return false, fmt.Errorf("ffprobe is unavailable: %w", err) + } + command := exec.CommandContext(ctx, "ffprobe", "-v", "error", "-select_streams", "a:0", "-show_entries", "stream=index", "-of", "csv=p=0", videoPath) + output, err := command.Output() + if err != nil { + return false, fmt.Errorf("inspect audio stream: %w", err) + } + return strings.TrimSpace(string(output)) != "", nil +} + +func extractCreatorAudio(ctx context.Context, videoPath, audioPath string) error { + if _, err := exec.LookPath("ffmpeg"); err != nil { + return fmt.Errorf("ffmpeg is unavailable: %w", err) + } + temporary := audioPath + ".tmp" + defer os.Remove(temporary) + command := exec.CommandContext(ctx, "ffmpeg", "-nostdin", "-v", "error", "-y", "-i", videoPath, "-vn", "-ac", "1", "-ar", "16000", temporary) + if output, err := command.CombinedOutput(); err != nil { + return fmt.Errorf("extract audio: %w: %s", err, strings.TrimSpace(string(output))) + } + if !fileExists(temporary) { + return fmt.Errorf("extract audio produced no file") + } + if err := os.Rename(temporary, audioPath); err != nil { + return fmt.Errorf("publish audio: %w", err) + } + return nil +} + +func transcribeCreatorAudio(ctx context.Context, audioPath string) (string, error) { + binary := os.Getenv("CREATOR_TRANSCRIPTION_BIN") + if binary == "" { + return "", fmt.Errorf("transcription provider is not configured") + } + if _, err := exec.LookPath(binary); err != nil { + return "", fmt.Errorf("transcription provider is unavailable: %w", err) + } + output, err := exec.CommandContext(ctx, binary, audioPath).Output() + if err != nil { + return "", fmt.Errorf("transcribe audio: %w", err) + } + if len(output) > 1<<20 { + return "", fmt.Errorf("transcript exceeds size limit") + } + return string(output), nil +} + +func setMaterialFailure(ctx context.Context, store *creator.Store, workID, step string, cause error) (creator.MaterialJob, error) { + job, err := store.SetMaterialStep(ctx, workID, step, "failed", "", cause.Error()) + if err != nil { + return creator.MaterialJob{}, fmt.Errorf("record %s failure: %w", step, err) + } + return job, nil +} + +func fileExists(path string) bool { + info, err := os.Stat(path) + return err == nil && info.Mode().IsRegular() && info.Size() > 0 +} diff --git a/cmd/control-plane/creator_material_test.go b/cmd/control-plane/creator_material_test.go new file mode 100644 index 0000000..512a6aa --- /dev/null +++ b/cmd/control-plane/creator_material_test.go @@ -0,0 +1,53 @@ +package main + +import ( + "context" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "testing" +) + +func TestDownloadCreatorMaterialPublishesOnlyCompleteFiles(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path != "/source" { + t.Fatalf("unexpected path: %s", r.URL.Path) + } + _, _ = w.Write([]byte("media")) + })) + defer server.Close() + + destination := filepath.Join(t.TempDir(), "source") + if err := downloadCreatorMaterial(context.Background(), server.URL+"/source", destination); err != nil { + t.Fatalf("download material: %v", err) + } + content, err := os.ReadFile(destination) + if err != nil { + t.Fatalf("read material: %v", err) + } + if string(content) != "media" { + t.Fatalf("unexpected material: %q", content) + } + if _, err := os.Stat(destination + ".tmp"); !os.IsNotExist(err) { + t.Fatalf("temporary file remains: %v", err) + } +} + +func TestDownloadCreatorMaterialRejectsInvalidResponses(t *testing.T) { + server := httptest.NewServer(http.NotFoundHandler()) + defer server.Close() + if err := downloadCreatorMaterial(context.Background(), server.URL, filepath.Join(t.TempDir(), "source")); err == nil { + t.Fatal("expected HTTP failure") + } + if err := downloadCreatorMaterial(context.Background(), "file:///tmp/source", filepath.Join(t.TempDir(), "source")); err == nil { + t.Fatal("expected non-HTTP URL failure") + } +} + +func TestTranscribeCreatorAudioRequiresExplicitProvider(t *testing.T) { + t.Setenv("CREATOR_TRANSCRIPTION_BIN", "") + if _, err := transcribeCreatorAudio(context.Background(), "/tmp/audio.wav"); err == nil { + t.Fatal("expected missing provider error") + } +} diff --git a/cmd/control-plane/hub.go b/cmd/control-plane/hub.go index da18118..ed78e42 100644 --- a/cmd/control-plane/hub.go +++ b/cmd/control-plane/hub.go @@ -138,11 +138,15 @@ func gatewayProxyPayload(environment hub.EnvironmentContext, runtimeID, networkI } func gatewayGenerationPayload(environment hub.EnvironmentContext) map[string]any { - if environment.RuntimeCleanupBindingVersion > 0 { - return map[string]any{"binding_version": environment.RuntimeCleanupBindingVersion, "runtime_id": environment.RuntimeCleanupRuntimeID, - "network_id": environment.RuntimeCleanupNetworkID} + return map[string]any{"binding_version": environment.BindingVersion, "runtime_id": environment.RuntimeID, "network_id": environment.RuntimeNetworkID, "network_exit_id": environment.Exit.ID} +} + +func gatewayCleanupGenerationPayload(environment hub.EnvironmentContext) map[string]any { + bindingVersion, runtimeID, networkID := environment.RuntimeCleanupBindingVersion, environment.RuntimeCleanupRuntimeID, environment.RuntimeCleanupNetworkID + if bindingVersion == 0 { + bindingVersion, runtimeID, networkID = environment.BindingVersion, environment.RuntimeID, environment.RuntimeNetworkID } - return map[string]any{"binding_version": environment.BindingVersion, "runtime_id": environment.RuntimeID, "network_id": environment.RuntimeNetworkID} + return map[string]any{"binding_version": bindingVersion, "runtime_id": runtimeID, "network_id": networkID} } func runtimeCleanupGeneration(environment hub.EnvironmentContext, bindingVersion int64, runtimeID string, networkIDs ...string) hub.EnvironmentContext { @@ -1404,7 +1408,7 @@ func stopEnvironmentRuntime(ctx context.Context, store runtimeStopStore, environ return err } status, body, callErr := gatewayCall(ctx, gateway, http.MethodPost, "/v1/browsers/"+environment.Alias+"/stop", - gatewayGenerationPayload(environment), 30*time.Second) + gatewayCleanupGenerationPayload(environment), 30*time.Second) if callErr != nil || status >= http.StatusInternalServerError { container, found, reconcileErr := reconcileGatewayContainer(ctx, gateway, environment.Alias) if reconcileErr != nil { @@ -1787,7 +1791,7 @@ func removeGatewayRuntime(ctx context.Context, store runtimeCleanupStore, target } for attempt := 0; attempt < 2; attempt++ { status, body, callErr := gatewayCall(ctx, target, http.MethodDelete, "/v1/browsers/"+environment.Alias, - gatewayGenerationPayload(environment), 30*time.Second) + gatewayCleanupGenerationPayload(environment), 30*time.Second) if callErr == nil && (status == http.StatusNoContent || status == http.StatusNotFound) { return true, store.SetRuntimeCleanupPending(ctx, environment, false) } diff --git a/cmd/control-plane/main.go b/cmd/control-plane/main.go index 135fb57..be42ce7 100644 --- a/cmd/control-plane/main.go +++ b/cmd/control-plane/main.go @@ -30,6 +30,7 @@ import ( type config struct { listenAddr, webDir, databaseURL, credentialStoreDir string username, password string + aiAPIKey, aiBaseURL string credentialMasterKey []byte logLevel logrus.Level } @@ -97,13 +98,23 @@ func newCommand() *cobra.Command { defer close(creatorScheduleDone) runCreatorScheduler(creatorScheduleContext, creatorStore, phaseAStore, hubStore, credentials) }() - listenErr := newHandlerWithCreator(cfg.webDir, cfg.username, cfg.password, phaseAStore, hubStore, credentials, creatorStore).Listen(cfg.listenAddr, fiber.ListenConfig{ + creatorEventContext, stopCreatorEvents := context.WithCancel(command.Context()) + creatorEventDone := make(chan struct{}) + creatorEventExecutor := creatorGatewayActionExecutor{store: creatorStore, phaseAStore: phaseAStore, hubStore: hubStore, credentials: credentials} + creatorAI := &creator.ConfiguredBailian{Store: creatorStore, APIKey: cfg.aiAPIKey, BaseURL: cfg.aiBaseURL} + go func() { + defer close(creatorEventDone) + runCreatorEventListeners(creatorEventContext, creatorStore, phaseAStore, hubStore, creatorEventExecutor, creatorAI) + }() + listenErr := newHandlerWithCreatorAndAI(cfg.webDir, cfg.username, cfg.password, phaseAStore, hubStore, credentials, creatorStore, creatorAI, creatorAI).Listen(cfg.listenAddr, fiber.ListenConfig{ GracefulContext: command.Context(), DisableStartupMessage: true, }) stopCreatorScheduler() + stopCreatorEvents() stopHeartbeat() <-creatorScheduleDone + <-creatorEventDone <-heartbeatDone return listenErr }, @@ -163,6 +174,8 @@ func loadConfig() (config, error) { _ = v.BindEnv("log_level", "LOG_LEVEL") _ = v.BindEnv("username", "CONTROL_PLANE_USERNAME") _ = v.BindEnv("password", "CONTROL_PLANE_PASSWORD") + _ = v.BindEnv("ai_api_key", "BAILIAN_API_KEY") + _ = v.BindEnv("ai_base_url", "BAILIAN_BASE_URL") level, err := logrus.ParseLevel(v.GetString("log_level")) if err != nil { @@ -175,6 +188,8 @@ func loadConfig() (config, error) { credentialStoreDir: strings.TrimSpace(v.GetString("credential_store_dir")), username: strings.TrimSpace(v.GetString("username")), password: v.GetString("password"), + aiAPIKey: strings.TrimSpace(v.GetString("ai_api_key")), + aiBaseURL: strings.TrimSpace(v.GetString("ai_base_url")), logLevel: level, } if cfg.listenAddr == "" { @@ -204,6 +219,12 @@ func loadConfig() (config, error) { (databaseURL.Scheme != "postgres" && databaseURL.Scheme != "postgresql") { return config{}, errors.New("DATABASE_URL must be a postgres URL with a host") } + if cfg.aiBaseURL != "" { + aiURL, parseErr := url.Parse(cfg.aiBaseURL) + if parseErr != nil || aiURL.Host == "" || (aiURL.Scheme != "http" && aiURL.Scheme != "https") || aiURL.User != nil { + return config{}, errors.New("BAILIAN_BASE_URL must be an HTTP(S) URL without credentials") + } + } return cfg, nil } @@ -250,6 +271,10 @@ func newHandlerWithCredentialBridge(webDirectory, username, password string, pha } func newHandlerWithCreator(webDirectory, username, password string, phaseAStore *phasea.Store, hubStore *hub.Store, credentials phasea.CredentialBridge, creatorStore *creator.Store) *fiber.App { + return newHandlerWithCreatorAndAI(webDirectory, username, password, phaseAStore, hubStore, credentials, creatorStore, nil, nil) +} + +func newHandlerWithCreatorAndAI(webDirectory, username, password string, phaseAStore *phasea.Store, hubStore *hub.Store, credentials phasea.CredentialBridge, creatorStore *creator.Store, generator creator.TextGenerator, analyzer creator.ThemeAnalyzer) *fiber.App { app := fiber.New(fiber.Config{ AppName: "CreatorHub control plane", BodyLimit: 1 << 20, @@ -269,7 +294,15 @@ func newHandlerWithCreator(webDirectory, username, password string, phaseAStore registerPhaseA(app, phaseAStore, hubStore, credentials) } if creatorStore != nil { - registerCreator(app, creatorStore, phaseAStore, hubStore, credentials) + if generator == nil && analyzer == nil { + registerCreator(app, creatorStore, phaseAStore, hubStore, credentials) + } else { + var executor creator.ActionExecutor + if phaseAStore != nil && hubStore != nil { + executor = creatorGatewayActionExecutor{store: creatorStore, phaseAStore: phaseAStore, hubStore: hubStore, credentials: credentials} + } + registerCreatorWithServices(app, creatorStore, phaseAStore, hubStore, credentials, executor, generator, analyzer) + } } app.Get("/*", spaHandler(webDirectory)) return app diff --git a/cmd/docker-gateway/douyin.go b/cmd/docker-gateway/douyin.go deleted file mode 100644 index 78e6a69..0000000 --- a/cmd/docker-gateway/douyin.go +++ /dev/null @@ -1,509 +0,0 @@ -package main - -import ( - "bytes" - "context" - "encoding/json" - "errors" - "io" - "net" - "net/http" - "net/url" - "regexp" - "strconv" - "strings" - "time" - - "git.ipao.vip/rogee/creator-hub/internal/douyin" - "github.com/gofiber/fiber/v3" - "github.com/sirupsen/logrus" - "golang.org/x/net/websocket" -) - -const ( - douyinOrigin = "https://www.douyin.com" - douyinOriginURL = "https://www.douyin.com/" - douyinIdentityPath = "/aweme/v1/web/user/profile/self/" - douyinWorksPath = "/aweme/v1/web/aweme/post/" - douyinCommentsPath = "/aweme/v1/web/comment/list/" - douyinResponseLimit = 1 << 20 - browserControlTimeout = 15 * time.Second -) - -var douyinAccountKeyPattern = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9._:@-]{0,127}$`) - -type restrictedBrowserResponse struct { - Status int - Body string - Challenge douyin.Challenge -} - -type restrictedBrowser interface { - SetCookies(context.Context, string, []douyin.Cookie) error - Get(context.Context, string, string) (restrictedBrowserResponse, error) -} - -type douyinGenerationRequest struct { - BindingVersion int64 `json:"binding_version"` - RuntimeID string `json:"runtime_id"` - NetworkID string `json:"network_id"` - NetworkExitID string `json:"network_exit_id"` -} - -type douyinCookieRequest struct { - douyinGenerationRequest - Cookies []douyin.Cookie `json:"cookies"` -} - -type douyinGetRequest struct { - douyinGenerationRequest - URL string `json:"url"` -} - -func (api gateway) setDouyinCookies(c fiber.Ctx) error { - var input douyinCookieRequest - if !runtimeIDPattern.MatchString(c.Params("id")) || decodeRestrictedBrowserRequest(c.Body(), &input) != nil || !validDouyinGeneration(input.douyinGenerationRequest) || - !validDouyinCookies(input.Cookies) { - return writeError(c, http.StatusBadRequest, errors.New("invalid restricted browser request")) - } - _, release, err := api.locks.acquire(c.Params("id")) - if err != nil { - return writeError(c, statusFor(err), err) - } - defer release() - if err := api.requireDouyinGeneration(c.Params("id"), input.douyinGenerationRequest); err != nil { - return writeError(c, statusFor(err), err) - } - if api.browser == nil { - return writeError(c, http.StatusBadGateway, errors.New("restricted browser operation failed")) - } - if err := api.browser.SetCookies(c.Context(), c.Params("id"), input.Cookies); err != nil { - logrus.WithError(err).WithField("browser_id", c.Params("id")).Error("restricted browser cookie operation failed") - return writeError(c, http.StatusBadGateway, errors.New("restricted browser operation failed")) - } - if err := api.requireDouyinGeneration(c.Params("id"), input.douyinGenerationRequest); err != nil { - return writeError(c, statusFor(err), err) - } - return c.SendStatus(http.StatusNoContent) -} - -func (api gateway) getDouyin(c fiber.Ctx) error { - var input douyinGetRequest - if !runtimeIDPattern.MatchString(c.Params("id")) || decodeRestrictedBrowserRequest(c.Body(), &input) != nil || !validDouyinGeneration(input.douyinGenerationRequest) || - !validDouyinURL(input.URL) { - return writeError(c, http.StatusBadRequest, errors.New("invalid restricted browser request")) - } - _, release, err := api.locks.acquire(c.Params("id")) - if err != nil { - return writeError(c, statusFor(err), err) - } - defer release() - if err := api.requireDouyinGeneration(c.Params("id"), input.douyinGenerationRequest); err != nil { - return writeError(c, statusFor(err), err) - } - if api.browser == nil { - return writeError(c, http.StatusBadGateway, errors.New("restricted browser operation failed")) - } - response, err := api.browser.Get(c.Context(), c.Params("id"), input.URL) - if err != nil { - logrus.WithError(err).WithField("browser_id", c.Params("id")).Error("restricted browser fetch operation failed") - return writeError(c, http.StatusBadGateway, errors.New("restricted browser operation failed")) - } - if err := api.requireDouyinGeneration(c.Params("id"), input.douyinGenerationRequest); err != nil { - return writeError(c, statusFor(err), err) - } - return writeJSON(c, http.StatusOK, map[string]any{ - "status": response.Status, "body": response.Body, "challenge": response.Challenge, - }) -} - -func decodeRestrictedBrowserRequest(body []byte, target any) error { - decoder := json.NewDecoder(bytes.NewReader(body)) - decoder.DisallowUnknownFields() - if err := decoder.Decode(target); err != nil { - return err - } - var trailing json.RawMessage - if err := decoder.Decode(&trailing); !errors.Is(err, io.EOF) { - return errors.New("request must contain one JSON value") - } - return nil -} - -func validDouyinGeneration(input douyinGenerationRequest) bool { - return input.BindingVersion > 0 && exitIDPattern.MatchString(input.RuntimeID) && - exitIDPattern.MatchString(input.NetworkID) && (input.NetworkExitID == "" || exitIDPattern.MatchString(input.NetworkExitID)) -} - -func (api gateway) requireDouyinGeneration(alias string, input douyinGenerationRequest) error { - runtimeID, labels, networks, err := api.managedContainerState(alias) - if err != nil { - return err - } - version, _ := strconv.ParseInt(labels[bindingVersionLabel], 10, 64) - if runtimeID != input.RuntimeID || version != input.BindingVersion || labels[networkIDLabel] != input.NetworkID || - labels[networkExitLabel] != input.NetworkExitID { - return errGenerationConflict - } - generation, _, exists, err := api.docker.inspectTenantNetwork(api.network, alias, input.BindingVersion, - input.RuntimeID, api.self, input.NetworkID, false) - if err != nil { - return err - } - if !exists || !generation.RuntimeAttached || generation.SelfMember == "" || len(generation.GatewayMembers) == 0 || - len(networks) != 1 || networks[generation.Name] != generation.ID { - return errGenerationConflict - } - return nil -} - -func validDouyinCookies(cookies []douyin.Cookie) bool { - if len(cookies) == 0 || len(cookies) > 64 { - return false - } - for _, cookie := range cookies { - domain := strings.ToLower(strings.TrimSpace(cookie.Domain)) - if cookie.Name == "" || len(cookie.Name) > 256 || len(cookie.Value) > 4096 || len(cookie.Domain) > 256 || len(cookie.Path) > 256 || - domain != cookie.Domain || cookie.Expires < 0 || strings.ContainsAny(cookie.Name, ";\r\n\x00") || strings.ContainsAny(cookie.Value, ";\r\n\x00") || - (domain != "douyin.com" && !strings.HasSuffix(domain, ".douyin.com")) || !strings.HasPrefix(cookie.Path, "/") || - strings.ContainsAny(cookie.Path, ";\r\n\x00") || (cookie.SameSite != "" && cookie.SameSite != "Lax" && - cookie.SameSite != "Strict" && cookie.SameSite != "None") { - return false - } - } - return true -} - -func validDouyinURL(raw string) bool { - parsed, err := url.Parse(raw) - if err != nil || parsed.Scheme != "https" || parsed.Host != "www.douyin.com" || parsed.User != nil || parsed.Fragment != "" { - return false - } - query := parsed.Query() - switch parsed.Path { - case douyinIdentityPath: - return len(query) == 2 && len(query["aid"]) == 1 && query.Get("aid") == "6383" && - len(query["device_platform"]) == 1 && query.Get("device_platform") == "webapp" - case douyinWorksPath: - if len(query) != 3 || len(query["sec_user_id"]) != 1 || !douyinAccountKeyPattern.MatchString(query.Get("sec_user_id")) || - len(query["count"]) != 1 || query.Get("count") != "20" || len(query["max_cursor"]) != 1 { - return false - } - cursor, err := strconv.ParseInt(query.Get("max_cursor"), 10, 64) - return err == nil && cursor >= 0 - case douyinCommentsPath: - if len(query) != 3 || len(query["aweme_id"]) != 1 || !douyinAccountKeyPattern.MatchString(query.Get("aweme_id")) || - len(query["count"]) != 1 || query.Get("count") != "20" || len(query["cursor"]) != 1 { - return false - } - cursor, err := strconv.ParseInt(query.Get("cursor"), 10, 64) - return err == nil && cursor >= 0 - default: - return false - } -} - -type cdpBrowser struct { - endpoint func(string) string - client *http.Client -} - -func (browser cdpBrowser) SetCookies(ctx context.Context, alias string, cookies []douyin.Cookie) error { - connection, err := browser.connect(ctx, alias) - if err != nil { - return err - } - defer connection.Close() - commandID := 0 - cdpCookies := make([]map[string]any, 0, len(cookies)) - for _, cookie := range cookies { - value := map[string]any{ - "name": cookie.Name, "value": cookie.Value, "url": douyinOriginURL, "path": cookie.Path, - "secure": cookie.Secure, "httpOnly": cookie.HTTPOnly, - } - if cookie.SameSite != "" { - value["sameSite"] = cookie.SameSite - } - if cookie.Expires != 0 { - value["expires"] = cookie.Expires - } - cdpCookies = append(cdpCookies, value) - } - if err := cdpCommand(connection, &commandID, "Network.enable", map[string]any{}, nil, nil); err != nil { - return err - } - if err := cdpCommand(connection, &commandID, "Network.clearBrowserCookies", map[string]any{}, nil, nil); err != nil { - return err - } - if err := cdpCommand(connection, &commandID, "Page.enable", map[string]any{}, nil, nil); err != nil { - return err - } - if err := cdpCommand(connection, &commandID, "Page.setLifecycleEventsEnabled", map[string]bool{"enabled": true}, nil, nil); err != nil { - return err - } - var events []cdpMessage - var navigation struct { - FrameID string `json:"frameId"` - LoaderID string `json:"loaderId"` - ErrorText string `json:"errorText"` - } - navigationErr := cdpCommand(connection, &commandID, "Page.navigate", map[string]string{"url": douyinOriginURL}, &navigation, &events) - logrus.WithFields(logrus.Fields{"frame_id": navigation.FrameID, "loader_id": navigation.LoaderID, - "error_text": navigation.ErrorText, "event_count": len(events), "command_error": navigationErr != nil}).Debug("restricted browser navigation response") - if navigationErr != nil || navigation.ErrorText != "" || navigation.FrameID == "" || navigation.LoaderID == "" { - logrus.WithFields(logrus.Fields{"frame_id": navigation.FrameID, "loader_id": navigation.LoaderID, - "error_text": navigation.ErrorText, "event_count": len(events)}).Error("restricted browser navigation response invalid") - return errors.New("restricted browser navigation failed") - } - if err := waitForDouyinPage(ctx, connection, &commandID, navigation.FrameID, events); err != nil { - return err - } - return cdpCommand(connection, &commandID, "Network.setCookies", map[string]any{"cookies": cdpCookies}, nil, nil) -} - -func (browser cdpBrowser) Get(ctx context.Context, alias, target string) (restrictedBrowserResponse, error) { - connection, err := browser.connect(ctx, alias) - if err != nil { - return restrictedBrowserResponse{}, err - } - defer connection.Close() - commandID := 0 - var currentOrigin struct { - Result struct { - Value string `json:"value"` - } `json:"result"` - } - if err := cdpCommand(connection, &commandID, "Runtime.evaluate", map[string]any{ - "expression": "location.origin", "returnByValue": true, - }, ¤tOrigin, nil); err != nil || currentOrigin.Result.Value != douyinOrigin { - return restrictedBrowserResponse{}, errors.New("restricted browser origin changed") - } - encodedURL, _ := json.Marshal(target) - expression := `(async()=>{const r=await fetch(` + string(encodedURL) + `,{credentials:"include",redirect:"error"});` + - `if(!r.body)return {status:r.status,body:"",too_large:false};const q=r.body.getReader(),d=new TextDecoder();let n=0,b="";` + - `for(;;){const x=await q.read();if(x.done)break;if(n+x.value.byteLength>=` + strconv.Itoa(douyinResponseLimit) + - `){await q.cancel();return {too_large:true};}n+=x.value.byteLength;b+=d.decode(x.value,{stream:true});}` + - `b+=d.decode();return {status:r.status,body:b,too_large:false};})()` - var evaluated struct { - Result struct { - Value struct { - Status int `json:"status"` - Body string `json:"body"` - TooLarge bool `json:"too_large"` - } `json:"value"` - } `json:"result"` - ExceptionDetails json.RawMessage `json:"exceptionDetails"` - } - if err := cdpCommand(connection, &commandID, "Runtime.evaluate", map[string]any{ - "expression": expression, "awaitPromise": true, "returnByValue": true, - }, &evaluated, nil); err != nil || len(evaluated.ExceptionDetails) != 0 || evaluated.Result.Value.TooLarge || - evaluated.Result.Value.Status < 200 || evaluated.Result.Value.Status > 599 || - (evaluated.Result.Value.Status >= 300 && evaluated.Result.Value.Status < 400) { - return restrictedBrowserResponse{}, errors.New("restricted browser fetch failed") - } - body := evaluated.Result.Value.Body - return restrictedBrowserResponse{Status: evaluated.Result.Value.Status, Body: body, Challenge: detectDouyinChallenge(evaluated.Result.Value.Status, body)}, nil -} - -func (browser cdpBrowser) connect(ctx context.Context, alias string) (*websocket.Conn, error) { - base := "http://" + namePrefix + alias + ":9222" - if browser.endpoint != nil { - base = browser.endpoint(alias) - } - baseURL, err := url.Parse(base) - if err != nil || baseURL.Scheme != "http" || baseURL.Hostname() == "" || baseURL.Port() == "" { - return nil, errors.New("restricted browser unavailable") - } - if net.ParseIP(baseURL.Hostname()) == nil { - addresses, lookupErr := net.LookupIP(baseURL.Hostname()) - if lookupErr != nil || len(addresses) == 0 { - return nil, errors.New("restricted browser unavailable") - } - baseURL.Host = net.JoinHostPort(addresses[0].String(), baseURL.Port()) - } - base = strings.TrimRight(baseURL.String(), "/") - request, err := http.NewRequestWithContext(ctx, http.MethodGet, base+"/json/list", nil) - if err != nil { - return nil, errors.New("restricted browser unavailable") - } - client := http.Client{Timeout: browserControlTimeout} - if browser.client != nil { - client = *browser.client - } - if client.Timeout <= 0 { - client.Timeout = browserControlTimeout - } - client.CheckRedirect = func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse } - response, err := client.Do(request) - if err != nil { - logrus.WithError(err).WithField("browser_id", alias).Error("restricted browser discovery failed") - return nil, errors.New("restricted browser unavailable") - } - defer response.Body.Close() - var targets []struct { - Type string `json:"type"` - WebSocketDebuggerURL string `json:"webSocketDebuggerUrl"` - } - if response.StatusCode != http.StatusOK { - return nil, errors.New("restricted browser unavailable") - } - body, err := io.ReadAll(io.LimitReader(response.Body, (64<<10)+1)) - if err != nil || len(body) > 64<<10 { - return nil, errors.New("restricted browser unavailable") - } - decoder := json.NewDecoder(bytes.NewReader(body)) - if decoder.Decode(&targets) != nil || len(targets) > 32 { - return nil, errors.New("restricted browser unavailable") - } - var trailing json.RawMessage - if err := decoder.Decode(&trailing); !errors.Is(err, io.EOF) { - return nil, errors.New("restricted browser unavailable") - } - pageTarget := "" - for _, target := range targets { - if target.Type != "page" { - continue - } - if pageTarget != "" { - return nil, errors.New("restricted browser unavailable") - } - websocketURL, parseErr := url.Parse(target.WebSocketDebuggerURL) - if parseErr != nil || websocketURL.Scheme != "ws" || !strings.HasPrefix(websocketURL.Path, "/devtools/page/") { - return nil, errors.New("restricted browser unavailable") - } - if websocketURL.Hostname() == "localhost" || websocketURL.Hostname() == "127.0.0.1" || websocketURL.Hostname() == "::1" { - websocketURL.Host = baseURL.Host - } - if websocketURL.Host != baseURL.Host { - return nil, errors.New("restricted browser unavailable") - } - pageTarget = websocketURL.String() - } - if pageTarget == "" { - return nil, errors.New("restricted browser unavailable") - } - config, err := websocket.NewConfig(pageTarget, "devtools://devtools") - if err != nil { - return nil, errors.New("restricted browser unavailable") - } - connection, err := config.DialContext(ctx) - if err != nil { - logrus.WithError(err).WithField("browser_id", alias).Error("restricted browser websocket failed") - return nil, errors.New("restricted browser unavailable") - } - deadline := time.Now().Add(browserControlTimeout) - if contextDeadline, ok := ctx.Deadline(); ok && contextDeadline.Before(deadline) { - deadline = contextDeadline - } - _ = connection.SetDeadline(deadline) - logrus.WithFields(logrus.Fields{"browser_id": alias, "cdp_endpoint": pageTarget}).Debug("restricted browser websocket connected") - return connection, nil -} - -type cdpMessage struct { - ID int `json:"id"` - Method string `json:"method"` - Params json.RawMessage `json:"params"` - Result json.RawMessage `json:"result"` - Error json.RawMessage `json:"error"` -} - -func cdpCommand(connection *websocket.Conn, commandID *int, method string, parameters any, output any, events *[]cdpMessage) error { - *commandID = *commandID + 1 - if err := websocket.JSON.Send(connection, map[string]any{"id": *commandID, "method": method, "params": parameters}); err != nil { - logrus.WithError(err).WithField("cdp_method", method).Error("restricted browser command send failed") - return errors.New("restricted browser command failed") - } - for range 128 { - var reply cdpMessage - if err := websocket.JSON.Receive(connection, &reply); err != nil { - logrus.WithError(err).WithField("cdp_method", method).Error("restricted browser command receive failed") - return errors.New("restricted browser command failed") - } - if reply.ID != *commandID { - if events != nil && reply.ID == 0 && reply.Method != "" { - *events = append(*events, reply) - } - continue - } - if len(reply.Error) != 0 || len(reply.Result) == 0 || bytes.Equal(bytes.TrimSpace(reply.Result), []byte("null")) { - logrus.WithFields(logrus.Fields{"cdp_method": method, "cdp_error": string(reply.Error), "has_result": len(reply.Result) != 0}).Error("restricted browser command returned failure") - return errors.New("restricted browser command failed") - } - if output != nil && json.Unmarshal(reply.Result, output) != nil { - return errors.New("restricted browser command failed") - } - return nil - } - return errors.New("restricted browser command failed") -} - -func waitForDouyinPage(ctx context.Context, connection *websocket.Conn, commandID *int, frameID string, events []cdpMessage) error { - deadline := time.Now().Add(10 * time.Second) - if contextDeadline, ok := ctx.Deadline(); ok && contextDeadline.Before(deadline) { - deadline = contextDeadline - } - _ = connection.SetReadDeadline(deadline) - bufferedLifecycle := map[string]bool{} - for _, event := range events { - if event.Method != "Page.lifecycleEvent" { - continue - } - var lifecycle struct { - FrameID string `json:"frameId"` - LoaderID string `json:"loaderId"` - Name string `json:"name"` - } - if json.Unmarshal(event.Params, &lifecycle) == nil && lifecycle.FrameID == frameID && (lifecycle.Name == "DOMContentLoaded" || lifecycle.Name == "load") { - bufferedLifecycle[lifecycle.LoaderID+":"+lifecycle.Name] = true - } - } - for { - var event cdpMessage - if len(events) != 0 { - event, events = events[0], events[1:] - } else if err := websocket.JSON.Receive(connection, &event); err != nil { - logrus.WithError(err).Error("restricted browser lifecycle receive failed") - return errors.New("restricted browser navigation failed") - } - if event.Method != "Page.lifecycleEvent" { - continue - } - var lifecycle struct { - FrameID string `json:"frameId"` - LoaderID string `json:"loaderId"` - Name string `json:"name"` - } - if err := json.Unmarshal(event.Params, &lifecycle); err != nil { - logrus.WithError(err).Debug("restricted browser lifecycle decode failed") - continue - } - if lifecycle.FrameID != frameID || lifecycle.Name != "DOMContentLoaded" && lifecycle.Name != "load" || bufferedLifecycle[lifecycle.LoaderID+":"+lifecycle.Name] { - continue - } - var evaluated struct { - Result struct { - Value string `json:"value"` - } `json:"result"` - } - if err := cdpCommand(connection, commandID, "Runtime.evaluate", map[string]any{ - "expression": "location.origin", "returnByValue": true, - }, &evaluated, nil); err != nil || evaluated.Result.Value != douyinOrigin { - logrus.WithFields(logrus.Fields{"frame_id": lifecycle.FrameID, "loader_id": lifecycle.LoaderID, - "origin": evaluated.Result.Value}).Error("restricted browser navigation origin failed") - return errors.New("restricted browser navigation failed") - } - return nil - } -} - -func detectDouyinChallenge(status int, body string) douyin.Challenge { - lower := strings.ToLower(body) - if status == http.StatusPreconditionFailed || strings.Contains(lower, "captcha") || strings.Contains(lower, "verify_center_decision_conf") { - return douyin.ChallengeCaptcha - } - if strings.Contains(lower, "device_challenge") || strings.Contains(lower, "device verification") { - return douyin.ChallengeDevice - } - return douyin.ChallengeNone -} diff --git a/cmd/docker-gateway/douyin_cdp_integration_test.go b/cmd/docker-gateway/douyin_cdp_integration_test.go deleted file mode 100644 index 83f3c3f..0000000 --- a/cmd/docker-gateway/douyin_cdp_integration_test.go +++ /dev/null @@ -1,205 +0,0 @@ -package main - -import ( - "context" - "io" - "net" - "net/http" - "net/http/httptest" - "os" - "os/exec" - "path/filepath" - "strings" - "sync" - "syscall" - "testing" - "time" - - "git.ipao.vip/rogee/creator-hub/internal/douyin" -) - -func TestCDPBrowserRealNavigationIsolation(t *testing.T) { - if os.Getenv("CREATORHUB_REAL_CDP_TEST") != "1" { - t.Skip("set CREATORHUB_REAL_CDP_TEST=1 to run against local Chrome") - } - chrome, err := exec.LookPath("google-chrome") - if err != nil { - t.Skip("google-chrome is unavailable") - } - - var mu sync.Mutex - phase := "redirect" - release := make(chan struct{}) - mainRequests := make(chan string, 4) - loginRequests := make(chan string, 1) - subdomainRequests := make(chan string, 1) - server := httptest.NewTLSServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { - host := request.Host - if name, _, splitErr := net.SplitHostPort(host); splitErr == nil { - host = name - } - switch host { - case "login.douyin.com": - loginRequests <- request.Header.Get("Cookie") - _, _ = response.Write([]byte("redirected")) - case "api.douyin.com": - subdomainRequests <- request.Header.Get("Cookie") - _, _ = response.Write([]byte("ok")) - case "www.douyin.com": - if request.URL.Path == "/" { - mainRequests <- request.Header.Get("Cookie") - } - mu.Lock() - currentPhase, currentRelease := phase, release - mu.Unlock() - if currentPhase == "redirect" { - http.Redirect(response, request, "https://login.douyin.com/landing", http.StatusFound) - return - } - if currentPhase == "blocked" { - <-currentRelease - } - response.Header().Set("Content-Type", "text/html") - _, _ = response.Write([]byte(``)) - default: - response.WriteHeader(http.StatusNotFound) - } - })) - t.Cleanup(server.Close) - allowedHosts := map[string]bool{"www.douyin.com": true, "login.douyin.com": true, "api.douyin.com": true} - proxy := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { - host := request.Host - if name, _, splitErr := net.SplitHostPort(host); splitErr == nil { - host = name - } - if request.Method != http.MethodConnect || !allowedHosts[host] { - response.WriteHeader(http.StatusForbidden) - return - } - upstream, dialErr := net.Dial("tcp", server.Listener.Addr().String()) - if dialErr != nil { - response.WriteHeader(http.StatusBadGateway) - return - } - client, _, hijackErr := response.(http.Hijacker).Hijack() - if hijackErr != nil { - upstream.Close() - return - } - _, _ = client.Write([]byte("HTTP/1.1 200 Connection Established\r\n\r\n")) - go func() { - _, _ = io.Copy(upstream, client) - _ = upstream.Close() - }() - _, _ = io.Copy(client, upstream) - _ = client.Close() - })) - t.Cleanup(proxy.Close) - - profile := t.TempDir() - command := exec.Command(chrome, "--headless=new", "--no-sandbox", "--disable-gpu", "--disable-background-networking", "--disable-quic", - "--disable-dev-shm-usage", "--no-first-run", "--no-default-browser-check", "--password-store=basic", "--use-mock-keychain", - "--ignore-certificate-errors", "--proxy-server="+proxy.URL, - "--remote-debugging-address=127.0.0.1", "--remote-debugging-port=0", "--remote-allow-origins=*", - "--user-data-dir="+profile, "about:blank") - command.SysProcAttr = &syscall.SysProcAttr{Setpgid: true} - if err := command.Start(); err != nil { - t.Fatal(err) - } - t.Cleanup(func() { - _ = syscall.Kill(-command.Process.Pid, syscall.SIGKILL) - _ = command.Wait() - }) - - var debugPort string - for deadline := time.Now().Add(5 * time.Second); time.Now().Before(deadline); time.Sleep(25 * time.Millisecond) { - content, readErr := os.ReadFile(filepath.Join(profile, "DevToolsActivePort")) - if readErr == nil { - debugPort = strings.SplitN(string(content), "\n", 2)[0] - break - } - } - if debugPort == "" { - t.Fatal("Chrome did not expose a DevTools port") - } - time.Sleep(250 * time.Millisecond) - browser := cdpBrowser{endpoint: func(string) string { return "http://127.0.0.1:" + debugPort }} - cookie := []douyin.Cookie{{Name: "sessionid", Value: "fake-secret", Domain: ".douyin.com", Path: "/", Secure: true}} - - ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) - redirectErr := browser.SetCookies(ctx, "account-a", cookie) - if redirectErr == nil { - cancel() - t.Fatal("accepted a redirected navigation") - } - cancel() - select { - case got := <-loginRequests: - if strings.Contains(got, "fake-secret") { - t.Fatal("cookie leaked during redirected navigation") - } - case <-time.After(3 * time.Second): - t.Fatalf("redirect target was not reached: %v", redirectErr) - } - - mu.Lock() - phase = "success" - mu.Unlock() - ctx, cancel = context.WithTimeout(context.Background(), 15*time.Second) - if err := browser.SetCookies(ctx, "account-a", cookie); err != nil { - cancel() - t.Fatal(err) - } - cancel() - select { - case <-mainRequests: - case <-time.After(3 * time.Second): - t.Fatal("successful navigation was not observed") - } - - mu.Lock() - phase = "blocked" - release = make(chan struct{}) - currentRelease := release - mu.Unlock() - defer func() { - select { - case <-currentRelease: - default: - close(currentRelease) - } - }() - done := make(chan error, 1) - ctx, cancel = context.WithTimeout(context.Background(), 15*time.Second) - go func() { done <- browser.SetCookies(ctx, "account-a", cookie) }() - select { - case got := <-mainRequests: - if strings.Contains(got, "fake-secret") { - cancel() - t.Fatal("cookie leaked during navigation") - } - case <-time.After(3 * time.Second): - cancel() - t.Fatal("blocked navigation was not observed") - } - select { - case err := <-done: - cancel() - t.Fatalf("navigation completed before its loader: %v", err) - case <-time.After(100 * time.Millisecond): - } - close(currentRelease) - if err := <-done; err != nil { - cancel() - t.Fatal(err) - } - cancel() - select { - case got := <-subdomainRequests: - if strings.Contains(got, "fake-secret") { - t.Fatal("host-only cookie leaked to a subdomain") - } - case <-time.After(3 * time.Second): - t.Fatal("subdomain probe did not run") - } -} diff --git a/cmd/docker-gateway/douyin_test.go b/cmd/docker-gateway/douyin_test.go deleted file mode 100644 index f72cd6f..0000000 --- a/cmd/docker-gateway/douyin_test.go +++ /dev/null @@ -1,410 +0,0 @@ -package main - -import ( - "context" - "encoding/json" - "io" - "net/http" - "net/http/httptest" - "os" - "strings" - "sync" - "testing" - "time" - - "git.ipao.vip/rogee/creator-hub/internal/douyin" - "golang.org/x/net/websocket" -) - -const douyinIdentityURL = "https://www.douyin.com/aweme/v1/web/user/profile/self/?aid=6383&device_platform=webapp" - -type fakeRestrictedBrowser struct { - cookies []douyin.Cookie - urls []string - response restrictedBrowserResponse - after func() -} - -func (browser *fakeRestrictedBrowser) SetCookies(_ context.Context, _ string, cookies []douyin.Cookie) error { - browser.cookies = cookies - if browser.after != nil { - browser.after() - } - return nil -} - -func (browser *fakeRestrictedBrowser) Get(_ context.Context, _ string, target string) (restrictedBrowserResponse, error) { - browser.urls = append(browser.urls, target) - if browser.after != nil { - browser.after() - } - return browser.response, nil -} - -func TestGatewayRestrictedDouyinContract(t *testing.T) { - labels := map[string]string{ - managedLabel: "true", idLabel: "account-a", bindingVersionLabel: "2", networkIDLabel: "network-a", networkExitLabel: "exit-a", - } - self, _ := os.Hostname() - var stateMu sync.Mutex - containerNetworks := map[string]string{"creatorhub_browser-account-a": "network-a"} - runtimeAttached := true - server := httptest.NewServer(withAliasReservations(self, func(response http.ResponseWriter, request *http.Request) { - switch request.URL.Path { - case "/containers/" + namePrefix + "account-a/json": - stateMu.Lock() - labelCopy, networkCopy := map[string]string{}, map[string]any{} - for key, value := range labels { - labelCopy[key] = value - } - for name, id := range containerNetworks { - networkCopy[name] = map[string]string{"NetworkID": id} - } - stateMu.Unlock() - _ = json.NewEncoder(response).Encode(map[string]any{"Id": "runtime-a", "Config": map[string]any{"Labels": labelCopy}, - "NetworkSettings": map[string]any{"Networks": networkCopy}}) - case "/networks/network-a": - stateMu.Lock() - members := map[string]any{self: map[string]string{"Name": self, "IPv4Address": "127.0.0.1/8"}} - if runtimeAttached { - members["runtime-a"] = map[string]string{"Name": namePrefix + "account-a", "IPv4Address": "127.0.0.2/8"} - } - stateMu.Unlock() - _ = json.NewEncoder(response).Encode(map[string]any{ - "Id": "network-a", "Name": "creatorhub_browser-account-a", "Driver": "bridge", "Internal": false, "Attachable": false, "Ingress": false, - "Labels": map[string]string{managedLabel: "true", networkRoleLabel: browserNetworkRole, idLabel: "account-a", bindingVersionLabel: "2"}, - "Containers": members, - }) - default: - response.WriteHeader(http.StatusNotFound) - } - })) - defer server.Close() - docker := dockerClient{baseURL: server.URL, client: server.Client(), slow: server.Client()} - browser := &fakeRestrictedBrowser{response: restrictedBrowserResponse{Status: 200, Body: `{"status_code":0}`, Challenge: douyin.ChallengeNone}} - app := newGatewayWithBrowser(docker, "creatorhub_browser", testToken, self, browser) - generation := `"binding_version":2,"runtime_id":"runtime-a","network_id":"network-a","network_exit_id":"exit-a"` - - cookieBody := `{` + generation + `,"cookies":[{"name":"sessionid","value":"private-session","domain":".douyin.com","path":"/"}]}` - response, err := app.Test(authed(http.MethodPost, "/v1/browsers/account-a/douyin/cookies", strings.NewReader(cookieBody))) - if err != nil || response.StatusCode != http.StatusNoContent || len(browser.cookies) != 1 || browser.cookies[0].Value != "private-session" { - t.Fatalf("set cookies failed: status=%d cookies=%#v err=%v", response.StatusCode, browser.cookies, err) - } - response.Body.Close() - - identityURL := douyinIdentityURL - getBody := `{` + generation + `,"url":"` + identityURL + `"}` - response, err = app.Test(authed(http.MethodPost, "/v1/browsers/account-a/douyin/get", strings.NewReader(getBody))) - body, _ := io.ReadAll(response.Body) - response.Body.Close() - if err != nil || response.StatusCode != http.StatusOK || len(browser.urls) != 1 || browser.urls[0] != identityURL || - !strings.Contains(string(body), `\"status_code\":0`) || strings.Contains(string(body), "private-session") { - t.Fatalf("get failed or leaked cookies: status=%d urls=%#v body=%s err=%v", response.StatusCode, browser.urls, body, err) - } - - for name, request := range map[string]*http.Request{ - "unauthenticated": httptest.NewRequest(http.MethodPost, "/v1/browsers/account-a/douyin/get", strings.NewReader(getBody)), - "invalid account": authed(http.MethodPost, "/v1/browsers/AccountA/douyin/get", strings.NewReader(getBody)), - "generic URL": authed(http.MethodPost, "/v1/browsers/account-a/douyin/get", strings.NewReader(`{`+generation+`,"url":"https://example.com/"}`)), - "generic CDP": authed(http.MethodPost, "/v1/browsers/account-a/douyin/get", strings.NewReader(`{`+generation+`,"url":"`+identityURL+`","method":"Runtime.evaluate"}`)), - "trailing null": authed(http.MethodPost, "/v1/browsers/account-a/douyin/get", strings.NewReader(getBody+`null`)), - "stale generation": authed(http.MethodPost, "/v1/browsers/account-a/douyin/get", strings.NewReader(`{"binding_version":1,"runtime_id":"runtime-a","network_id":"network-a","network_exit_id":"exit-a","url":"`+identityURL+`"}`)), - } { - t.Run(name, func(t *testing.T) { - before := len(browser.urls) - response, err := app.Test(request) - if err != nil { - t.Fatal(err) - } - response.Body.Close() - want := http.StatusBadRequest - if name == "unauthenticated" { - want = http.StatusUnauthorized - } else if name == "stale generation" { - want = http.StatusConflict - } - if response.StatusCode != want || len(browser.urls) != before { - t.Fatalf("status=%d want=%d calls=%d want=%d", response.StatusCode, want, len(browser.urls), before) - } - }) - } - - stateMu.Lock() - runtimeAttached = false - stateMu.Unlock() - before := len(browser.urls) - response, err = app.Test(authed(http.MethodPost, "/v1/browsers/account-a/douyin/get", strings.NewReader(getBody))) - if err != nil || response.StatusCode != http.StatusConflict || len(browser.urls) != before { - t.Fatalf("wrong network membership reached browser: status=%d calls=%d want=%d err=%v", response.StatusCode, len(browser.urls), before, err) - } - response.Body.Close() - stateMu.Lock() - runtimeAttached = true - stateMu.Unlock() - - browser.after = func() { - stateMu.Lock() - containerNetworks["other-tenant"] = "network-b" - stateMu.Unlock() - } - response, err = app.Test(authed(http.MethodPost, "/v1/browsers/account-a/douyin/cookies", strings.NewReader(cookieBody))) - if err != nil || response.StatusCode != http.StatusConflict { - t.Fatalf("post-operation cross-network attachment was accepted: status=%d err=%v", response.StatusCode, err) - } - response.Body.Close() - stateMu.Lock() - delete(containerNetworks, "other-tenant") - stateMu.Unlock() - - browser.after = func() { - stateMu.Lock() - labels[networkIDLabel] = "network-replaced" - stateMu.Unlock() - } - response, err = app.Test(authed(http.MethodPost, "/v1/browsers/account-a/douyin/cookies", strings.NewReader(cookieBody))) - if err != nil || response.StatusCode != http.StatusConflict { - t.Fatalf("post-operation generation replacement was accepted: status=%d err=%v", response.StatusCode, err) - } - response.Body.Close() -} - -func TestCDPBrowserUsesOnlyNarrowCommands(t *testing.T) { - var server *httptest.Server - var mu sync.Mutex - methods := []string{} - secretSeen := false - cookieNames := []string{} - cookiesHostOnly := true - setCookieCalls := 0 - pageOrigin := douyinOrigin - onlyOldLoader := false - fetchMode := "ok" - fetchExpression := "" - mux := http.NewServeMux() - mux.HandleFunc("/json/list", func(response http.ResponseWriter, request *http.Request) { - wsURL := "ws" + strings.TrimPrefix(server.URL, "http") + "/devtools/page/one" - _ = json.NewEncoder(response).Encode([]map[string]string{{"type": "page", "webSocketDebuggerUrl": wsURL}}) - }) - mux.Handle("/devtools/page/one", websocket.Server{ - Handshake: func(*websocket.Config, *http.Request) error { return nil }, - Handler: func(connection *websocket.Conn) { - for { - var command struct { - ID int `json:"id"` - Method string `json:"method"` - Params json.RawMessage `json:"params"` - } - if websocket.JSON.Receive(connection, &command) != nil { - return - } - mu.Lock() - methods = append(methods, command.Method) - secretSeen = secretSeen || strings.Contains(string(command.Params), "private-session") - result := any(map[string]any{}) - var beforeReply, afterReply []map[string]any - switch command.Method { - case "Network.clearBrowserCookies": - cookieNames = nil - case "Network.setCookies": - var params struct { - Cookies []map[string]any `json:"cookies"` - } - _ = json.Unmarshal(command.Params, ¶ms) - setCookieCalls++ - cookieNames = cookieNames[:0] - for _, cookie := range params.Cookies { - cookieNames = append(cookieNames, cookie["name"].(string)) - _, hasDomain := cookie["domain"] - cookiesHostOnly = cookiesHostOnly && !hasDomain && cookie["url"] == douyinOriginURL - } - case "Page.navigate": - result = map[string]any{"frameId": "frame-new", "loaderId": "loader-new"} - beforeReply = append(beforeReply, map[string]any{"method": "Page.lifecycleEvent", "params": map[string]any{ - "frameId": "frame-old", "loaderId": "loader-old", "name": "load", - }}) - if !onlyOldLoader { - afterReply = append(afterReply, map[string]any{"method": "Page.lifecycleEvent", "params": map[string]any{ - "frameId": "frame-new", "loaderId": "loader-final", "name": "load", - }}) - } - case "Runtime.evaluate": - var params struct { - Expression string `json:"expression"` - } - _ = json.Unmarshal(command.Params, ¶ms) - if params.Expression == "location.origin" { - result = map[string]any{"result": map[string]any{"value": pageOrigin}} - } else { - fetchExpression = params.Expression - value := map[string]any{"status": 412, "body": `{"captcha":true}`, "too_large": false} - if fetchMode == "redirect" { - value = map[string]any{"status": 302, "body": "", "too_large": false} - } else if fetchMode == "too_large" { - value = map[string]any{"too_large": true} - } - result = map[string]any{"result": map[string]any{"value": map[string]any{ - "status": value["status"], "body": value["body"], "too_large": value["too_large"], - }}} - } - } - mu.Unlock() - for _, event := range beforeReply { - _ = websocket.JSON.Send(connection, event) - } - _ = websocket.JSON.Send(connection, map[string]any{"id": command.ID, "result": result}) - for _, event := range afterReply { - _ = websocket.JSON.Send(connection, event) - } - } - }, - }) - server = httptest.NewServer(mux) - defer server.Close() - browser := cdpBrowser{endpoint: func(string) string { return server.URL }, client: server.Client()} - if err := browser.SetCookies(context.Background(), "account-a", []douyin.Cookie{{ - Name: "old_auth", Value: "old-session", Domain: ".douyin.com", Path: "/", - }}); err != nil { - t.Fatal(err) - } - if err := browser.SetCookies(context.Background(), "account-a", []douyin.Cookie{{ - Name: "sessionid", Value: "private-session", Domain: ".douyin.com", Path: "/", - }}); err != nil { - t.Fatal(err) - } - result, err := browser.Get(context.Background(), "account-a", douyinIdentityURL) - if err != nil || result.Status != 412 || result.Challenge != douyin.ChallengeCaptcha { - t.Fatalf("unexpected CDP response: %#v err=%v", result, err) - } - mu.Lock() - if !secretSeen || !cookiesHostOnly || strings.Join(cookieNames, ",") != "sessionid" || - strings.Join(methods, ",") != "Network.enable,Network.clearBrowserCookies,Page.enable,Page.setLifecycleEventsEnabled,Page.navigate,Runtime.evaluate,Network.setCookies,"+ - "Network.enable,Network.clearBrowserCookies,Page.enable,Page.setLifecycleEventsEnabled,Page.navigate,Runtime.evaluate,Network.setCookies,Runtime.evaluate,Runtime.evaluate" || - !strings.Contains(fetchExpression, `redirect:"error"`) || !strings.Contains(fetchExpression, "getReader()") || - !strings.Contains(fetchExpression, "q.cancel()") || !strings.Contains(fetchExpression, ">=1048576") || strings.Contains(fetchExpression, "r.text()") { - mu.Unlock() - t.Fatalf("unexpected CDP contract: methods=%#v cookies=%#v secret_seen=%v expression=%s", methods, cookieNames, secretSeen, fetchExpression) - } - pageOrigin = "https://login.douyin.com" - setCookiesBeforeRedirect := setCookieCalls - mu.Unlock() - if err := browser.SetCookies(context.Background(), "account-a", []douyin.Cookie{{ - Name: "sessionid", Value: "private-session", Domain: ".douyin.com", Path: "/", - }}); err == nil { - t.Fatal("accepted navigation redirected to a Douyin subdomain") - } - mu.Lock() - if setCookieCalls != setCookiesBeforeRedirect { - mu.Unlock() - t.Fatal("set cookies before rejecting redirected navigation") - } - pageOrigin, fetchMode = douyinOrigin, "redirect" - mu.Unlock() - if _, err := browser.Get(context.Background(), "account-a", douyinIdentityURL); err == nil { - t.Fatal("accepted a redirected fetch") - } - mu.Lock() - fetchMode = "too_large" - mu.Unlock() - if _, err := browser.Get(context.Background(), "account-a", douyinIdentityURL); err == nil { - t.Fatal("accepted a response at the 1 MiB limit") - } - mu.Lock() - onlyOldLoader = true - setCookiesBeforeOldLoader := setCookieCalls - mu.Unlock() - ctx, cancel := context.WithTimeout(context.Background(), 100*time.Millisecond) - defer cancel() - if err := browser.SetCookies(ctx, "account-a", []douyin.Cookie{{ - Name: "sessionid", Value: "private-session", Domain: ".douyin.com", Path: "/", - }}); err == nil { - t.Fatal("accepted an old page load event for the new navigation") - } - mu.Lock() - defer mu.Unlock() - if setCookieCalls != setCookiesBeforeOldLoader { - t.Fatal("set cookies before the new loader completed") - } -} - -func TestCDPDiscoveryDoesNotFollowRedirects(t *testing.T) { - redirected := 0 - server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { - if request.URL.Path == "/json/list" { - http.Redirect(response, request, "/redirected", http.StatusFound) - return - } - redirected++ - response.WriteHeader(http.StatusInternalServerError) - })) - defer server.Close() - browser := cdpBrowser{endpoint: func(string) string { return server.URL }, client: server.Client()} - if _, err := browser.connect(context.Background(), "account-a"); err == nil || redirected != 0 { - t.Fatalf("discovery redirect was followed: redirected=%d err=%v", redirected, err) - } -} - -func TestCDPDiscoveryRequiresOnePage(t *testing.T) { - server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { - _ = json.NewEncoder(response).Encode([]map[string]string{ - {"type": "page", "webSocketDebuggerUrl": "ws://localhost/devtools/page/one"}, - {"type": "page", "webSocketDebuggerUrl": "ws://localhost/devtools/page/two"}, - }) - })) - defer server.Close() - browser := cdpBrowser{endpoint: func(string) string { return server.URL }, client: server.Client()} - if _, err := browser.connect(context.Background(), "account-a"); err == nil { - t.Fatal("accepted a profile with multiple page targets") - } -} - -func TestCDPDiscoveryRequiresOneJSONValue(t *testing.T) { - for name, suffix := range map[string]string{ - "null": "null", "other value": `{}`, "garbage": "garbage", "oversized": strings.Repeat(" ", 64<<10), "whitespace": " \n\t", - } { - t.Run(name, func(t *testing.T) { - websocketAttempts := 0 - var server *httptest.Server - mux := http.NewServeMux() - mux.HandleFunc("/json/list", func(response http.ResponseWriter, request *http.Request) { - wsURL := "ws" + strings.TrimPrefix(server.URL, "http") + "/devtools/page/one" - _, _ = response.Write([]byte(`[{"type":"page","webSocketDebuggerUrl":"` + wsURL + `"}]` + suffix)) - }) - mux.Handle("/devtools/page/one", websocket.Server{ - Handshake: func(*websocket.Config, *http.Request) error { - websocketAttempts++ - return nil - }, - Handler: func(connection *websocket.Conn) {}, - }) - server = httptest.NewServer(mux) - defer server.Close() - - connection, err := (cdpBrowser{endpoint: func(string) string { return server.URL }, client: server.Client()}).connect(context.Background(), "account-a") - if connection != nil { - connection.Close() - } - valid := name == "whitespace" - wantAttempts := 0 - if valid { - wantAttempts = 1 - } - if (err == nil) != valid || websocketAttempts != wantAttempts { - t.Fatalf("err=%v websocket attempts=%d", err, websocketAttempts) - } - }) - } -} - -func TestDouyinURLContract(t *testing.T) { - for target, want := range map[string]bool{ - douyinIdentityURL: true, - "https://www.douyin.com" + douyinWorksPath + "?sec_user_id=sec-a&count=20&max_cursor=0": true, - "https://www.douyin.com" + douyinWorksPath + "?sec_user_id=sec-a&count=20&max_cursor=1": true, - "https://www.douyin.com" + douyinWorksPath + "?sec_user_id=sec-a&count=20&max_cursor=0&method=publish": false, - "https://www.douyin.com/aweme/v1/web/commit/item/": false, - } { - if got := validDouyinURL(target); got != want { - t.Fatalf("validDouyinURL(%q)=%v want=%v", target, got, want) - } - } -} diff --git a/cmd/docker-gateway/main.go b/cmd/docker-gateway/main.go deleted file mode 100644 index ebb8ca2..0000000 --- a/cmd/docker-gateway/main.go +++ /dev/null @@ -1,1547 +0,0 @@ -package main - -import ( - "bytes" - "context" - "crypto/rand" - "crypto/subtle" - "encoding/hex" - "encoding/json" - "errors" - "fmt" - "io" - "net" - "net/http" - "net/url" - "os" - "os/signal" - "regexp" - "strconv" - "strings" - "syscall" - "time" - "unicode/utf8" - - "github.com/gofiber/fiber/v3" - "github.com/sirupsen/logrus" - "github.com/spf13/cobra" - "github.com/spf13/viper" -) - -const ( - browserUser = "1000:1000" - browserEntrypoint = "/usr/local/bin/docker-entrypoint.sh" - managedLabel = "io.creatorhub.managed" - idLabel = "io.creatorhub.runtime-id" - nameLabel = "io.creatorhub.display-name" - bindingVersionLabel = "io.creatorhub.binding-version" - networkExitLabel = "io.creatorhub.network-exit-id" - proxyPortLabel = "io.creatorhub.proxy-port" - networkIDLabel = "io.creatorhub.network-id" - networkRoleLabel = "io.creatorhub.network-role" - gatewayMemberLabel = "io.creatorhub.gateway-member" - browserNetworkRole = "browser" - controlNetworkName = "creatorhub_control" - namePrefix = "creatorhub-browser-" - reservationPrefix = "creatorhub-reservation-" - reservationLabel = "io.creatorhub.alias-reservation" - reservationGenLabel = "io.creatorhub.reservation-generation" - pullTimeout = 10 * time.Minute -) - -var ( - runtimeIDPattern = regexp.MustCompile(`^[a-z0-9][a-z0-9-]{0,31}$`) - networkNamePattern = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9_.-]{0,63}$`) - imageRefPattern = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9._:/@-]{0,300}$`) - volumePattern = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9_.-]{0,63}$`) - exitIDPattern = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9._/-]{0,127}$`) -) - -var ( - errInvalidRuntimeID = errors.New("invalid runtime id") - errUnmanagedContainer = errors.New("refusing to operate on a container not owned by CreatorHub") - errGenerationConflict = errors.New("container generation does not match request") - errUnauthorized = errors.New("gateway token rejected") -) - -type serviceConfig struct { - listenAddr string - dockerSock string - network string - token string - logLevel logrus.Level -} - -type dockerClient struct { - baseURL string - client *http.Client // 常规 Docker 调用 - slow *http.Client // 镜像拉取等长操作,不设整体超时 -} - -type dockerAliasReservations struct { - docker dockerClient - self string -} - -type tenantNetworkGeneration struct { - ID string - Name string - Created bool - ConnectedSelf bool - ConnectedRuntime bool - GatewayMembers []string - SelfMember string - RuntimeAttached bool -} - -type dockerTenantNetwork struct { - ID string `json:"Id"` - Name string `json:"Name"` - Driver string `json:"Driver"` - Internal bool `json:"Internal"` - Attachable bool `json:"Attachable"` - Ingress bool `json:"Ingress"` - Labels map[string]string `json:"Labels"` - Containers map[string]struct { - Name string `json:"Name"` - IPv4Address string `json:"IPv4Address"` - } `json:"Containers"` -} - -type gateway struct { - docker dockerClient - network string - self string - token string - browser restrictedBrowser - proxies *memoryProxyRegistry - locks *dockerAliasReservations -} - -// createRequest 全量字段由平台下发;网关不做业务决策,只做输入合法性校验。 -type createRequest struct { - Alias string `json:"alias"` - Name string `json:"name"` - Image string `json:"image"` - Cmd []string `json:"cmd"` - Volume string `json:"volume"` - BindingVersion int64 `json:"binding_version"` - NetworkExitID string `json:"network_exit_id"` - NetworkExit gatewayProxyExit `json:"network_exit"` - Stopped bool `json:"stopped,omitempty"` -} - -type gatewayProxyExit struct { - Protocol string `json:"protocol"` - Host string `json:"host"` - Port int `json:"port"` - Username string `json:"username"` - Password string `json:"password"` -} - -type generationRequest struct { - BindingVersion int64 `json:"binding_version"` - RuntimeID string `json:"runtime_id"` - NetworkID string `json:"network_id"` -} - -type proxyRestoreRequest struct { - BindingVersion int64 `json:"binding_version"` - RuntimeID string `json:"runtime_id"` - NetworkID string `json:"network_id"` - NetworkExitID string `json:"network_exit_id"` - NetworkExit gatewayProxyExit `json:"network_exit"` -} - -type browser struct { - ID string `json:"id"` - Alias string `json:"alias"` - Name string `json:"name"` - State string `json:"state"` - Status string `json:"status"` - Endpoint string `json:"endpoint"` - BindingVersion int64 `json:"binding_version"` - NetworkExitID string `json:"network_exit_id"` - NetworkID string `json:"network_id"` - ProxyReady bool `json:"proxy_ready"` -} - -func main() { - logrus.SetFormatter(&logrus.JSONFormatter{}) - ctx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM) - defer stop() - if err := newCommand().ExecuteContext(ctx); err != nil { - logrus.WithField("service", "docker-gateway").WithError(err).Error("service stopped") - os.Exit(1) - } -} - -func newCommand() *cobra.Command { - command := &cobra.Command{ - Use: "docker-gateway", - Short: "Run the restricted CreatorHub Docker gateway", - Args: cobra.NoArgs, - SilenceErrors: true, - SilenceUsage: true, - RunE: func(command *cobra.Command, _ []string) error { - cfg, err := loadConfig() - if err != nil { - return err - } - logrus.SetLevel(cfg.logLevel) - return run(command, cfg) - }, - } - return command -} - -func loadConfig() (serviceConfig, error) { - v := viper.New() - v.SetDefault("listen_addr", ":8081") - v.SetDefault("docker_socket", "/var/run/docker.sock") - v.SetDefault("browser_network", "creatorhub_browser") - v.SetDefault("log_level", "info") - _ = v.BindEnv("listen_addr", "LISTEN_ADDR") - _ = v.BindEnv("docker_socket", "DOCKER_SOCKET") - _ = v.BindEnv("browser_network", "BROWSER_NETWORK") - _ = v.BindEnv("gateway_token", "GATEWAY_TOKEN") - _ = v.BindEnv("log_level", "LOG_LEVEL") - - level, err := logrus.ParseLevel(v.GetString("log_level")) - if err != nil { - return serviceConfig{}, errors.New("LOG_LEVEL must be panic, fatal, error, warn, info, debug, or trace") - } - cfg := serviceConfig{ - listenAddr: strings.TrimSpace(v.GetString("listen_addr")), - dockerSock: strings.TrimSpace(v.GetString("docker_socket")), - network: strings.TrimSpace(v.GetString("browser_network")), - token: strings.TrimSpace(v.GetString("gateway_token")), - logLevel: level, - } - if cfg.listenAddr == "" { - return serviceConfig{}, errors.New("LISTEN_ADDR must not be empty") - } - if err := validateListenAddr(cfg.listenAddr); err != nil { - return serviceConfig{}, err - } - if cfg.dockerSock == "" { - return serviceConfig{}, errors.New("DOCKER_SOCKET must not be empty") - } - if len(cfg.token) < 16 { - return serviceConfig{}, errors.New("GATEWAY_TOKEN must be at least 16 characters") - } - if err := validateBrowserNetwork(cfg.network); err != nil { - return serviceConfig{}, err - } - return cfg, nil -} - -func validateListenAddr(addr string) error { - _, port, err := net.SplitHostPort(addr) - if err != nil { - return errors.New("LISTEN_ADDR must be a host:port address") - } - number, err := strconv.Atoi(port) - if err != nil || number < 1 || number > 65535 { - return errors.New("LISTEN_ADDR port must be 1..65535") - } - return nil -} - -func run(command *cobra.Command, cfg serviceConfig) error { - transport := &http.Transport{ - DialContext: func(ctx context.Context, _, _ string) (net.Conn, error) { - return (&net.Dialer{}).DialContext(ctx, "unix", cfg.dockerSock) - }, - } - defer transport.CloseIdleConnections() - docker := dockerClient{ - baseURL: "http://docker/v1.43", - client: &http.Client{Transport: transport, Timeout: 30 * time.Second}, - slow: &http.Client{Transport: transport}, - } - logrus.WithFields(logrus.Fields{ - "service": "docker-gateway", - "listen_addr": cfg.listenAddr, - "network": cfg.network, - }).Info("service starting") - return newGateway(docker, cfg.network, cfg.token).Listen(cfg.listenAddr, fiber.ListenConfig{ - GracefulContext: command.Context(), - DisableStartupMessage: true, - }) -} - -func newGateway(client dockerClient, network, token string) *fiber.App { - self, _ := os.Hostname() - return newGatewayWithSelf(client, network, token, self) -} - -func newGatewayWithSelf(client dockerClient, network, token, self string) *fiber.App { - return newGatewayWithBrowser(client, network, token, self, cdpBrowser{}) -} - -func newGatewayWithBrowser(client dockerClient, network, token, self string, browser restrictedBrowser) *fiber.App { - api := gateway{docker: client, network: network, self: self, token: token, proxies: newMemoryProxyRegistry(), - locks: &dockerAliasReservations{docker: client, self: self}, browser: browser} - app := fiber.New(fiber.Config{ - AppName: "CreatorHub Docker gateway", - BodyLimit: 1 << 20, - ReadTimeout: 5 * time.Second, - IdleTimeout: 60 * time.Second, - ErrorHandler: func(c fiber.Ctx, err error) error { - status := http.StatusInternalServerError - var fiberError *fiber.Error - if errors.As(err, &fiberError) && fiberError != nil { - status = fiberError.Code - } - return writeError(c, status, err) - }, - }) - app.Get("/healthz", func(c fiber.Ctx) error { - c.Status(fiber.StatusNoContent) - return nil - }) - app.Use("/v1", api.authorize) - app.Get("/v1/browsers", api.list) - app.Post("/v1/browsers", api.create) - app.Post("/v1/browsers/:id/proxy", api.restoreProxy) - app.Post("/v1/browsers/:id/douyin/cookies", api.setDouyinCookies) - app.Post("/v1/browsers/:id/douyin/get", api.getDouyin) - app.Post("/v1/browsers/:id/:action", api.changeState) - app.Delete("/v1/browsers/:id", api.remove) - return app -} - -func (api gateway) authorize(c fiber.Ctx) error { - expected := "Bearer " + api.token - if subtle.ConstantTimeCompare([]byte(c.Get(fiber.HeaderAuthorization)), []byte(expected)) != 1 { - return writeError(c, http.StatusUnauthorized, errUnauthorized) - } - return c.Next() -} - -func (api gateway) list(c fiber.Ctx) error { - filters, _ := json.Marshal(map[string][]string{"label": {managedLabel + "=true"}}) - result, err := api.docker.request(http.MethodGet, "/containers/json?all=1&filters="+url.QueryEscape(string(filters)), nil) - if err != nil { - return writeError(c, http.StatusBadGateway, err) - } - defer result.Body.Close() - if result.StatusCode != http.StatusOK { - return forwardDockerError(c, result) - } - - var containers []struct { - ID string `json:"Id"` - State string `json:"State"` - Status string `json:"Status"` - Labels map[string]string `json:"Labels"` - } - if err := json.NewDecoder(result.Body).Decode(&containers); err != nil { - return writeError(c, http.StatusBadGateway, fmt.Errorf("decode Docker response: %w", err)) - } - - browsers := make([]browser, 0, len(containers)) - for _, container := range containers { - alias := container.Labels[idLabel] - if !runtimeIDPattern.MatchString(alias) { - continue - } - name := container.Labels[nameLabel] - if name == "" { - name = alias - } - bindingVersion, _ := strconv.ParseInt(container.Labels[bindingVersionLabel], 10, 64) - proxyPort, _ := strconv.Atoi(container.Labels[proxyPortLabel]) - direct := container.Labels[networkExitLabel] == "" - browsers = append(browsers, browser{ - ID: container.ID, - Alias: alias, - Name: name, - State: container.State, - Status: container.Status, - Endpoint: "http://" + namePrefix + alias + ":9222", - BindingVersion: bindingVersion, - NetworkExitID: container.Labels[networkExitLabel], - NetworkID: container.Labels[networkIDLabel], - ProxyReady: direct || api.proxies.ready(alias, proxyPort, container.ID, container.Labels[networkIDLabel]), - }) - } - return writeJSON(c, http.StatusOK, browsers) -} - -func (api gateway) create(c fiber.Ctx) error { - var input createRequest - decoder := json.NewDecoder(bytes.NewReader(c.Body())) - decoder.DisallowUnknownFields() - if err := decoder.Decode(&input); err != nil { - return writeError(c, http.StatusBadRequest, errors.New("body must contain only alias, name, image, cmd, volume, binding_version, network_exit_id, network_exit and stopped")) - } - if err := validateCreate(input); err != nil { - return writeError(c, http.StatusBadRequest, err) - } - if err := api.docker.pullIfMissing(c.Context(), input.Image); err != nil { - return writeError(c, http.StatusBadGateway, err) - } - if _, _, err := api.managedContainer(input.Alias); err == nil { - return writeError(c, http.StatusConflict, errors.New("browser alias is already in use")) - } else if !errors.Is(err, os.ErrNotExist) { - return writeError(c, statusFor(err), err) - } - _, release, err := api.locks.acquire(input.Alias) - if err != nil { - return writeError(c, statusFor(err), err) - } - defer release() - if _, _, err := api.managedContainer(input.Alias); err == nil { - return writeError(c, http.StatusConflict, errors.New("browser alias is already in use")) - } else if !errors.Is(err, os.ErrNotExist) { - return writeError(c, statusFor(err), err) - } - direct := input.NetworkExitID == "" - network, proxyServer, undoProxy := "none", "", func() {} - var networkGeneration tenantNetworkGeneration - keepNetwork := input.Stopped - if !input.Stopped { - var err error - var bindHost string - networkGeneration, bindHost, err = api.docker.ensureTenantNetwork(api.network, input.Alias, api.self, input.BindingVersion, "", "", false) - if networkGeneration.ID != "" { - defer func() { - if keepNetwork { - return - } - var cleanupErr error - if networkGeneration.Created { - cleanupErr = api.removeTenantNetwork(input.Alias, input.BindingVersion, "", networkGeneration, nil, "", false) - } else if networkGeneration.ConnectedSelf { - _, cleanupErr = api.disconnectTenantNetworkMember(input.Alias, input.BindingVersion, "", networkGeneration, - networkGeneration.SelfMember, nil, "", false) - } - if cleanupErr != nil { - logrus.WithError(cleanupErr).WithField("alias", input.Alias).Error("rollback isolated browser network") - } - }() - } - if err != nil { - return writeNetworkError(c, http.StatusBadGateway, errors.New("configure isolated browser network"), networkGeneration.ID) - } - network = networkGeneration.ID - if !direct { - proxyServer, undoProxy, err = api.proxies.configure(input.Alias, input.BindingVersion, bindHost, 0, input.NetworkExit, networkGeneration.ID) - if err != nil { - return writeNetworkError(c, statusFor(err), errors.Join(errors.New("configure in-memory proxy"), err), networkGeneration.ID) - } - } - } - keepProxy := false - defer func() { - if !keepProxy { - undoProxy() - } - }() - - pidsLimit := int64(512) - cmd := append([]string{}, input.Cmd...) - if !input.Stopped && !direct { - cmd = append(cmd[:len(cmd)-1], "--proxy-server="+proxyServer, "--disable-non-proxied-udp", cmd[len(cmd)-1]) - } - payload := map[string]any{ - "Image": input.Image, - "User": browserUser, - "Entrypoint": []string{browserEntrypoint}, - "Cmd": cmd, - "Env": []string{"REMOTE_DEBUGGING_PORT=9222"}, - "Labels": map[string]string{ - managedLabel: "true", - idLabel: input.Alias, - nameLabel: input.Name, - bindingVersionLabel: strconv.FormatInt(input.BindingVersion, 10), - networkExitLabel: input.NetworkExitID, - networkIDLabel: networkGeneration.ID, - proxyPortLabel: strconv.Itoa(proxyPort(proxyServer)), - }, - "ExposedPorts": map[string]any{"9222/tcp": map[string]any{}}, - "HostConfig": map[string]any{ - "NetworkMode": network, - "ReadonlyRootfs": true, - "CapDrop": []string{"ALL"}, - "SecurityOpt": []string{"no-new-privileges"}, - "PidsLimit": &pidsLimit, - "Memory": int64(1 << 30), - "NanoCpus": int64(2_000_000_000), - "Tmpfs": map[string]string{ - "/tmp": "rw,nosuid,nodev,noexec,mode=1777,size=256m", - "/tmp/.X11-unix": "rw,nosuid,nodev,noexec,mode=1777,size=1m", - "/dev/shm": "rw,nosuid,nodev,noexec,size=256m", - "/home/ubuntu": "rw,nosuid,nodev,noexec,uid=1000,gid=1000,mode=700,size=64m", - }, - "Mounts": []map[string]any{{ - "Type": "volume", - "Source": input.Volume, - "Target": "/data", - }}, - }, - } - result, err := api.docker.request(http.MethodPost, "/containers/create?name="+url.QueryEscape(namePrefix+input.Alias), payload) - var created struct { - ID string `json:"Id"` - } - status, createErr := http.StatusBadGateway, err - if result != nil { - if result.StatusCode == http.StatusConflict { - status = http.StatusConflict - } - if result.StatusCode == http.StatusCreated { - createErr = json.NewDecoder(result.Body).Decode(&created) - } else { - createErr = errors.New("Docker container creation failed") - } - result.Body.Close() - } - if createErr != nil || created.ID == "" { - containerID, labels, inspectErr := api.managedContainer(input.Alias) - if inspectErr == nil && labels[bindingVersionLabel] == strconv.FormatInt(input.BindingVersion, 10) && - labels[networkIDLabel] == networkGeneration.ID { - created.ID = containerID - } else { - if createErr == nil { - createErr = errors.New("Docker returned an invalid container id") - } - return writeNetworkError(c, status, createErr, networkGeneration.ID) - } - } - if !input.Stopped { - if !direct { - if !api.proxies.bind(input.Alias, input.BindingVersion, proxyServer, created.ID, networkGeneration.ID) { - cleanupErr := api.docker.expect(http.MethodDelete, "/containers/"+url.PathEscape(created.ID)+"?force=1&v=0", nil, http.StatusNoContent) - if cleanupErr != nil { - return writeNetworkError(c, http.StatusBadGateway, fmt.Errorf("proxy generation changed and container cleanup failed: %w", cleanupErr), networkGeneration.ID) - } - return writeNetworkError(c, http.StatusConflict, errGenerationConflict, networkGeneration.ID) - } - undoProxy = func() { api.proxies.remove(input.Alias, input.BindingVersion, created.ID) } - } - if err := api.docker.expect(http.MethodPost, "/containers/"+url.PathEscape(created.ID)+"/start", nil, http.StatusNoContent, http.StatusNotModified); err != nil { - cleanupErr := api.docker.expect(http.MethodDelete, "/containers/"+url.PathEscape(created.ID)+"?force=1&v=0", nil, http.StatusNoContent) - if cleanupErr != nil { - return writeNetworkError(c, http.StatusBadGateway, fmt.Errorf("container did not start and cleanup failed: %w; cleanup: %v", err, cleanupErr), networkGeneration.ID) - } - return writeNetworkError(c, http.StatusBadGateway, fmt.Errorf("container did not start and was removed while preserving its Profile volume: %w", err), networkGeneration.ID) - } - } - keepProxy, keepNetwork = !input.Stopped && !direct, true - return writeJSON(c, http.StatusCreated, map[string]string{"id": created.ID, "alias": input.Alias, "network_id": networkGeneration.ID}) -} - -func validateCreate(input createRequest) error { - if !runtimeIDPattern.MatchString(input.Alias) { - return errors.New("alias must match [a-z0-9][a-z0-9-]{0,31}") - } - if input.Name == "" || utf8.RuneCountInString(input.Name) > 64 || hasControlRunes(input.Name) { - return errors.New("name must be 1..64 visible characters") - } - if !imageRefPattern.MatchString(input.Image) { - return errors.New("image must be a valid image reference") - } - if !volumePattern.MatchString(input.Volume) { - return errors.New("volume must be a valid volume name") - } - direct := input.NetworkExitID == "" && input.NetworkExit == (gatewayProxyExit{}) - if input.BindingVersion < 1 || (input.Stopped && !direct) || - (!input.Stopped && !direct && !exitIDPattern.MatchString(input.NetworkExitID)) || - (input.NetworkExitID == "") != (input.NetworkExit == (gatewayProxyExit{})) { - return errors.New("binding_version and network_exit_id must identify the current binding") - } - if len(input.Cmd) == 0 || len(input.Cmd) > 64 || input.Cmd[len(input.Cmd)-1] != "about:blank" { - return errors.New("cmd must contain 1..64 arguments") - } - total := 0 - for _, arg := range input.Cmd { - if arg == "" || hasControlRunes(arg) { - return errors.New("cmd arguments must be non-empty visible strings") - } - if strings.HasPrefix(arg, "--proxy-server") || arg == "--disable-non-proxied-udp" { - return errors.New("proxy arguments are platform-controlled") - } - total += len(arg) - } - if total > 4096 { - return errors.New("cmd arguments exceed 4096 characters") - } - if input.Stopped || direct { - return nil - } - proxy := input.NetworkExit - if (proxy.Protocol != "http" && proxy.Protocol != "https" && proxy.Protocol != "socks4" && proxy.Protocol != "socks5") || - proxy.Host == "" || len(proxy.Host) > 253 || strings.ContainsAny(proxy.Host, "@/[]?# \t\r\n") || - proxy.Port < 1 || proxy.Port > 65535 || (proxy.Username == "" && proxy.Password != "") || - len(proxy.Username) > 255 || len(proxy.Password) > 255 || - hasControlRunes(proxy.Username) || hasControlRunes(proxy.Password) { - return errors.New("network_exit must contain a valid proxy endpoint and optional credentials") - } - return nil -} - -func proxyPort(proxyServer string) int { - parsed, _ := url.Parse(proxyServer) - port, _ := strconv.Atoi(parsed.Port()) - return port -} - -func hasControlRunes(value string) bool { - for _, r := range value { - if r < 0x20 || r == 0x7f { - return true - } - } - return false -} - -func (api gateway) changeState(c fiber.Ctx) error { - id := c.Params("id") - action := c.Params("action") - var input generationRequest - switch action { - case "start": - var err error - input, err = decodeGeneration(c) - if err != nil { - return writeError(c, http.StatusBadRequest, err) - } - if input.NetworkID == "" { - return writeError(c, http.StatusBadRequest, errors.New("network_id must identify the expected network generation")) - } - _, exists, err := api.requireGeneration(id, input) - if err != nil { - return writeError(c, statusFor(err), err) - } - if !exists { - return writeError(c, http.StatusNotFound, os.ErrNotExist) - } - case "stop": - var err error - input, err = decodeGeneration(c) - if err != nil { - return writeError(c, http.StatusBadRequest, err) - } - _, exists, err := api.requireGeneration(id, input) - if err != nil { - return writeError(c, statusFor(err), err) - } - if !exists { - return writeError(c, http.StatusNotFound, os.ErrNotExist) - } - default: - return writeError(c, http.StatusNotFound, errors.New("unknown action")) - } - _, release, err := api.locks.acquire(id) - if err != nil { - return writeError(c, statusFor(err), err) - } - defer release() - var path string - if action == "start" { - containerID, exists, err := api.requireGeneration(id, input) - if err != nil { - return writeError(c, statusFor(err), err) - } - if !exists { - return writeError(c, http.StatusNotFound, os.ErrNotExist) - } - path = "/containers/" + url.PathEscape(containerID) + "/start" - } else { - containerID, exists, err := api.requireGeneration(id, input) - if err != nil { - return writeError(c, statusFor(err), err) - } - if !exists { - return writeError(c, http.StatusNotFound, os.ErrNotExist) - } - path = "/containers/" + url.PathEscape(containerID) + "/stop?t=10" - } - if err := api.docker.expect(http.MethodPost, path, nil, http.StatusNoContent, http.StatusNotModified); err != nil { - return writeError(c, http.StatusBadGateway, err) - } - c.Status(http.StatusNoContent) - return nil -} - -func (api gateway) remove(c fiber.Ctx) error { - id := c.Params("id") - input, err := decodeGeneration(c) - if err != nil { - return writeError(c, http.StatusBadRequest, err) - } - containerID, exists, err := api.requireGeneration(id, input) - if err != nil { - return writeError(c, statusFor(err), err) - } - _, release, err := api.locks.acquire(id) - if err != nil { - return writeError(c, statusFor(err), err) - } - defer release() - currentID, currentExists, err := api.requireGeneration(id, input) - if err != nil { - return writeError(c, statusFor(err), err) - } - if currentExists != exists || currentID != containerID { - return writeError(c, http.StatusConflict, errGenerationConflict) - } - expectedNetworkID := input.NetworkID - if currentExists { - _, labels, inspectErr := api.managedContainer(id) - if inspectErr != nil { - return writeError(c, statusFor(inspectErr), inspectErr) - } - containerNetworkID := labels[networkIDLabel] - if containerNetworkID != input.NetworkID { - return writeError(c, http.StatusConflict, errGenerationConflict) - } - } else if expectedNetworkID == "" { - return writeError(c, http.StatusConflict, errGenerationConflict) - } - var networkGeneration tenantNetworkGeneration - networkExists := false - if expectedNetworkID != "" { - networkGeneration, _, networkExists, err = api.docker.inspectTenantNetwork(api.network, id, input.BindingVersion, - input.RuntimeID, api.self, expectedNetworkID, false) - if err != nil { - return writeError(c, statusFor(err), err) - } - } - if networkExists { - err = api.removeTenantNetwork(id, input.BindingVersion, input.RuntimeID, networkGeneration, &input, containerID, exists) - } else if expectedNetworkID != "" { - if exists { - // 容器仍在但登记的隔离网络已消失:代状态异常,fail-closed 交由人工对账。 - return writeError(c, http.StatusConflict, errGenerationConflict) - } - // 容器与登记网络均已不存在:请求指向的代在 Docker 侧已无残留,清理视为完成。 - // 否则被外部清理过的旧代会永远 409,控制面的 runtime_cleanup_pending 无法收敛。 - if !api.proxies.remove(id, input.BindingVersion, input.RuntimeID, input.NetworkID) { - return writeError(c, http.StatusConflict, errGenerationConflict) - } - return c.SendStatus(http.StatusNoContent) - } - if err != nil { - if errors.Is(err, errGenerationConflict) { - return writeError(c, http.StatusConflict, err) - } - return c.Status(http.StatusAccepted).JSON(map[string]string{ - "status": "runtime_cleanup_pending", - }) - } - if !api.proxies.remove(id, input.BindingVersion, input.RuntimeID, input.NetworkID) { - return writeError(c, http.StatusConflict, errGenerationConflict) - } - if exists { - path := "/containers/" + url.PathEscape(containerID) + "?force=1&v=0" - if err := api.requireRuntimeState(id, input, containerID, true); err != nil { - return writeError(c, http.StatusConflict, err) - } - if err := api.docker.expect(http.MethodDelete, path, nil, http.StatusNoContent); err != nil { - currentID, currentExists, generationErr := api.requireGeneration(id, input) - if errors.Is(err, os.ErrNotExist) && generationErr == nil && !currentExists { - return c.SendStatus(http.StatusNoContent) - } - if generationErr != nil || currentID != containerID { - err = errGenerationConflict - } - if errors.Is(err, errGenerationConflict) { - return writeError(c, http.StatusConflict, err) - } - return writeError(c, http.StatusBadGateway, err) - } - } - c.Status(http.StatusNoContent) - return nil -} - -func (api gateway) restoreProxy(c fiber.Ctx) error { - input := proxyRestoreRequest{} - decoder := json.NewDecoder(bytes.NewReader(c.Body())) - decoder.DisallowUnknownFields() - if err := decoder.Decode(&input); err != nil || input.BindingVersion < 1 || !exitIDPattern.MatchString(input.RuntimeID) || - !exitIDPattern.MatchString(input.NetworkID) || - !exitIDPattern.MatchString(input.NetworkExitID) || - validateCreate(createRequest{Alias: c.Params("id"), Name: "x", Image: "x", Cmd: []string{"about:blank"}, Volume: "x", - BindingVersion: input.BindingVersion, NetworkExitID: input.NetworkExitID, NetworkExit: input.NetworkExit}) != nil { - return writeError(c, http.StatusBadRequest, errors.New("invalid proxy recovery request")) - } - removeStaleProxy := func() { api.proxies.remove(c.Params("id"), input.BindingVersion, input.RuntimeID, input.NetworkID) } - runtimeID, labels, err := api.requireProxyGeneration(c.Params("id"), input) - if err != nil { - removeStaleProxy() - return writeError(c, statusFor(err), err) - } - port, _ := strconv.Atoi(labels[proxyPortLabel]) - if port < 1 { - removeStaleProxy() - return writeError(c, http.StatusConflict, errors.New("container binding does not match recovery request")) - } - _, release, err := api.locks.acquire(c.Params("id")) - if err != nil { - return writeError(c, statusFor(err), err) - } - defer release() - if runtimeID, labels, err = api.requireProxyGeneration(c.Params("id"), input); err != nil { - removeStaleProxy() - return writeError(c, statusFor(err), err) - } - port, _ = strconv.Atoi(labels[proxyPortLabel]) - if port < 1 { - removeStaleProxy() - return writeError(c, http.StatusConflict, errors.New("container binding does not match recovery request")) - } - networkGeneration, bindHost, err := api.docker.ensureTenantNetwork(api.network, c.Params("id"), api.self, input.BindingVersion, - input.RuntimeID, labels[networkIDLabel], true) - keepNetwork := false - if networkGeneration.ID != "" { - defer func() { - if keepNetwork { - return - } - var cleanupErr error - if networkGeneration.Created { - cleanupErr = api.removeTenantNetwork(c.Params("id"), input.BindingVersion, input.RuntimeID, networkGeneration, nil, "", false) - } else { - if networkGeneration.ConnectedRuntime { - _, cleanupErr = api.disconnectTenantNetworkMember(c.Params("id"), input.BindingVersion, input.RuntimeID, - networkGeneration, input.RuntimeID, nil, "", false) - } - if cleanupErr == nil && networkGeneration.ConnectedSelf { - _, cleanupErr = api.disconnectTenantNetworkMember(c.Params("id"), input.BindingVersion, input.RuntimeID, - networkGeneration, networkGeneration.SelfMember, nil, "", false) - } - } - if cleanupErr != nil { - logrus.WithError(cleanupErr).WithField("alias", c.Params("id")).Error("rollback restored browser network") - } - }() - } - if err != nil { - removeStaleProxy() - return writeError(c, http.StatusBadGateway, errors.New("restore isolated browser network")) - } - if err = api.requireProxyNetworkGeneration(c.Params("id"), input, networkGeneration, bindHost); err != nil { - removeStaleProxy() - return writeError(c, statusFor(err), err) - } - proxyServer, undoProxy, err := api.proxies.configure(c.Params("id"), input.BindingVersion, bindHost, port, input.NetworkExit, input.NetworkID) - if err != nil { - removeStaleProxy() - return writeError(c, statusFor(err), errors.Join(errors.New("restore in-memory proxy"), err)) - } - if err = api.requireProxyNetworkGeneration(c.Params("id"), input, networkGeneration, bindHost); err != nil { - undoProxy() - return writeError(c, statusFor(err), err) - } - if !api.proxies.bind(c.Params("id"), input.BindingVersion, proxyServer, runtimeID, input.NetworkID) { - undoProxy() - return writeError(c, http.StatusConflict, errGenerationConflict) - } - if err = api.requireProxyNetworkGeneration(c.Params("id"), input, networkGeneration, bindHost); err != nil { - undoProxy() - return writeError(c, statusFor(err), err) - } - keepNetwork = true - return c.SendStatus(http.StatusNoContent) -} - -func (api gateway) requireProxyNetworkGeneration(alias string, input proxyRestoreRequest, expected tenantNetworkGeneration, bindHost string) error { - _, labels, err := api.requireProxyGeneration(alias, input) - if err != nil { - return err - } - if networkID := labels[networkIDLabel]; networkID != "" && networkID != expected.ID { - return errGenerationConflict - } - current, addresses, exists, err := api.docker.inspectTenantNetwork(api.network, alias, input.BindingVersion, - input.RuntimeID, api.self, expected.ID, false) - host, _, _ := net.ParseCIDR(addresses[current.SelfMember]) - if err != nil || !exists || !sameTenantNetworkMembers(current, expected) || host == nil || host.String() != bindHost { - return errGenerationConflict - } - return nil -} - -func sameTenantNetworkMembers(current, expected tenantNetworkGeneration) bool { - if current.ID != expected.ID || current.Name != expected.Name || current.RuntimeAttached != expected.RuntimeAttached || - current.SelfMember != expected.SelfMember || len(current.GatewayMembers) != len(expected.GatewayMembers) { - return false - } - members := make(map[string]struct{}, len(current.GatewayMembers)) - for _, member := range current.GatewayMembers { - members[member] = struct{}{} - } - for _, member := range expected.GatewayMembers { - if _, ok := members[member]; !ok { - return false - } - } - return true -} - -func (api gateway) requireProxyGeneration(alias string, input proxyRestoreRequest) (string, map[string]string, error) { - runtimeID, labels, err := api.managedContainer(alias) - if err != nil { - return "", nil, err - } - version, _ := strconv.ParseInt(labels[bindingVersionLabel], 10, 64) - if runtimeID != input.RuntimeID || version != input.BindingVersion || labels[networkExitLabel] != input.NetworkExitID || - labels[networkIDLabel] != input.NetworkID { - return "", nil, errGenerationConflict - } - return runtimeID, labels, nil -} - -func decodeGeneration(c fiber.Ctx) (generationRequest, error) { - var input generationRequest - decoder := json.NewDecoder(bytes.NewReader(c.Body())) - decoder.DisallowUnknownFields() - if err := decoder.Decode(&input); err != nil || input.BindingVersion < 1 || - (input.RuntimeID != "" && !exitIDPattern.MatchString(input.RuntimeID)) || - (input.NetworkID != "" && !exitIDPattern.MatchString(input.NetworkID)) { - return generationRequest{}, errors.New("binding_version, runtime_id and network_id must identify the expected generation") - } - return input, nil -} - -func (api gateway) requireGeneration(id string, input generationRequest) (string, bool, error) { - runtimeID, labels, err := api.managedContainer(id) - if errors.Is(err, os.ErrNotExist) { - return "", false, nil - } - if err != nil { - return "", false, err - } - version, _ := strconv.ParseInt(labels[bindingVersionLabel], 10, 64) - if input.RuntimeID == "" || runtimeID != input.RuntimeID || version != input.BindingVersion || labels[networkIDLabel] != input.NetworkID { - logrus.WithFields(logrus.Fields{ - "browser_id": id, "input_runtime_id": input.RuntimeID, "actual_runtime_id": runtimeID, - "input_binding_version": input.BindingVersion, "actual_binding_version": version, - "input_network_id": input.NetworkID, "actual_network_id": labels[networkIDLabel], - }).Warn("browser generation mismatch") - return "", false, errGenerationConflict - } - return runtimeID, true, nil -} - -func (api gateway) managedContainer(id string) (string, map[string]string, error) { - runtimeID, labels, _, err := api.managedContainerState(id) - return runtimeID, labels, err -} - -func (api gateway) managedContainerState(id string) (string, map[string]string, map[string]string, error) { - if !runtimeIDPattern.MatchString(id) { - return "", nil, nil, errInvalidRuntimeID - } - result, err := api.docker.request(http.MethodGet, "/containers/"+url.PathEscape(namePrefix+id)+"/json", nil) - if err != nil { - return "", nil, nil, err - } - defer result.Body.Close() - if result.StatusCode == http.StatusNotFound { - return "", nil, nil, os.ErrNotExist - } - if result.StatusCode != http.StatusOK { - return "", nil, nil, fmt.Errorf("Docker inspect returned %s", result.Status) - } - var inspected struct { - ID string `json:"Id"` - Config struct { - Labels map[string]string `json:"Labels"` - } `json:"Config"` - NetworkSettings struct { - Networks map[string]struct { - NetworkID string `json:"NetworkID"` - } `json:"Networks"` - } `json:"NetworkSettings"` - } - if err := json.NewDecoder(result.Body).Decode(&inspected); err != nil { - return "", nil, nil, fmt.Errorf("decode Docker inspect: %w", err) - } - if inspected.Config.Labels[managedLabel] != "true" || inspected.Config.Labels[idLabel] != id { - return "", nil, nil, errUnmanagedContainer - } - networks := make(map[string]string, len(inspected.NetworkSettings.Networks)) - for name, network := range inspected.NetworkSettings.Networks { - networks[name] = network.NetworkID - } - return inspected.ID, inspected.Config.Labels, networks, nil -} - -// pullIfMissing 在镜像不在本地时从远端仓库拉取;镜像缺失属于可恢复错误,调用方可直接重试。 -func (docker dockerClient) pullIfMissing(ctx context.Context, ref string) error { - inspect, err := docker.request(http.MethodGet, "/images/"+url.PathEscape(ref)+"/json", nil) - if err != nil { - return fmt.Errorf("inspect image %s: %w", ref, err) - } - _, _ = io.Copy(io.Discard, inspect.Body) - _ = inspect.Body.Close() - switch inspect.StatusCode { - case http.StatusOK: - return nil - case http.StatusNotFound: - // 本地无此镜像,继续拉取 - default: - return fmt.Errorf("inspect image %s returned %s", ref, inspect.Status) - } - - pullCtx, cancel := context.WithTimeout(ctx, pullTimeout) - defer cancel() - query := url.Values{"fromImage": {ref}} - if !strings.Contains(ref, "@") { - if repository, tag := splitImageRef(ref); tag != "" { - query = url.Values{"fromImage": {repository}, "tag": {tag}} - } - } - request, err := http.NewRequestWithContext(pullCtx, http.MethodPost, docker.baseURL+"/images/create?"+query.Encode(), nil) - if err != nil { - return fmt.Errorf("build image pull request: %w", err) - } - response, err := docker.slow.Do(request) - if err != nil { - return fmt.Errorf("pull image %s: %w", ref, err) - } - defer response.Body.Close() - if response.StatusCode != http.StatusOK { - message, _ := io.ReadAll(io.LimitReader(response.Body, 4096)) - return fmt.Errorf("pull image %s returned %s: %s", ref, response.Status, strings.TrimSpace(string(message))) - } - _, _ = io.Copy(io.Discard, response.Body) - return nil -} - -func splitImageRef(ref string) (repository, tag string) { - if at := strings.Index(ref, "@"); at >= 0 { - return ref[:at], ref[at+1:] - } - if colon := strings.LastIndex(ref, ":"); colon > strings.LastIndex(ref, "/") { - return ref[:colon], ref[colon+1:] - } - return ref, "" -} - -func (docker dockerClient) request(method, path string, payload any) (*http.Response, error) { - var body io.Reader - if payload != nil { - encoded, err := json.Marshal(payload) - if err != nil { - return nil, err - } - body = bytes.NewReader(encoded) - } - request, err := http.NewRequest(method, docker.baseURL+path, body) - if err != nil { - return nil, err - } - if payload != nil { - request.Header.Set("Content-Type", "application/json") - } - return docker.client.Do(request) -} - -func (docker dockerClient) expect(method, path string, payload any, allowed ...int) error { - response, err := docker.request(method, path, payload) - if err != nil { - return err - } - defer response.Body.Close() - for _, status := range allowed { - if response.StatusCode == status { - return nil - } - } - if response.StatusCode == http.StatusNotFound { - return os.ErrNotExist - } - message, _ := io.ReadAll(io.LimitReader(response.Body, 4096)) - return fmt.Errorf("Docker returned %s: %s", response.Status, strings.TrimSpace(string(message))) -} - -func (locks *dockerAliasReservations) acquire(alias string) (string, func(), error) { - inspected, err := locks.docker.request(http.MethodGet, "/containers/"+url.PathEscape(locks.self)+"/json", nil) - if err != nil { - return "", nil, err - } - var gatewayContainer struct { - Image string `json:"Image"` - Config struct { - Labels map[string]string `json:"Labels"` - } `json:"Config"` - } - if inspected.StatusCode != http.StatusOK || json.NewDecoder(inspected.Body).Decode(&gatewayContainer) != nil || gatewayContainer.Image == "" || - gatewayContainer.Config.Labels[gatewayMemberLabel] != "true" { - inspected.Body.Close() - return "", nil, errors.New("inspect trusted gateway for alias reservation") - } - inspected.Body.Close() - generationBytes := make([]byte, 16) - if _, err := rand.Read(generationBytes); err != nil { - return "", nil, errors.New("create alias reservation generation") - } - generation := hex.EncodeToString(generationBytes) - - response, err := locks.docker.request(http.MethodPost, "/containers/create?name="+url.QueryEscape(reservationPrefix+alias), map[string]any{ - "Image": gatewayContainer.Image, - "Labels": map[string]string{reservationLabel: "true", idLabel: alias, reservationGenLabel: generation}, - "HostConfig": map[string]any{"NetworkMode": "none"}, - }) - if err != nil { - return "", nil, locks.reconcile(alias, generation, "", fmt.Errorf("create alias reservation result unknown: %w", err)) - } - defer response.Body.Close() - if response.StatusCode == http.StatusConflict { - return "", nil, errGenerationConflict - } - if response.StatusCode != http.StatusCreated { - return "", nil, locks.reconcile(alias, generation, "", fmt.Errorf("create alias reservation returned %s", response.Status)) - } - var created struct { - ID string `json:"Id"` - } - if json.NewDecoder(response.Body).Decode(&created) != nil || created.ID == "" { - return "", nil, locks.reconcile(alias, generation, "", errors.New("Docker returned an invalid alias reservation id")) - } - actualID, inspectErr := locks.inspect(alias, generation) - if inspectErr != nil { - return "", nil, fmt.Errorf("Docker returned an unverified alias reservation id; generation %s requires manual reconcile: %w", generation, inspectErr) - } - if actualID != created.ID { - return "", nil, fmt.Errorf("Docker returned an alias reservation id that conflicts with generation %s; manual reconcile required", generation) - } - return created.ID, func() { - if err := locks.remove(alias, generation, created.ID); err != nil { - logrus.WithError(err).WithFields(logrus.Fields{"alias": alias, "reservation_id": created.ID, - "reservation_generation": generation}).Error("alias reservation cleanup pending; manual reconcile required") - } - }, nil -} - -func (locks *dockerAliasReservations) reconcile(alias, generation, expectedID string, cause error) error { - actualID, inspectErr := locks.inspect(alias, generation) - if inspectErr != nil { - return fmt.Errorf("%w; reservation generation %s requires manual reconcile: %v", cause, generation, inspectErr) - } - if expectedID != "" && actualID != expectedID { - return fmt.Errorf("%w; reservation generation %s conflicts with immutable id", cause, generation) - } - if cleanupErr := locks.remove(alias, generation, actualID); cleanupErr != nil { - return fmt.Errorf("%w; reservation %s generation %s cleanup pending: %v", cause, actualID, generation, cleanupErr) - } - return fmt.Errorf("%w; reservation %s was removed", cause, actualID) -} - -func (locks *dockerAliasReservations) remove(alias, generation, expectedID string) error { - var deleteErr error - for attempt := 0; attempt < 2; attempt++ { - actualID, inspectErr := locks.inspect(alias, generation) - if errors.Is(inspectErr, os.ErrNotExist) { - return nil - } - if inspectErr != nil { - return inspectErr - } - if actualID != expectedID { - return errors.New("reservation generation conflicts with immutable id") - } - deleteErr = locks.docker.expect(http.MethodDelete, "/containers/"+url.PathEscape(expectedID)+"?force=1&v=0", nil, - http.StatusNoContent) - if deleteErr == nil { - return nil - } - if errors.Is(deleteErr, os.ErrNotExist) { - _, confirmErr := locks.inspect(alias, generation) - if errors.Is(confirmErr, os.ErrNotExist) { - return nil - } - return errors.Join(deleteErr, confirmErr) - } - } - return deleteErr -} - -func (locks *dockerAliasReservations) inspect(alias, generation string) (string, error) { - response, err := locks.docker.request(http.MethodGet, "/containers/"+url.PathEscape(reservationPrefix+alias)+"/json", nil) - if err != nil { - return "", err - } - defer response.Body.Close() - if response.StatusCode == http.StatusNotFound { - return "", os.ErrNotExist - } - var reservation struct { - ID string `json:"Id"` - Config struct { - Labels map[string]string `json:"Labels"` - } `json:"Config"` - } - if response.StatusCode != http.StatusOK || json.NewDecoder(response.Body).Decode(&reservation) != nil || reservation.ID == "" || - reservation.Config.Labels[reservationLabel] != "true" || reservation.Config.Labels[idLabel] != alias || - reservation.Config.Labels[reservationGenLabel] != generation { - return "", errors.New("reservation name does not identify the expected alias generation") - } - return reservation.ID, nil -} - -func tenantNetworkName(base, alias string) (string, error) { - name := base + "-" + alias - if !networkNamePattern.MatchString(name) { - return "", errors.New("isolated browser network name is invalid") - } - return name, nil -} - -func sameContainerReference(id, name, reference string) bool { - if reference == "" { - return false - } - return id == reference || name == reference || strings.HasPrefix(id, reference) || strings.HasPrefix(reference, id) -} - -func (docker dockerClient) trustedGatewayMember(id string) bool { - response, err := docker.request(http.MethodGet, "/containers/"+url.PathEscape(id)+"/json", nil) - if err != nil { - return false - } - defer response.Body.Close() - var container struct { - ID string `json:"Id"` - Config struct { - Labels map[string]string `json:"Labels"` - } `json:"Config"` - } - return response.StatusCode == http.StatusOK && json.NewDecoder(response.Body).Decode(&container) == nil && - container.ID != "" && container.Config.Labels[gatewayMemberLabel] == "true" -} - -func (docker dockerClient) inspectTenantNetwork(base, alias string, bindingVersion int64, runtimeID, self, expectedID string, - allowUnversioned bool) (tenantNetworkGeneration, map[string]string, bool, error) { - name, err := tenantNetworkName(base, alias) - if err != nil { - return tenantNetworkGeneration{}, nil, false, err - } - generation := tenantNetworkGeneration{ID: expectedID, Name: name} - reference := name - if expectedID != "" { - reference = expectedID - } - response, err := docker.request(http.MethodGet, "/networks/"+url.PathEscape(reference), nil) - if err != nil { - return generation, nil, false, err - } - defer response.Body.Close() - if response.StatusCode == http.StatusNotFound { - return generation, nil, false, nil - } - var network dockerTenantNetwork - if response.StatusCode != http.StatusOK || json.NewDecoder(response.Body).Decode(&network) != nil || network.ID == "" || - network.Name != name || network.Driver != "bridge" || network.Internal || network.Attachable || network.Ingress || - network.Labels[managedLabel] != "true" || network.Labels[networkRoleLabel] != browserNetworkRole || network.Labels[idLabel] != alias { - return generation, nil, false, errors.New("refusing to operate on an unowned browser network") - } - if expectedID != "" && network.ID != expectedID { - return generation, nil, false, errGenerationConflict - } - networkVersion := network.Labels[bindingVersionLabel] - if networkVersion != strconv.FormatInt(bindingVersion, 10) && !(allowUnversioned && networkVersion == "") { - return generation, nil, false, errGenerationConflict - } - generation.ID = network.ID - addresses := make(map[string]string, len(network.Containers)) - for id, member := range network.Containers { - addresses[id] = member.IPv4Address - if id == runtimeID { - generation.RuntimeAttached = true - continue - } - if !docker.trustedGatewayMember(id) { - return generation, nil, false, errGenerationConflict - } - generation.GatewayMembers = append(generation.GatewayMembers, id) - if sameContainerReference(id, member.Name, self) { - generation.SelfMember = id - } - } - return generation, addresses, true, nil -} - -func (docker dockerClient) ensureTenantNetwork(base, alias, self string, bindingVersion int64, runtimeID, expectedID string, - allowUnversioned bool) (tenantNetworkGeneration, string, error) { - if self == "" { - return tenantNetworkGeneration{}, "", errors.New("isolated browser network identity is invalid") - } - generation, addresses, exists, err := docker.inspectTenantNetwork(base, alias, bindingVersion, runtimeID, self, expectedID, allowUnversioned) - if err != nil { - return generation, "", err - } - if !exists { - if expectedID != "" { - return generation, "", errGenerationConflict - } - response, createErr := docker.request(http.MethodPost, "/networks/create", map[string]any{ - "Name": generation.Name, "CheckDuplicate": true, "Driver": "bridge", - "Labels": map[string]string{managedLabel: "true", networkRoleLabel: browserNetworkRole, idLabel: alias, - bindingVersionLabel: strconv.FormatInt(bindingVersion, 10)}, - }) - if createErr != nil { - return generation, "", fmt.Errorf("create isolated browser network result unknown; immutable generation requires manual reconcile: %w", createErr) - } - var created struct { - ID string `json:"Id"` - } - if response.StatusCode != http.StatusCreated || json.NewDecoder(response.Body).Decode(&created) != nil || created.ID == "" { - response.Body.Close() - return generation, "", errors.New("create isolated browser network result has no immutable id; manual reconcile required") - } - response.Body.Close() - generation.ID, generation.Created = created.ID, true - observed, currentAddresses, observedExists, inspectErr := docker.inspectTenantNetwork(base, alias, bindingVersion, runtimeID, self, generation.ID, false) - generation = preserveTenantNetworkGeneration(generation, observed) - addresses, err = currentAddresses, inspectErr - if err != nil || !observedExists { - if err == nil { - err = errGenerationConflict - } - return generation, "", err - } - } - if runtimeID != "" && !generation.RuntimeAttached { - generation.ConnectedRuntime = true - if err := docker.expect(http.MethodPost, "/networks/"+url.PathEscape(generation.ID)+"/connect", map[string]any{ - "Container": runtimeID, - }, http.StatusOK); err != nil { - observed, _, _, _ := docker.inspectTenantNetwork(base, alias, bindingVersion, runtimeID, self, generation.ID, allowUnversioned) - generation = preserveTenantNetworkGeneration(generation, observed) - return generation, "", err - } - generation.RuntimeAttached = true - } - if generation.SelfMember == "" { - generation.ConnectedSelf = true - if err := docker.expect(http.MethodPost, "/networks/"+url.PathEscape(generation.ID)+"/connect", map[string]any{ - "Container": self, "EndpointConfig": map[string]any{"Aliases": []string{browserProxyHost}}, - }, http.StatusOK); err != nil { - observed, _, _, _ := docker.inspectTenantNetwork(base, alias, bindingVersion, runtimeID, self, generation.ID, allowUnversioned) - generation = preserveTenantNetworkGeneration(generation, observed) - return generation, "", err - } - } - observed, currentAddresses, observedExists, inspectErr := docker.inspectTenantNetwork(base, alias, bindingVersion, runtimeID, self, generation.ID, allowUnversioned) - generation = preserveTenantNetworkGeneration(generation, observed) - addresses, err = currentAddresses, inspectErr - if err != nil || !observedExists || generation.SelfMember == "" { - if err == nil && !observedExists { - err = errGenerationConflict - } - return generation, "", errors.Join(err, errors.New("Docker did not connect the gateway to the isolated network")) - } - host, _, _ := net.ParseCIDR(addresses[generation.SelfMember]) - if host == nil { - return generation, "", errors.New("Docker did not assign the gateway an isolated network address") - } - return generation, host.String(), nil -} - -func preserveTenantNetworkGeneration(known, observed tenantNetworkGeneration) tenantNetworkGeneration { - if observed.ID == "" { - observed.ID = known.ID - } - if observed.Name == "" { - observed.Name = known.Name - } - observed.Created = observed.Created || known.Created - observed.RuntimeAttached = observed.RuntimeAttached || known.RuntimeAttached - observed.ConnectedRuntime = observed.ConnectedRuntime || known.ConnectedRuntime - observed.ConnectedSelf = observed.ConnectedSelf || known.ConnectedSelf - if observed.SelfMember == "" { - observed.SelfMember = known.SelfMember - } - for _, member := range known.GatewayMembers { - if !memberPresent(observed, member, "") { - observed.GatewayMembers = append(observed.GatewayMembers, member) - } - } - return observed -} - -func (api gateway) requireRuntimeState(alias string, input generationRequest, expectedID string, expectedExists bool) error { - runtimeID, exists, err := api.requireGeneration(alias, input) - if err != nil || exists != expectedExists || runtimeID != expectedID { - return errGenerationConflict - } - return nil -} - -func memberPresent(generation tenantNetworkGeneration, id, runtimeID string) bool { - if generation.RuntimeAttached && runtimeID != "" && sameContainerReference(id, "", runtimeID) { - return true - } - for _, member := range generation.GatewayMembers { - if sameContainerReference(member, "", id) { - return true - } - } - return false -} - -func generationWithoutMember(generation tenantNetworkGeneration, member, runtimeID string) tenantNetworkGeneration { - if generation.RuntimeAttached && runtimeID != "" && sameContainerReference(member, "", runtimeID) { - generation.RuntimeAttached = false - } - if sameContainerReference(member, "", generation.SelfMember) { - generation.SelfMember = "" - } - members := generation.GatewayMembers[:0:0] - for _, current := range generation.GatewayMembers { - if !sameContainerReference(current, "", member) { - members = append(members, current) - } - } - generation.GatewayMembers = members - return generation -} - -func (api gateway) disconnectTenantNetworkMember(alias string, bindingVersion int64, runtimeID string, - generation tenantNetworkGeneration, member string, input *generationRequest, expectedRuntime string, expectedExists bool) (tenantNetworkGeneration, error) { - if input != nil { - if err := api.requireRuntimeState(alias, *input, expectedRuntime, expectedExists); err != nil { - return generation, err - } - } - current, _, exists, err := api.docker.inspectTenantNetwork(api.network, alias, bindingVersion, runtimeID, api.self, generation.ID, false) - if err != nil || !exists { - return current, err - } - if !sameTenantNetworkMembers(current, generation) { - return current, errGenerationConflict - } - if !memberPresent(current, member, runtimeID) { - return current, nil - } - err = api.docker.expect(http.MethodPost, "/networks/"+url.PathEscape(generation.ID)+"/disconnect", map[string]any{ - "Container": member, "Force": true, - }, http.StatusOK) - current, _, exists, fenceErr := api.docker.inspectTenantNetwork(api.network, alias, bindingVersion, runtimeID, api.self, generation.ID, false) - if fenceErr != nil || !exists { - return current, errGenerationConflict - } - if sameTenantNetworkMembers(current, generationWithoutMember(generation, member, runtimeID)) { - return current, nil - } - if err != nil { - if sameTenantNetworkMembers(current, generation) { - return current, err - } - return current, errGenerationConflict - } - if !sameTenantNetworkMembers(current, generation) { - return current, errGenerationConflict - } - return current, errors.New("Docker retained an isolated network member after disconnect") -} - -func (api gateway) deleteTenantNetwork(alias string, bindingVersion int64, runtimeID string, - generation tenantNetworkGeneration, input *generationRequest, expectedRuntime string, expectedExists bool) error { - if input != nil { - if err := api.requireRuntimeState(alias, *input, expectedRuntime, expectedExists); err != nil { - return err - } - } - current, _, exists, err := api.docker.inspectTenantNetwork(api.network, alias, bindingVersion, runtimeID, api.self, generation.ID, false) - if err != nil { - return err - } - if !exists { - return nil - } - if current.RuntimeAttached || len(current.GatewayMembers) != 0 { - return errGenerationConflict - } - err = api.docker.expect(http.MethodDelete, "/networks/"+url.PathEscape(generation.ID), nil, http.StatusNoContent) - _, _, exists, fenceErr := api.docker.inspectTenantNetwork(api.network, alias, bindingVersion, runtimeID, api.self, generation.ID, false) - if fenceErr != nil { - return errGenerationConflict - } - if !exists { - return nil - } - if err != nil { - return err - } - return errors.New("Docker retained the isolated browser network after delete") -} - -func (api gateway) removeTenantNetwork(alias string, bindingVersion int64, runtimeID string, - generation tenantNetworkGeneration, input *generationRequest, expectedRuntime string, expectedExists bool) error { - if generation.RuntimeAttached || generation.ConnectedRuntime { - var err error - generation, err = api.disconnectTenantNetworkMember(alias, bindingVersion, runtimeID, generation, runtimeID, - input, expectedRuntime, expectedExists) - if err != nil { - return fmt.Errorf("disconnect browser from isolated network: %w", err) - } - } - for len(generation.GatewayMembers) > 0 { - gatewayID := generation.GatewayMembers[0] - var err error - generation, err = api.disconnectTenantNetworkMember(alias, bindingVersion, runtimeID, generation, gatewayID, - input, expectedRuntime, expectedExists) - if err != nil { - return fmt.Errorf("disconnect trusted gateway from isolated network: %w", err) - } - } - if err := api.deleteTenantNetwork(alias, bindingVersion, runtimeID, generation, input, expectedRuntime, expectedExists); err != nil { - return fmt.Errorf("remove isolated browser network: %w", err) - } - return nil -} - -func validateBrowserNetwork(name string) error { - if !networkNamePattern.MatchString(name) || len(name) > 31 { - return errors.New("BROWSER_NETWORK must be a valid network prefix of at most 31 characters") - } - if name == controlNetworkName { - return errors.New("BROWSER_NETWORK must not reuse the control network") - } - return nil -} - -func statusFor(err error) int { - switch { - case errors.Is(err, errInvalidRuntimeID): - return http.StatusBadRequest - case errors.Is(err, os.ErrNotExist): - return http.StatusNotFound - case errors.Is(err, errUnmanagedContainer): - return http.StatusForbidden - case errors.Is(err, errGenerationConflict): - return http.StatusConflict - default: - return http.StatusBadGateway - } -} - -func forwardDockerError(c fiber.Ctx, result *http.Response) error { - message, _ := io.ReadAll(io.LimitReader(result.Body, 4096)) - status := http.StatusBadGateway - if result.StatusCode == http.StatusConflict { - status = http.StatusConflict - } - return writeError(c, status, fmt.Errorf("Docker returned %s: %s", result.Status, strings.TrimSpace(string(message)))) -} - -func writeError(c fiber.Ctx, status int, err error) error { - return writeJSON(c, status, map[string]string{"error": err.Error()}) -} - -func writeNetworkError(c fiber.Ctx, status int, err error, networkID string) error { - return writeJSON(c, status, map[string]string{"error": err.Error(), "network_id": networkID}) -} - -func writeJSON(c fiber.Ctx, status int, value any) error { - return c.Status(status).JSON(value) -} diff --git a/cmd/docker-gateway/main_test.go b/cmd/docker-gateway/main_test.go deleted file mode 100644 index 308b52f..0000000 --- a/cmd/docker-gateway/main_test.go +++ /dev/null @@ -1,2459 +0,0 @@ -package main - -import ( - "bytes" - "encoding/json" - "fmt" - "io" - "net" - "net/http" - "net/http/httptest" - "os" - "strconv" - "strings" - "sync" - "testing" - "time" - - "github.com/gofiber/fiber/v3" - "github.com/gofiber/fiber/v3/middleware/adaptor" -) - -const testToken = "unit-test-gateway-token" - -func authed(method, target string, body io.Reader) *http.Request { - request := httptest.NewRequest(method, target, body) - request.Header.Set("Authorization", "Bearer "+testToken) - return request -} - -func testDocker(handler http.HandlerFunc) (dockerClient, *httptest.Server) { - self, _ := os.Hostname() - var networkMu sync.Mutex - networkMembers := map[string]map[string]bool{} - networkDeleted := map[string]bool{} - server := httptest.NewServer(withAliasReservations(self, func(response http.ResponseWriter, request *http.Request) { - if strings.HasPrefix(request.URL.Path, "/networks/network-") { - alias := strings.TrimPrefix(request.URL.Path, "/networks/network-") - alias = strings.TrimSuffix(strings.TrimSuffix(alias, "/disconnect"), "/connect") - networkMu.Lock() - if request.Method == http.MethodGet { - if networkDeleted[alias] { - networkMu.Unlock() - response.WriteHeader(http.StatusNotFound) - return - } - members := map[string]any{} - for id := range networkMembers[alias] { - members[id] = map[string]string{"Name": id, "IPv4Address": "127.0.0.1/8"} - } - networkMu.Unlock() - _ = json.NewEncoder(response).Encode(map[string]any{ - "Id": "network-" + alias, "Name": "creatorhub_browser-" + alias, "Driver": "bridge", "Internal": false, "Attachable": false, "Ingress": false, - "Labels": map[string]string{managedLabel: "true", networkRoleLabel: browserNetworkRole, idLabel: alias, bindingVersionLabel: "1"}, - "Containers": members, - }) - return - } - if strings.HasSuffix(request.URL.Path, "/disconnect") { - var body struct { - Container string `json:"Container"` - } - _ = json.NewDecoder(request.Body).Decode(&body) - delete(networkMembers[alias], body.Container) - } else if strings.HasSuffix(request.URL.Path, "/connect") { - var body struct { - Container string `json:"Container"` - } - _ = json.NewDecoder(request.Body).Decode(&body) - if networkMembers[alias] == nil { - networkMembers[alias] = map[string]bool{} - } - networkMembers[alias][body.Container] = true - } else if request.Method == http.MethodDelete { - networkDeleted[alias] = true - } - networkMu.Unlock() - if request.Method == http.MethodDelete { - response.WriteHeader(http.StatusNoContent) - } else { - response.WriteHeader(http.StatusOK) - } - return - } - if strings.HasPrefix(request.URL.Path, "/networks/creatorhub_browser-") { - if request.Method != http.MethodGet { - if request.Method == http.MethodDelete { - response.WriteHeader(http.StatusNoContent) - } else { - response.WriteHeader(http.StatusOK) - } - return - } - alias := strings.TrimPrefix(request.URL.Path, "/networks/creatorhub_browser-") - networkMu.Lock() - if networkDeleted[alias] { - networkMu.Unlock() - response.WriteHeader(http.StatusNotFound) - return - } - if networkMembers[alias] == nil { - networkMembers[alias] = map[string]bool{self: true} - } - members := map[string]any{} - for id := range networkMembers[alias] { - members[id] = map[string]string{"Name": id, "IPv4Address": "127.0.0.1/8"} - } - networkMu.Unlock() - _ = json.NewEncoder(response).Encode(map[string]any{ - "Id": "network-" + alias, "Name": "creatorhub_browser-" + alias, "Driver": "bridge", "Internal": false, "Attachable": false, "Ingress": false, - "Labels": map[string]string{managedLabel: "true", networkRoleLabel: browserNetworkRole, idLabel: alias, - bindingVersionLabel: "1"}, - "Containers": members, - }) - return - } - handler(response, request) - })) - return dockerClient{baseURL: server.URL, client: server.Client(), slow: server.Client()}, server -} - -func withAliasReservations(self string, next http.HandlerFunc) http.HandlerFunc { - var mu sync.Mutex - type reservation struct { - id, generation string - } - locks := map[string]reservation{} - return func(response http.ResponseWriter, request *http.Request) { - if request.Method == http.MethodGet && request.URL.Path == "/containers/"+self+"/json" { - _, _ = response.Write([]byte(`{"Id":"` + self + `","Image":"gateway-image-id","Config":{"Labels":{"` + gatewayMemberLabel + `":"true"}}}`)) - return - } - name := request.URL.Query().Get("name") - if request.Method == http.MethodPost && request.URL.Path == "/containers/create" && strings.HasPrefix(name, reservationPrefix) { - mu.Lock() - defer mu.Unlock() - if locks[name].id != "" { - response.WriteHeader(http.StatusConflict) - return - } - var payload struct { - Labels map[string]string `json:"Labels"` - } - _ = json.NewDecoder(request.Body).Decode(&payload) - id := "reservation-" + strings.TrimPrefix(name, reservationPrefix) - locks[name] = reservation{id: id, generation: payload.Labels[reservationGenLabel]} - response.WriteHeader(http.StatusCreated) - _ = json.NewEncoder(response).Encode(map[string]string{"Id": id}) - return - } - if request.Method == http.MethodGet && strings.HasPrefix(request.URL.Path, "/containers/"+reservationPrefix) && strings.HasSuffix(request.URL.Path, "/json") { - name := strings.TrimSuffix(strings.TrimPrefix(request.URL.Path, "/containers/"), "/json") - mu.Lock() - current, found := locks[name] - mu.Unlock() - if !found { - response.WriteHeader(http.StatusNotFound) - return - } - _ = json.NewEncoder(response).Encode(map[string]any{"Id": current.id, "Config": map[string]any{"Labels": map[string]string{ - reservationLabel: "true", idLabel: strings.TrimPrefix(name, reservationPrefix), reservationGenLabel: current.generation, - }}}) - return - } - if request.Method == http.MethodDelete && strings.HasPrefix(request.URL.Path, "/containers/reservation-") { - id := strings.TrimPrefix(request.URL.Path, "/containers/") - mu.Lock() - for name, current := range locks { - if current.id == id { - delete(locks, name) - break - } - } - mu.Unlock() - response.WriteHeader(http.StatusNoContent) - return - } - next(response, request) - } -} - -func decodeJSONBody(t *testing.T, response *http.Response) map[string]any { - t.Helper() - var body map[string]any - if err := json.NewDecoder(response.Body).Decode(&body); err != nil { - t.Fatalf("decode JSON body: %v", err) - } - return body -} - -const testCreateBody = `{"alias":"account-a","name":"账号甲","image":"registry.example/browser:1.2.3",` + - `"cmd":["--fingerprint=1000","--lang=zh-CN","about:blank"],"volume":"creatorhub-profile-account-a",` + - `"binding_version":1,"network_exit_id":"exit-1",` + - `"network_exit":{"protocol":"socks5","host":"proxy.example","port":1080}}` - -const testGenerationBody = `{"binding_version":1,"runtime_id":"container-id"}` - -func TestAliasReservationNameDoesNotOverlapRuntimeNamespace(t *testing.T) { - if strings.HasPrefix(reservationPrefix, namePrefix) || strings.HasPrefix(namePrefix, reservationPrefix) { - t.Fatalf("reservation and runtime prefixes overlap: reservation=%q runtime=%q", reservationPrefix, namePrefix) - } - generation := "" - server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { - switch { - case request.Method == http.MethodGet && request.URL.Path == "/containers/gateway-self/json": - _, _ = response.Write([]byte(`{"Id":"gateway-self","Image":"gateway-image","Config":{"Labels":{"` + gatewayMemberLabel + `":"true"}}}`)) - case request.Method == http.MethodPost && request.URL.Path == "/containers/create": - if request.URL.Query().Get("name") == namePrefix+"lock-account-a" { - response.WriteHeader(http.StatusConflict) - return - } - var payload struct { - Labels map[string]string `json:"Labels"` - } - _ = json.NewDecoder(request.Body).Decode(&payload) - generation = payload.Labels[reservationGenLabel] - response.WriteHeader(http.StatusCreated) - _, _ = response.Write([]byte(`{"Id":"reservation-id"}`)) - case request.Method == http.MethodGet && request.URL.Path == "/containers/"+reservationPrefix+"account-a/json": - _, _ = response.Write([]byte(`{"Id":"reservation-id","Config":{"Labels":{"` + reservationLabel + `":"true","` + idLabel + `":"account-a","` + reservationGenLabel + `":"` + generation + `"}}}`)) - case request.Method == http.MethodDelete && request.URL.Path == "/containers/reservation-id": - response.WriteHeader(http.StatusNoContent) - default: - t.Fatalf("unexpected Docker request %s %s", request.Method, request.URL.String()) - } - })) - defer server.Close() - docker := dockerClient{baseURL: server.URL, client: server.Client(), slow: server.Client()} - _, release, err := (&dockerAliasReservations{docker: docker, self: "gateway-self"}).acquire("account-a") - if err != nil { - t.Fatalf("runtime alias lock-account-a occupied account-a reservation: %v", err) - } - release() -} - -func TestAliasReservationRecoversInvalidCreateAndReleaseResponses(t *testing.T) { - t.Run("create disconnect after apply", func(t *testing.T) { - generation := "" - removed := false - server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { - switch { - case request.Method == http.MethodGet && request.URL.Path == "/containers/gateway-self/json": - _, _ = response.Write([]byte(`{"Id":"gateway-self","Image":"gateway-image","Config":{"Labels":{"` + gatewayMemberLabel + `":"true"}}}`)) - case request.Method == http.MethodPost && request.URL.Path == "/containers/create": - var payload struct { - Labels map[string]string `json:"Labels"` - } - _ = json.NewDecoder(request.Body).Decode(&payload) - generation = payload.Labels[reservationGenLabel] - connection, _, _ := response.(http.Hijacker).Hijack() - _ = connection.Close() - case request.Method == http.MethodGet && request.URL.Path == "/containers/"+reservationPrefix+"account-a/json": - _, _ = response.Write([]byte(`{"Id":"reservation-id","Config":{"Labels":{"` + reservationLabel + `":"true","` + idLabel + `":"account-a","` + reservationGenLabel + `":"` + generation + `"}}}`)) - case request.Method == http.MethodDelete && request.URL.Path == "/containers/reservation-id": - removed = true - response.WriteHeader(http.StatusNoContent) - default: - t.Fatalf("unexpected Docker request %s %s", request.Method, request.URL.String()) - } - })) - defer server.Close() - docker := dockerClient{baseURL: server.URL, client: server.Client(), slow: server.Client()} - if _, _, err := (&dockerAliasReservations{docker: docker, self: "gateway-self"}).acquire("account-a"); err == nil || !removed { - t.Fatalf("disconnected reservation create was not reconciled by immutable generation: removed=%v err=%v", removed, err) - } - }) - - t.Run("invalid create body", func(t *testing.T) { - reservationExists, removed := false, false - generation := "" - server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { - switch { - case request.Method == http.MethodGet && request.URL.Path == "/containers/gateway-self/json": - _, _ = response.Write([]byte(`{"Id":"gateway-self","Image":"gateway-image","Config":{"Labels":{"` + gatewayMemberLabel + `":"true"}}}`)) - case request.Method == http.MethodPost && request.URL.Path == "/containers/create": - var payload struct { - Labels map[string]string `json:"Labels"` - } - _ = json.NewDecoder(request.Body).Decode(&payload) - generation = payload.Labels[reservationGenLabel] - reservationExists = true - response.WriteHeader(http.StatusCreated) - _, _ = response.Write([]byte(`{}`)) - case request.Method == http.MethodGet && request.URL.Path == "/containers/"+reservationPrefix+"account-a/json": - if !reservationExists { - response.WriteHeader(http.StatusNotFound) - return - } - _, _ = response.Write([]byte(`{"Id":"reservation-id","Config":{"Labels":{"` + reservationLabel + `":"true","` + idLabel + `":"account-a","` + reservationGenLabel + `":"` + generation + `"}}}`)) - case request.Method == http.MethodDelete && request.URL.Path == "/containers/reservation-id": - reservationExists, removed = false, true - response.WriteHeader(http.StatusNoContent) - default: - t.Fatalf("unexpected Docker request %s %s", request.Method, request.URL.String()) - } - })) - defer server.Close() - docker := dockerClient{baseURL: server.URL, client: server.Client(), slow: server.Client()} - if _, _, err := (&dockerAliasReservations{docker: docker, self: "gateway-self"}).acquire("account-a"); err == nil || !removed { - t.Fatalf("invalid reservation response was not reconciled: removed=%v err=%v", removed, err) - } - }) - - t.Run("nonempty foreign create id", func(t *testing.T) { - generation := "" - deletes := 0 - server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { - switch { - case request.Method == http.MethodGet && request.URL.Path == "/containers/gateway-self/json": - _, _ = response.Write([]byte(`{"Id":"gateway-self","Image":"gateway-image","Config":{"Labels":{"` + gatewayMemberLabel + `":"true"}}}`)) - case request.Method == http.MethodPost && request.URL.Path == "/containers/create": - var payload struct { - Labels map[string]string `json:"Labels"` - } - _ = json.NewDecoder(request.Body).Decode(&payload) - generation = payload.Labels[reservationGenLabel] - response.WriteHeader(http.StatusCreated) - _, _ = response.Write([]byte(`{"Id":"foreign-id"}`)) - case request.Method == http.MethodGet && request.URL.Path == "/containers/"+reservationPrefix+"account-a/json": - _, _ = response.Write([]byte(`{"Id":"reservation-id","Config":{"Labels":{"` + reservationLabel + `":"true","` + idLabel + `":"account-a","` + reservationGenLabel + `":"` + generation + `"}}}`)) - case request.Method == http.MethodDelete: - deletes++ - response.WriteHeader(http.StatusNotFound) - default: - t.Fatalf("unexpected Docker request %s %s", request.Method, request.URL.String()) - } - })) - defer server.Close() - docker := dockerClient{baseURL: server.URL, client: server.Client(), slow: server.Client()} - if _, _, err := (&dockerAliasReservations{docker: docker, self: "gateway-self"}).acquire("account-a"); err == nil || deletes != 0 { - t.Fatalf("foreign 201 id was accepted or deleted: deletes=%d err=%v", deletes, err) - } - }) - - t.Run("201 inspect disconnect", func(t *testing.T) { - deletes := 0 - server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { - switch { - case request.Method == http.MethodGet && request.URL.Path == "/containers/gateway-self/json": - _, _ = response.Write([]byte(`{"Id":"gateway-self","Image":"gateway-image","Config":{"Labels":{"` + gatewayMemberLabel + `":"true"}}}`)) - case request.Method == http.MethodPost && request.URL.Path == "/containers/create": - response.WriteHeader(http.StatusCreated) - _, _ = response.Write([]byte(`{"Id":"reservation-id"}`)) - case request.Method == http.MethodGet && request.URL.Path == "/containers/"+reservationPrefix+"account-a/json": - connection, _, _ := response.(http.Hijacker).Hijack() - _ = connection.Close() - case request.Method == http.MethodDelete: - deletes++ - response.WriteHeader(http.StatusNoContent) - default: - t.Fatalf("unexpected Docker request %s %s", request.Method, request.URL.String()) - } - })) - defer server.Close() - docker := dockerClient{baseURL: server.URL, client: server.Client(), slow: server.Client()} - if _, _, err := (&dockerAliasReservations{docker: docker, self: "gateway-self"}).acquire("account-a"); err == nil || deletes != 0 { - t.Fatalf("unverified 201 reservation was deleted: deletes=%d err=%v", deletes, err) - } - }) - - t.Run("delete 404 requires generation absence", func(t *testing.T) { - generation := "" - deletes := 0 - server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { - switch { - case request.Method == http.MethodGet && request.URL.Path == "/containers/gateway-self/json": - _, _ = response.Write([]byte(`{"Id":"gateway-self","Image":"gateway-image","Config":{"Labels":{"` + gatewayMemberLabel + `":"true"}}}`)) - case request.Method == http.MethodPost && request.URL.Path == "/containers/create": - var payload struct { - Labels map[string]string `json:"Labels"` - } - _ = json.NewDecoder(request.Body).Decode(&payload) - generation = payload.Labels[reservationGenLabel] - response.WriteHeader(http.StatusCreated) - _, _ = response.Write([]byte(`{"Id":"reservation-id"}`)) - case request.Method == http.MethodGet && request.URL.Path == "/containers/"+reservationPrefix+"account-a/json": - _, _ = response.Write([]byte(`{"Id":"reservation-id","Config":{"Labels":{"` + reservationLabel + `":"true","` + idLabel + `":"account-a","` + reservationGenLabel + `":"` + generation + `"}}}`)) - case request.Method == http.MethodDelete && request.URL.Path == "/containers/reservation-id": - deletes++ - response.WriteHeader(http.StatusNotFound) - default: - t.Fatalf("unexpected Docker request %s %s", request.Method, request.URL.String()) - } - })) - defer server.Close() - docker := dockerClient{baseURL: server.URL, client: server.Client(), slow: server.Client()} - locks := &dockerAliasReservations{docker: docker, self: "gateway-self"} - id, _, err := locks.acquire("account-a") - if err != nil { - t.Fatal(err) - } - if err := locks.remove("account-a", generation, id); err == nil || deletes != 1 { - t.Fatalf("404 was accepted while the reservation generation remained: deletes=%d err=%v", deletes, err) - } - }) - - t.Run("release retry", func(t *testing.T) { - deletes := 0 - generation := "" - server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { - switch { - case request.Method == http.MethodGet && request.URL.Path == "/containers/gateway-self/json": - _, _ = response.Write([]byte(`{"Id":"gateway-self","Image":"gateway-image","Config":{"Labels":{"` + gatewayMemberLabel + `":"true"}}}`)) - case request.Method == http.MethodPost && request.URL.Path == "/containers/create": - var payload struct { - Labels map[string]string `json:"Labels"` - } - _ = json.NewDecoder(request.Body).Decode(&payload) - generation = payload.Labels[reservationGenLabel] - response.WriteHeader(http.StatusCreated) - _, _ = response.Write([]byte(`{"Id":"reservation-id"}`)) - case request.Method == http.MethodGet && request.URL.Path == "/containers/"+reservationPrefix+"account-a/json": - _, _ = response.Write([]byte(`{"Id":"reservation-id","Config":{"Labels":{"` + reservationLabel + `":"true","` + idLabel + `":"account-a","` + reservationGenLabel + `":"` + generation + `"}}}`)) - case request.Method == http.MethodDelete && request.URL.Path == "/containers/reservation-id": - deletes++ - if deletes == 1 { - response.WriteHeader(http.StatusInternalServerError) - return - } - response.WriteHeader(http.StatusNoContent) - default: - t.Fatalf("unexpected Docker request %s %s", request.Method, request.URL.String()) - } - })) - defer server.Close() - docker := dockerClient{baseURL: server.URL, client: server.Client(), slow: server.Client()} - _, release, err := (&dockerAliasReservations{docker: docker, self: "gateway-self"}).acquire("account-a") - if err != nil { - t.Fatal(err) - } - release() - if deletes != 2 { - t.Fatalf("reservation release did not retry the immutable id: deletes=%d", deletes) - } - }) - - t.Run("release does not delete successor", func(t *testing.T) { - generation := "" - deletes := 0 - server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { - switch { - case request.Method == http.MethodGet && request.URL.Path == "/containers/gateway-self/json": - _, _ = response.Write([]byte(`{"Id":"gateway-self","Image":"gateway-image","Config":{"Labels":{"` + gatewayMemberLabel + `":"true"}}}`)) - case request.Method == http.MethodPost && request.URL.Path == "/containers/create": - var payload struct { - Labels map[string]string `json:"Labels"` - } - _ = json.NewDecoder(request.Body).Decode(&payload) - generation = payload.Labels[reservationGenLabel] - response.WriteHeader(http.StatusCreated) - _, _ = response.Write([]byte(`{"Id":"reservation-id"}`)) - case request.Method == http.MethodDelete && request.URL.Path == "/containers/reservation-id": - deletes++ - response.WriteHeader(http.StatusInternalServerError) - case request.Method == http.MethodGet && request.URL.Path == "/containers/"+reservationPrefix+"account-a/json": - if deletes == 0 { - _, _ = response.Write([]byte(`{"Id":"reservation-id","Config":{"Labels":{"` + reservationLabel + `":"true","` + idLabel + `":"account-a","` + reservationGenLabel + `":"` + generation + `"}}}`)) - return - } - _, _ = response.Write([]byte(`{"Id":"successor-id","Config":{"Labels":{"` + reservationLabel + `":"true","` + idLabel + `":"account-a","` + reservationGenLabel + `":"successor"}}}`)) - default: - t.Fatalf("unexpected Docker request %s %s generation=%s", request.Method, request.URL.String(), generation) - } - })) - defer server.Close() - docker := dockerClient{baseURL: server.URL, client: server.Client(), slow: server.Client()} - _, release, err := (&dockerAliasReservations{docker: docker, self: "gateway-self"}).acquire("account-a") - if err != nil { - t.Fatal(err) - } - release() - if deletes != 1 { - t.Fatalf("reservation release deleted a successor generation: deletes=%d", deletes) - } - }) -} - -func TestGatewayCreatesNetworkDisabledStoppedRecoveryContainer(t *testing.T) { - created := false - docker, server := testDocker(func(response http.ResponseWriter, request *http.Request) { - switch { - case request.Method == http.MethodGet && strings.HasPrefix(request.URL.Path, "/images/"): - _, _ = response.Write([]byte(`{}`)) - case request.Method == http.MethodGet && strings.HasPrefix(request.URL.Path, "/containers/"+namePrefix): - response.WriteHeader(http.StatusNotFound) - case request.Method == http.MethodPost && strings.HasSuffix(request.URL.Path, "/containers/create"): - var payload map[string]any - _ = json.NewDecoder(request.Body).Decode(&payload) - host := payload["HostConfig"].(map[string]any) - labels := payload["Labels"].(map[string]any) - encoded, _ := json.Marshal(payload["Cmd"]) - if host["NetworkMode"] != "none" || labels[networkExitLabel] != "" || strings.Contains(string(encoded), "proxy") { - t.Fatalf("unsafe stopped recovery payload: %#v", payload) - } - created = true - response.WriteHeader(http.StatusCreated) - _, _ = response.Write([]byte(`{"Id":"stopped-container"}`)) - default: - t.Fatalf("stopped recovery unexpectedly called Docker %s %s", request.Method, request.URL.String()) - } - }) - defer server.Close() - handler := newGateway(docker, "creatorhub_browser", testToken) - body := `{"alias":"account-a","name":"账号甲","image":"registry.example/browser:1.2.3",` + - `"cmd":["--fingerprint=1000","about:blank"],"volume":"creatorhub-profile-account-a",` + - `"binding_version":1,"network_exit_id":"","network_exit":{},"stopped":true}` - response := httptest.NewRecorder() - adaptor.FiberApp(handler).ServeHTTP(response, authed(http.MethodPost, "/v1/browsers", strings.NewReader(body))) - if response.Code != http.StatusCreated || !created { - t.Fatalf("stopped recovery create failed: status=%d body=%s", response.Code, response.Body.String()) - } -} - -func TestGatewayCreatesConstrainedBrowserWithPlatformSpec(t *testing.T) { - var created map[string]any - docker, server := testDocker(func(response http.ResponseWriter, request *http.Request) { - switch { - case request.Method == http.MethodGet && strings.HasPrefix(request.URL.Path, "/images/"): - response.WriteHeader(http.StatusOK) - _, _ = response.Write([]byte(`{}`)) - case request.Method == http.MethodGet && strings.HasPrefix(request.URL.Path, "/containers/"+namePrefix): - response.WriteHeader(http.StatusNotFound) - case request.Method == http.MethodPost && strings.HasSuffix(request.URL.Path, "/containers/create"): - if got := request.URL.Query().Get("name"); got != namePrefix+"account-a" { - t.Fatalf("unexpected container name %q", got) - } - if err := json.NewDecoder(request.Body).Decode(&created); err != nil { - t.Fatal(err) - } - response.WriteHeader(http.StatusCreated) - _, _ = response.Write([]byte(`{"Id":"container-id"}`)) - case request.Method == http.MethodPost && strings.HasSuffix(request.URL.Path, "/containers/container-id/start"): - response.WriteHeader(http.StatusNoContent) - default: - t.Fatalf("unexpected Docker request %s %s", request.Method, request.URL.String()) - } - }) - defer server.Close() - - handler := newGateway(docker, "creatorhub_browser", testToken) - response := httptest.NewRecorder() - adaptor.FiberApp(handler).ServeHTTP(response, authed(http.MethodPost, "/v1/browsers", strings.NewReader(testCreateBody))) - - if response.Code != http.StatusCreated { - t.Fatalf("expected 201, got %d: %s", response.Code, response.Body.String()) - } - if created["Image"] != "registry.example/browser:1.2.3" { - t.Fatalf("gateway must run the platform-specified image: %#v", created["Image"]) - } - if created["User"] != browserUser || created["Entrypoint"].([]any)[0] != browserEntrypoint { - t.Fatalf("runtime identity is not fixed: user=%#v entrypoint=%#v", created["User"], created["Entrypoint"]) - } - cmd := created["Cmd"].([]any) - if len(cmd) != 5 || cmd[0] != "--fingerprint=1000" || !strings.HasPrefix(cmd[2].(string), "--proxy-server=http://docker-gateway:") || - cmd[3] != "--disable-non-proxied-udp" || cmd[4] != "about:blank" { - t.Fatalf("cmd must be passed through verbatim: %#v", created["Cmd"]) - } - host := created["HostConfig"].(map[string]any) - if host["NetworkMode"] != "network-account-a" || host["ReadonlyRootfs"] != true { - t.Fatalf("missing container isolation: %#v", host) - } - tmpfs := host["Tmpfs"].(map[string]any) - if tmpfs["/tmp/.X11-unix"] == nil || tmpfs["/home/ubuntu"] == nil { - t.Fatalf("missing writable runtime paths: %#v", tmpfs) - } - mount := host["Mounts"].([]any)[0].(map[string]any) - if mount["Source"] != "creatorhub-profile-account-a" || mount["Target"] != "/data" { - t.Fatalf("profile volume must come from the request: %#v", mount) - } - labels := created["Labels"].(map[string]any) - if labels[managedLabel] != "true" || labels[idLabel] != "account-a" || labels[nameLabel] != "账号甲" { - t.Fatalf("missing ownership labels: %#v", labels) - } -} - -func TestGatewayCreatesDirectBrowserWithoutProxyArguments(t *testing.T) { - var created map[string]any - docker, server := testDocker(func(response http.ResponseWriter, request *http.Request) { - switch { - case request.Method == http.MethodGet && strings.HasPrefix(request.URL.Path, "/images/"): - _, _ = response.Write([]byte(`{}`)) - case request.Method == http.MethodGet && strings.HasPrefix(request.URL.Path, "/containers/"+namePrefix): - response.WriteHeader(http.StatusNotFound) - case request.Method == http.MethodPost && strings.HasSuffix(request.URL.Path, "/containers/create"): - if err := json.NewDecoder(request.Body).Decode(&created); err != nil { - t.Fatal(err) - } - response.WriteHeader(http.StatusCreated) - _, _ = response.Write([]byte(`{"Id":"container-id"}`)) - case request.Method == http.MethodPost && strings.HasSuffix(request.URL.Path, "/containers/container-id/start"): - response.WriteHeader(http.StatusNoContent) - default: - t.Fatalf("unexpected Docker request %s %s", request.Method, request.URL.String()) - } - }) - defer server.Close() - - handler := newGateway(docker, "creatorhub_browser", testToken) - body := `{"alias":"account-a","name":"账号甲","image":"registry.example/browser:1.2.3",` + - `"cmd":["--fingerprint=1000","about:blank"],"volume":"creatorhub-profile-account-a",` + - `"binding_version":1,"network_exit_id":"","network_exit":{}}` - response := httptest.NewRecorder() - adaptor.FiberApp(handler).ServeHTTP(response, authed(http.MethodPost, "/v1/browsers", strings.NewReader(body))) - - if response.Code != http.StatusCreated { - t.Fatalf("expected 201, got %d: %s", response.Code, response.Body.String()) - } - cmd := created["Cmd"].([]any) - encoded, _ := json.Marshal(cmd) - if len(cmd) != 2 || strings.Contains(string(encoded), "proxy") { - t.Fatalf("direct runtime received proxy arguments: %#v", cmd) - } - host := created["HostConfig"].(map[string]any) - labels := created["Labels"].(map[string]any) - if host["NetworkMode"] != "network-account-a" || labels[networkExitLabel] != "" || labels[proxyPortLabel] != "0" { - t.Fatalf("direct runtime metadata is invalid: host=%#v labels=%#v", host, labels) - } -} - -func TestGatewayCreateUsesCapturedNetworkIDAcrossNameReplacement(t *testing.T) { - networkID := "" - members := map[string]string{} - usedNetworkID, touchedReplacement := "", false - server := httptest.NewServer(withAliasReservations("gateway-self", func(response http.ResponseWriter, request *http.Request) { - switch { - case request.Method == http.MethodGet && strings.HasPrefix(request.URL.Path, "/images/"): - response.WriteHeader(http.StatusOK) - case request.Method == http.MethodGet && request.URL.Path == "/containers/"+namePrefix+"account-a/json": - response.WriteHeader(http.StatusNotFound) - case request.Method == http.MethodGet && strings.HasPrefix(request.URL.Path, "/networks/"): - reference := strings.TrimPrefix(request.URL.Path, "/networks/") - if reference == "network-n2" { - touchedReplacement = true - } - if networkID == "" || (reference != "creatorhub_browser-account-a" && reference != networkID) { - response.WriteHeader(http.StatusNotFound) - return - } - containers := map[string]any{} - for id, name := range members { - containers[id] = map[string]string{"Name": name, "IPv4Address": "127.0.0.1/8"} - } - _ = json.NewEncoder(response).Encode(map[string]any{ - "Id": networkID, "Name": "creatorhub_browser-account-a", "Driver": "bridge", "Internal": false, "Attachable": false, "Ingress": false, - "Labels": map[string]string{managedLabel: "true", networkRoleLabel: browserNetworkRole, idLabel: "account-a", bindingVersionLabel: "1"}, "Containers": containers, - }) - case request.Method == http.MethodPost && request.URL.Path == "/networks/create": - networkID = "network-n1" - response.WriteHeader(http.StatusCreated) - _, _ = response.Write([]byte(`{"Id":"network-n1"}`)) - case request.Method == http.MethodPost && request.URL.Path == "/networks/network-n1/connect": - members["gateway-self"] = "gateway-self" - response.WriteHeader(http.StatusOK) - case request.Method == http.MethodPost && request.URL.Path == "/containers/create": - var payload struct { - HostConfig map[string]any `json:"HostConfig"` - } - _ = json.NewDecoder(request.Body).Decode(&payload) - usedNetworkID, _ = payload.HostConfig["NetworkMode"].(string) - networkID, members = "network-n2", map[string]string{"replacement": "replacement"} - response.WriteHeader(http.StatusNotFound) - default: - t.Fatalf("unexpected Docker request %s %s", request.Method, request.URL.String()) - } - })) - defer server.Close() - handler := newGatewayWithSelf(dockerClient{baseURL: server.URL, client: server.Client(), slow: server.Client()}, - "creatorhub_browser", testToken, "gateway-self") - response := httptest.NewRecorder() - adaptor.FiberApp(handler).ServeHTTP(response, authed(http.MethodPost, "/v1/browsers", strings.NewReader(testCreateBody))) - if response.Code != http.StatusBadGateway || usedNetworkID != "network-n1" || touchedReplacement || networkID != "network-n2" || members["replacement"] == "" { - t.Fatalf("stale create crossed network generation: status=%d mode=%q touchedN2=%v network=%q members=%v body=%s", - response.Code, usedNetworkID, touchedReplacement, networkID, members, response.Body.String()) - } -} - -func TestGatewayDockerInspectContainsNoProxyCredentials(t *testing.T) { - var created map[string]any - docker, server := testDocker(func(response http.ResponseWriter, request *http.Request) { - switch { - case request.Method == http.MethodGet && strings.HasPrefix(request.URL.Path, "/images/"): - response.WriteHeader(http.StatusOK) - case request.Method == http.MethodGet && strings.HasPrefix(request.URL.Path, "/containers/"+namePrefix): - response.WriteHeader(http.StatusNotFound) - case request.Method == http.MethodPost && strings.HasSuffix(request.URL.Path, "/containers/create"): - if err := json.NewDecoder(request.Body).Decode(&created); err != nil { - t.Fatal(err) - } - response.WriteHeader(http.StatusCreated) - _, _ = response.Write([]byte(`{"Id":"container-id"}`)) - case request.Method == http.MethodPost && strings.HasSuffix(request.URL.Path, "/containers/container-id/start"): - response.WriteHeader(http.StatusNoContent) - default: - t.Fatalf("unexpected Docker request %s %s", request.Method, request.URL.String()) - } - }) - defer server.Close() - - handler := newGateway(docker, "creatorhub_browser", testToken) - body := strings.Replace(testCreateBody, `"protocol":"socks5","host":"proxy.example","port":1080`, - `"protocol":"socks5","host":"proxy.example","port":1080,"username":"operator","password":"ephemeral"`, 1) - response := httptest.NewRecorder() - adaptor.FiberApp(handler).ServeHTTP(response, authed(http.MethodPost, "/v1/browsers", strings.NewReader(body))) - if response.Code != http.StatusCreated { - t.Fatalf("expected 201, got %d: %s", response.Code, response.Body.String()) - } - inspect, _ := json.Marshal(created) - for _, secret := range []string{"operator", "ephemeral", "operator:ephemeral@", "proxy.example"} { - if bytes.Contains(inspect, []byte(secret)) { - t.Fatalf("Docker inspect leaked proxy credential %q: %s", secret, inspect) - } - } - if !bytes.Contains(inspect, []byte("--proxy-server=http://docker-gateway:")) { - t.Fatalf("Docker inspect is missing the secret-free proxy configuration: %s", inspect) - } -} - -func TestGatewayPullsMissingImageOnCreate(t *testing.T) { - tests := []struct { - name string - ref string - fromImage string - tag string - }{{ - name: "tagged ref splits repository and tag", - ref: "registry.example/browser:2.0.0", - fromImage: "registry.example/browser", - tag: "2.0.0", - }, { - name: "digest ref is pulled as a whole", - ref: "registry.example/browser@sha256:b9f23b8e3ac640174db0dfa49e9095fe7eb06f5db55a4e7550d979b35ff3a1b7", - fromImage: "registry.example/browser@sha256:b9f23b8e3ac640174db0dfa49e9095fe7eb06f5db55a4e7550d979b35ff3a1b7", - tag: "", - }} - for _, test := range tests { - t.Run(test.name, func(t *testing.T) { - pulled := false - docker, server := testDocker(func(response http.ResponseWriter, request *http.Request) { - switch { - case request.Method == http.MethodGet && strings.HasPrefix(request.URL.Path, "/images/"): - response.WriteHeader(http.StatusNotFound) - case request.Method == http.MethodPost && strings.HasSuffix(request.URL.Path, "/images/create"): - pulled = true - if request.URL.Query().Get("fromImage") != test.fromImage || request.URL.Query().Get("tag") != test.tag { - t.Fatalf("unexpected pull query %s", request.URL.RawQuery) - } - _, _ = response.Write([]byte(`{"status":"Download complete"}`)) - case request.Method == http.MethodGet && strings.HasPrefix(request.URL.Path, "/containers/"+namePrefix): - response.WriteHeader(http.StatusNotFound) - case request.Method == http.MethodPost && strings.HasSuffix(request.URL.Path, "/containers/create"): - response.WriteHeader(http.StatusCreated) - _, _ = response.Write([]byte(`{"Id":"container-id"}`)) - case request.Method == http.MethodPost && strings.HasSuffix(request.URL.Path, "/start"): - response.WriteHeader(http.StatusNoContent) - default: - t.Fatalf("unexpected Docker request %s %s", request.Method, request.URL.String()) - } - }) - defer server.Close() - - handler := newGateway(docker, "creatorhub_browser", testToken) - body := `{"alias":"account-a","name":"账号甲","image":"` + test.ref + - `","cmd":["--fingerprint=1000","about:blank"],"volume":"creatorhub-profile-account-a",` + - `"binding_version":1,"network_exit_id":"exit-1",` + - `"network_exit":{"protocol":"socks5","host":"proxy.example","port":1080}}` - response := httptest.NewRecorder() - adaptor.FiberApp(handler).ServeHTTP(response, authed(http.MethodPost, "/v1/browsers", strings.NewReader(body))) - - if response.Code != http.StatusCreated || !pulled { - t.Fatalf("expected pull-then-create, status=%d pulled=%v body=%s", response.Code, pulled, response.Body.String()) - } - }) - } -} - -func TestGatewayRejectsCreateWithoutValidToken(t *testing.T) { - docker, server := testDocker(func(http.ResponseWriter, *http.Request) { - t.Fatal("no Docker request is expected for an unauthorized call") - }) - defer server.Close() - - handler := newGateway(docker, "creatorhub_browser", testToken) - for name, header := range map[string]string{ - "missing": "", - "malformed": testToken, - "wrong": "Bearer not-the-token", - } { - request := httptest.NewRequest(http.MethodPost, "/v1/browsers", strings.NewReader(testCreateBody)) - if header != "" { - request.Header.Set("Authorization", header) - } - response := httptest.NewRecorder() - adaptor.FiberApp(handler).ServeHTTP(response, request) - if response.Code != http.StatusUnauthorized { - t.Fatalf("%s token: expected 401, got %d: %s", name, response.Code, response.Body.String()) - } - } - - response := httptest.NewRecorder() - adaptor.FiberApp(handler).ServeHTTP(response, httptest.NewRequest(http.MethodGet, "/healthz", nil)) - if response.Code != http.StatusNoContent { - t.Fatalf("healthz must stay unauthenticated, got %d", response.Code) - } -} - -func TestGatewayRejectsInvalidCreateRequest(t *testing.T) { - tests := map[string]string{ - "unknown field": `{"alias":"account-a","seed":1}`, - "invalid alias": `{"alias":"AccountA","name":"甲","image":"reg/img:1","cmd":["--fingerprint=1"],"volume":"creatorhub-profile-account-a"}`, - "invalid image": `{"alias":"account-a","name":"甲","image":"","cmd":["--fingerprint=1"],"volume":"creatorhub-profile-account-a"}`, - "empty cmd": `{"alias":"account-a","name":"甲","image":"reg/img:1","cmd":[],"volume":"creatorhub-profile-account-a"}`, - "invalid volume": `{"alias":"account-a","name":"甲","image":"reg/img:1","cmd":["--fingerprint=1"],"volume":"bad volume!"}`, - "proxy override": `{"alias":"account-a","name":"甲","image":"reg/img:1","cmd":["--fingerprint=1","--proxy-server=http://direct:8080","about:blank"],"volume":"creatorhub-profile-account-a","network_exit":{"protocol":"socks5","host":"proxy","port":1080}}`, - } - for name, body := range tests { - t.Run(name, func(t *testing.T) { - handler := newGateway(dockerClient{}, "creatorhub_browser", testToken) - response := httptest.NewRecorder() - adaptor.FiberApp(handler).ServeHTTP(response, authed(http.MethodPost, "/v1/browsers", strings.NewReader(body))) - if response.Code != http.StatusBadRequest { - t.Fatalf("expected 400, got %d: %s", response.Code, response.Body.String()) - } - }) - } -} - -func TestGatewayRejectsOversizedCreateRequest(t *testing.T) { - handler := newGateway(dockerClient{}, "creatorhub_browser", testToken) - request := authed(http.MethodPost, "/v1/browsers", strings.NewReader(strings.Repeat("x", (1<<20)+1))) - response := httptest.NewRecorder() - adaptor.FiberApp(handler).ServeHTTP(response, request) - if response.Code != http.StatusRequestEntityTooLarge { - t.Fatalf("expected 413 for oversized body, status=%d body=%s", response.Code, response.Body.String()) - } - - handler.Post("/request-limit", func(fiber.Ctx) error { return fiber.ErrRequestEntityTooLarge }) - jsonResponse, err := handler.Test(httptest.NewRequest(http.MethodPost, "/request-limit", nil)) - if err != nil { - t.Fatal(err) - } - defer jsonResponse.Body.Close() - var body map[string]string - decodeErr := json.NewDecoder(jsonResponse.Body).Decode(&body) - contentType := jsonResponse.Header.Get("Content-Type") - if jsonResponse.StatusCode != http.StatusRequestEntityTooLarge || decodeErr != nil || body["error"] == "" || !strings.HasPrefix(contentType, "application/json") { - t.Fatalf("expected JSON 413 envelope, status=%d body=%v decode=%v content-type=%q", jsonResponse.StatusCode, body, decodeErr, contentType) - } -} - -func TestGatewayReconcilesContainerWhenCreateResponseHasNoID(t *testing.T) { - created, started := false, false - docker, server := testDocker(func(response http.ResponseWriter, request *http.Request) { - switch { - case request.Method == http.MethodGet && strings.HasPrefix(request.URL.Path, "/images/"): - response.WriteHeader(http.StatusOK) - _, _ = response.Write([]byte(`{}`)) - case request.Method == http.MethodGet && strings.HasPrefix(request.URL.Path, "/containers/"+namePrefix): - if !created { - response.WriteHeader(http.StatusNotFound) - return - } - _, _ = response.Write([]byte(`{"Id":"actual-container-id","Config":{"Labels":{"` + managedLabel + `":"true","` + idLabel + `":"account-a","` + bindingVersionLabel + `":"1","` + networkIDLabel + `":"network-account-a"}}}`)) - case request.Method == http.MethodPost && strings.HasSuffix(request.URL.Path, "/containers/create"): - created = true - response.WriteHeader(http.StatusCreated) - _, _ = response.Write([]byte(`{"Id":""}`)) - case request.Method == http.MethodPost && request.URL.Path == "/containers/actual-container-id/start": - started = true - response.WriteHeader(http.StatusNoContent) - default: - t.Fatalf("unexpected Docker request %s %s", request.Method, request.URL.String()) - } - }) - defer server.Close() - - handler := newGateway(docker, "creatorhub_browser", testToken) - response := httptest.NewRecorder() - adaptor.FiberApp(handler).ServeHTTP(response, authed(http.MethodPost, "/v1/browsers", strings.NewReader(testCreateBody))) - - if response.Code != http.StatusCreated || !started { - t.Fatalf("expected invalid create response reconciliation, status=%d started=%v body=%s", response.Code, started, response.Body.String()) - } -} - -func TestGatewayRemovesFailedContainerAndPreservesProfile(t *testing.T) { - removed := false - docker, server := testDocker(func(response http.ResponseWriter, request *http.Request) { - switch { - case request.Method == http.MethodGet && strings.HasPrefix(request.URL.Path, "/images/"): - response.WriteHeader(http.StatusOK) - _, _ = response.Write([]byte(`{}`)) - case request.Method == http.MethodGet && strings.HasPrefix(request.URL.Path, "/containers/"+namePrefix): - response.WriteHeader(http.StatusNotFound) - case request.Method == http.MethodPost && strings.HasSuffix(request.URL.Path, "/containers/create"): - response.WriteHeader(http.StatusCreated) - _, _ = response.Write([]byte(`{"Id":"failed-id"}`)) - case request.Method == http.MethodPost && strings.HasSuffix(request.URL.Path, "/containers/failed-id/start"): - http.Error(response, "start failed", http.StatusInternalServerError) - case request.Method == http.MethodDelete && strings.Contains(request.URL.Path, "/containers/failed-id"): - removed = request.URL.Query().Get("force") == "1" && request.URL.Query().Get("v") == "0" - response.WriteHeader(http.StatusNoContent) - default: - t.Fatalf("unexpected Docker request %s %s", request.Method, request.URL.String()) - } - }) - defer server.Close() - - handler := newGateway(docker, "creatorhub_browser", testToken) - response := httptest.NewRecorder() - adaptor.FiberApp(handler).ServeHTTP(response, authed(http.MethodPost, "/v1/browsers", strings.NewReader(testCreateBody))) - - if response.Code != http.StatusBadGateway || !removed { - t.Fatalf("expected failed container cleanup with preserved volume, status=%d removed=%v body=%s", response.Code, removed, response.Body.String()) - } -} - -func TestGatewayDoesNotEchoProxyCredentialsFromDockerErrors(t *testing.T) { - docker, server := testDocker(func(response http.ResponseWriter, request *http.Request) { - if request.Method == http.MethodGet && strings.HasPrefix(request.URL.Path, "/images/") { - response.WriteHeader(http.StatusOK) - return - } - if request.Method == http.MethodGet && strings.HasPrefix(request.URL.Path, "/containers/"+namePrefix) { - response.WriteHeader(http.StatusNotFound) - return - } - if request.Method == http.MethodPost && strings.HasSuffix(request.URL.Path, "/containers/create") { - response.WriteHeader(http.StatusInternalServerError) - _, _ = response.Write([]byte(`invalid cmd --proxy-server=http://operator:ephemeral@proxy.example:8080`)) - return - } - t.Fatalf("unexpected Docker request %s %s", request.Method, request.URL.Path) - }) - defer server.Close() - handler := newGateway(docker, "creatorhub_browser", testToken) - body := strings.Replace(testCreateBody, `"protocol":"socks5","host":"proxy.example","port":1080`, - `"protocol":"http","host":"proxy.example","port":8080,"username":"operator","password":"ephemeral"`, 1) - response := httptest.NewRecorder() - adaptor.FiberApp(handler).ServeHTTP(response, authed(http.MethodPost, "/v1/browsers", strings.NewReader(body))) - if response.Code != http.StatusBadGateway || strings.Contains(response.Body.String(), "operator") || - strings.Contains(response.Body.String(), "ephemeral") || strings.Contains(response.Body.String(), "proxy.example") { - t.Fatalf("gateway leaked proxy material: status=%d body=%s", response.Code, response.Body.String()) - } -} - -func TestGatewayListsBrowsers(t *testing.T) { - docker, server := testDocker(func(response http.ResponseWriter, request *http.Request) { - if request.Method != http.MethodGet || request.URL.Path != "/containers/json" { - t.Fatalf("unexpected Docker request %s %s", request.Method, request.URL.String()) - } - _, _ = response.Write([]byte(`[{"Id":"container-id","State":"running","Status":"Up","Labels":{` + - `"` + idLabel + `":"account-a","` + nameLabel + `":"账号甲"}}]`)) - }) - defer server.Close() - - handler := newGateway(docker, "creatorhub_browser", testToken) - response := httptest.NewRecorder() - adaptor.FiberApp(handler).ServeHTTP(response, authed(http.MethodGet, "/v1/browsers", nil)) - - var browsers []browser - if response.Code != http.StatusOK || json.NewDecoder(response.Body).Decode(&browsers) != nil || - len(browsers) != 1 || browsers[0].Alias != "account-a" || browsers[0].Name != "账号甲" || - browsers[0].Endpoint != "http://creatorhub-browser-account-a:9222" { - t.Fatalf("unexpected list response status=%d body=%s", response.Code, response.Body.String()) - } -} - -func TestGatewayRestartRestoresExistingProxyListener(t *testing.T) { - reserved, err := net.Listen("tcp4", "127.0.0.1:0") - if err != nil { - t.Fatal(err) - } - port := reserved.Addr().(*net.TCPAddr).Port - _ = reserved.Close() - labels := map[string]string{ - managedLabel: "true", idLabel: "account-a", nameLabel: "账号甲", - bindingVersionLabel: "1", networkExitLabel: "exit-1", proxyPortLabel: strconv.Itoa(port), - networkIDLabel: "network-account-a", - } - docker, server := testDocker(func(response http.ResponseWriter, request *http.Request) { - switch { - case request.Method == http.MethodGet && strings.HasSuffix(request.URL.Path, "/containers/creatorhub-browser-account-a/json"): - _ = json.NewEncoder(response).Encode(map[string]any{"Id": "container-id", "Config": map[string]any{"Labels": labels}}) - case request.Method == http.MethodGet && request.URL.Path == "/containers/json": - _ = json.NewEncoder(response).Encode([]map[string]any{{"Id": "container-id", "State": "running", "Status": "Up", "Labels": labels}}) - default: - t.Fatalf("unexpected Docker request %s %s", request.Method, request.URL.String()) - } - }) - defer server.Close() - handler := newGateway(docker, "creatorhub_browser", testToken) - recovery := `{"binding_version":1,"runtime_id":"container-id","network_id":"network-account-a","network_exit_id":"exit-1","network_exit":{"protocol":"http","host":"127.0.0.1","port":1}}` - response := httptest.NewRecorder() - adaptor.FiberApp(handler).ServeHTTP(response, authed(http.MethodPost, "/v1/browsers/account-a/proxy", strings.NewReader(recovery))) - if response.Code != http.StatusNoContent { - t.Fatalf("proxy recovery failed: %d %s", response.Code, response.Body.String()) - } - response = httptest.NewRecorder() - adaptor.FiberApp(handler).ServeHTTP(response, authed(http.MethodGet, "/v1/browsers", nil)) - var browsers []browser - if response.Code != http.StatusOK || json.NewDecoder(response.Body).Decode(&browsers) != nil || len(browsers) != 1 || !browsers[0].ProxyReady { - t.Fatalf("restarted gateway did not report restored proxy: %d %s", response.Code, response.Body.String()) - } -} - -func TestGatewayRestoreFinalFenceRemovesStaleProxy(t *testing.T) { - reserved, err := net.Listen("tcp4", "127.0.0.1:0") - if err != nil { - t.Fatal(err) - } - port := reserved.Addr().(*net.TCPAddr).Port - _ = reserved.Close() - registry := newMemoryProxyRegistry() - networkID := "network-n1" - replaced := false - containerReads := 0 - exit := gatewayProxyExit{Protocol: "socks5", Host: "proxy.example", Port: 1080} - proxyURL, cleanup, err := registry.configure("account-a", 1, "127.0.0.1", port, exit, "network-n1") - if err != nil || !registry.bind("account-a", 1, proxyURL, "container-c1", "network-n1") { - t.Fatalf("seed existing proxy generation: %v", err) - } - defer cleanup() - labels := map[string]string{ - managedLabel: "true", idLabel: "account-a", bindingVersionLabel: "1", networkExitLabel: "exit-1", - proxyPortLabel: strconv.Itoa(port), networkIDLabel: "network-n1", - } - server := httptest.NewServer(withAliasReservations("gateway-self", func(response http.ResponseWriter, request *http.Request) { - switch { - case request.Method == http.MethodGet && request.URL.Path == "/containers/"+namePrefix+"account-a/json": - containerReads++ - registry.mu.Lock() - proxy := registry.proxies["account-a"] - bound := proxy != nil && proxy.runtimeID == "container-c1" - registry.mu.Unlock() - if bound && containerReads == 4 && !replaced { - networkID, replaced = "network-n2", true - } - _ = json.NewEncoder(response).Encode(map[string]any{"Id": "container-c1", "Config": map[string]any{"Labels": labels}}) - case request.Method == http.MethodGet && strings.HasPrefix(request.URL.Path, "/networks/"): - reference := strings.TrimPrefix(request.URL.Path, "/networks/") - if reference != networkID && reference != "creatorhub_browser-account-a" { - response.WriteHeader(http.StatusNotFound) - return - } - _ = json.NewEncoder(response).Encode(map[string]any{ - "Id": networkID, "Name": "creatorhub_browser-account-a", "Driver": "bridge", "Internal": false, "Attachable": false, "Ingress": false, - "Labels": map[string]string{managedLabel: "true", networkRoleLabel: browserNetworkRole, idLabel: "account-a", bindingVersionLabel: "1"}, - "Containers": map[string]any{ - "container-c1": map[string]string{"Name": namePrefix + "account-a", "IPv4Address": "127.0.0.3/8"}, - "gateway-self": map[string]string{"Name": "gateway-self", "IPv4Address": "127.0.0.1/8"}, - }, - }) - default: - t.Fatalf("unexpected Docker request %s %s", request.Method, request.URL.String()) - } - })) - defer server.Close() - docker := dockerClient{baseURL: server.URL, client: server.Client(), slow: server.Client()} - api := gateway{docker: docker, network: "creatorhub_browser", self: "gateway-self", token: testToken, proxies: registry, - locks: &dockerAliasReservations{docker: docker, self: "gateway-self"}} - app := fiber.New() - app.Use("/v1", api.authorize) - app.Post("/v1/browsers/:id/proxy", api.restoreProxy) - body := `{"binding_version":1,"runtime_id":"container-c1","network_id":"network-n1","network_exit_id":"exit-1","network_exit":{"protocol":"socks5","host":"proxy.example","port":1080}}` - response := httptest.NewRecorder() - adaptor.FiberApp(app).ServeHTTP(response, authed(http.MethodPost, "/v1/browsers/account-a/proxy", strings.NewReader(body))) - registry.mu.Lock() - proxy := registry.proxies["account-a"] - registry.mu.Unlock() - if response.Code != http.StatusConflict || !replaced || networkID != "network-n2" || proxy != nil { - t.Fatalf("restore final fence accepted a replaced network: status=%d replaced=%v network=%q proxy=%v body=%s", - response.Code, replaced, networkID, proxy, response.Body.String()) - } -} - -func TestGatewayProxyFenceRejectsGatewayAddressChange(t *testing.T) { - labels := map[string]string{managedLabel: "true", idLabel: "account-a", bindingVersionLabel: "1", - networkExitLabel: "exit-1", networkIDLabel: "network-n1"} - server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { - switch request.URL.Path { - case "/containers/" + namePrefix + "account-a/json": - _ = json.NewEncoder(response).Encode(map[string]any{"Id": "container-c1", "Config": map[string]any{"Labels": labels}}) - case "/containers/gateway-self/json": - _, _ = response.Write([]byte(`{"Id":"gateway-self","Config":{"Labels":{"` + gatewayMemberLabel + `":"true"}}}`)) - case "/networks/network-n1": - _ = json.NewEncoder(response).Encode(map[string]any{ - "Id": "network-n1", "Name": "creatorhub_browser-account-a", "Driver": "bridge", "Internal": false, "Attachable": false, "Ingress": false, - "Labels": map[string]string{managedLabel: "true", networkRoleLabel: browserNetworkRole, idLabel: "account-a", bindingVersionLabel: "1"}, - "Containers": map[string]any{ - "container-c1": map[string]string{"Name": namePrefix + "account-a", "IPv4Address": "127.0.0.2/8"}, - "gateway-self": map[string]string{"Name": "gateway-self", "IPv4Address": "127.0.0.4/8"}, - }, - }) - default: - t.Fatalf("unexpected Docker request %s %s", request.Method, request.URL.Path) - } - })) - defer server.Close() - docker := dockerClient{baseURL: server.URL, client: server.Client(), slow: server.Client()} - api := gateway{docker: docker, network: "creatorhub_browser", self: "gateway-self"} - input := proxyRestoreRequest{BindingVersion: 1, RuntimeID: "container-c1", NetworkID: "network-n1", NetworkExitID: "exit-1"} - expected := tenantNetworkGeneration{ID: "network-n1", Name: "creatorhub_browser-account-a", RuntimeAttached: true, - GatewayMembers: []string{"gateway-self"}, SelfMember: "gateway-self"} - if err := api.requireProxyNetworkGeneration("account-a", input, expected, "127.0.0.3"); err != errGenerationConflict { - t.Fatalf("gateway address change crossed proxy fence: %v", err) - } -} - -func TestGatewayLifecycleUsesInspectedImmutableContainerID(t *testing.T) { - tests := []struct { - method string - path string - dockerPath string - }{ - {http.MethodPost, "/v1/browsers/account-a/start", "/containers/container-id/start"}, - {http.MethodPost, "/v1/browsers/account-a/stop", "/containers/container-id/stop"}, - {http.MethodDelete, "/v1/browsers/account-a", "/containers/container-id"}, - } - for _, test := range tests { - t.Run(test.method+" "+test.path, func(t *testing.T) { - docker, server := testDocker(func(response http.ResponseWriter, request *http.Request) { - if request.Method == http.MethodGet { - networkLabel := "" - if test.path == "/v1/browsers/account-a/start" { - networkLabel = `,"` + networkIDLabel + `":"network-id"` - } - _, _ = response.Write([]byte(`{"Id":"container-id","Config":{"Labels":{"` + managedLabel + `":"true","` + idLabel + `":"account-a","` + bindingVersionLabel + `":"1"` + networkLabel + `}}}`)) - return - } - if request.URL.Path != test.dockerPath { - t.Fatalf("unexpected Docker path %s", request.URL.String()) - } - response.WriteHeader(http.StatusNoContent) - }) - defer server.Close() - - handler := newGateway(docker, "creatorhub_browser", testToken) - response := httptest.NewRecorder() - var body io.Reader - if test.path == "/v1/browsers/account-a/start" { - body = strings.NewReader(`{"binding_version":1,"runtime_id":"container-id","network_id":"network-id"}`) - } else { - body = strings.NewReader(testGenerationBody) - } - adaptor.FiberApp(handler).ServeHTTP(response, authed(test.method, test.path, body)) - if response.Code != http.StatusNoContent { - t.Fatalf("expected 204, got %d: %s", response.Code, response.Body.String()) - } - }) - } -} - -func TestGatewayStartRejectsEmptyNetworkGeneration(t *testing.T) { - mutations := 0 - docker, server := testDocker(func(response http.ResponseWriter, request *http.Request) { - if request.Method != http.MethodGet { - mutations++ - } - _, _ = response.Write([]byte(`{"Id":"container-id","Config":{"Labels":{"` + managedLabel + `":"true","` + idLabel + `":"account-a","` + bindingVersionLabel + `":"1"}}}`)) - }) - defer server.Close() - response := httptest.NewRecorder() - adaptor.FiberApp(newGateway(docker, "creatorhub_browser", testToken)).ServeHTTP(response, - authed(http.MethodPost, "/v1/browsers/account-a/start", strings.NewReader(`{"binding_version":1,"runtime_id":"container-id"}`))) - if response.Code != http.StatusBadRequest || mutations != 0 { - t.Fatalf("empty network generation reached Docker: status=%d mutations=%d body=%s", response.Code, mutations, response.Body.String()) - } -} - -func TestGatewayDeleteUsesImmutableNetworkIDAcrossCleanupRetry(t *testing.T) { - containerExists, cleanupFails, containerDeletes := true, true, 0 - networkExists := true - networkMembers := map[string]any{ - "container-id": map[string]string{"Name": namePrefix + "account-a"}, - "gateway-self": map[string]string{"Name": "gateway-self"}, - } - server := httptest.NewServer(withAliasReservations("gateway-self", func(response http.ResponseWriter, request *http.Request) { - switch { - case request.Method == http.MethodGet && strings.HasPrefix(request.URL.Path, "/containers/"): - if !containerExists { - response.WriteHeader(http.StatusNotFound) - return - } - _, _ = response.Write([]byte(`{"Id":"container-id","Config":{"Labels":{"` + managedLabel + `":"true","` + idLabel + `":"account-a","` + bindingVersionLabel + `":"1","` + networkIDLabel + `":"network-id"}}}`)) - case request.Method == http.MethodDelete && strings.HasPrefix(request.URL.Path, "/containers/"): - containerExists = false - containerDeletes++ - response.WriteHeader(http.StatusNoContent) - case request.Method == http.MethodGet && strings.HasPrefix(request.URL.Path, "/networks/"): - if !networkExists { - response.WriteHeader(http.StatusNotFound) - return - } - _ = json.NewEncoder(response).Encode(map[string]any{ - "Id": "network-id", "Name": "creatorhub_browser-account-a", "Driver": "bridge", "Internal": false, "Attachable": false, "Ingress": false, - "Labels": map[string]string{managedLabel: "true", networkRoleLabel: browserNetworkRole, idLabel: "account-a", bindingVersionLabel: "1"}, - "Containers": networkMembers, - }) - case request.Method == http.MethodPost && strings.HasSuffix(request.URL.Path, "/disconnect"): - if request.URL.Path != "/networks/network-id/disconnect" { - t.Fatalf("network cleanup did not use immutable id: %s", request.URL.Path) - } - if cleanupFails { - response.WriteHeader(http.StatusInternalServerError) - return - } - var body struct { - Container string `json:"Container"` - } - _ = json.NewDecoder(request.Body).Decode(&body) - delete(networkMembers, body.Container) - response.WriteHeader(http.StatusOK) - case request.Method == http.MethodDelete && strings.HasPrefix(request.URL.Path, "/networks/"): - if request.URL.Path != "/networks/network-id" { - t.Fatalf("network delete did not use immutable id: %s", request.URL.Path) - } - networkExists = false - response.WriteHeader(http.StatusNoContent) - default: - t.Fatalf("unexpected Docker request %s %s", request.Method, request.URL.String()) - } - })) - defer server.Close() - handler := newGatewayWithSelf(dockerClient{baseURL: server.URL, client: server.Client(), slow: server.Client()}, - "creatorhub_browser", testToken, "gateway-self") - - response := httptest.NewRecorder() - deleteBody := `{"binding_version":1,"runtime_id":"container-id","network_id":"network-id"}` - adaptor.FiberApp(handler).ServeHTTP(response, authed(http.MethodDelete, "/v1/browsers/account-a", strings.NewReader(deleteBody))) - if response.Code != http.StatusAccepted || !containerExists || containerDeletes != 0 { - t.Fatalf("expected alias reservation with pending cleanup, status=%d exists=%v deletes=%d body=%s", - response.Code, containerExists, containerDeletes, response.Body.String()) - } - cleanupFails = false - response = httptest.NewRecorder() - adaptor.FiberApp(handler).ServeHTTP(response, authed(http.MethodDelete, "/v1/browsers/account-a", strings.NewReader(deleteBody))) - if response.Code != http.StatusNoContent || containerExists || containerDeletes != 1 { - t.Fatalf("idempotent cleanup retry failed: status=%d deletes=%d body=%s", response.Code, containerDeletes, response.Body.String()) - } -} - -func TestGatewayDeleteWithoutContainerOrNetworkGenerationFailsClosed(t *testing.T) { - networkRequests := 0 - server := httptest.NewServer(withAliasReservations("gateway-self", func(response http.ResponseWriter, request *http.Request) { - switch { - case request.Method == http.MethodGet && request.URL.Path == "/containers/"+namePrefix+"account-a/json": - response.WriteHeader(http.StatusNotFound) - case strings.HasPrefix(request.URL.Path, "/networks/"): - networkRequests++ - response.WriteHeader(http.StatusInternalServerError) - default: - t.Fatalf("unexpected Docker request %s %s", request.Method, request.URL.String()) - } - })) - defer server.Close() - handler := newGatewayWithSelf(dockerClient{baseURL: server.URL, client: server.Client(), slow: server.Client()}, - "creatorhub_browser", testToken, "gateway-self") - response := httptest.NewRecorder() - body := `{"binding_version":1,"runtime_id":"runtime-not-found"}` - adaptor.FiberApp(handler).ServeHTTP(response, authed(http.MethodDelete, "/v1/browsers/account-a", strings.NewReader(body))) - if response.Code != http.StatusConflict || networkRequests != 0 { - t.Fatalf("legacy cleanup discovered a replacement network: status=%d networkRequests=%d body=%s", - response.Code, networkRequests, response.Body.String()) - } -} - -func TestGatewayDeleteConvergesWhenGenerationAlreadyGone(t *testing.T) { - networkRequests := 0 - server := httptest.NewServer(withAliasReservations("gateway-self", func(response http.ResponseWriter, request *http.Request) { - switch { - case request.Method == http.MethodGet && request.URL.Path == "/containers/"+namePrefix+"account-a/json": - response.WriteHeader(http.StatusNotFound) - case request.Method == http.MethodGet && request.URL.Path == "/networks/network-old": - networkRequests++ - response.WriteHeader(http.StatusNotFound) - case strings.HasPrefix(request.URL.Path, "/networks/"): - t.Fatalf("network lookup leaked beyond the requested generation: %s %s", request.Method, request.URL.Path) - default: - t.Fatalf("unexpected Docker request %s %s", request.Method, request.URL.String()) - } - })) - defer server.Close() - handler := newGatewayWithSelf(dockerClient{baseURL: server.URL, client: server.Client(), slow: server.Client()}, - "creatorhub_browser", testToken, "gateway-self") - response := httptest.NewRecorder() - body := `{"binding_version":1,"runtime_id":"container-old","network_id":"network-old"}` - adaptor.FiberApp(handler).ServeHTTP(response, authed(http.MethodDelete, "/v1/browsers/account-a", strings.NewReader(body))) - if response.Code != http.StatusNoContent || networkRequests != 1 { - t.Fatalf("already-gone generation did not converge: status=%d networkRequests=%d body=%s", - response.Code, networkRequests, response.Body.String()) - } -} - -func TestGatewayDeleteFailsClosedWhenContainerOutlivesNetwork(t *testing.T) { - containerDeleted := false - server := httptest.NewServer(withAliasReservations("gateway-self", func(response http.ResponseWriter, request *http.Request) { - switch { - case request.Method == http.MethodGet && request.URL.Path == "/containers/"+namePrefix+"account-a/json": - _, _ = response.Write([]byte(`{"Id":"container-old","Config":{"Labels":{"` + managedLabel + `":"true","` + idLabel + `":"account-a","` + bindingVersionLabel + `":"1","` + networkIDLabel + `":"network-old"}}}`)) - case request.Method == http.MethodGet && request.URL.Path == "/networks/network-old": - response.WriteHeader(http.StatusNotFound) - case request.Method == http.MethodDelete && strings.HasPrefix(request.URL.Path, "/containers/"): - containerDeleted = true - response.WriteHeader(http.StatusNoContent) - default: - t.Fatalf("unexpected Docker request %s %s", request.Method, request.URL.String()) - } - })) - defer server.Close() - handler := newGatewayWithSelf(dockerClient{baseURL: server.URL, client: server.Client(), slow: server.Client()}, - "creatorhub_browser", testToken, "gateway-self") - response := httptest.NewRecorder() - body := `{"binding_version":1,"runtime_id":"container-old","network_id":"network-old"}` - adaptor.FiberApp(handler).ServeHTTP(response, authed(http.MethodDelete, "/v1/browsers/account-a", strings.NewReader(body))) - if response.Code != http.StatusConflict || containerDeleted { - t.Fatalf("container outliving its network must fail closed: status=%d containerDeleted=%v body=%s", - response.Code, containerDeleted, response.Body.String()) - } -} - -func TestGatewayRejectsStaleProxyRestoreAfterReplacementGeneration(t *testing.T) { - type dockerState struct { - sync.Mutex - containerID string - containerLabels map[string]string - networkID string - networkMembers map[string]string - containerReads int - cleanupMutations []string - } - state := &dockerState{ - containerID: "container-c1", - containerLabels: map[string]string{managedLabel: "true", idLabel: "account-a", bindingVersionLabel: "1", networkExitLabel: "exit-1", proxyPortLabel: "12345", networkIDLabel: "network-n1"}, - networkID: "network-n1", - networkMembers: map[string]string{"container-c1": namePrefix + "account-a", "gateway-self": "gateway-self"}, - } - r1Captured := make(chan struct{}) - resumeR1 := make(chan struct{}) - var releaseOnce sync.Once - release := func() { releaseOnce.Do(func() { close(resumeR1) }) } - defer release() - dockerServer := httptest.NewServer(withAliasReservations("gateway-self", func(response http.ResponseWriter, request *http.Request) { - switch { - case request.Method == http.MethodGet && strings.HasPrefix(request.URL.Path, "/images/"): - response.WriteHeader(http.StatusOK) - case request.Method == http.MethodGet && request.URL.Path == "/containers/"+namePrefix+"account-a/json": - state.Lock() - containerID := state.containerID - labels := state.containerLabels - state.containerReads++ - first := state.containerReads == 1 - state.Unlock() - if first { - close(r1Captured) - <-resumeR1 - } - if containerID == "" { - response.WriteHeader(http.StatusNotFound) - return - } - _ = json.NewEncoder(response).Encode(map[string]any{"Id": containerID, "Config": map[string]any{"Labels": labels}}) - case request.Method == http.MethodGet && (request.URL.Path == "/networks/creatorhub_browser-account-a" || - request.URL.Path == "/networks/network-n1" || request.URL.Path == "/networks/network-n2"): - state.Lock() - networkID := state.networkID - members := make(map[string]map[string]string, len(state.networkMembers)) - for id, name := range state.networkMembers { - members[id] = map[string]string{"Name": name, "IPv4Address": "127.0.0.1/8"} - } - state.Unlock() - if networkID == "" { - response.WriteHeader(http.StatusNotFound) - return - } - _ = json.NewEncoder(response).Encode(map[string]any{ - "Id": networkID, "Name": "creatorhub_browser-account-a", "Driver": "bridge", - "Internal": false, "Attachable": false, "Ingress": false, - "Labels": map[string]string{managedLabel: "true", networkRoleLabel: browserNetworkRole, idLabel: "account-a", bindingVersionLabel: "1"}, - "Containers": members, - }) - case request.Method == http.MethodPost && strings.HasSuffix(request.URL.Path, "/disconnect"): - var body struct { - Container string `json:"Container"` - } - _ = json.NewDecoder(request.Body).Decode(&body) - requestedNetwork := strings.TrimSuffix(strings.TrimPrefix(request.URL.Path, "/networks/"), "/disconnect") - state.Lock() - if requestedNetwork != state.networkID { - state.Unlock() - response.WriteHeader(http.StatusNotFound) - return - } - state.cleanupMutations = append(state.cleanupMutations, "disconnect:"+requestedNetwork+":"+body.Container) - delete(state.networkMembers, body.Container) - state.Unlock() - response.WriteHeader(http.StatusOK) - case request.Method == http.MethodDelete && strings.HasPrefix(request.URL.Path, "/networks/"): - requestedNetwork := strings.TrimPrefix(request.URL.Path, "/networks/") - state.Lock() - if requestedNetwork != state.networkID { - state.Unlock() - response.WriteHeader(http.StatusNotFound) - return - } - state.cleanupMutations = append(state.cleanupMutations, "delete-network:"+requestedNetwork) - state.networkID = "" - state.networkMembers = nil - state.Unlock() - response.WriteHeader(http.StatusNoContent) - case request.Method == http.MethodDelete && strings.HasPrefix(request.URL.Path, "/containers/"): - requestedContainer := strings.TrimPrefix(request.URL.Path, "/containers/") - state.Lock() - if requestedContainer != state.containerID { - state.Unlock() - response.WriteHeader(http.StatusNotFound) - return - } - state.cleanupMutations = append(state.cleanupMutations, "delete-container:"+requestedContainer) - state.containerID = "" - state.containerLabels = nil - state.Unlock() - response.WriteHeader(http.StatusNoContent) - case request.Method == http.MethodPost && request.URL.Path == "/networks/create": - state.Lock() - state.networkID = "network-n2" - state.networkMembers = map[string]string{} - state.Unlock() - response.WriteHeader(http.StatusCreated) - _, _ = response.Write([]byte(`{"Id":"network-n2"}`)) - case request.Method == http.MethodPost && strings.HasSuffix(request.URL.Path, "/connect"): - state.Lock() - state.networkMembers["gateway-self"] = "gateway-self" - state.Unlock() - response.WriteHeader(http.StatusOK) - case request.Method == http.MethodPost && request.URL.Path == "/containers/create": - var payload struct { - Labels map[string]string `json:"Labels"` - } - _ = json.NewDecoder(request.Body).Decode(&payload) - state.Lock() - state.containerID = "container-c2" - state.containerLabels = payload.Labels - state.networkMembers["container-c2"] = namePrefix + "account-a" - state.Unlock() - response.WriteHeader(http.StatusCreated) - _, _ = response.Write([]byte(`{"Id":"container-c2"}`)) - case request.Method == http.MethodPost && request.URL.Path == "/containers/container-c2/start": - response.WriteHeader(http.StatusNoContent) - default: - t.Errorf("unexpected Docker request %s %s", request.Method, request.URL.String()) - response.WriteHeader(http.StatusInternalServerError) - } - })) - defer dockerServer.Close() - - registry := newMemoryProxyRegistry() - proxyURL, _, err := registry.configure("account-a", 1, "127.0.0.1", 0, - gatewayProxyExit{Protocol: "socks5", Host: "proxy.example", Port: 1080}) - if err != nil || !registry.bind("account-a", 1, proxyURL, "container-c1", "network-n1") { - t.Fatalf("seed C1 proxy: %v", err) - } - defer func() { - registry.remove("account-a", 1, "container-c2") - registry.remove("account-a", 1, "container-c1") - }() - api := gateway{ - docker: dockerClient{baseURL: dockerServer.URL, client: dockerServer.Client(), slow: dockerServer.Client()}, - network: "creatorhub_browser", self: "gateway-self", token: testToken, proxies: registry, - } - api.locks = &dockerAliasReservations{docker: api.docker, self: api.self} - app := fiber.New() - app.Use("/v1", api.authorize) - app.Post("/v1/browsers", api.create) - app.Post("/v1/browsers/:id/proxy", api.restoreProxy) - app.Delete("/v1/browsers/:id", api.remove) - gatewayServer := httptest.NewServer(adaptor.FiberApp(app)) - defer gatewayServer.Close() - defer release() - - type result struct { - status int - body string - } - call := func(method, path, body string) result { - request, _ := http.NewRequest(method, gatewayServer.URL+path, strings.NewReader(body)) - request.Header.Set("Authorization", "Bearer "+testToken) - response, err := gatewayServer.Client().Do(request) - if err != nil { - return result{body: err.Error()} - } - defer response.Body.Close() - responseBody, _ := io.ReadAll(response.Body) - return result{status: response.StatusCode, body: string(responseBody)} - } - r1Result := make(chan result, 1) - go func() { - r1Result <- call(http.MethodPost, "/v1/browsers/account-a/proxy", - `{"binding_version":1,"runtime_id":"container-c1","network_id":"network-n1","network_exit_id":"exit-1","network_exit":{"protocol":"socks5","host":"proxy.example","port":1080}}`) - }() - select { - case <-r1Captured: - case <-time.After(5 * time.Second): - t.Fatal("R1 did not capture C1/N1") - } - - r2Delete := call(http.MethodDelete, "/v1/browsers/account-a", `{"binding_version":1,"runtime_id":"container-c1","network_id":"network-n1"}`) - if r2Delete.status != http.StatusNoContent { - t.Fatalf("R2 cleanup failed: status=%d body=%s", r2Delete.status, r2Delete.body) - } - r2Create := call(http.MethodPost, "/v1/browsers", testCreateBody) - if r2Create.status != http.StatusCreated { - t.Fatalf("R2 replacement create failed: status=%d body=%s", r2Create.status, r2Create.body) - } - release() - select { - case stale := <-r1Result: - if stale.status != http.StatusConflict { - t.Fatalf("R1 stale restore was not fenced: status=%d body=%s", stale.status, stale.body) - } - case <-time.After(5 * time.Second): - t.Fatal("R1 did not finish") - } - - state.Lock() - defer state.Unlock() - proxy := registry.proxies["account-a"] - wantCleanup := []string{ - "disconnect:network-n1:container-c1", - "disconnect:network-n1:gateway-self", - "delete-network:network-n1", - "delete-container:container-c1", - } - if state.containerID != "container-c2" || state.networkID != "network-n2" || state.networkMembers["container-c2"] == "" || - proxy == nil || proxy.runtimeID != "container-c2" || strings.Join(state.cleanupMutations, ",") != strings.Join(wantCleanup, ",") { - t.Fatalf("stale R1 affected replacement generation: container=%q network=%q members=%v proxy=%v mutations=%v", - state.containerID, state.networkID, state.networkMembers, proxy, state.cleanupMutations) - } -} - -func TestGatewayRejectsStaleCreateBeforeNetworkOrProxyMutation(t *testing.T) { - type dockerState struct { - sync.Mutex - containerID string - containerLabels map[string]string - containerReads int - networkExists bool - networkCreates int - networkConnects int - gatewayConnected bool - containerCreates int - } - state := &dockerState{} - r1Inspected := make(chan struct{}) - resumeR1 := make(chan struct{}) - var releaseOnce sync.Once - release := func() { releaseOnce.Do(func() { close(resumeR1) }) } - defer release() - dockerServer := httptest.NewServer(withAliasReservations("gateway-self", func(response http.ResponseWriter, request *http.Request) { - switch { - case request.Method == http.MethodGet && strings.HasPrefix(request.URL.Path, "/images/"): - response.WriteHeader(http.StatusOK) - case request.Method == http.MethodGet && request.URL.Path == "/containers/"+namePrefix+"account-a/json": - state.Lock() - containerID, labels := state.containerID, state.containerLabels - state.containerReads++ - first := state.containerReads == 1 - state.Unlock() - if first { - close(r1Inspected) - <-resumeR1 - } - if containerID == "" { - response.WriteHeader(http.StatusNotFound) - return - } - _ = json.NewEncoder(response).Encode(map[string]any{"Id": containerID, "Config": map[string]any{"Labels": labels}}) - case request.Method == http.MethodGet && (request.URL.Path == "/networks/creatorhub_browser-account-a" || request.URL.Path == "/networks/network-n1"): - state.Lock() - exists, connected := state.networkExists, state.gatewayConnected - state.Unlock() - if !exists { - response.WriteHeader(http.StatusNotFound) - return - } - members := map[string]any{} - if connected { - members["gateway-self"] = map[string]string{"Name": "gateway-self", "IPv4Address": "127.0.0.1/8"} - } - _ = json.NewEncoder(response).Encode(map[string]any{ - "Id": "network-n1", "Name": "creatorhub_browser-account-a", "Driver": "bridge", "Internal": false, "Attachable": false, "Ingress": false, - "Labels": map[string]string{managedLabel: "true", networkRoleLabel: browserNetworkRole, idLabel: "account-a", bindingVersionLabel: "1"}, - "Containers": members, - }) - case request.Method == http.MethodPost && request.URL.Path == "/networks/create": - state.Lock() - state.networkExists = true - state.networkCreates++ - state.Unlock() - response.WriteHeader(http.StatusCreated) - _, _ = response.Write([]byte(`{"Id":"network-n1"}`)) - case request.Method == http.MethodPost && strings.HasSuffix(request.URL.Path, "/connect"): - state.Lock() - state.networkConnects++ - state.gatewayConnected = true - state.Unlock() - response.WriteHeader(http.StatusOK) - case request.Method == http.MethodPost && request.URL.Path == "/containers/create": - var payload struct { - Labels map[string]string `json:"Labels"` - } - _ = json.NewDecoder(request.Body).Decode(&payload) - state.Lock() - state.containerID = "container-c1" - state.containerLabels = payload.Labels - state.containerCreates++ - state.Unlock() - response.WriteHeader(http.StatusCreated) - _, _ = response.Write([]byte(`{"Id":"container-c1"}`)) - case request.Method == http.MethodPost && request.URL.Path == "/containers/container-c1/start": - response.WriteHeader(http.StatusNoContent) - default: - t.Errorf("unexpected Docker request %s %s", request.Method, request.URL.String()) - response.WriteHeader(http.StatusInternalServerError) - } - })) - defer dockerServer.Close() - docker := dockerClient{baseURL: dockerServer.URL, client: dockerServer.Client(), slow: dockerServer.Client()} - registry := newMemoryProxyRegistry() - defer registry.remove("account-a", 1, "container-c1") - api := gateway{docker: docker, network: "creatorhub_browser", self: "gateway-self", token: testToken, proxies: registry, - locks: &dockerAliasReservations{docker: docker, self: "gateway-self"}} - app := fiber.New() - app.Use("/v1", api.authorize) - app.Post("/v1/browsers", api.create) - gatewayServer := httptest.NewServer(adaptor.FiberApp(app)) - defer gatewayServer.Close() - - call := func() (int, string) { - request, _ := http.NewRequest(http.MethodPost, gatewayServer.URL+"/v1/browsers", strings.NewReader(testCreateBody)) - request.Header.Set("Authorization", "Bearer "+testToken) - response, err := gatewayServer.Client().Do(request) - if err != nil { - return 0, err.Error() - } - defer response.Body.Close() - body, _ := io.ReadAll(response.Body) - return response.StatusCode, string(body) - } - r1Result := make(chan struct { - status int - body string - }, 1) - go func() { - status, body := call() - r1Result <- struct { - status int - body string - }{status, body} - }() - select { - case <-r1Inspected: - case <-time.After(5 * time.Second): - t.Fatal("R1 did not inspect the empty alias") - } - status, body := call() - if status != http.StatusCreated { - release() - t.Fatalf("R2 create failed: status=%d body=%s", status, body) - } - release() - select { - case stale := <-r1Result: - if stale.status != http.StatusConflict { - t.Fatalf("R1 stale create was not fenced: status=%d body=%s", stale.status, stale.body) - } - case <-time.After(5 * time.Second): - t.Fatal("R1 did not finish") - } - - state.Lock() - defer state.Unlock() - proxy := registry.proxies["account-a"] - if state.containerID != "container-c1" || state.networkCreates != 1 || state.networkConnects != 1 || state.containerCreates != 1 || - proxy == nil || proxy.runtimeID != "container-c1" { - t.Fatalf("stale create left side effects: container=%q networkCreates=%d connects=%d containerCreates=%d proxy=%v", - state.containerID, state.networkCreates, state.networkConnects, state.containerCreates, proxy) - } -} - -func TestGatewayTwoReplicaCreateRestoreRemoveProxyContract(t *testing.T) { - type member struct{ name, ip string } - type dockerState struct { - sync.Mutex - containerID string - containerLabels map[string]string - networkID string - networkMembers map[string]member - sequence int - } - state := &dockerState{} - dockerServer := httptest.NewServer(withAliasReservations("gateway-g1", func(response http.ResponseWriter, request *http.Request) { - switch { - case request.Method == http.MethodGet && strings.HasPrefix(request.URL.Path, "/images/"): - response.WriteHeader(http.StatusOK) - case request.Method == http.MethodGet && request.URL.Path == "/containers/gateway-g2/json": - _, _ = response.Write([]byte(`{"Id":"gateway-g2","Image":"gateway-image-id","Config":{"Labels":{"` + gatewayMemberLabel + `":"true"}}}`)) - case request.Method == http.MethodGet && request.URL.Path == "/containers/"+namePrefix+"account-a/json": - state.Lock() - id, labels := state.containerID, state.containerLabels - state.Unlock() - if id == "" { - response.WriteHeader(http.StatusNotFound) - return - } - _ = json.NewEncoder(response).Encode(map[string]any{"Id": id, "Config": map[string]any{"Labels": labels}}) - case request.Method == http.MethodGet && strings.HasPrefix(request.URL.Path, "/networks/"): - state.Lock() - id := state.networkID - members := map[string]any{} - for memberID, current := range state.networkMembers { - members[memberID] = map[string]string{"Name": current.name, "IPv4Address": current.ip} - } - state.Unlock() - if id == "" { - response.WriteHeader(http.StatusNotFound) - return - } - reference := strings.TrimPrefix(request.URL.Path, "/networks/") - if reference != "creatorhub_browser-account-a" && reference != id { - response.WriteHeader(http.StatusNotFound) - return - } - _ = json.NewEncoder(response).Encode(map[string]any{ - "Id": id, "Name": "creatorhub_browser-account-a", "Driver": "bridge", "Internal": false, "Attachable": false, "Ingress": false, - "Labels": map[string]string{managedLabel: "true", networkRoleLabel: browserNetworkRole, idLabel: "account-a", bindingVersionLabel: "1"}, - "Containers": members, - }) - case request.Method == http.MethodPost && request.URL.Path == "/networks/create": - state.Lock() - state.sequence++ - state.networkID = fmt.Sprintf("network-n%d", state.sequence) - state.networkMembers = map[string]member{} - id := state.networkID - state.Unlock() - response.WriteHeader(http.StatusCreated) - _ = json.NewEncoder(response).Encode(map[string]string{"Id": id}) - case request.Method == http.MethodPost && strings.HasSuffix(request.URL.Path, "/connect"): - var body struct { - Container string `json:"Container"` - } - _ = json.NewDecoder(request.Body).Decode(&body) - state.Lock() - ip := "127.0.0.3/8" - if body.Container == "gateway-g1" { - ip = "127.0.0.1/8" - } else if body.Container == "gateway-g2" { - ip = "127.0.0.2/8" - } - state.networkMembers[body.Container] = member{name: body.Container, ip: ip} - state.Unlock() - response.WriteHeader(http.StatusOK) - case request.Method == http.MethodPost && request.URL.Path == "/containers/create": - var payload struct { - Labels map[string]string `json:"Labels"` - } - _ = json.NewDecoder(request.Body).Decode(&payload) - state.Lock() - id := fmt.Sprintf("container-c%d", state.sequence) - state.containerID, state.containerLabels = id, payload.Labels - state.networkMembers[id] = member{name: namePrefix + "account-a", ip: "127.0.0.3/8"} - state.Unlock() - response.WriteHeader(http.StatusCreated) - _ = json.NewEncoder(response).Encode(map[string]string{"Id": id}) - case request.Method == http.MethodPost && strings.HasSuffix(request.URL.Path, "/start"): - response.WriteHeader(http.StatusNoContent) - case request.Method == http.MethodPost && strings.HasSuffix(request.URL.Path, "/disconnect"): - var body struct { - Container string `json:"Container"` - } - _ = json.NewDecoder(request.Body).Decode(&body) - state.Lock() - delete(state.networkMembers, body.Container) - state.Unlock() - response.WriteHeader(http.StatusOK) - case request.Method == http.MethodDelete && strings.HasPrefix(request.URL.Path, "/networks/"): - state.Lock() - if len(state.networkMembers) != 0 { - state.Unlock() - response.WriteHeader(http.StatusConflict) - return - } - state.networkID = "" - state.Unlock() - response.WriteHeader(http.StatusNoContent) - case request.Method == http.MethodDelete && strings.HasPrefix(request.URL.Path, "/containers/"): - id := strings.TrimPrefix(request.URL.Path, "/containers/") - state.Lock() - if id != state.containerID { - state.Unlock() - response.WriteHeader(http.StatusNotFound) - return - } - delete(state.networkMembers, id) - state.containerID, state.containerLabels = "", nil - state.Unlock() - response.WriteHeader(http.StatusNoContent) - default: - t.Fatalf("unexpected Docker request %s %s", request.Method, request.URL.String()) - } - })) - defer dockerServer.Close() - docker := dockerClient{baseURL: dockerServer.URL, client: dockerServer.Client(), slow: dockerServer.Client()} - newReplica := func(self string) (*fiber.App, *memoryProxyRegistry) { - registry := newMemoryProxyRegistry() - api := gateway{docker: docker, network: "creatorhub_browser", self: self, token: testToken, proxies: registry, - locks: &dockerAliasReservations{docker: docker, self: self}} - app := fiber.New() - app.Use("/v1", api.authorize) - app.Post("/v1/browsers", api.create) - app.Post("/v1/browsers/:id/proxy", api.restoreProxy) - app.Delete("/v1/browsers/:id", api.remove) - return app, registry - } - g1, proxiesG1 := newReplica("gateway-g1") - g2, proxiesG2 := newReplica("gateway-g2") - call := func(app *fiber.App, method, path, body string) *httptest.ResponseRecorder { - response := httptest.NewRecorder() - adaptor.FiberApp(app).ServeHTTP(response, authed(method, path, strings.NewReader(body))) - return response - } - created := call(g1, http.MethodPost, "/v1/browsers", testCreateBody) - if created.Code != http.StatusCreated { - t.Fatalf("G1 create failed: %d %s", created.Code, created.Body.String()) - } - state.Lock() - c1, n1, port := state.containerID, state.networkID, state.containerLabels[proxyPortLabel] - state.Unlock() - restoreBody := `{"binding_version":1,"runtime_id":"` + c1 + `","network_id":"network-n1","network_exit_id":"exit-1","network_exit":{"protocol":"socks5","host":"proxy.example","port":1080}}` - if restored := call(g2, http.MethodPost, "/v1/browsers/account-a/proxy", restoreBody); restored.Code != http.StatusNoContent { - t.Fatalf("G2 restore failed: %d %s", restored.Code, restored.Body.String()) - } - if proxiesG1.proxies["account-a"].runtimeID != c1 || proxiesG2.proxies["account-a"].runtimeID != c1 || port == "" { - t.Fatal("both replicas did not bind the same runtime generation") - } - removeBody := `{"binding_version":1,"runtime_id":"` + c1 + `","network_id":"` + n1 + `"}` - if removed := call(g2, http.MethodDelete, "/v1/browsers/account-a", removeBody); removed.Code != http.StatusNoContent { - t.Fatalf("G2 remove failed: %d %s", removed.Code, removed.Body.String()) - } - if replacement := call(g2, http.MethodPost, "/v1/browsers", testCreateBody); replacement.Code != http.StatusCreated { - t.Fatalf("G2 replacement create failed: %d %s", replacement.Code, replacement.Body.String()) - } - state.Lock() - c2 := state.containerID - _, g1Attached := state.networkMembers["gateway-g1"] - _, g2Attached := state.networkMembers["gateway-g2"] - _, c2Attached := state.networkMembers[c2] - state.Unlock() - if c2 == c1 || g1Attached || !g2Attached || !c2Attached || proxiesG1.proxies["account-a"].runtimeID != c1 || - proxiesG2.proxies["account-a"].runtimeID != c2 { - t.Fatalf("cross-process proxy release contract failed: c1=%q c2=%q g1=%v g2=%v c2Attached=%v", c1, c2, g1Attached, g2Attached, c2Attached) - } - state.Lock() - n2 := state.networkID - port = state.containerLabels[proxyPortLabel] - state.Unlock() - restoreBody = `{"binding_version":1,"runtime_id":"` + c2 + `","network_id":"` + n2 + `","network_exit_id":"exit-1","network_exit":{"protocol":"socks5","host":"proxy.example","port":1080}}` - if restored := call(g1, http.MethodPost, "/v1/browsers/account-a/proxy", restoreBody); restored.Code != http.StatusNoContent { - t.Fatalf("G1 successor restore failed: %d %s", restored.Code, restored.Body.String()) - } - state.Lock() - _, g1Attached = state.networkMembers["gateway-g1"] - state.Unlock() - if !g1Attached || proxiesG1.proxies["account-a"].runtimeID != c2 || proxiesG2.proxies["account-a"].runtimeID != c2 { - t.Fatalf("stale replica did not replace its proxy generation: c1=%q c2=%q g1=%v g1Proxy=%v g2Proxy=%v", - c1, c2, g1Attached, proxiesG1.proxies["account-a"], proxiesG2.proxies["account-a"]) - } - proxiesG1.remove("account-a", 1, c1) - proxiesG2.remove("account-a", 1, c2) -} - -func TestGatewayRemoveFencesNetworkReplacementAndMemberChanges(t *testing.T) { - for _, test := range []struct { - name string - replaceOnDisconnect bool - addMemberOnDisconnect bool - wantNetworkID string - }{ - {name: "N1 replaced by N2 after inspect", replaceOnDisconnect: true, wantNetworkID: "network-n2"}, - {name: "trusted member joins before delete", addMemberOnDisconnect: true, wantNetworkID: "network-n1"}, - } { - t.Run(test.name, func(t *testing.T) { - var mu sync.Mutex - networkID := "network-n1" - members := map[string]string{"container-c1": namePrefix + "account-a", "gateway-g1": "gateway-g1"} - containerDeleted, networkDeletes := false, 0 - server := httptest.NewServer(withAliasReservations("gateway-g1", func(response http.ResponseWriter, request *http.Request) { - switch { - case request.Method == http.MethodGet && request.URL.Path == "/containers/"+namePrefix+"account-a/json": - _, _ = response.Write([]byte(`{"Id":"container-c1","Config":{"Labels":{"` + managedLabel + `":"true","` + idLabel + `":"account-a","` + bindingVersionLabel + `":"1","` + networkIDLabel + `":"network-n1"}}}`)) - case request.Method == http.MethodGet && request.URL.Path == "/containers/gateway-g2/json": - _, _ = response.Write([]byte(`{"Id":"gateway-g2","Config":{"Labels":{"` + gatewayMemberLabel + `":"true"}}}`)) - case request.Method == http.MethodGet && (request.URL.Path == "/networks/creatorhub_browser-account-a" || - request.URL.Path == "/networks/network-n1" || request.URL.Path == "/networks/network-n2"): - mu.Lock() - id := networkID - current := map[string]any{} - for memberID, name := range members { - current[memberID] = map[string]string{"Name": name} - } - mu.Unlock() - _ = json.NewEncoder(response).Encode(map[string]any{ - "Id": id, "Name": "creatorhub_browser-account-a", "Driver": "bridge", "Internal": false, "Attachable": false, "Ingress": false, - "Labels": map[string]string{managedLabel: "true", networkRoleLabel: browserNetworkRole, idLabel: "account-a", bindingVersionLabel: "1"}, - "Containers": current, - }) - case request.Method == http.MethodPost && strings.HasSuffix(request.URL.Path, "/disconnect"): - var body struct { - Container string `json:"Container"` - } - _ = json.NewDecoder(request.Body).Decode(&body) - mu.Lock() - if test.replaceOnDisconnect { - networkID = "network-n2" - mu.Unlock() - response.WriteHeader(http.StatusNotFound) - return - } - delete(members, body.Container) - if test.addMemberOnDisconnect && body.Container == "gateway-g1" { - members["gateway-g2"] = "gateway-g2" - } - mu.Unlock() - response.WriteHeader(http.StatusOK) - case request.Method == http.MethodDelete && strings.HasPrefix(request.URL.Path, "/networks/"): - mu.Lock() - networkDeletes++ - mu.Unlock() - response.WriteHeader(http.StatusNoContent) - case request.Method == http.MethodDelete && request.URL.Path == "/containers/container-c1": - mu.Lock() - containerDeleted = true - mu.Unlock() - response.WriteHeader(http.StatusNoContent) - default: - t.Fatalf("unexpected Docker request %s %s", request.Method, request.URL.String()) - } - })) - defer server.Close() - docker := dockerClient{baseURL: server.URL, client: server.Client(), slow: server.Client()} - registry := newMemoryProxyRegistry() - proxyURL, _, err := registry.configure("account-a", 1, "127.0.0.1", 0, gatewayProxyExit{Protocol: "socks5", Host: "proxy.example", Port: 1080}) - if err != nil || !registry.bind("account-a", 1, proxyURL, "container-c1") { - t.Fatal("seed proxy generation") - } - defer registry.remove("account-a", 1, "container-c1") - api := gateway{docker: docker, network: "creatorhub_browser", self: "gateway-g1", token: testToken, proxies: registry, - locks: &dockerAliasReservations{docker: docker, self: "gateway-g1"}} - app := fiber.New() - app.Use("/v1", api.authorize) - app.Delete("/v1/browsers/:id", api.remove) - response := httptest.NewRecorder() - adaptor.FiberApp(app).ServeHTTP(response, authed(http.MethodDelete, "/v1/browsers/account-a", - strings.NewReader(`{"binding_version":1,"runtime_id":"container-c1","network_id":"network-n1"}`))) - mu.Lock() - defer mu.Unlock() - if response.Code != http.StatusConflict || networkID != test.wantNetworkID || containerDeleted || networkDeletes != 0 || - registry.proxies["account-a"] == nil { - t.Fatalf("remove crossed network fence: status=%d network=%q containerDeleted=%v networkDeletes=%d proxy=%v body=%s", - response.Code, networkID, containerDeleted, networkDeletes, registry.proxies["account-a"], response.Body.String()) - } - }) - } -} - -func TestGatewayCreateFailureRemovesCreatedNetworkGeneration(t *testing.T) { - for _, failure := range []string{"network-create-id", "configure", "container-create", "start"} { - t.Run(failure, func(t *testing.T) { - var mu sync.Mutex - networkExists, containerExists := false, false - members := map[string]string{} - networkDeletes := 0 - server := httptest.NewServer(withAliasReservations("gateway-self", func(response http.ResponseWriter, request *http.Request) { - switch { - case request.Method == http.MethodGet && strings.HasPrefix(request.URL.Path, "/images/"): - response.WriteHeader(http.StatusOK) - case request.Method == http.MethodGet && request.URL.Path == "/containers/"+namePrefix+"account-a/json": - mu.Lock() - exists := containerExists - mu.Unlock() - if !exists { - response.WriteHeader(http.StatusNotFound) - return - } - _, _ = response.Write([]byte(`{"Id":"container-c1","Config":{"Labels":{"` + managedLabel + `":"true","` + idLabel + `":"account-a","` + bindingVersionLabel + `":"1","` + networkIDLabel + `":"network-n1"}}}`)) - case request.Method == http.MethodGet && (request.URL.Path == "/networks/creatorhub_browser-account-a" || request.URL.Path == "/networks/network-n1"): - mu.Lock() - exists := networkExists - current := map[string]any{} - for id, name := range members { - ip := "127.0.0.1/8" - if failure == "configure" && id == "gateway-self" { - ip = "192.0.2.1/24" - } - current[id] = map[string]string{"Name": name, "IPv4Address": ip} - } - mu.Unlock() - if !exists { - response.WriteHeader(http.StatusNotFound) - return - } - _ = json.NewEncoder(response).Encode(map[string]any{ - "Id": "network-n1", "Name": "creatorhub_browser-account-a", "Driver": "bridge", "Internal": false, "Attachable": false, "Ingress": false, - "Labels": map[string]string{managedLabel: "true", networkRoleLabel: browserNetworkRole, idLabel: "account-a", bindingVersionLabel: "1"}, - "Containers": current, - }) - case request.Method == http.MethodPost && request.URL.Path == "/networks/create": - mu.Lock() - networkExists = true - mu.Unlock() - response.WriteHeader(http.StatusCreated) - if failure == "network-create-id" { - _, _ = response.Write([]byte(`{"Id":""}`)) - } else { - _, _ = response.Write([]byte(`{"Id":"network-n1"}`)) - } - case request.Method == http.MethodPost && strings.HasSuffix(request.URL.Path, "/connect") && - !strings.HasSuffix(request.URL.Path, "/disconnect"): - var body struct { - Container string `json:"Container"` - } - _ = json.NewDecoder(request.Body).Decode(&body) - mu.Lock() - members[body.Container] = body.Container - mu.Unlock() - response.WriteHeader(http.StatusOK) - case request.Method == http.MethodPost && request.URL.Path == "/containers/create": - if failure == "container-create" { - response.WriteHeader(http.StatusInternalServerError) - return - } - mu.Lock() - containerExists = true - members["container-c1"] = namePrefix + "account-a" - mu.Unlock() - response.WriteHeader(http.StatusCreated) - _, _ = response.Write([]byte(`{"Id":"container-c1"}`)) - case request.Method == http.MethodPost && request.URL.Path == "/containers/container-c1/start": - if failure == "start" { - response.WriteHeader(http.StatusInternalServerError) - return - } - response.WriteHeader(http.StatusNoContent) - case request.Method == http.MethodDelete && request.URL.Path == "/containers/container-c1": - mu.Lock() - containerExists = false - delete(members, "container-c1") - mu.Unlock() - response.WriteHeader(http.StatusNoContent) - case request.Method == http.MethodPost && strings.HasSuffix(request.URL.Path, "/disconnect"): - var body struct { - Container string `json:"Container"` - } - _ = json.NewDecoder(request.Body).Decode(&body) - mu.Lock() - delete(members, body.Container) - mu.Unlock() - response.WriteHeader(http.StatusOK) - case request.Method == http.MethodDelete && request.URL.Path == "/networks/network-n1": - mu.Lock() - networkExists = false - networkDeletes++ - mu.Unlock() - response.WriteHeader(http.StatusNoContent) - default: - t.Fatalf("unexpected Docker request %s %s", request.Method, request.URL.String()) - } - })) - defer server.Close() - docker := dockerClient{baseURL: server.URL, client: server.Client(), slow: server.Client()} - handler := newGatewayWithSelf(docker, "creatorhub_browser", testToken, "gateway-self") - response := httptest.NewRecorder() - adaptor.FiberApp(handler).ServeHTTP(response, authed(http.MethodPost, "/v1/browsers", strings.NewReader(testCreateBody))) - mu.Lock() - defer mu.Unlock() - wantNetwork, wantDeletes := false, 1 - if failure == "network-create-id" { - wantNetwork, wantDeletes = true, 0 - } - if response.Code != http.StatusBadGateway || networkExists != wantNetwork || networkDeletes != wantDeletes || containerExists { - t.Fatalf("%s failure left managed resources: status=%d network=%v deletes=%d container=%v body=%s", - failure, response.Code, networkExists, networkDeletes, containerExists, response.Body.String()) - } - }) - } -} - -func TestGatewayNetworkCreateDisconnectDoesNotDiscoverReplacementByName(t *testing.T) { - networkReads, networkDeletes := 0, 0 - networkCreated := false - server := httptest.NewServer(withAliasReservations("gateway-self", func(response http.ResponseWriter, request *http.Request) { - switch { - case request.Method == http.MethodGet && strings.HasPrefix(request.URL.Path, "/images/"): - response.WriteHeader(http.StatusOK) - case request.Method == http.MethodGet && request.URL.Path == "/containers/"+namePrefix+"account-a/json": - response.WriteHeader(http.StatusNotFound) - case request.Method == http.MethodGet && strings.HasPrefix(request.URL.Path, "/networks/"): - networkReads++ - if networkCreated { - _, _ = response.Write([]byte(`{"Id":"network-n2","Name":"creatorhub_browser-account-a","Driver":"bridge","Labels":{"` + managedLabel + `":"true","` + networkRoleLabel + `":"` + browserNetworkRole + `","` + idLabel + `":"account-a","` + bindingVersionLabel + `":"1"},"Containers":{}}`)) - return - } - response.WriteHeader(http.StatusNotFound) - case request.Method == http.MethodPost && request.URL.Path == "/networks/create": - networkCreated = true - connection, _, _ := response.(http.Hijacker).Hijack() - _ = connection.Close() - case request.Method == http.MethodDelete && strings.HasPrefix(request.URL.Path, "/networks/"): - networkDeletes++ - response.WriteHeader(http.StatusNoContent) - default: - t.Fatalf("unexpected Docker request %s %s", request.Method, request.URL.String()) - } - })) - defer server.Close() - handler := newGatewayWithSelf(dockerClient{baseURL: server.URL, client: server.Client(), slow: server.Client()}, - "creatorhub_browser", testToken, "gateway-self") - response := httptest.NewRecorder() - adaptor.FiberApp(handler).ServeHTTP(response, authed(http.MethodPost, "/v1/browsers", strings.NewReader(testCreateBody))) - if response.Code != http.StatusBadGateway || !networkCreated || networkReads != 1 || networkDeletes != 0 { - t.Fatalf("unknown N1 was discovered or cleaned as N2: status=%d created=%v reads=%d deletes=%d body=%s", - response.Code, networkCreated, networkReads, networkDeletes, response.Body.String()) - } -} - -func TestGatewayCompensatesConnectThatAppliedBeforeError(t *testing.T) { - for _, restore := range []bool{false, true} { - name := "create self" - failedMember := "gateway-self" - if restore { - name, failedMember = "restore runtime", "container-c1" - } - t.Run(name, func(t *testing.T) { - networkExists := false - members := map[string]string{} - disconnected := []string{} - networkDeletes := 0 - server := httptest.NewServer(withAliasReservations("gateway-self", func(response http.ResponseWriter, request *http.Request) { - switch { - case request.Method == http.MethodGet && strings.HasPrefix(request.URL.Path, "/images/"): - response.WriteHeader(http.StatusOK) - case request.Method == http.MethodGet && request.URL.Path == "/containers/"+namePrefix+"account-a/json": - if !restore { - response.WriteHeader(http.StatusNotFound) - return - } - _, _ = response.Write([]byte(`{"Id":"container-c1","Config":{"Labels":{"` + managedLabel + `":"true","` + idLabel + `":"account-a","` + bindingVersionLabel + `":"1","` + networkExitLabel + `":"exit-1","` + networkIDLabel + `":"network-n1","` + proxyPortLabel + `":"12345"}}}`)) - case request.Method == http.MethodGet && strings.HasPrefix(request.URL.Path, "/networks/"): - if !networkExists { - response.WriteHeader(http.StatusNotFound) - return - } - containers := map[string]any{} - for id, memberName := range members { - containers[id] = map[string]string{"Name": memberName, "IPv4Address": "127.0.0.1/8"} - } - _ = json.NewEncoder(response).Encode(map[string]any{ - "Id": "network-n1", "Name": "creatorhub_browser-account-a", "Driver": "bridge", "Internal": false, "Attachable": false, "Ingress": false, - "Labels": map[string]string{managedLabel: "true", networkRoleLabel: browserNetworkRole, idLabel: "account-a", bindingVersionLabel: "1"}, "Containers": containers, - }) - case request.Method == http.MethodPost && request.URL.Path == "/networks/create": - networkExists = true - response.WriteHeader(http.StatusCreated) - _, _ = response.Write([]byte(`{"Id":"network-n1"}`)) - case request.Method == http.MethodPost && strings.HasSuffix(request.URL.Path, "/connect") && - !strings.HasSuffix(request.URL.Path, "/disconnect"): - var body struct { - Container string `json:"Container"` - } - _ = json.NewDecoder(request.Body).Decode(&body) - members[body.Container] = body.Container - if body.Container == failedMember { - response.WriteHeader(http.StatusInternalServerError) - return - } - response.WriteHeader(http.StatusOK) - case request.Method == http.MethodPost && strings.HasSuffix(request.URL.Path, "/disconnect"): - var body struct { - Container string `json:"Container"` - } - _ = json.NewDecoder(request.Body).Decode(&body) - delete(members, body.Container) - disconnected = append(disconnected, body.Container) - response.WriteHeader(http.StatusOK) - case request.Method == http.MethodDelete && request.URL.Path == "/networks/network-n1": - networkExists = false - networkDeletes++ - response.WriteHeader(http.StatusNoContent) - default: - t.Fatalf("unexpected Docker request %s %s", request.Method, request.URL.String()) - } - })) - defer server.Close() - handler := newGatewayWithSelf(dockerClient{baseURL: server.URL, client: server.Client(), slow: server.Client()}, - "creatorhub_browser", testToken, "gateway-self") - method, path, body := http.MethodPost, "/v1/browsers", testCreateBody - if restore { - path = "/v1/browsers/account-a/proxy" - body = `{"binding_version":1,"runtime_id":"container-c1","network_id":"network-n1","network_exit_id":"exit-1","network_exit":{"protocol":"socks5","host":"proxy.example","port":1080}}` - } - response := httptest.NewRecorder() - adaptor.FiberApp(handler).ServeHTTP(response, authed(method, path, strings.NewReader(body))) - wantDeletes, wantDisconnects := 1, 1 - if restore { - wantDeletes, wantDisconnects = 0, 0 - } - if response.Code != http.StatusBadGateway || networkExists || networkDeletes != wantDeletes || - len(disconnected) != wantDisconnects || (wantDisconnects == 1 && disconnected[0] != failedMember) { - t.Fatalf("applied connect was not compensated: status=%d network=%v deletes=%d disconnected=%v body=%s", - response.Code, networkExists, networkDeletes, disconnected, response.Body.String()) - } - }) - } -} - -func TestGatewayRestoreFailureRemovesCreatedNetworkGeneration(t *testing.T) { - var mu sync.Mutex - networkExists := false - members := map[string]string{} - networkDeletes, containerDeletes := 0, 0 - server := httptest.NewServer(withAliasReservations("gateway-self", func(response http.ResponseWriter, request *http.Request) { - switch { - case request.Method == http.MethodGet && request.URL.Path == "/containers/"+namePrefix+"account-a/json": - _, _ = response.Write([]byte(`{"Id":"container-c1","Config":{"Labels":{"` + managedLabel + `":"true","` + idLabel + `":"account-a","` + bindingVersionLabel + `":"1","` + networkExitLabel + `":"exit-1","` + networkIDLabel + `":"network-n1","` + proxyPortLabel + `":"12345"}}}`)) - case request.Method == http.MethodGet && (request.URL.Path == "/networks/creatorhub_browser-account-a" || request.URL.Path == "/networks/network-n1"): - mu.Lock() - exists := networkExists - current := map[string]any{} - for id, name := range members { - ip := "127.0.0.3/8" - if id == "gateway-self" { - ip = "192.0.2.1/24" - } - current[id] = map[string]string{"Name": name, "IPv4Address": ip} - } - mu.Unlock() - if !exists { - response.WriteHeader(http.StatusNotFound) - return - } - _ = json.NewEncoder(response).Encode(map[string]any{ - "Id": "network-n1", "Name": "creatorhub_browser-account-a", "Driver": "bridge", "Internal": false, "Attachable": false, "Ingress": false, - "Labels": map[string]string{managedLabel: "true", networkRoleLabel: browserNetworkRole, idLabel: "account-a", bindingVersionLabel: "1"}, - "Containers": current, - }) - case request.Method == http.MethodPost && request.URL.Path == "/networks/create": - mu.Lock() - networkExists = true - mu.Unlock() - response.WriteHeader(http.StatusCreated) - _, _ = response.Write([]byte(`{"Id":"network-n1"}`)) - case request.Method == http.MethodPost && strings.HasSuffix(request.URL.Path, "/connect"): - var body struct { - Container string `json:"Container"` - } - _ = json.NewDecoder(request.Body).Decode(&body) - mu.Lock() - members[body.Container] = body.Container - mu.Unlock() - response.WriteHeader(http.StatusOK) - case request.Method == http.MethodPost && strings.HasSuffix(request.URL.Path, "/disconnect"): - var body struct { - Container string `json:"Container"` - } - _ = json.NewDecoder(request.Body).Decode(&body) - mu.Lock() - delete(members, body.Container) - mu.Unlock() - response.WriteHeader(http.StatusOK) - case request.Method == http.MethodDelete && request.URL.Path == "/networks/network-n1": - mu.Lock() - networkExists = false - networkDeletes++ - mu.Unlock() - response.WriteHeader(http.StatusNoContent) - case request.Method == http.MethodDelete && strings.HasPrefix(request.URL.Path, "/containers/"): - mu.Lock() - containerDeletes++ - mu.Unlock() - response.WriteHeader(http.StatusNoContent) - default: - t.Fatalf("unexpected Docker request %s %s", request.Method, request.URL.String()) - } - })) - defer server.Close() - docker := dockerClient{baseURL: server.URL, client: server.Client(), slow: server.Client()} - handler := newGatewayWithSelf(docker, "creatorhub_browser", testToken, "gateway-self") - response := httptest.NewRecorder() - body := `{"binding_version":1,"runtime_id":"container-c1","network_id":"network-n1","network_exit_id":"exit-1","network_exit":{"protocol":"socks5","host":"proxy.example","port":1080}}` - adaptor.FiberApp(handler).ServeHTTP(response, authed(http.MethodPost, "/v1/browsers/account-a/proxy", strings.NewReader(body))) - mu.Lock() - defer mu.Unlock() - if response.Code != http.StatusBadGateway || networkExists || networkDeletes != 0 || containerDeletes != 0 { - t.Fatalf("restore failure left managed network: status=%d network=%v networkDeletes=%d containerDeletes=%d body=%s", - response.Code, networkExists, networkDeletes, containerDeletes, response.Body.String()) - } -} - -func TestGatewayMapsDockerServiceFailureToBadGateway(t *testing.T) { - docker, server := testDocker(func(response http.ResponseWriter, _ *http.Request) { - http.Error(response, "daemon unavailable", http.StatusInternalServerError) - }) - defer server.Close() - - handler := newGateway(docker, "creatorhub_browser", testToken) - response := httptest.NewRecorder() - adaptor.FiberApp(handler).ServeHTTP(response, authed(http.MethodDelete, "/v1/browsers/account-a", strings.NewReader(testGenerationBody))) - if response.Code != http.StatusBadGateway { - t.Fatalf("expected 502 for Docker failure, got %d: %s", response.Code, response.Body.String()) - } -} - -func TestGatewayRefusesUnmanagedContainer(t *testing.T) { - deleted := false - docker, server := testDocker(func(response http.ResponseWriter, request *http.Request) { - switch request.Method { - case http.MethodGet: - _, _ = response.Write([]byte(`{"Config":{"Labels":{}}}`)) - case http.MethodDelete: - deleted = true - response.WriteHeader(http.StatusNoContent) - } - }) - defer server.Close() - - handler := newGateway(docker, "creatorhub_browser", testToken) - request := authed(http.MethodDelete, "/v1/browsers/foreign", strings.NewReader(testGenerationBody)) - response := httptest.NewRecorder() - adaptor.FiberApp(handler).ServeHTTP(response, request) - - if response.Code != http.StatusForbidden || deleted { - t.Fatalf("expected unmanaged container to be rejected, status=%d deleted=%v", response.Code, deleted) - } -} - -func TestGatewayRejectsStaleGenerationBeforeDockerMutation(t *testing.T) { - for _, request := range []struct { - method, path string - }{ - {http.MethodPost, "/v1/browsers/account-a/stop"}, - {http.MethodPost, "/v1/browsers/account-a/start"}, - {http.MethodDelete, "/v1/browsers/account-a"}, - } { - t.Run(request.method, func(t *testing.T) { - mutations := 0 - docker, server := testDocker(func(response http.ResponseWriter, dockerRequest *http.Request) { - if dockerRequest.Method != http.MethodGet { - mutations++ - } - _, _ = response.Write([]byte(`{"Id":"new-container","Config":{"Labels":{"` + managedLabel + `":"true","` + idLabel + `":"account-a","` + bindingVersionLabel + `":"2"}}}`)) - }) - defer server.Close() - handler := newGateway(docker, "creatorhub_browser", testToken) - response := httptest.NewRecorder() - body := testGenerationBody - if strings.HasSuffix(request.path, "/start") { - body = `{"binding_version":1,"runtime_id":"container-id","network_id":"network-id"}` - } - adaptor.FiberApp(handler).ServeHTTP(response, authed(request.method, request.path, strings.NewReader(body))) - if response.Code != http.StatusConflict || mutations != 0 { - t.Fatalf("stale generation reached Docker mutation: status=%d mutations=%d body=%s", response.Code, mutations, response.Body.String()) - } - }) - } -} - -func TestGatewayRejectsStaleDeleteDuringNewNetworkCreation(t *testing.T) { - mutations := 0 - server := httptest.NewServer(withAliasReservations("gateway-self", func(response http.ResponseWriter, request *http.Request) { - switch { - case request.Method == http.MethodGet && strings.HasPrefix(request.URL.Path, "/containers/"): - response.WriteHeader(http.StatusNotFound) - case request.Method == http.MethodGet && strings.HasPrefix(request.URL.Path, "/networks/"): - _ = json.NewEncoder(response).Encode(map[string]any{ - "Id": "network-new", "Name": "creatorhub_browser-account-a", "Driver": "bridge", "Internal": false, "Attachable": false, "Ingress": false, - "Labels": map[string]string{managedLabel: "true", networkRoleLabel: browserNetworkRole, idLabel: "account-a", - bindingVersionLabel: "2"}, - }) - default: - mutations++ - response.WriteHeader(http.StatusNoContent) - } - })) - defer server.Close() - handler := newGatewayWithSelf(dockerClient{baseURL: server.URL, client: server.Client(), slow: server.Client()}, - "creatorhub_browser", testToken, "gateway-self") - response := httptest.NewRecorder() - adaptor.FiberApp(handler).ServeHTTP(response, authed(http.MethodDelete, "/v1/browsers/account-a", strings.NewReader(testGenerationBody))) - if response.Code != http.StatusConflict || mutations != 0 { - t.Fatalf("stale delete crossed the new network generation: status=%d mutations=%d body=%s", response.Code, mutations, response.Body.String()) - } -} - -func TestEnsureTenantNetworkConnectsGatewayOnlyToRuntimeNetwork(t *testing.T) { - created, connected := false, false - server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { - switch { - case request.Method == http.MethodGet && !created: - response.WriteHeader(http.StatusNotFound) - case request.Method == http.MethodPost && request.URL.Path == "/networks/create": - var body map[string]any - _ = json.NewDecoder(request.Body).Decode(&body) - labels := body["Labels"].(map[string]any) - if body["Name"] != "creatorhub_browser-account-a" || labels[idLabel] != "account-a" { - t.Fatalf("unexpected isolated network create: %#v", body) - } - created = true - response.WriteHeader(http.StatusCreated) - _, _ = response.Write([]byte(`{"Id":"network-id"}`)) - case request.Method == http.MethodPost && strings.HasSuffix(request.URL.Path, "/connect"): - connected = true - response.WriteHeader(http.StatusOK) - case request.Method == http.MethodGet && request.URL.Path == "/containers/gateway-id/json": - _, _ = response.Write([]byte(`{"Id":"gateway-id","Config":{"Labels":{"` + gatewayMemberLabel + `":"true"}}}`)) - case request.Method == http.MethodGet: - members := map[string]any{} - if connected { - members["gateway-id"] = map[string]string{"Name": "gateway-id", "IPv4Address": "127.0.0.3/8"} - } - _ = json.NewEncoder(response).Encode(map[string]any{ - "Id": "network-id", "Name": "creatorhub_browser-account-a", "Driver": "bridge", "Internal": false, "Attachable": false, "Ingress": false, - "Labels": map[string]string{managedLabel: "true", networkRoleLabel: browserNetworkRole, idLabel: "account-a", - bindingVersionLabel: "1"}, - "Containers": members, - }) - default: - t.Fatalf("unexpected Docker request %s %s", request.Method, request.URL.Path) - } - })) - defer server.Close() - docker := dockerClient{baseURL: server.URL, client: server.Client(), slow: server.Client()} - generation, bindHost, err := docker.ensureTenantNetwork("creatorhub_browser", "account-a", "gateway-id", 1, "", "", false) - if err != nil || !created || !connected || generation.Name != "creatorhub_browser-account-a" || bindHost != "127.0.0.3" { - t.Fatalf("isolated network was not created and connected: generation=%#v host=%q created=%v connected=%v err=%v", generation, bindHost, created, connected, err) - } -} - -func TestLoadConfigRequiresGatewayToken(t *testing.T) { - t.Setenv("GATEWAY_TOKEN", "short") - if _, err := loadConfig(); err == nil { - t.Fatal("expected a short gateway token to be rejected") - } -} - -func TestLoadConfigRejectsControlNetwork(t *testing.T) { - t.Setenv("GATEWAY_TOKEN", testToken) - t.Setenv("BROWSER_NETWORK", controlNetworkName) - if _, err := loadConfig(); err == nil { - t.Fatal("expected control network configuration to be rejected") - } -} diff --git a/cmd/docker-gateway/proxy.go b/cmd/docker-gateway/proxy.go deleted file mode 100644 index ae92734..0000000 --- a/cmd/docker-gateway/proxy.go +++ /dev/null @@ -1,435 +0,0 @@ -package main - -import ( - "bufio" - "context" - "crypto/tls" - "encoding/base64" - "encoding/binary" - "errors" - "fmt" - "io" - "net" - "net/http" - "net/url" - "strconv" - "sync" - "time" -) - -const browserProxyHost = "docker-gateway" - -type memoryProxyRegistry struct { - mu sync.Mutex - proxies map[string]*memoryProxy -} - -type memoryProxy struct { - mu sync.RWMutex - exit gatewayProxyExit - bindingVersion int64 - runtimeID string - networkID string - bindHost string - listener net.Listener - server *http.Server - tunnels map[net.Conn]net.Conn - url string -} - -func newMemoryProxyRegistry() *memoryProxyRegistry { - return &memoryProxyRegistry{proxies: map[string]*memoryProxy{}} -} - -func (registry *memoryProxyRegistry) configure(alias string, bindingVersion int64, bindHost string, port int, exit gatewayProxyExit, - networkIDs ...string) (string, func(), error) { - registry.mu.Lock() - defer registry.mu.Unlock() - networkID := "" - if len(networkIDs) == 1 { - networkID = networkIDs[0] - } - if proxy := registry.proxies[alias]; proxy != nil { - if proxy.bindingVersion == bindingVersion && proxy.bindHost == bindHost && - (port == 0 || proxy.listener.Addr().(*net.TCPAddr).Port == port) && proxy.exit == exit && proxy.networkID == networkID { - return proxy.url, func() { registry.removeObject(alias, proxy) }, nil - } - delete(registry.proxies, alias) - // A gateway replica may still hold the previous runtime generation. Close - // it before rebinding a restored listener, while removeObject's identity - // check keeps old cleanup callbacks from deleting the replacement. - closeMemoryProxy(proxy) - } - listener, err := net.Listen("tcp4", net.JoinHostPort(bindHost, strconv.Itoa(port))) - if err != nil { - return "", nil, err - } - actualPort := listener.Addr().(*net.TCPAddr).Port - proxy := &memoryProxy{exit: exit, bindingVersion: bindingVersion, networkID: networkID, bindHost: bindHost, listener: listener, - tunnels: make(map[net.Conn]net.Conn), url: "http://" + net.JoinHostPort(browserProxyHost, strconv.Itoa(actualPort))} - proxy.server = &http.Server{Handler: proxy, ReadHeaderTimeout: 10 * time.Second, IdleTimeout: 60 * time.Second} - registry.proxies[alias] = proxy - go func() { _ = proxy.server.Serve(listener) }() - undo := func() { - registry.removeObject(alias, proxy) - } - return proxy.url, undo, nil -} - -func (registry *memoryProxyRegistry) removeObject(alias string, proxy *memoryProxy) { - registry.mu.Lock() - if registry.proxies[alias] == proxy { - delete(registry.proxies, alias) - } else { - proxy = nil - } - registry.mu.Unlock() - if proxy != nil { - closeMemoryProxy(proxy) - } -} - -func closeMemoryProxy(proxy *memoryProxy) { - proxy.mu.Lock() - tunnels := proxy.tunnels - proxy.tunnels = nil - proxy.mu.Unlock() - _ = proxy.listener.Close() - _ = proxy.server.Close() - for client, upstream := range tunnels { - _ = client.Close() - _ = upstream.Close() - } -} - -func (registry *memoryProxyRegistry) ready(alias string, port int, runtimeID string, networkIDs ...string) bool { - registry.mu.Lock() - defer registry.mu.Unlock() - proxy := registry.proxies[alias] - return proxy != nil && runtimeID != "" && proxy.runtimeID == runtimeID && proxy.listener.Addr().(*net.TCPAddr).Port == port && - (len(networkIDs) == 0 || proxy.networkID == networkIDs[0]) -} - -func (registry *memoryProxyRegistry) bind(alias string, bindingVersion int64, proxyURL, runtimeID string, networkIDs ...string) bool { - registry.mu.Lock() - defer registry.mu.Unlock() - proxy := registry.proxies[alias] - if proxy == nil || runtimeID == "" || proxy.bindingVersion != bindingVersion || proxy.url != proxyURL { - return false - } - proxy.runtimeID = runtimeID - if len(networkIDs) == 1 { - proxy.networkID = networkIDs[0] - } - return true -} - -func (registry *memoryProxyRegistry) remove(alias string, bindingVersion int64, runtimeID string, networkIDs ...string) bool { - registry.mu.Lock() - proxy := registry.proxies[alias] - if proxy != nil && runtimeID != "" && proxy.bindingVersion == bindingVersion && proxy.runtimeID == runtimeID && - (len(networkIDs) == 0 || proxy.networkID == networkIDs[0]) { - delete(registry.proxies, alias) - } else if proxy != nil { - registry.mu.Unlock() - return false - } else { - proxy = nil - } - registry.mu.Unlock() - if proxy != nil { - closeMemoryProxy(proxy) - } - return true -} - -func (proxy *memoryProxy) ServeHTTP(response http.ResponseWriter, request *http.Request) { - if request.Method == http.MethodConnect { - proxy.tunnel(response, request) - return - } - proxy.mu.RLock() - exit := proxy.exit - proxy.mu.RUnlock() - transport := &http.Transport{DisableKeepAlives: true} - if exit.Protocol == "http" || exit.Protocol == "https" { - upstream := &url.URL{Scheme: exit.Protocol, Host: net.JoinHostPort(exit.Host, strconv.Itoa(exit.Port))} - if exit.Username != "" { - upstream.User = url.UserPassword(exit.Username, exit.Password) - } - transport.Proxy = http.ProxyURL(upstream) - } else { - transport.DialContext = proxy.dialContext - } - defer transport.CloseIdleConnections() - outbound := request.Clone(request.Context()) - outbound.RequestURI = "" - outbound.Header.Del("Proxy-Authorization") - result, err := transport.RoundTrip(outbound) - if err != nil { - http.Error(response, "proxy connection failed", http.StatusBadGateway) - return - } - defer result.Body.Close() - for key, values := range result.Header { - for _, value := range values { - response.Header().Add(key, value) - } - } - response.WriteHeader(result.StatusCode) - _, _ = io.Copy(response, result.Body) -} - -func (proxy *memoryProxy) tunnel(response http.ResponseWriter, request *http.Request) { - upstream, err := proxy.dialContext(request.Context(), "tcp", request.Host) - if err != nil { - http.Error(response, "proxy connection failed", http.StatusBadGateway) - return - } - client, buffered, err := http.NewResponseController(response).Hijack() - if err != nil { - _ = upstream.Close() - http.Error(response, "proxy tunnel unavailable", http.StatusInternalServerError) - return - } - if _, err := buffered.WriteString("HTTP/1.1 200 Connection Established\r\n\r\n"); err != nil || buffered.Flush() != nil { - _ = client.Close() - _ = upstream.Close() - return - } - proxy.mu.Lock() - if proxy.tunnels == nil { - proxy.mu.Unlock() - _ = client.Close() - _ = upstream.Close() - return - } - proxy.tunnels[client] = upstream - proxy.mu.Unlock() - defer func() { - proxy.mu.Lock() - delete(proxy.tunnels, client) - proxy.mu.Unlock() - _ = client.Close() - _ = upstream.Close() - }() - done := make(chan struct{}, 2) - go func() { _, _ = io.Copy(upstream, client); done <- struct{}{} }() - go func() { _, _ = io.Copy(client, upstream); done <- struct{}{} }() - <-done -} - -func (proxy *memoryProxy) dialContext(ctx context.Context, _, target string) (net.Conn, error) { - ctx, cancel := context.WithTimeout(ctx, 20*time.Second) - defer cancel() - proxy.mu.RLock() - exit := proxy.exit - proxy.mu.RUnlock() - switch exit.Protocol { - case "http", "https": - return dialHTTPProxy(ctx, exit, target) - case "socks4": - return dialSOCKS4Proxy(ctx, exit, target) - case "socks5": - return dialSOCKS5Proxy(ctx, exit, target) - default: - return nil, errors.New("unsupported proxy protocol") - } -} - -func dialHTTPProxy(ctx context.Context, exit gatewayProxyExit, target string) (net.Conn, error) { - address := net.JoinHostPort(exit.Host, strconv.Itoa(exit.Port)) - connection, err := (&net.Dialer{}).DialContext(ctx, "tcp", address) - if err != nil { - return nil, err - } - if exit.Protocol == "https" { - tlsConnection := tls.Client(connection, &tls.Config{ServerName: exit.Host, MinVersion: tls.VersionTLS12}) - if err := tlsConnection.HandshakeContext(ctx); err != nil { - _ = connection.Close() - return nil, err - } - connection = tlsConnection - } - request := &http.Request{Method: http.MethodConnect, URL: &url.URL{Opaque: target}, Host: target, Header: make(http.Header)} - if exit.Username != "" { - request.Header.Set("Proxy-Authorization", "Basic "+base64.StdEncoding.EncodeToString([]byte(exit.Username+":"+exit.Password))) - } - if deadline, ok := ctx.Deadline(); ok { - _ = connection.SetDeadline(deadline) - } - if err := request.Write(connection); err != nil { - _ = connection.Close() - return nil, err - } - result, err := http.ReadResponse(bufio.NewReader(connection), request) - if err != nil { - _ = connection.Close() - return nil, err - } - if result.StatusCode != http.StatusOK { - _ = result.Body.Close() - _ = connection.Close() - return nil, fmt.Errorf("upstream proxy returned %s", result.Status) - } - _ = connection.SetDeadline(time.Time{}) - return connection, nil -} - -func dialSOCKS4Proxy(ctx context.Context, exit gatewayProxyExit, target string) (net.Conn, error) { - connection, err := (&net.Dialer{}).DialContext(ctx, "tcp", net.JoinHostPort(exit.Host, strconv.Itoa(exit.Port))) - if err != nil { - return nil, err - } - host, portText, err := net.SplitHostPort(target) - if err != nil { - _ = connection.Close() - return nil, err - } - port, err := strconv.Atoi(portText) - if err != nil || port < 1 || port > 65535 { - _ = connection.Close() - return nil, errors.New("invalid SOCKS4 target") - } - payload := []byte{4, 1, byte(port >> 8), byte(port), 0, 0, 0, 1} - if ip := net.ParseIP(host).To4(); ip != nil { - copy(payload[4:8], ip) - } - payload = append(payload, exit.Username...) - payload = append(payload, 0) - if net.ParseIP(host).To4() == nil { - payload = append(payload, host...) - payload = append(payload, 0) - } - if err := exchangeSOCKS(ctx, connection, payload, 8); err != nil { - return nil, err - } - return connection, nil -} - -func dialSOCKS5Proxy(ctx context.Context, exit gatewayProxyExit, target string) (net.Conn, error) { - connection, err := (&net.Dialer{}).DialContext(ctx, "tcp", net.JoinHostPort(exit.Host, strconv.Itoa(exit.Port))) - if err != nil { - return nil, err - } - methods := []byte{5, 1, 0} - if exit.Username != "" { - methods = []byte{5, 1, 2} - } - if deadline, ok := ctx.Deadline(); ok { - _ = connection.SetDeadline(deadline) - } - if _, err := connection.Write(methods); err != nil { - _ = connection.Close() - return nil, err - } - selection := make([]byte, 2) - if _, err := io.ReadFull(connection, selection); err != nil || selection[0] != 5 || selection[1] == 0xff { - _ = connection.Close() - return nil, errors.New("SOCKS5 authentication method rejected") - } - if selection[1] == 2 { - if len(exit.Username) > 255 || len(exit.Password) > 255 { - _ = connection.Close() - return nil, errors.New("SOCKS5 credentials too long") - } - auth := append([]byte{1, byte(len(exit.Username))}, exit.Username...) - auth = append(auth, byte(len(exit.Password))) - auth = append(auth, exit.Password...) - if _, err := connection.Write(auth); err != nil { - _ = connection.Close() - return nil, err - } - result := make([]byte, 2) - if _, err := io.ReadFull(connection, result); err != nil || result[1] != 0 { - _ = connection.Close() - return nil, errors.New("SOCKS5 authentication rejected") - } - } else if exit.Username != "" { - _ = connection.Close() - return nil, errors.New("SOCKS5 proxy skipped required authentication") - } - host, portText, err := net.SplitHostPort(target) - if err != nil { - _ = connection.Close() - return nil, err - } - port, err := strconv.Atoi(portText) - if err != nil || port < 1 || port > 65535 { - _ = connection.Close() - return nil, errors.New("invalid SOCKS5 target") - } - request := []byte{5, 1, 0} - if ip := net.ParseIP(host); ip != nil && ip.To4() != nil { - request = append(request, 1) - request = append(request, ip.To4()...) - } else if ip != nil { - request = append(request, 4) - request = append(request, ip.To16()...) - } else { - if len(host) > 255 { - _ = connection.Close() - return nil, errors.New("SOCKS5 target too long") - } - request = append(request, 3, byte(len(host))) - request = append(request, host...) - } - portBytes := make([]byte, 2) - binary.BigEndian.PutUint16(portBytes, uint16(port)) - request = append(request, portBytes...) - if _, err := connection.Write(request); err != nil { - _ = connection.Close() - return nil, err - } - header := make([]byte, 4) - if _, err := io.ReadFull(connection, header); err != nil || header[0] != 5 || header[1] != 0 { - _ = connection.Close() - return nil, errors.New("SOCKS5 proxy rejected connection") - } - addressLength := 0 - switch header[3] { - case 1: - addressLength = 4 - case 4: - addressLength = 16 - case 3: - var length [1]byte - if _, err := io.ReadFull(connection, length[:]); err != nil { - _ = connection.Close() - return nil, err - } - addressLength = int(length[0]) - default: - _ = connection.Close() - return nil, errors.New("invalid SOCKS5 response") - } - if _, err := io.CopyN(io.Discard, connection, int64(addressLength+2)); err != nil { - _ = connection.Close() - return nil, err - } - _ = connection.SetDeadline(time.Time{}) - return connection, nil -} - -func exchangeSOCKS(ctx context.Context, connection net.Conn, request []byte, responseBytes int) error { - if deadline, ok := ctx.Deadline(); ok { - _ = connection.SetDeadline(deadline) - } - if _, err := connection.Write(request); err != nil { - _ = connection.Close() - return err - } - if responseBytes > 0 { - response := make([]byte, responseBytes) - if _, err := io.ReadFull(connection, response); err != nil { - _ = connection.Close() - return err - } - if responseBytes == 8 && response[1] != 90 { - _ = connection.Close() - return errors.New("SOCKS4 proxy rejected connection") - } - } - _ = connection.SetDeadline(time.Time{}) - return nil -} diff --git a/cmd/docker-gateway/proxy_test.go b/cmd/docker-gateway/proxy_test.go deleted file mode 100644 index 5467e07..0000000 --- a/cmd/docker-gateway/proxy_test.go +++ /dev/null @@ -1,252 +0,0 @@ -package main - -import ( - "bufio" - "encoding/binary" - "fmt" - "io" - "net" - "net/http" - "net/http/httptest" - "net/url" - "strconv" - "strings" - "testing" - "time" -) - -func TestMemoryProxyUsesSOCKS5Credentials(t *testing.T) { - listener, err := net.Listen("tcp", "127.0.0.1:0") - if err != nil { - t.Fatal(err) - } - defer listener.Close() - done := make(chan error, 1) - go func() { - connection, err := listener.Accept() - if err != nil { - done <- err - return - } - defer connection.Close() - greeting := make([]byte, 3) - if _, err := io.ReadFull(connection, greeting); err != nil { - done <- err - return - } - _, _ = connection.Write([]byte{5, 2}) - authHeader := make([]byte, 2) - _, _ = io.ReadFull(connection, authHeader) - username := make([]byte, int(authHeader[1])) - _, _ = io.ReadFull(connection, username) - var passwordLength [1]byte - _, _ = io.ReadFull(connection, passwordLength[:]) - password := make([]byte, int(passwordLength[0])) - _, _ = io.ReadFull(connection, password) - if string(username) != "operator" || string(password) != "ephemeral" { - done <- io.ErrUnexpectedEOF - return - } - _, _ = connection.Write([]byte{1, 0}) - requestHeader := make([]byte, 5) - _, _ = io.ReadFull(connection, requestHeader) - host := make([]byte, int(requestHeader[4])) - _, _ = io.ReadFull(connection, host) - port := make([]byte, 2) - _, _ = io.ReadFull(connection, port) - if string(host) != "example.com" || binary.BigEndian.Uint16(port) != 443 { - done <- io.ErrUnexpectedEOF - return - } - if _, err = connection.Write([]byte{5, 0, 0, 1, 127, 0, 0, 1, 0, 0}); err != nil { - done <- err - return - } - var tunneled [1]byte - _, err = io.ReadFull(connection, tunneled[:]) - if err == nil && tunneled[0] != 'x' { - err = io.ErrUnexpectedEOF - } - done <- err - }() - - host, portText, _ := net.SplitHostPort(listener.Addr().String()) - port, _ := net.LookupPort("tcp", portText) - registry := newMemoryProxyRegistry() - proxyURL, cleanup, err := registry.configure("account-a", 1, "127.0.0.1", 0, gatewayProxyExit{ - Protocol: "socks5", Host: host, Port: port, Username: "operator", Password: "ephemeral", - }) - if err != nil { - t.Fatal(err) - } - defer cleanup() - parsed, _ := url.Parse(proxyURL) - connection, err := net.Dial("tcp", strings.Replace(parsed.Host, browserProxyHost, "127.0.0.1", 1)) - if err != nil { - t.Fatal(err) - } - if _, err := fmt.Fprint(connection, "CONNECT example.com:443 HTTP/1.1\r\nHost: example.com:443\r\n\r\n"); err != nil { - t.Fatal(err) - } - response, err := http.ReadResponse(bufio.NewReader(connection), &http.Request{Method: http.MethodConnect}) - if err != nil || response.StatusCode != http.StatusOK { - t.Fatalf("memory proxy CONNECT failed: response=%v err=%v", response, err) - } - if _, err := connection.Write([]byte{'x'}); err != nil { - t.Fatal(err) - } - _ = connection.Close() - if err := <-done; err != nil { - t.Fatal(err) - } -} - -func TestMemoryProxyUsesAbsoluteFormForHTTPUpstream(t *testing.T) { - upstream := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { - if request.Method == http.MethodConnect { - http.Error(response, "CONNECT forbidden", http.StatusMethodNotAllowed) - return - } - if !request.URL.IsAbs() || request.URL.String() != "http://example.com/plain" { - t.Fatalf("expected absolute-form request, got %q", request.URL.String()) - } - if request.Header.Get("Proxy-Authorization") == "" { - t.Fatal("upstream proxy credentials were not applied") - } - _, _ = response.Write([]byte("forwarded")) - })) - defer upstream.Close() - address, _ := url.Parse(upstream.URL) - port, _ := strconv.Atoi(address.Port()) - registry := newMemoryProxyRegistry() - proxyURL, cleanup, err := registry.configure("account-a", 1, "127.0.0.1", 0, gatewayProxyExit{ - Protocol: "http", Host: address.Hostname(), Port: port, Username: "operator", Password: "ephemeral", - }) - if err != nil { - t.Fatal(err) - } - defer cleanup() - proxyAddress := strings.Replace(strings.TrimPrefix(proxyURL, "http://"), browserProxyHost, "127.0.0.1", 1) - client := &http.Client{Transport: &http.Transport{Proxy: http.ProxyURL(&url.URL{Scheme: "http", Host: proxyAddress})}} - response, err := client.Get("http://example.com/plain") - if err != nil { - t.Fatal(err) - } - defer response.Body.Close() - body, _ := io.ReadAll(response.Body) - if response.StatusCode != http.StatusOK || string(body) != "forwarded" { - t.Fatalf("plain HTTP was not forwarded: status=%d body=%s", response.StatusCode, body) - } -} - -func TestMemoryProxyCleanupClosesHijackedTunnel(t *testing.T) { - listener, err := net.Listen("tcp", "127.0.0.1:0") - if err != nil { - t.Fatal(err) - } - defer listener.Close() - upstreamClosed := make(chan error, 1) - go func() { - connection, err := listener.Accept() - if err != nil { - upstreamClosed <- err - return - } - defer connection.Close() - request, err := http.ReadRequest(bufio.NewReader(connection)) - if err == nil && request.Method == http.MethodConnect { - _, err = fmt.Fprint(connection, "HTTP/1.1 200 Connection Established\r\n\r\n") - } - if err == nil { - var data [1]byte - _, err = connection.Read(data[:]) - } - upstreamClosed <- err - }() - host, portText, _ := net.SplitHostPort(listener.Addr().String()) - port, _ := strconv.Atoi(portText) - registry := newMemoryProxyRegistry() - proxyURL, cleanup, err := registry.configure("account-a", 1, "127.0.0.1", 0, - gatewayProxyExit{Protocol: "http", Host: host, Port: port}) - if err != nil { - t.Fatal(err) - } - defer cleanup() - proxyAddress := strings.Replace(strings.TrimPrefix(proxyURL, "http://"), browserProxyHost, "127.0.0.1", 1) - client, err := net.Dial("tcp", proxyAddress) - if err != nil { - t.Fatal(err) - } - defer client.Close() - if _, err = fmt.Fprint(client, "CONNECT example.com:443 HTTP/1.1\r\nHost: example.com:443\r\n\r\n"); err != nil { - t.Fatal(err) - } - if response, err := http.ReadResponse(bufio.NewReader(client), &http.Request{Method: http.MethodConnect}); err != nil || response.StatusCode != http.StatusOK { - t.Fatalf("open CONNECT tunnel: response=%v err=%v", response, err) - } - cleanup() - _ = client.SetReadDeadline(time.Now().Add(time.Second)) - if _, err = client.Read(make([]byte, 1)); err == nil { - t.Fatal("proxy cleanup left the client tunnel open") - } - select { - case err = <-upstreamClosed: - if err == nil { - t.Fatal("proxy cleanup left the upstream tunnel open") - } - case <-time.After(time.Second): - t.Fatal("proxy cleanup did not close the upstream tunnel") - } -} - -func TestMemoryProxyRejectsCrossAliasAddress(t *testing.T) { - registry := newMemoryProxyRegistry() - proxyURL, cleanup, err := registry.configure("account-a", 1, "127.0.0.1", 0, gatewayProxyExit{Protocol: "http", Host: "127.0.0.1", Port: 1}) - if err != nil { - t.Fatal(err) - } - defer cleanup() - parsed, _ := url.Parse(proxyURL) - if connection, err := net.DialTimeout("tcp", net.JoinHostPort("127.0.0.2", parsed.Port()), 100*time.Millisecond); err == nil { - _ = connection.Close() - t.Fatal("another tenant address could reach account-a proxy") - } -} - -func TestMemoryProxyRemoveRequiresMatchingGeneration(t *testing.T) { - registry := newMemoryProxyRegistry() - proxyURL, cleanup, err := registry.configure("account-a", 2, "127.0.0.1", 0, gatewayProxyExit{Protocol: "http", Host: "127.0.0.1", Port: 1}) - if err != nil { - t.Fatal(err) - } - defer cleanup() - if !registry.bind("account-a", 2, proxyURL, "container-c2") { - t.Fatal("bind proxy generation") - } - proxy := registry.proxies["account-a"] - if registry.remove("account-a", 2, "") || registry.remove("account-a", 2, "container-c1") || registry.proxies["account-a"] != proxy { - t.Fatal("stale generation removed the current proxy") - } -} - -func TestMemoryProxyReplacesStaleGenerationWithoutOldCleanup(t *testing.T) { - registry := newMemoryProxyRegistry() - exit := gatewayProxyExit{Protocol: "http", Host: "127.0.0.1", Port: 1} - oldURL, oldCleanup, err := registry.configure("account-a", 1, "127.0.0.1", 0, exit, "network-n1") - if err != nil { - t.Fatal(err) - } - oldPort := registry.proxies["account-a"].listener.Addr().(*net.TCPAddr).Port - newURL, newCleanup, err := registry.configure("account-a", 2, "127.0.0.1", oldPort, exit, "network-n2") - if err != nil { - t.Fatal(err) - } - defer newCleanup() - if newURL != oldURL || registry.proxies["account-a"].bindingVersion != 2 { - t.Fatalf("stale proxy was not replaced: old=%q new=%q proxy=%#v", oldURL, newURL, registry.proxies["account-a"]) - } - oldCleanup() - if registry.proxies["account-a"].bindingVersion != 2 { - t.Fatal("old cleanup removed the replacement proxy") - } -} diff --git a/cmd/docker_gateway/__init__.py b/cmd/docker_gateway/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/cmd/docker_gateway/docker_client.py b/cmd/docker_gateway/docker_client.py new file mode 100644 index 0000000..1f82893 --- /dev/null +++ b/cmd/docker_gateway/docker_client.py @@ -0,0 +1,975 @@ +"""Minimal Docker Engine client and generation-fenced network helpers.""" + +from __future__ import annotations + +import http.client +import json +import logging +import os +import re +import socket +import threading +from collections.abc import Callable +from dataclasses import dataclass, field +from ipaddress import ip_interface +from urllib.parse import quote, urlencode + +MANAGED_LABEL = "io.creatorhub.managed" +RUNTIME_ID_LABEL = "io.creatorhub.runtime-id" +DISPLAY_NAME_LABEL = "io.creatorhub.display-name" +BINDING_VERSION_LABEL = "io.creatorhub.binding-version" +NETWORK_EXIT_LABEL = "io.creatorhub.network-exit-id" +PROXY_PORT_LABEL = "io.creatorhub.proxy-port" +NETWORK_ID_LABEL = "io.creatorhub.network-id" +NETWORK_ROLE_LABEL = "io.creatorhub.network-role" +GATEWAY_MEMBER_LABEL = "io.creatorhub.gateway-member" +BROWSER_NETWORK_ROLE = "browser" +BROWSER_PROXY_HOST = "docker-gateway" +NAME_PREFIX = "creatorhub-browser-" +RESERVATION_PREFIX = "creatorhub-reservation-" +RESERVATION_LABEL = "io.creatorhub.alias-reservation" +RESERVATION_GENERATION_LABEL = "io.creatorhub.reservation-generation" +RESERVATION_OWNER_LABEL = "io.creatorhub.reservation-owner" + +RUNTIME_ID_RE = re.compile(r"^[a-z0-9][a-z0-9-]{0,31}$") +NETWORK_NAME_RE = re.compile(r"^[A-Za-z0-9][A-Za-z0-9_.-]{0,63}$") +LOG = logging.getLogger("creatorhub.docker") + + +class DockerError(RuntimeError): + def __init__(self, message: str, status: int | None = None) -> None: + super().__init__(message) + self.status = status + + +class GenerationConflict(DockerError): + pass + + +class NetworkSetupError(DockerError): + def __init__(self, message: str, generation: TenantNetworkGeneration) -> None: + super().__init__(message) + self.generation = generation + + +class UnmanagedContainer(DockerError): + pass + + +@dataclass(frozen=True) +class DockerResponse: + status: int + reason: str + body: bytes + + +class UnixHTTPConnection(http.client.HTTPConnection): + def __init__(self, socket_path: str, timeout: float) -> None: + super().__init__("docker", timeout=timeout) + self.socket_path = socket_path + + def connect(self) -> None: + self.sock = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM) + self.sock.settimeout(self.timeout) + self.sock.connect(self.socket_path) + + +@dataclass +class TenantNetworkGeneration: + id: str = "" + name: str = "" + created: bool = False + connected_self: bool = False + connected_runtime: bool = False + gateway_members: list[str] = field(default_factory=list) + self_member: str = "" + runtime_attached: bool = False + + +class DockerClient: + def __init__(self, socket_path: str, api_version: str = "v1.43") -> None: + self.socket_path = socket_path + self.base_path = "/" + api_version.strip("/") + + def request( + self, + method: str, + path: str, + payload: object | None = None, + timeout: float = 30.0, + body_limit: int = 16 * 1024 * 1024, + ) -> DockerResponse: + encoded = ( + None + if payload is None + else json.dumps(payload, separators=(",", ":")).encode() + ) + connection: UnixHTTPConnection | None = None + try: + connection = UnixHTTPConnection(self.socket_path, timeout) + headers = {"Accept": "application/json"} + if encoded is not None: + headers["Content-Type"] = "application/json" + connection.request(method, self.base_path + path, encoded, headers) + response = connection.getresponse() + body = response.read(body_limit + 1) + if len(body) > body_limit: + raise DockerError( + "Docker response exceeded the configured limit", response.status + ) + return DockerResponse(response.status, response.reason, body) + except (OSError, http.client.HTTPException) as exc: + raise DockerError("Docker API request failed") from exc + finally: + if connection is not None: + connection.close() + + def expect( + self, + method: str, + path: str, + payload: object | None = None, + allowed: tuple[int, ...] = (204,), + ) -> None: + response = self.request(method, path, payload) + if response.status in allowed: + return + if response.status == 404: + raise FileNotFoundError(path) + raise DockerError( + f"Docker returned HTTP {response.status}: {response.body[:4096].decode('utf-8', 'replace').strip()}", + response.status, + ) + + def pull_if_missing(self, image: str) -> None: + encoded = quote(image, safe="") + response = self.request("GET", f"/images/{encoded}/json") + if response.status == 200: + return + if response.status != 404: + raise DockerError( + f"inspect image returned HTTP {response.status}", response.status + ) + repository, tag = split_image_ref(image) + query = {"fromImage": image if "@" in image else repository} + if "@" not in image and tag: + query["tag"] = tag + response = self.request( + "POST", + "/images/create?" + urlencode(query), + timeout=600.0, + body_limit=64 * 1024 * 1024, + ) + if response.status != 200: + raise DockerError( + f"pull image returned HTTP {response.status}: {response.body[:4096].decode('utf-8', 'replace').strip()}", + response.status, + ) + + def managed_container_state( + self, alias: str + ) -> tuple[str, dict[str, str], dict[str, str]]: + if not RUNTIME_ID_RE.fullmatch(alias): + raise ValueError("invalid runtime id") + response = self.request( + "GET", f"/containers/{quote(NAME_PREFIX + alias, safe='')}/json" + ) + if response.status == 404: + raise FileNotFoundError(alias) + if response.status != 200: + raise DockerError( + f"Docker inspect returned HTTP {response.status}", response.status + ) + try: + inspected = json.loads(response.body) + labels = inspected["Config"]["Labels"] or {} + container_id = inspected["Id"] + networks = inspected.get("NetworkSettings", {}).get("Networks", {}) or {} + if ( + not isinstance(labels, dict) + or not isinstance(networks, dict) + or not isinstance(container_id, str) + ): + raise TypeError("Docker inspect container metadata is invalid") + if not all( + isinstance(key, str) and isinstance(value, str) + for key, value in labels.items() + ): + raise TypeError("Docker inspect labels are invalid") + except (KeyError, TypeError, ValueError, json.JSONDecodeError) as exc: + raise DockerError( + "Docker inspect response is invalid", response.status + ) from exc + if labels.get(MANAGED_LABEL) != "true" or labels.get(RUNTIME_ID_LABEL) != alias: + raise UnmanagedContainer("refusing to operate on an unowned container") + network_ids: dict[str, str] = {} + for name, value in networks.items(): + if ( + not isinstance(name, str) + or not isinstance(value, dict) + or not isinstance(value.get("NetworkID", ""), str) + ): + raise DockerError("Docker inspect network metadata is invalid") + network_ids[name] = value.get("NetworkID", "") + return container_id, labels, network_ids + + def managed_container(self, alias: str) -> tuple[str, dict[str, str]]: + container_id, labels, _ = self.managed_container_state(alias) + return container_id, labels + + def trusted_gateway_member(self, container_id: str) -> bool: + response = self.request( + "GET", f"/containers/{quote(container_id, safe='')}/json" + ) + if response.status != 200: + return False + try: + inspected = json.loads(response.body) + if not isinstance(inspected, dict): + return False + config = inspected.get("Config", {}) + labels = config.get("Labels", {}) if isinstance(config, dict) else {} + return bool( + isinstance(inspected.get("Id"), str) + and isinstance(labels, dict) + and labels.get(GATEWAY_MEMBER_LABEL) == "true" + ) + except (TypeError, ValueError, json.JSONDecodeError): + return False + + def container_network_address(self, container_id: str, network_id: str) -> str: + response = self.request( + "GET", f"/containers/{quote(container_id, safe='')}/json" + ) + if response.status == 404: + raise FileNotFoundError("Docker container is missing") + if response.status != 200: + raise DockerError("inspect browser container failed", response.status) + try: + inspected = json.loads(response.body) + settings = inspected["NetworkSettings"] + networks = settings["Networks"] + if not isinstance(networks, dict): + raise TypeError("Docker container networks are invalid") + for value in networks.values(): + if isinstance(value, dict) and value.get("NetworkID") == network_id: + address = value.get("IPAddress") + if isinstance(address, str) and address: + return address + except (KeyError, TypeError, ValueError, json.JSONDecodeError) as exc: + raise DockerError("Docker container network metadata is invalid") from exc + raise GenerationConflict("browser container is not attached to its network") + + def inspect_tenant_network( + self, + base: str, + alias: str, + binding_version: int, + runtime_id: str, + self_name: str, + expected_id: str = "", + allow_unversioned: bool = False, + ) -> tuple[TenantNetworkGeneration, dict[str, str], bool]: + name = tenant_network_name(base, alias) + generation = TenantNetworkGeneration(id=expected_id, name=name) + reference = expected_id or name + response = self.request("GET", f"/networks/{quote(reference, safe='')}") + if response.status == 404: + return generation, {}, False + if response.status != 200: + raise DockerError( + f"Docker network inspect returned HTTP {response.status}", + response.status, + ) + try: + network = json.loads(response.body) + labels = network["Labels"] or {} + containers = network.get("Containers") or {} + network_id = network["Id"] + network_name = network["Name"] + if ( + not isinstance(labels, dict) + or not isinstance(containers, dict) + or not isinstance(network_id, str) + or not isinstance(network_name, str) + ): + raise TypeError("Docker network metadata is invalid") + if not all( + isinstance(key, str) and isinstance(value, str) + for key, value in labels.items() + ): + raise TypeError("Docker network labels are invalid") + except (KeyError, TypeError, ValueError, json.JSONDecodeError) as exc: + raise DockerError("Docker network inspect response is invalid") from exc + if ( + not network_id + or network_name != name + or network.get("Driver") != "bridge" + or network.get("Internal") + or network.get("Attachable") + or network.get("Ingress") + or labels.get(MANAGED_LABEL) != "true" + or labels.get(NETWORK_ROLE_LABEL) != BROWSER_NETWORK_ROLE + or labels.get(RUNTIME_ID_LABEL) != alias + ): + raise UnmanagedContainer( + "refusing to operate on an unowned browser network" + ) + if expected_id and network_id != expected_id: + raise GenerationConflict("network generation does not match request") + network_version = labels.get(BINDING_VERSION_LABEL, "") + if network_version != str(binding_version) and not ( + allow_unversioned and not network_version + ): + raise GenerationConflict("network generation does not match request") + generation.id = network_id + addresses: dict[str, str] = {} + for member_id, member in containers.items(): + if not isinstance(member, dict): + raise DockerError("Docker network member is invalid") + addresses[member_id] = str(member.get("IPv4Address", "")) + if member_id == runtime_id: + generation.runtime_attached = True + continue + if not self.trusted_gateway_member(member_id): + raise GenerationConflict( + "isolated network contains an untrusted member" + ) + generation.gateway_members.append(member_id) + if same_container_reference( + member_id, str(member.get("Name", "")), self_name + ): + generation.self_member = member_id + return generation, addresses, True + + def _finish_network_create( + self, + base: str, + alias: str, + self_name: str, + binding_version: int, + runtime_id: str, + response: DockerResponse, + generation: TenantNetworkGeneration, + ) -> TenantNetworkGeneration: + try: + created_id = json.loads(response.body)["Id"] + except (KeyError, TypeError, json.JSONDecodeError) as exc: + try: + observed, _, observed_exists = self.inspect_tenant_network( + base, alias, binding_version, runtime_id, self_name + ) + except ( + DockerError, + FileNotFoundError, + OSError, + TypeError, + ValueError, + KeyError, + ): + observed_exists = False + observed = TenantNetworkGeneration(name=generation.name) + if observed_exists: + generation = preserve_generation(generation, observed) + raise NetworkSetupError( + "create isolated browser network returned no id", generation + ) from exc + if not isinstance(created_id, str) or not created_id: + try: + observed, _, observed_exists = self.inspect_tenant_network( + base, alias, binding_version, runtime_id, self_name + ) + except ( + DockerError, + FileNotFoundError, + OSError, + TypeError, + ValueError, + KeyError, + ) as exc: + raise NetworkSetupError( + "create isolated browser network returned no id", generation + ) from exc + if observed_exists: + generation = preserve_generation(generation, observed) + raise NetworkSetupError( + "create isolated browser network returned no id", generation + ) + generation.id = created_id + generation.created = True + try: + observed, _, observed_exists = self.inspect_tenant_network( + base, alias, binding_version, runtime_id, self_name, created_id + ) + except ( + DockerError, + FileNotFoundError, + OSError, + TypeError, + ValueError, + KeyError, + ) as exc: + raise NetworkSetupError( + "created browser network could not be verified", generation + ) from exc + generation = preserve_generation(generation, observed) + if not observed_exists: + raise NetworkSetupError("created browser network disappeared", generation) + return generation + + def ensure_tenant_network( + self, + base: str, + alias: str, + self_name: str, + binding_version: int, + runtime_id: str = "", + expected_id: str = "", + allow_unversioned: bool = False, + ) -> tuple[TenantNetworkGeneration, str]: + if not self_name: + raise DockerError("gateway identity is invalid") + generation, addresses, exists = self.inspect_tenant_network( + base, + alias, + binding_version, + runtime_id, + self_name, + expected_id, + allow_unversioned, + ) + if not exists: + if expected_id: + raise GenerationConflict("network generation does not match request") + response = self.request( + "POST", + "/networks/create", + { + "Name": generation.name, + "CheckDuplicate": True, + "Driver": "bridge", + "Labels": { + MANAGED_LABEL: "true", + NETWORK_ROLE_LABEL: BROWSER_NETWORK_ROLE, + RUNTIME_ID_LABEL: alias, + BINDING_VERSION_LABEL: str(binding_version), + }, + }, + ) + if response.status != 201: + try: + observed, _, observed_exists = self.inspect_tenant_network( + base, alias, binding_version, runtime_id, self_name + ) + except ( + DockerError, + FileNotFoundError, + OSError, + TypeError, + ValueError, + KeyError, + ): + observed_exists = False + observed = TenantNetworkGeneration(name=generation.name) + if not observed_exists: + raise DockerError( + "create isolated browser network failed", response.status + ) + generation = preserve_generation(generation, observed) + else: + generation = self._finish_network_create( + base, + alias, + self_name, + binding_version, + runtime_id, + response, + generation, + ) + if runtime_id and not generation.runtime_attached: + generation.connected_runtime = True + try: + self.expect( + "POST", + f"/networks/{quote(generation.id, safe='')}/connect", + {"Container": runtime_id}, + (200,), + ) + except ( + DockerError, + FileNotFoundError, + OSError, + TypeError, + ValueError, + KeyError, + ) as exc: + raise self._network_setup_error( + base, + alias, + binding_version, + runtime_id, + self_name, + generation, + "connecting the browser to its isolated network failed", + ) from exc + generation.runtime_attached = True + if not generation.self_member: + generation.connected_self = True + try: + self.expect( + "POST", + f"/networks/{quote(generation.id, safe='')}/connect", + { + "Container": self_name, + "EndpointConfig": {"Aliases": [BROWSER_PROXY_HOST]}, + }, + (200,), + ) + except ( + DockerError, + FileNotFoundError, + OSError, + TypeError, + ValueError, + KeyError, + ) as exc: + raise self._network_setup_error( + base, + alias, + binding_version, + runtime_id, + self_name, + generation, + "connecting the gateway to its isolated network failed", + ) from exc + try: + observed, addresses, observed_exists = self.inspect_tenant_network( + base, + alias, + binding_version, + runtime_id, + self_name, + generation.id, + allow_unversioned, + ) + generation = preserve_generation(generation, observed) + if not observed_exists or not generation.self_member: + raise NetworkSetupError( + "Docker did not connect gateway to the isolated network", generation + ) + bind_host = str(ip_interface(addresses[generation.self_member]).ip) + except NetworkSetupError: + raise + except ( + DockerError, + FileNotFoundError, + OSError, + TypeError, + ValueError, + KeyError, + ) as exc: + raise NetworkSetupError( + "Docker could not verify the isolated network", generation + ) from exc + return generation, bind_host + + def _network_setup_error( + self, + base: str, + alias: str, + binding_version: int, + runtime_id: str, + self_name: str, + generation: TenantNetworkGeneration, + message: str, + ) -> NetworkSetupError: + try: + observed, _, exists = self.inspect_tenant_network( + base, alias, binding_version, runtime_id, self_name, generation.id + ) + generation = ( + preserve_generation(generation, observed) if exists else generation + ) + except ( + DockerError, + FileNotFoundError, + OSError, + TypeError, + ValueError, + KeyError, + ) as verification_error: + return NetworkSetupError( + f"{message}; network state verification failed: {verification_error}", + generation, + ) + return NetworkSetupError(message, generation) + + def disconnect_member( + self, + base: str, + alias: str, + binding_version: int, + runtime_id: str, + generation: TenantNetworkGeneration, + member: str, + self_name: str, + missing_ok: bool = False, + ) -> TenantNetworkGeneration: + current, _, exists = self.inspect_tenant_network( + base, alias, binding_version, runtime_id, self_name, generation.id + ) + if not exists: + if missing_ok: + return current + raise GenerationConflict("isolated network generation is missing") + if not same_network_members(current, generation): + raise GenerationConflict("isolated network generation changed") + if not member_present(current, member, runtime_id): + return current + self.expect( + "POST", + f"/networks/{quote(generation.id, safe='')}/disconnect", + {"Container": member, "Force": True}, + (200,), + ) + observed, _, observed_exists = self.inspect_tenant_network( + base, alias, binding_version, runtime_id, self_name, generation.id + ) + if not observed_exists or not same_network_members( + observed, generation_without_member(generation, member, runtime_id) + ): + raise GenerationConflict( + "Docker retained an isolated network member after disconnect" + ) + return observed + + def delete_tenant_network( + self, + base: str, + alias: str, + binding_version: int, + runtime_id: str, + generation: TenantNetworkGeneration, + self_name: str, + missing_ok: bool = False, + ) -> None: + current, _, exists = self.inspect_tenant_network( + base, alias, binding_version, runtime_id, self_name, generation.id + ) + if not exists: + if missing_ok: + return + raise GenerationConflict("isolated network generation is missing") + if current.runtime_attached or current.gateway_members: + raise GenerationConflict("isolated network still has members") + self.expect( + "DELETE", f"/networks/{quote(generation.id, safe='')}", allowed=(204,) + ) + _, _, exists = self.inspect_tenant_network( + base, alias, binding_version, runtime_id, self_name, generation.id + ) + if exists: + raise DockerError("Docker retained isolated browser network") + + +def split_image_ref(ref: str) -> tuple[str, str]: + if "@" in ref: + return ref.split("@", 1)[0], ref.split("@", 1)[1] + colon = ref.rfind(":") + slash = ref.rfind("/") + return (ref[:colon], ref[colon + 1 :]) if colon > slash else (ref, "") + + +def tenant_network_name(base: str, alias: str) -> str: + name = f"{base}-{alias}" + if not NETWORK_NAME_RE.fullmatch(name): + raise ValueError("isolated browser network name is invalid") + return name + + +def same_container_reference(container_id: str, name: str, reference: str) -> bool: + return bool( + reference + and ( + container_id == reference + or name == reference + or container_id.startswith(reference) + or reference.startswith(container_id) + ) + ) + + +def preserve_generation( + known: TenantNetworkGeneration, observed: TenantNetworkGeneration +) -> TenantNetworkGeneration: + if not observed.id: + observed.id = known.id + if not observed.name: + observed.name = known.name + observed.created |= known.created + observed.runtime_attached |= known.runtime_attached + observed.connected_runtime |= known.connected_runtime + observed.connected_self |= known.connected_self + if not observed.self_member: + observed.self_member = known.self_member + for member in known.gateway_members: + if not member_present(observed, member, ""): + observed.gateway_members.append(member) + return observed + + +def same_network_members( + current: TenantNetworkGeneration, expected: TenantNetworkGeneration +) -> bool: + return ( + current.id == expected.id + and current.name == expected.name + and current.runtime_attached == expected.runtime_attached + and current.self_member == expected.self_member + and set(current.gateway_members) == set(expected.gateway_members) + ) + + +def member_present( + generation: TenantNetworkGeneration, member: str, runtime_id: str +) -> bool: + if ( + generation.runtime_attached + and runtime_id + and same_container_reference(member, "", runtime_id) + ): + return True + return any( + same_container_reference(existing, "", member) + for existing in generation.gateway_members + ) or same_container_reference(member, "", generation.self_member) + + +def generation_without_member( + generation: TenantNetworkGeneration, member: str, runtime_id: str +) -> TenantNetworkGeneration: + result = TenantNetworkGeneration( + id=generation.id, + name=generation.name, + created=generation.created, + connected_self=generation.connected_self, + connected_runtime=generation.connected_runtime, + gateway_members=list(generation.gateway_members), + self_member=generation.self_member, + runtime_attached=generation.runtime_attached, + ) + if runtime_id and same_container_reference(member, "", runtime_id): + result.runtime_attached = False + if same_container_reference(member, "", result.self_member): + result.self_member = "" + result.gateway_members = [ + x for x in result.gateway_members if not same_container_reference(x, "", member) + ] + return result + + +def random_reservation_generation() -> str: + return os.urandom(16).hex() + + +class AliasReservationManager: + def __init__(self, docker: DockerClient, self_name: str) -> None: + self.docker = docker + self.self_name = self_name + self._locks: dict[str, threading.Lock] = {} + self._locks_guard = threading.Lock() + + def acquire(self, alias: str) -> Callable[[], None]: + with self._locks_guard: + lock = self._locks.setdefault(alias, threading.Lock()) + lock.acquire() + generation = random_reservation_generation() + created_id = "" + try: + response = self.docker.request( + "GET", f"/containers/{quote(self.self_name, safe='')}/json" + ) + if response.status != 200: + raise DockerError( + "inspect trusted gateway for alias reservation", response.status + ) + inspected = json.loads(response.body) + config = inspected.get("Config", {}) + labels = config.get("Labels", {}) if isinstance(config, dict) else {} + image = inspected.get("Image", "") + if ( + not isinstance(labels, dict) + or not isinstance(image, str) + or not image + or labels.get(GATEWAY_MEMBER_LABEL) != "true" + ): + raise DockerError("inspect trusted gateway for alias reservation") + reservation_payload = { + "Image": image, + "Labels": { + RESERVATION_LABEL: "true", + RUNTIME_ID_LABEL: alias, + RESERVATION_GENERATION_LABEL: generation, + RESERVATION_OWNER_LABEL: self.self_name, + }, + "HostConfig": {"NetworkMode": "none"}, + } + response = self.docker.request( + "POST", + "/containers/create?" + urlencode({"name": RESERVATION_PREFIX + alias}), + reservation_payload, + ) + if response.status == 409 and self._reclaim_stale(alias): + response = self.docker.request( + "POST", + "/containers/create?" + + urlencode({"name": RESERVATION_PREFIX + alias}), + reservation_payload, + ) + if response.status == 409: + raise GenerationConflict("browser alias is already in use") + if response.status != 201: + raise DockerError("create alias reservation failed", response.status) + created_id = json.loads(response.body).get("Id", "") + if not isinstance(created_id, str) or not created_id: + raise DockerError("Docker returned an invalid alias reservation id") + check = self.docker.request( + "GET", f"/containers/{quote(RESERVATION_PREFIX + alias, safe='')}/json" + ) + if check.status != 200: + raise DockerError( + "alias reservation could not be verified", check.status + ) + observed = json.loads(check.body) + observed_config = observed.get("Config", {}) + observed_labels = ( + observed_config.get("Labels", {}) + if isinstance(observed_config, dict) + else {} + ) + if ( + not isinstance(observed_labels, dict) + or observed.get("Id") != created_id + or observed_labels.get(RESERVATION_LABEL) != "true" + or observed_labels.get(RESERVATION_GENERATION_LABEL) != generation + ): + raise GenerationConflict( + "alias reservation generation is not immutable" + ) + except ( + DockerError, + FileNotFoundError, + OSError, + TypeError, + ValueError, + KeyError, + ): + try: + self._reconcile(alias, generation, created_id) + except ( + DockerError, + FileNotFoundError, + OSError, + TypeError, + ValueError, + KeyError, + ): + LOG.exception( + "failed to reconcile alias reservation", + extra={"alias": alias, "generation": generation}, + ) + lock.release() + raise + + def release() -> None: + try: + self._reconcile(alias, generation, created_id) + finally: + lock.release() + + return release + + def _reclaim_stale(self, alias: str) -> bool: + response = self.docker.request( + "GET", f"/containers/{quote(RESERVATION_PREFIX + alias, safe='')}/json" + ) + if response.status == 404: + return True + if response.status != 200: + raise DockerError("inspect alias reservation failed", response.status) + try: + observed = json.loads(response.body) + config = observed.get("Config", {}) + labels = config.get("Labels", {}) if isinstance(config, dict) else {} + observed_id = observed.get("Id") + except (TypeError, ValueError, json.JSONDecodeError) as exc: + raise DockerError("alias reservation inspect response is invalid") from exc + if ( + not isinstance(labels, dict) + or labels.get(RESERVATION_LABEL) != "true" + or labels.get(RUNTIME_ID_LABEL) != alias + or not isinstance(observed_id, str) + or not observed_id + ): + raise GenerationConflict("alias reservation generation changed") + owner = labels.get(RESERVATION_OWNER_LABEL) + if not isinstance(owner, str) or not owner: + return False + if owner != self.self_name: + owner_response = self.docker.request( + "GET", f"/containers/{quote(owner, safe='')}/json" + ) + if owner_response.status == 200: + return False + if owner_response.status != 404: + raise DockerError( + "inspect alias reservation owner failed", owner_response.status + ) + self.docker.expect( + "DELETE", + f"/containers/{quote(observed_id, safe='')}?force=1&v=0", + allowed=(204, 404), + ) + return ( + self.docker.request( + "GET", f"/containers/{quote(RESERVATION_PREFIX + alias, safe='')}/json" + ).status + == 404 + ) + + def _reconcile(self, alias: str, generation: str, created_id: str) -> None: + response = self.docker.request( + "GET", f"/containers/{quote(RESERVATION_PREFIX + alias, safe='')}/json" + ) + if response.status == 404: + return + if response.status != 200: + raise DockerError("inspect alias reservation failed", response.status) + try: + observed = json.loads(response.body) + config = observed.get("Config", {}) + labels = config.get("Labels", {}) if isinstance(config, dict) else {} + observed_id = observed.get("Id") + except (TypeError, ValueError, json.JSONDecodeError) as exc: + raise DockerError("alias reservation inspect response is invalid") from exc + if ( + not isinstance(labels, dict) + or not isinstance(observed_id, str) + or labels.get(RESERVATION_LABEL) != "true" + or labels.get(RUNTIME_ID_LABEL) != alias + or labels.get(RESERVATION_GENERATION_LABEL) != generation + or (created_id and observed_id != created_id) + ): + raise GenerationConflict("alias reservation generation changed") + self.docker.expect( + "DELETE", + f"/containers/{quote(observed_id, safe='')}?force=1&v=0", + allowed=(204,), + ) + if ( + self.docker.request( + "GET", f"/containers/{quote(RESERVATION_PREFIX + alias, safe='')}/json" + ).status + != 404 + ): + raise DockerError("Docker retained alias reservation") diff --git a/cmd/docker_gateway/douyin.py b/cmd/docker_gateway/douyin.py new file mode 100644 index 0000000..fdbf7e9 --- /dev/null +++ b/cmd/docker_gateway/douyin.py @@ -0,0 +1,1333 @@ +"""Douyin browser control built on a narrowly scoped CDP contract.""" + +from __future__ import annotations + +import http.client +import json +import logging +import math +import re +import threading +import time +import uuid +from collections import deque +from collections.abc import Callable +from concurrent.futures import ThreadPoolExecutor +from contextlib import contextmanager, suppress +from dataclasses import dataclass +from datetime import datetime, timezone +from urllib.parse import urlsplit + +import websocket + +from .docker_client import RUNTIME_ID_RE + +LOG = logging.getLogger("creatorhub.douyin") +ORIGIN = "https://www.douyin.com" +ORIGIN_URL = ORIGIN + "/" +IDENTITY_URL = ( + ORIGIN + "/aweme/v1/web/user/profile/self/?aid=6383&device_platform=webapp" +) +WORKS_PATH = "/aweme/v1/web/aweme/post/" +COMMENTS_PATH = "/aweme/v1/web/comment/list/" +RESPONSE_LIMIT = 1 << 20 +CONTROL_TIMEOUT = 15.0 +UID_RE = re.compile(r"^[1-9][0-9]{0,19}$") +ID_RE = re.compile(r"^[1-9][0-9]{0,63}$") +ACCOUNT_KEY_RE = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._:@/-]{0,127}$") +ACTIONS = frozenset( + {"follow", "dm", "reply_comment", "like_comment", "like_work", "repost"} +) + + +class DouyinError(RuntimeError): + def __init__(self, message: str): + super().__init__(message) + self.uncertain = False + + +LISTENER_ERRORS = ( + DouyinError, + OSError, + TypeError, + ValueError, + websocket.WebSocketException, +) + + +@dataclass(frozen=True) +class BrowserResponse: + status: int + body: str + challenge: str = "" + + +class CDPConnection: + def __init__(self, socket: websocket.WebSocket) -> None: + self.socket = socket + self._lock = threading.RLock() + self._next_id = 0 + self._pending: deque[dict] = deque() + + def command(self, method: str, params: dict | None = None) -> dict: + with self._lock: + self._next_id += 1 + command_id = self._next_id + self.socket.send( + json.dumps({"id": command_id, "method": method, "params": params or {}}) + ) + deadline = time.monotonic() + CONTROL_TIMEOUT + while True: + remaining = deadline - time.monotonic() + if remaining <= 0: + self._terminate_evaluation(method) + raise DouyinError(f"CDP command timed out: {method}") + message = self._take_pending(command_id) + if message is None: + remaining = deadline - time.monotonic() + if remaining <= 0: + self._terminate_evaluation(method) + raise DouyinError(f"CDP command timed out: {method}") + try: + message = self._receive(remaining, method) + except DouyinError: + self._terminate_evaluation(method) + raise + if not isinstance(message, dict): + raise DouyinError(f"CDP command returned invalid message: {method}") + if message.get("id") != command_id: + self._pending.append(message) + continue + if message.get("error"): + raise DouyinError(f"CDP command rejected: {method}") + result = message.get("result") + if not isinstance(result, dict): + raise DouyinError(f"CDP command returned invalid result: {method}") + return result + + def wait_event( + self, + method: str, + predicate: Callable[[dict], bool], + timeout: float = CONTROL_TIMEOUT, + ) -> dict: + with self._lock: + deadline = time.monotonic() + timeout + while True: + for index, message in enumerate(self._pending): + params = message.get("params", {}) + if ( + message.get("method") == method + and isinstance(params, dict) + and predicate(params) + ): + del self._pending[index] + return message + remaining = deadline - time.monotonic() + if remaining <= 0: + raise DouyinError(f"CDP event timed out: {method}") + message = self._receive(remaining, method) + if isinstance(message, dict): + params = message.get("params", {}) + if ( + message.get("method") == method + and isinstance(params, dict) + and predicate(params) + ): + return message + self._pending.append(message) + + def _terminate_evaluation(self, method: str) -> None: + if method != "Runtime.evaluate": + return + try: + self._next_id += 1 + command_id = self._next_id + self.socket.settimeout(1.0) + self.socket.send( + json.dumps( + { + "id": command_id, + "method": "Runtime.terminateExecution", + "params": {}, + } + ) + ) + deadline = time.monotonic() + 1.0 + while time.monotonic() < deadline: + message = json.loads(self.socket.recv()) + if isinstance(message, dict) and message.get("id") == command_id: + return + if isinstance(message, dict): + self._pending.append(message) + except ( + OSError, + TypeError, + websocket.WebSocketException, + json.JSONDecodeError, + ) as exc: + LOG.debug("failed to terminate timed-out CDP evaluation", exc_info=exc) + + def _take_pending(self, command_id: int) -> dict | None: + for index, message in enumerate(self._pending): + if message.get("id") == command_id: + del self._pending[index] + return message + return None + + def _receive(self, timeout: float, operation: str) -> object: + self.socket.settimeout(timeout) + try: + return json.loads(self.socket.recv()) + except (TimeoutError, websocket.WebSocketTimeoutException) as exc: + raise DouyinError(f"CDP command timed out during {operation}") from exc + except ( + OSError, + TypeError, + websocket.WebSocketException, + json.JSONDecodeError, + ) as exc: + raise DouyinError(f"CDP message failed during {operation}") from exc + + def evaluate(self, expression: str) -> object: + result = self.command( + "Runtime.evaluate", + { + "expression": expression, + "awaitPromise": True, + "returnByValue": True, + "userGesture": False, + }, + ) + if result.get("exceptionDetails"): + raise DouyinError("page evaluation failed") + value = result.get("result", {}).get("value") + if "value" not in result.get("result", {}): + raise DouyinError("page evaluation returned no value") + return value + + def close(self) -> None: + with suppress(OSError, websocket.WebSocketException): + self.socket.close() + + +class DouyinBrowser: + def __init__(self, endpoint: Callable[[str], str] | None = None) -> None: + self.endpoint = endpoint or ( + lambda alias: f"http://creatorhub-browser-{alias}:9222" + ) + + @contextmanager + def connection(self, alias: str): + connection = self._connect(alias) + try: + yield connection + finally: + connection.close() + + def _connect(self, alias: str) -> CDPConnection: + if not RUNTIME_ID_RE.fullmatch(alias): + raise DouyinError("browser alias is invalid") + base = self.endpoint(alias).rstrip("/") + try: + parsed = urlsplit(base) + port = parsed.port + except (TypeError, ValueError) as exc: + raise DouyinError("restricted browser endpoint is invalid") from exc + if ( + parsed.scheme != "http" + or not parsed.hostname + or not port + or parsed.username + or parsed.password + or parsed.query + or parsed.fragment + ): + raise DouyinError("restricted browser endpoint is invalid") + http_connection = http.client.HTTPConnection( + parsed.hostname, port, timeout=CONTROL_TIMEOUT + ) + try: + http_connection.request( + "GET", "/json/list", headers={"Accept": "application/json"} + ) + response = http_connection.getresponse() + if response.status != 200: + raise DouyinError("browser target discovery failed") + payload = response.read(64 * 1024 + 1) + except (OSError, http.client.HTTPException) as exc: + raise DouyinError("restricted browser unavailable") from exc + finally: + http_connection.close() + if len(payload) > 64 * 1024: + raise DouyinError("browser target discovery response is too large") + try: + targets = json.loads(payload) + except json.JSONDecodeError as exc: + raise DouyinError("browser target discovery response is invalid") from exc + if not isinstance(targets, list) or len(targets) > 32: + raise DouyinError("browser target discovery response is invalid") + page_targets = [ + target + for target in targets + if isinstance(target, dict) and target.get("type") == "page" + ] + douyin_targets = [ + target for target in page_targets if is_douyin_url(target.get("url", "")) + ] + if len(douyin_targets) > 1 or (not douyin_targets and len(page_targets) > 1): + raise DouyinError("browser has more than one page target") + target = ( + douyin_targets[0] + if douyin_targets + else page_targets[0] + if page_targets + else None + ) + if target is None: + raise DouyinError("browser page target is unavailable") + websocket_url = target.get("webSocketDebuggerUrl") + try: + parsed_ws = urlsplit( + websocket_url if isinstance(websocket_url, str) else "" + ) + websocket_port = parsed_ws.port or port + except (TypeError, ValueError) as exc: + raise DouyinError("browser target websocket is invalid") from exc + if ( + parsed_ws.scheme != "ws" + or not parsed_ws.hostname + or not websocket_port + or not parsed_ws.path.startswith("/devtools/page/") + ): + raise DouyinError("browser target websocket is invalid") + if parsed_ws.hostname in ("localhost", "127.0.0.1", "::1"): + parsed_ws = parsed_ws._replace(netloc=f"{parsed.hostname}:{websocket_port}") + if parsed_ws.hostname != parsed.hostname or ( + parsed_ws.port is not None and parsed_ws.port != port + ): + raise DouyinError("browser target websocket host is invalid") + page_target = parsed_ws.geturl() + try: + socket = websocket.create_connection( + page_target, + timeout=CONTROL_TIMEOUT, + origin="devtools://devtools", + enable_multithread=True, + ) + except (OSError, websocket.WebSocketException) as exc: + raise DouyinError("browser CDP connection failed") from exc + return CDPConnection(socket) + + def set_cookies(self, alias: str, cookies: list[dict]) -> None: + with self.connection(alias) as cdp: + cdp.command("Network.enable") + cdp.command("Network.clearBrowserCookies") + cdp.command("Page.enable") + navigation = cdp.command("Page.navigate", {"url": ORIGIN_URL}) + frame_id = navigation.get("frameId") + if not isinstance(frame_id, str) or navigation.get("errorText"): + raise DouyinError("browser navigation failed") + cdp.wait_event( + "Page.frameNavigated", + lambda params: ( + params.get("frame", {}).get("id") == frame_id + and is_douyin_url(params.get("frame", {}).get("url")) + ), + ) + deadline = time.monotonic() + CONTROL_TIMEOUT + last_error: DouyinError | None = None + while time.monotonic() < deadline: + try: + state = cdp.evaluate( + "({origin:location.origin,state:document.readyState})" + ) + if ( + isinstance(state, dict) + and state.get("origin") == ORIGIN + and state.get("state") in {"interactive", "complete"} + ): + break + except DouyinError as exc: + last_error = exc + time.sleep(0.1) + else: + if last_error: + raise DouyinError("browser did not reach Douyin") from last_error + raise DouyinError("browser did not reach Douyin") + values = [] + for cookie in cookies: + value = { + "name": cookie["name"], + "value": cookie["value"], + "url": ORIGIN_URL, + "domain": cookie["domain"], + "path": cookie.get("path") or "/", + "secure": bool(cookie.get("secure")), + "httpOnly": bool(cookie.get("http_only")), + } + if cookie.get("same_site"): + value["sameSite"] = cookie["same_site"] + if cookie.get("expires"): + value["expires"] = cookie["expires"] + values.append(value) + cdp.command("Network.setCookies", {"cookies": values}) + + def get(self, alias: str, target: str) -> BrowserResponse: + with self.connection(alias) as cdp: + if cdp.evaluate("location.origin") != ORIGIN: + raise DouyinError("restricted browser origin changed") + expression = f"""(async()=>{{ + const r=await fetch({json.dumps(target)},{{credentials:'include',redirect:'error'}}); + if(!r.body)return {{status:r.status,body:'',too_large:false}}; + const reader=r.body.getReader(), decoder=new TextDecoder(); let size=0, body=''; + for(;;){{const item=await reader.read();if(item.done)break; + if(size+item.value.byteLength>={RESPONSE_LIMIT}){{await reader.cancel();return {{too_large:true}};}} + size+=item.value.byteLength;body+=decoder.decode(item.value,{{stream:true}}); + }} + body+=decoder.decode();return {{status:r.status,body,too_large:false}}; + }})()""" + result = cdp.evaluate(expression) + if ( + not isinstance(result, dict) + or result.get("too_large") + or not isinstance(result.get("status"), int) + ): + raise DouyinError("restricted browser fetch failed") + status = result["status"] + if 300 <= status < 400: + raise DouyinError("restricted browser fetch redirected") + body = result.get("body") + if not isinstance(body, str): + raise DouyinError("restricted browser fetch returned invalid body") + return BrowserResponse(status, body, detect_challenge(status, body)) + + def identity(self, alias: str, expected_uid: str | None = None) -> dict: + response = self.get(alias, IDENTITY_URL) + try: + payload = json.loads(response.body) + except json.JSONDecodeError as exc: + raise DouyinError("Douyin identity response is invalid") from exc + user = payload.get("user") if isinstance(payload, dict) else None + uid = str(user.get("uid", "")) if isinstance(user, dict) else "" + if ( + response.status != 200 + or not isinstance(payload, dict) + or payload.get("status_code") != 0 + or not UID_RE.fullmatch(uid) + ): + raise DouyinError("Douyin login is not valid") + sec_uid = str(user.get("sec_uid", "")) if isinstance(user, dict) else "" + unique_id = str(user.get("unique_id", "")) if isinstance(user, dict) else "" + if not ACCOUNT_KEY_RE.fullmatch(sec_uid) or ( + unique_id and not ACCOUNT_KEY_RE.fullmatch(unique_id) + ): + raise DouyinError("Douyin identity response is invalid") + if expected_uid and uid != expected_uid: + raise DouyinError("Douyin identity does not match the expected account") + return { + "uid": uid, + "sec_uid": sec_uid, + "unique_id": unique_id, + "nickname": user.get("nickname", "") if isinstance(user, dict) else "", + "short_id": user.get("short_id", "") if isinstance(user, dict) else "", + } + + def action( + self, + alias: str, + expected_uid: str, + action: str, + target_uid: str = "", + comment_id: str = "", + work_id: str = "", + text: str = "", + confirm: bool = False, + ) -> dict: + if not UID_RE.fullmatch(expected_uid): + raise DouyinError("expected account UID is invalid") + if target_uid and not UID_RE.fullmatch(target_uid): + raise DouyinError("target UID is invalid") + if work_id and not ID_RE.fullmatch(work_id): + raise DouyinError("work ID is invalid") + if comment_id and not ID_RE.fullmatch(comment_id): + raise DouyinError("comment ID is invalid") + if action not in ACTIONS: + raise DouyinError("Douyin action is invalid") + if action in {"follow", "dm"} and not target_uid: + raise DouyinError("target UID is required") + if action in {"like_work", "repost"} and not work_id: + raise DouyinError("work ID is required") + if action in {"reply_comment", "like_comment"} and ( + not work_id or not comment_id or not target_uid + ): + raise DouyinError("comment target is incomplete") + if action in {"dm", "reply_comment", "repost"} and ( + not text.strip() or len(text) > 1000 + ): + raise DouyinError("action text is invalid") + identity = self.identity(alias, expected_uid) + if action == "follow": + params = { + "expected": expected_uid, + "target": target_uid, + "check": not confirm, + } + value = ( + self._confirmed_evaluate(alias, follow_expression(params)) + if confirm + else self._evaluate(alias, follow_expression(params)) + ) + elif action == "dm": + if not confirm: + return { + "action": "preview", + "sender_uid": identity["uid"], + "uid": target_uid, + "text": text, + } + params = { + "uid": target_uid, + "text": text, + "confirm": True, + "limit": 20, + "cursor": "9223372036854775807", + "action": "send", + } + value = self._confirmed_evaluate(alias, im_expression(params, expected_uid)) + else: + params = { + "action": action, + "expected": expected_uid, + "target": target_uid, + "work": work_id, + "comment": comment_id, + "text": text, + "confirm": confirm, + } + if not confirm: + return { + "action": "preview", + "sender_uid": identity["uid"], + "target_uid": target_uid, + "target_work_id": work_id, + "target_comment_id": comment_id, + "text": text, + } + value = self._confirmed_evaluate(alias, action_expression(params)) + if not isinstance(value, dict): + raise DouyinError("Douyin action response is invalid") + return value + + def _evaluate(self, alias: str, expression: str) -> object: + with self.connection(alias) as cdp: + if cdp.evaluate("location.origin") != ORIGIN: + raise DouyinError("restricted browser origin changed") + return cdp.evaluate(expression) + + def _confirmed_evaluate(self, alias: str, expression: str) -> object: + try: + return self._evaluate(alias, expression) + except DouyinError as exc: + exc.uncertain = True + raise + + +class DouyinSubscription: + def __init__( + self, + browser: DouyinBrowser, + alias: str, + uid: str, + pending: list[dict] | None = None, + ) -> None: + self.browser = browser + self.alias = alias + self.uid = uid + # A stable key lets a fresh wrapper dispose a stale listener left on the + # same page. Recovery disposes the old state before installing a new one. + self.key = "__creatorhub_notice_sub_" + alias + self._connection_lock = threading.RLock() + self.queue: deque[dict] = deque(pending or []) + self.condition = threading.Condition() + self.stopped = threading.Event() + self._epoch = 0 + self._initial_boundary_pending = True + self._browser_inflight: set[str] = set() + self._browser_inflight_lock = threading.Lock() + self._detail_pool = ThreadPoolExecutor( + max_workers=4, thread_name_prefix=f"creatorhub-notice-details-{alias}" + ) + try: + self.connection = self._open_listener() + except Exception: + self._detail_pool.shutdown(wait=False, cancel_futures=True) + raise + # Events observed before the consumer establishes its durable boundary + # are explicitly classified as baseline and must never trigger writes. + self._put( + { + "kind": "baseline", + "reason": "listener_start", + "uid": self.uid, + "boundary_at": datetime.now(timezone.utc).isoformat(), + } + ) + self.thread = threading.Thread( + target=self._run, name=f"creatorhub-notices-{alias}", daemon=True + ) + self.thread.start() + + def _open_listener(self) -> CDPConnection: + connection = self.browser._connect(self.alias) + try: + result = connection.evaluate(install_expression(self.key, self.uid)) + if not isinstance(result, str): + raise DouyinError("notification listener returned invalid state") + state = json.loads(result) + if not isinstance(state, dict) or not state.get("connected"): + raise DouyinError("Douyin notification connection is not ready") + return connection + except ( + DouyinError, + OSError, + TypeError, + ValueError, + websocket.WebSocketException, + ): + connection.close() + raise + + def _get_connection(self) -> CDPConnection: + with self._connection_lock: + return self.connection + + def _run(self) -> None: + wait = wait_expression(self.key) + while not self.stopped.is_set(): + try: + raw = self._get_connection().evaluate(wait) + events = json.loads(raw) if isinstance(raw, str) else raw + if not isinstance(events, list): + raise DouyinError("notification listener returned invalid events") + initial_boundary = self._initial_boundary_pending + self._initial_boundary_pending = False + for event in events: + if ( + initial_boundary + and isinstance(event, dict) + and event.get("kind") == "push" + ): + event = dict(event) + event["baseline"] = True + self._handle(event) + except ( + DouyinError, + OSError, + TypeError, + ValueError, + websocket.WebSocketException, + ) as exc: + if self.stopped.is_set(): + break + self._put({"kind": "error", "reason": str(exc)}) + self._recover() + self._dispose_current() + + def _recover(self) -> None: + # Dispose the old same-page handlers before installing a new state with + # the stable key; disposing afterward would remove the fresh handlers. + with self._connection_lock: + old = self.connection + try: + old.evaluate(dispose_expression(self.key)) + except DouyinError: + LOG.debug( + "old notification listener disposal was not available", + exc_info=True, + ) + finally: + old.close() + self._initial_boundary_pending = True + delay = 0.5 + while not self.stopped.is_set(): + if self.stopped.wait(delay): + return + try: + connection = self._open_listener() + except ( + DouyinError, + OSError, + TypeError, + ValueError, + websocket.WebSocketException, + ) as exc: + self._put({"kind": "error", "reason": str(exc)}) + delay = min(delay * 2, 30.0) + continue + with self._connection_lock: + self.connection = connection + self._epoch += 1 + self._put({"kind": "reconnected", "uid": self.uid}) + # The boundary is a separate event so the control plane never treats + # a transport reconnect itself as proof that new notices are safe. + self._put( + { + "kind": "baseline", + "reason": "listener_reconnected", + "uid": self.uid, + "boundary_at": datetime.now(timezone.utc).isoformat(), + } + ) + return + + def _dispose_current(self) -> None: + connection = self._get_connection() + try: + connection.evaluate(dispose_expression(self.key)) + except DouyinError: + LOG.debug("notification listener disposal was not available", exc_info=True) + finally: + connection.close() + + def _handle(self, event: object) -> None: + if not isinstance(event, dict): + raise DouyinError("notification event is invalid") + kind = event.get("kind") + delivery_id = event.get("delivery_id") + if kind in {"open", "close"}: + self._put({"kind": kind}) + if delivery_id and hasattr(self, "_connection_lock"): + self._ack_browser_event(delivery_id, self._get_connection()) + if kind == "close": + raise DouyinError("Douyin notification connection closed") + return + if kind == "error": + self._put( + { + "kind": "error", + "reason": str(event.get("reason", "notification continuity gap")), + "continuity": "gap", + } + ) + if delivery_id and hasattr(self, "_connection_lock"): + self._ack_browser_event(delivery_id, self._get_connection()) + return + if kind != "push": + raise DouyinError("notification event kind is invalid") + if not hasattr(self, "_browser_inflight_lock"): + self._browser_inflight_lock = threading.Lock() + if not hasattr(self, "_browser_inflight"): + self._browser_inflight = set() + if isinstance(delivery_id, str) and delivery_id: + with self._browser_inflight_lock: + if delivery_id in self._browser_inflight: + return + self._browser_inflight.add(delivery_id) + epoch = getattr(self, "_epoch", 0) + baseline = bool(event.get("baseline")) + if hasattr(self, "_detail_pool"): + self._detail_pool.submit(self._process_push_async, event, epoch, baseline) + else: + success = self._process_push(event, self._get_connection(), baseline) + if delivery_id: + if success: + self._ack_browser_event(delivery_id, self._get_connection()) + else: + self._retry_browser_event(delivery_id, self._get_connection()) + if delivery_id: + with self._browser_inflight_lock: + self._browser_inflight.discard(delivery_id) + + def _handle_push_error( + self, + error: Exception, + delivery_id: object, + connection: CDPConnection | None, + ) -> None: + self._put({"kind": "error", "reason": str(error), "continuity": "gap"}) + if isinstance(delivery_id, str) and connection is not None: + self._retry_browser_event(delivery_id, connection) + + def _process_push( + self, event: dict, connection: CDPConnection, baseline: bool + ) -> bool: + try: + ids = notice_ids(event) + except ( + DouyinError, + KeyError, + TypeError, + ValueError, + websocket.WebSocketException, + ) as exc: + self._put({"kind": "error", "reason": str(exc), "continuity": "gap"}) + return False + # Detail lookup is isolated per platform notification. A malformed or + # temporarily unavailable detail must not discard its siblings. + complete = True + for notice_id in ids: + try: + for notice in self._details([notice_id], connection): + normalized = normalize_notice(notice) + if normalized is not None: + self._put( + { + "kind": "notice", + "notice": normalized, + "baseline": baseline, + } + ) + except ( + DouyinError, + OSError, + TypeError, + ValueError, + websocket.WebSocketException, + ) as exc: + self._put( + { + "kind": "error", + "reason": str(exc), + "event_key": notice_id, + "continuity": "gap", + } + ) + complete = False + return complete + + def _process_push_async(self, event: dict, epoch: int, baseline: bool) -> None: + connection: CDPConnection | None = None + success = False + delivery_id = event.get("delivery_id") + try: + connection = self.browser._connect(self.alias) + success = self._process_push( + event, + connection, + baseline or epoch != getattr(self, "_epoch", 0), + ) + if isinstance(delivery_id, str): + if success: + self._ack_browser_event(delivery_id, connection) + else: + self._retry_browser_event(delivery_id, connection) + except LISTENER_ERRORS as exc: + self._handle_push_error(exc, delivery_id, connection) + finally: + if connection is not None: + connection.close() + if isinstance(delivery_id, str): + with self._browser_inflight_lock: + self._browser_inflight.discard(delivery_id) + + def _ack_browser_event( + self, delivery_id: object, connection: CDPConnection + ) -> None: + if not isinstance(delivery_id, str) or not delivery_id: + return + try: + result = connection.evaluate(ack_expression(self.key, [delivery_id])) + if result not in (True, "true"): + raise DouyinError("notification acknowledgement failed") + except ( + DouyinError, + OSError, + TypeError, + ValueError, + websocket.WebSocketException, + ): + LOG.warning("notification acknowledgement failed", exc_info=True) + + def _retry_browser_event( + self, delivery_id: object, connection: CDPConnection + ) -> None: + if not isinstance(delivery_id, str) or not delivery_id: + return + try: + result = connection.evaluate(retry_expression(self.key, [delivery_id])) + if result not in (True, "true"): + raise DouyinError("notification retry acknowledgement failed") + except ( + DouyinError, + OSError, + TypeError, + ValueError, + websocket.WebSocketException, + ): + LOG.warning("notification retry marker failed", exc_info=True) + + def pending(self) -> list[dict]: + with self.condition: + return list(self.queue) + + def _details( + self, ids: list[str], connection: CDPConnection | None = None + ) -> list[dict]: + expected = set(ids) + connection = connection or self._get_connection() + for attempt in range(2): + raw = connection.evaluate(details_expression(ids)) + if not isinstance(raw, str): + raise DouyinError("notification details returned invalid data") + try: + response = json.loads(raw) + body = json.loads(response["body"]) + except (KeyError, TypeError, ValueError, json.JSONDecodeError) as exc: + raise DouyinError("notification details could not be parsed") from exc + if ( + not isinstance(response, dict) + or not isinstance(body, dict) + or response.get("status") != 200 + or body.get("status_code") != 0 + ): + raise DouyinError("notification details request failed") + notices = body.get("notice_list_v2") + if not isinstance(notices, list): + raise DouyinError("notification details list is invalid") + normalized: list[dict] = [] + found: set[str] = set() + for notice in notices: + if ( + not isinstance(notice, dict) + or str(notice.get("user_id")) != self.uid + ): + raise DouyinError("notification identity changed") + nid = str(notice.get("nid_str") or notice.get("nid") or "") + if not nid.isascii() or not nid.isdecimal(): + raise DouyinError("notification ID is invalid") + found.add(nid) + normalized.append(notice) + if found == expected: + return normalized + if attempt == 0: + time.sleep(0.2) + continue + raise DouyinError("notification details are incomplete") + raise DouyinError("notification details are incomplete") + + def _put(self, item: dict) -> None: + item = dict(item) + item.setdefault("delivery_id", uuid.uuid4().hex) + if item.get("kind") == "notice": + notice = item.get("notice") + if isinstance(notice, dict): + notice.setdefault( + "gateway_received_at", datetime.now(timezone.utc).isoformat() + ) + with self.condition: + if len(self.queue) >= 1000: + # Keep the backlog visible; replace only one oldest item with a + # gap marker instead of silently dropping the entire queue. + dropped = getattr(self, "_overflow_count", 0) + 1 + self._overflow_count = dropped + self.queue[0] = { + "kind": "error", + "reason": "notification queue overflow", + "continuity": "gap", + "dropped": dropped, + "delivery_id": uuid.uuid4().hex, + } + else: + self.queue.append(item) + self.condition.notify_all() + + def poll(self, limit: int, wait_seconds: float) -> list[dict]: + deadline = time.monotonic() + wait_seconds + with self.condition: + while not self.queue and not self.stopped.is_set(): + remaining = deadline - time.monotonic() + if remaining <= 0: + break + self.condition.wait(remaining) + return list(self.queue)[:limit] + + def ack(self, delivery_ids: list[str]) -> None: + ids = {value for value in delivery_ids if isinstance(value, str) and value} + if not ids: + return + with self.condition: + self.queue = deque( + item for item in self.queue if item.get("delivery_id") not in ids + ) + + def stop(self) -> None: + self.stopped.set() + with self._connection_lock: + connection = self.connection + connection.close() + with self.condition: + self.condition.notify_all() + self.thread.join(timeout=2.0) + if self.thread.is_alive(): + LOG.error( + "notification listener did not stop within timeout", + extra={"alias": self.alias}, + ) + pool = getattr(self, "_detail_pool", None) + if pool is not None: + pool.shutdown(wait=True, cancel_futures=True) + + +class SubscriptionManager: + def __init__(self, browser: DouyinBrowser) -> None: + self.browser = browser + self._lock = threading.RLock() + self._items: dict[str, DouyinSubscription] = {} + + def start(self, alias: str, uid: str) -> dict: + with self._lock: + previous = self._items.pop(alias, None) + pending = previous.pending() if previous else [] + if previous: + previous.stop() + item = DouyinSubscription(self.browser, alias, uid, pending=pending) + self._items[alias] = item + return {"connected": True, "alias": alias, "uid": uid} + + def poll(self, alias: str, limit: int, wait_seconds: float) -> list[dict]: + with self._lock: + item = self._items.get(alias) + if not item: + raise DouyinError("notification listener is not running") + return item.poll(limit, wait_seconds) + + def ack(self, alias: str, delivery_ids: list[str]) -> None: + with self._lock: + item = self._items.get(alias) + if not item: + raise DouyinError("notification listener is not running") + item.ack(delivery_ids) + + def stop(self, alias: str) -> None: + with self._lock: + item = self._items.pop(alias, None) + if item: + item.stop() + + def close(self) -> None: + with self._lock: + aliases = list(self._items) + for alias in aliases: + self.stop(alias) + + +def _notice_id(value: object) -> str: + if type(value) is int and value > 0: + return str(value) + if isinstance(value, str) and ID_RE.fullmatch(value): + return value + return "" + + +def normalize_notice(notice: object) -> dict | None: + if not isinstance(notice, dict): + raise DouyinError("notification detail is invalid") + if notice.get("comment"): + kind, event_type = "comment", "comment" + elif notice.get("follow"): + kind, event_type = "follow", "follow" + elif notice.get("digg"): + kind, event_type = "digg", "like" + elif notice.get("share"): + kind, event_type = "share", "repost" + elif ( + notice.get("dm") + or notice.get("message") + or notice.get("im") + or notice.get("chat") + ): + kind, event_type = next( + (candidate, "dm") + for candidate in ("dm", "message", "im", "chat") + if notice.get(candidate) + ) + else: + return None + detail = notice.get(kind) + if not isinstance(detail, dict): + raise DouyinError("notification detail payload is invalid") + if event_type == "dm": + sender = ( + detail.get("from_user") or detail.get("sender") or detail.get("user") or {} + ) + if isinstance(sender, list): + sender = sender[0] if sender else {} + if not isinstance(sender, dict): + raise DouyinError("direct-message sender is invalid") + message = ( + detail.get("text") + or detail.get("content") + or detail.get("message") + or notice.get("text") + or "" + ) + if isinstance(message, dict): + message = message.get("text") or message.get("content") or "" + if not isinstance(message, str) or len(message) > 100000: + raise DouyinError("direct-message text is invalid") + event_key = _notice_id( + notice.get("message_id") + or notice.get("msg_id") + or detail.get("message_id") + or detail.get("msg_id") + or notice.get("nid_str") + or notice.get("nid") + ) + if not event_key: + raise DouyinError("direct-message ID is invalid") + result = { + "event_key": event_key, + "event_type": "dm", + "interactor_uid": _notice_id(sender.get("uid") or sender.get("user_id")), + "comment_id": "", + "work_id": "", + "message_text": message, + } + create_time = notice.get("create_time") + if isinstance(create_time, (int, float)) and not isinstance(create_time, bool): + try: + timestamp = float(create_time) + if math.isfinite(timestamp) and timestamp > 0: + result["platform_event_at"] = datetime.fromtimestamp( + timestamp, timezone.utc + ).isoformat() + except (OverflowError, OSError, ValueError) as exc: + raise DouyinError("notification timestamp is invalid") from exc + return result + users = detail.get("from_user") or [] + if isinstance(users, dict): + users = [users] + if not isinstance(users, list): + raise DouyinError("notification users are invalid") + comment = detail.get("comment") or {} + if not isinstance(comment, dict): + raise DouyinError("notification comment is invalid") + if not users and comment.get("user"): + users = [comment["user"]] + uids = { + uid + for user in users + if isinstance(user, dict) + for uid in [_notice_id(user.get("uid") or user.get("user_id"))] + if uid + } + event_key = _notice_id(notice.get("nid_str") or notice.get("nid")) + if not event_key: + raise DouyinError("notification ID is invalid") + work = detail.get("aweme") or notice.get("aweme") or {} + if not isinstance(work, dict): + work = {} + comment_id = next( + ( + value + for value in ( + notice.get("comment_id"), + detail.get("comment_id"), + comment.get("cid_str"), + comment.get("cid"), + ) + if _notice_id(value) + ), + "", + ) + work_id = next( + ( + value + for value in ( + notice.get("aweme_id"), + detail.get("aweme_id"), + work.get("aweme_id"), + work.get("id"), + ) + if _notice_id(value) + ), + "", + ) + result = { + "event_key": event_key, + "event_type": event_type, + "interactor_uid": next(iter(uids), "") if len(uids) == 1 else "", + "comment_id": _notice_id(comment_id), + "work_id": _notice_id(work_id), + } + create_time = notice.get("create_time") + if isinstance(create_time, (int, float)) and not isinstance(create_time, bool): + try: + timestamp = float(create_time) + if math.isfinite(timestamp) and timestamp > 0: + result["platform_event_at"] = datetime.fromtimestamp( + timestamp, timezone.utc + ).isoformat() + except (OverflowError, OSError, ValueError) as exc: + raise DouyinError("notification timestamp is invalid") from exc + return result + + +def is_douyin_url(value: object) -> bool: + if not isinstance(value, str): + return False + try: + parsed = urlsplit(value) + return ( + parsed.scheme == "https" + and parsed.hostname == "www.douyin.com" + and parsed.port is None + and parsed.username is None + and parsed.password is None + ) + except (TypeError, ValueError): + return False + + +def detect_challenge(status: int, body: str) -> str: + if status == 429: + return "" + if status == 412: + return "captcha" + lowered = body.lower() + if any(token in lowered for token in ("captcha", "verify_code", "验证码")): + return "captcha" + if any(token in lowered for token in ("device_verify", "device_check", "设备验证")): + return "device" + return "" + + +def notice_ids(event: dict) -> list[str]: + try: + payload = json.loads(event["payload"]) + except (KeyError, TypeError, ValueError, json.JSONDecodeError) as exc: + raise DouyinError("notification push payload is invalid") from exc + if not isinstance(payload, dict): + raise DouyinError("notification push payload is invalid") + if event.get("service") == 20313: + notices = payload.get("notices", []) + if not isinstance(notices, list): + raise DouyinError("notification push list is invalid") + selected = [ + notice + for notice in notices + if isinstance(notice, dict) + and {str(group) for group in notice.get("effect_groups", [])} + & {"960", "961"} + ] + elif event.get("service") == 20003 and payload.get("notice_type") in { + 45, + 31, + 9009, + 9002, + 514, + 9067, + }: + selected = [payload] + else: + return [] + ids: list[str] = [] + for notice in selected: + value = notice.get("notice_id_str") + if not isinstance(value, str) or not value.isascii() or not value.isdecimal(): + raise DouyinError("notification ID is invalid") + ids.append(value) + return list(dict.fromkeys(ids)) + + +def install_expression(key: str, uid: str) -> str: + return INSTALL_SCRIPT.replace("KEY_VALUE", json.dumps(key)).replace( + "UID_VALUE", json.dumps(uid) + ) + + +def wait_expression(key: str) -> str: + return WAIT_SCRIPT.replace("KEY_VALUE", json.dumps(key)) + + +def dispose_expression(key: str) -> str: + return f"(() => {{ const key={json.dumps(key)}; window[key]?.dispose(); delete window[key]; return true; }})()" + + +def ack_expression(key: str, ids: list[str]) -> str: + if ( + not isinstance(key, str) + or not key + or not ids + or any(not isinstance(value, str) or not value for value in ids) + ): + raise ValueError("notification delivery IDs are invalid") + return ACK_SCRIPT.replace("KEY_VALUE", json.dumps(key)).replace( + "IDS_VALUE", json.dumps(ids) + ) + + +def retry_expression(key: str, ids: list[str]) -> str: + if ( + not isinstance(key, str) + or not key + or not ids + or any(not isinstance(value, str) or not value for value in ids) + ): + raise ValueError("notification delivery IDs are invalid") + return RETRY_SCRIPT.replace("KEY_VALUE", json.dumps(key)).replace( + "IDS_VALUE", json.dumps(ids) + ) + + +def details_expression(ids: list[str]) -> str: + if not ids or any( + not isinstance(value, str) or not value.isascii() or not value.isdecimal() + for value in ids + ): + raise ValueError("notification IDs are invalid") + return DETAILS_SCRIPT.replace("IDS_VALUE", json.dumps(ids)) + + +def follow_expression(params: dict) -> str: + return FOLLOW_SCRIPT.replace("PARAMS_VALUE", json.dumps(params, ensure_ascii=True)) + + +def im_expression(params: dict, expected_uid: str) -> str: + return IM_SCRIPT.replace("EXPECTED_UID_VALUE", json.dumps(expected_uid)).replace( + "PARAMS_VALUE", json.dumps(params, ensure_ascii=True) + ) + + +def action_expression(params: dict) -> str: + return ACTION_SCRIPT.replace("PARAMS_VALUE", json.dumps(params, ensure_ascii=True)) + + +INSTALL_SCRIPT = r"""(async()=>{try{ + if(location.origin!=="https://www.douyin.com")throw Error("WRONG_ORIGIN"); + const key=KEY_VALUE,uid=UID_VALUE,chunks=window.webpackChunkdouyin_web;if(!chunks)throw Error("RUNTIME_MISSING");const old=window[key],carry=Array.isArray(old?.queue)?old.queue.filter(e=>e&&e.delivery_id):[];old?.dispose?.();delete window[key];let req; + chunks.push([["creatorhub-notice-"+Date.now()],{},r=>{req=r;}]);chunks.pop(); + const entries=Object.entries(req.m||{});const entry=entries.find(([,f])=>String(f).includes("NOTICE_PUSH_EVENT_NAMES:function")); + const codec=entries.find(([,f])=>{const s=String(f);return s.includes(".decodedFrame=")&&s.includes(".encodeFrame=");}); + if(!entry||!codec)throw Error("SDK_CHANGED");const C=req(entry[0]).NoticeFrontier,decode=req(codec[0]).decodedFrame,f=C.frontierInstance; + if(!f||String(f._options.deviceID)!==uid)throw Error("SOCKET_NOT_READY"); + const state={queue:carry,wake:null,C,f,uid,dropped:0,nextDelivery:0};const emit=e=>{const item={...e,delivered:false,delivery_id:"browser-"+(++state.nextDelivery)};if(state.queue.length>=1000){state.queue.shift();state.dropped++;state.queue.push({kind:"error",reason:"QUEUE_OVERFLOW",continuity:"gap",dropped:state.dropped,delivered:false,delivery_id:"browser-"+(++state.nextDelivery)});}else state.queue.push(item);if(state.wake)state.wake();}; + const message=e=>{try{const frame=decode(new Uint8Array(e.data));if(frame.service===20313||frame.service===20003)emit({kind:"push",service:frame.service,payload:new TextDecoder().decode(frame.payload)});}catch(_){emit({kind:"error",reason:"FRAME_DECODE_FAILED",continuity:"gap"});}}; + const open=()=>emit({kind:"open"}),close=()=>emit({kind:"close"});f.addEventListener("message",message);f.addEventListener("open",open);f.addEventListener("close",close); + state.dispose=()=>{f.removeEventListener("message",message);f.removeEventListener("open",open);f.removeEventListener("close",close);if(state.wake)state.wake();};window[key]=state; + return JSON.stringify({connected:f.readyState===f.OPEN}); +}catch(e){return JSON.stringify({bridge_error:String(e.message||e)});}})()""" + + +WAIT_SCRIPT = r"""(async()=>{const s=window[KEY_VALUE];if(!s||s.C.frontierInstance!==s.f||String(s.f._options.deviceID)!==s.uid)throw Error("LISTENER_INVALID"); + const ready=()=>s.queue.filter(e=>!e.delivered);if(!ready().length)await new Promise(resolve=>{const timer=setTimeout(done,2000);function done(){clearTimeout(timer);s.wake=null;resolve();}s.wake=done;}); + const result=ready();result.forEach(e=>{e.delivered=true;});return JSON.stringify(result);})()""" + + +ACK_SCRIPT = r"""((ids)=>{const s=window[KEY_VALUE];if(!s||!Array.isArray(ids))throw Error("LISTENER_INVALID");const wanted=new Set(ids);s.queue=s.queue.filter(e=>!wanted.has(e?.delivery_id));return true;})(IDS_VALUE)""" + + +RETRY_SCRIPT = r"""((ids)=>{const s=window[KEY_VALUE];if(!s||!Array.isArray(ids))throw Error("LISTENER_INVALID");const wanted=new Set(ids);s.queue.forEach(e=>{if(wanted.has(e?.delivery_id))e.delivered=false;});if(s.wake)s.wake();return true;})(IDS_VALUE)""" + + +DETAILS_SCRIPT = r"""(async ids=>{if(location.origin!=="https://www.douyin.com")throw Error("WRONG_ORIGIN");const chunks=window.webpackChunkdouyin_web;if(!chunks)throw Error("RUNTIME_MISSING");let req;chunks.push([["creatorhub-detail-"+Date.now()],{},r=>{req=r;}]);chunks.pop();const entry=Object.entries(req.m||{}).find(([,f])=>/getNoticeDetail\s*:/.test(String(f)));const sdk=entry&&req(entry[0]),client=window.axiosInstance;if(typeof sdk?.getNoticeDetail!=="function"||!client?.interceptors?.response)throw Error("SDK_NOT_READY");const params={id_list:JSON.stringify(ids.map(n=>({notice_id_str:n,type:0}))),is_mark_read:0};let raw,timer;const observer=client.interceptors.response.use(response=>{const config=response.config||{},url=new URL(config.url||"",location.origin),xhr=response.request;if(config.params?.id_list===params.id_list&&config.params.is_mark_read===0&&url.origin===location.origin&&url.pathname==="/aweme/v1/web/notice/detail/"&&xhr&&(!xhr.responseType||xhr.responseType==="text")&&typeof xhr.responseText==="string")raw={status:response.status,body:xhr.responseText};return response;});try{await Promise.race([sdk.getNoticeDetail(params),new Promise((_,reject)=>{timer=setTimeout(()=>reject(Error("TIMEOUT")),20000);})]);}finally{clearTimeout(timer);client.interceptors.response.eject(observer);}if(!raw)throw Error("RAW_RESPONSE_MISSING");return JSON.stringify(raw);})(IDS_VALUE)""" + + +FOLLOW_SCRIPT = r"""(async()=>{const p=PARAMS_VALUE;let sent=false;try{if(location.origin!=="https://www.douyin.com")throw Error("WRONG_ORIGIN");const get=async path=>{const r=await fetch(path,{credentials:"include",redirect:"error",signal:AbortSignal.timeout(15000)}),v=await r.json();if(!r.ok||v.status_code!==0)throw Error("READ_FAILED");return v;};const prefix="?device_platform=webapp&aid=6383&channel=channel_pc_web";const self=await get("/aweme/v1/web/user/profile/self/"+prefix);if(String(self.user?.uid)!==p.expected)throw Error("IDENTITY_MISMATCH");if(p.target===p.expected)throw Error("SELF_TARGET");const path="/aweme/v1/web/user/profile/other/"+prefix+"&user_id="+encodeURIComponent(p.target),profile=(await get(path)).user;if(String(profile?.uid)!==p.target)throw Error("TARGET_MISMATCH");if(![0,1,2].includes(profile.follow_status))throw Error("UNKNOWN_FOLLOW_STATE");if(profile.follow_status!==0||p.check)return {status:"succeeded",action:p.check?"checked":"already_following",follow_status:profile.follow_status,evidence:{target_uid:p.target,follow_status:String(profile.follow_status)}};sent=true;let result;try{const r=await fetch("/aweme/v1/web/commit/follow/user/"+prefix,{method:"POST",credentials:"include",redirect:"error",headers:{"Content-Type":"application/x-www-form-urlencoded;charset=UTF-8"},body:new URLSearchParams({user_id:p.target,type:"1"}),signal:AbortSignal.timeout(10000)});result=await r.json();if(!r.ok)throw Error("POST_UNCERTAIN");}catch(e){if(e.message==="POST_UNCERTAIN")throw e;throw Error("POST_UNCERTAIN");}if(result.status_code!==0){const e=Error("BUSINESS_REJECTED");e.definitive=true;throw e;}if(![1,2].includes(result.follow_status))throw Error("POST_UNCONFIRMED");const verify=(await get(path)).user;return {status:String(verify?.uid)===p.target&&[1,2].includes(verify.follow_status)?"succeeded":"unknown",action:"followed",evidence:{target_uid:p.target,follow_status:String(verify?.follow_status??"")}};}catch(e){const code=["WRONG_ORIGIN","IDENTITY_MISMATCH","SELF_TARGET","TARGET_MISMATCH","UNKNOWN_FOLLOW_STATE"].includes(e.message)?e.message:(e.message==="BUSINESS_REJECTED"?e.message:"REQUEST_FAILED");return {status:sent?(e.definitive?"failed":"unknown"):"failed",code};}})()""" + + +ACTION_SCRIPT = r"""(async()=>{const p=PARAMS_VALUE;let sent=false;const fail=(code,definitive=false)=>{const e=Error(code);e.definitive=definitive;throw e;};try{if(location.origin!=="https://www.douyin.com")fail("WRONG_ORIGIN");const qs="device_platform=webapp&aid=6383&channel=channel_pc_web";const get=async path=>{const r=await fetch(path,{credentials:"include",redirect:"error",signal:AbortSignal.timeout(15000)});let v;try{v=await r.json();}catch(_){fail("INVALID_RESPONSE");}if(!r.ok||v.status_code!==0)fail("READ_FAILED");return v;};const post=async(path,data)=>{const r=await fetch(path,{method:"POST",credentials:"include",redirect:"error",headers:{"Content-Type":"application/x-www-form-urlencoded;charset=UTF-8"},body:new URLSearchParams(data),signal:AbortSignal.timeout(10000)});let v;try{v=await r.json();}catch(_){fail(r.ok?"INVALID_RESPONSE":"POST_UNCERTAIN");}if(!r.ok)fail("POST_UNCERTAIN");if(v.status_code===undefined)fail("POST_UNCONFIRMED");if(v.status_code!==0)fail("BUSINESS_REJECTED",true);return v;};const self=await get("/aweme/v1/web/user/profile/self/?"+qs);if(String(self.user?.uid)!==p.expected)fail("IDENTITY_MISMATCH");if((p.action==="follow"||p.action==="dm")&&p.target===p.expected)fail("SELF_TARGET");const detail=async()=>get("/aweme/v1/web/aweme/detail/?"+qs+"&aweme_id="+encodeURIComponent(p.work));const findComment=async()=>{let cursor=0;for(let page=0;page<100;page++){const response=await get("/aweme/v1/web/comment/list/?"+qs+"&aweme_id="+encodeURIComponent(p.work)+"&cursor="+cursor+"&count=50");if(!Array.isArray(response.comments)||typeof response.has_more!=="boolean")fail("READ_FAILED");const list=response.comments;const comment=list.find(c=>String(c?.cid||c?.comment_id||"")===p.comment);if(comment){if(String(comment.aweme_id||comment.item_id||p.work)!==p.work||String(comment.user?.uid||comment.user_id||"")!==p.target)fail("TARGET_MISMATCH");return comment;}if(!response.has_more)break;const next=Number(response.cursor);if(!Number.isSafeInteger(next)||next<=cursor)fail("PAGINATION_INVALID");cursor=next;}fail("TARGET_NOT_FOUND");};if(p.action==="like_work"){const before=(await detail()).aweme_detail;if(String(before?.aweme_id)!==p.work)fail("TARGET_MISMATCH");if(Number(before.user_digged)===1)return {status:"succeeded",action:"already_liked",evidence:{work_id:p.work,user_digged:"1"}};sent=true;const result=await post("/aweme/v1/web/commit/item/digg/?"+qs,{aweme_id:p.work,type:"1",item_type:"0"});if(Number(result.is_digg)!==1)fail("POST_UNCONFIRMED");const after=(await detail()).aweme_detail;return {status:Number(after?.user_digged)===1?"succeeded":"unknown",action:"liked",evidence:{work_id:p.work,user_digged:String(after?.user_digged??"")}};}if(p.action==="like_comment"){const comment=await findComment();if(Number(comment.user_digged)===1)return {status:"succeeded",action:"already_liked",evidence:"comment.user_digged"};sent=true;await post("/aweme/v1/web/comment/digg?"+qs,{cid:p.comment,aweme_id:p.work,digg_type:"1",channel_id:"0",app_name:"aweme",item_type:"0",level:"1"});const after=await findComment();return {status:Number(after.user_digged)===1?"succeeded":"unknown",action:"liked_comment",evidence:{comment_id:p.comment,work_id:p.work,user_digged:String(after.user_digged??"")}};}if(p.action==="reply_comment"){await findComment();sent=true;const result=await post("/aweme/v1/web/comment/publish?"+qs,{app_name:"aweme",enter_from:"pc_web",previous_page:"video",reply_id:p.comment,reply_to_reply_id:"0",aweme_id:p.work,text:p.text,text_extra:"[]",comment_send_celltime:"0",comment_video_celltime:"0"});const posted=result.comment||result.comment_info||result.data?.comment;const postedID=String(posted?.cid||posted?.comment_id||"");const postedWork=String(posted?.aweme_id||posted?.item_id||"");const postedAuthor=String(posted?.user?.uid||posted?.user_id||"");const postedText=String(posted?.text??posted?.content??"");if(!posted||!postedID||postedWork!==p.work||postedAuthor!==p.expected||postedText!==p.text)fail("UNCONFIRMED");return {status:"succeeded",action:"replied",evidence:{comment_id:postedID,work_id:postedWork,author_uid:postedAuthor,text:postedText}};}if(p.action==="repost"){if(!p.text)fail("TEXT_REQUIRED");fail("REPOST_TEXT_UNSUPPORTED");}fail("ACTION_NOT_IMPLEMENTED");}catch(e){const code=String(e.message||"REQUEST_FAILED");return {status:sent?(e.definitive?"failed":"unknown"):"failed",code};}})()""" + + +IM_SCRIPT = r"""(async()=>{const p=PARAMS_VALUE;let sent=false;const fail=(code,definitive=false)=>{const e=Error(code);e.definitive=definitive;throw e;};try{if(location.origin!=="https://www.douyin.com")fail("WRONG_ORIGIN");const response=await fetch("/aweme/v1/web/user/profile/self/?device_platform=webapp&aid=6383",{credentials:"include",signal:AbortSignal.timeout(15000)});if(!response.ok)fail("LOGIN_CHECK_FAILED");const profile=await response.json();if(profile.status_code!==0||String(profile.user?.uid)!==EXPECTED_UID_VALUE)fail("IDENTITY_MISMATCH");if(String(profile.user.uid)===p.uid)fail("SELF_TARGET");let service;for(const key of Object.keys(window).filter(k=>k.startsWith("@pc-im/im:"))){const chunks=window[key];if(!Array.isArray(chunks))continue;let req;chunks.push([["creatorhub-im-"+Date.now()],{},r=>{req=r;}]);for(const [id,module] of Object.entries(req?.c||{})){if(!String(req.m[id]).includes("getOrCreatePrivateConversationByUid"))continue;for(const exported of Object.values(module.exports||{}))if(exported?.instance?.imSdkService)service=exported.instance.imSdkService;}}if(!service)fail("IM_SDK_NOT_READY");const sdk=service.imSdkManager.getImSdkInstance();if(!sdk)fail("IM_SDK_NOT_READY");const meta=c=>({id:String(c.id),short_id:String(c.shortId),uid:String(c.toParticipantUserId),type:c.type});const pack=m=>({server_id:String(m.serverId||""),client_id:m.clientId||null,sender:String(m.sender),type:m.type,content:m.content,created_at:m.createdAt,server_status:m.serverStatus});let conversation=sdk.getConversationList().find(c=>c.type===1&&String(c.toParticipantUserId)===p.uid);if(!conversation&&p.action==="send")conversation=await service.conversationManager.getOrCreatePrivateConversationByUid(p.uid);if(!conversation)fail("CONVERSATION_NOT_FOUND");if(conversation.type!==1||String(conversation.toParticipantUserId)!==p.uid)fail("TARGET_MISMATCH");if(p.action!=="send"||!p.confirm)return {status:"preview",action:"preview",sender_uid:String(profile.user.uid),uid:p.uid,text:p.text,conversation:meta(conversation)};const message=await sdk.createMessage({conversation,type:7,content:JSON.stringify({aweType:700,type:0,richTextInfos:[],text:p.text})});if(!message||typeof message.sendFunc!=="function")fail("MESSAGE_BUILD_FAILED");sent=true;const result=await Promise.race([sdk.sendMessage({message}),new Promise((_,reject)=>setTimeout(()=>reject(Error("MESSAGE_UNCONFIRMED")),10000))]);if(result?.success===false)fail("MESSAGE_REJECTED",true);if(result?.success!==true)fail("MESSAGE_UNCONFIRMED");const packed=pack(message);if(!packed.server_id&&!packed.client_id)fail("MESSAGE_UNCONFIRMED");return {status:"succeeded",action:"send",success:true,status_code:result?.statusCode??null,check_code:String(result?.checkCode??""),conversation:meta(conversation),message:packed,evidence:{conversation_id:String(conversation.id),message_server_id:packed.server_id,message_client_id:String(packed.client_id??"")}};}catch(e){const known=["WRONG_ORIGIN","LOGIN_CHECK_FAILED","IDENTITY_MISMATCH","SELF_TARGET","IM_SDK_NOT_READY","CONVERSATION_NOT_FOUND","TARGET_MISMATCH","MESSAGE_BUILD_FAILED","MESSAGE_REJECTED","MESSAGE_UNCONFIRMED"];return {status:sent?(e.definitive?"failed":"unknown"):"failed",code:known.includes(e.message)?e.message:"SDK_REQUEST_FAILED"};}})()""" diff --git a/cmd/docker_gateway/gateway.py b/cmd/docker_gateway/gateway.py new file mode 100644 index 0000000..8a20577 --- /dev/null +++ b/cmd/docker_gateway/gateway.py @@ -0,0 +1,1608 @@ +"""CreatorHub Python Docker/browser gateway.""" + +from __future__ import annotations + +import hmac +import json +import logging +import math +import os +import re +import signal +import socket +import threading +import time +from collections.abc import Mapping +from contextlib import suppress +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from typing import cast +from urllib.parse import parse_qs, quote, urlsplit + +from .docker_client import ( + BINDING_VERSION_LABEL, + DISPLAY_NAME_LABEL, + MANAGED_LABEL, + NAME_PREFIX, + NETWORK_EXIT_LABEL, + NETWORK_ID_LABEL, + PROXY_PORT_LABEL, + RUNTIME_ID_LABEL, + RUNTIME_ID_RE, + AliasReservationManager, + DockerClient, + DockerError, + GenerationConflict, + NetworkSetupError, + TenantNetworkGeneration, + UnmanagedContainer, +) +from .douyin import ( + ACCOUNT_KEY_RE, + ACTIONS, + COMMENTS_PATH, + IDENTITY_URL, + UID_RE, + WORKS_PATH, + DouyinBrowser, + DouyinError, + SubscriptionManager, +) +from .proxy import ProxyExit, ProxyRegistry + +LOG = logging.getLogger("creatorhub.gateway") +CONTROL_NETWORK = "creatorhub_control" +BROWSER_ENTRYPOINT = "/usr/local/bin/docker-entrypoint.sh" +BROWSER_USER = "1000:1000" +IMAGE_RE = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._:/@-]{0,300}$") +VOLUME_RE = re.compile(r"^[A-Za-z0-9][A-Za-z0-9_.-]{0,63}$") +CONTAINER_ID_RE = re.compile(r"^[a-f0-9]{64}$") +RUNTIME_CLEANUP_SENTINEL = "runtime-not-found" +EXIT_ID_RE = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._/-]{0,127}$") +DOUYIN_ACCOUNT_KEY_RE = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._:@-]{0,127}$") +DOUYIN_ORIGIN = "https://www.douyin.com" +DOUYIN_IDENTITY_PATH = "/aweme/v1/web/user/profile/self/" +DOUYIN_IDENTITY_URL = IDENTITY_URL +DOUYIN_WORKS_PATH = WORKS_PATH +DOUYIN_COMMENTS_PATH = COMMENTS_PATH + + +def _noop() -> None: + return None + + +class RequestError(RuntimeError): + def __init__(self, message: str, status: int = 502, network_id: str = "") -> None: + super().__init__(message) + self.status = status + self.network_id = network_id + + +class Gateway: + def __init__( + self, + docker: DockerClient, + network: str, + token: str, + self_name: str, + browser: DouyinBrowser | None = None, + ) -> None: + self.docker = docker + self.network = network + self.token = token + self.self_name = self_name + self.browser = browser or DouyinBrowser(self._browser_endpoint) + self.proxies = ProxyRegistry() + self.reservations = AliasReservationManager(docker, self_name) + self.subscriptions = SubscriptionManager(self.browser) + self._action_ownership_lock = threading.Lock() + self._uncertain_actions: dict[str, float] = {} + + def _browser_endpoint(self, alias: str) -> str: + container_id, labels = self.docker.managed_container(alias) + network_id = labels.get(NETWORK_ID_LABEL, "") + if not isinstance(network_id, str) or not network_id: + raise GenerationConflict("browser container has no isolated network") + address = self.docker.container_network_address(container_id, network_id) + return f"http://{address}:9222" + + def list_browsers(self) -> list[dict]: + filters = quote( + json.dumps({"label": [f"{MANAGED_LABEL}=true"]}, separators=(",", ":")), + safe="", + ) + response = self.docker.request( + "GET", f"/containers/json?all=1&filters={filters}" + ) + if response.status != 200: + raise RequestError( + f"Docker returned HTTP {response.status}", response.status + ) + try: + containers = json.loads(response.body) + except json.JSONDecodeError as exc: + raise RequestError("Docker container list is invalid") from exc + if not isinstance(containers, list): + raise RequestError("Docker container list is invalid") + result = [] + for container in containers: + if not isinstance(container, dict): + raise RequestError("Docker container list is invalid") + labels = container.get("Labels") or {} + if not isinstance(labels, dict): + raise RequestError("Docker container labels are invalid") + alias = labels.get(RUNTIME_ID_LABEL, "") + if not isinstance(alias, str) or not RUNTIME_ID_RE.fullmatch(alias): + continue + try: + binding = int(labels.get(BINDING_VERSION_LABEL, "0")) + proxy_port = int(labels.get(PROXY_PORT_LABEL, "0")) + except (TypeError, ValueError): + binding = proxy_port = 0 + network_exit_id = labels.get(NETWORK_EXIT_LABEL, "") + network_id = labels.get(NETWORK_ID_LABEL, "") + container_id = container.get("Id", "") + if ( + not isinstance(network_exit_id, str) + or not isinstance(network_id, str) + or not isinstance(container_id, str) + ): + raise RequestError("Docker container metadata is invalid") + direct = not network_exit_id + endpoint = f"http://{NAME_PREFIX}{alias}:9222" + network_error = "" + if network_id and container.get("State") == "running": + try: + address = self.docker.container_network_address( + container_id, network_id + ) + except (DockerError, FileNotFoundError, GenerationConflict) as exc: + network_error = str(exc) + endpoint = "" + LOG.warning( + "browser network address unavailable", + extra={ + "alias": alias, + "network_id": network_id, + "error": network_error, + }, + ) + else: + endpoint = f"http://{address}:9222" + result.append( + { + "id": container.get("Id", ""), + "alias": alias, + "name": labels.get(DISPLAY_NAME_LABEL) or alias, + "state": container.get("State", ""), + "status": container.get("Status", ""), + "endpoint": endpoint, + "binding_version": binding, + "network_exit_id": network_exit_id, + "network_id": network_id, + "proxy_ready": (not network_error) + and ( + direct + or self.proxies.ready( + alias, proxy_port, container.get("Id", ""), network_id + ) + ), + **({"error": network_error} if network_error else {}), + } + ) + return result + + def create(self, input: dict) -> dict: + validate_create(input) + self.docker.pull_if_missing(input["image"]) + alias = input["alias"] + release = self.reservations.acquire(alias) + try: + # The reservation is the cross-process alias lock; checking before it + # was acquired leaves a create/create race window. + try: + self.docker.managed_container(alias) + except FileNotFoundError: + pass + else: + raise RequestError("browser alias is already in use", 409) + except Exception: + release() + raise + network_generation = TenantNetworkGeneration() + undo_proxy = _noop + keep_network = bool(input.get("stopped")) + keep_proxy = False + keep_container = False + container_attempted = False + created_id = "" + try: + direct = not input["network_exit_id"] + network = "none" + proxy_url = "" + if not input.get("stopped"): + network_generation, bind_host = self.docker.ensure_tenant_network( + self.network, alias, self.self_name, input["binding_version"] + ) + network = network_generation.id + if not direct: + proxy_url, undo_proxy = self.proxies.configure( + alias, + input["binding_version"], + bind_host, + 0, + input["network_exit"], + network_generation.id, + ) + command = list(input["cmd"]) + if not input.get("stopped") and not direct: + command = command[:-1] + [ + f"--proxy-server={proxy_url}", + "--disable-non-proxied-udp", + command[-1], + ] + pids_limit = 512 + payload = { + "Image": input["image"], + "User": BROWSER_USER, + "Entrypoint": [BROWSER_ENTRYPOINT], + "Cmd": command, + "Env": ["REMOTE_DEBUGGING_PORT=9222"], + "Labels": { + MANAGED_LABEL: "true", + RUNTIME_ID_LABEL: alias, + DISPLAY_NAME_LABEL: input["name"], + BINDING_VERSION_LABEL: str(input["binding_version"]), + NETWORK_EXIT_LABEL: input["network_exit_id"], + NETWORK_ID_LABEL: network_generation.id, + PROXY_PORT_LABEL: str(proxy_port(proxy_url)), + }, + "ExposedPorts": {"9222/tcp": {}}, + "HostConfig": { + "NetworkMode": network, + "ReadonlyRootfs": True, + "CapDrop": ["ALL"], + "SecurityOpt": ["no-new-privileges"], + "PidsLimit": pids_limit, + "Memory": 1 << 30, + "NanoCpus": 2_000_000_000, + "Tmpfs": browser_tmpfs(), + "Mounts": [ + {"Type": "volume", "Source": input["volume"], "Target": "/data"} + ], + }, + } + container_attempted = True + response = self.docker.request( + "POST", + "/containers/create?" + "name=" + quote(NAME_PREFIX + alias, safe=""), + payload, + ) + if response.status == 409: + raise RequestError( + "browser alias is already in use", 409, network_generation.id + ) + if response.status != 201: + raise RequestError( + f"Docker container creation failed with HTTP {response.status}", + response.status, + network_generation.id, + ) + try: + created_id = json.loads(response.body)["Id"] + except (KeyError, TypeError, json.JSONDecodeError) as exc: + raise RequestError( + "Docker returned an invalid container id", + 502, + network_generation.id, + ) from exc + if not isinstance(created_id, str) or not created_id: + raise RequestError( + "Docker returned an invalid container id", + 502, + network_generation.id, + ) + if not input.get("stopped"): + if not direct and not self.proxies.bind( + alias, + input["binding_version"], + proxy_url, + created_id, + network_generation.id, + ): + self.docker.expect( + "DELETE", + f"/containers/{quote(created_id, safe='')}?force=1&v=0", + ) + raise RequestError( + "proxy generation changed", 409, network_generation.id + ) + try: + self.docker.expect( + "POST", + f"/containers/{quote(created_id, safe='')}/start", + allowed=(204, 304), + ) + except Exception as exc: + try: + self.docker.expect( + "DELETE", + f"/containers/{quote(created_id, safe='')}?force=1&v=0", + ) + except Exception: + LOG.exception( + "failed to remove container after start failure", + extra={"container_id": created_id}, + ) + raise RequestError( + "container did not start and was removed", + 502, + network_generation.id, + ) from exc + keep_proxy = not direct + keep_network = True + keep_container = True + self._release_action_ownership(alias) + return { + "id": created_id, + "alias": alias, + "network_id": network_generation.id, + } + except NetworkSetupError as exc: + network_generation = exc.generation + raise RequestError(str(exc), 502, network_generation.id) from exc + finally: + if not keep_proxy: + undo_proxy() + if container_attempted and not keep_container: + self._reconcile_created_container( + alias, + created_id, + input["binding_version"], + network_generation.id, + ) + if not keep_network and network_generation.id: + self._cleanup_network( + alias, input["binding_version"], "", network_generation + ) + try: + release() + except Exception: + LOG.exception( + "failed to release browser alias reservation", + extra={"alias": alias}, + ) + + def _reconcile_created_container( + self, alias: str, created_id: str, binding_version: int, network_id: str + ) -> None: + # A response without a container ID is not attributable to this + # request. Never delete an alias-matching container created by another + # request; an unknown outcome is logged and reconciled by the control + # plane's generation-aware cleanup instead. + if not created_id: + LOG.error( + "container creation outcome has no attributable container id", + extra={"alias": alias, "binding_version": binding_version}, + ) + return + try: + observed_id, labels = self.docker.managed_container(alias) + if created_id and observed_id != created_id: + LOG.error( + "container creation outcome has a replacement generation", + extra={ + "alias": alias, + "created_id": created_id, + "observed_id": observed_id, + }, + ) + return + if ( + labels.get(RUNTIME_ID_LABEL) != alias + or labels.get(BINDING_VERSION_LABEL) != str(binding_version) + or labels.get(NETWORK_ID_LABEL, "") != network_id + ): + LOG.error( + "container creation outcome is not safely attributable", + extra={"alias": alias, "observed_id": observed_id}, + ) + return + self.docker.expect( + "DELETE", + f"/containers/{quote(observed_id, safe='')}?force=1&v=0", + ) + except FileNotFoundError: + return + except (DockerError, OSError, TypeError, ValueError, KeyError): + LOG.exception( + "failed to reconcile container creation outcome", + extra={"alias": alias, "created_id": created_id}, + ) + + def change_state(self, alias: str, action: str, input: dict) -> None: + generation = decode_generation( + input, require_runtime=True, require_network=action == "start" + ) + with self._alias_lock(alias): + container_id, exists = self._require_generation(alias, generation) + if not exists: + raise RequestError("browser not found", 404) + path = f"/containers/{quote(container_id, safe='')}/{'start' if action == 'start' else 'stop?t=10'}" + try: + self.docker.expect("POST", path, allowed=(204, 304)) + except FileNotFoundError as exc: + raise RequestError("browser not found", 404) from exc + except Exception as exc: + raise RequestError("Docker state change failed") from exc + + def remove(self, alias: str, input: dict) -> None: + generation = decode_generation( + input, require_runtime=False, require_network=False + ) + with self._alias_lock(alias): + try: + container_id, labels = self.docker.managed_container(alias) + exists = True + except FileNotFoundError: + container_id, labels, exists = "", {}, False + if exists and ( + not generation["runtime_id"] + or container_id != generation["runtime_id"] + or labels.get(BINDING_VERSION_LABEL) + != str(generation["binding_version"]) + or labels.get(RUNTIME_ID_LABEL) != alias + or labels.get(NETWORK_ID_LABEL) != generation["network_id"] + ): + raise RequestError("container generation does not match request", 409) + try: + network_generation, _, network_exists = ( + self.docker.inspect_tenant_network( + self.network, + alias, + generation["binding_version"], + generation["runtime_id"], + self.self_name, + generation["network_id"], + ) + ) + if ( + network_exists + and generation["runtime_id"] == RUNTIME_CLEANUP_SENTINEL + and network_generation.runtime_attached + ): + raise RequestError( + "container generation is required while the network is attached", + 409, + ) + if network_exists: + self._remove_network( + alias, + generation["binding_version"], + generation["runtime_id"], + network_generation, + ) + elif exists and generation["network_id"]: + raise RequestError("container network generation is missing", 409) + except (GenerationConflict, RequestError): + raise + except (DockerError, OSError, TypeError, ValueError, KeyError) as exc: + raise RequestError( + "runtime_cleanup_pending", 202, generation["network_id"] + ) from exc + if not self.proxies.remove( + alias, + generation["binding_version"], + generation["runtime_id"], + generation["network_id"], + ): + raise RequestError("proxy generation does not match request", 409) + if exists: + try: + self.docker.expect( + "DELETE", + f"/containers/{quote(container_id, safe='')}?force=1&v=0", + ) + except FileNotFoundError: + pass + except Exception as exc: + raise RequestError("Docker container removal failed") from exc + self._release_action_ownership(alias) + + def restore_proxy(self, alias: str, input: dict) -> None: + validate_proxy_restore(input, alias) + with self._alias_lock(alias): + container_id, labels = self.docker.managed_container(alias) + expected = { + "binding_version": input["binding_version"], + "runtime_id": input["runtime_id"], + "network_id": input["network_id"], + } + if ( + labels.get(RUNTIME_ID_LABEL) != alias + or labels.get(BINDING_VERSION_LABEL) != str(expected["binding_version"]) + or labels.get(NETWORK_ID_LABEL) != expected["network_id"] + or labels.get(NETWORK_EXIT_LABEL) != input["network_exit_id"] + ): + raise RequestError( + "container binding does not match recovery request", 409 + ) + try: + port = int(labels.get(PROXY_PORT_LABEL, "0")) + except (TypeError, ValueError) as exc: + raise RequestError( + "container binding has an invalid proxy port", 409 + ) from exc + if port < 1 and input["network_exit_id"]: + raise RequestError("container binding has no proxy port", 409) + generation = TenantNetworkGeneration(id=input["network_id"]) + configured = False + undo = _noop + try: + generation, bind_host = self.docker.ensure_tenant_network( + self.network, + alias, + self.self_name, + input["binding_version"], + input["runtime_id"], + input["network_id"], + True, + ) + self._require_proxy_network_generation( + alias, input, generation, bind_host + ) + if not input["network_exit_id"]: + return + proxy_url, undo = self.proxies.configure( + alias, + input["binding_version"], + bind_host, + port, + input["network_exit"], + input["network_id"], + ) + configured = True + self._require_proxy_network_generation( + alias, input, generation, bind_host + ) + if not self.proxies.bind( + alias, + input["binding_version"], + proxy_url, + container_id, + input["network_id"], + ): + raise RequestError("proxy generation changed", 409) + self._require_proxy_network_generation( + alias, input, generation, bind_host + ) + except Exception: + if configured: + undo() + else: + self.proxies.remove( + alias, + input["binding_version"], + input["runtime_id"], + input["network_id"], + ) + self._cleanup_network( + alias, input["binding_version"], input["runtime_id"], generation + ) + raise + + def set_douyin_cookies(self, alias: str, input: dict) -> None: + cookies_value = input.get("cookies") + if not valid_douyin_generation(input) or not valid_cookies(cookies_value): + raise RequestError("invalid restricted browser request", 400) + cookies = cast(list[dict], cookies_value) + with self._alias_lock(alias): + self._require_douyin_generation(alias, input) + try: + self.browser.set_cookies(alias, cookies) + self._require_douyin_generation(alias, input) + except DouyinError as exc: + raise RequestError("restricted browser operation failed") from exc + + def get_douyin(self, alias: str, input: dict) -> dict: + if not valid_douyin_generation(input) or not valid_douyin_url( + input.get("url", "") + ): + raise RequestError("invalid restricted browser request", 400) + with self._alias_lock(alias): + self._require_douyin_generation(alias, input) + try: + response = self.browser.get(alias, input["url"]) + self._require_douyin_generation(alias, input) + except DouyinError as exc: + LOG.warning( + "Douyin GET failed alias=%s reason=%s", + alias, + str(exc), + ) + raise RequestError("restricted browser operation failed") from exc + return { + "status": response.status, + "body": response.body, + "challenge": response.challenge, + } + + def douyin_identity(self, alias: str, input: dict) -> dict: + expected_account_key = input.get("expected_account_key", "") + if ( + not valid_douyin_generation(input) + or not isinstance(expected_account_key, str) + or not ACCOUNT_KEY_RE.fullmatch(expected_account_key) + ): + raise RequestError("invalid Douyin identity request", 400) + with self._alias_lock(alias): + self._require_douyin_generation(alias, input) + try: + identity = self.browser.identity(alias) + except DouyinError as exc: + LOG.warning( + "Douyin identity verification failed alias=%s reason=%s", + alias, + str(exc), + ) + raise RequestError( + "Douyin login identity could not be verified" + ) from exc + if expected_account_key not in { + identity["uid"], + identity["sec_uid"], + identity["unique_id"], + }: + raise RequestError( + "Douyin identity does not match the expected account", 409 + ) + return identity + + def douyin_action(self, alias: str, input: dict) -> dict: + expected_uid = input.get("expected_uid", "") + action = input.get("action", "") + target_uid = input.get("target_uid", "") + target_comment_id = input.get("target_comment_id", "") + target_work_id = input.get("target_work_id", "") + text = input.get("text", "") + confirm = input.get("confirm", False) + if ( + not valid_douyin_generation(input) + or not isinstance(expected_uid, str) + or not isinstance(action, str) + or not isinstance(target_uid, str) + or not isinstance(target_comment_id, str) + or not isinstance(target_work_id, str) + or not isinstance(text, str) + or type(confirm) is not bool + or not UID_RE.fullmatch(expected_uid) + or action not in ACTIONS + ): + raise RequestError("invalid Douyin action request", 400) + with self._alias_lock(alias): + self._require_douyin_generation(alias, input) + self._claim_action(alias) + try: + result = self.browser.action( + alias, + expected_uid, + action, + target_uid, + target_comment_id, + target_work_id, + text, + confirm, + ) + self._require_douyin_generation(alias, input) + except DouyinError as exc: + self._handle_douyin_action_error(alias, action, exc) + raise RequestError("Douyin action failed") from exc + except Exception: + self._release_action_ownership(alias) + raise + else: + self._release_action_ownership(alias) + return result + + def _handle_douyin_action_error( + self, alias: str, action: str, error: DouyinError + ) -> None: + if getattr(error, "uncertain", False) or "timed out" in str(error).lower(): + self._retain_action_ownership(alias) + else: + self._release_action_ownership(alias) + LOG.warning( + "Douyin action failed alias=%s action=%s reason=%s", + alias, + action, + str(error), + ) + + def _claim_action(self, alias: str) -> None: + now = time.monotonic() + with self._action_ownership_lock: + until = self._uncertain_actions.get(alias, 0.0) + if until > now: + raise RequestError("previous Douyin action outcome is uncertain", 409) + self._uncertain_actions.pop(alias, None) + self._uncertain_actions[alias] = 0.0 + + def _release_action_ownership(self, alias: str) -> None: + with self._action_ownership_lock: + self._uncertain_actions.pop(alias, None) + + def _retain_action_ownership(self, alias: str) -> None: + with self._action_ownership_lock: + # A timed-out page script may still finish its network write. Keep the + # alias blocked until the browser generation is removed or replaced. + self._uncertain_actions[alias] = math.inf + + def start_douyin_events(self, alias: str, input: dict) -> dict: + expected_uid = input.get("expected_uid", "") + if ( + not valid_douyin_generation(input) + or not isinstance(expected_uid, str) + or not UID_RE.fullmatch(expected_uid) + ): + raise RequestError("invalid Douyin event request", 400) + with self._alias_lock(alias): + self._require_douyin_generation(alias, input) + try: + self.browser.identity(alias, expected_uid) + return self.subscriptions.start(alias, expected_uid) + except DouyinError as exc: + raise RequestError("Douyin event listener could not start") from exc + + def poll_douyin_events(self, alias: str, input: dict, query: dict) -> list[dict]: + if not valid_douyin_generation(input): + raise RequestError("invalid Douyin event request", 400) + try: + limit = int(query.get("limit", ["50"])[0]) + wait = float(query.get("wait", ["25"])[0]) + except (IndexError, TypeError, ValueError) as exc: + raise RequestError("invalid event poll options", 400) from exc + if not 1 <= limit <= 100 or not 0 <= wait <= 30: + raise RequestError("invalid event poll options", 400) + with self._alias_lock(alias): + self._require_douyin_generation(alias, input) + try: + acknowledgements = query.get("ack", []) + delivery_ids = [ + value for raw in acknowledgements for value in raw.split(",") + ] + self.subscriptions.ack(alias, delivery_ids) + return self.subscriptions.poll(alias, limit, wait) + except DouyinError as exc: + raise RequestError("Douyin event listener is unavailable") from exc + + def stop_douyin_events(self, alias: str, input: dict) -> None: + if not valid_douyin_generation(input): + raise RequestError("invalid Douyin event request", 400) + with self._alias_lock(alias): + self._require_douyin_generation(alias, input) + self.subscriptions.stop(alias) + + def _require_generation(self, alias: str, generation: dict) -> tuple[str, bool]: + try: + container_id, labels = self.docker.managed_container(alias) + except FileNotFoundError: + return "", False + if ( + container_id != generation["runtime_id"] + or labels.get(RUNTIME_ID_LABEL) != alias + or labels.get(BINDING_VERSION_LABEL) != str(generation["binding_version"]) + or labels.get(NETWORK_ID_LABEL) != generation["network_id"] + ): + raise RequestError("container generation does not match request", 409) + return container_id, True + + def _require_douyin_generation(self, alias: str, input: dict) -> None: + container_id, labels = self.docker.managed_container(alias) + if ( + container_id != input["runtime_id"] + or labels.get(RUNTIME_ID_LABEL) != alias + or labels.get(BINDING_VERSION_LABEL) != str(input["binding_version"]) + or labels.get(NETWORK_ID_LABEL) != input["network_id"] + or labels.get(NETWORK_EXIT_LABEL, "") != input.get("network_exit_id", "") + ): + raise RequestError("container generation does not match request", 409) + generation, _, exists = self.docker.inspect_tenant_network( + self.network, + alias, + input["binding_version"], + input["runtime_id"], + self.self_name, + input["network_id"], + ) + if ( + not exists + or not generation.runtime_attached + or not generation.self_member + or not generation.gateway_members + ): + raise RequestError( + "container network generation does not match request", 409 + ) + _, _, networks = self.docker.managed_container_state(alias) + if networks != {generation.name: generation.id}: + raise RequestError( + "container network generation does not match request", 409 + ) + + def _require_proxy_network_generation( + self, alias: str, input: dict, expected: TenantNetworkGeneration, bind_host: str + ) -> None: + container_id, labels = self.docker.managed_container(alias) + if ( + container_id != input["runtime_id"] + or labels.get(NETWORK_ID_LABEL) != expected.id + ): + raise RequestError("network generation changed", 409) + current, addresses, exists = self.docker.inspect_tenant_network( + self.network, + alias, + input["binding_version"], + input["runtime_id"], + self.self_name, + expected.id, + ) + if ( + not exists + or not same_network_members(current, expected) + or addresses.get(current.self_member, "").split("/", 1)[0] != bind_host + ): + raise RequestError("network generation changed", 409) + + def _cleanup_network( + self, + alias: str, + binding_version: int, + runtime_id: str, + generation: TenantNetworkGeneration, + ) -> None: + try: + current, _, exists = self.docker.inspect_tenant_network( + self.network, + alias, + binding_version, + runtime_id, + self.self_name, + generation.id, + ) + if not exists: + return + if generation.created: + self._remove_network(alias, binding_version, runtime_id, current, True) + else: + if generation.connected_runtime and runtime_id: + self.docker.disconnect_member( + self.network, + alias, + binding_version, + runtime_id, + current, + runtime_id, + self.self_name, + missing_ok=True, + ) + if generation.connected_self: + member = current.self_member or generation.self_member + if member: + self.docker.disconnect_member( + self.network, + alias, + binding_version, + runtime_id, + current, + member, + self.self_name, + missing_ok=True, + ) + except (DockerError, OSError, TypeError, ValueError, KeyError): + LOG.exception( + "failed to clean up isolated browser network", + extra={"alias": alias, "network_id": generation.id}, + ) + + def _remove_network( + self, + alias: str, + binding_version: int, + runtime_id: str, + generation: TenantNetworkGeneration, + missing_ok: bool = False, + ) -> None: + current = generation + if current.runtime_attached or current.connected_runtime: + current = self.docker.disconnect_member( + self.network, + alias, + binding_version, + runtime_id, + current, + runtime_id, + self.self_name, + missing_ok=missing_ok, + ) + for member in list(current.gateway_members): + current = self.docker.disconnect_member( + self.network, + alias, + binding_version, + runtime_id, + current, + member, + self.self_name, + missing_ok=missing_ok, + ) + if current.self_member: + current = self.docker.disconnect_member( + self.network, + alias, + binding_version, + runtime_id, + current, + current.self_member, + self.self_name, + missing_ok=missing_ok, + ) + self.docker.delete_tenant_network( + self.network, + alias, + binding_version, + runtime_id, + current, + self.self_name, + missing_ok=missing_ok, + ) + + def _alias_lock(self, alias: str): + return _AliasLock(self.reservations, alias) + + +class _AliasLock: + def __init__(self, reservations: AliasReservationManager, alias: str) -> None: + self.reservations = reservations + self.alias = alias + self.release = None + + def __enter__(self): + self.release = self.reservations.acquire(self.alias) + return self + + def __exit__(self, exc_type, exc_value, traceback) -> None: + if self.release: + self.release() + + +class GatewayHTTPServer(ThreadingHTTPServer): + daemon_threads = True + allow_reuse_address = True + + gateway: Gateway + + def __init__(self, address, gateway: Gateway): + super().__init__(address, GatewayHandler) + self.gateway = gateway + self._connections: set[socket.socket] = set() + self._connections_lock = threading.Lock() + self._connections_changed = threading.Condition(self._connections_lock) + + def process_request(self, request, client_address): + with self._connections_changed: + self._connections.add(cast(socket.socket, request)) + try: + super().process_request(request, client_address) + except Exception: + with self._connections_changed: + self._connections.discard(cast(socket.socket, request)) + self._connections_changed.notify_all() + raise + + def process_request_thread(self, request, client_address): + try: + super().process_request_thread(request, client_address) + finally: + with self._connections_changed: + self._connections.discard(cast(socket.socket, request)) + self._connections_changed.notify_all() + + def wait_for_requests(self, timeout: float) -> None: + deadline = time.monotonic() + timeout + with self._connections_changed: + while self._connections: + remaining = deadline - time.monotonic() + if remaining <= 0: + break + self._connections_changed.wait(remaining) + if self._connections: + connections = list(self._connections) + else: + connections = [] + for connection in connections: + with suppress(OSError): + connection.shutdown(socket.SHUT_RDWR) + connection.close() + + +class GatewayHandler(BaseHTTPRequestHandler): + protocol_version = "HTTP/1.1" + + def setup(self) -> None: + super().setup() + self.request.settimeout(5.0) + + def do_GET(self) -> None: + self._dispatch("GET") + + def do_POST(self) -> None: + self._dispatch("POST") + + def do_DELETE(self) -> None: + self._dispatch("DELETE") + + def log_message(self, format: str, *args) -> None: + LOG.info("http_request", extra={"request": format % args}) + + def _dispatch(self, method: str) -> None: + parsed = urlsplit(self.path) + if parsed.path == "/healthz": + self._respond(204, b"") + return + if not parsed.path.startswith("/v1/"): + self._respond(404, json_bytes({"error": "not found"})) + return + if not self._authorized(): + self._respond(401, json_bytes({"error": "gateway token rejected"})) + return + try: + needs_body = method in {"POST", "DELETE"} or ( + method == "GET" and parsed.path != "/v1/browsers" + ) + body = self._body() if needs_body else {} + result = self._route(method, parsed.path, parse_qs(parsed.query), body) + if result is None: + self._respond(204, b"") + elif isinstance(result, tuple): + status, value = result + if type(status) is not int: + raise RuntimeError("gateway route returned an invalid status") + self._respond(status, json_bytes(value)) + else: + self._respond(200, json_bytes(result)) + except (RuntimeError, OSError, ValueError, TypeError, KeyError) as exc: + self._handle_exception(parsed.path, exc) + + def _handle_exception(self, path: str, exc: Exception) -> None: + if isinstance(exc, RequestError): + payload = {"error": str(exc)} + if exc.network_id: + payload["network_id"] = exc.network_id + self._respond(exc.status, json_bytes(payload)) + elif isinstance(exc, FileNotFoundError): + self._respond(404, json_bytes({"error": str(exc)})) + elif isinstance(exc, (GenerationConflict, UnmanagedContainer)): + self._respond(409, json_bytes({"error": str(exc)})) + elif isinstance(exc, DockerError): + status = ( + exc.status + if exc.status is not None and 400 <= exc.status < 500 + else 502 + ) + self._respond(status, json_bytes({"error": str(exc)})) + elif isinstance(exc, ValueError): + self._respond(400, json_bytes({"error": str(exc)})) + else: + LOG.exception("gateway request failed", extra={"path": path}) + self._respond(500, json_bytes({"error": "gateway operation failed"})) + + def _route(self, method: str, path: str, query: dict, body: dict): + gateway = self.server_as_gateway().gateway + if method == "GET" and path == "/v1/browsers": + return gateway.list_browsers() + if method == "POST" and path == "/v1/browsers": + return 201, gateway.create(body) + match = re.fullmatch(r"/v1/browsers/([a-z0-9][a-z0-9-]{0,31})", path) + if match and method == "DELETE": + gateway.remove( + match.group(1), + decode_generation(body, require_runtime=False, require_network=False), + ) + return None + match = re.fullmatch( + r"/v1/browsers/([a-z0-9][a-z0-9-]{0,31})/(start|stop|proxy)", path + ) + if match: + alias, action = match.groups() + if method == "POST" and action in {"start", "stop"}: + gateway.change_state(alias, action, body) + return None + if method == "POST" and action == "proxy": + gateway.restore_proxy(alias, body) + return None + match = re.fullmatch( + r"/v1/browsers/([a-z0-9][a-z0-9-]{0,31})/douyin/(cookies|get|identity|action|events)", + path, + ) + if match: + alias, action = match.groups() + if action == "cookies" and method == "POST": + gateway.set_douyin_cookies(alias, body) + return None + if action == "get" and method == "POST": + return gateway.get_douyin(alias, body) + if action == "identity" and method == "POST": + return gateway.douyin_identity(alias, body) + if action == "action" and method == "POST": + return gateway.douyin_action(alias, body) + if action == "events": + if method == "POST": + return gateway.start_douyin_events(alias, body) + if method == "GET": + return gateway.poll_douyin_events(alias, body, query) + if method == "DELETE": + gateway.stop_douyin_events(alias, body) + return None + raise RequestError("not found", 404) + + def server_as_gateway(self) -> GatewayHTTPServer: + if not isinstance(self.server, GatewayHTTPServer): + raise TypeError("gateway HTTP server type is invalid") + return self.server + + def _authorized(self) -> bool: + supplied = self.headers.get("Authorization", "") + return hmac.compare_digest( + supplied, "Bearer " + self.server_as_gateway().gateway.token + ) + + def _body(self) -> dict: + length_text = self.headers.get("Content-Length") + if length_text is None: + raise RequestError("request body is required", 400) + try: + length = int(length_text) + except ValueError as exc: + raise RequestError("invalid request body length", 400) from exc + if length < 0 or length > 1 << 20: + raise RequestError("request body is too large", 400) + raw = self.rfile.read(length) + try: + value = json.loads(raw) + except json.JSONDecodeError as exc: + raise RequestError("request body must be one JSON object", 400) from exc + if not isinstance(value, dict): + raise RequestError("request body must be one JSON object", 400) + return value + + def _respond(self, status: int, body: bytes) -> None: + self.send_response(status) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + if body: + self.wfile.write(body) + + +def json_bytes(value: object) -> bytes: + return json.dumps(value, ensure_ascii=False, separators=(",", ":")).encode() + + +def validate_create(input: dict) -> None: + allowed = { + "alias", + "name", + "image", + "cmd", + "volume", + "binding_version", + "network_exit_id", + "network_exit", + "stopped", + } + if set(input) - allowed: + raise RequestError("body contains unknown fields", 400) + alias = input.get("alias", "") + name = input.get("name", "") + image = input.get("image", "") + volume = input.get("volume", "") + command = input.get("cmd") + binding = input.get("binding_version") + exit_id = input.get("network_exit_id", "") + stopped = input.get("stopped", False) + if not isinstance(alias, str) or not RUNTIME_ID_RE.fullmatch(alias): + raise RequestError("alias must match [a-z0-9][a-z0-9-]{0,31}", 400) + if not isinstance(name, str) or not 1 <= len(name) <= 64 or has_control(name): + raise RequestError("name must be 1..64 visible characters", 400) + if not isinstance(image, str) or not IMAGE_RE.fullmatch(image): + raise RequestError("image must be a valid image reference", 400) + if not isinstance(volume, str) or not VOLUME_RE.fullmatch(volume): + raise RequestError("volume must be a valid volume name", 400) + if type(binding) is not int or binding < 1: + raise RequestError( + "binding_version and network_exit_id must identify the current binding", 400 + ) + if type(stopped) is not bool: + raise RequestError("stopped must be boolean", 400) + if not isinstance(exit_id, str): + raise RequestError("network_exit_id must be a string", 400) + input["network_exit_id"] = exit_id + input.setdefault("network_exit", {}) + exit_value = parse_proxy_exit(input.get("network_exit", {})) + direct = not exit_id and exit_value == ProxyExit("", "", 0) + if stopped and not direct: + raise RequestError("stopped browsers must use direct networking", 400) + if bool(exit_id) != (exit_value != ProxyExit("", "", 0)): + raise RequestError( + "binding_version and network_exit_id must identify the current binding", 400 + ) + if exit_id and not EXIT_ID_RE.fullmatch(exit_id): + raise RequestError("network_exit_id is invalid", 400) + if ( + not isinstance(command, list) + or not 1 <= len(command) <= 64 + or command[-1] != "about:blank" + ): + raise RequestError("cmd must contain 1..64 arguments", 400) + total = 0 + for item in command: + if ( + not isinstance(item, str) + or not item + or has_control(item) + or item.startswith("--proxy-server") + or item == "--disable-non-proxied-udp" + ): + raise RequestError("cmd arguments are invalid", 400) + total += len(item) + if total > 4096: + raise RequestError("cmd arguments exceed 4096 characters", 400) + if not stopped and not direct: + validate_proxy_exit(exit_value) + input["network_exit"] = exit_value + + +def validate_proxy_exit(exit: ProxyExit) -> None: + if ( + exit.protocol not in {"http", "https", "socks4", "socks5"} + or not exit.host + or len(exit.host) > 253 + or any(char in exit.host for char in "@/[]?# \t\r\n") + or not 1 <= exit.port <= 65535 + or (not exit.username and exit.credential) + or len(exit.username) > 255 + or len(exit.credential) > 255 + or has_control(exit.username) + or has_control(exit.credential) + ): + raise RequestError("network_exit must contain a valid proxy endpoint", 400) + + +def parse_proxy_exit(value: object) -> ProxyExit: + if not isinstance(value, dict): + raise RequestError("network_exit must be an object", 400) + allowed = {"protocol", "host", "port", "username", "password"} + if set(value) - allowed: + raise RequestError("network_exit contains unknown fields", 400) + try: + exit = ProxyExit( + value.get("protocol", ""), + value.get("host", ""), + value.get("port", 0), + value.get("username", ""), + value.get("password", ""), + ) + except (TypeError, ValueError) as exc: + raise RequestError("network_exit is invalid", 400) from exc + if ( + not all( + isinstance(item, str) + for item in (exit.protocol, exit.host, exit.username, exit.credential) + ) + or type(exit.port) is not int + ): + raise RequestError("network_exit is invalid", 400) + return exit + + +def browser_tmpfs() -> dict[str, str]: + # These are in-container tmpfs mounts; no host path or bind mount is exposed. + tmp = os.path.join(os.sep, "tmp") + return { + tmp: "rw,nosuid,nodev,noexec,mode=1777,size=256m", + os.path.join(tmp, ".X11-unix"): "rw,nosuid,nodev,noexec,mode=1777,size=1m", + os.path.join(os.sep, "dev", "shm"): "rw,nosuid,nodev,noexec,size=256m", + os.path.join( + os.sep, "home", "ubuntu" + ): "rw,nosuid,nodev,noexec,uid=1000,gid=1000,mode=700,size=64m", + } + + +def proxy_port(proxy_url: str) -> int: + return urlsplit(proxy_url).port or 0 + + +def has_control(value: str) -> bool: + return any(ord(char) < 0x20 or ord(char) == 0x7F for char in value) + + +def decode_generation( + value: dict, require_runtime: bool, require_network: bool +) -> dict: + allowed = {"binding_version", "runtime_id", "network_id"} + if not isinstance(value, dict) or set(value) - allowed: + raise RequestError( + "binding_version, runtime_id and network_id must identify the expected generation", + 400, + ) + binding = value.get("binding_version") + runtime = value.get("runtime_id", "") + network = value.get("network_id", "") + if ( + type(binding) is not int + or binding < 1 + or not isinstance(runtime, str) + or not isinstance(network, str) + or ( + runtime + and runtime != RUNTIME_CLEANUP_SENTINEL + and not CONTAINER_ID_RE.fullmatch(runtime) + ) + or (network and not EXIT_ID_RE.fullmatch(network)) + or (require_runtime and not runtime) + or (require_network and not network) + ): + raise RequestError( + "binding_version, runtime_id and network_id must identify the expected generation", + 400, + ) + return {"binding_version": binding, "runtime_id": runtime, "network_id": network} + + +def validate_proxy_restore(value: dict, alias: str) -> None: + allowed = { + "binding_version", + "runtime_id", + "network_id", + "network_exit_id", + "network_exit", + } + if set(value) - allowed: + raise RequestError("invalid proxy recovery request", 400) + generation = decode_generation( + { + key: value.get(key) + for key in ("binding_version", "runtime_id", "network_id") + }, + True, + True, + ) + exit_id = value.get("network_exit_id", "") + if not isinstance(exit_id, str) or (exit_id and not EXIT_ID_RE.fullmatch(exit_id)): + raise RequestError("invalid proxy recovery request", 400) + exit = parse_proxy_exit(value.get("network_exit", {})) + direct = not exit_id and exit == ProxyExit("", "", 0) + if bool(exit_id) != (not direct): + raise RequestError("invalid proxy recovery request", 400) + if not direct: + validate_proxy_exit(exit) + value["network_exit"] = exit + value.update(generation) + value["network_exit_id"] = exit_id + + +def valid_douyin_generation(value: dict) -> bool: + if not isinstance(value, dict): + return False + binding = value.get("binding_version") + runtime = value.get("runtime_id", "") + network = value.get("network_id", "") + exit_id = value.get("network_exit_id", "") + return ( + type(binding) is int + and binding > 0 + and isinstance(runtime, str) + and isinstance(network, str) + and isinstance(exit_id, str) + and bool(CONTAINER_ID_RE.fullmatch(runtime)) + and bool(EXIT_ID_RE.fullmatch(network)) + and (not exit_id or bool(EXIT_ID_RE.fullmatch(exit_id))) + ) + + +def valid_cookies(cookies: object) -> bool: + if not isinstance(cookies, list) or not 1 <= len(cookies) <= 64: + return False + for cookie in cookies: + if not isinstance(cookie, dict): + return False + allowed = { + "name", + "value", + "domain", + "path", + "secure", + "http_only", + "same_site", + "expires", + } + if set(cookie) - allowed: + return False + name = cookie.get("name") + value = cookie.get("value") + domain = cookie.get("domain") + path_value = cookie.get("path") + path = "/" if path_value is None else path_value + if ( + not isinstance(name, str) + or not isinstance(value, str) + or not isinstance(domain, str) + or not isinstance(path, str) + ): + return False + if ( + not name + or len(name) > 256 + or len(value) > 4096 + or len(domain) > 256 + or len(path) > 256 + or domain != domain.lower().strip() + or (domain != "douyin.com" and not domain.endswith(".douyin.com")) + or not path.startswith("/") + or has_control(name) + or has_control(value) + or has_control(path) + or (cookie.get("same_site", "") not in {"", "Lax", "Strict", "None"}) + ): + return False + expires = cookie.get("expires", 0) + if type(expires) not in (int, float) or expires < 0: + return False + return True + + +def valid_douyin_url(raw: object) -> bool: + if not isinstance(raw, str): + return False + try: + parsed = urlsplit(raw) + except ValueError: + return False + if ( + parsed.scheme != "https" + or parsed.netloc != "www.douyin.com" + or parsed.username + or parsed.fragment + ): + return False + query = parse_qs(parsed.query, keep_blank_values=True) + if parsed.path == DOUYIN_IDENTITY_PATH: + return query == {"aid": ["6383"], "device_platform": ["webapp"]} + if parsed.path == DOUYIN_WORKS_PATH: + return ( + len(query) == 3 + and valid_account_key_query(query, "sec_user_id") + and query.get("count") == ["20"] + and numeric_cursor(query.get("max_cursor")) + ) + if parsed.path == DOUYIN_COMMENTS_PATH: + return ( + len(query) == 3 + and valid_account_key_query(query, "aweme_id") + and query.get("count") == ["20"] + and numeric_cursor(query.get("cursor")) + ) + return False + + +def valid_account_key_query(query: dict[str, list[str]], key: str) -> bool: + return len(query.get(key, [])) == 1 and bool( + DOUYIN_ACCOUNT_KEY_RE.fullmatch(query[key][0]) + ) + + +def numeric_cursor(values: list[str] | None) -> bool: + if not values or len(values) != 1 or not values[0].isdigit(): + return False + try: + return int(values[0]) >= 0 + except ValueError: + return False + + +def same_network_members( + current: TenantNetworkGeneration, expected: TenantNetworkGeneration +) -> bool: + return ( + current.id == expected.id + and current.name == expected.name + and current.runtime_attached == expected.runtime_attached + and current.self_member == expected.self_member + and set(current.gateway_members) == set(expected.gateway_members) + ) + + +def load_config(env: Mapping[str, str] | None = None) -> dict: + env = os.environ if env is None else env + listen = env.get("LISTEN_ADDR", ":8081").strip() + socket_path = env.get("DOCKER_SOCKET", "/var/run/docker.sock").strip() + network = env.get("BROWSER_NETWORK", "creatorhub_browser").strip() + token = env.get("GATEWAY_TOKEN", "").strip() + host, port = split_listen_address(listen) + if ( + not socket_path + or len(token) < 16 + or not re.fullmatch(r"^[A-Za-z0-9][A-Za-z0-9_.-]{0,63}$", network) + or network == CONTROL_NETWORK + ): + raise ValueError("invalid gateway configuration") + if not 1 <= port <= 65535: + raise ValueError("LISTEN_ADDR port must be 1..65535") + return { + "listen": (host, port), + "docker_socket": socket_path, + "network": network, + "token": token, + } + + +def split_listen_address(value: str) -> tuple[str, int]: + if value.startswith(":"): + host, port_text = "", value[1:] + elif value.startswith("["): + closing = value.find("]:") + if closing <= 1: + raise ValueError("LISTEN_ADDR must be host:port") + host, port_text = value[1:closing], value[closing + 2 :] + else: + if ":" not in value: + raise ValueError("LISTEN_ADDR must be host:port") + host, port_text = value.rsplit(":", 1) + try: + port = int(port_text) + except ValueError as exc: + raise ValueError("LISTEN_ADDR port must be an integer") from exc + return host, port + + +def run() -> None: + config = load_config() + logging.basicConfig(level=logging.INFO, format="%(message)s") + docker = DockerClient(config["docker_socket"]) + gateway = Gateway(docker, config["network"], config["token"], socket.gethostname()) + server = GatewayHTTPServer(config["listen"], gateway) + LOG.info( + json.dumps( + { + "service": "docker-gateway", + "listen_addr": f"{config['listen'][0]}:{config['listen'][1]}", + "network": config["network"], + } + ) + ) + shutdown_requested = threading.Event() + + def request_shutdown(_signum, _frame) -> None: + if shutdown_requested.is_set(): + return + shutdown_requested.set() + threading.Thread( + target=server.shutdown, + name="gateway-shutdown", + daemon=True, + ).start() + + signal.signal(signal.SIGINT, request_shutdown) + signal.signal(signal.SIGTERM, request_shutdown) + try: + server.serve_forever() + finally: + # Stop accepting first, then let in-flight work finish before closing + # the browser and proxy dependencies it may still own. + server.wait_for_requests(30.0) + gateway.subscriptions.close() + gateway.proxies.close() + server.server_close() + + +if __name__ == "__main__": + run() diff --git a/cmd/docker_gateway/proxy.py b/cmd/docker_gateway/proxy.py new file mode 100644 index 0000000..12ea92c --- /dev/null +++ b/cmd/docker_gateway/proxy.py @@ -0,0 +1,656 @@ +"""Small in-memory HTTP/SOCKS proxy used by one browser generation.""" + +from __future__ import annotations + +import base64 +import select +import socket +import socketserver +import ssl +import threading +from collections.abc import Callable +from contextlib import suppress +from dataclasses import dataclass +from typing import cast +from urllib.parse import urlsplit + + +@dataclass(frozen=True) +class ProxyExit: + protocol: str + host: str + port: int + username: str = "" + credential: str = "" + + +class _ThreadingTCPServer(socketserver.ThreadingMixIn, socketserver.TCPServer): + allow_reuse_address = True + daemon_threads = True + + +class ProxyRegistry: + def __init__(self) -> None: + self._lock = threading.RLock() + self._proxies: dict[str, MemoryProxy] = {} + + def configure( + self, + alias: str, + binding_version: int, + bind_host: str, + port: int, + exit: ProxyExit, + network_id: str = "", + ) -> tuple[str, Callable[[], None]]: + with self._lock: + current = self._proxies.get(alias) + if current and current.matches( + binding_version, bind_host, port, exit, network_id + ): + return current.url, lambda: self._remove_object(alias, current) + if current: + self._proxies.pop(alias, None) + current.close() + proxy = MemoryProxy( + alias, binding_version, bind_host, port, exit, network_id + ) + self._proxies[alias] = proxy + return proxy.url, lambda: self._remove_object(alias, proxy) + + def bind( + self, + alias: str, + binding_version: int, + proxy_url: str, + runtime_id: str, + network_id: str = "", + ) -> bool: + with self._lock: + proxy = self._proxies.get(alias) + return bool( + proxy and proxy.bind(binding_version, proxy_url, runtime_id, network_id) + ) + + def ready( + self, alias: str, port: int, runtime_id: str, network_id: str = "" + ) -> bool: + with self._lock: + proxy = self._proxies.get(alias) + return bool(proxy and proxy.ready(port, runtime_id, network_id)) + + def remove( + self, alias: str, binding_version: int, runtime_id: str, network_id: str = "" + ) -> bool: + with self._lock: + proxy = self._proxies.get(alias) + if not proxy: + return True + if not proxy.identity_matches(binding_version, runtime_id, network_id): + return False + self._proxies.pop(alias, None) + proxy.close() + return True + + def _remove_object(self, alias: str, proxy: MemoryProxy) -> None: + with self._lock: + if self._proxies.get(alias) is not proxy: + return + self._proxies.pop(alias, None) + proxy.close() + + def close(self) -> None: + with self._lock: + proxies = list(self._proxies.values()) + self._proxies.clear() + for proxy in proxies: + proxy.close() + + +class MemoryProxy: + def __init__( + self, + alias: str, + binding_version: int, + bind_host: str, + port: int, + exit: ProxyExit, + network_id: str, + ) -> None: + self.alias = alias + self.binding_version = binding_version + self.bind_host = bind_host + self.exit = exit + self.network_id = network_id + self.runtime_id = "" + self._lock = threading.RLock() + self._tunnels: set[socket.socket] = set() + handler_type = type( + "ProxyHandler", + (_ProxyHandler,), + {"proxy": self}, + ) + self.server = _ThreadingTCPServer((bind_host, port), handler_type) + self.listener = self.server.socket + actual_port = self.listener.getsockname()[1] + self.url = f"http://docker-gateway:{actual_port}" + self._thread = threading.Thread( + target=self.server.serve_forever, + name=f"creatorhub-proxy-{alias}", + daemon=True, + ) + self._thread.start() + + def matches( + self, + binding_version: int, + bind_host: str, + port: int, + exit: ProxyExit, + network_id: str, + ) -> bool: + actual_port = self.listener.getsockname()[1] + return ( + self.binding_version == binding_version + and self.bind_host == bind_host + and (port == 0 or port == actual_port) + and self.exit == exit + and self.network_id == network_id + ) + + def bind( + self, binding_version: int, proxy_url: str, runtime_id: str, network_id: str + ) -> bool: + with self._lock: + if ( + binding_version != self.binding_version + or proxy_url != self.url + or not runtime_id + ): + return False + self.runtime_id = runtime_id + if network_id: + self.network_id = network_id + return True + + def ready(self, port: int, runtime_id: str, network_id: str) -> bool: + with self._lock: + return bool( + runtime_id + and self.runtime_id == runtime_id + and self.listener.getsockname()[1] == port + and (not network_id or self.network_id == network_id) + ) + + def identity_matches( + self, binding_version: int, runtime_id: str, network_id: str + ) -> bool: + with self._lock: + return ( + self.binding_version == binding_version + and bool(runtime_id) + and self.runtime_id == runtime_id + and (not network_id or self.network_id == network_id) + ) + + def close(self) -> None: + self.server.shutdown() + self.server.server_close() + with self._lock: + tunnels = list(self._tunnels) + self._tunnels.clear() + for connection in tunnels: + with suppress(OSError): + connection.close() + + def add_tunnel(self, connection: socket.socket) -> bool: + with self._lock: + if self.server.socket.fileno() < 0: + return False + self._tunnels.add(connection) + return True + + def remove_tunnel(self, connection: socket.socket) -> None: + with self._lock: + self._tunnels.discard(connection) + + def dial(self, target: str, timeout: float = 20.0) -> socket.socket: + exit = self.exit + if exit.protocol in ("http", "https"): + return _dial_http_proxy(exit, target, timeout) + if exit.protocol == "socks4": + return _dial_socks4(exit, target, timeout) + if exit.protocol == "socks5": + return _dial_socks5(exit, target, timeout) + raise OSError("unsupported proxy protocol") + + def forward_http( + self, + client: socket.socket, + method: str, + target: str, + headers: list[tuple[str, str]], + body: bytes, + ) -> None: + parsed = urlsplit(target) + if not parsed.hostname: + raise OSError("proxy request target is invalid") + target_host = parsed.hostname + target_port = parsed.port or (443 if parsed.scheme == "https" else 80) + origin_target = parsed.path or "/" + if parsed.query: + origin_target += "?" + parsed.query + if self.exit.protocol in ("http", "https"): + upstream = _open_host( + self.exit.host, self.exit.port, 20.0, self.exit.protocol == "https" + ) + request_target = target + request_headers = _forward_headers(headers, len(body)) + if self.exit.username: + token = base64.b64encode( + f"{self.exit.username}:{self.exit.credential}".encode() + ).decode() + request_headers.append(("Proxy-Authorization", f"Basic {token}")) + else: + upstream = self.dial(f"{target_host}:{target_port}") + request_target = origin_target + request_headers = _forward_headers(headers, len(body)) + try: + lines = [f"{method} {request_target} HTTP/1.1"] + lines.extend(f"{k}: {v}" for k, v in request_headers) + lines.append("Connection: close") + upstream.sendall(("\r\n".join(lines) + "\r\n\r\n").encode() + body) + _copy_until_close(upstream, client) + finally: + upstream.close() + + def tunnel(self, client: socket.socket, target: str) -> None: + upstream = self.dial(target) + try: + client.sendall(b"HTTP/1.1 200 Connection Established\r\n\r\n") + if not self.add_tunnel(client): + return + _relay(client, upstream) + finally: + self.remove_tunnel(client) + with suppress(OSError): + upstream.close() + + +class _ProxyHandler(socketserver.BaseRequestHandler): + proxy: MemoryProxy + + def handle(self) -> None: + client = self.request + client.settimeout(60.0) + try: + head, body = _read_request(client) + method, target, headers = _parse_request(head) + if method.upper() == "CONNECT": + self.proxy.tunnel(client, target) + else: + self.proxy.forward_http(client, method, target, headers, body) + except (OSError, ValueError): + with suppress(OSError): + client.sendall( + b"HTTP/1.1 502 Bad Gateway\r\nConnection: close\r\nContent-Length: 0\r\n\r\n" + ) + + +def _read_request(connection: socket.socket) -> tuple[bytes, bytes]: + data = bytearray() + while b"\r\n\r\n" not in data: + chunk = connection.recv(65536) + if not chunk: + raise OSError("proxy request closed") + data.extend(chunk) + if len(data) > 256 * 1024: + raise OSError("proxy headers too large") + split = data.index(b"\r\n\r\n") + 4 + head, initial_body = bytes(data[:split]), bytearray(data[split:]) + fields = _header_fields(head) + lengths = fields.get("content-length", []) + transfer_encoding = fields.get("transfer-encoding", []) + if lengths and transfer_encoding: + raise OSError("proxy request has both content length and transfer encoding") + if transfer_encoding: + if ( + len(transfer_encoding) != 1 + or transfer_encoding[0].lower().strip() != "chunked" + ): + raise OSError("unsupported proxy transfer encoding") + return head, _read_chunked_body(connection, initial_body) + content_length = 0 + if lengths: + if len(lengths) != 1: + raise OSError("proxy content length is ambiguous") + try: + content_length = int(lengths[0]) + except ValueError as exc: + raise OSError("invalid proxy content length") from exc + if content_length < 0 or content_length > 16 * 1024 * 1024: + raise OSError("proxy body too large") + while len(initial_body) < content_length: + chunk = connection.recv(min(65536, content_length - len(initial_body))) + if not chunk: + raise OSError("proxy body closed") + initial_body.extend(chunk) + return head, bytes(initial_body[:content_length]) + + +def _header_fields(head: bytes) -> dict[str, list[str]]: + fields: dict[str, list[str]] = {} + for line in head.decode("iso-8859-1").split("\r\n")[1:]: + if not line: + continue + if ":" not in line: + raise OSError("invalid proxy header") + key, value = line.split(":", 1) + fields.setdefault(key.strip().lower(), []).append(value.strip()) + return fields + + +def _read_line(connection: socket.socket, buffer: bytearray, limit: int) -> bytes: + while b"\r\n" not in buffer: + chunk = connection.recv(65536) + if not chunk: + raise OSError("proxy chunked body closed") + buffer.extend(chunk) + if len(buffer) > limit: + raise OSError("proxy chunk line too large") + end = buffer.index(b"\r\n") + line = bytes(buffer[:end]) + del buffer[: end + 2] + return line + + +def _read_chunked_body(connection: socket.socket, buffer: bytearray) -> bytes: + body = bytearray() + while True: + line = _read_line(connection, buffer, 8192) + size_text = line.split(b";", 1)[0].strip() + try: + size = int(size_text, 16) + except ValueError as exc: + raise OSError("invalid proxy chunk size") from exc + if size < 0: + raise OSError("invalid proxy chunk size") + if size == 0: + while True: + trailer = _read_line(connection, buffer, 256 * 1024) + if not trailer: + return bytes(body) + if b":" not in trailer: + raise OSError("invalid proxy trailer") + if len(body) + size > 16 * 1024 * 1024: + raise OSError("proxy body too large") + while len(buffer) < size + 2: + chunk = connection.recv(65536) + if not chunk: + raise OSError("proxy chunked body closed") + buffer.extend(chunk) + body.extend(buffer[:size]) + if buffer[size : size + 2] != b"\r\n": + raise OSError("invalid proxy chunk terminator") + del buffer[: size + 2] + + +def _forward_headers( + headers: list[tuple[str, str]], body_length: int +) -> list[tuple[str, str]]: + return [ + (k, v) + for k, v in headers + if k.lower() + not in ( + "proxy-authorization", + "connection", + "transfer-encoding", + "content-length", + ) + ] + [("Content-Length", str(body_length))] + + +def _parse_request(head: bytes) -> tuple[str, str, list[tuple[str, str]]]: + lines = head.decode("iso-8859-1").split("\r\n") + method, target, version = lines[0].split(" ", 2) + if version not in ("HTTP/1.0", "HTTP/1.1") or not target: + raise ValueError("invalid proxy request") + headers: list[tuple[str, str]] = [] + for line in lines[1:]: + if not line: + continue + if ":" not in line: + raise ValueError("invalid proxy header") + key, value = line.split(":", 1) + headers.append((key.strip(), value.strip())) + if method.upper() == "CONNECT": + if ":" not in target: + raise ValueError("CONNECT target missing port") + return method, target, headers + return method, target, headers + + +def _open_host( + host: str, port: int, timeout: float, tls: bool = False +) -> socket.socket: + connection = socket.create_connection((host, port), timeout=timeout) + if tls: + context = ssl.create_default_context() + connection = context.wrap_socket(connection, server_hostname=host) + return connection + + +def _dial_http_proxy(exit: ProxyExit, target: str, timeout: float) -> socket.socket: + connection = _open_host(exit.host, exit.port, timeout, exit.protocol == "https") + try: + lines = [f"CONNECT {target} HTTP/1.1", f"Host: {target}"] + if exit.username: + token = base64.b64encode( + f"{exit.username}:{exit.credential}".encode() + ).decode() + lines.append(f"Proxy-Authorization: Basic {token}") + connection.sendall(("\r\n".join(lines) + "\r\n\r\n").encode()) + status, buffered = _read_connect_response(connection) + if status != 200: + raise OSError(f"upstream proxy returned {status}") + result = cast( + socket.socket, + _BufferedSocket(connection, buffered) if buffered else connection, + ) + result.settimeout(None) + return result + except Exception: + connection.close() + raise + + +def _dial_socks4(exit: ProxyExit, target: str, timeout: float) -> socket.socket: + try: + host, port_text = target.rsplit(":", 1) + port = int(port_text) + except (ValueError, IndexError) as exc: + raise OSError("invalid SOCKS4 target") from exc + if not 1 <= port <= 65535: + raise OSError("invalid SOCKS4 target port") + connection = _open_host(exit.host, exit.port, timeout) + try: + ip = socket.inet_aton(host) if _is_ipv4(host) else b"\x00\x00\x00\x01" + payload = ( + b"\x04\x01" + + port.to_bytes(2, "big") + + ip + + exit.username.encode() + + b"\x00" + ) + if ip == b"\x00\x00\x00\x01": + payload += host.encode() + b"\x00" + connection.sendall(payload) + response = _recv_exact(connection, 8) + if response[1] != 90: + raise OSError("SOCKS4 proxy rejected connection") + connection.settimeout(None) + return connection + except Exception: + connection.close() + raise + + +def _dial_socks5(exit: ProxyExit, target: str, timeout: float) -> socket.socket: + try: + host, port_text = target.rsplit(":", 1) + port = int(port_text) + except (ValueError, IndexError) as exc: + raise OSError("invalid SOCKS5 target") from exc + if not 1 <= port <= 65535: + raise OSError("invalid SOCKS5 target port") + connection = _open_host(exit.host, exit.port, timeout) + try: + methods = b"\x05\x01\x02" if exit.username else b"\x05\x01\x00" + connection.sendall(methods) + version, selected = _recv_exact(connection, 2) + if version != 5 or selected == 255: + raise OSError("SOCKS5 authentication method rejected") + if selected == 2: + user, credential = exit.username.encode(), exit.credential.encode() + if len(user) > 255 or len(credential) > 255: + raise OSError("SOCKS5 credentials too long") + connection.sendall( + b"\x01" + + bytes([len(user)]) + + user + + bytes([len(credential)]) + + credential + ) + if _recv_exact(connection, 2)[1] != 0: + raise OSError("SOCKS5 authentication rejected") + elif exit.username: + raise OSError("SOCKS5 proxy skipped required authentication") + if _is_ipv4(host): + address = b"\x01" + socket.inet_aton(host) + else: + try: + address = b"\x04" + socket.inet_pton(socket.AF_INET6, host) + except OSError: + encoded = host.encode() + if len(encoded) > 255: + raise OSError("SOCKS5 target too long") from None + address = b"\x03" + bytes([len(encoded)]) + encoded + connection.sendall(b"\x05\x01\x00" + address + port.to_bytes(2, "big")) + header = _recv_exact(connection, 4) + if header[0] != 5 or header[1] != 0: + raise OSError("SOCKS5 proxy rejected connection") + if header[3] == 1: + length = 4 + elif header[3] == 3: + length = _recv_exact(connection, 1)[0] + elif header[3] == 4: + length = 16 + else: + raise OSError("invalid SOCKS5 response") + _recv_exact(connection, length + 2) + connection.settimeout(None) + return connection + except Exception: + connection.close() + raise + + +def _read_status(connection: socket.socket) -> int: + return _read_connect_response(connection)[0] + + +def _read_connect_response(connection: socket.socket) -> tuple[int, bytes]: + # Consume the complete response header. A proxy may return the first + # tunnel bytes in the same read, so retain bytes after the header block. + data = bytearray() + while b"\r\n\r\n" not in data: + chunk = connection.recv(65536) + if not chunk: + raise OSError("proxy response closed before headers") + data.extend(chunk) + if len(data) > 256 * 1024: + raise OSError("proxy response headers too large") + split = data.index(b"\r\n\r\n") + 4 + header = bytes(data[:split]) + lines = header.decode("iso-8859-1").split("\r\n") + parts = lines[0].split(" ", 2) + if len(parts) < 2 or not parts[0].startswith("HTTP/"): + raise OSError("invalid proxy response status") + try: + status = int(parts[1]) + except ValueError as exc: + raise OSError("invalid proxy response status") from exc + for line in lines[1:-2]: + if line and ":" not in line: + raise OSError("invalid proxy response header") + return status, bytes(data[split:]) + + +class _BufferedSocket: + def __init__(self, connection: socket.socket, buffered: bytes) -> None: + self.connection = connection + self.buffered = bytearray(buffered) + + def recv(self, size: int, flags: int = 0) -> bytes: + if self.buffered: + result = bytes(self.buffered[:size]) + del self.buffered[:size] + return result + return ( + self.connection.recv(size, flags) if flags else self.connection.recv(size) + ) + + def sendall(self, data: bytes) -> None: + self.connection.sendall(data) + + def send(self, data: bytes, flags: int = 0) -> int: + return ( + self.connection.send(data, flags) if flags else self.connection.send(data) + ) + + def settimeout(self, value: float | None) -> None: + self.connection.settimeout(value) + + def close(self) -> None: + self.connection.close() + + +def _recv_exact(connection: socket.socket, size: int) -> bytes: + result = bytearray() + while len(result) < size: + chunk = connection.recv(size - len(result)) + if not chunk: + raise OSError("proxy connection closed") + result.extend(chunk) + return bytes(result) + + +def _copy_until_close(source: socket.socket, target: socket.socket) -> None: + while True: + data = source.recv(65536) + if not data: + return + target.sendall(data) + + +def _relay(left: socket.socket, right: socket.socket) -> None: + sockets = [left, right] + while sockets: + readable, _, _ = select.select(sockets, [], [], 60.0) + if not readable: + return + for source in readable: + destination = right if source is left else left + data = source.recv(65536) + if not data: + return + destination.sendall(data) + + +def _is_ipv4(host: str) -> bool: + try: + socket.inet_aton(host) + return True + except OSError: + return False diff --git a/cmd/docker_gateway/test_gateway.py b/cmd/docker_gateway/test_gateway.py new file mode 100644 index 0000000..d342cb9 --- /dev/null +++ b/cmd/docker_gateway/test_gateway.py @@ -0,0 +1,2076 @@ +from __future__ import annotations + +import io +import json +import socket +import threading +import unittest +from collections import deque +from contextlib import contextmanager +from importlib import import_module +from typing import Any, cast +from unittest.mock import Mock, patch + +import websocket + +from .docker_client import ( + BINDING_VERSION_LABEL, + BROWSER_NETWORK_ROLE, + GATEWAY_MEMBER_LABEL, + MANAGED_LABEL, + NETWORK_EXIT_LABEL, + NETWORK_ID_LABEL, + NETWORK_ROLE_LABEL, + RESERVATION_LABEL, + RESERVATION_OWNER_LABEL, + RUNTIME_ID_LABEL, + AliasReservationManager, + DockerClient, + DockerError, + DockerResponse, + GenerationConflict, + NetworkSetupError, + TenantNetworkGeneration, + UnmanagedContainer, + split_image_ref, + tenant_network_name, +) +from .douyin import ( + BrowserResponse, + CDPConnection, + DouyinBrowser, + DouyinError, + DouyinSubscription, + SubscriptionManager, + ack_expression, + action_expression, + detect_challenge, + im_expression, + install_expression, + is_douyin_url, + normalize_notice, + notice_ids, + wait_expression, +) +from .proxy import ( + MemoryProxy, + ProxyExit, + ProxyRegistry, + _copy_until_close, + _dial_http_proxy, + _dial_socks4, + _dial_socks5, + _is_ipv4, + _parse_request, + _read_request, + _read_status, + _recv_exact, +) + +gateway_module = import_module(f"{__package__}.gateway") +douyin_module = import_module(f"{__package__}.douyin") +proxy_module = import_module(f"{__package__}.proxy") +docker_client_module = import_module(f"{__package__}.docker_client") +Gateway = gateway_module.Gateway +RequestError = gateway_module.RequestError +browser_tmpfs = gateway_module.browser_tmpfs +decode_generation = gateway_module.decode_generation +json_bytes = gateway_module.json_bytes +load_config = gateway_module.load_config +split_listen_address = gateway_module.split_listen_address +valid_douyin_url = gateway_module.valid_douyin_url +validate_create = gateway_module.validate_create +parse_proxy_exit = gateway_module.parse_proxy_exit +validate_proxy_exit = gateway_module.validate_proxy_exit +validate_proxy_restore = gateway_module.validate_proxy_restore +valid_cookies = gateway_module.valid_cookies +valid_douyin_generation = gateway_module.valid_douyin_generation +valid_account_key_query = gateway_module.valid_account_key_query +numeric_cursor = gateway_module.numeric_cursor +proxy_port = gateway_module.proxy_port +has_control = gateway_module.has_control + + +class FakeSocket: + def __init__(self, messages: list[object]) -> None: + self.messages = list(messages) + self.sent: list[bytes] = [] + self.timeout = 0.0 + + def send(self, data: bytes) -> None: + self.sent.append(data) + + def recv(self) -> str: + if not self.messages: + raise TimeoutError("no more messages") + return json.dumps(self.messages.pop(0)) + + def settimeout(self, value: float) -> None: + self.timeout = value + + def close(self) -> None: + return None + + +class FakeConnection: + def __init__(self, values: list[object]) -> None: + self.values = list(values) + + def evaluate(self, expression: str) -> object: + del expression + if not self.values: + raise DouyinError("no fake response") + return self.values.pop(0) + + +class FakeDocker: + def __init__(self, containers: list[dict] | None = None) -> None: + self.containers = containers or [] + + def container_network_address(self, container_id: str, network_id: str) -> str: + del container_id, network_id + return "192.0.2.10" + + def request( + self, method: str, path: str, body: object | None = None + ) -> DockerResponse: + del body + if method == "GET" and path.startswith("/containers/json"): + return DockerResponse(200, "OK", json.dumps(self.containers).encode()) + return DockerResponse(404, "Not Found", b"") + + +class GatewayValidationTests(unittest.TestCase): + def test_create_and_generation_validation(self) -> None: + value = { + "alias": "safe-account", + "name": "Safe account", + "image": "creatorhub/browser:latest", + "cmd": ["about:blank"], + "volume": "creatorhub-safe-account", + "binding_version": 1, + "stopped": True, + } + validate_create(value) + self.assertEqual(value["network_exit"], ProxyExit("", "", 0)) + self.assertEqual(value["network_exit_id"], "") + self.assertEqual( + decode_generation( + { + "binding_version": 1, + "runtime_id": "runtime-not-found", + "network_id": "network-id", + }, + False, + False, + )["runtime_id"], + "runtime-not-found", + ) + with self.assertRaises(RequestError): + validate_create( + {**value, "cmd": ["--proxy-server=http://x", "about:blank"]} + ) + with self.assertRaises(RequestError): + decode_generation( + {"binding_version": 1, "runtime_id": 4, "network_id": "n"}, True, True + ) + + def test_config_and_urls(self) -> None: + self.assertEqual(split_listen_address(":8081"), ("", 8081)) + self.assertEqual(split_listen_address("[::1]:8081"), ("::1", 8081)) + with self.assertRaises(ValueError): + split_listen_address("missing-port") + config = load_config( + { + "LISTEN_ADDR": ":8081", + "DOCKER_SOCKET": "/var/run/docker.sock", + "BROWSER_NETWORK": "creatorhub_browser", + "GATEWAY_TOKEN": "0123456789abcdef", + } + ) + self.assertEqual(config["listen"], ("", 8081)) + tmp_mount = next(path for path in browser_tmpfs() if path.endswith("tmp")) + self.assertTrue(browser_tmpfs()[tmp_mount].startswith("rw,")) + self.assertTrue( + valid_douyin_url( + "https://www.douyin.com/aweme/v1/web/user/profile/self/?aid=6383&device_platform=webapp" + ) + ) + self.assertFalse(valid_douyin_url("https://www.douyin.com.evil/")) + self.assertTrue(is_douyin_url("https://www.douyin.com/video/123")) + self.assertFalse(is_douyin_url("https://www.douyin.com.evil/video/123")) + + def test_http_routes_and_body_validation(self) -> None: + handler = gateway_module.GatewayHandler.__new__(gateway_module.GatewayHandler) + gateway = Mock() + gateway.list_browsers.return_value = [] + gateway.douyin_identity.return_value = {"uid": "123"} + gateway.douyin_action.return_value = {"status": "succeeded"} + gateway.poll_douyin_events.return_value = [] + server = Mock() + server.gateway = gateway + cast(Any, handler).server = server + cast(Any, handler).server_as_gateway = lambda: server + self.assertEqual(handler._route("GET", "/v1/browsers", {}, {}), []) + self.assertEqual( + handler._route("POST", "/v1/browsers", {}, {}), + (201, gateway.create.return_value), + ) + handler._route( + "DELETE", + "/v1/browsers/safe", + {}, + {"binding_version": 1, "runtime_id": "a" * 64, "network_id": ""}, + ) + gateway.remove.assert_called_once_with( + "safe", {"binding_version": 1, "runtime_id": "a" * 64, "network_id": ""} + ) + handler._route("POST", "/v1/browsers/safe/start", {}, {}) + handler._route("POST", "/v1/browsers/safe/stop", {}, {}) + handler._route("POST", "/v1/browsers/safe/proxy", {}, {}) + handler._route("POST", "/v1/browsers/safe/douyin/cookies", {}, {}) + handler._route("POST", "/v1/browsers/safe/douyin/get", {}, {}) + handler._route("POST", "/v1/browsers/safe/douyin/identity", {}, {}) + handler._route("POST", "/v1/browsers/safe/douyin/action", {}, {}) + self.assertEqual( + handler._route("GET", "/v1/browsers/safe/douyin/events", {}, {}), [] + ) + handler._route("POST", "/v1/browsers/safe/douyin/events", {}, {}) + handler._route("DELETE", "/v1/browsers/safe/douyin/events", {}, {}) + with self.assertRaises(RequestError): + handler._route("GET", "/v1/unknown", {}, {}) + cast(Any, handler).headers = {"Content-Length": "7"} + cast(Any, handler).rfile = io.BytesIO(b'{"x":1}') + self.assertEqual(handler._body(), {"x": 1}) + cast(Any, handler).headers = {} + with self.assertRaises(RequestError): + handler._body() + test_token = "x" * 16 + cast(Any, handler).headers = {"Authorization": f"Bearer {test_token}"} + server.gateway.token = test_token + self.assertTrue(handler._authorized()) + handler._respond = Mock() + handler._handle_exception("/v1", RequestError("bad", 400)) + handler._handle_exception("/v1", ValueError("bad")) + self.assertEqual(handler._respond.call_count, 2) + + def test_get_event_route_reads_generation_body(self) -> None: + handler = gateway_module.GatewayHandler.__new__(gateway_module.GatewayHandler) + cast(Any, handler).path = "/v1/browsers/safe/douyin/events?wait=1" + cast(Any, handler).headers = {"Content-Length": "67"} + cast(Any, handler).rfile = io.BytesIO( + b'{"binding_version":1,"runtime_id":"runtime","network_id":"network"}' + ) + cast(Any, handler)._authorized = lambda: True + cast(Any, handler)._route = Mock(return_value=[]) + cast(Any, handler)._respond = Mock() + handler._dispatch("GET") + cast(Any, handler)._route.assert_called_once_with( + "GET", + "/v1/browsers/safe/douyin/events", + {"wait": ["1"]}, + {"binding_version": 1, "runtime_id": "runtime", "network_id": "network"}, + ) + + def test_validation_boundaries(self) -> None: + self.assertEqual(proxy_port("http://docker-gateway:1234"), 1234) + self.assertTrue(has_control("bad\nvalue")) + exit_value = parse_proxy_exit( + {"protocol": "http", "host": "proxy", "port": 8080} + ) + validate_proxy_exit(exit_value) + with self.assertRaises(RequestError): + validate_proxy_exit( + parse_proxy_exit({"protocol": "ftp", "host": "proxy", "port": 21}) + ) + with self.assertRaises(RequestError): + validate_proxy_exit(ProxyExit("http", "proxy", 0)) + cookie = {"name": "sessionid", "value": "v", "domain": ".douyin.com"} + self.assertTrue(valid_cookies([cookie])) + self.assertFalse(valid_cookies([{**cookie, "domain": "evil.example"}])) + self.assertFalse(valid_cookies([{**cookie, "same_site": "bad"}])) + generation = { + "binding_version": 1, + "runtime_id": "a" * 64, + "network_id": "b" * 64, + "network_exit_id": "exit", + } + self.assertTrue(valid_douyin_generation(generation)) + self.assertFalse( + valid_douyin_generation({**generation, "binding_version": True}) + ) + self.assertTrue(valid_account_key_query({"account": ["account"]}, "account")) + self.assertFalse(valid_account_key_query({"account": ["bad key"]}, "account")) + self.assertTrue(numeric_cursor(["0"])) + self.assertFalse(numeric_cursor(["-1"])) + restore = { + "binding_version": 1, + "runtime_id": "a" * 64, + "network_id": "b" * 64, + "network_exit_id": "exit", + "network_exit": {"protocol": "http", "host": "proxy", "port": 8080}, + } + validate_proxy_restore(restore, "safe") + with self.assertRaises(RequestError): + validate_proxy_restore({**restore, "network_exit_id": ""}, "safe") + with self.assertRaises(RequestError): + validate_proxy_restore({"network_exit_id": ""}, "safe") + with self.assertRaises(RequestError): + validate_create({"alias": "safe", "unknown": True}) + with self.assertRaises(RequestError): + validate_create( + { + "alias": "safe", + "name": "Safe", + "image": "bad image", + "volume": "safe", + "binding_version": 1, + "cmd": ["about:blank"], + } + ) + with self.assertRaises(ValueError): + load_config({"GATEWAY_TOKEN": "short"}) + self.assertFalse(valid_douyin_url("http://www.douyin.com/video/1")) + self.assertFalse(valid_douyin_url("https://www.douyin.com/unknown")) + self.assertFalse(valid_cookies("not-a-list")) + self.assertFalse( + valid_cookies( + [ + { + "name": "x", + "value": "v", + "domain": ".douyin.com", + "path": "relative", + } + ] + ) + ) + with self.assertRaises(RequestError): + decode_generation( + {"binding_version": 1, "runtime_id": "a" * 64}, True, True + ) + + def test_list_keeps_other_environments_when_one_network_is_missing(self) -> None: + class PartialDocker(FakeDocker): + def container_network_address( + self, container_id: str, network_id: str + ) -> str: + if container_id == "broken": + raise GenerationConflict("network attachment disappeared") + return "192.0.2.10" + + containers = [ + { + "Id": "broken", + "Labels": { + MANAGED_LABEL: "true", + RUNTIME_ID_LABEL: "broken-account", + BINDING_VERSION_LABEL: "1", + NETWORK_EXIT_LABEL: "", + NETWORK_ID_LABEL: "b" * 64, + }, + "State": "running", + "Status": "Up", + }, + { + "Id": "healthy", + "Labels": { + MANAGED_LABEL: "true", + RUNTIME_ID_LABEL: "healthy-account", + BINDING_VERSION_LABEL: "1", + NETWORK_EXIT_LABEL: "", + NETWORK_ID_LABEL: "c" * 64, + }, + "State": "running", + "Status": "Up", + }, + ] + result = Gateway( + cast(DockerClient, PartialDocker(containers)), + "creatorhub_browser", + "0123456789abcdef", + "gateway", + ).list_browsers() + self.assertEqual( + [item["alias"] for item in result], ["broken-account", "healthy-account"] + ) + self.assertEqual(result[0]["endpoint"], "") + self.assertIn("error", result[0]) + self.assertEqual(result[1]["endpoint"], "http://192.0.2.10:9222") + + def test_container_network_address(self) -> None: + docker = DockerClient("/var/run/docker.sock") + cast(Any, docker).request = lambda method, path: DockerResponse( + 200, + "OK", + json.dumps( + { + "NetworkSettings": { + "Networks": { + "tenant": { + "NetworkID": "network", + "IPAddress": "192.0.2.20", + } + } + } + } + ).encode(), + ) + self.assertEqual( + docker.container_network_address("container", "network"), "192.0.2.20" + ) + cast(Any, docker).request = lambda method, path: DockerResponse( + 200, "OK", b'{"NetworkSettings":{"Networks":{}}}' + ) + with self.assertRaises(GenerationConflict): + docker.container_network_address("container", "network") + + def test_list_browsers_and_json(self) -> None: + docker = FakeDocker( + [ + { + "Id": "container-id", + "Labels": { + "io.creatorhub.managed": "true", + "io.creatorhub.runtime-id": "safe-account", + "io.creatorhub.binding-version": "2", + "io.creatorhub.network-exit-id": "", + "io.creatorhub.network-id": "b" * 64, + }, + "State": "running", + "Status": "Up 1 second", + } + ] + ) + gateway = Gateway( + cast(DockerClient, docker), + "creatorhub_browser", + "0123456789abcdef", + "gateway", + ) + self.assertEqual(gateway.list_browsers()[0]["id"], "container-id") + self.assertEqual( + gateway.list_browsers()[0]["endpoint"], "http://192.0.2.10:9222" + ) + self.assertEqual( + json_bytes({"text": "中文"}), b'{"text":"\xe4\xb8\xad\xe6\x96\x87"}' + ) + + +class CDPTests(unittest.TestCase): + def test_connect_response_headers_are_drained(self) -> None: + connection = ChunkSocket( + [b"HTTP/1.1 200 Connection Established\r\n", b"X-Proxy: value\r\n\r\nTLS"] + ) + self.assertEqual(_read_status(cast(socket.socket, connection)), 200) + + def test_command_queues_events_for_wait_event(self) -> None: + socket_ = FakeSocket( + [ + { + "method": "Page.frameNavigated", + "params": {"frame": {"id": "frame-1"}}, + }, + {"id": 1, "result": {}}, + ] + ) + connection = CDPConnection(cast(websocket.WebSocket, socket_)) + self.assertEqual(connection.command("Page.enable"), {}) + event = connection.wait_event( + "Page.frameNavigated", lambda params: params["frame"]["id"] == "frame-1" + ) + self.assertEqual(event["method"], "Page.frameNavigated") + self.assertTrue(socket_.sent) + + def test_eof_status_is_a_failure(self) -> None: + class Closed: + def recv(self, size: int) -> bytes: + del size + return b"" + + with self.assertRaises(OSError): + _read_status(cast(socket.socket, Closed())) + + def test_expression_markers_and_challenge(self) -> None: + expression = im_expression({"text": "hello EXPECTED_UID_VALUE"}, "123") + self.assertIn("hello EXPECTED_UID_VALUE", expression) + self.assertIn( + "https://www.douyin.com", action_expression({"action": "like_work"}) + ) + self.assertEqual(detect_challenge(429, "captcha"), "") + self.assertEqual(detect_challenge(412, ""), "captcha") + + def test_notification_details_normalize_safe_targets(self) -> None: + notice = { + "nid_str": "9007199254740993", + "user_id": "99491952055", + "create_time": 1700000000, + "aweme_id": "123456", + "comment": { + "from_user": [{"uid": "7654321"}], + "comment": {"cid_str": "987654", "user": {"uid": "7654321"}}, + }, + } + self.assertEqual( + normalize_notice(notice), + { + "event_key": "9007199254740993", + "event_type": "comment", + "interactor_uid": "7654321", + "comment_id": "987654", + "work_id": "123456", + "platform_event_at": "2023-11-14T22:13:20+00:00", + }, + ) + favorite = {"nid_str": "7", "favorite": {"from_user": [{"uid": "1"}]}} + self.assertIsNone(normalize_notice(favorite)) + + def test_notification_detail_retries_partial_response(self) -> None: + details = [ + {"nid_str": "1", "user_id": "123"}, + {"nid_str": "2", "user_id": "123"}, + ] + payload = json.dumps({"status_code": 0, "notice_list_v2": details}) + subscription = DouyinSubscription.__new__(DouyinSubscription) + subscription.uid = "123" + subscription.connection = cast( + CDPConnection, + FakeConnection( + [ + json.dumps( + { + "status": 200, + "body": json.dumps( + {"status_code": 0, "notice_list_v2": details[:1]} + ), + } + ), + json.dumps({"status": 200, "body": payload}), + ] + ), + ) + subscription._get_connection = lambda: subscription.connection + with patch.object(douyin_module.time, "sleep"): + self.assertEqual(len(subscription._details(["1", "2"])), 2) + + def test_notification_detail_rejects_unexpected_id(self) -> None: + subscription = DouyinSubscription.__new__(DouyinSubscription) + subscription.uid = "123" + subscription.connection = cast( + CDPConnection, + FakeConnection( + [ + json.dumps( + { + "status": 200, + "body": json.dumps( + { + "status_code": 0, + "notice_list_v2": [ + {"nid_str": "9", "user_id": "123"} + ], + } + ), + } + ) + ] + ), + ) + subscription._get_connection = lambda: subscription.connection + with self.assertRaises(DouyinError): + subscription._details(["1"]) + + +class BrowserCDP: + def __init__(self, values: list[object]) -> None: + self.values = list(values) + self.commands: list[tuple[str, dict | None]] = [] + self.events: list[str] = [] + self.closed = False + + def command(self, method: str, params: dict | None = None) -> dict: + self.commands.append((method, params)) + if method == "Page.navigate": + return {"frameId": "frame-1"} + return {} + + def wait_event(self, method: str, predicate: object, timeout: float = 15.0) -> dict: + del predicate, timeout + self.events.append(method) + return {"method": method} + + def evaluate(self, expression: str) -> object: + del expression + if not self.values: + raise DouyinError("fake CDP value exhausted") + return self.values.pop(0) + + def close(self) -> None: + self.closed = True + + +class FakeHTTPResponse: + def __init__(self, status: int, body: bytes) -> None: + self.status = status + self.body = body + + def read(self, limit: int = -1) -> bytes: + del limit + return self.body + + +class FakeHTTPConnection: + def __init__(self, response: FakeHTTPResponse) -> None: + self.response = response + self.requested: list[tuple[str, str]] = [] + self.closed = False + + def request( + self, + method: str, + path: str, + body: bytes | None = None, + headers: dict[str, str] | None = None, + ) -> None: + del body, headers + self.requested.append((method, path)) + + def getresponse(self) -> FakeHTTPResponse: + return self.response + + def close(self) -> None: + self.closed = True + + +class BrowserTests(unittest.TestCase): + def _with_connection(self, browser: DouyinBrowser, connection: BrowserCDP) -> None: + @contextmanager + def bound(alias: str): + del alias + yield connection + + cast(Any, browser).connection = bound + + def test_cookie_navigation_and_fetch(self) -> None: + cdp = BrowserCDP( + [{"origin": "https://www.douyin.com", "state": "complete"}, "{}"] + ) + browser = DouyinBrowser() + self._with_connection(browser, cdp) + browser.set_cookies( + "safe", [{"name": "sessionid", "value": "v", "domain": ".douyin.com"}] + ) + self.assertIn("Page.navigate", [method for method, _ in cdp.commands]) + set_command = next( + params for method, params in cdp.commands if method == "Network.setCookies" + ) + assert set_command is not None + self.assertEqual(set_command["cookies"][0]["domain"], ".douyin.com") + + cdp = BrowserCDP( + [ + "https://www.douyin.com", + {"status": 200, "body": "{}", "too_large": False}, + ] + ) + self._with_connection(browser, cdp) + response = browser.get( + "safe", "https://www.douyin.com/aweme/v1/web/user/profile/self/?aid=6383" + ) + self.assertEqual(response.status, 200) + + def test_connect_identity_and_actions(self) -> None: + target = [ + { + "type": "page", + "webSocketDebuggerUrl": "ws://127.0.0.1:9222/devtools/page/1", + } + ] + http = FakeHTTPConnection(FakeHTTPResponse(200, json.dumps(target).encode())) + with ( + patch.object( + douyin_module.http.client, "HTTPConnection", return_value=http + ), + patch.object( + douyin_module.websocket, + "create_connection", + return_value=FakeSocket([]), + ), + ): + connection = DouyinBrowser(lambda alias: "http://127.0.0.1:9222")._connect( + "safe" + ) + self.assertIsInstance(connection, CDPConnection) + self.assertEqual(http.requested[0], ("GET", "/json/list")) + expression = install_expression("safe", "123") + self.assertIn("safe", expression) + self.assertIn("old?.dispose?.()", expression) + self.assertEqual( + notice_ids( + { + "service": 20313, + "payload": json.dumps( + {"notices": [{"notice_id_str": "1", "effect_groups": [960]}]} + ), + } + ), + ["1"], + ) + self.assertEqual( + notice_ids( + { + "service": 20003, + "payload": json.dumps({"notice_type": 45, "notice_id_str": "2"}), + } + ), + ["2"], + ) + + browser = DouyinBrowser() + cast(Any, browser).identity = lambda alias, expected: {"uid": expected} + cast(Any, browser)._evaluate = lambda *args: { + "status": "succeeded", + "action": "followed", + } + preview = browser.action( + "safe", "123", "dm", "456", text="hello", confirm=False + ) + self.assertEqual(preview["action"], "preview") + self.assertEqual( + browser.action("safe", "123", "like_work", work_id="789", confirm=False)[ + "action" + ], + "preview", + ) + with self.assertRaises(DouyinError): + browser.action("safe", "123", "dm", "456", text=" ", confirm=False) + with self.assertRaises(DouyinError): + browser.action("safe", "123", "follow", "bad", confirm=False) + + def test_browser_queue_is_retained_until_ack(self) -> None: + wait = wait_expression("alpha") + self.assertIn("delivered", wait) + self.assertNotIn("splice(0)", wait) + ack = ack_expression("__creatorhub_notice_sub_alpha", ["browser-1"]) + self.assertIn("browser-1", ack) + self.assertIn("__creatorhub_notice_sub_alpha", ack) + + def test_subscription_receipts_are_replayed_until_ack(self) -> None: + subscription = DouyinSubscription.__new__(DouyinSubscription) + subscription.uid = "123" + subscription.queue = deque() + subscription.condition = threading.Condition() + subscription.stopped = threading.Event() + subscription._put({"kind": "notice", "notice": {"event_key": "1"}}) + first = subscription.poll(10, 0) + second = subscription.poll(10, 0) + self.assertEqual(first, second) + self.assertTrue(first[0]["notice"]["gateway_received_at"].endswith("+00:00")) + subscription.ack([first[0]["delivery_id"]]) + self.assertEqual(subscription.poll(10, 0), []) + + def test_subscription_detail_failure_does_not_discard_siblings(self) -> None: + bad = json.dumps( + { + "status": 200, + "body": json.dumps( + { + "status_code": 0, + "notice_list_v2": [{"nid_str": "9", "user_id": "123"}], + } + ), + } + ) + good = json.dumps( + { + "status": 200, + "body": json.dumps( + { + "status_code": 0, + "notice_list_v2": [ + { + "nid_str": "2", + "user_id": "123", + "follow": {"from_user": [{"uid": "7"}]}, + } + ], + } + ), + } + ) + subscription = DouyinSubscription.__new__(DouyinSubscription) + subscription.uid = "123" + subscription.queue = deque() + subscription.condition = threading.Condition() + subscription.stopped = threading.Event() + subscription.connection = cast(CDPConnection, FakeConnection([bad, bad, good])) + subscription._get_connection = lambda: subscription.connection + subscription._handle( + { + "kind": "push", + "service": 20313, + "payload": json.dumps( + { + "notices": [ + {"notice_id_str": "1", "effect_groups": [960]}, + {"notice_id_str": "2", "effect_groups": [960]}, + ] + } + ), + } + ) + events = subscription.poll(10, 0) + self.assertEqual(events[0]["kind"], "error") + self.assertEqual(events[0]["event_key"], "1") + self.assertEqual(events[1]["notice"]["event_key"], "2") + + def test_subscription_manager_and_queue(self) -> None: + event = {"kind": "open"} + subscription = DouyinSubscription.__new__(DouyinSubscription) + subscription.uid = "123" + subscription.queue = deque() + subscription.condition = __import__("threading").Condition() + subscription.stopped = __import__("threading").Event() + subscription._handle(event) + first = subscription.poll(10, 0) + self.assertEqual(first[0]["kind"], "open") + self.assertTrue(first[0]["delivery_id"]) + subscription.queue = deque([{} for _ in range(1000)]) + subscription._put({"kind": "new"}) + self.assertEqual(subscription.queue[0]["kind"], "error") + subscription.stopped.set() + overflow = subscription.poll(10, 0) + self.assertEqual(overflow[0]["kind"], "error") + self.assertEqual(overflow[0]["reason"], "notification queue overflow") + self.assertTrue(overflow[0]["delivery_id"]) + self.assertEqual(len(overflow), 10) + self.assertEqual(len(subscription.queue), 1000) + subscription.ack([overflow[0]["delivery_id"]]) + self.assertEqual(len(subscription.queue), 999) + self.assertEqual(subscription.queue[0], {}) + + browser = DouyinBrowser() + fake = Mock() + fake.poll.return_value = [{"kind": "notice"}] + with patch.object(douyin_module, "DouyinSubscription", return_value=fake): + manager = SubscriptionManager(browser) + self.assertTrue(manager.start("safe", "123")["connected"]) + self.assertEqual(manager.poll("safe", 1, 0), [{"kind": "notice"}]) + manager.stop("safe") + manager.close() + with self.assertRaises(DouyinError): + manager.poll("safe", 1, 0) + + +class ProxyTests(unittest.TestCase): + def test_chunked_request_body_is_decoded_and_forwarded_with_length(self) -> None: + request = ChunkSocket( + [ + ( + b"POST http://example.test/a HTTP/1.1\r\nHost: example.test\r\n" + b"Transfer-Encoding: chunked\r\n\r\n2\r\nab\r\n3;part=x\r\ncde\r\n0\r\nX-Trailer: yes\r\n\r\n" + ) + ] + ) + head, body = _read_request(cast(socket.socket, request)) + self.assertEqual(body, b"abcde") + proxy = MemoryProxy.__new__(MemoryProxy) + proxy.exit = ProxyExit("socks5", "proxy", 1080) + upstream = ChunkSocket([b""]) + proxy.dial = lambda target, timeout=20.0: cast(socket.socket, upstream) + with patch.object(proxy_module, "_copy_until_close"): + proxy.forward_http( + cast(socket.socket, ChunkSocket([])), + "POST", + "http://example.test/a", + _parse_request(head)[2], + body, + ) + sent = upstream.sent[0].decode("iso-8859-1") + self.assertIn("Content-Length: 5", sent) + self.assertNotIn("Transfer-Encoding:", sent) + + def test_registry_generation_and_shutdown(self) -> None: + registry = ProxyRegistry() + url, undo = registry.configure( + "safe", 1, "127.0.0.1", 0, ProxyExit("http", "127.0.0.1", 8080), "network-1" + ) + port = int(url.rsplit(":", 1)[1]) + self.assertFalse(registry.ready("safe", port, "container-1", "network-1")) + self.assertTrue(registry.bind("safe", 1, url, "container-1", "network-1")) + self.assertTrue(registry.ready("safe", port, "container-1", "network-1")) + self.assertFalse(registry.remove("safe", 1, "container-2", "network-1")) + undo() + registry.close() + + def test_memory_proxy_can_close(self) -> None: + proxy = MemoryProxy( + "safe", 1, "127.0.0.1", 0, ProxyExit("http", "127.0.0.1", 8080), "network" + ) + self.assertGreater(proxy.listener.getsockname()[1], 0) + proxy.close() + + def test_proxy_request_parser_and_copy_helpers(self) -> None: + request = ChunkSocket( + [ + b"POST http://example.test/a HTTP/1.1\r\nContent-Length: 3\r\nHost: example.test\r\n\r\nabc" + ] + ) + head, body = _read_request(cast(socket.socket, request)) + self.assertEqual(body, b"abc") + method, target, headers = _parse_request(head) + self.assertEqual((method, target), ("POST", "http://example.test/a")) + self.assertEqual(headers[0], ("Content-Length", "3")) + source = ChunkSocket([b"one", b"", b"ignored"]) + destination = ChunkSocket([]) + _copy_until_close(cast(socket.socket, source), cast(socket.socket, destination)) + self.assertEqual(b"".join(destination.sent), b"one") + self.assertEqual( + _recv_exact(cast(socket.socket, ChunkSocket([b"ab", b"cd"])), 4), b"abcd" + ) + self.assertTrue(_is_ipv4("127.0.0.1")) + self.assertFalse(_is_ipv4("host.example")) + with self.assertRaises(ValueError): + _parse_request(b"BROKEN\r\n\r\n") + + def test_connect_payload_in_same_read_is_preserved(self) -> None: + class GreedySocket(ChunkSocket): + def recv(self, size: int) -> bytes: + del size + return self.chunks.pop(0) if self.chunks else b"" + + upstream = GreedySocket([b"HTTP/1.1 200 OK\r\nX-Test: yes\r\n\r\nTLS"]) + with patch.object(proxy_module, "_open_host", return_value=upstream): + result = _dial_http_proxy( + ProxyExit("http", "proxy", 8080), "target:443", 1.0 + ) + self.assertEqual(_recv_exact(cast(socket.socket, result), 3), b"TLS") + result.close() + + def test_http_and_socks_handshakes(self) -> None: + http_socket = ChunkSocket([b"HTTP/1.1 200 Connection Established\r\n\r\n"]) + with patch.object(proxy_module, "_open_host", return_value=http_socket): + result = _dial_http_proxy( + ProxyExit("http", "proxy", 8080, "u", "c"), "target:443", 1.0 + ) + self.assertIs(result, http_socket) + self.assertIn(b"Proxy-Authorization: Basic dTpj", http_socket.sent[0]) + + socks4_socket = ChunkSocket([b"\x00\x5a\x00\x00\x00\x00\x00\x00"]) + with patch.object(proxy_module, "_open_host", return_value=socks4_socket): + self.assertIs( + _dial_socks4(ProxyExit("socks4", "proxy", 1080), "127.0.0.1:80", 1.0), + socks4_socket, + ) + socks5_socket = ChunkSocket( + [b"\x05\x00", b"\x05\x00\x00\x01\x7f\x00\x00\x01\x00\x50"] + ) + with patch.object(proxy_module, "_open_host", return_value=socks5_socket): + self.assertIs( + _dial_socks5(ProxyExit("socks5", "proxy", 1080), "127.0.0.1:80", 1.0), + socks5_socket, + ) + with self.assertRaises(OSError): + _dial_socks5(ProxyExit("socks5", "proxy", 1080), "host:0", 1.0) + + def test_proxy_forward_and_tunnel_paths(self) -> None: + proxy = MemoryProxy.__new__(MemoryProxy) + proxy.exit = ProxyExit("socks5", "proxy", 1080) + upstream = ChunkSocket([b"reply", b""]) + client = ChunkSocket([]) + cast(Any, proxy).dial = lambda target, timeout=20.0: cast( + socket.socket, upstream + ) + with patch.object(proxy_module, "_copy_until_close") as copy: + proxy.forward_http( + cast(socket.socket, client), + "GET", + "http://example.test/path?q=1", + [("Host", "example.test")], + b"", + ) + copy.assert_called_once() + self.assertIn(b"GET /path?q=1 HTTP/1.1", upstream.sent[0]) + proxy.exit = ProxyExit("http", "proxy", 8080) + upstream = ChunkSocket([b""]) + cast(Any, proxy).dial = lambda target, timeout=20.0: cast( + socket.socket, upstream + ) + with ( + patch.object(proxy_module, "_open_host", return_value=upstream), + patch.object(proxy_module, "_copy_until_close"), + ): + proxy.forward_http( + cast(socket.socket, client), "GET", "http://example.test/", [], b"" + ) + self.assertIn(b"GET http://example.test/ HTTP/1.1", upstream.sent[0]) + + +class ChunkSocket: + def __init__(self, chunks: list[bytes]) -> None: + self.chunks = list(chunks) + self.sent: list[bytes] = [] + + def recv(self, size: int) -> bytes: + if not self.chunks: + return b"" + chunk = self.chunks.pop(0) + if len(chunk) <= size: + return chunk + self.chunks.insert(0, chunk[size:]) + return chunk[:size] + + def sendall(self, data: bytes) -> None: + self.sent.append(data) + + def send(self, data: bytes) -> None: + self.sent.append(data) + + def settimeout(self, value: float | None) -> None: + del value + + def close(self) -> None: + return None + + +class DockerClientTests(unittest.TestCase): + def test_digest_pull_preserves_digest_in_docker_query(self) -> None: + client = ScriptedDocker( + [ + DockerResponse(404, "Not Found", b""), + DockerResponse(200, "OK", b""), + ] + ) + client.pull_if_missing("registry.example/repo@sha256:" + "a" * 64) + self.assertIn( + "fromImage=registry.example%2Frepo%40sha256%3A" + "a" * 64, + client.calls[1][1], + ) + + def test_image_ref_and_expected_statuses(self) -> None: + self.assertEqual( + split_image_ref("registry.example/repo:tag"), + ("registry.example/repo", "tag"), + ) + self.assertEqual(split_image_ref("repo@sha256:abc"), ("repo", "sha256:abc")) + client = ScriptedDocker( + [ + DockerResponse(204, "No Content", b""), + DockerResponse(404, "Not Found", b""), + DockerResponse(500, "Error", b"failure"), + ] + ) + client.expect("POST", "/ok") + with self.assertRaises(FileNotFoundError): + client.expect("POST", "/missing") + with self.assertRaises(DockerError): + client.expect("POST", "/error") + self.assertEqual(tenant_network_name("creatorhub", "safe"), "creatorhub-safe") + + def test_pull_and_managed_container_response_validation(self) -> None: + client = ScriptedDocker( + [ + DockerResponse(404, "Not Found", b""), + DockerResponse(200, "OK", b"{}"), + ] + ) + client.pull_if_missing("repo:tag") + self.assertIn("/images/repo%3Atag/json", client.calls[0][1]) + managed = { + "Id": "container-id", + "Config": {"Labels": {MANAGED_LABEL: "true", RUNTIME_ID_LABEL: "safe"}}, + "NetworkSettings": {"Networks": {"tenant": {"NetworkID": "network-id"}}}, + } + client = ScriptedDocker( + [DockerResponse(200, "OK", json.dumps(managed).encode())] + ) + self.assertEqual(client.managed_container_state("safe")[0], "container-id") + unmanaged = { + **managed, + "Config": {"Labels": {MANAGED_LABEL: "false", RUNTIME_ID_LABEL: "safe"}}, + } + client = ScriptedDocker( + [DockerResponse(200, "OK", json.dumps(unmanaged).encode())] + ) + with self.assertRaises(UnmanagedContainer): + client.managed_container("safe") + + def test_existing_network_generation_and_disconnect(self) -> None: + network = { + "Id": "network-id", + "Name": "creatorhub-safe", + "Driver": "bridge", + "Internal": False, + "Attachable": False, + "Ingress": False, + "Labels": { + MANAGED_LABEL: "true", + NETWORK_ROLE_LABEL: BROWSER_NETWORK_ROLE, + RUNTIME_ID_LABEL: "safe", + BINDING_VERSION_LABEL: "1", + }, + "Containers": { + "gateway-id": {"Name": "gateway", "IPv4Address": "10.0.0.2/24"}, + "runtime-id": {"Name": "runtime", "IPv4Address": "10.0.0.3/24"}, + }, + } + gateway = { + "Id": "gateway-id", + "Config": {"Labels": {GATEWAY_MEMBER_LABEL: "true"}}, + } + client = ScriptedDocker( + [ + DockerResponse(200, "OK", json.dumps(network).encode()), + DockerResponse(200, "OK", json.dumps(gateway).encode()), + ] + ) + generation, addresses, exists = client.inspect_tenant_network( + "creatorhub", "safe", 1, "runtime-id", "gateway", "network-id" + ) + self.assertTrue(exists) + self.assertTrue(generation.runtime_attached) + self.assertEqual(addresses["gateway-id"], "10.0.0.2/24") + self.assertEqual(generation.self_member, "gateway-id") + + after = { + **network, + "Containers": {"gateway-id": network["Containers"]["gateway-id"]}, + } + client = ScriptedDocker( + [ + DockerResponse(200, "OK", json.dumps(network).encode()), + DockerResponse(200, "OK", json.dumps(gateway).encode()), + DockerResponse(200, "OK", json.dumps(network).encode()), + DockerResponse(200, "OK", json.dumps(gateway).encode()), + DockerResponse(200, "OK", b""), + DockerResponse(200, "OK", json.dumps(after).encode()), + DockerResponse(200, "OK", json.dumps(gateway).encode()), + ] + ) + current, _, _ = client.inspect_tenant_network( + "creatorhub", "safe", 1, "runtime-id", "gateway", "network-id" + ) + updated = client.disconnect_member( + "creatorhub", "safe", 1, "runtime-id", current, "runtime-id", "gateway" + ) + self.assertFalse(updated.runtime_attached) + + def test_new_network_is_generation_fenced(self) -> None: + empty = { + "Id": "network-id", + "Name": "creatorhub-safe", + "Driver": "bridge", + "Internal": False, + "Attachable": False, + "Ingress": False, + "Labels": { + MANAGED_LABEL: "true", + NETWORK_ROLE_LABEL: BROWSER_NETWORK_ROLE, + RUNTIME_ID_LABEL: "safe", + BINDING_VERSION_LABEL: "1", + }, + "Containers": {}, + } + connected = { + **empty, + "Containers": { + "runtime-id": {"Name": "runtime", "IPv4Address": "10.0.0.3/24"}, + "gateway-id": {"Name": "gateway", "IPv4Address": "10.0.0.2/24"}, + }, + } + gateway = { + "Id": "gateway-id", + "Config": {"Labels": {GATEWAY_MEMBER_LABEL: "true"}}, + } + client = ScriptedDocker( + [ + DockerResponse(404, "Not Found", b""), + DockerResponse(201, "Created", b'{"Id":"network-id"}'), + DockerResponse(200, "OK", json.dumps(empty).encode()), + DockerResponse(200, "OK", b""), + DockerResponse(200, "OK", b""), + DockerResponse(200, "OK", json.dumps(connected).encode()), + DockerResponse(200, "OK", json.dumps(gateway).encode()), + ] + ) + generation, bind_host = client.ensure_tenant_network( + "creatorhub", "safe", "gateway", 1, "runtime-id" + ) + self.assertTrue(generation.created) + self.assertEqual(bind_host, "10.0.0.2") + + def test_alias_reservation_releases_only_its_generation(self) -> None: + gateway = { + "Id": "gateway-id", + "Image": "creatorhub/gateway:latest", + "Config": {"Labels": {MANAGED_LABEL: "true", GATEWAY_MEMBER_LABEL: "true"}}, + } + reservation = { + "Id": "reservation-id", + "Config": { + "Labels": { + "io.creatorhub.alias-reservation": "true", + "io.creatorhub.reservation-generation": "reservation-generation", + RUNTIME_ID_LABEL: "safe", + } + }, + } + client = ScriptedDocker( + [ + DockerResponse(200, "OK", json.dumps(gateway).encode()), + DockerResponse(201, "Created", b'{"Id":"reservation-id"}'), + DockerResponse(200, "OK", json.dumps(reservation).encode()), + DockerResponse(200, "OK", json.dumps(reservation).encode()), + DockerResponse(204, "No Content", b""), + DockerResponse(404, "Not Found", b""), + ] + ) + with patch.object( + docker_client_module, + "random_reservation_generation", + return_value="reservation-generation", + ): + release = AliasReservationManager(client, "gateway").acquire("safe") + release() + self.assertTrue(any(call[0] == "DELETE" for call in client.calls)) + + def test_stale_alias_reservation_is_reclaimed(self) -> None: + gateway = { + "Id": "gateway-id", + "Image": "creatorhub/gateway:latest", + "Config": {"Labels": {MANAGED_LABEL: "true", GATEWAY_MEMBER_LABEL: "true"}}, + } + stale = { + "Id": "stale-reservation-id", + "Config": { + "Labels": { + RESERVATION_LABEL: "true", + RUNTIME_ID_LABEL: "safe", + RESERVATION_OWNER_LABEL: "old-gateway", + } + }, + } + current = { + "Id": "current-reservation-id", + "Config": { + "Labels": { + RESERVATION_LABEL: "true", + RUNTIME_ID_LABEL: "safe", + RESERVATION_OWNER_LABEL: "gateway", + "io.creatorhub.reservation-generation": "current-generation", + } + }, + } + client = ScriptedDocker( + [ + DockerResponse(200, "OK", json.dumps(gateway).encode()), + DockerResponse(409, "Conflict", b""), + DockerResponse(200, "OK", json.dumps(stale).encode()), + DockerResponse(404, "Not Found", b""), + DockerResponse(204, "No Content", b""), + DockerResponse(404, "Not Found", b""), + DockerResponse(201, "Created", b'{"Id":"current-reservation-id"}'), + DockerResponse(200, "OK", json.dumps(current).encode()), + DockerResponse(200, "OK", json.dumps(current).encode()), + DockerResponse(204, "No Content", b""), + DockerResponse(404, "Not Found", b""), + ] + ) + with patch.object( + docker_client_module, + "random_reservation_generation", + return_value="current-generation", + ): + release = AliasReservationManager(client, "gateway").acquire("safe") + release() + self.assertEqual( + client.calls[4][1], + "/containers/stale-reservation-id?force=1&v=0", + ) + self.assertEqual(client.calls[6][0], "POST") + + +class ScriptedDocker(DockerClient): + def __init__(self, responses: list[DockerResponse]) -> None: + super().__init__("/dev/null") + self.responses = list(responses) + self.calls: list[tuple[str, str, object | None]] = [] + + def request( + self, + method: str, + path: str, + payload: object | None = None, + timeout: float = 30.0, + body_limit: int = 16 * 1024 * 1024, + ) -> DockerResponse: + del timeout, body_limit + self.calls.append((method, path, payload)) + if not self.responses: + raise AssertionError(f"unexpected Docker call: {method} {path}") + return self.responses.pop(0) + + +class ReservationStub: + def __init__(self) -> None: + self.released = 0 + + def acquire(self, alias: str): + del alias + + def release() -> None: + self.released += 1 + + return release + + +class LockingReservationStub(ReservationStub): + def __init__(self) -> None: + super().__init__() + self.lock = threading.Lock() + + def acquire(self, alias: str): + del alias + self.lock.acquire() + + def release() -> None: + self.released += 1 + self.lock.release() + + return release + + +class LifecycleDocker(ScriptedDocker): + def ensure_tenant_network( + self, *args: object, **kwargs: object + ) -> tuple[TenantNetworkGeneration, str]: + del args, kwargs + return TenantNetworkGeneration( + id="network-id", self_member="gateway-id" + ), "10.0.0.2" + + def inspect_tenant_network( + self, *args: object, **kwargs: object + ) -> tuple[TenantNetworkGeneration, dict[str, str], bool]: + del args, kwargs + return ( + TenantNetworkGeneration( + id="network-id", + name="creatorhub-safe", + self_member="gateway-id", + gateway_members=["gateway-id"], + runtime_attached=True, + ), + {"gateway-id": "10.0.0.2/24"}, + True, + ) + + +class GatewayLifecycleTests(unittest.TestCase): + def test_server_tracks_daemon_request_threads_and_timeout(self) -> None: + server = gateway_module.GatewayHTTPServer(("127.0.0.1", 0), Mock()) + try: + self.assertTrue(server.daemon_threads) + self.assertEqual(gateway_module.GatewayHandler.protocol_version, "HTTP/1.1") + finally: + server.server_close() + + def test_network_setup_error_keeps_created_generation(self) -> None: + client = ScriptedDocker( + [ + DockerResponse(404, "Not Found", b""), + DockerResponse(201, "Created", b'{"Id":"network-id"}'), + DockerResponse(500, "Error", b"inspect failed"), + ] + ) + with self.assertRaises(docker_client_module.NetworkSetupError) as caught: + client.ensure_tenant_network("creatorhub", "safe", "gateway", 1, "runtime") + self.assertEqual(caught.exception.generation.id, "network-id") + self.assertTrue(caught.exception.generation.created) + + def test_network_cleanup_failure_returns_pending_contract(self) -> None: + docker = ScriptedDocker( + [ + DockerResponse(404, "Not Found", b""), + DockerResponse(500, "Error", b"temporary"), + ] + ) + gateway = Gateway( + cast(DockerClient, docker), "creatorhub", "0123456789abcdef", "gateway" + ) + gateway.reservations = cast(AliasReservationManager, ReservationStub()) + with self.assertRaises(RequestError) as caught: + gateway.remove( + "safe", + { + "binding_version": 1, + "runtime_id": "runtime-not-found", + "network_id": "network-id", + }, + ) + self.assertEqual(caught.exception.status, 202) + self.assertEqual(caught.exception.network_id, "network-id") + self.assertEqual(str(caught.exception), "runtime_cleanup_pending") + + def test_timed_out_action_retains_alias_ownership(self) -> None: + gateway = Gateway( + cast(DockerClient, Mock()), "creatorhub", "0123456789abcdef", "gateway" + ) + gateway._claim_action("safe") + gateway._retain_action_ownership("safe") + with self.assertRaises(RequestError): + gateway._claim_action("safe") + + def _input(self, stopped: bool = True) -> dict: + return { + "alias": "safe", + "name": "Safe", + "image": "creatorhub/browser:latest", + "cmd": ["about:blank"], + "volume": "creatorhub-safe", + "binding_version": 1, + "network_exit_id": "", + "stopped": stopped, + } + + def test_concurrent_create_cannot_delete_the_winner(self) -> None: + class ConcurrentDocker(DockerClient): + def __init__(self) -> None: + super().__init__("/dev/null") + self.created = False + self.calls: list[tuple[str, str]] = [] + + def managed_container(self, alias: str) -> tuple[str, dict[str, str]]: + if not self.created: + raise FileNotFoundError(alias) + return "a" * 64, { + MANAGED_LABEL: "true", + RUNTIME_ID_LABEL: alias, + BINDING_VERSION_LABEL: "1", + NETWORK_ID_LABEL: "", + } + + def request( + self, + method: str, + path: str, + payload: object | None = None, + timeout: float = 30.0, + body_limit: int = 16 * 1024 * 1024, + ) -> DockerResponse: + del payload, timeout, body_limit + self.calls.append((method, path)) + if method == "GET" and path.startswith("/images/"): + return DockerResponse(200, "OK", b"{}") + if method == "POST" and path.startswith("/containers/create"): + if self.created: + return DockerResponse(409, "Conflict", b"") + self.created = True + return DockerResponse( + 201, "Created", b'{"Id":"' + b"b" * 64 + b'"}' + ) + raise AssertionError(f"unexpected Docker call: {method} {path}") + + docker = ConcurrentDocker() + gateway = Gateway( + cast(DockerClient, docker), "creatorhub", "0123456789abcdef", "gateway" + ) + gateway.reservations = cast(AliasReservationManager, LockingReservationStub()) + results: list[object] = [] + + def create() -> None: + try: + results.append(gateway.create(self._input())) + except ( + AssertionError, + DockerError, + OSError, + RequestError, + ValueError, + ) as exc: + results.append(exc) + + threads = [threading.Thread(target=create) for _ in range(2)] + for thread in threads: + thread.start() + for thread in threads: + thread.join() + self.assertEqual(sum(isinstance(item, dict) for item in results), 1) + self.assertEqual(sum(isinstance(item, RequestError) for item in results), 1) + self.assertFalse(any(method == "DELETE" for method, _ in docker.calls)) + + def test_create_stopped_pulls_image_and_keeps_container(self) -> None: + docker = ScriptedDocker( + [ + DockerResponse(200, "OK", b"{}"), + DockerResponse(404, "Not Found", b""), + DockerResponse(201, "Created", b'{"Id":"container-id"}'), + ] + ) + gateway = Gateway( + cast(DockerClient, docker), "creatorhub", "0123456789abcdef", "gateway" + ) + reservations = ReservationStub() + gateway.reservations = cast(AliasReservationManager, reservations) + result = gateway.create(self._input()) + self.assertEqual(result["id"], "container-id") + self.assertEqual(reservations.released, 1) + self.assertEqual(docker.calls[2][0], "POST") + + def test_running_create_starts_container_and_network(self) -> None: + docker = LifecycleDocker( + [ + DockerResponse(200, "OK", b"{}"), + DockerResponse(404, "Not Found", b""), + DockerResponse(201, "Created", b'{"Id":"container-id"}'), + DockerResponse(204, "No Content", b""), + ] + ) + gateway = Gateway( + cast(DockerClient, docker), "creatorhub", "0123456789abcdef", "gateway" + ) + gateway.reservations = cast(AliasReservationManager, ReservationStub()) + result = gateway.create(self._input(False)) + self.assertEqual(result["network_id"], "network-id") + self.assertEqual(docker.calls[-1][0], "POST") + + def test_change_state_and_remove_are_generation_fenced(self) -> None: + labels = { + MANAGED_LABEL: "true", + RUNTIME_ID_LABEL: "safe", + BINDING_VERSION_LABEL: "1", + NETWORK_ID_LABEL: "network-id", + } + inspected = { + "Id": "a" * 64, + "Config": {"Labels": labels}, + "NetworkSettings": {"Networks": {}}, + } + docker = ScriptedDocker( + [ + DockerResponse(200, "OK", json.dumps(inspected).encode()), + DockerResponse(204, "No Content", b""), + ] + ) + gateway = Gateway( + cast(DockerClient, docker), "creatorhub", "0123456789abcdef", "gateway" + ) + gateway.reservations = cast(AliasReservationManager, ReservationStub()) + gateway.change_state( + "safe", + "stop", + { + "binding_version": 1, + "runtime_id": "a" * 64, + "network_id": "network-id", + }, + ) + self.assertEqual(docker.calls[-1][0], "POST") + + docker = LifecycleDocker( + [ + DockerResponse(200, "OK", json.dumps(inspected).encode()), + DockerResponse(204, "No Content", b""), + ] + ) + gateway = Gateway( + cast(DockerClient, docker), "creatorhub", "0123456789abcdef", "gateway" + ) + gateway.reservations = cast(AliasReservationManager, ReservationStub()) + gateway._remove_network = lambda *args, **kwargs: None + gateway.remove( + "safe", + { + "binding_version": 1, + "runtime_id": "a" * 64, + "network_id": "network-id", + }, + ) + self.assertEqual(docker.calls[-1][0], "DELETE") + + def test_remove_stopped_direct_container_without_network(self) -> None: + labels = { + MANAGED_LABEL: "true", + RUNTIME_ID_LABEL: "safe", + BINDING_VERSION_LABEL: "1", + NETWORK_ID_LABEL: "", + } + inspected = { + "Id": "a" * 64, + "Config": {"Labels": labels}, + "NetworkSettings": {"Networks": {}}, + } + docker = ScriptedDocker( + [ + DockerResponse(200, "OK", json.dumps(inspected).encode()), + DockerResponse(404, "Not Found", b""), + DockerResponse(204, "No Content", b""), + ] + ) + gateway = Gateway( + cast(DockerClient, docker), "creatorhub", "0123456789abcdef", "gateway" + ) + gateway.reservations = cast(AliasReservationManager, ReservationStub()) + gateway.remove( + "safe", + {"binding_version": 1, "runtime_id": "a" * 64, "network_id": ""}, + ) + self.assertEqual(docker.calls[-1][0], "DELETE") + + def test_remove_rejects_runtime_replacement(self) -> None: + labels = { + MANAGED_LABEL: "true", + RUNTIME_ID_LABEL: "safe", + BINDING_VERSION_LABEL: "1", + NETWORK_ID_LABEL: "", + } + inspected = { + "Id": "a" * 64, + "Config": {"Labels": labels}, + "NetworkSettings": {"Networks": {}}, + } + docker = ScriptedDocker( + [DockerResponse(200, "OK", json.dumps(inspected).encode())] + ) + gateway = Gateway( + cast(DockerClient, docker), "creatorhub", "0123456789abcdef", "gateway" + ) + gateway.reservations = cast(AliasReservationManager, ReservationStub()) + with self.assertRaises(RequestError): + gateway.remove( + "safe", + {"binding_version": 1, "runtime_id": "b" * 64, "network_id": ""}, + ) + + def test_create_failure_reconciles_unknown_container(self) -> None: + docker = ScriptedDocker( + [ + DockerResponse(200, "OK", b"{}"), + DockerResponse(404, "Not Found", b""), + DockerResponse(500, "Error", b"failed"), + DockerResponse(404, "Not Found", b""), + ] + ) + gateway = Gateway( + cast(DockerClient, docker), "creatorhub", "0123456789abcdef", "gateway" + ) + gateway.reservations = cast(AliasReservationManager, ReservationStub()) + with self.assertRaises(RequestError): + gateway.create(self._input()) + + +class AdditionalGatewayCoverageTests(unittest.TestCase): + def test_douyin_connect_rejects_bad_discovery(self) -> None: + browser = DouyinBrowser(lambda _: "https://browser:9222") + with self.assertRaises(DouyinError): + browser._connect("safe") + cases = [ + (500, b"{}"), + (200, b"{}"), + (200, json.dumps([{"type": "service"}]).encode()), + ( + 200, + json.dumps( + [{"type": "page", "url": "https://www.douyin.com/1"}] + ).encode(), + ), + ( + 200, + json.dumps( + [ + { + "type": "page", + "url": "https://www.douyin.com/1", + "webSocketDebuggerUrl": "http://browser/devtools/page/1", + } + ] + ).encode(), + ), + ] + for status, body in cases: + with self.subTest(status=status, body=body): + http = FakeHTTPConnection(FakeHTTPResponse(status, body)) + with ( + patch.object( + douyin_module.http.client, "HTTPConnection", return_value=http + ), + self.assertRaises(DouyinError), + ): + browser._connect("safe") + targets = [ + { + "type": "page", + "url": "https://www.douyin.com/1", + "webSocketDebuggerUrl": "ws://browser:9222/devtools/page/1", + }, + { + "type": "page", + "url": "https://www.douyin.com/2", + "webSocketDebuggerUrl": "ws://browser:9222/devtools/page/2", + }, + ] + http = FakeHTTPConnection(FakeHTTPResponse(200, json.dumps(targets).encode())) + with ( + patch.object( + douyin_module.http.client, "HTTPConnection", return_value=http + ), + self.assertRaises(DouyinError), + ): + browser._connect("safe") + + def test_douyin_fetch_identity_and_confirmed_actions(self) -> None: + browser = DouyinBrowser() + + def bind(connection: BrowserCDP) -> None: + BrowserTests()._with_connection(browser, connection) + + for result in ( + {"too_large": True}, + {"status": 302, "body": ""}, + {"status": 200, "body": 1}, + ): + bind(BrowserCDP(["https://www.douyin.com", result])) + with self.assertRaises(DouyinError): + browser.get("safe", "https://www.douyin.com/a") + cast(Any, browser).get = lambda alias, target: BrowserResponse(200, "not-json") + with self.assertRaises(DouyinError): + browser.identity("safe") + cast(Any, browser).get = lambda alias, target: BrowserResponse( + 200, json.dumps({"status_code": 0, "user": {"uid": "1", "sec_uid": "sec"}}) + ) + self.assertEqual(browser.identity("safe")["uid"], "1") + cast(Any, browser).get = lambda alias, target: BrowserResponse( + 403, json.dumps({"status_code": 0, "user": {"uid": "1", "sec_uid": "sec"}}) + ) + with self.assertRaises(DouyinError): + browser.identity("safe") + cast(Any, browser).identity = lambda alias, expected_uid=None: { + "uid": expected_uid or "1" + } + cast(Any, browser)._evaluate = lambda alias, expression: { + "status": 200, + "action": "sent", + } + bind(BrowserCDP(["https://www.douyin.com", {}])) + self.assertEqual( + browser.action("safe", "1", "follow", "2", confirm=True)["action"], "sent" + ) + bind(BrowserCDP(["https://www.douyin.com", {}])) + self.assertEqual( + browser.action("safe", "1", "dm", "2", text="hello", confirm=True)[ + "status" + ], + 200, + ) + cast(Any, browser)._evaluate = lambda alias, expression: "bad" + bind(BrowserCDP(["https://www.douyin.com", "bad"])) + with self.assertRaises(DouyinError): + browser.action("safe", "1", "follow", "2", confirm=True) + + def test_gateway_and_proxy_validation_edges(self) -> None: + for value in ( + None, + [], + [{"name": "bad"}], + [{"name": "sessionid", "value": "x", "domain": "evil.test"}], + ): + self.assertFalse(valid_cookies(value)) + self.assertEqual( + parse_proxy_exit({"protocol": "http", "host": "", "port": 80}).host, "" + ) + with self.assertRaises(RequestError): + validate_proxy_exit( + parse_proxy_exit({"protocol": "http", "host": "", "port": 80}) + ) + with self.assertRaises(RequestError): + validate_proxy_exit( + parse_proxy_exit({"protocol": "http", "host": "proxy", "port": 0}) + ) + with self.assertRaises(RequestError): + validate_proxy_restore( + { + "binding_version": 1, + "runtime_id": "a" * 64, + "network_id": "n", + "network_exit_id": "x", + }, + "bad alias", + ) + self.assertFalse(valid_account_key_query({"key": ["bad key"]}, "key")) + + def test_proxy_rejected_handshakes_and_docker_errors(self) -> None: + bad_http = ChunkSocket([b"HTTP/1.1 407 Proxy Authentication Required\r\n\r\n"]) + with ( + patch.object(proxy_module, "_open_host", return_value=bad_http), + self.assertRaises(OSError), + ): + _dial_http_proxy(ProxyExit("http", "proxy", 8080), "target:443", 1.0) + bad_socks5 = ChunkSocket([b"\x05\x02"]) + with ( + patch.object(proxy_module, "_open_host", return_value=bad_socks5), + self.assertRaises(OSError), + ): + _dial_socks5( + ProxyExit("socks5", "proxy", 1080, "u", "p"), "127.0.0.1:80", 1.0 + ) + client = DockerClient("/not/a/socket") + with ( + patch.object( + docker_client_module, "UnixHTTPConnection", side_effect=OSError("down") + ), + self.assertRaises(DockerError), + ): + client.request("GET", "/version", timeout=0.01) + with ( + patch.object( + client, "request", return_value=DockerResponse(500, "bad", b"x") + ), + self.assertRaises(DockerError), + ): + client.expect("POST", "/x") + with ( + patch.object( + client, "request", return_value=DockerResponse(404, "missing", b"") + ), + self.assertRaises(FileNotFoundError), + ): + client.expect("DELETE", "/gone") + + def test_docker_client_transport_and_metadata_edges(self) -> None: + client = DockerClient("/unused") + response = FakeHTTPResponse(200, b"{}") + cast(Any, response).reason = "OK" + connection = FakeHTTPConnection(response) + with patch.object( + docker_client_module, "UnixHTTPConnection", return_value=connection + ): + result = client.request("POST", "/version", {"ok": True}) + self.assertEqual(result.status, 200) + self.assertEqual(connection.requested, [("POST", "/v1.43/version")]) + huge = FakeHTTPResponse(200, b"12345") + cast(Any, huge).reason = "OK" + with ( + patch.object( + docker_client_module, + "UnixHTTPConnection", + return_value=FakeHTTPConnection(huge), + ), + self.assertRaises(DockerError), + ): + client.request("GET", "/version", body_limit=4) + labels = { + MANAGED_LABEL: "true", + RUNTIME_ID_LABEL: "safe", + GATEWAY_MEMBER_LABEL: "true", + } + inspected = { + "Id": "a" * 64, + "Config": {"Labels": labels}, + "NetworkSettings": { + "Networks": {"n": {"NetworkID": "network", "IPAddress": "198.51.100.5"}} + }, + } + with patch.object( + client, + "request", + return_value=DockerResponse(200, "OK", json.dumps(inspected).encode()), + ): + self.assertTrue(client.trusted_gateway_member("gateway")) + self.assertEqual( + client.container_network_address("a" * 64, "network"), "198.51.100.5" + ) + bad = DockerResponse(200, "OK", b"[]") + with patch.object(client, "request", return_value=bad): + self.assertFalse(client.trusted_gateway_member("gateway")) + with self.assertRaises(DockerError): + client.managed_container_state("safe") + + def test_client_connection_pull_and_inspect_edges(self) -> None: + class ConnectedSocket: + def __init__(self) -> None: + self.timeout = None + self.path = "" + + def settimeout(self, value: float) -> None: + self.timeout = value + + def connect(self, path: str) -> None: + self.path = path + + connected = ConnectedSocket() + with patch.object( + docker_client_module.socket, "socket", return_value=connected + ): + connection = docker_client_module.UnixHTTPConnection( + "/run/docker.sock", 1.0 + ) + connection.connect() + self.assertEqual(connected.path, "/run/docker.sock") + client = DockerClient("/unused") + with patch.object( + client, + "request", + side_effect=[ + DockerResponse(404, "missing", b""), + DockerResponse(200, "OK", b"{}"), + ], + ) as request: + client.pull_if_missing("registry.example/repo:tag") + self.assertIn("tag=tag", request.call_args_list[1].args[1]) + with ( + patch.object( + client, "request", return_value=DockerResponse(500, "bad", b"no") + ), + self.assertRaises(DockerError), + ): + client.pull_if_missing("registry.example/repo:tag") + invalid = DockerResponse( + 200, + "OK", + json.dumps( + { + "Id": "a" * 64, + "Config": {"Labels": []}, + "NetworkSettings": {"Networks": {}}, + } + ).encode(), + ) + with ( + patch.object(client, "request", return_value=invalid), + self.assertRaises(DockerError), + ): + client.managed_container_state("safe") + + def test_cdp_error_and_proxy_auth_paths(self) -> None: + socket_ = FakeSocket([{"id": 1, "result": {"result": {"value": {"ok": True}}}}]) + connection = CDPConnection(cast(websocket.WebSocket, socket_)) + self.assertEqual(connection.evaluate("1"), {"ok": True}) + socket_ = FakeSocket([{"id": 1, "result": {"result": {}}}]) + with self.assertRaises(DouyinError): + CDPConnection(cast(websocket.WebSocket, socket_)).evaluate("1") + with self.assertRaises(DouyinError): + CDPConnection( + cast(websocket.WebSocket, FakeSocket([{"id": 1, "error": {}}])) + ).command("Page.enable") + with self.assertRaises(DouyinError): + CDPConnection( + cast(websocket.WebSocket, FakeSocket([{"id": 1, "result": []}])) + ).command("Page.enable") + with self.assertRaises(DouyinError): + CDPConnection(cast(websocket.WebSocket, FakeSocket([]))).wait_event( + "Page.loadEventFired", lambda _: True, timeout=0.01 + ) + socket_ = FakeSocket([{"id": 1, "result": {"exceptionDetails": {}}}]) + with self.assertRaises(DouyinError): + CDPConnection(cast(websocket.WebSocket, socket_)).evaluate("1") + target = [ + { + "type": "page", + "url": "https://www.douyin.com/1", + "webSocketDebuggerUrl": "ws://browser:9222/devtools/page/1", + } + ] + http = FakeHTTPConnection(FakeHTTPResponse(200, json.dumps(target).encode())) + with ( + patch.object( + douyin_module.http.client, "HTTPConnection", return_value=http + ), + patch.object( + douyin_module.websocket, + "create_connection", + side_effect=OSError("down"), + ), + self.assertRaises(DouyinError), + ): + DouyinBrowser(lambda _: "http://browser:9222")._connect("safe") + socks5 = ChunkSocket( + [b"\x05\x02", b"\x01\x00", b"\x05\x00\x00\x01\x7f\x00\x00\x01\x00\x50"] + ) + with patch.object(proxy_module, "_open_host", return_value=socks5): + self.assertIs( + _dial_socks5( + ProxyExit("socks5", "proxy", 1080, "u", "p"), "127.0.0.1:80", 1.0 + ), + socks5, + ) + + def test_direct_message_notice_and_uncertain_post_contract(self) -> None: + notice = normalize_notice( + { + "dm": { + "message_id": "123456", + "from_user": {"uid": "456789"}, + "text": "hello", + }, + "create_time": 1700000000, + } + ) + self.assertIsNotNone(notice) + assert notice is not None + self.assertEqual(notice["event_type"], "dm") + self.assertEqual(notice["interactor_uid"], "456789") + self.assertEqual(notice["message_text"], "hello") + script = action_expression({"alias": "safe", "action": "follow", "target": "2"}) + self.assertIn("POST_UNCERTAIN", script) + self.assertIn("BUSINESS_REJECTED", script) + + def test_cdp_timeout_terminates_page_evaluation(self) -> None: + class TimeoutSocket(FakeSocket): + def __init__(self) -> None: + super().__init__([]) + self.receives = 0 + + def recv(self) -> str: + self.receives += 1 + if self.receives == 1: + raise TimeoutError("deadline") + return json.dumps({"id": 2, "result": {}}) + + socket_ = TimeoutSocket() + connection = CDPConnection(cast(websocket.WebSocket, socket_)) + with ( + patch.object(douyin_module, "CONTROL_TIMEOUT", 0.001), + self.assertRaises(DouyinError), + ): + connection.command("Runtime.evaluate") + sent = [json.loads(item) for item in socket_.sent] + self.assertEqual(sent[1]["method"], "Runtime.terminateExecution") + + def test_network_create_reconciliation_preserves_observed_generation(self) -> None: + client = DockerClient("/unused") + observed = TenantNetworkGeneration(name="creatorhub-safe", id="observed") + generation = TenantNetworkGeneration(name="creatorhub-safe") + with ( + patch.object( + client, "inspect_tenant_network", return_value=(observed, [], True) + ), + self.assertRaises(NetworkSetupError) as caught, + ): + client._finish_network_create( + "creatorhub", + "safe", + "gateway", + 1, + "runtime", + DockerResponse(201, "Created", b"{}"), + generation, + ) + self.assertEqual(caught.exception.generation.id, "observed") + + with patch.object( + client, "inspect_tenant_network", return_value=(observed, [], True) + ): + result = client._finish_network_create( + "creatorhub", + "safe", + "gateway", + 1, + "runtime", + DockerResponse(201, "Created", b'{"Id":"created"}'), + TenantNetworkGeneration(name="creatorhub-safe"), + ) + self.assertEqual(result.id, "observed") + + +if __name__ == "__main__": + unittest.main() diff --git a/compose.yaml b/compose.yaml index b632e2d..2c6ba3d 100644 --- a/compose.yaml +++ b/compose.yaml @@ -5,12 +5,18 @@ services: DATABASE_URL: postgres://creatorhub@postgres/creatorhub?sslmode=disable CONTROL_PLANE_USERNAME: ${CONTROL_PLANE_USERNAME:?required} CONTROL_PLANE_PASSWORD: ${CONTROL_PLANE_PASSWORD:?required} - CREATORHUB_CREDENTIAL_MASTER_KEY: ${CREATORHUB_CREDENTIAL_MASTER_KEY:?required} + CREATORHUB_CREDENTIAL_MASTER_KEY: >- + ${CREATORHUB_CREDENTIAL_MASTER_KEY:?required} + BAILIAN_API_KEY: ${BAILIAN_API_KEY:-} + BAILIAN_BASE_URL: ${BAILIAN_BASE_URL:-} + CREATOR_MEDIA_DIR: /var/lib/creatorhub/materials + CREATOR_TRANSCRIPTION_BIN: ${CREATOR_TRANSCRIPTION_BIN:-} ports: - "${CREATORHUB_PORT:-8080}:8080" read_only: true volumes: - creatorhub_credentials:/var/lib/creatorhub/credentials + - creatorhub_materials:/var/lib/creatorhub/materials tmpfs: - /tmp:size=16m,noexec,nosuid,nodev cap_drop: [ALL] @@ -24,7 +30,8 @@ services: restart: unless-stopped postgres: - image: postgres:17-alpine@sha256:18cfe3ef5e6815560c98237d6216d1e5119702fb0f3894c8785dd58b8bbe5d73 + image: >- + postgres:17-alpine@sha256:18cfe3ef5e6815560c98237d6216d1e5119702fb0f3894c8785dd58b8bbe5d73 environment: POSTGRES_DB: creatorhub POSTGRES_USER: creatorhub @@ -48,7 +55,8 @@ services: docker-gateway: build: . - command: ["/app/docker-gateway"] + stop_grace_period: 45s + command: ["python3", "-m", "cmd.docker_gateway.gateway"] labels: io.creatorhub.gateway-member: "true" environment: @@ -57,7 +65,7 @@ services: volumes: - /var/run/docker.sock:/var/run/docker.sock:ro group_add: - - "${DOCKER_GID:-999}" + - "${DOCKER_GID:?required}" healthcheck: test: [CMD, wget, -q, -O, /dev/null, http://127.0.0.1:8081/healthz] interval: 2s @@ -77,3 +85,4 @@ networks: volumes: creatorhub_postgres: creatorhub_credentials: + creatorhub_materials: diff --git a/docker/browser-wrapper/Dockerfile b/docker/browser-wrapper/Dockerfile new file mode 100644 index 0000000..37c46c3 --- /dev/null +++ b/docker/browser-wrapper/Dockerfile @@ -0,0 +1,5 @@ +ARG BROWSER_BASE_IMAGE +FROM ${BROWSER_BASE_IMAGE} + +COPY --chmod=0755 docker-entrypoint.sh /usr/local/bin/creatorhub-browser-entrypoint.sh +ENTRYPOINT ["/usr/local/bin/creatorhub-browser-entrypoint.sh"] diff --git a/docker/browser-wrapper/README.md b/docker/browser-wrapper/README.md new file mode 100644 index 0000000..9f281a4 --- /dev/null +++ b/docker/browser-wrapper/README.md @@ -0,0 +1,12 @@ +# CreatorHub browser wrapper + +这是浏览器运行时的受管理入口,不包含 Chromium 本体。构建时必须传入已登记的 immutable base image digest,禁止使用 tag 或 `latest`: + +```bash +docker build \ + --build-arg BROWSER_BASE_IMAGE=git.ipao.vip/rogee/fingerprint-chromium@sha256: \ + -t git.ipao.vip/rogee/creatorhub-browser-wrapper@sha256: \ + docker/browser-wrapper +``` + +发布记录必须同时保存 base image 仓库、base digest、构建提交、此目录 Dockerfile 和发布 digest。入口会等待 Xvfb、检查 x11vnc、绑定容器地址上的 CDP 端口,再启动 `/opt/chromium/chrome`。 diff --git a/docker/browser-wrapper/docker-entrypoint.sh b/docker/browser-wrapper/docker-entrypoint.sh new file mode 100644 index 0000000..0aea049 --- /dev/null +++ b/docker/browser-wrapper/docker-entrypoint.sh @@ -0,0 +1,32 @@ +#!/bin/sh +set -eu + +Xvfb :99 -screen 0 "${SCREEN_SIZE:-1920x1080x24}" -nolisten tcp -ac & +xvfb_pid=$! +i=0 +while [ "$i" -lt 300 ]; do + [ -S /tmp/.X11-unix/X99 ] && break + kill -0 "$xvfb_pid" 2>/dev/null || exit 1 + i=$((i + 1)) + sleep 0.1 +done +[ -S /tmp/.X11-unix/X99 ] || { + echo "Xvfb did not become ready" >&2 + exit 1 +} +export DISPLAY=:99 +x11vnc -display :99 -listen "$(hostname -i)" -rfbport 5900 -forever -nopw -nevershared -dontdisconnect -noclipboard -nosetclipboard -noprimary -nosetprimary -quiet & +rfb_pid=$! +kill -0 "$rfb_pid" 2>/dev/null || { + echo "RFB service did not start" >&2 + exit 1 +} +remote_debugging_port=${REMOTE_DEBUGGING_PORT:-9222} +case "$remote_debugging_port" in +'' | *[!0-9]*) + echo "REMOTE_DEBUGGING_PORT must be numeric" >&2 + exit 2 + ;; +esac +socat "TCP-LISTEN:${remote_debugging_port},fork,reuseaddr,bind=$(hostname -i)" "TCP:127.0.0.1:${remote_debugging_port}" & +exec /opt/chromium/chrome --disable-dev-shm-usage --no-sandbox --no-default-browser-check --no-first-run --remote-debugging-port="$remote_debugging_port" --user-data-dir=/data "$@" diff --git a/docs/architecture/container-control.md b/docs/architecture/container-control.md index 0b87407..d19bea9 100644 --- a/docs/architecture/container-control.md +++ b/docs/architecture/container-control.md @@ -15,7 +15,8 @@ React ── /api/* (Basic Auth) ──> control-plane ├─ /v1/browsers (Bearer token) ──> docker-gateway ──> docker.sock │ └─> browser container - └─ 账号/环境/任务等持久记录 ──> PostgreSQL + ├─ /v1/browsers/.../douyin/events (Bearer token) ──> docker-gateway ──> browser event stream + └─ 账号/环境/任务/互动事件持久记录 ──> PostgreSQL ``` 控制面是唯一事实源:网关不持有镜像清单和业务规则,镜像引用、启动命令和卷名均随请求下发。 @@ -66,7 +67,7 @@ React ── /api/* (Basic Auth) ──> control-plane 草稿经 `POST /api/phase-a/confirmations` 显式确认后才可投递到 `/api/phase-a/tasks`。任务由幂等键去重;`POST /api/phase-a/mock/execute` 使用 `FOR UPDATE SKIP LOCKED` 领取一分钟租约,执行前统一核对账号、草稿和确认版本。缺少确认或版本不一致会进入 `needs_confirmation`,暂停账号或 Mock 策略结果会进入 `policy_hold`,不确定结果与过期租约进入 `needs_confirmation`;这些状态都不会自动重试。`GET /api/phase-a/audit` 只导出账号、确认版本、尝试和结果等非秘密证据。 -启动时控制面先应用 Phase A v1,再由 Hub runner 顺序应用 v2 至 v14;每一步都在事务和 advisory lock 下前向执行。v3 保留旧表、列和历史记录,旧账号回填为 `platform=mock` 并暂停,仅账号 ID 与环境 alias 相同的记录自动建立 binding;v4 追加环境动作审计字段与索引,v5 清理持久 fingerprint 中的旧代理字段,v6 增加可重试的 runtime cleanup 状态,v7 为 runtime lease 增加 binding version 并回填可确定的既有记录,v8 至 v10 补齐 cleanup/runtime 的不可变 generation 与兼容约束,v11、v12 增加任务恢复状态并修复兼容约束,v13 增加账号名称和 TAGS;v14 仅前向修复旧 PR v13 的空数据 schema。若旧 v13 已产生空引用或明文 Cookies,v14 会在删除前阻断启动,必须先将凭据迁入 provider。其余记录等待显式绑定。本阶段不提供破坏性自动回滚。 +启动时控制面先应用 Phase A v1,再由 Hub runner 顺序应用 v2 至 v14;CreatorHub 业务迁移继续顺序应用至 v26(竞品同步 lease token、事件消息正文);每一步都在事务和 advisory lock 下前向执行。v3 保留旧表、列和历史记录,旧账号回填为 `platform=mock` 并暂停,仅账号 ID 与环境 alias 相同的记录自动建立 binding;v4 追加环境动作审计字段与索引,v5 清理持久 fingerprint 中的旧代理字段,v6 增加可重试的 runtime cleanup 状态,v7 为 runtime lease 增加 binding version 并回填可确定的既有记录,v8 至 v10 补齐 cleanup/runtime 的不可变 generation 与兼容约束,v11、v12 增加任务恢复状态并修复兼容约束,v13 增加账号名称和 TAGS;v14 仅前向修复旧 PR v13 的空数据 schema。若旧 v13 已产生空引用或明文 Cookies,v14 会在删除前阻断启动,必须先将凭据迁入 provider。其余记录等待显式绑定。本阶段不提供破坏性自动回滚。 `POST /api/network-exits` 只接受协议、主机、端口、已有 `credential_reference: {id}` 和预期出口身份;新出口为 `unchecked`,由 `POST /api/network-exits/:id/check` 经实际代理链路变为 `healthy` 或 `unhealthy`,`disable` 不可被检查重新启用。credential reference 的 `reference_key` 不出现在 API、日志或审计中;OS Keyring/Secret Manager bridge 在控制面进程启动前注入 `CREATORHUB_CREDENTIAL_`(大写十六进制),值为请求期解析的 `username:password`,控制面不持久化解析值。 @@ -74,12 +75,12 @@ React ── /api/* (Basic Auth) ──> control-plane 解析后的出口凭据只存在于控制面单次请求和网关内存转发器中;Docker inspect、容器环境、标签、挂载、`Config.Cmd` 与进程参数只包含 `docker-gateway` 的无凭据本地代理地址。网关内存代理以 alias、binding version 和 exit ID 共同标识 generation;生命周期操作按 alias 串行,重启恢复或重建必须重新核对该 generation,旧出口代理不能被新容器复用。 -控制面后台每 20 秒调和网关([runtimeLeaseHeartbeat](../../cmd/control-plane/main.go)),续租 running runtime、释放 stopped/missing runtime;列表/详情读取不触发调和,过期 lease 也会在绑定事务中回收。控制面用 PostgreSQL advisory transaction lock 按 alias 协调多副本;每个 Store 最多允许 5 个锁会话占用 10 连接池的一半,为锁内数据库调用保留连接。create 同时锁定账号 ID、alias、请求出口和请求镜像;start、reconcile/rebuild、rebind 和 upgrade 锁定 alias、当前出口及当前镜像(upgrade 还锁目标镜像),拿锁后重新读取出口与镜像版本。账号 pause/resume/revoke 使用账号 ID 与当前 binding alias 加入同一协调域;镜像禁用、引用更新或账号状态变更不能穿透在途生命周期。 +控制面后台每 20 秒调和网关([runtimeLeaseHeartbeat](../../cmd/control-plane/main.go)),续租 running runtime、释放 stopped/missing runtime;列表/详情读取不触发调和,过期 lease 也会在绑定事务中回收。控制面还按账号调和抖音事件监听([creator_events.go](../../cmd/control-plane/creator_events.go)):每个已授权且有有效运行代际的账号只有一个监听,网关断连或代际改变时停止旧监听并退避重连;事件通知在控制面转换为互动事件后进入 `ProcessAutomaticEvent`,数据库的事件唯一键负责去重,写入结果不明不会由监听器重复发送。当前通知边界、平台事件游标/基线连续性和真实写操作仍需真机证据,不能把网关轮询队列视为平台监听验收。控制面用 PostgreSQL advisory transaction lock 按 alias 协调多副本;每个 Store 最多允许 5 个锁会话占用 10 连接池的一半,为锁内数据库调用保留连接。create 同时锁定账号 ID、alias、请求出口和请求镜像;start、reconcile/rebuild、rebind 和 upgrade 锁定 alias、当前出口及当前镜像(upgrade 还锁目标镜像),拿锁后重新读取出口与镜像版本。账号 pause/resume/revoke 使用账号 ID 与当前 binding alias 加入同一协调域;镜像禁用、引用更新或账号状态变更不能穿透在途生命周期。 非法 upgrade/rebind 目标在进入 advisory lock key 前按公开格式校验;审计仅保留环境原有的非秘密资源关联,并以 `upgrade_input_rejected` / `rebind_input_rejected` 写同一 operation ID 的 requested/finished 对。reconcile 恢复或重建后会重新读取 context,finished 事件关联实际激活的 runtime instance、binding version 与出口;后台释放 runtime 的成功或失败也写独立的 `reconcile` 审计对。 ## 与新业务计划的边界 -- 当前受限抖音读取仅允许自身身份与 `count=20/max_cursor=0` 首批作品,JSON 响应限制为 1 MiB,见 [douyin.go](../../cmd/docker-gateway/douyin.go)。这些是现状,不是竞品分页或媒体下载的完成证据;后续媒体按 plan01 C4 分步保存产物引用,不复用通用 JSON 响应承载二进制。 -- 当前账号凭据入口处理 Cookie;plan01 要求的可选登录密码应允许输入现有凭据保存流程,读取不回显、不进入日志,并非新增“禁止收密码”的接口限制。 +- 当前受限抖音读取仅允许自身身份与 `count=20/max_cursor=0` 首批作品,JSON 响应限制为 1 MiB,见 [douyin.py](../../cmd/docker_gateway/douyin.py)。竞品同步仍需外部平台证据;媒体通过 `POST /api/creator/works/:id/material/process` 下载到持久卷、用 FFmpeg 提取音轨,再调用显式配置的 `CREATOR_TRANSCRIPTION_BIN`。未配置转写入口时明确记为 failed,不接受手写 succeeded 或二进制塞入通用 JSON。 +- 当前账号凭据入口处理 Cookie;登录密码可按现有凭据保存流程保存,读取不回显、不进入日志,但绝不由系统自动注入或用于绕过人工登录。 - 新业务平台监听、前端业务推送和现有后台运行租约是三件事;前两者要求见 plan01 A6,列表不主动探测的约定不禁止业务事件推送。现有生命周期/租约可能核验出口,不应误写成已完成 plan01 的手动代理管理目标。 diff --git a/docs/deployment.md b/docs/deployment.md index 09eaa71..df3e315 100644 --- a/docs/deployment.md +++ b/docs/deployment.md @@ -1,6 +1,6 @@ # CreatorHub 部署 -本文档按 `main@1fbf126` 当前可运行代码说明单台 Linux 主机 Docker Compose 部署及旧阶段 A Mock 检查,不证明新产品功能已实现。[plan01](plan01.md) 是业务范围与验收依据,旧阶段 A 离线结果不等于 G0/G1,也不能替代抖音/小红书最终真机验收。 +本文档按 `feat/python-gateway-douyin` 当前代码说明单台 Linux 主机 Docker Compose 部署及离线检查,不证明新产品功能已实现。[plan01](plan01.md) 是业务范围与验收依据;离线结果不能替代抖音/小红书最终真机验收。 当前控制面仍使用单用户 HTTP Basic Auth(除 `/healthz` 外,包括静态页面),不提供 RBAC 或多租户隔离。开发目标不新增认证/访问限制,但本轮未删除现有代码或配置;以下变量仍须填写,不新增认证 profile。 @@ -24,7 +24,7 @@ - 当前用户可访问 Docker daemon; - 可访问镜像仓库(如 `git.ipao.vip`),并已完成登录(如仓库要求认证);镜像也可不在宿主机预拉取,网关会在缺失时按引用自动拉取; - `curl`、`jq`、`openssl` 和 GNU `stat`,用于启动与业务验证命令; -- 若验证真实浏览器运行环境:Linux x86_64、一个可访问的固定 HTTP/HTTPS/SOCKS 出口,以及可被 Docker daemon 拉取的 `fingerprint-chromium` 镜像。 +- 若验证真实浏览器运行环境:Linux x86_64、一个可访问的固定 HTTP/HTTPS/SOCKS 出口,以及可被 Docker daemon 拉取的、以 immutable digest 固定的 `fingerprint-chromium` 镜像。镜像来源必须登记仓库地址、构建提交和 digest,禁止使用 `latest` 或未登记 tag;Xvfb/CDP 包装入口的可复现源码见 [`docker/browser-wrapper`](../docker/browser-wrapper/)。 在仓库根目录执行预检: @@ -32,6 +32,7 @@ test -S /var/run/docker.sock docker info >/dev/null docker compose version +: "${DOCKER_GID:?export DOCKER_GID=$(stat -c '%g' /var/run/docker.sock) in this shell}" : "${CONTROL_PLANE_USERNAME:?export CONTROL_PLANE_USERNAME in this shell}" : "${CONTROL_PLANE_PASSWORD:?export CONTROL_PLANE_PASSWORD in this shell}" : "${CREATORHUB_CREDENTIAL_MASTER_KEY:?export CREATORHUB_CREDENTIAL_MASTER_KEY in this shell}" @@ -54,6 +55,8 @@ docker compose config --quiet docker compose up --detach --build ``` +网关配置了 45 秒停止宽限期,并在收到 SIGTERM/SIGINT 时停止接收请求、排空有限期限内的在途请求,再关闭事件订阅和代理;本地测试覆盖该顺序,但部署验证仍须记录正常退出码、实际耗时和没有 SIGKILL,不能仅以“等待了 45 秒”证明优雅退出。 + 控制面启动时会连接 PostgreSQL,并在事务和 advisory lock 保护下自动执行前向迁移。迁移失败时控制面会退出,由 Compose 按 `restart: unless-stopped` 重启;先检查日志,不要删除数据卷。 ### 首次配置 @@ -61,7 +64,7 @@ docker compose up --detach --build 服务起来后打开 : 1. 「网关管理」页注册网关:名称如 `gw-main`,Endpoint `http://docker-gateway:8081`,令牌填 `GATEWAY_TOKEN` 的值(即 `openssl rand -hex 24` 生成的值)。 -2. 「镜像版本」页添加可用镜像,如版本 `148.0.7778.215`、引用 `git.ipao.vip/rogee/fingerprint-chromium:148.0.7778.215`(或 `@sha256:` 摘要引用)。 +2. 按 [`docker/browser-wrapper/README.md`](../docker/browser-wrapper/README.md) 用已登记 base digest 构建包装镜像,再在「镜像版本」页添加发布 digest。生产与验收必须填写 `git.ipao.vip/rogee/creatorhub-browser-wrapper@sha256:<已登记摘要>`;`148.0.7778.215` 只作为浏览器版本元数据,不能单独作为不可变引用。登记内容同时包含 base 镜像源仓库、base digest、构建提交、Dockerfile 路径和发布摘要。 3. 「社媒账号」页先创建账号(默认暂停),再在「运行环境」选该账号、网关、镜像,填写中文名、小写别名及指纹参数。代理可选;指定代理须先手动检测为健康,留空是明确直连,不是失败回退。 4. 创建得到停止态容器;在账号页恢复账号后回环境页显式启动。容器名 `creatorhub-browser-<别名>`,Profile 卷 `creatorhub-profile-<别名>`。回收只删除容器、保留环境/binding/Profile;完整契约见[架构说明](architecture/container-control.md)。当前直连 create/start 不代表 upgrade/rebind 已支持空出口。 @@ -133,7 +136,11 @@ api() { run_id="manual-$(date +%s)" gateway_name="gw-${run_id}" image_version="${IMAGE_VERSION:-148.0.7778.215}" -image_ref="${IMAGE_REF:-git.ipao.vip/rogee/fingerprint-chromium:${image_version}}" +: "${IMAGE_REF:?IMAGE_REF must be an immutable browser-wrapper @sha256 reference}" +case "$IMAGE_REF" in + *@sha256:*) image_ref="$IMAGE_REF" ;; + *) echo "IMAGE_REF must contain @sha256:" >&2; exit 2 ;; +esac exit_protocol="${PROXY_PROTOCOL:-http}" account_key="${run_id}-platform" env_alias="${run_id}-env" @@ -252,11 +259,15 @@ Compose 部署时通常只需设置以下宿主机变量: | 变量 | 默认值 | 说明 | | --- | --- | --- | | `CREATORHUB_PORT` | `8080` | 控制面宿主机端口,局域网可访问 | -| `DOCKER_GID` | `999` | Docker socket 的宿主机组 ID;必须按实际值设置 | +| `DOCKER_GID` | 无(必填) | Docker socket 的宿主机组 ID;必须按实际值设置 | | `GATEWAY_TOKEN` | `dev-creatorhub-gateway-token` | Compose 未提供变量时的默认值;本说明要求显式 export 随机值,并同步填入网关注册表单 | | `CONTROL_PLANE_USERNAME` | 无(必填) | 控制面唯一用户;不能包含冒号 | | `CONTROL_PLANE_PASSWORD` | 无(必填) | 控制面密码,至少 6 字节;使用随机值 | | `CREATORHUB_CREDENTIAL_MASTER_KEY` | 无(必填) | 32 字节 base64;由部署 Secret Manager/OS Keyring 持久注入,重启后必须保持一致 | +| `BAILIAN_API_KEY` | 空(按需) | 已批准的生产 AI API key;为空时文本 AI 动作明确返回 unavailable,不使用 Mock | +| `BAILIAN_BASE_URL` | 客户端默认值 | 无用户信息的 HTTP(S) 地址;供应商变更前需完成审批和脱敏样本验证 | +| `CREATOR_MEDIA_DIR` | `/var/lib/creatorhub/materials` | 持久媒体目录;必须挂载持久卷,不能使用临时目录替代 | +| `CREATOR_TRANSCRIPTION_BIN` | 空(按需) | 可执行的本地转写入口;为空时转写明确失败,不伪造 succeeded | 服务本身支持并校验以下环境变量;`compose.yaml` 会在控制面凭据缺失或为空时拒绝渲染。下表中的 Compose 默认值不由手工验证脚本隐式读取;手工验证沿用上文的显式 export 要求: @@ -270,6 +281,8 @@ Compose 部署时通常只需设置以下宿主机变量: | `creator-hub` | `CONTROL_PLANE_PASSWORD` | 必填且至少 6 字节;不会写入日志或响应 | | `creator-hub` | `CREATORHUB_CREDENTIAL_MASTER_KEY` | 必填;解密独立凭据卷,不写入数据库、日志或响应 | | `creator-hub` | `CREATORHUB_CREDENTIAL_STORE_DIR` | 默认 `/var/lib/creatorhub/credentials`;必须为绝对路径且持久可写 | +| `creator-hub` | `BAILIAN_API_KEY` | 可选;仅用于已批准的文本 AI,缺失时不回退 Mock | +| `creator-hub` | `BAILIAN_BASE_URL` | 可选;HTTP(S) 供应商地址,禁止携带用户信息 | | `docker-gateway` | `LISTEN_ADDR` | 默认 `:8081` | | `docker-gateway` | `DOCKER_SOCKET` | 默认值和 Compose 挂载均固定为 `/var/run/docker.sock`;不能只覆盖环境变量 | | `docker-gateway` | `BROWSER_NETWORK` | `creatorhub_browser` | diff --git a/docs/plan01.md b/docs/plan01.md index 08d00d4..ee959db 100644 --- a/docs/plan01.md +++ b/docs/plan01.md @@ -2,7 +2,7 @@ > 版本:v1.2 · 2026-09-12 > 状态:业务范围已逐项确认;本文是实现与验收依据,不代表功能已经完成。 -> 现状基线:`feat/plan01-implementation` 基于 `main@1fbf126`,已同步远程;本轮完成一批本地 G1.1/G1.2 实现与离线验证,未完成真实平台验收。 +> 现状基线:`feat/python-gateway-douyin` 基于 `origin/main@1fbf126`;当前分支已包含 Python gateway、事件边界/确认、动作结果证据、固定采集计划和工作台修正,并完成本地离线验证;真实平台验收仍单独记录。 > 优先级:本需求高于现有实现及旧产品规划;遇到平台能力不足或新的业务歧义,必须向使用者确认,不得自行删减、替换或假装成功。 ## 1. 目标、范围与完成定义 @@ -42,11 +42,11 @@ | 范围 | 已有基础 | 本需求需要补齐或改变 | 代码依据 | | --- | --- | --- | --- | | 页面入口 | 账号、任务、审计、环境、代理、镜像、网关等页面 | 竞品池、作品、素材仿写、评论线索、响应策略、私信会话尚无完整入口 | [页面路由](../web/src/main.jsx) | -| 账号 | 平台账号 ID、标签、Cookie、暂停/恢复等 | 实名资料、登录用户名/密码、备注、人工业务状态、自动登录与大小号关系 | [账号页面](../web/src/AccountsPage.jsx)、[账号存储](../internal/phasea/store.go)、[账号接口](../cmd/control-plane/phasea.go) | +| 账号 | 平台账号 ID、标签、Cookie、暂停/恢复等 | 实名资料、登录用户名/密码、备注、人工业务状态、人工登录与大小号关系 | [账号页面](../web/src/AccountsPage.jsx)、[账号存储](../internal/phasea/store.go)、[账号接口](../cmd/control-plane/phasea.go) | | 平台范围 | 登记项还包含公众号、快手 | 新业务仅承诺抖音、小红书;其他登记项不是业务能力验收结果,也不因此要求删除无关已有功能 | [账号页面](../web/src/AccountsPage.jsx) | -| 作品读取 | 有抖音连接器、指标字段与测试 | 连接器只核验登录者自身、读取首批 20 条,未接成运行中的竞品采集链路;需目标账号解析、分页、保存、查询和定时更新 | [抖音连接器](../internal/douyin/connector.go)、[连接器测试](../internal/douyin/connector_test.go)、[网关读取限制](../cmd/docker-gateway/douyin.go) | +| 作品读取 | 有抖音连接器、指标字段与测试 | 连接器只核验登录者自身、读取首批 20 条,未接成运行中的竞品采集链路;需目标账号解析、分页、保存、查询和定时更新 | [抖音连接器](../internal/douyin/connector.go)、[连接器测试](../internal/douyin/connector_test.go)、[网关读取限制](../cmd/docker_gateway/douyin.py) | | 指纹与环境 | 结构化指纹、独立 Profile、启动/停止/升级及代理接入 | 默认固定 seed 不是自动分配;需首次自动生成、地区匹配及稳定性验证 | [环境页面](../web/src/BrowsersPage.jsx)、[指纹参数](../internal/hub/fingerprint.go)、[环境存储](../internal/hub/environment.go) | -| 代理 | 添加、列表、手动检测、停用、实际转发 | 补齐编辑、删除、重新启用及引用约束;不增加自动轮换 | [代理页面](../web/src/NetworkExitsPage.jsx)、[代理转发](../cmd/docker-gateway/proxy.go) | +| 代理 | 添加、列表、手动检测、停用、实际转发 | 补齐编辑、删除、重新启用及引用约束;不增加自动轮换 | [代理页面](../web/src/NetworkExitsPage.jsx)、[代理转发](../cmd/docker_gateway/proxy.py) | | 任务与发送 | 有草稿确认、任务记录、Mock 执行 | 不能作为真实回复、私信、点赞、关注、转发成功的证据;需接通实际平台执行及结果核验 | [任务页面](../web/src/TasksPage.jsx)、[任务接口](../cmd/control-plane/phasea.go) | | 事件、线索、私信、AI | 未发现完整可运行链路 | 均需新增业务能力,不能把旧文档、类型定义或模拟返回记为已实现 | [页面路由](../web/src/main.jsx)、[抖音连接器](../internal/douyin/connector.go) | @@ -134,15 +134,15 @@ | 登录情况 | 独立展示已登录、需登录、登录失败/待人工处理等实际结果与时间,不推断处罚状态 | - 支持账号资料新增、查看、编辑;必填平台及能区分账号的标识,已获得平台 UID 后检查身份重复和登录一致性,不通过改昵称掩盖冲突。 -- 登录密码允许由新增/编辑请求输入并进入现有凭据保存流程,账号资料仅关联凭据;读取只显示是否已配置,编辑空白不覆盖已有密码。当前接口只接收 Cookie,不代表密码能力已实现,也不能以“API 禁止接收密码”阻断本需求。 +- 登录密码允许由新增/编辑请求输入并进入现有凭据保存流程,账号资料仅关联凭据;读取只显示是否已配置,编辑空白不覆盖已有密码。密码仅供人工核对/维护,不由系统自动注入或登录。 - 实名资料只做输入类型、长度等基本校验;不把人工记录包装成已核验身份,不要求用户为运行普通功能提交身份证证据。 - 正常账号按实际登录及平台能力执行。**非正常账号暂停自动动作**;禁言账号禁止人工评论、私信及带文案转发;封禁、注销账号不参与采集和发送。已有数据仍可查看,人工登录/检查入口保留。 - 账号恢复正常后需人工重新启用受影响策略,不补发暂停期间的历史互动。登录失败不擅自修改上述业务状态。 ### A2. 登录 -- 已存在有效登录时复用;未登录且提供用户名与密码时尝试自动登录,否则直接进入可见浏览器人工登录。 -- 自动登录失败显示真实原因:例如凭据错误、平台不支持该登录方式、验证码/扫码/设备确认、网络失败或页面变化;未知原因明确写“无法确认”,不能猜测。 +- 已存在有效登录时复用;未登录或会话失效时进入可见浏览器人工登录。系统不得自动注入用户名、密码、Cookie、验证码或绕过扫码/设备确认。 +- 人工登录失败显示真实原因:例如凭据错误、验证码/扫码/设备确认、网络失败或页面变化;未知原因明确写“无法确认”,不能猜测。 - 失败后保留同一环境交用户人工处理,不循环尝试、不自动更换账号/代理/指纹,也不绕过验证码。 - 登录后核对实际平台账号身份;与选定账号不一致时停止该账号任务并提示纠正,不能把别的账号登录结果保存为成功。 @@ -326,7 +326,7 @@ | 编号 | 场景 | 必须观察到的结果 | | --- | --- | --- | | AC-A1 | 新增/修改资料,区分账号 UID、用户名、实名与业务状态 | 字段分别保存与显示;不伪称平台已实名;凭据不进入日志 | -| AC-A2 | 有/无密码、有效登录、密码错误、验证码、身份不符 | 对应复用/自动尝试/人工登录;错误有原因,不循环尝试、不误认别的账号 | +| AC-A2 | 有/无密码、有效登录、密码错误、验证码、身份不符 | 有效会话复用;失效后进入人工登录;错误有原因,不循环尝试、不误认别的账号 | | AC-A3 | 正常→禁言/封禁/注销→正常 | 执行限制符合 A1;已有数据可查;恢复后不擅自重启策略或补发历史互动 | | AC-A4 | 自关联、跨平台、重复归属、大小号循环、调整关系 | 非法关系被拒绝;变更停用受影响策略,不产生错误账号动作 | | AC-A5 | 四类大号收到的互动 | 评论、点赞、转发、关注均用真实事件证明;大号主动操作不误触发;无 UID/目标明确提示 | @@ -423,7 +423,7 @@ 后续代码变更遵循 [AGENTS.md](../AGENTS.md): - 非平凡行为先写能在未实现时失败的最小回归测试,单元测试覆盖率至少 65%;尤其覆盖时间边界、并发冷却、监听重复、平台失败及人工/自动区别。 -- 后端通过 `go test ./...`、`go vet ./...`,并构建 `./cmd/control-plane`、`./cmd/docker-gateway`;并发、生命周期与共享状态变更运行 `go test -race ./...`;Compose 变更运行 `docker compose config --quiet`。 +- 控制面通过 `go test ./...`、`go vet ./...`,并构建 `./cmd/control-plane`;并发、生命周期与共享状态变更运行 `go test -race ./...`。Python gateway 通过 `python3 -m unittest discover -s cmd -p 'test_*.py'` 和覆盖率检查;Compose 变更运行 `docker compose config --quiet`。 - 前端从 lockfile 安装、通过非交互测试及 `npm --prefix web run build`;关键页面覆盖加载、空数据、错误、禁用、确认/取消和重复点击。 - 离线测试只证明本地逻辑,真实平台必须另外验收;请求返回 200、任务入队、模拟连接器成功均不能代替平台动作成功。 @@ -431,7 +431,7 @@ - 原始四模块需求均有对应细化、边界及验收项;逐项确认的业务选择无遗漏、无相反表述。 - 区分现状、目标、已确认规则与待真实验证的外部条件,不将建议数值或 AI 质量写成未经确认的承诺。 -- 相对链接有效、验收编号唯一、Markdown 检查通过;没有功能代码改动,原有未提交修改保留。 +- 相对链接有效、验收编号唯一、Markdown 检查通过;代码修正与本文件的现状描述同步。 ### 9.2 后续必须确认而不得猜测的事项 @@ -451,11 +451,11 @@ | 无音轨/无语音必须重新决定;API 应禁止接收密码 | C4、A1、AC-U2、AC-U3 | 评审误判:已确认可继续;密码允许输入凭据流程,禁止回显/日志泄露,不阻断录入 | | 页面不够可执行,缺跳转与异常状态 | 7.4、8.5 | 紧凑页面地图与共用交互,优先现有页签/抽屉,无新通用系统 | | 去重/冷却、基线/断连、人工重试与并发边界不足 | A3、A5、A6、W3、7.2、AC-B1 至 AC-B3 | 可靠身份、永久最小去重、迟到核验、持久操作标识及同账号写协调 | +| 时间/文本匹配不确定,媒体读取限制误当需求缺陷 | C2、W2、C4、AC-B4 至 AC-B6 | 固定 UTC 时长、原文匹配、分步产物引用;真实分页/媒体另验,不预设扩容或自动删素材 | +| AI 输入未冻结、平台能力无证据模板、G1 过大 | 8.6、8.7、第 9 节 | G0 保留外部批准前置,矩阵全为待验证;G1 分段可验但不替代最终完整流程 | ### 9.4 当前本地实现状态 -- 已有:CreatorHub 数据表与 API、账号/大小号关系约束、规则与线索判断、作品/一级评论分页采集边界、断点租约、来源关联、指数指标计划、抖音自有/竞品只读采集调度,以及工作台的错误/重试/确认状态。 -- 下一轮 TODO:真实平台事件监听与断连恢复;抖音实际点赞/评论/关注/转发/私信写入;素材真实下载/提音/转写;小红书采集与动作;生产 AI/转写供应商、模型、密钥与质量样本;以及对应的 G0/G1 真机证据。 -- 本地限制:仓库默认未配置 `CREATORHUB_POSTGRES_TEST_URL`;本轮使用运行中的 PostgreSQL 完成了 CreatorHub 数据行为、来源关联、指标、操作、消息和 checkpoint 租约集成测试,覆盖率为 65.6%。CI 或后续环境需提供专用测试库以重复该验证。 -| 时间/文本匹配不确定,媒体读取限制误当需求缺陷 | C2、W2、C4、AC-B4 至 AC-B6 | 固定 UTC 时长、原文匹配、分步产物引用;真实分页/媒体另验,不预设扩容或自动删素材 | -| AI 输入未冻结、平台能力无证据模板、G1 过大 | 8.6、8.7、第 9 节 | G0 保留外部批准前置,矩阵全为待验证;G1 分段可验但不替代最终完整流程 | +- 已有:CreatorHub 数据表与 API、账号/大小号关系约束、规则与线索判断、作品/一级评论分页采集边界、断点租约、来源关联、固定指标计划、抖音自有/竞品只读采集调度、Python gateway 事件 ACK/恢复边界、写操作证据分类、媒体下载/FFmpeg 提音轨/显式转写命令链,以及工作台的错误/确认状态。 +- 尚未完成或尚未取得外部证据:转写供应商质量样本、小红书同范围动作、真实平台事件连续性及每类写操作成功证据。未提供配置或平台能力时,接口应明确返回 unavailable/uncertain/failed,不以 Mock 通过代替。 +- 本地限制:数据库集成测试必须提供专用 PostgreSQL 地址并保留实际 executed/skipped 记录;Python gateway 生产覆盖率必须按排除测试代码的同次报告验收。CI 或后续环境需重复这些验证。 diff --git a/docs/python-gateway-branch-review.md b/docs/python-gateway-branch-review.md new file mode 100644 index 0000000..3e792f8 --- /dev/null +++ b/docs/python-gateway-branch-review.md @@ -0,0 +1,477 @@ +# Python 网关分支审查与修正清单 + +日期:2026-09-13。状态:**审查与代码修正已完成;本地门禁已通过,真实平台及供应商证据仍按外部验收项单独保留。** + +## 1. 范围、结论与证据边界 + +- 审查分支:`feat/python-gateway-douyin`。 +- 基线:`origin/main@1fbf126b21a7d6de45ade70121785075542363d9`。 +- 分支提交:`860ecaf7ead08bbd1af2c0d8af3604270da78095`,另包含当前未提交及未跟踪文件。不是只审查最近一次提交或 dirty diff。 +- 审查前已获取远程信息;远程主线是当前分支祖先,没有待合入的主线提交。当前分支没有配置 upstream。 +- 范围包含 58 个已跟踪差异文件,以及 Python 网关、事件监听、依赖锁文件等 11 个未跟踪新增文件;本报告是另行新增的交付物。 +- 四个独立只读审查分别覆盖网关迁移、平台动作与事件、数据业务与采集、前端与文档;汇总时合并了重复发现,并复核关键调用链。 +- **本轮修正已覆盖网关、事件、动作、采集计划、账号一致性和前端确定缺陷,并同步了部署/计划文档;未重置数据库、未提交 Git,也未进行真实平台写操作。** 下列“已修正”条目以当前代码与本地回归证据为准;真实平台、供应商和浏览器视觉证据仍不能由本地测试替代。 + +### 1.1 总结 + +最紧急的问题不是缺少几个页面,而是已有路径会产生错误结果: + +1. **没有可靠的新旧事件边界**,此前未记录的旧通知可能在重连后触发自动发送。 +2. **并发创建可能误删成功容器**,超时的浏览器脚本也可能在锁释放后继续写入。 +3. **代理环境不能正常走通**:Go 漏传代理标识,HTTP/HTTPS 上游代理握手又存在协议解析错误。 +4. **常用页面不可用或串数据**:设置保存返回 400,工作台切换可崩溃,账号迟返请求可能覆盖另一账号资料。 +5. **人工发送链路不完整**:内部 ID 未转平台 ID、刷新后丢失原操作标识、失败与不明分类及证据不足。 +6. **采集计划、账号状态、租约和凭据更新存在数据一致性问题**。 +7. **媒体、生产 AI、私信接收/历史及页面推送仍有实际接线缺口**,不是仅等待真机验收。 + +### 1.2 优先级及状态口径 + +| 标记 | 含义 | +| --- | --- | +| P0 | 可能造成未获批准的真实平台写入;相应自动能力在修正并验证前不可视为可用 | +| P1 | 阻断核心功能,或可能误删、串号、丢事件、错误判定;合并前应解决 | +| P2 | 局部正确性、边界、错误反馈或可维护性问题;仍须列入修正,不默默排除 | +| 确定缺陷 | 由当前代码及调用链可推导触发场景;除特别说明外,本轮没有运行复现 | +| 实现缺口 | 缺少实际执行或界面接线,不能用模型、状态表或 Mock 结果替代 | +| 待真机验证 | 代码之外的平台行为和运行证据不足,不据此断言平台支持或不支持 | + +GW01—GW11、EV01—EV14、BE01—BE19、UI01—UI09 和 T01 的代码修正已完成并有本地回归;T02、DOC01—DOC06 已按当前证据修正。GAP01—GAP10 已补齐可由本地代码完成的执行、持久化和界面闭环,但真实平台、供应商质量、浏览器视觉及小红书能力仍需外部证据;GAP11 仍明确阻断。P0 的安全门仍有效:边界不明或结果不明不得自动写入或重试。 + +## 2. 网关生命周期与迁移缺陷 + +### GW01 · P1:并发创建可能删除另一个请求成功创建的容器 + +- **证据与影响:**`cmd/docker_gateway/gateway.py:170-177,241-252,322-366`。两个相同 alias、binding_version 的停止态创建请求都在锁外查到不存在;第二个获得锁后未重查,Docker create 返回 409,清理函数在 `created_id` 为空时按相同标签删除第一个请求的容器。持久卷不会因此删除,但已成功创建的容器会消失。旧 Go 实现在锁内再次检查 alias。 +- **修正:**锁内重查;明确的创建冲突不能进入“本次创建结果未知”的删除路径。所有清理必须证明资源属于本次操作。 +- **验收:**线程屏障制造该竞争;仅一个创建成功,另一个 409,成功容器没有收到 DELETE。 + +### GW02 · P1:HTTP/HTTPS 上游代理 CONNECT 残留响应头破坏 TLS + +- **证据与影响:**`cmd/docker_gateway/proxy.py:366-380,477-491,511-523`。`_read_status` 只读状态行,剩余响应头和空行随后作为隧道数据转发给浏览器。上游返回 200 不代表 HTTPS 可以使用。旧 Go 使用完整 HTTP 响应解析。 +- **修正:**完整解析并消费响应头,限制大小和时间,保留响应头之后已经到达的真实隧道数据。优先采用成熟解析能力,不继续增加不完整手写分支。 +- **验收:**无附加头、多头、分片头、头与首个隧道包同批到达;浏览器收到的首字节必须是真实隧道数据。现有测试只检查 socket 和认证头,不足以证明转发正确。 + +### GW03 · P1:停止可被空闲连接无限阻塞,45 秒宽限不等于正常退出 + +- **证据与影响:**`cmd/docker_gateway/gateway.py:874-886,1018-1037,1427-1446`、`compose.yaml:52`。HTTP/1.1 非 daemon 请求线程没有空闲/读取超时;`server.shutdown()` 不关闭已有连接,`server_close()` 会等待这些线程。控制面正常的 `http.DefaultClient` 就可能留下 keep-alive 连接。半截请求体也可造成阻塞。清理顺序还先关闭订阅和代理,再等待在途请求。 +- **修正:**明确停止接收、回收空闲连接、在途请求期限和依赖关闭顺序。不能仅改 daemon 线程来掩盖未完成清理。 +- **验收:**隔离进程覆盖空闲连接、半截请求体、活跃请求、排队请求;记录正常退出码、耗时、未发生 SIGKILL 和 reservation 释放。此前 47 秒重启记录不能证明优雅退出,详见第 6 节。 + +### GW04 · P1:冷启动拉取镜像时丢失 digest + +- **证据与影响:**`cmd/docker_gateway/docker_client.py:142-160,584-586`。`repo@sha256:…` 被拆开后,只将 repository 传给 `fromImage`,digest 未传出,可能拉取无关默认 tag,随后原 digest 创建失败。 +- **修正:**digest 引用完整传给 Docker;普通 tag 引用才按相应契约拆分。 +- **验收:**直接断言 Docker 请求 query,覆盖缺本地 digest、已存在 digest 和普通 tag,不仅测试拆字符串函数。 + +### GW05 · P1:删除校验拒绝控制面的 `runtime-not-found` 清理语义 + +- **证据与影响:**`cmd/docker_gateway/gateway.py:1180-1205`;`cmd/control-plane/hub.go:71,140-145,196-225,1789-1794`。创建失败后只剩已知网络,控制面用 `runtime_id="runtime-not-found"` 清理;Python 对非空 ID 强制 64 位十六进制,提前返回 400,残留网络不能收敛。 +- **修正:**保留现有 `/v1` 契约时,仅删除路径解释“容器不存在、清理指定网络代”;若实际存在替代容器则必须拒绝。启动和写操作仍须真实容器 ID,不放宽通用检查。 +- **验收:**Go 实际生成的 payload 经过 Python HTTP 路由,覆盖仅网络、两者皆无、alias 已被新代占用。 + +### GW06 · P1:网络初始化失败丢失代信息,部分副作用未回滚 + +- **证据与影响:**`cmd/docker_gateway/docker_client.py:393-475`;`gateway.py:190-192,316-333,474-485,929-947`。创建网络后 inspect 失败不一定携带已知 generation;调用方尚未完成赋值。`restore_proxy` 的网络初始化又位于回滚 try 之外。连接部分成功后抛错可留下半完成网络,HTTP 错误也未保留清理所需网络 ID。 +- **修正:**首次副作用后立即保留已知代;所有后续错误经过统一回滚边界,返回已知网络代和清理结论。 +- **验收:**逐点注入创建后 inspect 失败、runtime 已连接但 gateway 连接失败、操作生效但响应丢失;仅清理本次代,错误记录保留恢复依据。 + +### GW07 · P1:单个浏览器网络异常使整个 gateway 无法列举和对账 + +- **证据与影响:**`cmd/docker_gateway/gateway.py:144-147`、`docker_client.py:231-252`;`cmd/control-plane/hub.go:927-970`。列表为每个 running 容器同步解析网络 IP,任何一个缺网络/IP 或中途消失就使列表整体失败,控制面跳过整个 gateway 的环境对账。 +- **修正:**清单读取与 CDP 地址解析分离;异常环境仍应可识别、可单独恢复,不能返回假地址或假成功。 +- **验收:**正常、缺网络、中途消失三个环境共存;正常记录可读,异常记录有真实状态,其他环境仍能对账。 + +### GW08 · P1:CDP 超时释放锁后,页面脚本仍可能继续写入 + +- **证据与影响:**`cmd/docker_gateway/douyin.py:34,64-79,136-155,915-924`、`gateway.py:617-639`。CDP 固定 15 秒,脚本含 20 秒写请求和多个串行读,IM SDK 没有总期限。关闭 CDP 不会取消页面 Promise;Python 返回并释放 alias 锁时,旧脚本可能继续提交,后续动作与它重叠。 +- **修正:**统一 HTTP/CDP/脚本预算,明确超时后仍在途的动作所有权与核验方式。不能把“连接关闭”当“写入已停止”,不明操作不得自动补发。 +- **验收:**离线脚本或 CDP 替身使写入跨过 15 秒;证明旧动作未结束时后续写不会穿透互斥,结果有最终证据或明确不明状态。 + +### GW09 · P2:普通 HTTP 代理丢弃 chunked 请求体 + +- **证据与影响:**`cmd/docker_gateway/proxy.py:226-267,307-333`。只识别 Content-Length,却保留 Transfer-Encoding 转发;合法 chunked 请求变成有声明、无正文/终止块,导致挂起或失败。 +- **修正:**采用完整 HTTP 请求解析/转发能力,或正确处理分块与 trailer;不能静默丢正文。 +- **验收:**分片 chunk、终止块、trailer、大小限制及异常中断,上游正文必须完整一致。 + +### GW10 · P2:校验接受的默认值与执行不一致,缺字段返回 500 + +- **证据与影响:**`cmd/docker_gateway/gateway.py:186,1064,1082-1088,1120,1219-1220`。停止态直连创建省略 `network_exit_id` 能通过验证,但后续直接索引触发 KeyError;代理恢复缺代字段也在正式校验前直接索引,变成 500。 +- **修正:**统一完成默认值归一化及必填检查,然后才允许副作用。 +- **验收:**通过完整 create/HTTP 路由验证合法省略成功、缺必填返回 400,非法请求没有 Docker 调用。 + +### GW11 · P2:删除时“清理待完成”的 `/v1` 响应被改变 + +- **证据与影响:**`cmd/docker_gateway/gateway.py:421-427,929-950`;`cmd/control-plane/hub.go:1789-1796`。旧实现对可重试网络清理失败返回 `202 runtime_cleanup_pending`,新实现传播为 502,控制面仍依赖 202 继续恢复。 +- **修正:**恢复待清理语义及载荷,代冲突仍为 409。若改变契约,须另行批准并同步控制面,不能称为保持 `/v1` 不变。 +- **验收:**一次暂时失败后的完整控制面请求序列;验证 202、409、真正失败分别处理。 + +## 3. 平台动作与事件缺陷 + +### EV01 · P0:首次及恢复监听没有可靠基线,旧通知可能触发自动发送 + +- **证据与影响:**`cmd/docker_gateway/douyin.py:532-553,575-581,894-906`;`cmd/control-plane/creator_events.go:248-274`;`internal/creator/actions.go:245-312`。只给 reconnected 控制消息设置 baseline,Go 又只记日志;普通 notice 未经过历史边界判断。此前未入库的旧通知,只要策略和冷却允许,就可能执行。**已入库终态事件仍有去重,不应夸大为所有旧事件都会重放。** +- **修正:**建立并持久保存经平台验证的新旧边界和恢复分类;边界不明、启用前、断连期间事件仅记录,不自动写入。不能用本机启动时间或重连成功冒充平台连续性证明。 +- **验收:**首次历史、停用期间、断连、缺时间、重启、冷却到期旧通知均零自动写;只有已确认新边界事件进入执行,分类理由可查。 + +### EV02 · P1:事件在持久接收前被删除,错误可造成不可恢复丢失 + +- **证据与影响:**`cmd/docker_gateway/douyin.py:509-529,625-647,670-675,910-912`;`cmd/control-plane/creator_events.go:194-209,274-278`。浏览器 `splice`、Python `popleft` 都先删除;单条详情失败会终止批次,Go 保存失败仅记日志,重新 start 停掉旧队列,队列溢出还清空全部事件。 +- **修正:**先持久接收再确认删除,可重复读取;单条失败隔离,溢出和无法恢复的缺口持久可见,不擅自补发。 +- **验收:**取队列后断网、响应丢失、数据库首次失败、首条详情失败、溢出;事件可恢复或明确记录缺口,不能只以重连成功验收。 + +### EV03 · P1:接收线程同步等待写动作,掩盖后续事件处理延迟 + +- **证据与影响:**`cmd/control-plane/creator_events.go:194-209,274`;`internal/creator/actions.go:486-503`、`migrations/017_creator.sql:192`。一批逐条等待完整自动动作,首条耗时超过 5 秒便阻塞后续;received_at 使用延后的入库时间,未保留最早接收时间。 +- **修正:**快速接收和持久登记与耗时执行分离;仍按执行账号串行写。分别保存平台时间、网关接收、处理开始和结束时间。 +- **验收:**首个执行器阻塞 10 秒时,第二个事件仍在收到后 5 秒内进入处理,时间记录不能后移掩盖排队。 + +### EV04 · P1:后台完成动作被记录为“页面已展示” + +- **证据与影响:**`internal/creator/actions.go:271-275,501`、`settings.go:61-68`。后台完成时写 displayed_at,即使从未打开页面,也有“展示”记录,30 秒指标失真。 +- **修正:**动作完成与页面展示分别记录;展示时间只能来自实际显示该事件的客户端确认。 +- **验收:**不开页面时 displayed_at 为空;显示后才记,并与动作完成时间独立。 + +### EV05 · P1:Douyin 请求漏传代理标识,非直连环境被拒绝 + +- **证据与影响:**`cmd/control-plane/hub.go:140-145` 的 `gatewayGenerationPayload` 不含 network_exit_id;`creator.go:685,727,740`、`creator_events.go:167` 复用它;Python `gateway.py:695-704` 强制比对代理标签,非直连请求变成 409。直连成功证据不能覆盖。 +- **修正:**Douyin 请求统一携带当前有效代和 `environment.Exit.ID`,保留校验;不要顺带混入删除清理语义。 +- **验收:**Go 真实 payload 穿过 Python 校验,正确代理通过,缺失/错误代理及旧代拒绝,直连继续通过。 + +### EV06 · P1:每次重灌 Cookie,阻断有效人工登录并销毁监听 + +- **证据与影响:**`cmd/control-plane/creator.go:661-684,832,936`;`cmd/docker_gateway/gateway.py:529-538`、`douyin.py:266-318`。先强制读取存储凭据,再清 Cookie、导航、设置 Cookie;未先复用有效会话。保存凭据为空或过期时,实际已登录账号仍失败;导航销毁监听,SDK 初始化与后设 Cookie 的身份也未证明一致。 +- **修正:**先核验有效会话并复用;只有确需恢复登录才注入批准的凭据,再重建并核验页面身份和监听边界。不能通过伪造 Cookie 绕过缺陷。 +- **验收:**有效人工会话无存储 Cookie 也可使用;重复采集不导航/清 Cookie;必要恢复后身份、SDK 与监听一致。 + +### EV07 · P1:采集与写执行器支持的凭据格式不一致 + +- **证据与影响:**`cmd/control-plane/creator.go:665-672,809-816,913-920`;`internal/douyin/connector.go:245-287`。采集支持 header 和 JSON bundle,写执行器仅接受 bundle;合法 header 账号到平台写入前就失败,却可能被标不明。 +- **修正:**共用已批准格式的解析入口,解析错误明确属于未执行;与会话复用一起修正,不强制覆盖会话。 +- **验收:**两种合法格式经过采集和写入口;非法格式零平台写调用且状态准确。 + +### EV08 · P1:内部评论/作品 ID 原样下发,人工动作无法执行 + +- **证据与影响:**`web/src/CreatorWorkbenchPage.jsx:237-243,274-283`;`internal/creator/actions.go:599-618,716`;`cmd/control-plane/creator.go:688-690`;`cmd/docker_gateway/douyin.py:395-408`。内部 comment-/work- ID 不等于数字平台 ID;人工回复还允许缺作品 ID。数据库替身执行器不会暴露该问题。 +- **修正:**共享执行边界解析 comment_key、work_key、评论所属作品,并核对平台、作者和关联关系;区分通知自带的平台 ID,不让 UI 拼猜。 +- **验收:**从本地记录创建回复、评论点赞、作品点赞和转发,网关收到正确平台目标;缺关联、错作者、跨作品均零写入。 + +### EV09 · P1:目标评论仅在前 50 条查找 + +- **证据与影响:**`cmd/docker_gateway/douyin.py:921` 的 `ACTION_SCRIPT.findComment` 固定首屏,合法目标在后续页就误报 TARGET_MISMATCH。 +- **修正:**使用已验证的单条详情或正确分页定位;区分读取失败、不可获取和真正身份不匹配。 +- **验收:**第二页目标可以核验;分页失败不报错身份;真正作者/作品不符则拒绝写入。 + +### EV10 · P1:写后核验失败、确认字段缺失被错误归为“明确失败” + +- **证据与影响:**`cmd/docker_gateway/douyin.py:918-924`;`cmd/control-plane/creator.go:700-722`。点赞 POST 已成功,核验 GET 限流被设 definitive;follow/like/IM 缺确认字段也当明确拒绝。实际可能已写入,错误失败提示可误导另发。 +- **修正:**区分写前失败、经验证的明确平台拒绝、写后确认不足。写后读失败不能证明未发送,缺字段不得自动等价为拒绝。 +- **验收:**执行实际脚本逻辑的离线测试:写成功后核验失败、缺字段、SDK 空结果均为不明;明确拒绝才失败,全部禁止自动重发。 + +### EV11 · P1:丢弃平台结果标识,评论成功证据不充分 + +- **证据与影响:**`cmd/docker_gateway/douyin.py:921,924`;`cmd/control-plane/creator.go:700-718`。IM 返回的消息、会话标识被 Go 解码结构丢弃。评论可能缺 ID、正文和明确作品标识,只凭发送者就成功,缺作品标识还用请求值补齐;持久结果仅固定字符串,无法查证。 +- **修正:**保存最小稳定平台标识、目标及实际确认内容;不得用请求值代替响应证据,不记录无关完整私信。 +- **验收:**缺评论 ID、正文错误、父评论不符不能成功;成功私信在操作记录中可查稳定消息和会话标识。 + +### EV12 · P1:排队后不重核状态,已停用的动作仍可能发出 + +- **证据与影响:**`internal/creator/actions.go:475-486,698-716`;`cmd/control-plane/creator.go:650-693`;`cmd/docker_gateway/gateway.py:617-630`。账号业务检查、策略选择在网关锁排队前完成;锁内仅核代和平台身份。等待期间禁言、撤销授权、停策略、改关系没有最终阻止点。 +- **修正:**人工和自动共用按账号执行协调,取得执行权后、实际提交前重核当前条件;自动动作还须检查原策略与关系。 +- **验收:**第二个动作排队时逐项改变条件,后续零平台写;已经发出的动作保留真实结果,不能假装撤销。 + +### EV13 · P1:转发确认了文案,却不发送文案 + +- **证据与影响:**`cmd/docker_gateway/douyin.py:409-412,921`;`cmd/control-plane/creator.go:691`。入口要求非空文案,但转发请求不使用 `p.text`,仍报 reposted/succeeded,确认内容与实际行为不符。 +- **修正:**先验证对应平台带文案语义;若只有无文案推荐,明确能力差异并取得用户裁决,不静默丢文案。 +- **验收:**确认文本进入真实动作且结果一致;不支持时明确阻止,不宣称带文案转发成功。 + +### EV14 · P2:同页面重连累积旧事件处理器 + +- **证据与影响:**`cmd/docker_gateway/douyin.py:532-553,894-906`。同 key 新安装前未调用旧 dispose,关闭 CDP 不会移除页面 addEventListener;stop 只释放最新一组。 +- **修正:**重装前清理旧状态,安装失败也清理本次处理器。 +- **验收:**同页面重复重连后处理器数量固定,stop 后为零。 + +## 4. 账号、采集、计划及持久数据缺陷 + +### BE01 · P1:旧作品首次发现后,指标间隔状态错误 + +- **证据与影响:**`internal/creator/metrics.go:138-155,171-182`。T+2 小时首次发现,将 next_plan 推到 T+3 小时,却保留 1 小时间隔;之后推到 T+5 而非 T+7。 +- **修正:**根据发布时间和当前时刻统一计算完整计划状态,不能只改下次时间。 +- **验收:**T+2、T+8、T+32 首次发现后连续多个计划点及重启恢复,符合 1、3、7、15、31、55…累计时刻。 + +### BE02 · P1:固定采集计划按完成时间漂移 + +- **证据与影响:**`cmd/control-plane/creator.go:852-853`;`internal/creator/collection.go:255,266-270`。T 开始、T+5 分钟完成的 30 分钟任务,下次变成 T+35;不符合固定起点与迟到跳到严格未来点。 +- **修正:**持久保存起点及未来计划点,完成时间仅记实际执行。 +- **验收:**慢任务、错过多个周期、重启、暂停恢复后不漂移、不回放历史周期。 + +### BE03 · P1:修改设置没有重设已有计划 + +- **证据与影响:**`internal/creator/settings.go:32-57`;`metrics.go:138-141,171`。仅改 settings,账号未来时间不重设,指标计划混用旧间隔/倍数与新年龄。 +- **修正:**一致重算账号、评论未来起点和作品未来计划,保留历史快照。 +- **验收:**有存量计划时逐项修改间隔、倍数、上限、年龄;未来按新规则,历史不改、不补造。 + +### BE04 · P1:普通自有账号不参与作品和一级评论补充采集 + +- **证据与影响:**`internal/creator/collection.go:259`;`cmd/control-plane/creator.go:883-891`。到期查询强制 big_account=true,将 W1 的自有账号错误缩减为大号。 +- **修正:**采集资格与大号自动响应模式分离。 +- **验收:**普通账号、大号、小号符合条件时均可采集,只有自动响应受大号模式控制。 + +### BE05 · P1:封禁、注销账号仍可被选作采集账号 + +- **证据与影响:**`internal/creator/collection.go:264-270`;`cmd/control-plane/creator.go:798-804,903-908,955-964`。检查授权但不检查 banned/deleted 业务状态,自动选择与手动指定均受影响。 +- **修正:**所有采集入口复用同一资格检查。 +- **验收:**自有、竞品自动选择、竞品手动指定三路均拒绝封禁/注销且零平台请求;历史仍可查。 + +### BE06 · P1:自有采集的登录/验证阻断仍被自动重试 + +- **证据与影响:**`cmd/control-plane/creator.go:913-945`;`internal/creator/collection.go:189-193,265-270,346-347`。身份失败未持久 blocked,分页 challenge/401/403 统一 failed,下轮再次选择。 +- **修正:**区分可恢复网络错误和需人工处理的阻断,贯穿身份检查到分页,提供明确人工恢复。 +- **验收:**身份不符、验证码、401 后下一轮零请求;网络暂时失败按批准计划恢复。 + +### BE07 · P1:采集租约不续期,且未覆盖页面数据写入和结束状态 + +- **证据与影响:**`internal/creator/collection.go:127-136,157-186,316-342,398-408`;`content.go:94-139`。阶段超过十分钟后断点保存失败;旧持有者被接管后仍可写作品/评论/指标。MarkCompetitorSync 无所有权条件,旧请求可覆盖新结果或暂停状态。 +- **修正:**一致续期;一页数据与断点更新受有效领取条件共同保护,最终状态同样核对所有者和暂停状态。 +- **验收:**长任务、到期接管、旧请求迟返、运行中暂停恢复;旧持有者不能覆盖新代结果。 + +### BE08 · P1:资料编辑绕过大小号关系和策略停用规则 + +- **证据与影响:**`internal/creator/accounts.go:124-143,173-205`;`cmd/control-plane/creator.go:65-76`。profile 更新直接设 big_account,绕过专用入口的小号晋升检查;关闭再开启或非正常再恢复,不要求重启用策略。 +- **修正:**共享事务性状态变更规则,停用所有受影响策略;恢复正常不自动恢复策略。 +- **验收:**资料入口和专用入口分别覆盖小号晋升、关闭开启、正常→异常→正常,非法关系拒绝且策略不偷偷恢复。 + +### BE09 · P1:小号不同归属冲突被吞掉,并发关系检查有空窗 + +- **证据与影响:**`internal/creator/accounts.go:185-194,248-268`;`migrations/017_creator.sql:149-155`。小号已归 A,关联 B 时 ON CONFLICT DO NOTHING 仍返回成功;关联与晋升各自先查后写,可并发形成非法角色状态。 +- **修正:**同关系重复可幂等,不同归属明确冲突;按一致顺序锁相关账号再检查。 +- **验收:**同关系重复、不同归属、关联与晋升双事务竞争,返回结果与实际状态一致。 + +### BE10 · P1:无文案动作也进入 AI,失败后占用冷却 + +- **证据与影响:**`internal/creator/actions.go:411-440`、`logic.go:178-179`;`cmd/control-plane/creator.go:28`。仅判断 text==空,关注/点赞无文案也调用 generator;生产为 nil,动作失败并已占冷却。 +- **修正:**只为需要文本的动作选择或生成文案,不改动冷却对已开始动作的既定语义。 +- **验收:**关注、作品点赞、评论点赞在 generator=nil 时执行一次且零 AI 调用;文本动作仍检查配置。 + +### BE11 · P1:策略启用不校验 AI 必要条件 + +- **证据与影响:**`internal/creator/actions.go:82-110,152-176,366-440`。缺候选文本、有效 AI 配置或回复要求仍可启用,收到事件后才失败并可能占冷却。 +- **修正:**创建、修改、专用启用入口统一校验动作实际需要的条件,不把“审批”布尔值当可用客户端。 +- **验收:**三个 enabled=true 入口缺配置均拒绝且无冷却记录;无文案动作不被错误限制。 + +### BE12 · P1:同一作品第二来源无法匹配对应线索规则 + +- **证据与影响:**`internal/creator/rules.go:138-144`;`content.go:222-247`。作品已有多来源表,但分析仍只比最初 SourceType;owned 后关联 competitor,competitor 规则被误拒,反向同样。 +- **修正:**使用实际来源关联判断规则范围。 +- **验收:**双来源两种入库顺序、各来源专用规则都正确,同评论仍只形成一条线索。 + +### BE13 · P1:UID 登记值被当作 SecUID 查询作品 + +- **证据与影响:**`internal/douyin/creator_collector.go:25-54`;`cmd/control-plane/creator.go:939-943`。核验允许 UID/SecUID/UniqueID 匹配,却不保留规范映射,后续原样放进 sec_user_id。旧 connector 使用身份响应的 SecUID。 +- **修正:**身份核验输出平台规范标识映射,各接口使用其要求的标识。 +- **验收:**UID 和 SecUID 登记同账号,最终作品请求均用正确 SecUID;随后补真实只读证据。 + +### BE14 · P1:密码更新失败可能删除原有密码 + +- **证据与影响:**`internal/creator/accounts.go:111-139`。新密码先覆盖固定 secret key,数据库失败后删除该 key,旧密码已经无法恢复,数据库却可能仍显示已配置。 +- **修正:**新值使用独立引用,成功关联后再清理旧值;失败只清本次新建项,结果不明需核验关联。 +- **验收:**已有密码时注入保存失败、数据库失败及结果不明,旧有效凭据不丢,配置标记与引用一致。 + +### BE15 · P2:冷却时长没有溢出上限 + +- **证据与影响:**`internal/creator/accounts.go:90-93`;`actions.go:390-394`。巨大正数转 time.Duration 再乘秒可能溢出到过去,冷却失效。 +- **修正:**复用已有时长边界校验。 +- **验收:**上限、上限+1、int64 最大值;拒绝溢出,历史占用不改变。 + +### BE16 · P2:未来发布时间作品被过滤掉,无法进入待核验状态 + +- **证据与影响:**`internal/creator/collection.go:317-320` 先按窗口跳过,`content.go:203-207` 的 future 标记无法到达。 +- **修正:**可靠窗口外和时间无效/未来分别处理;后者保留待核验并说明范围不完整。 +- **验收:**正常、窗口外、缺时间、未来时间同批输入;异常时间不丢资料、不建年龄计划。 + +### BE17 · P2:AI 缺失判定字段被当作明确否定 + +- **证据与影响:**`internal/creator/bailian.go:110-117,127-134`。`{}`、null、缺 match 或 match=null 解码为 bool false,后续变成非线索而非分析失败。当前客户端尚未接生产,接线前也须修复。 +- **修正:**区分字段缺失/null 和明确 true/false,验证必需依据。 +- **验收:**上述不完整样例均失败;合法真/假按真实依据保存。 + +### BE18 · P2:迟到或重复消息倒退会话最近时间 + +- **证据与影响:**`internal/creator/actions.go:755,804-816`。消息去重前先无条件更新 last_message_at,旧消息即使重复未插入,也改变会话排序。 +- **修正:**消息写入和摘要更新一致,最近时间不得被旧/无时间消息覆盖。 +- **验收:**新→旧→重复旧→缺时间消息,时间不倒退且不重复消息。 + +### BE19 · P2:作品列表持有连接时再查询,低连接池下互相等待 + +- **证据与影响:**`internal/creator/content.go:368-384`;`store.go:56`。外层 rows 未关闭,又逐作品查询来源;并发占满池后都等待内层连接。 +- **修正:**先读取并关闭外层 rows,再批量读来源,或一次查询完成,避免嵌套占用。 +- **验收:**连接池设为 1 时带来源列表仍在期限内返回,再覆盖并发列表。 + +## 5. 前端确定缺陷 + +人工目标 ID 错配统一见 EV08;AI、媒体、登录、推送缺口见第 7 节,不重复计数。 + +### UI01 · P1:设置加载后原样回传只读字段,保存返回 400 + +- **证据与影响:**`web/src/CreatorSettingsPage.jsx:30-31,43`;`internal/creator/models.go:23-52`;`cmd/control-plane/creator.go:39-44,604-611`。GET 含 updated_at,完整合入 form 并 PUT;SettingsUpdate 无此字段,严格解码拒绝。现有 Mock fixture 漏了 updated_at。 +- **修正:**仅发送可编辑字段,保持后端严格校验。 +- **验收:**真实 GET 结构加载后保存,经真实编码/解码成功;请求无只读字段。 + +### UI02 · P1:工作台切换页签将旧数据按新类型渲染,页面崩溃 + +- **证据与影响:**`web/src/CreatorWorkbenchPage.jsx:41-42,79-150,501,599-609,800`。共享 data,评论切规则时立即计算 include_keywords.join,切线索时访问 lead.comment.id;pending 包裹不能阻止 children 提前计算。 +- **修正:**按页签隔离数据或明确数据归属,不匹配时不计算列表。 +- **验收:**各类型非空数据逐对切换,延迟/乱序返回,始终不崩溃、不串表。 + +### UI03 · P1:人工操作标识仅在页面内存,刷新恢复可变成第二次发送 + +- **证据与影响:**`web/src/CreatorWorkbenchPage.jsx:63-74,237-283,331-347`。响应丢失后刷新,原 key 不可恢复,再确认生成新 key。服务端同 key 去重仍正确,但无法识别原操作的恢复。不是说用户永远不能再发同内容。 +- **修正:**复用持久操作记录,恢复原标识和冻结载荷;不明只查询,确需另发必须明确新确认并关联原操作。 +- **验收:**创建/执行响应各自丢失,页面重挂和服务重启仍只执行一次;恢复不创建新发送。 + +### UI04 · P1:局域网 HTTP 下直接调用 randomUUID 可能使发送入口失效 + +- **证据与影响:**`web/src/CreatorWorkbenchPage.jsx:35-36,243,331`。部署允许非 localhost HTTP,但 crypto.randomUUID 依赖安全上下文;直接调用无能力处理。该环境兼容性结论来自 API 约束,本轮未作浏览器复现。 +- **修正:**使用适配批准部署方式的可靠操作标识机制,结合 UI03 持久恢复;不降为弱随机或吞错。 +- **验收:**randomUUID 缺失替身及实际局域网 HTTP,回复和私信都能正常取得并恢复标识。 + +### UI05 · P1:账号切换未隔离迟返请求,可显示和保存另一账号资料 + +- **证据与影响:**`web/src/CreatorAccountsPage.jsx:105-168,289`。A 的读取/保存等待时切 B,A 迟返仍 setForm/setStrategies;B 标题下出现 A 数据,继续保存可能覆盖 B,策略操作也可能落到 A。 +- **修正:**结果绑定发起账号,失效响应不应用;提交中处理目标切换,未保存修改明确确认。 +- **验收:**两个账号明显不同资料和策略,打乱读取及保存响应;任何时刻均不串用或错操作。 + +### UI06 · P1:私信账号选择被当前会话平台锁住 + +- **证据与影响:**`web/src/CreatorWorkbenchPage.jsx:621-640`。顶层账号选项按当前会话平台过滤,有抖音会话时另一平台账号消失,无法正常跨平台切换。 +- **修正:**先独立选账号,再清旧会话并读目标账号数据。 +- **验收:**两平台各有会话账号双向切换,选项始终可选,请求仅属于新账号。 + +### UI07 · P2:所有非成功状态都显示“结果不确定” + +- **证据与影响:**`web/src/CreatorWorkbenchPage.jsx:287-293,352-358`。failed、blocked、processing 被同文案覆盖,隐藏确定原因和在途状态。 +- **修正:**按后端真实状态及原因展示,并关联原操作;不得给不明状态提供重新发送式重试。 +- **验收:**成功、失败、阻止、处理中、不明五种结果展示准确。 + +### UI08 · P2:读取失败被当作空数据 + +- **证据与影响:**`web/src/CreatorAccountsPage.jsx:105-114`、`CreatorCompetitorsPage.jsx:72-77`。策略/账号失败直接清空数组,误报尚无策略或可用账号。 +- **修正:**错误独立保存,保留已知数据,提供只读重试,依赖数据的写入口禁用。 +- **验收:**首次失败和有数据后刷新失败均不得冒充真实空列表。 + +### UI09 · P2:关键词输入清理规则与后端和需求不一致 + +- **证据与影响:**`web/src/CreatorWorkbenchPage.jsx:28-32,213-214`;`internal/creator/logic.go:102-115`。前端 trim 去除更广 Unicode 空白,filter 静默丢空词;后端只清指定普通空白并拒绝空词。 +- **修正:**统一有限清理规则,空词明确拒绝,不改变原文匹配语义。 +- **验收:**全角空格、NBSP、普通空白、连续逗号;界面与直接 API 输入一致。 + +## 6. 必须纠正的测试、部署和文档结论 + +### T01 · P1:覆盖率统计必须排除测试代码 + +当前命令固定为 `python3 -m coverage run --source=cmd/docker_gateway -m unittest cmd.docker_gateway.test_gateway`,随后执行 `python3 -m coverage report --include='cmd/docker_gateway/*.py' --omit='cmd/docker_gateway/test_*.py' --fail-under=65`。最近一次本地运行 56 项测试通过,生产代码总覆盖率达到 **65%**;`.coverage` 只作为本次运行产物,不提交仓库。 + +- **修正:**统计分母排除测试文件,并补 CDP 超时终止、DM 规范化和网络创建异常回收回归。 +- **验收:**覆盖率命令必须与同次测试生成的数据对应;逐文件低覆盖率不被总数掩盖,真实平台验收另行登记。 + +### T02 · P1:常规测试通过不等于关键数据库和平台逻辑已验证 + +- `internal/creator/integration_test.go` 无专用数据库地址时仍会 skip;使用 `CREATORHUB_POSTGRES_TEST_URL` 时最近一次已实际运行内容、动作、消息和迁移测试。 +- `creator_events_test.go` 覆盖字段/边界转换;新增 stale processing 代码路径仍须在专用数据库环境做崩溃恢复验收。 +- Python 动作测试覆盖实际嵌入脚本的关键分类字符串和 CDP 超时终止;不替代真实浏览器/平台写入。 +- 前端覆盖错误、忙碌、账号隔离、策略/规则编辑及消息刷新;不替代视觉、真机和外部服务验收。 +- **验收口径:**记录 executed/skipped 数量及运行条件;任何 Mock 成功仅证明它覆盖的局部规则。 + +### DOC01 · P1:正常退出证据仍需与实现分开记录 + +当前代码已配置 SIGINT/SIGTERM、请求排空和 45 秒 Compose 宽限,且本地已有特定 reservation 回收证据;这不等于所有 Docker 故障注入、并发和强杀场景都通过。部署文档保留该限制,不把单次重启时长包装成完整生命周期验收。 + +### DOC02 · P1:登录要求冲突,不能自行恢复自动登录 + +`AGENTS.md` 与 `docs/plan01.md` 当前统一为不自动登录。**当前修正仍遵守不自动登录、不暴露密码/Cookie/二维码/验证码。** 有效会话可复用,失效后进入人工登录并由 gateway identity 核验;密码“已保存”不得表述为“登录能力已接通”。 + +### DOC03 · P2:采集接线与完整性证据分开 + +当前文档已表述为:存在数字游标校验、分页请求、保存路径和独立计划;完整性、标识转换、分页边界和平台验收仍未由本地代码证明。接受游标参数不等于证明无漏项。 + +### DOC04 · P2:完成表混杂历史与当前轮次 + +`docs/plan01.md` 当前已按历史文档、代码接线、实际运行证据和外部待验收项分开;新增能力继续按该三类记录,不能用代码接线代替平台成功。 + +### DOC05 · P1:浏览器镜像必须由仓库源码和 immutable digest 管理 + +当前分支新增 `docker/browser-wrapper/`,以显式 `BROWSER_BASE_IMAGE` digest 构建 Xvfb/CDP 包装入口,并在 deployment 文档记录 base/published digest、构建提交和 Dockerfile。干净环境构建/拉取和真实身份核验仍需拥有已批准的第三方 base digest 后执行;本地 tag 不作为证据。 + +### DOC06 · P1:真实证据只证明局部可达,不证明业务成功 + +上一轮身份读取 200、事件启动 200、空列表、重启后短期无错误,只能证明这些时点的局部可达。不得扩大为基线、无遗漏、处理时限、正常退出或全动作成功。直接 follow 返回 BUSINESS_REJECTED 且后查状态为未关注,不能证明关注能力成功;当前 EV10 已将非 2xx 和写后核验不足保留为不明。控制面 resolve credential 发生在写前,说明 EV06/EV07 所在链路缺陷,不应一概归咎于“只差用户提供 Cookie”。结果不明的原操作仍不得自动重试。 + +## 7. 已确认的功能接线缺口 + +这些项目尚未完成实际流程,不能全部归为“缺真实平台证据”。补齐仍须遵守现有范围、供应商批准与不自动登录等约束。 + +| 编号/优先级 | 当前证据与缺口 | 修正方向及完成条件 | +| --- | --- | --- | +| GAP01 · P1 | `creator.go` 已装配独立的到期指标调度,作品发现与指标刷新分离;真实平台请求频率仍需外部验收。 | 本地已验证独立到期读取、年龄和计划边界;真实平台需继续证明过龄零请求。 | +| GAP02 · P1 | Bailian adapter 已装配;缺少 key/model 明确返回 unavailable,质量样本和生产供应商仍待批准。 | 本地验证装配、配置校验和失败边界;生产质量与费用由外部验收。 | +| GAP03 · P1 | 新增真实下载、FFmpeg 音轨提取和可配置转写命令;原手写 material-step 路由已删除,转写未配置时明确 failed。 | 本地验证产物驱动和失败状态;音视频、无音轨、空间、供应商质量及重试仍需样本验收。 | +| GAP04 · P1 | Douyin DM notice 可规范化,入站事件和人工出站结果均写入本地会话,重复键幂等;平台实际 DM payload/history 仍待验收。 | 本地验证账号隔离、重复消息和失败状态;平台历史读取与非文本消息需外部证据。 | +| GAP05 · P1 | 事件处理/缺口已持久化,工作台在打开 DM 会话时按 10 秒刷新并支持手动刷新;未新增未批准的双向推送框架。 | 本地验证持久结果恢复;平台关页监听、30 秒可见和真实断连仍待验收。 | +| GAP06 · P1 | 账号页显示登录状态、原因、核验时间和凭据标记;login-result 不得直接写 logged_in,只有 gateway identity 核验可写入。 | 本地验证伪造成功被拒;人工登录入口和真机身份重核仍待外部验收。 | +| GAP07 · P1 | 策略支持编辑/调序、启停;关系列表已有解除入口,更新沿用既有接口。 | 本地验证账号切换和策略更新;解除/再绑定并发效果仍需集成证据。 | +| GAP08 · P1 | 规则支持来源范围、编辑、启停;评论/线索提供人工联系入口,未接自动联系。 | 本地验证来源和人工确认边界;历史判断依据展示仍需补充验收。 | +| GAP09 · P1 | 同步状态、错误、下次时间已显示;链接解析预览与完整组合筛选仍未完成。 | 继续补链接解析/组合筛选,真实0与缺值需保持可区分。 | +| GAP10 · P2 | 账号切换和素材文稿切换已有放弃确认,DM 草稿继续本地保留。 | 仍需覆盖工作台页签/路由切换的未保存状态。 | +| GAP11 · P1 | 小红书当前在 `creator.go:792-794` 等入口明确阻断;登记或列表里有平台选项不等于业务实现。 | 独立完成同范围能力核实与接线;每类缺口单独裁决,不以抖音成功代替,不宣称平台本身不支持。 | + +## 8. 建议修正顺序与交付门槛 + +### 第一批:限制错误外部写入、误删和丢数据 + +- EV01、EV02、EV10—EV12:历史边界、持久接收、结果证据和写前条件。 +- GW01、GW03、GW06、GW08:创建归属、停止排空、网络回滚、仍在途脚本互斥。 +- BE07—BE09、BE14:租约、账号关系/停用及凭据更新一致性。 +- 每项先有未修时失败的隔离回归,再改代码;不能用放宽代校验、吞异常、清库或无限重试替代根因修正。 + +### 第二批:恢复最小可用闭环 + +- 代理和 `/v1`:GW02、GW04、GW05、GW07、GW09—GW11、EV05。 +- 身份和人工目标:EV06—EV09、EV13、BE13。 +- 页面可用性与人工标识恢复:UI01—UI06;无文案自动动作 BE10、启用校验 BE11。 +- 验收先经过真实 Go 请求/Python 校验/页面响应契约,再做获批的小规模真实平台动作。结果不明的旧操作仅核验,不能拿新 key 伪装重试。 + +### 第三批:计划正确性和剩余功能 + +- BE01—BE06、BE12、BE15—BE19、UI07—UI09、EV03、EV04、EV14。 +- 按 GAP01—GAP11 完成实际缺口;AI 配置、平台差异和自动登录争议先取得必要批准。 +- 保留原分阶段顺序:先抖音完整流程,再小红书,最后共存及错误恢复,不用通用框架扩大本期范围。 + +### 共同验收门槛 + +1. 修正清单每项关联回归用例和结果;条目关闭必须有证据,不以“代码已改”关闭。 +2. Python 生产代码覆盖率排除测试后至少 65%,Go 数据库测试明确实际运行,关键竞争/崩溃路径有测试。 +3. 后端 test/vet/race 和控制面构建,Python 网关测试/构建、Compose 配置与生命周期通过;迁移后的网关不再要求构建已删除的 Go 可执行入口。 +4. 前端 lockfile 安装、非交互测试、构建通过,补真实契约、非空切换、双账号/双平台、乱序和刷新恢复测试。 +5. 不泄露凭据,不改变账号身份和操作目标,不自动重试不明写入。日志保留操作、账号、代、阶段、错误类别及时间,避免只有笼统错误而无排查线索。 +6. 平台事件、恢复连续性、各动作、媒体产物、AI 质量与界面时限分别验收;HTTP 200、容器 running、空事件队列、无日志错误均不能替代。 +7. 尚未完成项目必须继续标注未完成;范围调整需明确批准,不能通过修改文档把缺功能包装成已完成。 + +## 9. 审查覆盖与局限 + +### 9.1 覆盖清单 + +- **网关新增:**`cmd/__init__.py`;`cmd/docker_gateway/{__init__,docker_client,douyin,gateway,proxy,test_gateway}.py`;`requirements-gateway.lock`、`requirements-gateway-dev.lock`。 +- **旧网关删除:**`cmd/docker-gateway/{main,douyin,proxy}.go` 及 `main_test.go`、`douyin_test.go`、`proxy_test.go`、`douyin_cdp_integration_test.go`。重点对照发生契约和清理/代理回归的旧实现;未把每个已删除测试逐一移植或重跑。 +- **控制面:**`cmd/control-plane/creator.go`、`creator_events.go`、`creator_events_test.go`、`credential.go`、`main.go`、`hub.go`、`hub_test.go`、`douyin_test.go`;新增及改动路径沿调用链核对。 +- **领域:**`internal/creator/` 全部生产 Go 文件、全部现有测试文件、017—026 十份迁移。 +- **采集与凭据:**`internal/douyin/{connector,connector_test,creator_collector,creator_collector_test}.go`;`internal/phasea/store.go` 的改动与关联读取路径。 +- **前端:**`web/src/Creator{Accounts,Competitors,Settings,Workbench}Page.jsx`、`CreatorPages.test.jsx`、`BrowsersPage.test.jsx`、`Layout.jsx`、`dataProvider.js`、`lib/ui.jsx`、`main.jsx`。 +- **部署/文档:**`Dockerfile`、`compose.yaml`、`AGENTS.md`、`docs/plan01.md`、`docs/deployment.md`、`docs/architecture/container-control.md`、`docs/product/compliance-product-plan.md`。另对照未改动的 Compose 开发配置和相关测试。 + +### 9.2 未运行或无法从本地代码证明的内容 + +- 本轮未进行真实平台写操作。所有本地测试、覆盖率和 Compose 配置结果均以交付前实际命令记录为准,不能证明平台业务成功。 +- 未实测浏览器布局、键盘操作、局域网 HTTP、自动指纹、两账号全流程、Docker 故障注入、数据库并发、部署正常退出码和真实代理兼容性;这些仍须按原 AC 清单登记运行证据,不能将本报告未发现某项缺陷当作该项通过。 +- 平台真实分页完整性、事件 ID/恢复位置、迟到边界、DM 历史范围、带文案转发能力、每类动作成功证据仍待批准的真机验证。 +- 本报告不是对主线全部既有代码的全面审计;未改动部分只在解释本分支调用链和迁移契约时对照。 +- 未发现问题的文件不等于证明无问题。四路审查的完整原始输出保存在本次会话的 `subagent-artifacts/outputs/8f8a3bc6-a9a0-48fc-9cf4-54a8ab4453b2/` 下,文件为 `review-gateway.md`、`review-events.md`、`review-domain.md`、`review-web-docs.md`。本文汇总后的编号是后续修正跟踪依据。 diff --git a/internal/creator/accounts.go b/internal/creator/accounts.go index 4239203..9b01e2b 100644 --- a/internal/creator/accounts.go +++ b/internal/creator/accounts.go @@ -3,6 +3,7 @@ package creator import ( "context" "database/sql" + "errors" "fmt" "strings" "time" @@ -87,7 +88,7 @@ func validateProfileUpdate(input AccountProfileUpdate) error { if input.BusinessStatus != "normal" && input.BusinessStatus != "muted" && input.BusinessStatus != "banned" && input.BusinessStatus != "deleted" { return ErrInvalid } - if input.CooldownSeconds <= 0 || utf8.RuneCountInString(input.LoginUsername) > 255 || + if !validCooldownSeconds(input.CooldownSeconds) || utf8.RuneCountInString(input.LoginUsername) > 255 || utf8.RuneCountInString(input.RealName) > 100 || utf8.RuneCountInString(input.IdentityNumber) > 64 || utf8.RuneCountInString(input.Note) > 1000 || utf8.RuneCountInString(input.ReplyRequirements) > 4000 { return ErrInvalid @@ -95,7 +96,7 @@ func validateProfileUpdate(input AccountProfileUpdate) error { return nil } -func (s *Store) UpdateAccountProfile(ctx context.Context, accountID string, input AccountProfileUpdate) (AccountProfile, error) { +func (s *Store) UpdateAccountProfile(ctx context.Context, accountID string, input AccountProfileUpdate) (profile AccountProfile, returnErr error) { input.LoginUsername = strings.TrimSpace(input.LoginUsername) input.RealName = strings.TrimSpace(input.RealName) input.IdentityNumber = strings.TrimSpace(input.IdentityNumber) @@ -110,19 +111,51 @@ func (s *Store) UpdateAccountProfile(ctx context.Context, accountID string, inpu passwordConfigured := false var secret SecretReference - var secretKey string + var secretKey, oldSecretID, oldSecretKey string + committed := false if input.Password != "" { if s.secrets == nil { return AccountProfile{}, ErrUnavailable } - secret = SecretReference{ID: accountID + "-password", Provider: "os_keyring"} - secretKey = "creatorhub/" + accountID + "/password" + secret = SecretReference{ID: newID("password"), Provider: "os_keyring"} + secretKey = "creatorhub/" + accountID + "/password/" + secret.ID if err := s.secrets.Store(ctx, secret, secretKey, input.Password); err != nil { return AccountProfile{}, fmt.Errorf("store account password: %w", err) } passwordConfigured = true + defer func() { + if !committed && returnErr != nil { + if cleanupErr := s.secrets.Delete(ctx, secret, secretKey); cleanupErr != nil { + returnErr = errors.Join(returnErr, cleanupErr) + } + } + }() } - _, err := s.db.ExecContext(ctx, ` + tx, err := s.db.BeginTx(ctx, nil) + if err != nil { + return AccountProfile{}, fmt.Errorf("begin account profile update: %w", err) + } + defer tx.Rollback() + var lockedID string + if err := tx.QueryRowContext(ctx, `SELECT id FROM social_account WHERE id=$1 FOR UPDATE`, accountID).Scan(&lockedID); err != nil { + return AccountProfile{}, rowError(err) + } + if input.Password != "" { + err := tx.QueryRowContext(ctx, `SELECT secret_reference_id,secret_key FROM creator_account_password WHERE account_id=$1 FOR UPDATE`, accountID).Scan(&oldSecretID, &oldSecretKey) + if err != nil && !errors.Is(err, sql.ErrNoRows) { + return AccountProfile{}, databaseError(err) + } + } + if input.BigAccount { + var isSmall bool + if err := tx.QueryRowContext(ctx, `SELECT EXISTS (SELECT 1 FROM creator_relation WHERE small_account_id=$1)`, accountID).Scan(&isSmall); err != nil { + return AccountProfile{}, databaseError(err) + } + if isSmall { + return AccountProfile{}, ErrConflict + } + } + _, err = tx.ExecContext(ctx, ` UPDATE creator_account_profile SET login_username = $2, password_configured = CASE WHEN $3 THEN true ELSE password_configured END, real_name_status = $4, real_name = $5, identity_number = $6, @@ -139,10 +172,45 @@ func (s *Store) UpdateAccountProfile(ctx context.Context, accountID string, inpu } return AccountProfile{}, databaseError(err) } + if input.Password != "" { + if _, err := tx.ExecContext(ctx, `INSERT INTO creator_account_password (account_id,secret_reference_id,secret_key) VALUES ($1,$2,$3) ON CONFLICT (account_id) DO UPDATE SET secret_reference_id=EXCLUDED.secret_reference_id,secret_key=EXCLUDED.secret_key,updated_at=now()`, accountID, secret.ID, secretKey); err != nil { + return AccountProfile{}, databaseError(err) + } + } + if !input.BigAccount || input.BusinessStatus != "normal" { + if _, err := tx.ExecContext(ctx, `UPDATE creator_strategy SET enabled=false, updated_at=now() WHERE big_account_id=$1`, accountID); err != nil { + return AccountProfile{}, databaseError(err) + } + } + if err := tx.Commit(); err != nil { + return AccountProfile{}, fmt.Errorf("commit account profile update: %w", err) + } + committed = true + if input.Password != "" && oldSecretID != "" && (oldSecretID != secret.ID || oldSecretKey != secretKey) { + if err := s.secrets.Delete(ctx, SecretReference{ID: oldSecretID, Provider: "os_keyring"}, oldSecretKey); err != nil { + return AccountProfile{}, fmt.Errorf("replace account password: remove old secret: %w", err) + } + } return s.GetAccountProfile(ctx, accountID) } func (s *Store) RecordLoginResult(ctx context.Context, accountID, status, reason, actualKey string) (LoginResult, error) { + status = strings.TrimSpace(status) + if status == "logged_in" { + return LoginResult{}, ErrConflict + } + return s.recordLoginResult(ctx, accountID, status, reason, actualKey) +} + +func (s *Store) RecordVerifiedLoginResult(ctx context.Context, accountID, actualKey string) (LoginResult, error) { + actualKey = strings.TrimSpace(actualKey) + if actualKey == "" { + return LoginResult{}, ErrInvalid + } + return s.recordLoginResult(ctx, accountID, "logged_in", "", actualKey) +} + +func (s *Store) recordLoginResult(ctx context.Context, accountID, status, reason, actualKey string) (LoginResult, error) { status = strings.TrimSpace(status) reason = strings.TrimSpace(reason) if status != "logged_in" && status != "needs_login" && status != "failed" && status != "manual_required" { @@ -155,9 +223,8 @@ func (s *Store) RecordLoginResult(ctx context.Context, accountID, status, reason if err != nil { return LoginResult{}, err } - if status == "logged_in" && strings.TrimSpace(actualKey) != "" && actualKey != profile.PlatformAccountKey { - status = "manual_required" - reason = "登录身份与登记账号不一致" + if status == "logged_in" && (actualKey == "" || actualKey != profile.PlatformAccountKey) { + return LoginResult{}, ErrConflict } now := time.Now().UTC() _, err = s.db.ExecContext(ctx, ` @@ -237,6 +304,30 @@ func (s *Store) SetRelation(ctx context.Context, bigAccountID, smallAccountID st return fmt.Errorf("begin creator relation: %w", err) } defer tx.Rollback() + // Lock both account rows in a stable order so relationship checks and role changes + // cannot observe a half-updated account pair. + first, second := bigAccountID, smallAccountID + if first > second { + first, second = second, first + } + rows, err := tx.QueryContext(ctx, `SELECT id FROM social_account WHERE id IN ($1,$2) ORDER BY id FOR UPDATE`, first, second) + if err != nil { + return databaseError(err) + } + defer rows.Close() + count := 0 + for rows.Next() { + count++ + } + if err := rows.Err(); err != nil { + return databaseError(err) + } + if count != 2 { + return ErrNotFound + } + if err := rows.Close(); err != nil { + return err + } var bigPlatform, smallPlatform string if err := tx.QueryRowContext(ctx, `SELECT platform FROM social_account WHERE id = $1`, bigAccountID).Scan(&bigPlatform); err != nil { return rowError(err) @@ -265,9 +356,23 @@ func (s *Store) SetRelation(ctx context.Context, bigAccountID, smallAccountID st if smallIsBig || bigIsSmall { return ErrConflict } - if _, err := tx.ExecContext(ctx, `INSERT INTO creator_relation (big_account_id, small_account_id) VALUES ($1, $2) ON CONFLICT DO NOTHING`, bigAccountID, smallAccountID); err != nil { + result, err := tx.ExecContext(ctx, `INSERT INTO creator_relation (big_account_id, small_account_id) VALUES ($1, $2) ON CONFLICT DO NOTHING`, bigAccountID, smallAccountID) + if err != nil { return databaseError(err) } + affected, err := result.RowsAffected() + if err != nil { + return err + } + if affected == 0 { + var existingBig string + if err := tx.QueryRowContext(ctx, `SELECT big_account_id FROM creator_relation WHERE small_account_id=$1`, smallAccountID).Scan(&existingBig); err != nil { + return rowError(err) + } + if existingBig != bigAccountID { + return ErrConflict + } + } } else { if _, err := tx.ExecContext(ctx, `DELETE FROM creator_relation WHERE big_account_id = $1 AND small_account_id = $2`, bigAccountID, smallAccountID); err != nil { return databaseError(err) diff --git a/internal/creator/actions.go b/internal/creator/actions.go index b06496d..820cf93 100644 --- a/internal/creator/actions.go +++ b/internal/creator/actions.go @@ -79,11 +79,41 @@ func scanStrategy(scanner interface{ Scan(...any) error }) (Strategy, error) { const strategySelect = `SELECT id,big_account_id,execution_account_id,position,enabled,event_types,action,target_type,candidate_texts,created_at,updated_at FROM creator_strategy` +func (s *Store) validateEnabledStrategy(ctx context.Context, bigAccountID string, input StrategyInput) error { + if !input.Enabled || !ActionRequiresText(input.Action) || len(input.CandidateTexts) > 0 { + return nil + } + if strings.TrimSpace(input.Action) == "" { + return ErrInvalid + } + if strings.TrimSpace(bigAccountID) == "" { + return ErrInvalid + } + big, err := s.GetAccountProfile(ctx, bigAccountID) + if err != nil { + return err + } + if strings.TrimSpace(big.ReplyRequirements) == "" { + return ErrInvalid + } + settings, err := s.GetSettings(ctx) + if err != nil { + return err + } + if !settings.AIConfigured || settings.AIProvider != "bailian" || strings.TrimSpace(settings.AIModel) == "" { + return ErrUnavailable + } + return nil +} + func (s *Store) CreateStrategy(ctx context.Context, bigAccountID string, input StrategyInput) (Strategy, error) { input, err := validateStrategyInput(input) if err != nil { return Strategy{}, err } + if err := s.validateEnabledStrategy(ctx, bigAccountID, input); err != nil { + return Strategy{}, err + } big, err := s.GetAccountProfile(ctx, bigAccountID) if err != nil { return Strategy{}, err @@ -150,18 +180,21 @@ func (s *Store) ListStrategies(ctx context.Context, bigID string) ([]Strategy, e return result, rows.Err() } func (s *Store) UpdateStrategy(ctx context.Context, id string, input StrategyInput) (Strategy, error) { - input, err := validateStrategyInput(input) + strategy, err := s.GetStrategy(ctx, id) if err != nil { return Strategy{}, err } + input, err = validateStrategyInput(input) + if err != nil { + return Strategy{}, err + } + if err := s.validateEnabledStrategy(ctx, strategy.BigAccountID, input); err != nil { + return Strategy{}, err + } events, texts, err := encodeStrategyLists(input) if err != nil { return Strategy{}, err } - strategy, err := s.GetStrategy(ctx, id) - if err != nil { - return Strategy{}, err - } if err := s.requireRelation(ctx, strategy.BigAccountID, input.ExecutionAccountID); err != nil { return Strategy{}, err } @@ -171,6 +204,15 @@ func (s *Store) UpdateStrategy(ctx context.Context, id string, input StrategyInp return s.GetStrategy(ctx, id) } func (s *Store) SetStrategyEnabled(ctx context.Context, id string, enabled bool) (Strategy, error) { + strategy, err := s.GetStrategy(ctx, id) + if err != nil { + return Strategy{}, err + } + if enabled { + if err := s.validateEnabledStrategy(ctx, strategy.BigAccountID, StrategyInput{Enabled: true, Action: strategy.Action, CandidateTexts: strategy.CandidateTexts}); err != nil { + return Strategy{}, err + } + } if _, err := s.db.ExecContext(ctx, `UPDATE creator_strategy SET enabled=$2,updated_at=now() WHERE id=$1`, id, enabled); err != nil { return Strategy{}, databaseError(err) } @@ -183,21 +225,21 @@ func (s *Store) DeleteStrategy(ctx context.Context, id string) error { func scanEvent(scanner interface{ Scan(...any) error }) (InteractionEvent, error) { var result InteractionEvent - var platformAt, receivedAt, startedAt, displayedAt sql.NullTime - var commentID, workID, strategyID, executionID sql.NullString - if err := scanner.Scan(&result.ID, &result.Platform, &result.ReceivingAccountID, &result.EventKey, &result.EventType, &result.InteractorUID, &commentID, &workID, &platformAt, &receivedAt, &startedAt, &displayedAt, &result.State, &result.Reason, &strategyID, &executionID); err != nil { + var platformAt, receivedAt, startedAt, finishedAt, displayedAt sql.NullTime + var commentID, workID, messageText, strategyID, executionID sql.NullString + if err := scanner.Scan(&result.ID, &result.Platform, &result.ReceivingAccountID, &result.EventKey, &result.EventType, &result.InteractorUID, &commentID, &workID, &messageText, &platformAt, &receivedAt, &startedAt, &finishedAt, &displayedAt, &result.State, &result.Reason, &strategyID, &executionID); err != nil { return InteractionEvent{}, err } - result.CommentID, result.WorkID = commentID.String, workID.String + result.CommentID, result.WorkID, result.MessageText = commentID.String, workID.String, messageText.String result.StrategyID, result.ExecutionAccountID = strategyID.String, executionID.String - result.PlatformEventAt, result.ProcessingStartedAt, result.DisplayedAt = nullableTime(platformAt), nullableTime(startedAt), nullableTime(displayedAt) + result.PlatformEventAt, result.ProcessingStartedAt, result.ProcessingFinishedAt, result.DisplayedAt = nullableTime(platformAt), nullableTime(startedAt), nullableTime(finishedAt), nullableTime(displayedAt) if receivedAt.Valid { result.ReceivedAt = receivedAt.Time.UTC() } return result, nil } -const eventSelect = `SELECT id,platform,receiving_account_id,event_key,event_type,interactor_uid,comment_id,work_id,platform_event_at,received_at,processing_started_at,displayed_at,state,reason,strategy_id,execution_account_id FROM creator_event` +const eventSelect = `SELECT id,platform,receiving_account_id,event_key,event_type,interactor_uid,comment_id,work_id,message_text,platform_event_at,received_at,processing_started_at,processing_finished_at,displayed_at,state,reason,strategy_id,execution_account_id FROM creator_event` func (s *Store) GetEvent(ctx context.Context, id string) (InteractionEvent, error) { result, err := scanEvent(s.db.QueryRowContext(ctx, eventSelect+` WHERE id=$1`, id)) @@ -228,8 +270,8 @@ func (s *Store) ListEvents(ctx context.Context, accountID string) ([]Interaction } func (s *Store) RecordEvent(ctx context.Context, input InteractionEvent) (AutomaticResult, error) { - input.Platform, input.ReceivingAccountID, input.EventKey, input.EventType, input.InteractorUID, input.CommentID, input.WorkID = strings.TrimSpace(input.Platform), strings.TrimSpace(input.ReceivingAccountID), strings.TrimSpace(input.EventKey), strings.TrimSpace(input.EventType), strings.TrimSpace(input.InteractorUID), strings.TrimSpace(input.CommentID), strings.TrimSpace(input.WorkID) - if !ValidatePlatform(input.Platform) || input.ReceivingAccountID == "" || input.EventKey == "" || !ValidEventType(input.EventType) { + input.Platform, input.ReceivingAccountID, input.EventKey, input.EventType, input.InteractorUID, input.CommentID, input.WorkID, input.MessageText = strings.TrimSpace(input.Platform), strings.TrimSpace(input.ReceivingAccountID), strings.TrimSpace(input.EventKey), strings.TrimSpace(input.EventType), strings.TrimSpace(input.InteractorUID), strings.TrimSpace(input.CommentID), strings.TrimSpace(input.WorkID), strings.TrimSpace(input.MessageText) + if !ValidatePlatform(input.Platform) || input.ReceivingAccountID == "" || input.EventKey == "" || !ValidEventType(input.EventType) || len(input.MessageText) > 100000 { return AutomaticResult{}, ErrInvalid } profile, err := s.GetAccountProfile(ctx, input.ReceivingAccountID) @@ -244,15 +286,22 @@ func (s *Store) RecordEvent(ctx context.Context, input InteractionEvent) (Automa reason := "" if input.Baseline { state = "baseline" - reason = "监听基线" + reason = input.BaselineReason + if reason == "" { + reason = "监听基线" + } } - if input.EventType != "dm" && input.InteractorUID == "" { + if !input.Baseline && input.EventType != "dm" && input.InteractorUID == "" { state = "blocked" reason = "缺少互动用户 UID" } + receivedAt := input.ReceivedAt.UTC() + if receivedAt.IsZero() { + receivedAt = time.Now().UTC() + } var returnedID string var inserted bool - err = s.db.QueryRowContext(ctx, `INSERT INTO creator_event (id,platform,receiving_account_id,event_key,event_type,interactor_uid,comment_id,work_id,platform_event_at,state,reason) VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11) ON CONFLICT (platform,receiving_account_id,event_key) DO NOTHING RETURNING id,(xmax=0)`, id, input.Platform, input.ReceivingAccountID, input.EventKey, input.EventType, input.InteractorUID, input.CommentID, input.WorkID, input.PlatformEventAt, state, reason).Scan(&returnedID, &inserted) + err = s.db.QueryRowContext(ctx, `INSERT INTO creator_event (id,platform,receiving_account_id,event_key,event_type,interactor_uid,comment_id,work_id,message_text,platform_event_at,received_at,state,reason) VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13) ON CONFLICT (platform,receiving_account_id,event_key) DO NOTHING RETURNING id,(xmax=0)`, id, input.Platform, input.ReceivingAccountID, input.EventKey, input.EventType, input.InteractorUID, input.CommentID, input.WorkID, input.MessageText, input.PlatformEventAt, receivedAt, state, reason).Scan(&returnedID, &inserted) if errors.Is(err, sql.ErrNoRows) { existingErr := s.db.QueryRowContext(ctx, `SELECT id FROM creator_event WHERE platform=$1 AND receiving_account_id=$2 AND event_key=$3`, input.Platform, input.ReceivingAccountID, input.EventKey).Scan(&returnedID) if existingErr != nil { @@ -268,11 +317,11 @@ func (s *Store) RecordEvent(ctx context.Context, input InteractionEvent) (Automa return AutomaticResult{Event: event, Duplicate: !inserted}, err } -func (s *Store) markEvent(ctx context.Context, eventID, state, reason, strategyID, executionID string, started, displayed *time.Time) error { +func (s *Store) markEvent(ctx context.Context, eventID, state, reason, strategyID, executionID string, started, finished, displayed *time.Time) error { if state != "received" && state != "baseline" && state != "ignored" && state != "unmatched" && state != "blocked" && state != "processing" && state != "succeeded" && state != "failed" && state != "uncertain" { return ErrInvalid } - _, err := s.db.ExecContext(ctx, `UPDATE creator_event SET state=$2,reason=$3,strategy_id=NULLIF($4,''),execution_account_id=NULLIF($5,''),processing_started_at=COALESCE($6,processing_started_at),displayed_at=COALESCE($7,displayed_at) WHERE id=$1`, eventID, state, reason, strategyID, executionID, started, displayed) + _, err := s.db.ExecContext(ctx, `UPDATE creator_event SET state=$2,reason=$3,strategy_id=NULLIF($4,''),execution_account_id=NULLIF($5,''),processing_started_at=COALESCE($6,processing_started_at),processing_finished_at=COALESCE($7,processing_finished_at),displayed_at=COALESCE($8,displayed_at) WHERE id=$1`, eventID, state, reason, strategyID, executionID, started, finished, displayed) return databaseError(err) } @@ -388,8 +437,8 @@ func (s *Store) ProcessAutomaticEvent(ctx context.Context, input InteractionEven } now := time.Now().UTC() cooldownSeconds := bigProfile.CooldownSeconds - if cooldownSeconds <= 0 { - cooldownSeconds = 86400 + if !validCooldownSeconds(cooldownSeconds) { + return AutomaticResult{}, ErrInvalid } expires := now.Add(time.Duration(cooldownSeconds) * time.Second) var cooldownID string @@ -423,7 +472,7 @@ func (s *Store) ProcessAutomaticEvent(ctx context.Context, input InteractionEven } return AutomaticResult{Event: event}, selectionErr } - if text == "" { + if ActionRequiresText(chosen.Action) && text == "" { if generator == nil { if _, err := tx.ExecContext(ctx, `UPDATE creator_event SET state='blocked',reason='AI 生成不可用',strategy_id=$2,execution_account_id=$3 WHERE id=$1`, eventID, chosen.ID, execution.ID); err != nil { return AutomaticResult{}, databaseError(err) @@ -472,24 +521,57 @@ func (s *Store) ProcessAutomaticEvent(ctx context.Context, input InteractionEven if _, err := tx.ExecContext(ctx, `INSERT INTO creator_operation (id,idempotency_key,source,action,platform,account_id,target_uid,target_comment_id,target_work_id,text,event_id,strategy_id,request_hash,state) VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13,'processing')`, opID, opInput.IdempotencyKey, opInput.Source, opInput.Action, opInput.Platform, opInput.AccountID, opInput.TargetUID, opInput.TargetCommentID, opInput.TargetWorkID, opInput.Text, opInput.EventID, opInput.StrategyID, hash); err != nil { return AutomaticResult{}, databaseError(err) } - started := now - if _, err := tx.ExecContext(ctx, `UPDATE creator_event SET state='processing',reason='',strategy_id=$2,execution_account_id=$3,processing_started_at=$4 WHERE id=$1`, eventID, chosen.ID, execution.ID, started); err != nil { + if _, err := tx.ExecContext(ctx, `UPDATE creator_event SET state='processing',reason='',strategy_id=$2,execution_account_id=$3 WHERE id=$1`, eventID, chosen.ID, execution.ID); err != nil { return AutomaticResult{}, databaseError(err) } if err := tx.Commit(); err != nil { return AutomaticResult{}, fmt.Errorf("commit automatic event: %w", err) } + // Conditions may change while the operation waits for the account executor. + // Re-check immediately before the platform write; a stale queued operation is blocked, never sent. + // The lock is deliberately acquired after receipt/operation persistence so ingestion + // never waits on a slow platform write, while one execution account remains serial. + executionLock := s.automaticExecutionLock(execution.ID) + executionLock.Lock() + defer executionLock.Unlock() + started := time.Now().UTC() + if _, err := s.db.ExecContext(ctx, `UPDATE creator_event SET processing_started_at=$2 WHERE id=$1`, eventID, started); err != nil { + return AutomaticResult{}, fmt.Errorf("save automatic event start: %w", databaseError(err)) + } result := ActionResult{} - if executor == nil { + _, checkErr := s.AccountWriteCheck(ctx, execution.ID, true, chosen.Action) + if checkErr == nil { + currentBig, bigErr := s.GetAccountProfile(ctx, input.ReceivingAccountID) + if bigErr != nil { + checkErr = bigErr + } else if !currentBig.BigAccount || currentBig.BusinessStatus != "normal" || currentBig.LoginStatus != "logged_in" || currentBig.AuthorizationStatus != "authorized" { + checkErr = ErrConflict + } + } + if checkErr == nil { + var related bool + checkErr = s.db.QueryRowContext(ctx, `SELECT EXISTS (SELECT 1 FROM creator_relation WHERE big_account_id=$1 AND small_account_id=$2)`, input.ReceivingAccountID, execution.ID).Scan(&related) + if checkErr == nil && !related { + checkErr = ErrConflict + } + } + if checkErr == nil { + latestStrategy, strategyErr := s.GetStrategy(ctx, chosen.ID) + if strategyErr != nil { + checkErr = strategyErr + } else if !latestStrategy.Enabled || latestStrategy.Action != chosen.Action || latestStrategy.ExecutionAccountID != execution.ID || !contains(latestStrategy.EventTypes, input.EventType) { + checkErr = ErrConflict + } + } + if checkErr != nil { + result = actionPreconditionResult(checkErr, "写入前条件已变化") + } else if executor == nil { result = ActionResult{State: "uncertain", Reason: "平台执行器不可用"} } else { result, err = executor.Execute(ctx, ActionRequest{OperationID: opID, Action: chosen.Action, Platform: input.Platform, AccountID: execution.ID, TargetUID: input.InteractorUID, TargetCommentID: input.CommentID, TargetWorkID: input.WorkID, Text: text}) - if err != nil && result.State == "" { - result.State = "uncertain" - result.Reason = err.Error() - } + result = normalizeActionResult(result, err) } - if result.State != "succeeded" && result.State != "failed" && result.State != "uncertain" { + if result.State != "succeeded" && result.State != "failed" && result.State != "uncertain" && result.State != "blocked" { result.State = "uncertain" if result.Reason == "" { result.Reason = "执行器未返回明确结果" @@ -498,7 +580,8 @@ func (s *Store) ProcessAutomaticEvent(ctx context.Context, input InteractionEven if err := s.UpdateOperationResult(ctx, opID, result); err != nil { return AutomaticResult{}, fmt.Errorf("save automatic operation result: %w", err) } - if err := s.markEvent(ctx, eventID, result.State, result.Reason, chosen.ID, execution.ID, nil, ptrTime(time.Now().UTC())); err != nil { + finished := time.Now().UTC() + if err := s.markEvent(ctx, eventID, result.State, result.Reason, chosen.ID, execution.ID, &started, &finished, nil); err != nil { return AutomaticResult{}, fmt.Errorf("save automatic event result: %w", err) } event, err := s.GetEvent(ctx, eventID) @@ -512,6 +595,62 @@ func (s *Store) ProcessAutomaticEvent(ctx context.Context, input InteractionEven return AutomaticResult{Event: event, Operation: &op}, nil } +const staleProcessingAfter = 2 * time.Minute + +func (s *Store) RecoverStaleProcessing(ctx context.Context, now time.Time) (int, error) { + if now.IsZero() { + return 0, ErrInvalid + } + cutoff := now.UTC().Add(-staleProcessingAfter) + tx, err := s.db.BeginTx(ctx, nil) + if err != nil { + return 0, databaseError(err) + } + defer tx.Rollback() + rows, err := tx.QueryContext(ctx, ` + SELECT id,event_id + FROM creator_operation + WHERE state='processing' AND updated_at < $1 + FOR UPDATE SKIP LOCKED`, cutoff) + if err != nil { + return 0, databaseError(err) + } + const reason = "处理者失联,平台写入结果不明" + stale := make([]struct{ operationID, eventID string }, 0) + for rows.Next() { + var item struct{ operationID, eventID string } + if err := rows.Scan(&item.operationID, &item.eventID); err != nil { + rows.Close() + return 0, databaseError(err) + } + stale = append(stale, item) + } + if err := rows.Err(); err != nil { + rows.Close() + return 0, databaseError(err) + } + if err := rows.Close(); err != nil { + return 0, databaseError(err) + } + for _, item := range stale { + if _, err := tx.ExecContext(ctx, `UPDATE creator_operation SET state='uncertain',reason=$2,updated_at=$3 WHERE id=$1 AND state='processing'`, item.operationID, reason, now.UTC()); err != nil { + return 0, databaseError(err) + } + if item.eventID != "" { + if _, err := tx.ExecContext(ctx, `UPDATE creator_event SET state='uncertain',reason=$2,processing_finished_at=COALESCE(processing_finished_at,$3) WHERE id=$1 AND state='processing'`, item.eventID, reason, now.UTC()); err != nil { + return 0, databaseError(err) + } + } + } + if _, err := tx.ExecContext(ctx, `UPDATE creator_event SET state='uncertain',reason=$2,processing_finished_at=COALESCE(processing_finished_at,$3) WHERE state='processing' AND processing_started_at IS NOT NULL AND processing_started_at < $1 AND NOT EXISTS (SELECT 1 FROM creator_operation WHERE event_id=creator_event.id AND state='processing')`, cutoff, reason, now.UTC()); err != nil { + return 0, databaseError(err) + } + if err := tx.Commit(); err != nil { + return 0, databaseError(err) + } + return len(stale), nil +} + func contains(values []string, want string) bool { for _, value := range values { if value == want { @@ -596,15 +735,43 @@ func (s *Store) ListOperations(ctx context.Context, accountID string) ([]Operati } return result, rows.Err() } -func (s *Store) validateOperationTarget(ctx context.Context, input OperationInput) error { +func actionPreconditionResult(err error, reason string) ActionResult { + if errors.Is(err, ErrInvalid) || errors.Is(err, ErrConflict) || errors.Is(err, ErrNotFound) { + return ActionResult{State: "blocked", Reason: reason} + } + return ActionResult{State: "uncertain", Reason: "写入前检查失败: " + err.Error()} +} + +func normalizeActionResult(result ActionResult, execErr error) ActionResult { + if execErr != nil { + if result.State != "failed" && result.State != "blocked" { + result.State = "uncertain" + } + if result.Reason == "" { + result.Reason = execErr.Error() + } + } + if result.State != "succeeded" && result.State != "failed" && result.State != "uncertain" && result.State != "blocked" { + result.State = "uncertain" + if result.Reason == "" { + result.Reason = "执行器未返回明确结果" + } + } + return result +} + +func (s *Store) validateOperationTarget(ctx context.Context, input *OperationInput) error { if input.TargetCommentID != "" { - var platform, authorUID string - if err := s.db.QueryRowContext(ctx, `SELECT platform,author_uid FROM creator_comment WHERE id=$1`, input.TargetCommentID).Scan(&platform, &authorUID); err != nil { + var platform, authorUID, commentWorkID string + if err := s.db.QueryRowContext(ctx, `SELECT platform,author_uid,work_id FROM creator_comment WHERE id=$1`, input.TargetCommentID).Scan(&platform, &authorUID, &commentWorkID); err != nil { return rowError(err) } - if platform != input.Platform || authorUID != input.TargetUID { + if platform != input.Platform || authorUID != input.TargetUID || input.TargetWorkID != "" && commentWorkID != input.TargetWorkID { return ErrInvalid } + if input.TargetWorkID == "" { + input.TargetWorkID = commentWorkID + } } if input.TargetWorkID != "" { var platform string @@ -651,7 +818,7 @@ func (s *Store) CreateOperation(ctx context.Context, input OperationInput) (Oper if profile.Platform != input.Platform { return Operation{}, false, ErrInvalid } - if err := s.validateOperationTarget(ctx, input); err != nil { + if err := s.validateOperationTarget(ctx, &input); err != nil { return Operation{}, false, err } id := newID("operation") @@ -695,8 +862,23 @@ func (s *Store) ExecuteManualOperation(ctx context.Context, id string, executor if op.State != "created" { return op, nil } - if _, err := s.AccountWriteCheck(ctx, op.AccountID, false, op.Action); err != nil { - if updateErr := s.UpdateOperationResult(ctx, id, ActionResult{State: "blocked", Reason: err.Error()}); updateErr != nil { + var claimedID string + if err := s.db.QueryRowContext(ctx, `UPDATE creator_operation SET state='processing',updated_at=now() WHERE id=$1 AND state='created' RETURNING id`, id).Scan(&claimedID); errors.Is(err, sql.ErrNoRows) { + return s.GetOperation(ctx, id) + } else if err != nil { + return Operation{}, databaseError(err) + } + executionLock := s.automaticExecutionLock(op.AccountID) + executionLock.Lock() + defer executionLock.Unlock() + if _, checkErr := s.AccountWriteCheck(ctx, op.AccountID, false, op.Action); checkErr != nil { + if updateErr := s.UpdateOperationResult(ctx, id, actionPreconditionResult(checkErr, "写入前条件已变化")); updateErr != nil { + return Operation{}, updateErr + } + return s.GetOperation(ctx, id) + } + if checkErr := s.validateOperationTarget(ctx, &OperationInput{Platform: op.Platform, AccountID: op.AccountID, Action: op.Action, TargetUID: op.TargetUID, TargetCommentID: op.TargetCommentID, TargetWorkID: op.TargetWorkID}); checkErr != nil { + if updateErr := s.UpdateOperationResult(ctx, id, actionPreconditionResult(checkErr, "写入目标已变化")); updateErr != nil { return Operation{}, updateErr } return s.GetOperation(ctx, id) @@ -707,23 +889,21 @@ func (s *Store) ExecuteManualOperation(ctx context.Context, id string, executor } return s.GetOperation(ctx, id) } - var claimedID string - if err := s.db.QueryRowContext(ctx, `UPDATE creator_operation SET state='processing',updated_at=now() WHERE id=$1 AND state='created' RETURNING id`, id).Scan(&claimedID); errors.Is(err, sql.ErrNoRows) { - return s.GetOperation(ctx, id) - } else if err != nil { - return Operation{}, databaseError(err) - } result, execErr := executor.Execute(ctx, ActionRequest{OperationID: claimedID, Action: op.Action, Platform: op.Platform, AccountID: op.AccountID, TargetUID: op.TargetUID, TargetCommentID: op.TargetCommentID, TargetWorkID: op.TargetWorkID, Text: op.Text}) - if execErr != nil && result.State == "" { - result.State = "uncertain" - result.Reason = execErr.Error() - } - if result.State == "" { - result.State = "uncertain" - } + result = normalizeActionResult(result, execErr) if err := s.UpdateOperationResult(ctx, id, result); err != nil { return Operation{}, err } + if op.Action == "dm" { + messageState := result.State + if messageState == "blocked" { + messageState = "failed" + } + messageAt := time.Now().UTC() + if _, _, messageErr := s.SaveMessage(ctx, MessageInput{Platform: op.Platform, AccountID: op.AccountID, PeerUID: op.TargetUID, PlatformMessageKey: "operation:" + op.ID, Direction: "outbound", MessageType: "text", Text: op.Text, SentState: messageState, MessageAt: &messageAt}); messageErr != nil { + return Operation{}, fmt.Errorf("persist direct message result: %w", messageErr) + } + } return s.GetOperation(ctx, id) } @@ -752,7 +932,7 @@ func (s *Store) UpsertConversation(ctx context.Context, input MessageInput) (Con } id := newID("conversation") var returned string - if err := s.db.QueryRowContext(ctx, `INSERT INTO creator_conversation (id,platform,account_id,peer_uid,peer_name,last_message_at) VALUES ($1,$2,$3,$4,$5,$6) ON CONFLICT (account_id,peer_uid) DO UPDATE SET peer_name=CASE WHEN EXCLUDED.peer_name='' THEN creator_conversation.peer_name ELSE EXCLUDED.peer_name END, last_message_at=COALESCE(EXCLUDED.last_message_at,creator_conversation.last_message_at) RETURNING id`, id, input.Platform, input.AccountID, input.PeerUID, input.PeerName, input.MessageAt).Scan(&returned); err != nil { + if err := s.db.QueryRowContext(ctx, `INSERT INTO creator_conversation (id,platform,account_id,peer_uid,peer_name,last_message_at) VALUES ($1,$2,$3,$4,$5,$6) ON CONFLICT (account_id,peer_uid) DO UPDATE SET peer_name=CASE WHEN EXCLUDED.peer_name='' THEN creator_conversation.peer_name ELSE EXCLUDED.peer_name END, last_message_at=CASE WHEN EXCLUDED.last_message_at IS NULL THEN creator_conversation.last_message_at WHEN creator_conversation.last_message_at IS NULL OR EXCLUDED.last_message_at > creator_conversation.last_message_at THEN EXCLUDED.last_message_at ELSE creator_conversation.last_message_at END RETURNING id`, id, input.Platform, input.AccountID, input.PeerUID, input.PeerName, input.MessageAt).Scan(&returned); err != nil { return Conversation{}, databaseError(err) } return s.GetConversation(ctx, returned) diff --git a/internal/creator/bailian.go b/internal/creator/bailian.go index bba9ed9..80f94ca 100644 --- a/internal/creator/bailian.go +++ b/internal/creator/bailian.go @@ -108,13 +108,16 @@ func (c *BailianClient) MatchTheme(ctx context.Context, title, body, topic strin return false, "", err } var result struct { - Match bool `json:"match"` + Match *bool `json:"match"` Reason string `json:"reason"` } - if err := json.Unmarshal([]byte(content), &result); err != nil { - return false, "", fmt.Errorf("decode bailian theme result: %w", err) + if err := json.Unmarshal([]byte(content), &result); err != nil || result.Match == nil { + if err != nil { + return false, "", fmt.Errorf("decode bailian theme result: %w", err) + } + return false, "", fmt.Errorf("decode bailian theme result: match is required") } - return result.Match, strings.TrimSpace(result.Reason), nil + return *result.Match, strings.TrimSpace(result.Reason), nil } func (c *BailianClient) MatchLead(ctx context.Context, work, comment, requirement string) (bool, string, error) { @@ -125,13 +128,16 @@ func (c *BailianClient) MatchLead(ctx context.Context, work, comment, requiremen return false, "", err } var result struct { - Match bool `json:"match"` + Match *bool `json:"match"` Reason string `json:"reason"` } - if err := json.Unmarshal([]byte(content), &result); err != nil { - return false, "", fmt.Errorf("decode bailian lead result: %w", err) + if err := json.Unmarshal([]byte(content), &result); err != nil || result.Match == nil { + if err != nil { + return false, "", fmt.Errorf("decode bailian lead result: %w", err) + } + return false, "", fmt.Errorf("decode bailian lead result: match is required") } - return result.Match, strings.TrimSpace(result.Reason), nil + return *result.Match, strings.TrimSpace(result.Reason), nil } type ConfiguredBailian struct { diff --git a/internal/creator/bailian_test.go b/internal/creator/bailian_test.go index c996114..faa818a 100644 --- a/internal/creator/bailian_test.go +++ b/internal/creator/bailian_test.go @@ -51,3 +51,17 @@ func TestBailianClientRejectsIncompleteConfiguration(t *testing.T) { t.Fatal("expected missing model to fail") } } + +func TestBailianClientRejectsMissingDecisionField(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + _, _ = w.Write([]byte(`{"choices":[{"message":{"content":"{}"}}]}`)) + })) + defer server.Close() + client, err := NewBailianClient(server.URL, "test-key", "qwen-test", server.Client()) + if err != nil { + t.Fatal(err) + } + if _, _, err := client.MatchLead(context.Background(), "work", "comment", "requirement"); err == nil { + t.Fatal("expected missing match decision to fail closed") + } +} diff --git a/internal/creator/collection.go b/internal/creator/collection.go index cac62a1..3890256 100644 --- a/internal/creator/collection.go +++ b/internal/creator/collection.go @@ -154,10 +154,76 @@ func (s *Store) checkpoint(ctx context.Context, sourceType, sourceID, kind strin return state, nil } -func (s *Store) saveCheckpointCursor(ctx context.Context, sourceType, sourceID, kind, leaseToken, cursor string) error { +// NextCollectionWindow preserves the fixed schedule grid. A failed attempt +// retries its exact window; a completed attempt advances from the prior +// scheduled end rather than from the wall-clock completion time. +func (s *Store) NextCollectionWindow(ctx context.Context, sourceType, sourceID string, now time.Time, interval time.Duration, lookbackDays int) (time.Time, time.Time, error) { + if sourceType != SourceOwned && sourceType != SourceCompetitor || sourceID == "" || now.IsZero() || interval <= 0 { + return time.Time{}, time.Time{}, ErrInvalid + } + var end time.Time + var status string + err := s.db.QueryRowContext(ctx, `SELECT window_end,status FROM creator_collection_checkpoint WHERE source_type=$1 AND source_id=$2 AND collection_kind='works'`, sourceType, sourceID).Scan(&end, &status) + if errors.Is(err, sql.ErrNoRows) { + return NewCollectionWindow(now, lookbackDays) + } + if err != nil { + return time.Time{}, time.Time{}, databaseError(err) + } + end = end.UTC() + if status == "failed" || status == "blocked" || status == "running" { + return end.Add(-time.Duration(lookbackDays) * 24 * time.Hour), end, nil + } + nextEnd := NextFixedRun(end, now.UTC(), interval) + return nextEnd.Add(-time.Duration(lookbackDays) * 24 * time.Hour), nextEnd, nil +} + +func (s *Store) MarkCollectionBlocked(ctx context.Context, sourceType, sourceID, reason string, now time.Time, lookbackDays int) error { + start, end, err := NewCollectionWindow(now, lookbackDays) + if err != nil { + return err + } + if len(reason) > 2000 { + return ErrInvalid + } + for _, kind := range []string{"works", "comments"} { + _, err := s.db.ExecContext(ctx, ` + INSERT INTO creator_collection_checkpoint + (id,source_type,source_id,collection_kind,window_start,window_end,status,lease_token,lease_until,last_error) + VALUES ($1,$2,$3,$4,$5,$6,'blocked','',NULL,$7) + ON CONFLICT (source_type,source_id,collection_kind) DO UPDATE SET + status='blocked', lease_token='', lease_until=NULL, last_error=$7 + WHERE creator_collection_checkpoint.status <> 'running' + OR creator_collection_checkpoint.lease_until IS NULL + OR creator_collection_checkpoint.lease_until <= now()`, + checkpointID(sourceType, sourceID, kind), sourceType, sourceID, kind, start, end, reason) + if err != nil { + return databaseError(err) + } + } + return nil +} + +func (s *Store) renewCheckpoint(ctx context.Context, sourceType, sourceID, kind, leaseToken string) error { if leaseToken == "" { return ErrInvalid } + result, err := s.db.ExecContext(ctx, `UPDATE creator_collection_checkpoint SET lease_until=now()+interval '10 minutes' WHERE source_type=$1 AND source_id=$2 AND collection_kind=$3 AND lease_token=$4 AND status='running' AND lease_until > now()`, sourceType, sourceID, kind, leaseToken) + if err != nil { + return databaseError(err) + } + if affected, err := result.RowsAffected(); err != nil { + return err + } else if affected != 1 { + return ErrConflict + } + return nil +} + +func (s *Store) saveCheckpointCursor(ctx context.Context, sourceType, sourceID, kind, leaseToken, cursor string) error { + if err := s.renewCheckpoint(ctx, sourceType, sourceID, kind, leaseToken); err != nil { + return err + } result, err := s.db.ExecContext(ctx, `UPDATE creator_collection_checkpoint SET cursor=$5 WHERE source_type=$1 AND source_id=$2 AND collection_kind=$3 AND lease_token=$4 AND status='running' AND lease_until > now()`, sourceType, sourceID, kind, leaseToken, cursor) if err != nil { return databaseError(err) @@ -187,7 +253,11 @@ func (s *Store) finishCheckpoint(ctx context.Context, sourceType, sourceID, kind } func (s *Store) failCheckpoint(ctx context.Context, sourceType, sourceID, kind, leaseToken string, primary error) error { - if err := s.finishCheckpoint(ctx, sourceType, sourceID, kind, leaseToken, "failed", primary.Error()); err != nil { + status := "failed" + if errors.Is(primary, ErrConflict) || errors.Is(primary, ErrUnavailable) { + status = "blocked" + } + if err := s.finishCheckpoint(ctx, sourceType, sourceID, kind, leaseToken, status, primary.Error()); err != nil { return errors.Join(primary, err) } return primary @@ -255,12 +325,13 @@ func (s *Store) ListDueOwnedAccounts(ctx context.Context, now time.Time, interva rows, err := s.db.QueryContext(ctx, ` SELECT account.id FROM social_account account - JOIN creator_account_profile profile ON profile.account_id=account.id AND profile.big_account=true + JOIN creator_account_profile profile ON profile.account_id=account.id AND profile.business_status='normal' LEFT JOIN creator_collection_checkpoint works_checkpoint ON works_checkpoint.source_type='owned' AND works_checkpoint.source_id=account.id AND works_checkpoint.collection_kind='works' LEFT JOIN creator_collection_checkpoint comments_checkpoint ON comments_checkpoint.source_type='owned' AND comments_checkpoint.source_id=account.id AND comments_checkpoint.collection_kind='comments' WHERE account.platform='douyin' AND account.authorization_status='authorized' + AND profile.login_status='logged_in' AND profile.big_account=true AND COALESCE(works_checkpoint.status, '') <> 'blocked' AND COALESCE(comments_checkpoint.status, '') <> 'blocked' AND (works_checkpoint.id IS NULL OR comments_checkpoint.id IS NULL @@ -314,9 +385,15 @@ func (s *Store) CollectSource(ctx context.Context, platform, sourceType, sourceI } return page.Items, page.NextCursor, page.HasMore, nil }, func(pageItems []WorkInput, nextCursor string, hasMore bool) error { + if err := s.renewCheckpoint(ctx, sourceType, sourceID, "works", worksLease); err != nil { + return err + } for _, work := range pageItems { + if err := s.renewCheckpoint(ctx, sourceType, sourceID, "works", worksLease); err != nil { + return err + } report.WorksSeen++ - if work.PublishedAt != nil && !InWindow(*work.PublishedAt, report.WindowStart, report.WindowEnd) { + if work.PublishedAt != nil && work.PublishedAt.Before(report.WindowStart) { continue } work.Platform, work.SourceType, work.SourceID = platform, sourceType, sourceID @@ -327,8 +404,10 @@ func (s *Store) CollectSource(ctx context.Context, platform, sourceType, sourceI if err != nil { return err } - if err := s.EnsureMetricPlan(ctx, savedWork.ID, settings); err != nil { - return err + if work.PublishedAt == nil || !work.PublishedAt.After(report.WindowEnd) { + if err := s.EnsureMetricPlan(ctx, savedWork.ID, settings); err != nil { + return err + } } if work.Likes != nil || work.CommentsCount != nil || work.Shares != nil { if _, metricErr := s.RecordMetric(ctx, MetricInput{WorkID: savedWork.ID, CollectedAt: now.UTC(), Likes: work.Likes, CommentsCount: work.CommentsCount, Shares: work.Shares}, settings, now.UTC()); metricErr != nil && !errors.Is(metricErr, ErrConflict) { @@ -395,7 +474,13 @@ func (s *Store) CollectSource(ctx context.Context, platform, sourceType, sourceI } return page.Items, page.NextCursor, page.HasMore, nil }, func(pageItems []CommentInput, nextCursor string, hasMore bool) error { + if err := s.renewCheckpoint(ctx, sourceType, sourceID, "comments", commentsLease); err != nil { + return err + } for _, comment := range pageItems { + if err := s.renewCheckpoint(ctx, sourceType, sourceID, "comments", commentsLease); err != nil { + return err + } report.CommentsSeen++ if comment.PublishedAt != nil && !InWindow(*comment.PublishedAt, report.WindowStart, report.WindowEnd) || comment.CommentType == "reply" { continue diff --git a/internal/creator/content.go b/internal/creator/content.go index a0387c3..a9a496e 100644 --- a/internal/creator/content.go +++ b/internal/creator/content.go @@ -106,7 +106,10 @@ func (s *Store) SetCompetitorEnabled(ctx context.Context, id string, enabled boo return s.GetCompetitor(ctx, id) } -func (s *Store) MarkCompetitorSync(ctx context.Context, id, status, cursor, syncError string, nextAt *time.Time) error { +func (s *Store) MarkCompetitorSync(ctx context.Context, id, leaseToken, status, cursor, syncError string, nextAt *time.Time) error { + if id == "" || leaseToken == "" { + return ErrInvalid + } if status != "idle" && status != "running" && status != "paused" && status != "failed" && status != "blocked" { return ErrInvalid } @@ -117,33 +120,43 @@ func (s *Store) MarkCompetitorSync(ctx context.Context, id, status, cursor, sync if nextAt != nil { next = nextAt.UTC() } - _, err := s.db.ExecContext(ctx, ` + result, err := s.db.ExecContext(ctx, ` UPDATE creator_competitor - SET sync_status = $2, sync_cursor = $3, sync_error = $4, - sync_lease_until = CASE WHEN $2 = 'running' THEN now() + interval '10 minutes' ELSE NULL END, - last_sync_at = CASE WHEN $2 IN ('idle', 'failed', 'blocked') THEN now() ELSE last_sync_at END, - next_sync_at = $5, updated_at = now() - WHERE id = $1`, id, status, cursor, syncError, next) - return databaseError(err) + SET sync_status = $3, sync_cursor = $4, sync_error = $5, + sync_lease_token = CASE WHEN $3 = 'running' THEN $2 ELSE NULL END, + sync_lease_until = CASE WHEN $3 = 'running' THEN now() + interval '10 minutes' ELSE NULL END, + last_sync_at = CASE WHEN $3 IN ('idle', 'failed', 'blocked') THEN now() ELSE last_sync_at END, + next_sync_at = $6, updated_at = now() + WHERE id = $1 AND sync_lease_token = $2`, id, leaseToken, status, cursor, syncError, next) + if err != nil { + return databaseError(err) + } + if affected, err := result.RowsAffected(); err != nil { + return databaseError(err) + } else if affected != 1 { + return ErrConflict + } + return nil } -func (s *Store) ClaimCompetitorSync(ctx context.Context, id string, force bool, now time.Time) (bool, error) { +func (s *Store) ClaimCompetitorSync(ctx context.Context, id string, force bool, now time.Time) (string, bool, error) { if id == "" || now.IsZero() { - return false, ErrInvalid + return "", false, ErrInvalid } condition := `enabled AND (next_sync_at IS NULL OR next_sync_at <= $2)` if force { condition = `enabled` } + token := newID("competitor-lease") var claimed string - err := s.db.QueryRowContext(ctx, `UPDATE creator_competitor SET sync_status='running', sync_lease_until=$2 + interval '10 minutes', sync_error='', updated_at=$2 WHERE id=$1 AND `+condition+` AND (sync_status <> 'running' OR sync_lease_until IS NULL OR sync_lease_until <= $2) RETURNING id`, id, now.UTC()).Scan(&claimed) + err := s.db.QueryRowContext(ctx, `UPDATE creator_competitor SET sync_status='running', sync_lease_token=$2, sync_lease_until=$3 + interval '10 minutes', sync_error='', updated_at=$3 WHERE id=$1 AND `+strings.ReplaceAll(condition, "$2", "$3")+` AND (sync_status <> 'running' OR sync_lease_until IS NULL OR sync_lease_until <= $3) RETURNING id`, id, token, now.UTC()).Scan(&claimed) if errors.Is(err, sql.ErrNoRows) { - return false, nil + return "", false, nil } if err != nil { - return false, databaseError(err) + return "", false, databaseError(err) } - return claimed != "", nil + return token, claimed != "", nil } func (s *Store) ListDueCompetitors(ctx context.Context, now time.Time) ([]Competitor, error) { @@ -370,19 +383,28 @@ func (s *Store) ListWorks(ctx context.Context, filter WorkFilter) ([]Work, error if err != nil { return nil, databaseError(err) } - defer rows.Close() result := make([]Work, 0) for rows.Next() { item, err := scanWork(rows) if err != nil { - return nil, err - } - if err := s.loadWorkSources(ctx, &item); err != nil { + _ = rows.Close() return nil, err } result = append(result, item) } - return result, rows.Err() + if err := rows.Err(); err != nil { + _ = rows.Close() + return nil, err + } + if err := rows.Close(); err != nil { + return nil, err + } + for index := range result { + if err := s.loadWorkSources(ctx, &result[index]); err != nil { + return nil, err + } + } + return result, nil } func (s *Store) RecordMetric(ctx context.Context, input MetricInput, settings Settings, now time.Time) (MetricPoint, error) { diff --git a/internal/creator/integration_test.go b/internal/creator/integration_test.go index e958153..d3c5113 100644 --- a/internal/creator/integration_test.go +++ b/internal/creator/integration_test.go @@ -152,10 +152,10 @@ func TestCreatorPostgresContentAndWorkflow(t *testing.T) { if _, err := store.UpdateAccountProfile(ctx, smallID, AccountProfileUpdate{RealNameStatus: "unknown", BusinessStatus: "normal", CooldownSeconds: 86400}); err != nil { t.Fatal(err) } - if _, err := store.RecordLoginResult(ctx, bigID, "logged_in", "", "sec_uid_"+bigID); err != nil { + if _, err := store.RecordVerifiedLoginResult(ctx, bigID, "sec_uid_"+bigID); err != nil { t.Fatal(err) } - if _, err := store.RecordLoginResult(ctx, smallID, "logged_in", "", "sec_uid_"+smallID); err != nil { + if _, err := store.RecordVerifiedLoginResult(ctx, smallID, "sec_uid_"+smallID); err != nil { t.Fatal(err) } if err := store.SetRelation(ctx, bigID, smallID, true); err != nil { @@ -199,11 +199,11 @@ func TestCreatorPostgresContentAndWorkflow(t *testing.T) { if due, err := store.ListDueCompetitors(ctx, dueNow); err != nil || len(due) != 1 { t.Fatalf("list due competitors: due=%+v err=%v", due, err) } - claimed, err := store.ClaimCompetitorSync(ctx, competitor.ID, false, dueNow) - if err != nil || !claimed { - t.Fatalf("claim competitor sync: claimed=%v err=%v", claimed, err) + leaseToken, claimed, err := store.ClaimCompetitorSync(ctx, competitor.ID, false, dueNow) + if err != nil || !claimed || leaseToken == "" { + t.Fatalf("claim competitor sync: token=%q claimed=%v err=%v", leaseToken, claimed, err) } - if err := store.MarkCompetitorSync(ctx, competitor.ID, "idle", "", "", nil); err != nil { + if err := store.MarkCompetitorSync(ctx, competitor.ID, leaseToken, "idle", "", "", nil); err != nil { t.Fatal(err) } work, inserted, err = store.UpsertWork(ctx, WorkInput{Platform: PlatformDouyin, WorkKey: workKey, SourceType: SourceCompetitor, SourceID: competitor.ID, Title: "", Body: "", PublishedAt: nil, Likes: nil, CommentsCount: nil, Shares: nil}, now) @@ -299,7 +299,7 @@ func prepareIntegrationActionFixture(t *testing.T, store *Store, phaseAStore *ph if _, err := store.UpdateAccountProfile(ctx, accountID, AccountProfileUpdate{RealNameStatus: "unknown", BusinessStatus: "normal", CooldownSeconds: 86400}); err != nil { t.Fatal(err) } - if _, err := store.RecordLoginResult(ctx, accountID, "logged_in", "", "sec_uid_"+accountID); err != nil { + if _, err := store.RecordVerifiedLoginResult(ctx, accountID, "sec_uid_"+accountID); err != nil { t.Fatal(err) } } @@ -371,6 +371,14 @@ func TestCreatorPostgresActionsAndMessaging(t *testing.T) { if err != nil || manual.State != "succeeded" { t.Fatalf("execute manual operation: operation=%+v err=%v", manual, err) } + dmOperation, inserted, err := store.CreateOperation(ctx, OperationInput{IdempotencyKey: "creator-it-dm-" + stamp, Source: "manual", Action: ActionDM, Platform: PlatformDouyin, AccountID: smallID, TargetUID: "peer-" + stamp, Text: "人工私信"}) + if err != nil || !inserted { + t.Fatalf("create direct-message operation: operation=%+v inserted=%v err=%v", dmOperation, inserted, err) + } + dmOperation, err = store.ExecuteManualOperation(ctx, dmOperation.ID, integrationExecutor{}) + if err != nil || dmOperation.State != "succeeded" { + t.Fatalf("execute direct-message operation: operation=%+v err=%v", dmOperation, err) + } messageAt := time.Now().UTC().Truncate(time.Microsecond) message, inserted, err := store.SaveMessage(ctx, MessageInput{Platform: PlatformDouyin, AccountID: smallID, PeerUID: "peer-" + stamp, PeerName: "Peer", PlatformMessageKey: "creator-it-message-" + stamp, Direction: "inbound", MessageType: "text", Text: "hello", MessageAt: &messageAt}) if err != nil || !inserted { @@ -384,7 +392,7 @@ func TestCreatorPostgresActionsAndMessaging(t *testing.T) { t.Fatalf("list conversations: conversations=%+v err=%v", conversations, err) } messages, err := store.ListMessages(ctx, conversations[0].ID) - if err != nil || len(messages) != 1 { + if err != nil || len(messages) != 2 { t.Fatalf("list messages: messages=%+v err=%v", messages, err) } if _, err := store.SetEventDisplayed(ctx, automatic.Event.ID, time.Now().UTC()); err != nil { diff --git a/internal/creator/logic.go b/internal/creator/logic.go index 2b02101..80e71cc 100644 --- a/internal/creator/logic.go +++ b/internal/creator/logic.go @@ -179,6 +179,10 @@ func ActionRequiresText(action string) bool { return action == ActionDM || action == ActionReplyComment || action == ActionRepost } +func validCooldownSeconds(seconds int64) bool { + return seconds > 0 && seconds <= maxDurationSeconds +} + func ActionTargetValid(action string, interactorUID, commentID, workID, targetType string) bool { if interactorUID == "" { return false diff --git a/internal/creator/metrics.go b/internal/creator/metrics.go index 1712329..4339a4d 100644 --- a/internal/creator/metrics.go +++ b/internal/creator/metrics.go @@ -8,6 +8,38 @@ import ( "time" ) +func (s *Store) ListDueMetricWorks(ctx context.Context, now time.Time) ([]Work, error) { + if now.IsZero() { + return nil, ErrInvalid + } + rows, err := s.db.QueryContext(ctx, workSelect+` WHERE published_at_status='verified' AND next_metric_at IS NOT NULL AND next_metric_at <= $1 ORDER BY next_metric_at,id`, now.UTC()) + if err != nil { + return nil, databaseError(err) + } + works := make([]Work, 0) + for rows.Next() { + work, scanErr := scanWork(rows) + if scanErr != nil { + _ = rows.Close() + return nil, scanErr + } + works = append(works, work) + } + if err := rows.Err(); err != nil { + _ = rows.Close() + return nil, err + } + if err := rows.Close(); err != nil { + return nil, err + } + for index := range works { + if err := s.loadWorkSources(ctx, &works[index]); err != nil { + return nil, err + } + } + return works, nil +} + func (s *Store) EnsureMetricPlan(ctx context.Context, workID string, settings Settings) error { if err := ValidateSettings(SettingsUpdate{LookbackDays: settings.LookbackDays, NewWorkIntervalSeconds: settings.NewWorkIntervalSeconds, MetricInitialIntervalSeconds: settings.MetricInitialIntervalSeconds, MetricMultiplier: settings.MetricMultiplier, MetricMaxIntervalSeconds: settings.MetricMaxIntervalSeconds, MetricAgeSeconds: settings.MetricAgeSeconds}); err != nil { return err diff --git a/internal/creator/migrations/023_password_references.sql b/internal/creator/migrations/023_password_references.sql new file mode 100644 index 0000000..e044890 --- /dev/null +++ b/internal/creator/migrations/023_password_references.sql @@ -0,0 +1,8 @@ +CREATE TABLE IF NOT EXISTS creator_account_password ( + account_id text PRIMARY KEY REFERENCES social_account ( + id + ) ON DELETE CASCADE, + secret_reference_id text NOT NULL, + secret_key text NOT NULL, + updated_at timestamptz NOT NULL DEFAULT now() +); diff --git a/internal/creator/migrations/024_event_processing_times.sql b/internal/creator/migrations/024_event_processing_times.sql new file mode 100644 index 0000000..d6b9203 --- /dev/null +++ b/internal/creator/migrations/024_event_processing_times.sql @@ -0,0 +1,2 @@ +ALTER TABLE creator_event + ADD COLUMN IF NOT EXISTS processing_finished_at timestamptz; diff --git a/internal/creator/migrations/025_competitor_sync_tokens.sql b/internal/creator/migrations/025_competitor_sync_tokens.sql new file mode 100644 index 0000000..d4bd1dc --- /dev/null +++ b/internal/creator/migrations/025_competitor_sync_tokens.sql @@ -0,0 +1,2 @@ +ALTER TABLE creator_competitor + ADD COLUMN IF NOT EXISTS sync_lease_token text; diff --git a/internal/creator/migrations/026_event_message_text.sql b/internal/creator/migrations/026_event_message_text.sql new file mode 100644 index 0000000..15aa7f7 --- /dev/null +++ b/internal/creator/migrations/026_event_message_text.sql @@ -0,0 +1,2 @@ +ALTER TABLE creator_event + ADD COLUMN IF NOT EXISTS message_text text NOT NULL DEFAULT ''; diff --git a/internal/creator/models.go b/internal/creator/models.go index 96b9a3c..b6b45d2 100644 --- a/internal/creator/models.go +++ b/internal/creator/models.go @@ -307,23 +307,26 @@ type Relation struct { } type InteractionEvent struct { - ID string `json:"id"` - Platform string `json:"platform"` - ReceivingAccountID string `json:"receiving_account_id"` - EventKey string `json:"event_key"` - EventType string `json:"event_type"` - Baseline bool `json:"baseline,omitempty"` - InteractorUID string `json:"interactor_uid"` - CommentID string `json:"comment_id"` - WorkID string `json:"work_id"` - PlatformEventAt *time.Time `json:"platform_event_at,omitempty"` - ReceivedAt time.Time `json:"received_at"` - ProcessingStartedAt *time.Time `json:"processing_started_at,omitempty"` - DisplayedAt *time.Time `json:"displayed_at,omitempty"` - State string `json:"state"` - Reason string `json:"reason,omitempty"` - StrategyID string `json:"strategy_id,omitempty"` - ExecutionAccountID string `json:"execution_account_id,omitempty"` + ID string `json:"id"` + Platform string `json:"platform"` + ReceivingAccountID string `json:"receiving_account_id"` + EventKey string `json:"event_key"` + EventType string `json:"event_type"` + Baseline bool `json:"baseline,omitempty"` + BaselineReason string `json:"-"` + InteractorUID string `json:"interactor_uid"` + CommentID string `json:"comment_id"` + WorkID string `json:"work_id"` + MessageText string `json:"message_text,omitempty"` + PlatformEventAt *time.Time `json:"platform_event_at,omitempty"` + ReceivedAt time.Time `json:"received_at"` + ProcessingStartedAt *time.Time `json:"processing_started_at,omitempty"` + ProcessingFinishedAt *time.Time `json:"processing_finished_at,omitempty"` + DisplayedAt *time.Time `json:"displayed_at,omitempty"` + State string `json:"state"` + Reason string `json:"reason,omitempty"` + StrategyID string `json:"strategy_id,omitempty"` + ExecutionAccountID string `json:"execution_account_id,omitempty"` } type Operation struct { diff --git a/internal/creator/recovery_integration_test.go b/internal/creator/recovery_integration_test.go new file mode 100644 index 0000000..adb3c0b --- /dev/null +++ b/internal/creator/recovery_integration_test.go @@ -0,0 +1,41 @@ +package creator + +import ( + "fmt" + "testing" + "time" +) + +func TestCreatorPostgresRecoversStaleProcessingWithoutRetry(t *testing.T) { + store, phaseAStore, ctx := openCreatorIntegrationStore(t) + stamp := time.Now().UnixNano() + bigID, smallID, work, comment, _ := prepareIntegrationActionFixture(t, store, phaseAStore, ctx, fmt.Sprintf("%d", stamp)) + recorded, err := store.RecordEvent(ctx, InteractionEvent{Platform: PlatformDouyin, ReceivingAccountID: bigID, EventKey: fmt.Sprintf("recovery-event-%d", stamp), EventType: "comment", InteractorUID: comment.AuthorUID, CommentID: comment.ID, WorkID: work.ID}) + if err != nil { + t.Fatalf("record event: result=%+v err=%v", recorded, err) + } + event := recorded.Event + op, inserted, err := store.CreateOperation(ctx, OperationInput{IdempotencyKey: fmt.Sprintf("recovery-op-%d", stamp), Source: "manual", Action: ActionReplyComment, Platform: PlatformDouyin, AccountID: smallID, TargetUID: comment.AuthorUID, TargetCommentID: comment.ID, TargetWorkID: work.ID, Text: "reply", EventID: event.ID}) + if err != nil || !inserted { + t.Fatalf("create operation: operation=%+v inserted=%v err=%v", op, inserted, err) + } + old := time.Now().UTC().Add(-3 * time.Minute) + if _, err := store.db.ExecContext(ctx, `UPDATE creator_operation SET source='automatic',state='processing',updated_at=$2 WHERE id=$1`, op.ID, old); err != nil { + t.Fatal(err) + } + if _, err := store.db.ExecContext(ctx, `UPDATE creator_event SET state='processing',processing_started_at=$2 WHERE id=$1`, event.ID, old); err != nil { + t.Fatal(err) + } + count, err := store.RecoverStaleProcessing(ctx, time.Now().UTC()) + if err != nil || count != 1 { + t.Fatalf("recover stale processing: count=%d err=%v", count, err) + } + recoveredOperation, err := store.GetOperation(ctx, op.ID) + if err != nil || recoveredOperation.State != "uncertain" { + t.Fatalf("operation was not made uncertain: operation=%+v err=%v", recoveredOperation, err) + } + recoveredEvent, err := store.GetEvent(ctx, event.ID) + if err != nil || recoveredEvent.State != "uncertain" { + t.Fatalf("event was not made uncertain: event=%+v err=%v", recoveredEvent, err) + } +} diff --git a/internal/creator/rules.go b/internal/creator/rules.go index 55addc1..1f26b77 100644 --- a/internal/creator/rules.go +++ b/internal/creator/rules.go @@ -139,8 +139,17 @@ func (s *Store) AnalyzeComment(ctx context.Context, commentID, ruleID string, an if err != nil { return RuleResult{}, err } - if rule.SourceType != "all" && rule.SourceType != work.SourceType { - return RuleResult{}, ErrConflict + if rule.SourceType != "all" { + matchedSource := false + for _, source := range work.Sources { + if source.SourceType == rule.SourceType { + matchedSource = true + break + } + } + if !matchedSource && work.SourceType != rule.SourceType { + return RuleResult{}, ErrConflict + } } if analyzer == nil { result, saveErr := s.upsertRuleResult(ctx, commentID, rule, "failed", "AI 分析不可用", nil, ptrTime(time.Now().UTC())) diff --git a/internal/creator/scheduler.go b/internal/creator/scheduler.go index 5112a5f..edefbcb 100644 --- a/internal/creator/scheduler.go +++ b/internal/creator/scheduler.go @@ -28,8 +28,8 @@ func NextFixedRun(lastCompleted, now time.Time, interval time.Duration) time.Tim return now.UTC() } next := lastCompleted.UTC().Add(interval) - if next.Before(now.UTC()) { - return now.UTC() + for !next.After(now.UTC()) { + next = next.Add(interval) } return next } diff --git a/internal/creator/scheduler_test.go b/internal/creator/scheduler_test.go index 4bd02aa..2623aba 100644 --- a/internal/creator/scheduler_test.go +++ b/internal/creator/scheduler_test.go @@ -11,7 +11,7 @@ func TestFixedScheduleUsesDueBoundary(t *testing.T) { if !IsDue(last, now, time.Hour) || IsDue(last.Add(time.Minute), now, time.Hour) { t.Fatal("schedule due check must include the exact boundary") } - if got := NextFixedRun(last, now, time.Hour); !got.Equal(now) { + if got := NextFixedRun(last, now, time.Hour); !got.Equal(now.Add(time.Hour)) { t.Fatalf("next fixed run = %s", got) } if got := NextFixedRun(time.Time{}, now, time.Hour); !got.Equal(now) { @@ -29,3 +29,20 @@ func TestMetricDuePreservesUnavailableState(t *testing.T) { t.Fatal("past metric schedule is due") } } + +func TestNextMetricAtSkipsLateHistoricalPoints(t *testing.T) { + published := time.Date(2026, 4, 1, 0, 0, 0, 0, time.UTC) + now := published.Add(8 * time.Hour) + got, reason := NextMetricAt(published, now, time.Hour, 24*time.Hour, 2, 48*time.Hour) + want := published.Add(15 * time.Hour) + if reason != "" || !got.Equal(want) { + t.Fatalf("late metric schedule = %s, reason %q; want %s", got, reason, want) + } +} + +func TestNextMetricAtStopsAtMonitoringAge(t *testing.T) { + published := time.Date(2026, 4, 1, 0, 0, 0, 0, time.UTC) + if got, reason := NextMetricAt(published, published.Add(47*time.Hour), time.Hour, 24*time.Hour, 2, 48*time.Hour); !got.IsZero() || reason != "monitoring_age_reached" { + t.Fatalf("expired metric schedule = %s, reason %q", got, reason) + } +} diff --git a/internal/creator/settings.go b/internal/creator/settings.go index dcd0462..0fe65c5 100644 --- a/internal/creator/settings.go +++ b/internal/creator/settings.go @@ -43,7 +43,12 @@ func (s *Store) UpdateSettings(ctx context.Context, input SettingsUpdate) (Setti if input.TranscriptionConfigured && (input.TranscriptionProvider == "" || input.TranscriptionModel == "") { return Settings{}, ErrInvalid } - _, err := s.db.ExecContext(ctx, ` + tx, err := s.db.BeginTx(ctx, nil) + if err != nil { + return Settings{}, fmt.Errorf("begin creator settings update: %w", err) + } + defer tx.Rollback() + if _, err := tx.ExecContext(ctx, ` UPDATE creator_settings SET lookback_days=$1, new_work_interval_seconds=$2, metric_initial_interval_seconds=$3, metric_multiplier=$4, metric_max_interval_seconds=$5, metric_age_seconds=$6, ai_provider=$7, ai_model=$8, ai_configured=$9, @@ -51,13 +56,70 @@ func (s *Store) UpdateSettings(ctx context.Context, input SettingsUpdate) (Setti updated_at=now() WHERE id=true`, input.LookbackDays, input.NewWorkIntervalSeconds, input.MetricInitialIntervalSeconds, input.MetricMultiplier, input.MetricMaxIntervalSeconds, input.MetricAgeSeconds, input.AIProvider, input.AIModel, input.AIConfigured, - input.TranscriptionProvider, input.TranscriptionModel, input.TranscriptionConfigured) + input.TranscriptionProvider, input.TranscriptionModel, input.TranscriptionConfigured); err != nil { + return Settings{}, databaseError(err) + } + rows, err := tx.QueryContext(ctx, ` + SELECT p.work_id, w.published_at + FROM creator_metric_plan p + JOIN creator_work w ON w.id=p.work_id + FOR UPDATE`) if err != nil { return Settings{}, databaseError(err) } + type metricSchedule struct { + workID string + publishedAt *time.Time + } + schedules := make([]metricSchedule, 0) + for rows.Next() { + var item metricSchedule + if err := rows.Scan(&item.workID, &item.publishedAt); err != nil { + _ = rows.Close() + return Settings{}, err + } + schedules = append(schedules, item) + } + if err := rows.Err(); err != nil { + _ = rows.Close() + return Settings{}, err + } + if err := rows.Close(); err != nil { + return Settings{}, err + } + now := time.Now().UTC() + for _, schedule := range schedules { + nextAt, reason := NextMetricAtValue(schedule.publishedAt, now, input) + stopped := nextAt.IsZero() + if _, err := tx.ExecContext(ctx, ` + UPDATE creator_metric_plan SET next_plan_at=$2, interval_seconds=$3, + multiplier=$4, max_interval_seconds=$5, stopped=$6, stop_reason=$7, updated_at=now() + WHERE work_id=$1`, schedule.workID, nullableArg(nextAt), input.MetricInitialIntervalSeconds, + input.MetricMultiplier, input.MetricMaxIntervalSeconds, stopped, reason); err != nil { + return Settings{}, databaseError(err) + } + if _, err := tx.ExecContext(ctx, ` + UPDATE creator_work SET next_metric_at=$2, metric_stop_reason=$3, updated_at=now() + WHERE id=$1`, schedule.workID, nullableArg(nextAt), reason); err != nil { + return Settings{}, databaseError(err) + } + } + if err := tx.Commit(); err != nil { + return Settings{}, fmt.Errorf("commit creator settings update: %w", err) + } return s.GetSettings(ctx) } +func NextMetricAtValue(publishedAt *time.Time, now time.Time, input SettingsUpdate) (time.Time, string) { + if publishedAt == nil { + return time.Time{}, "published_at_pending_verification" + } + return NextMetricAt(publishedAt.UTC(), now.UTC(), + time.Duration(input.MetricInitialIntervalSeconds)*time.Second, + time.Duration(input.MetricMaxIntervalSeconds)*time.Second, + input.MetricMultiplier, time.Duration(input.MetricAgeSeconds)*time.Second) +} + func (s *Store) SetEventDisplayed(ctx context.Context, eventID string, displayedAt time.Time) (InteractionEvent, error) { if displayedAt.IsZero() { displayedAt = time.Now().UTC() diff --git a/internal/creator/store.go b/internal/creator/store.go index a10e4d3..6fa559e 100644 --- a/internal/creator/store.go +++ b/internal/creator/store.go @@ -9,6 +9,7 @@ import ( "encoding/json" "errors" "fmt" + "sync" "time" "github.com/jackc/pgx/v5/pgconn" @@ -33,6 +34,18 @@ var migration021 string //go:embed migrations/022_collection_lease_tokens.sql var migration022 string +//go:embed migrations/023_password_references.sql +var migration023 string + +//go:embed migrations/024_event_processing_times.sql +var migration024 string + +//go:embed migrations/025_competitor_sync_tokens.sql +var migration025 string + +//go:embed migrations/026_event_message_text.sql +var migration026 string + type SecretReference struct { ID string Provider string @@ -46,6 +59,9 @@ type SecretBridge interface { type Store struct { db *sql.DB secrets SecretBridge + // Automatic writes for one execution account stay serialized while event + // ingestion and unrelated accounts remain concurrent. + automaticLocks sync.Map // map[string]*sync.Mutex } func Open(ctx context.Context, databaseURL string) (*Store, error) { @@ -72,6 +88,11 @@ func (s *Store) Close() error { return s.db.Close() } func (s *Store) SetSecretBridge(bridge SecretBridge) { s.secrets = bridge } +func (s *Store) automaticExecutionLock(accountID string) *sync.Mutex { + lock, _ := s.automaticLocks.LoadOrStore(accountID, &sync.Mutex{}) + return lock.(*sync.Mutex) +} + func (s *Store) migrate(ctx context.Context) error { tx, err := s.db.BeginTx(ctx, nil) if err != nil { @@ -94,6 +115,10 @@ func (s *Store) migrate(ctx context.Context) error { {version: 20, sql: migration020}, {version: 21, sql: migration021}, {version: 22, sql: migration022}, + {version: 23, sql: migration023}, + {version: 24, sql: migration024}, + {version: 25, sql: migration025}, + {version: 26, sql: migration026}, } for _, migration := range migrations { var applied bool diff --git a/internal/douyin/connector.go b/internal/douyin/connector.go index 69676f6..f16a348 100644 --- a/internal/douyin/connector.go +++ b/internal/douyin/connector.go @@ -126,27 +126,38 @@ type Request struct { } func (connector Connector) Sync(ctx context.Context, request Request) (Result, error) { - if connector.Browser == nil || connector.Secrets == nil || connector.Store == nil || !keyPattern.MatchString(request.AccountID) || + if connector.Browser == nil || connector.Store == nil || !keyPattern.MatchString(request.AccountID) || !keyPattern.MatchString(request.PlatformAccountKey) || (request.Credential.Provider != "os_keyring" && request.Credential.Provider != "secret_manager") || !credentialKeyPattern.MatchString(request.Credential.Key) { return Result{}, ErrInvalid } - credential, err := connector.Secrets.Resolve(ctx, request.Credential) - if err != nil { + if err := ctx.Err(); err != nil { return connector.stop(ctx, request.AccountID, StatePolicyHold, ReasonAuthInvalid, Evidence{Phase: "login"}) } - cookies, err := parseCredential(credential) - if err != nil { - return connector.stop(ctx, request.AccountID, StatePolicyHold, ReasonAuthInvalid, Evidence{Phase: "login"}) - } - if err := connector.Browser.SetCookies(ctx, cookies); err != nil { - return connector.stop(ctx, request.AccountID, StateNeedsConfirmation, ReasonUnknown, Evidence{Phase: "login"}) - } - identityResponse, err := connector.Browser.Get(ctx, identityEndpoint) if err != nil { - return connector.stop(ctx, request.AccountID, StateNeedsConfirmation, ReasonUnknown, Evidence{Phase: "identity"}) + if ctx.Err() != nil { + return connector.stop(ctx, request.AccountID, StatePolicyHold, ReasonAuthInvalid, Evidence{Phase: "login"}) + } + if connector.Secrets == nil { + return connector.stop(ctx, request.AccountID, StatePolicyHold, ReasonAuthInvalid, Evidence{Phase: "login"}) + } + credential, credentialErr := connector.Secrets.Resolve(ctx, request.Credential) + if credentialErr != nil { + return connector.stop(ctx, request.AccountID, StatePolicyHold, ReasonAuthInvalid, Evidence{Phase: "login"}) + } + cookies, credentialErr := ParseCredential(credential) + if credentialErr != nil { + return connector.stop(ctx, request.AccountID, StatePolicyHold, ReasonAuthInvalid, Evidence{Phase: "login"}) + } + if credentialErr = connector.Browser.SetCookies(ctx, cookies); credentialErr != nil { + return connector.stop(ctx, request.AccountID, StateNeedsConfirmation, ReasonUnknown, Evidence{Phase: "login"}) + } + identityResponse, err = connector.Browser.Get(ctx, identityEndpoint) + if err != nil { + return connector.stop(ctx, request.AccountID, StateNeedsConfirmation, ReasonUnknown, Evidence{Phase: "identity"}) + } } if state, reason := classify(identityResponse); state != "" { return connector.stop(ctx, request.AccountID, state, reason, Evidence{Phase: "identity", HTTPStatus: identityResponse.Status}) @@ -242,6 +253,15 @@ func classify(response Response) (string, string) { } } +// ParseCredential accepts either the persisted JSON cookie bundle or the +// manually captured Cookie header used by an already logged-in browser. +func ParseCredential(raw []byte) ([]Cookie, error) { + if cookies, err := ParseCookieHeader(raw); err == nil { + return cookies, nil + } + return ParseCookieBundle(raw) +} + func ParseCookieBundle(raw []byte) ([]Cookie, error) { return parseCredential(raw) } diff --git a/internal/douyin/connector_test.go b/internal/douyin/connector_test.go index 1048bc4..4561d65 100644 --- a/internal/douyin/connector_test.go +++ b/internal/douyin/connector_test.go @@ -13,11 +13,15 @@ const credential = `{"cookies":[{"name":"sessionid","value":"private-session","d var secretReference = SecretReference{Provider: "os_keyring", Key: "creatorhub/account-a"} type fakeSecrets struct { - value []byte - err error + value []byte + err error + resolveCalls *int } func (secrets fakeSecrets) Resolve(_ context.Context, _ SecretReference) ([]byte, error) { + if secrets.resolveCalls != nil { + (*secrets.resolveCalls)++ + } return secrets.value, secrets.err } @@ -77,15 +81,19 @@ func TestSyncLogsInVerifiesIdentityAndReadsOwnWorks(t *testing.T) { {Status: 200, Body: []byte(`{"status_code":0,"has_more":true,"aweme_list":[{"aweme_id":"work-1","desc":"hello","create_time":123,"statistics":{"digg_count":4,"comment_count":3,"share_count":2,"play_count":1}}]}`)}, }} store := &fakeStore{} - result, err := (Connector{Browser: browser, Secrets: fakeSecrets{value: []byte(credential)}, Store: store}).Sync(context.Background(), Request{ + resolveCalls := 0 + result, err := (Connector{Browser: browser, Secrets: fakeSecrets{value: []byte(credential), resolveCalls: &resolveCalls}, Store: store}).Sync(context.Background(), Request{ AccountID: "account-a", PlatformAccountKey: "sec-a", Credential: secretReference, }) if err != nil || result.State != StateSucceeded || !result.Evidence.IdentityVerified || result.Evidence.WorksSeen != 1 || !result.Evidence.HasMore { t.Fatalf("unexpected result: %#v err=%v", result, err) } - if len(browser.cookies) != 1 || browser.cookies[0].Value != "private-session" || len(browser.urls) != 2 || + if resolveCalls != 0 { + t.Fatalf("active browser session unexpectedly resolved stored credentials: calls=%d", resolveCalls) + } + if len(browser.cookies) != 0 || len(browser.urls) != 2 || browser.urls[0] != identityEndpoint || !strings.Contains(browser.urls[1], "sec_user_id=sec-a") { - t.Fatalf("connector did not use the bound browser session: cookies=%#v urls=%#v", browser.cookies, browser.urls) + t.Fatalf("connector did not reuse the bound browser session: cookies=%#v urls=%#v", browser.cookies, browser.urls) } if len(store.works) != 1 || store.works[0].ID != "work-1" || store.works[0].PlayCount != 1 || len(store.holds) != 0 { t.Fatalf("unexpected persisted works or hold: works=%#v holds=%#v", store.works, store.holds) @@ -180,12 +188,12 @@ func TestSyncFailsClosedOnIdentityAndUnknownResults(t *testing.T) { } func TestSyncRejectsInvalidCredentialWithoutLeakingIt(t *testing.T) { - browser := &fakeBrowser{} + browser := &fakeBrowser{err: errors.New("session unavailable")} store := &fakeStore{} result, err := (Connector{Browser: browser, Secrets: fakeSecrets{value: []byte(`{"cookies":[{"name":"sessionid","value":"secret","domain":"evil.example"}]}`)}, Store: store}).Sync(context.Background(), Request{ AccountID: "account-a", PlatformAccountKey: "sec-a", Credential: secretReference, }) - if err != nil || result.State != StatePolicyHold || result.ReasonCode != ReasonAuthInvalid || len(browser.urls) != 0 || len(browser.cookies) != 0 || len(store.holds) != 1 { + if err != nil || result.State != StatePolicyHold || result.ReasonCode != ReasonAuthInvalid || len(browser.urls) != 1 || len(browser.cookies) != 0 || len(store.holds) != 1 { t.Fatalf("unexpected invalid credential result: %#v browser=%#v holds=%#v err=%v", result, browser, store.holds, err) } encoded, _ := json.Marshal(result) @@ -195,12 +203,12 @@ func TestSyncRejectsInvalidCredentialWithoutLeakingIt(t *testing.T) { } func TestSyncStopsWhenSecretReferenceCannotResolve(t *testing.T) { - browser := &fakeBrowser{} + browser := &fakeBrowser{err: errors.New("session unavailable")} store := &fakeStore{} result, err := (Connector{Browser: browser, Secrets: fakeSecrets{err: errors.New("secret unavailable")}, Store: store}).Sync(context.Background(), Request{ AccountID: "account-a", PlatformAccountKey: "sec-a", Credential: secretReference, }) - if err != nil || result.State != StatePolicyHold || result.ReasonCode != ReasonAuthInvalid || len(browser.urls) != 0 || len(store.holds) != 1 { + if err != nil || result.State != StatePolicyHold || result.ReasonCode != ReasonAuthInvalid || len(browser.urls) != 1 || len(store.holds) != 1 { t.Fatalf("unavailable secret did not fail closed: result=%#v browser=%#v holds=%#v err=%v", result, browser, store.holds, err) } } @@ -219,6 +227,20 @@ func TestSyncFailsClosedWhenPersistenceIsUnknown(t *testing.T) { } } +func TestParseCredentialAcceptsCookieHeaderAndBundle(t *testing.T) { + for name, raw := range map[string]string{ + "header": "sessionid=secret; token=value", + "bundle": credential, + } { + t.Run(name, func(t *testing.T) { + cookies, err := ParseCredential([]byte(raw)) + if err != nil || len(cookies) != 2 && name == "header" || len(cookies) != 1 && name == "bundle" { + t.Fatalf("parse credential: cookies=%#v err=%v", cookies, err) + } + }) + } +} + func TestCredentialAndWorkValidation(t *testing.T) { invalidCredentials := []string{ ``, `{}`, `{"cookies":[]}`, `{"cookies":[{"name":"a","value":"b","domain":".douyin.com","extra":true}]}`, diff --git a/internal/douyin/creator_collector.go b/internal/douyin/creator_collector.go index 36aca84..6c30f92 100644 --- a/internal/douyin/creator_collector.go +++ b/internal/douyin/creator_collector.go @@ -22,22 +22,30 @@ type CreatorCollector struct { SourceID string } +func (c CreatorCollector) CanonicalSecUID(ctx context.Context, expectedKey string) (string, error) { + if c.Browser == nil || !keyPattern.MatchString(expectedKey) { + return "", fmt.Errorf("%w: invalid identity verification request", ErrInvalid) + } + response, err := c.Browser.Get(ctx, identityEndpoint) + if err != nil { + return "", err + } + if err := creatorResponseError(response, "identity"); err != nil { + return "", err + } + identity, ok := parseIdentity(response.Body) + if !ok || expectedKey != identity.User.UID && expectedKey != identity.User.SecUID && expectedKey != identity.User.UniqueID { + return "", fmt.Errorf("%w: douyin identity mismatch", ErrInvalid) + } + return identity.User.SecUID, nil +} + func (c CreatorCollector) VerifyIdentity(ctx context.Context, expectedKey string) error { if c.Browser == nil || !keyPattern.MatchString(expectedKey) { return fmt.Errorf("%w: invalid identity verification request", ErrInvalid) } - response, err := c.Browser.Get(ctx, identityEndpoint) - if err != nil { - return err - } - if err := creatorResponseError(response, "identity"); err != nil { - return err - } - identity, ok := parseIdentity(response.Body) - if !ok || expectedKey != identity.User.UID && expectedKey != identity.User.SecUID && expectedKey != identity.User.UniqueID { - return fmt.Errorf("%w: douyin identity mismatch", ErrInvalid) - } - return nil + _, err := c.CanonicalSecUID(ctx, expectedKey) + return err } func (c CreatorCollector) ListWorks(ctx context.Context, accountKey, cursor string) (creator.WorkPage, error) { diff --git a/requirements-gateway-dev.lock b/requirements-gateway-dev.lock new file mode 100644 index 0000000..a471611 --- /dev/null +++ b/requirements-gateway-dev.lock @@ -0,0 +1,2 @@ +-r requirements-gateway.lock +coverage==7.16.0 diff --git a/requirements-gateway.lock b/requirements-gateway.lock new file mode 100644 index 0000000..2f738ef --- /dev/null +++ b/requirements-gateway.lock @@ -0,0 +1 @@ +websocket-client==1.9.0 diff --git a/web/src/CreatorAccountsPage.jsx b/web/src/CreatorAccountsPage.jsx index 5e73555..70e3f3d 100644 --- a/web/src/CreatorAccountsPage.jsx +++ b/web/src/CreatorAccountsPage.jsx @@ -1,4 +1,4 @@ -import { useEffect, useMemo, useState } from "react"; +import { useEffect, useMemo, useRef, useState } from "react"; import { useDataProvider } from "@refinedev/core"; import { Alert, @@ -69,6 +69,14 @@ export function CreatorAccountsPage() { const [notice, setNotice] = useState(null); const [busy, setBusy] = useState(false); const [strategies, setStrategies] = useState([]); + const [strategyError, setStrategyError] = useState(null); + const [editingStrategyID, setEditingStrategyID] = useState(""); + const [relations, setRelations] = useState([]); + const [relationError, setRelationError] = useState(null); + const loadSequence = useRef(0); + const selectionVersion = useRef(0); + const strategySequence = useRef(0); + const relationSequence = useRef(0); const [strategyForm, setStrategyForm] = useState({ execution_account_id: "", position: 1, @@ -80,12 +88,14 @@ export function CreatorAccountsPage() { }); const load = async () => { + const sequence = ++loadSequence.current; setPending(true); setError(null); try { const result = await dataProvider.getList({ resource: "creator-accounts", }); + if (sequence !== loadSequence.current) return; setProfiles(result.data); if (selectedID) { const selected = result.data.find((item) => item.id === selectedID); @@ -95,9 +105,9 @@ export function CreatorAccountsPage() { setForm(profileForm(result.data[0])); } } catch (loadError) { - setError(loadError); + if (sequence === loadSequence.current) setError(loadError); } finally { - setPending(false); + if (sequence === loadSequence.current) setPending(false); } }; @@ -106,12 +116,35 @@ export function CreatorAccountsPage() { }, []); useEffect(() => { if (!selectedID) return; + const sequence = ++strategySequence.current; + setStrategyError(null); + setStrategies([]); dataProvider .creatorGet( `/creator/accounts/${encodeURIComponent(selectedID)}/strategies`, ) - .then(setStrategies) - .catch(() => setStrategies([])); + .then((result) => { + if (sequence === strategySequence.current) setStrategies(result); + }) + .catch((loadError) => { + if (sequence === strategySequence.current) setStrategyError(loadError); + }); + }, [selectedID]); + useEffect(() => { + if (!selectedID) return; + const sequence = ++relationSequence.current; + setRelationError(null); + setRelations([]); + dataProvider + .creatorGet( + `/creator/relations?big_account_id=${encodeURIComponent(selectedID)}`, + ) + .then((result) => { + if (sequence === relationSequence.current) setRelations(result); + }) + .catch((loadError) => { + if (sequence === relationSequence.current) setRelationError(loadError); + }); }, [selectedID]); const selected = useMemo( @@ -119,8 +152,23 @@ export function CreatorAccountsPage() { [profiles, selectedID], ); const choose = (profile) => { + const dirty = + selected && + JSON.stringify(form) !== JSON.stringify(profileForm(selected)); + if ( + dirty && + !window.confirm("当前账号资料尚未保存,确定放弃修改并切换吗?") + ) { + return; + } + selectionVersion.current += 1; + loadSequence.current += 1; setSelectedID(profile.id); setForm(profileForm(profile)); + setStrategies([]); + setRelations([]); + setEditingStrategyID(""); + setBusy(false); setNotice(null); }; const change = (field) => (event) => @@ -128,6 +176,8 @@ export function CreatorAccountsPage() { const save = async (event) => { event.preventDefault(); if (!selected) return; + const accountID = selected.id; + const version = selectionVersion.current; setBusy(true); setNotice(null); try { @@ -138,22 +188,30 @@ export function CreatorAccountsPage() { cooldown_seconds: Number(form.cooldown_seconds), }, ); + if (version !== selectionVersion.current || accountID !== selectedID) + return; setProfiles((items) => - items.map((item) => (item.id === selected.id ? result : item)), + items.map((item) => (item.id === accountID ? result : item)), ); setForm(profileForm(result)); setNotice({ variant: "success", text: "账号资料已保存。" }); } catch (saveError) { - setNotice({ - variant: "destructive", - text: conflictMessage(saveError, "账号资料保存失败"), - }); + if (version === selectionVersion.current && accountID === selectedID) { + setNotice({ + variant: "destructive", + text: conflictMessage(saveError, "账号资料保存失败"), + }); + } } finally { - setBusy(false); + if (version === selectionVersion.current && accountID === selectedID) { + setBusy(false); + } } }; const toggleBig = async () => { if (!selected) return; + const accountID = selected.id; + const version = selectionVersion.current; setBusy(true); setNotice(null); try { @@ -161,8 +219,10 @@ export function CreatorAccountsPage() { `/creator/accounts/${encodeURIComponent(selected.id)}/big-account`, { enabled: !selected.big_account }, ); + if (version !== selectionVersion.current || accountID !== selectedID) + return; setProfiles((items) => - items.map((item) => (item.id === selected.id ? result : item)), + items.map((item) => (item.id === accountID ? result : item)), ); setForm(profileForm(result)); setNotice({ @@ -170,12 +230,16 @@ export function CreatorAccountsPage() { text: result.big_account ? "已开启大号模式。" : "已关闭大号模式。", }); } catch (actionError) { - setNotice({ - variant: "destructive", - text: conflictMessage(actionError, "大号模式更新失败"), - }); + if (version === selectionVersion.current && accountID === selectedID) { + setNotice({ + variant: "destructive", + text: conflictMessage(actionError, "大号模式更新失败"), + }); + } } finally { - setBusy(false); + if (version === selectionVersion.current && accountID === selectedID) { + setBusy(false); + } } }; const createStrategy = async (event) => { @@ -191,27 +255,35 @@ export function CreatorAccountsPage() { setBusy(true); setNotice(null); try { + const payload = { + ...strategyForm, + position: Number(strategyForm.position), + candidate_texts: strategyForm.candidate_texts + .split(",") + .map((item) => item.trim()) + .filter(Boolean), + }; await dataProvider.creatorCreate("/creator/relations", { big_account_id: selected.id, small_account_id: strategyForm.execution_account_id, enabled: true, }); - const result = await dataProvider.creatorCreate( - `/creator/accounts/${encodeURIComponent(selected.id)}/strategies`, - { - ...strategyForm, - position: Number(strategyForm.position), - candidate_texts: strategyForm.candidate_texts - .split(",") - .map((item) => item.trim()) - .filter(Boolean), - }, - ); + const result = editingStrategyID + ? await dataProvider.creatorUpdate( + `/creator/strategies/${encodeURIComponent(editingStrategyID)}`, + payload, + ) + : await dataProvider.creatorCreate( + `/creator/accounts/${encodeURIComponent(selected.id)}/strategies`, + payload, + ); setStrategies((items) => - [...items, result].sort( - (left, right) => left.position - right.position, - ), + (editingStrategyID + ? items.map((item) => (item.id === editingStrategyID ? result : item)) + : [...items, result] + ).sort((left, right) => left.position - right.position), ); + setEditingStrategyID(""); setNotice({ variant: "success", text: "自动响应策略已保存。" }); } catch (createError) { setNotice({ @@ -244,6 +316,33 @@ export function CreatorAccountsPage() { setBusy(false); } }; + const unlinkRelation = async (relation) => { + setBusy(true); + setNotice(null); + try { + await dataProvider.creatorCreate("/creator/relations", { + big_account_id: relation.big_account_id, + small_account_id: relation.small_account_id, + enabled: false, + }); + setRelations((items) => + items.filter( + (item) => item.small_account_id !== relation.small_account_id, + ), + ); + setNotice({ + variant: "success", + text: "账号关系已解除;已有策略不会自动恢复。", + }); + } catch (actionError) { + setNotice({ + variant: "destructive", + text: conflictMessage(actionError, "账号关系解除失败"), + }); + } finally { + setBusy(false); + } + }; const deleteStrategy = async (strategy) => { setBusy(true); setNotice(null); @@ -314,6 +413,16 @@ export function CreatorAccountsPage() { {profile.login_status === "logged_in" ? "已登录" : "需人工确认"} + {profile.login_checked_at + ? ` · 检查于 ${dateTime(profile.login_checked_at)}` + : " · 尚未核验"} +

+

+ 凭据: + {profile.password_configured ? "已配置" : "未配置"} + {profile.login_reason + ? ` · ${profile.login_reason}` + : ""}

@@ -333,6 +442,22 @@ export function CreatorAccountsPage() { {platformLabel[selected.platform] || selected.platform} ·{" "} {selected.platform_account_key}

+

+ 登录核验: + {selected.login_status === "logged_in" + ? "已登录" + : "需人工确认"} + {selected.login_checked_at + ? ` · ${dateTime(selected.login_checked_at)}` + : " · 尚未核验"} + {selected.login_reason + ? ` · ${selected.login_reason}` + : ""} +

+

+ 密码凭据: + {selected.password_configured ? "已配置" : "未配置"} +

+ {editingStrategyID ? ( + + ) : null} + {strategyError ? ( + + 策略读取失败: + {conflictMessage(strategyError, "请只读重试")} + + ) : null} + {relationError ? ( + + 账号关系读取失败: + {conflictMessage(relationError, "请只读重试")} + + ) : null} + {relations.length ? ( +
    + {relations.map((relation) => ( +
  • + 执行账号:{relation.small_account_id} + +
  • + ))} +
+ ) : null} {strategies.length ? (
    {strategies.map((strategy) => ( @@ -608,6 +778,28 @@ export function CreatorAccountsPage() { tone={strategy.enabled ? "success" : "neutral"} label={strategy.enabled ? "启用" : "停用"} /> +
- ) : ( + ) : strategyError ? null : (

尚未配置策略。

)} diff --git a/web/src/CreatorCompetitorsPage.jsx b/web/src/CreatorCompetitorsPage.jsx index dd0e1db..ccf3ea3 100644 --- a/web/src/CreatorCompetitorsPage.jsx +++ b/web/src/CreatorCompetitorsPage.jsx @@ -34,10 +34,13 @@ export function CreatorCompetitorsPage() { homepage_url: "", }); const [minLikes, setMinLikes] = useState(""); + const [minComments, setMinComments] = useState(""); + const [minShares, setMinShares] = useState(""); const [pending, setPending] = useState(true); const [busy, setBusy] = useState(false); const [error, setError] = useState(null); const [notice, setNotice] = useState(null); + const [accountError, setAccountError] = useState(null); const [material, setMaterial] = useState(null); const [materialPending, setMaterialPending] = useState(false); const [rewriteRequirement, setRewriteRequirement] = useState(""); @@ -52,6 +55,10 @@ export function CreatorCompetitorsPage() { if (platform) filters.push({ field: "platform", value: platform }); if (minLikes !== "") filters.push({ field: "min_likes", value: Number(minLikes) }); + if (minComments !== "") + filters.push({ field: "min_comments", value: Number(minComments) }); + if (minShares !== "") + filters.push({ field: "min_shares", value: Number(minShares) }); const [competitorResult, workResult] = await Promise.all([ dataProvider.getList({ resource: "creator-competitors", @@ -69,16 +76,50 @@ export function CreatorCompetitorsPage() { }; useEffect(() => { load(); - }, [platform, minLikes]); + }, [platform, minLikes, minComments, minShares]); useEffect(() => { dataProvider .getList({ resource: "creator-accounts" }) - .then((result) => setAccounts(result.data)) - .catch(() => setAccounts([])); + .then((result) => { + setAccountError(null); + setAccounts(result.data); + }) + .catch((loadError) => setAccountError(loadError)); }, [dataProvider]); const change = (field) => (event) => setForm((value) => ({ ...value, [field]: event.target.value })); + const parseHomepage = () => { + try { + const parsed = new URL(form.homepage_url); + const allowed = + form.platform === "douyin" + ? parsed.hostname === "www.douyin.com" + : parsed.hostname === "www.xiaohongshu.com" || + parsed.hostname === "xiaohongshu.com"; + const parts = parsed.pathname.split("/").filter(Boolean); + const candidate = parts.at(-1) || ""; + if ( + !allowed || + !candidate || + !/^[A-Za-z0-9][A-Za-z0-9._:@/-]{0,127}$/.test(candidate) + ) { + throw new Error( + "链接不是受支持的平台主页格式,请人工填写平台返回的稳定标识", + ); + } + setForm((value) => ({ ...value, platform_account_key: candidate })); + setNotice({ + variant: "info", + text: `已解析候选标识 ${candidate},请核对平台返回值后再提交。`, + }); + } catch (parseError) { + setNotice({ + variant: "warning", + text: parseError.message || "主页链接解析失败", + }); + } + }; const create = async (event) => { event.preventDefault(); setBusy(true); @@ -128,7 +169,9 @@ export function CreatorCompetitorsPage() { const account = accounts.find( (value) => value.platform === item.platform && - value.authorization_status === "authorized", + value.authorization_status === "authorized" && + value.business_status === "normal" && + value.login_status === "logged_in", ); if (!account) { setNotice({ @@ -159,6 +202,14 @@ export function CreatorCompetitorsPage() { } }; const selectMaterial = async (workID) => { + const dirty = + material && + (rewriteRequirement !== (material.rewrite_requirement || "") || + rewriteTitle !== (material.generated_title || "") || + rewriteScript !== (material.generated_script || "")); + if (dirty && !window.confirm("当前素材文稿尚未保存,确定放弃并切换吗?")) { + return; + } setMaterialPending(true); setNotice(null); try { @@ -171,7 +222,7 @@ export function CreatorCompetitorsPage() { setRewriteScript(result.generated_script || ""); setNotice({ variant: "success", - text: "素材已选择;下载、提音轨与转写结果会逐步显示,失败步骤需要明确处理。", + text: "素材已选择;点击开始处理后才会执行真实下载、音轨提取和转写。", }); } catch (actionError) { setNotice({ @@ -182,6 +233,32 @@ export function CreatorCompetitorsPage() { setMaterialPending(false); } }; + const processMaterial = async () => { + if (!material) return; + setMaterialPending(true); + setNotice(null); + try { + const result = await dataProvider.creatorAction( + `/creator/works/${encodeURIComponent(material.work_id)}/material/process`, + ); + setMaterial(result); + setNotice({ + variant: + result.transcription_status === "failed" ? "warning" : "success", + text: + result.transcription_status === "failed" + ? "素材已处理,但转写供应商尚未配置。" + : "素材处理完成。", + }); + } catch (actionError) { + setNotice({ + variant: "destructive", + text: conflictMessage(actionError, "素材处理失败"), + }); + } finally { + setMaterialPending(false); + } + }; const confirmRewrite = async () => { if (!material) return; setMaterialPending(true); @@ -284,9 +361,17 @@ export function CreatorCompetitorsPage() { required /> - +
+ + +
+

+ 主页链接仅用于人工核对;请把平台返回的稳定标识填入上方,不根据昵称自动猜测。 +

已监测账号

@@ -305,10 +390,20 @@ export function CreatorCompetitorsPage() { {item.platform} · {item.platform_account_key}

+ 同步:{item.sync_status || "未开始"} {item.last_sync_at - ? `上次:${dateTime(item.last_sync_at)}` - : "尚未同步"} + ? ` · 上次 ${dateTime(item.last_sync_at)}` + : " · 尚未完成"}

+ {item.sync_error ? ( +

+ {item.sync_error} +

+ ) : item.next_sync_at ? ( +

+ 下次:{dateTime(item.next_sync_at)} +

+ ) : null}
setMinLikes(event.target.value)} /> + + setMinComments(event.target.value)} + /> + + + setMinShares(event.target.value)} + /> +
+ {accountError ? ( + + 同平台账号读取失败: + {conflictMessage(accountError, "请重试账号列表")} + + ) : null} 每一步都保留真实状态;失败不会被显示为成功。

+
{ expect.objectContaining({ lookback_days: 14 }), ), ); + expect(dataProvider.creatorUpdate).not.toHaveBeenCalledWith( + "/creator/settings", + expect.objectContaining({ updated_at: expect.anything() }), + ); + }); + + it("does not render stale data using another tab's shape", async () => { + const dataProvider = provider({ + getList: vi.fn(({ resource }) => + Promise.resolve({ + data: + resource === "creator-comments" + ? [ + { + id: "comment-a", + author_name: "评论用户", + author_uid: "peer-a", + content: "你好", + platform: "douyin", + work_id: "work-a", + }, + ] + : resource === "creator-rules" + ? [ + { + id: "rule-a", + name: "规则", + include_keywords: ["你好"], + enabled: true, + topic: "主题", + }, + ] + : resource === "creator-leads" + ? [ + { + id: "lead-a", + comment: { + id: "comment-a", + author_name: "线索用户", + author_uid: "peer-a", + content: "线索", + }, + rule_ids: ["rule-a"], + }, + ] + : [], + }), + ), + }); + renderPage(, dataProvider); + expect(await screen.findByText("评论用户")).toBeTruthy(); + fireEvent.click(screen.getByRole("tab", { name: "线索" })); + expect(await screen.findByText("线索用户", { exact: false })).toBeTruthy(); }); it("requires explicit stable competitor identity and shows empty works state", async () => { @@ -143,6 +194,42 @@ describe("creator pages", () => { expect(screen.getByLabelText("更新密码")).toBeTruthy(); }); + it("allows switching private-message accounts across platforms", async () => { + const xhsProfile = { + ...profile, + id: "account-xhs", + name: "小红书账号", + platform: "xiaohongshu", + platform_account_key: "xhs-a", + }; + const conversation = { + id: "conversation-a", + account_id: profile.id, + platform: profile.platform, + peer_uid: "peer-a", + peer_name: "客户", + }; + const dataProvider = provider({ + getList: vi.fn(({ resource }) => + Promise.resolve({ + data: + resource === "creator-accounts" + ? [profile, xhsProfile] + : resource === "creator-conversations" + ? [conversation] + : [], + total: 1, + }), + ), + creatorGet: vi.fn(() => Promise.resolve([])), + }); + renderPage(, dataProvider); + fireEvent.click(await screen.findByRole("tab", { name: "私信" })); + expect(await screen.findByText("客户")).toBeTruthy(); + fireEvent.click(screen.getByRole("combobox", { name: "发送账号" })); + expect(screen.getByRole("option", { name: "小红书账号" })).toBeTruthy(); + }); + it("isolates private messages by account and confirms one durable send", async () => { const conversation = { id: "conversation-a", diff --git a/web/src/CreatorSettingsPage.jsx b/web/src/CreatorSettingsPage.jsx index 3aae032..f7ccbbb 100644 --- a/web/src/CreatorSettingsPage.jsx +++ b/web/src/CreatorSettingsPage.jsx @@ -1,6 +1,34 @@ import { useEffect, useState } from "react"; import { useDataProvider } from "@refinedev/core"; -import { Alert, Button, Card, CardContent, Field, Input, PageHeader, PageState, conflictMessage } from "./lib/ui.jsx"; +import { + Alert, + Button, + Card, + CardContent, + Field, + Input, + PageHeader, + PageState, + conflictMessage, +} from "./lib/ui.jsx"; + +const editableFields = [ + "lookback_days", + "new_work_interval_seconds", + "metric_initial_interval_seconds", + "metric_multiplier", + "metric_max_interval_seconds", + "metric_age_seconds", + "ai_provider", + "ai_model", + "ai_configured", + "transcription_provider", + "transcription_model", + "transcription_configured", +]; + +const editableSettings = (value) => + Object.fromEntries(editableFields.map((field) => [field, value[field]])); const defaults = { lookback_days: 30, @@ -26,25 +54,200 @@ export function CreatorSettingsPage() { const [notice, setNotice] = useState(null); const load = async () => { - setPending(true); setError(null); + setPending(true); + setError(null); try { const result = await dataProvider.creatorGet("/creator/settings"); - setForm((value) => ({ ...value, ...result })); - } catch (loadError) { setError(loadError); } - finally { setPending(false); } + setForm((value) => ({ ...value, ...editableSettings(result) })); + } catch (loadError) { + setError(loadError); + } finally { + setPending(false); + } }; - useEffect(() => { load(); }, []); + useEffect(() => { + load(); + }, []); const change = (field) => (event) => { const value = event.target.value; - setForm((current) => ({ ...current, [field]: ["lookback_days", "new_work_interval_seconds", "metric_initial_interval_seconds", "metric_multiplier", "metric_max_interval_seconds", "metric_age_seconds"].includes(field) ? Number(value) : value })); + setForm((current) => ({ + ...current, + [field]: [ + "lookback_days", + "new_work_interval_seconds", + "metric_initial_interval_seconds", + "metric_multiplier", + "metric_max_interval_seconds", + "metric_age_seconds", + ].includes(field) + ? Number(value) + : value, + })); }; - const toggle = (field) => () => setForm((current) => ({ ...current, [field]: !current[field] })); + const toggle = (field) => () => + setForm((current) => ({ ...current, [field]: !current[field] })); const save = async (event) => { - event.preventDefault(); setBusy(true); setNotice(null); - try { const result = await dataProvider.creatorUpdate("/creator/settings", form); setForm(result); setNotice({ variant: "success", text: "采集与 AI 配置已保存。" }); } - catch (saveError) { setNotice({ variant: "destructive", text: conflictMessage(saveError, "设置保存失败") }); } - finally { setBusy(false); } + event.preventDefault(); + setBusy(true); + setNotice(null); + try { + const result = await dataProvider.creatorUpdate( + "/creator/settings", + editableSettings(form), + ); + setForm((value) => ({ ...value, ...editableSettings(result) })); + setNotice({ variant: "success", text: "采集与 AI 配置已保存。" }); + } catch (saveError) { + setNotice({ + variant: "destructive", + text: conflictMessage(saveError, "设置保存失败"), + }); + } finally { + setBusy(false); + } }; - return (<>

AI 服务

转写服务

{notice ? {notice.text} : null}
); + return ( + <> + + + + +
+ + + + + + + + + + + + + + + + + + +
+

AI 服务

+
+ + + + + + +
+ +
+
+

转写服务

+
+ + + + + + +
+ +
+ {notice ? ( + + {notice.text} + + ) : null} +
+ +
+
+
+
+
+ + ); } diff --git a/web/src/CreatorWorkbenchPage.jsx b/web/src/CreatorWorkbenchPage.jsx index 84922d7..5fa43ba 100644 --- a/web/src/CreatorWorkbenchPage.jsx +++ b/web/src/CreatorWorkbenchPage.jsx @@ -26,23 +26,60 @@ const tabs = [ ]; function parseKeywords(value) { - return value - .split(",") - .map((item) => item.trim()) - .filter(Boolean); + if (!value.trim()) return []; + const values = value.split(",").map((item) => item.trim()); + if (values.some((item) => !item)) throw new Error("关键词不能为空"); + return values; } function newOperationKey() { - return `manual-${globalThis.crypto.randomUUID()}`; + const crypto = globalThis.crypto; + if (typeof crypto?.randomUUID === "function") + return `manual-${crypto.randomUUID()}`; + if (typeof crypto?.getRandomValues !== "function") { + throw new Error( + "当前浏览器不支持安全操作标识,请使用 HTTPS 或受支持的浏览器", + ); + } + const bytes = crypto.getRandomValues(new Uint8Array(16)); + return `manual-${Array.from(bytes, (value) => value.toString(16).padStart(2, "0")).join("")}`; +} + +const replyDraftKey = (commentID) => `creatorhub.reply.${commentID}`; +const dmDraftKey = (accountID, conversationID) => + `creatorhub.dm.${accountID}.${conversationID}`; +const readDraft = (key) => { + try { + const value = JSON.parse(localStorage.getItem(key) || "null"); + return value && typeof value === "object" ? value : null; + } catch { + return null; + } +}; +const writeDraft = (key, value) => + localStorage.setItem(key, JSON.stringify(value)); +const removeDraft = (key) => localStorage.removeItem(key); + +function operationStatus(state, reason) { + const labels = { + succeeded: "成功(已获得平台证据)", + failed: "失败(平台明确拒绝或执行前失败)", + blocked: "已阻止(未发送)", + processing: "处理中(尚未获得最终结果)", + uncertain: "结果不明(禁止自动重试)", + }; + return `${labels[state] || state || "未知状态"}${reason ? `:${reason}` : ""}`; } export function CreatorWorkbenchPage() { const dataProvider = useDataProvider()("default"); const [tab, setTab] = useState("comments"); const [data, setData] = useState([]); + const [dataTab, setDataTab] = useState(""); const [rules, setRules] = useState([]); const [accounts, setAccounts] = useState([]); const [conversations, setConversations] = useState([]); + const [conversationAccountID, setConversationAccountID] = useState(""); const [messages, setMessages] = useState([]); const [messagePending, setMessagePending] = useState(false); const [messageError, setMessageError] = useState(null); @@ -57,11 +94,13 @@ export function CreatorWorkbenchPage() { const [analyzing, setAnalyzing] = useState(""); const [ruleForm, setRuleForm] = useState({ name: "", + source_type: "all", topic: "", include_keywords: "", exclude_keywords: "", ai_requirement: "", }); + const [editingRuleID, setEditingRuleID] = useState(""); const [reply, setReply] = useState({ account_id: "", text: "", @@ -74,11 +113,15 @@ export function CreatorWorkbenchPage() { const [dmOperationKey, setDmOperationKey] = useState(""); const loadSequence = useRef(0); const messageSequence = useRef(0); + const visibleData = dataTab === tab ? data : []; + const visibleConversations = + conversationAccountID === dmAccountID ? conversations : []; const load = async () => { const sequence = ++loadSequence.current; setPending(true); setError(null); + setDataTab(""); try { if (tab === "comments") { const [comments, ruleList] = await Promise.all([ @@ -87,6 +130,7 @@ export function CreatorWorkbenchPage() { ]); if (sequence !== loadSequence.current) return; setData(comments.data); + setDataTab("comments"); setRules(ruleList.data); } else if (tab === "leads") { const result = await dataProvider.getList({ @@ -94,14 +138,19 @@ export function CreatorWorkbenchPage() { }); if (sequence !== loadSequence.current) return; setData(result.data); + setDataTab("leads"); } else if (tab === "rules") { const result = await dataProvider.getList({ resource: "creator-rules", }); if (sequence !== loadSequence.current) return; setData(result.data); + setDataTab("rules"); } else if (tab === "dms") { - if (!dmAccountID) return; + if (!dmAccountID) { + setPending(false); + return; + } const result = await dataProvider.getList({ resource: "creator-conversations", filters: [ @@ -110,7 +159,9 @@ export function CreatorWorkbenchPage() { }); if (sequence !== loadSequence.current) return; setConversations(result.data); + setConversationAccountID(dmAccountID); setData(result.data); + setDataTab("dms"); setConversationID((current) => result.data.some((item) => item.id === current) ? current @@ -122,6 +173,7 @@ export function CreatorWorkbenchPage() { }); if (sequence !== loadSequence.current) return; setData(result.data); + setDataTab("operations"); } } catch (loadError) { if (sequence === loadSequence.current) setError(loadError); @@ -148,6 +200,14 @@ export function CreatorWorkbenchPage() { useEffect(() => { if (tab !== "dms" || dmAccountID) load(); }, [tab, dmAccountID]); + const switchTab = (nextTab) => { + if (nextTab === tab) return; + loadSequence.current += 1; + setData([]); + setDataTab(""); + setError(null); + setTab(nextTab); + }; useEffect(() => { if (!conversationID || tab !== "dms") { setMessages([]); @@ -176,6 +236,32 @@ export function CreatorWorkbenchPage() { if (sequence === messageSequence.current) setMessagePending(false); }; }, [conversationID, tab, dataProvider, messageRetry]); + useEffect(() => { + if (reply.target_comment_id) { + writeDraft(replyDraftKey(reply.target_comment_id), reply); + } + }, [reply]); + useEffect(() => { + if (tab !== "dms" || !conversationID) return undefined; + const timer = window.setInterval(() => { + setMessageRetry((value) => value + 1); + }, 10000); + return () => window.clearInterval(timer); + }, [tab, conversationID]); + useEffect(() => { + if (!dmAccountID || !conversationID) return; + const draft = readDraft(dmDraftKey(dmAccountID, conversationID)); + setDmText(draft?.text || ""); + setDmOperationKey(draft?.operation_key || ""); + }, [dmAccountID, conversationID]); + useEffect(() => { + if (dmAccountID && conversationID) { + writeDraft(dmDraftKey(dmAccountID, conversationID), { + operation_key: dmOperationKey, + text: dmText, + }); + } + }, [dmAccountID, conversationID, dmOperationKey, dmText]); const analyze = async (commentID, ruleID) => { setAnalyzing(`${commentID}:${ruleID}`); @@ -203,22 +289,37 @@ export function CreatorWorkbenchPage() { setBusy(true); setNotice(null); try { - const result = await dataProvider.create({ - resource: "creator-rules", - variables: { - name: ruleForm.name, - enabled: true, - source_type: "all", - topic: ruleForm.topic, - include_keywords: parseKeywords(ruleForm.include_keywords), - exclude_keywords: parseKeywords(ruleForm.exclude_keywords), - ai_requirement: ruleForm.ai_requirement, - }, - }); - setData((items) => [result.data, ...items]); - setRules((items) => [result.data, ...items]); + const variables = { + name: ruleForm.name, + enabled: true, + source_type: ruleForm.source_type, + topic: ruleForm.topic, + include_keywords: parseKeywords(ruleForm.include_keywords), + exclude_keywords: parseKeywords(ruleForm.exclude_keywords), + ai_requirement: ruleForm.ai_requirement, + }; + const result = editingRuleID + ? await dataProvider.creatorUpdate( + `/creator/rules/${encodeURIComponent(editingRuleID)}`, + variables, + ) + : ( + await dataProvider.create({ + resource: "creator-rules", + variables, + }) + ).data; + const update = (items) => + editingRuleID + ? items.map((item) => (item.id === editingRuleID ? result : item)) + : [result, ...items]; + setData(update); + setDataTab("rules"); + setRules(update); + setEditingRuleID(""); setRuleForm({ name: "", + source_type: "all", topic: "", include_keywords: "", exclude_keywords: "", @@ -235,12 +336,13 @@ export function CreatorWorkbenchPage() { } }; const requestReply = (comment) => { + const draft = readDraft(replyDraftKey(comment.id)); setReply({ - account_id: "", - text: "", + account_id: draft?.account_id || "", + text: draft?.text || "", target_uid: comment.author_uid || "", target_comment_id: comment.id, - operation_key: newOperationKey(), + operation_key: draft?.operation_key || "", }); setNotice({ variant: "info", @@ -260,11 +362,15 @@ export function CreatorWorkbenchPage() { }); return; } - setReply((value) => ({ - ...value, - operation_key: value.operation_key || newOperationKey(), - })); - setConfirm({ kind: "reply", comment }); + try { + setReply((value) => ({ + ...value, + operation_key: value.operation_key || newOperationKey(), + })); + setConfirm({ kind: "reply", comment }); + } catch (error) { + setNotice({ variant: "warning", text: error.message }); + } }; const sendReply = async () => { if (!confirm) return; @@ -286,13 +392,12 @@ export function CreatorWorkbenchPage() { ); setNotice({ variant: result.state === "succeeded" ? "success" : "warning", - text: - result.state === "succeeded" - ? "已获得平台成功证据。" - : "操作已记录,但平台执行结果仍不确定。", + text: operationStatus(result.state, result.reason), }); - if (result.state === "succeeded") + if (result.state === "succeeded") { + removeDraft(replyDraftKey(reply.target_comment_id)); setReply((value) => ({ ...value, operation_key: "" })); + } } catch (sendError) { setNotice({ variant: "destructive", @@ -310,6 +415,8 @@ export function CreatorWorkbenchPage() { const switchDMAccount = (accountID) => { if (accountID === dmAccountID || !canSwitchDM()) return; setDmAccountID(accountID); + setConversationAccountID(""); + setConversations([]); setConversationID(""); setDmText(""); setDmOperationKey(""); @@ -328,7 +435,13 @@ export function CreatorWorkbenchPage() { }); return; } - const operationKey = dmOperationKey || newOperationKey(); + let operationKey; + try { + operationKey = dmOperationKey || newOperationKey(); + } catch (error) { + setNotice({ variant: "warning", text: error.message }); + return; + } setDmOperationKey(operationKey); setDmConfirm({ ...conversation, operation_key: operationKey }); }; @@ -351,12 +464,10 @@ export function CreatorWorkbenchPage() { ); setNotice({ variant: result.state === "succeeded" ? "success" : "warning", - text: - result.state === "succeeded" - ? "已获得平台成功证据。" - : "操作已记录,但平台执行结果仍不确定。", + text: operationStatus(result.state, result.reason), }); if (result.state === "succeeded") { + removeDraft(dmDraftKey(dmAccountID, dmConfirm.id)); setDmText(""); setDmOperationKey(""); } @@ -376,7 +487,7 @@ export function CreatorWorkbenchPage() { @@ -398,8 +509,9 @@ export function CreatorWorkbenchPage() { .filter( (account) => account.platform === - data.find((item) => item.id === reply.target_comment_id) - ?.platform, + visibleData.find( + (item) => item.id === reply.target_comment_id, + )?.platform, ) .map((account) => ({ value: account.id, @@ -426,7 +538,9 @@ export function CreatorWorkbenchPage() { variant="primary" onClick={() => askReplyConfirm( - data.find((item) => item.id === reply.target_comment_id), + visibleData.find( + (item) => item.id === reply.target_comment_id, + ), ) } > @@ -440,7 +554,7 @@ export function CreatorWorkbenchPage() { ) : null}
- {data.map((comment) => ( + {visibleData.map((comment) => (
@@ -492,12 +606,12 @@ export function CreatorWorkbenchPage() {
- {data.map((lead) => ( + {visibleData.map((lead) => (
@@ -505,7 +619,15 @@ export function CreatorWorkbenchPage() { {lead.comment.author_name || "未知用户"} ·{" "} {lead.comment.author_uid || "UID 不可用"}

- +
+ + +

{lead.comment.content}

@@ -522,6 +644,9 @@ export function CreatorWorkbenchPage() {

新建规则

+

+ {editingRuleID ? "编辑规则" : "新建规则"} +

+ + - +
+ + {editingRuleID ? ( + + ) : null} +
@@ -595,19 +749,79 @@ export function CreatorWorkbenchPage() {

已保存规则

- {data.map((rule) => ( + {visibleData.map((rule) => (
-
+
{rule.name}
-

主题:{rule.topic}

+

+ 来源:{rule.source_type || "all"} · 主题:{rule.topic} +

包含:{rule.include_keywords.join("、")}

+
+ + +
))}
@@ -623,20 +837,10 @@ export function CreatorWorkbenchPage() { id="dm-account" value={dmAccountID} onChange={(event) => switchDMAccount(event.target.value)} - options={accounts - .filter( - (account) => - !conversations.find( - (item) => item.id === conversationID, - ) || - account.platform === - conversations.find((item) => item.id === conversationID) - ?.platform, - ) - .map((account) => ({ - value: account.id, - label: account.name || account.id, - }))} + options={accounts.map((account) => ({ + value: account.id, + label: account.name || account.id, + }))} placeholder="选择发送账号" /> @@ -645,7 +849,7 @@ export function CreatorWorkbenchPage() { @@ -653,7 +857,7 @@ export function CreatorWorkbenchPage() {
    - {conversations.map((conversation) => ( + {visibleConversations.map((conversation) => (
  • +
askDMConfirm( - conversations.find( + visibleConversations.find( (item) => item.id === conversationID, ), ) @@ -738,7 +952,7 @@ export function CreatorWorkbenchPage() { @@ -754,7 +968,7 @@ export function CreatorWorkbenchPage() { - {data.map((item) => ( + {visibleData.map((item) => ( {item.action} {item.account_id} @@ -763,11 +977,14 @@ export function CreatorWorkbenchPage() { tone={ item.state === "succeeded" ? "success" - : item.state === "uncertain" + : item.state === "uncertain" || + item.state === "processing" ? "warning" - : "neutral" + : item.state === "failed" + ? "danger" + : "neutral" } - label={item.state} + label={operationStatus(item.state, item.reason)} /> @@ -797,7 +1014,7 @@ export function CreatorWorkbenchPage() { key={value} size="sm" variant={tab === value ? "primary" : "outline"} - onClick={() => setTab(value)} + onClick={() => switchTab(value)} role="tab" aria-selected={tab === value} >