Files
creator-hub/internal/account/store.go
T
rogee eab193007c
douyin-release-gate / verify (push) Failing after 3m43s
refactor(accounts): 移除「授权状态」概念——账号收敛为启用/暂停单一状态机
产品已收敛为自有账号管理,撤销授权在 UI 无入口、状态恒为 authorized,属废弃语义:
- migration 1045:social_account DROP authorization_status/revoked_at/authorization_kind
- 删除 revoke API 路由与 RevokeAccount/disableAccount 状态机分支(PauseAccount 独立)
- accountRunnable/就绪判定/采集过滤/登录校验删除 authorization_status 检查
- EnvironmentContext/AccountProfile 契约删字段;前端删「已授权/已撤销」展示与 readiness 分支
- 测试同步:revoke 流程/409 用例删除,seed 语句去列;dev 库测试遗留 revoked 账号待 UI 删除
2026-09-30 14:18:34 +08:00

526 lines
18 KiB
Go

package account
import (
"context"
"crypto/rand"
"database/sql"
"encoding/hex"
"encoding/json"
"errors"
"net/http"
"regexp"
"strings"
"time"
"unicode/utf8"
"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgconn"
"github.com/jackc/pgx/v5/pgtype"
_ "github.com/jackc/pgx/v5/stdlib"
)
var (
ErrConflict = errors.New("resource conflicts with existing state")
ErrInvalid = errors.New("invalid phase A input")
ErrNotFound = errors.New("resource not found")
ErrAccountCreationUnknown = errors.New("account creation result is unknown")
idPattern = regexp.MustCompile(`^[a-z0-9][a-z0-9-]{0,31}$`)
refPattern = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9._/-]{0,127}$`)
platformKeyPattern = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9._:@/-]{0,127}$`)
credentialKeyPattern = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9._-]{0,63}/[A-Za-z0-9][A-Za-z0-9._/-]{0,126}$`)
eventPattern = regexp.MustCompile(`^[a-z0-9_]{1,64}$`)
)
type Store struct {
db *sql.DB
accountCommit func(*sql.Tx) error
}
type Account struct {
ID string `json:"id"`
Name string `json:"name"`
Platform string `json:"platform"`
PlatformAccountKey string `json:"platform_account_key"`
Tags []string `json:"tags"`
Cookies string `json:"-"`
CredentialReference CredentialReference `json:"-"`
CredentialKey string `json:"-"`
RuntimeStatus string `json:"runtime_status"`
Version int64 `json:"version"`
}
type CredentialReference struct {
ID string
Provider string
}
type CredentialBridge interface {
// Store may fail after a partial write; Delete must be idempotent for compensation.
Store(context.Context, CredentialReference, string, string) error
Delete(context.Context, CredentialReference, string) error
}
type CredentialResolver interface {
Resolve(context.Context, CredentialReference, string) ([]byte, error)
}
type ReadinessError struct {
Reason string
Unavailable bool
}
func (e *ReadinessError) Error() string { return e.Reason }
type AuditEvent struct {
ID int64 `json:"id"`
EventType string `json:"event_type"`
AccountID string `json:"account_id,omitempty"`
BrowserEnvAlias string `json:"browser_env_alias,omitempty"`
NetworkExitID string `json:"network_exit_id,omitempty"`
BindingVersion int64 `json:"binding_version,omitempty"`
Actor string `json:"actor,omitempty"`
ReasonCode string `json:"reason_code,omitempty"`
OperationID string `json:"operation_id,omitempty"`
Action string `json:"action,omitempty"`
Outcome string `json:"outcome,omitempty"`
Details json.RawMessage `json:"details"`
CreatedAt time.Time `json:"created_at"`
}
type AuditFilter struct {
AccountID, BrowserEnvAlias, NetworkExitID, EventType string
From, To *time.Time
Page, PageSize int
}
type AuditPage struct {
Data []AuditEvent `json:"data"`
Total int `json:"total"`
Page int `json:"page"`
PageSize int `json:"page_size"`
}
func Open(ctx context.Context, databaseURL string) (*Store, error) {
db, err := sql.Open("pgx", databaseURL)
if err != nil {
return nil, errors.New("open phase A database")
}
db.SetMaxOpenConns(10)
db.SetMaxIdleConns(2)
db.SetConnMaxIdleTime(5 * time.Minute)
if err := db.PingContext(ctx); err != nil {
db.Close()
return nil, errors.New("connect to phase A database")
}
store := &Store{db: db}
return store, nil
}
func (s *Store) Close() error { return s.db.Close() }
func (s *Store) Ping(ctx context.Context) error { return s.db.PingContext(ctx) }
func (s *Store) CreateAccount(ctx context.Context, account Account, credentials CredentialBridge) (err error) {
if account.Tags == nil {
account.Tags = []string{}
}
if !validAccount(account) || credentials == nil {
return ErrInvalid
}
// 空凭据(扫码登录场景)不写 keyring:凭据留待后续登录/同步链路补齐
stored := false
if account.Cookies != "" {
if err := credentials.Store(ctx, account.CredentialReference, account.CredentialKey, account.Cookies); err != nil {
storeErr := errors.New("store account credential")
if cleanupErr := credentials.Delete(context.WithoutCancel(ctx), account.CredentialReference, account.CredentialKey); cleanupErr != nil {
storeErr = errors.Join(storeErr, errors.New("delete incomplete account credential"))
}
return storeErr
}
stored = true
}
defer func() {
if stored && err != nil && !errors.Is(err, ErrAccountCreationUnknown) {
if cleanupErr := credentials.Delete(context.WithoutCancel(ctx), account.CredentialReference, account.CredentialKey); cleanupErr != nil {
err = errors.Join(err, errors.New("delete orphaned account credential"))
}
}
}()
tx, err := s.db.BeginTx(ctx, nil)
if err != nil {
return errors.New("begin account transaction")
}
defer tx.Rollback()
if _, err := tx.ExecContext(ctx, `
INSERT INTO social_account
(account_id, credential_provider, credential_key, name, platform, platform_account_key, tags, status)
VALUES ($1, $2, $3, $4, $5, $6, $7, 'paused')`, account.ID,
account.CredentialReference.Provider, account.CredentialKey,
account.Name, account.Platform, account.PlatformAccountKey, account.Tags); err != nil {
return publicDatabaseError(err)
}
if err := appendAudit(ctx, tx, "account_created", "account_created", account.ID, map[string]string{"platform": account.Platform}); err != nil {
return err
}
var commitErr error
if s.accountCommit != nil {
commitErr = s.accountCommit(tx)
} else {
commitErr = tx.Commit()
}
if commitErr == nil {
return nil
}
if !commitKnownRolledBack(commitErr) {
return ErrAccountCreationUnknown
}
return errors.New("commit account transaction")
}
func commitKnownRolledBack(err error) bool {
if errors.Is(err, pgx.ErrTxCommitRollback) {
return true
}
var postgresError *pgconn.PgError
return errors.As(err, &postgresError)
}
func (s *Store) ListAccounts(ctx context.Context) ([]Account, error) {
rows, err := s.db.QueryContext(ctx, `
SELECT account.account_id, account.name, account.platform, account.platform_account_key, account.tags,
account.status, account.version
FROM social_account account
ORDER BY account.created_at, account.account_id`)
if err != nil {
return nil, errors.New("read accounts")
}
defer rows.Close()
accounts := []Account{}
for rows.Next() {
account, err := scanAccount(rows)
if err != nil {
return nil, err
}
accounts = append(accounts, account)
}
return accounts, rows.Err()
}
func (s *Store) GetAccount(ctx context.Context, id string) (Account, error) {
if !idPattern.MatchString(id) {
return Account{}, ErrInvalid
}
return scanAccount(s.db.QueryRowContext(ctx, `
SELECT account.account_id, account.name, account.platform, account.platform_account_key, account.tags,
account.status, account.version
FROM social_account account
WHERE account.account_id = $1`, id))
}
func (s *Store) ResolveAccountCredential(ctx context.Context, id string, resolver CredentialResolver) ([]byte, error) {
if !idPattern.MatchString(id) || resolver == nil {
return nil, ErrInvalid
}
reference := CredentialReference{ID: id, Provider: ""}
var key string
if err := s.db.QueryRowContext(ctx, `
SELECT credential_provider, credential_key
FROM social_account
WHERE account_id = $1`, id).Scan(&reference.Provider, &key); err != nil {
return nil, rowError(err)
}
value, err := resolver.Resolve(ctx, reference, key)
if err != nil {
return nil, err
}
if len(value) == 0 {
return nil, ErrInvalid
}
return value, nil
}
type accountScanner interface{ Scan(...any) error }
func scanAccount(row accountScanner) (Account, error) {
var account Account
var tags pgtype.FlatArray[string]
if err := row.Scan(&account.ID, &account.Name, &account.Platform, &account.PlatformAccountKey, pgtype.NewMap().SQLScanner(&tags),
&account.RuntimeStatus, &account.Version); err != nil {
return Account{}, rowError(err)
}
account.Tags = []string(tags)
return account, nil
}
func validAccount(account Account) bool {
if !idPattern.MatchString(account.ID) || strings.TrimSpace(account.Name) != account.Name || account.Name == "" ||
!utf8.ValidString(account.Name) || utf8.RuneCountInString(account.Name) > 128 ||
!platformKeyPattern.MatchString(account.PlatformAccountKey) || len(account.Tags) > 20 ||
len(account.Cookies) > 8192 || !refPattern.MatchString(account.CredentialReference.ID) ||
!credentialKeyPattern.MatchString(account.CredentialKey) ||
(account.CredentialReference.Provider != "os_keyring" && account.CredentialReference.Provider != "secret_manager") {
return false
}
switch account.Platform {
case "douyin":
default:
return false
}
for _, tag := range account.Tags {
if tag == "" || strings.TrimSpace(tag) != tag || !utf8.ValidString(tag) || utf8.RuneCountInString(tag) > 32 {
return false
}
}
// 空凭据合法(扫码登录场景);非空时才校验 Cookie Header 格式
if account.Cookies != "" {
if _, err := http.ParseCookie(account.Cookies); err != nil {
return false
}
}
return true
}
func (s *Store) PauseAccount(ctx context.Context, accountID string) error {
if !idPattern.MatchString(accountID) {
return ErrInvalid
}
tx, err := s.db.BeginTx(ctx, nil)
if err != nil {
return errors.New("begin account state transaction")
}
defer tx.Rollback()
var version int64
var runtimeStatus string
if err := tx.QueryRowContext(ctx, `SELECT version, status FROM social_account WHERE account_id = $1 FOR UPDATE`, accountID).
Scan(&version, &runtimeStatus); err != nil {
return rowError(err)
}
unchanged := runtimeStatus == "paused"
if !unchanged {
if err := tx.QueryRowContext(ctx, `
UPDATE social_account
SET status = 'paused', paused_at = now(),
version = version + 1, updated_at = now()
WHERE account_id = $1 RETURNING version`, accountID).Scan(&version); err != nil {
return errors.New("change account state")
}
if err := appendAudit(ctx, tx, "account_paused", "account_paused", accountID, map[string]any{
"account_version": version,
}); err != nil {
return err
}
}
return commit(tx)
}
func (s *Store) ResumeAccount(ctx context.Context, accountID string) error {
if !idPattern.MatchString(accountID) {
return ErrInvalid
}
tx, err := s.db.BeginTx(ctx, nil)
if err != nil {
return errors.New("begin resume transaction")
}
defer tx.Rollback()
var runtimeStatus string
if err := tx.QueryRowContext(ctx, `SELECT status FROM social_account WHERE account_id = $1 FOR UPDATE`, accountID).
Scan(&runtimeStatus); err != nil {
return rowError(err)
}
var ready bool
if err := tx.QueryRowContext(ctx, `
SELECT EXISTS (
SELECT 1 FROM browser_env environment
LEFT JOIN network_exit network ON network.id = environment.exit_id
JOIN social_account account ON account.id = environment.account_id
WHERE account.account_id = $1 AND (environment.exit_id IS NULL OR network.health_status = 'healthy')
AND NOT environment.runtime_cleanup_pending
AND environment.runtime_id IS NULL
)`, accountID).Scan(&ready); err != nil {
return errors.New("validate account binding")
}
if !ready {
return ErrConflict
}
if runtimeStatus == "active" {
return commit(tx)
}
var version int64
if err := tx.QueryRowContext(ctx, `
UPDATE social_account SET status = 'active', paused_at = NULL, version = version + 1, updated_at = now()
WHERE account_id = $1 RETURNING version`, accountID).Scan(&version); err != nil {
return errors.New("resume account")
}
if err := appendAudit(ctx, tx, "account_resumed", "account_resumed", accountID, map[string]any{"account_version": version}); err != nil {
return err
}
return commit(tx)
}
func (s *Store) Audit(ctx context.Context) ([]AuditEvent, error) {
page, err := s.ListAudit(ctx, AuditFilter{Page: 1, PageSize: 1000})
return page.Data, err
}
func (s *Store) ListAudit(ctx context.Context, filter AuditFilter) (AuditPage, error) {
if (filter.AccountID != "" && !idPattern.MatchString(filter.AccountID)) ||
(filter.BrowserEnvAlias != "" && !refPattern.MatchString(filter.BrowserEnvAlias)) ||
(filter.NetworkExitID != "" && !refPattern.MatchString(filter.NetworkExitID)) ||
(filter.EventType != "" && !eventPattern.MatchString(filter.EventType)) || filter.Page < 1 ||
filter.PageSize < 1 || filter.PageSize > 1000 || (filter.From != nil && filter.To != nil && filter.From.After(*filter.To)) {
return AuditPage{}, ErrInvalid
}
var from, to any
if filter.From != nil {
from = *filter.From
}
if filter.To != nil {
to = *filter.To
}
var total int
err := s.db.QueryRowContext(ctx, `
SELECT count(*) FROM audit_event
WHERE ($1 = '' OR account_id = (SELECT account.id FROM social_account account WHERE account.account_id = $1))
AND ($2 = '' OR browser_env_id = (SELECT environment.id FROM browser_env environment WHERE environment.alias = $2))
AND ($3 = '' OR exit_id = (SELECT network.id FROM network_exit network WHERE network.exit_id = $3)) AND ($4 = '' OR event_type = $4)
AND ($5::timestamptz IS NULL OR created_at >= $5) AND ($6::timestamptz IS NULL OR created_at <= $6)`,
filter.AccountID, filter.BrowserEnvAlias, filter.NetworkExitID, filter.EventType, from, to).Scan(&total)
if err != nil {
return AuditPage{}, errors.New("count audit events")
}
rows, err := s.db.QueryContext(ctx, `
SELECT audit.id, audit.event_type, COALESCE(account.account_id, ''),
COALESCE(environment.alias, ''), COALESCE(network.exit_id, ''), audit.binding_version, audit.actor, audit.reason_code,
audit.operation_id, audit.action, audit.outcome, audit.details, audit.created_at
FROM audit_event audit
LEFT JOIN social_account account ON account.id = audit.account_id
LEFT JOIN browser_env environment ON environment.id = audit.browser_env_id
LEFT JOIN network_exit network ON network.id = audit.exit_id
WHERE ($1 = '' OR audit.account_id = (SELECT account.id FROM social_account account WHERE account.account_id = $1))
AND ($2 = '' OR audit.browser_env_id = (SELECT environment.id FROM browser_env environment WHERE environment.alias = $2))
AND ($3 = '' OR audit.exit_id = (SELECT network.id FROM network_exit network WHERE network.exit_id = $3)) AND ($4 = '' OR audit.event_type = $4)
AND ($5::timestamptz IS NULL OR audit.created_at >= $5) AND ($6::timestamptz IS NULL OR audit.created_at <= $6)
ORDER BY audit.created_at DESC, audit.id DESC LIMIT $7 OFFSET $8`, filter.AccountID,
filter.BrowserEnvAlias, filter.NetworkExitID, filter.EventType, from, to, filter.PageSize, (filter.Page-1)*filter.PageSize)
if err != nil {
return AuditPage{}, errors.New("read audit events")
}
defer rows.Close()
events := []AuditEvent{}
for rows.Next() {
var event AuditEvent
var accountID, browserEnvAlias, networkExitID sql.NullString
var actor, reasonCode, operationID, action, outcome sql.NullString
var bindingVersion sql.NullInt64
if err := rows.Scan(&event.ID, &event.EventType, &accountID,
&browserEnvAlias, &networkExitID, &bindingVersion, &actor, &reasonCode,
&operationID, &action, &outcome, &event.Details, &event.CreatedAt); err != nil {
return AuditPage{}, errors.New("decode audit event")
}
event.AccountID = accountID.String
event.BrowserEnvAlias, event.NetworkExitID = browserEnvAlias.String, networkExitID.String
event.BindingVersion = bindingVersion.Int64
event.Actor, event.ReasonCode = actor.String, reasonCode.String
event.OperationID, event.Action, event.Outcome = operationID.String, action.String, outcome.String
event.Details = safeDetails(event.Details)
events = append(events, event)
}
if err := rows.Err(); err != nil {
return AuditPage{}, errors.New("read audit events")
}
return AuditPage{Data: events, Total: total, Page: filter.Page, PageSize: filter.PageSize}, nil
}
func safeDetails(raw json.RawMessage) json.RawMessage {
var value any
if json.Unmarshal(raw, &value) != nil {
return json.RawMessage(`{}`)
}
value = allowDetails(value)
encoded, err := json.Marshal(value)
if err != nil {
return json.RawMessage(`{}`)
}
return encoded
}
var allowedDetailKeys = map[string]bool{
"platform": true, "account_version": true,
"state": true, "worker_id": true,
}
func allowDetails(value any) any {
switch value := value.(type) {
case map[string]any:
for key, child := range value {
if !allowedDetailKeys[strings.ToLower(key)] {
delete(value, key)
continue
}
value[key] = allowDetails(child)
}
case []any:
for index, child := range value {
value[index] = allowDetails(child)
}
}
return value
}
func appendAudit(ctx context.Context, tx *sql.Tx, eventType, reasonCode, accountID string, details any) error {
if details == nil {
details = map[string]any{}
}
encoded, err := json.Marshal(details)
if err != nil {
return errors.New("encode audit details")
}
_, err = tx.ExecContext(ctx, `
INSERT INTO audit_event
(event_type, account_id,
browser_env_id, exit_id, binding_version, actor, reason_code, details)
SELECT $1, account.id,
environment.id, environment.exit_id, environment.version, 'local-user', $2, $4
FROM (VALUES (1)) AS singleton(value)
LEFT JOIN social_account account ON account.account_id = NULLIF($3, '')
LEFT JOIN browser_env environment ON environment.account_id = account.id`,
eventType, reasonCode, accountID, encoded)
if err != nil {
return errors.New("append audit event")
}
return nil
}
func newID() string {
var value [16]byte
_, _ = rand.Read(value[:])
return hex.EncodeToString(value[:])
}
func NewAccountID() string { return "account-" + newID()[:24] }
func commit(tx *sql.Tx) error {
if err := tx.Commit(); err != nil {
return errors.New("commit transaction")
}
return nil
}
func rowError(err error) error {
if errors.Is(err, sql.ErrNoRows) {
return ErrNotFound
}
return publicDatabaseError(err)
}
func publicDatabaseError(err error) error {
if err == nil {
return nil
}
var postgresError *pgconn.PgError
if errors.As(err, &postgresError) && (postgresError.Code == "23505" || postgresError.Code == "23503" || postgresError.Code == "23514") {
return ErrConflict
}
return errors.New("phase A persistence operation failed")
}