Files
sub2api/backend/internal/service/content_moderation_cyber_test.go
T
李建琦 6d655c9903
Release / update-version (push) Has been cancelled
Release / build-frontend (push) Has been cancelled
Release / release (push) Has been cancelled
Release / sync-version-file (push) Has been cancelled
CI / shell (push) Canceled after 0s
CI / test (push) Canceled after 0s
CI / frontend (push) Canceled after 0s
CI / golangci-lint (push) Canceled after 0s
Security Scan / backend-security (push) Canceled after 0s
Security Scan / frontend-security (push) Canceled after 0s
Sub2API v1.0 - AI API 网关(二开初始版本,基于上游 Wei-Shaw/sub2api)
2026-08-21 18:30:13 +08:00

437 lines
14 KiB
Go

package service
import (
"context"
"errors"
"strings"
"sync"
"testing"
"time"
"github.com/Wei-Shaw/sub2api/internal/pkg/pagination"
"github.com/stretchr/testify/require"
)
// cyberOrderingTestRepo records the sequence of repo calls to verify F7 ordering.
type cyberOrderingTestRepo struct {
mu sync.Mutex
calls []string
emailSents []bool // EmailSent value captured at each CreateLog call
}
func (r *cyberOrderingTestRepo) CreateLog(ctx context.Context, log *ContentModerationLog) error {
r.mu.Lock()
defer r.mu.Unlock()
r.calls = append(r.calls, "create")
if log != nil {
r.emailSents = append(r.emailSents, log.EmailSent)
log.ID = 1 // simulate DB-assigned ID so UpdateLogEmailSent guard passes
}
return nil
}
func (r *cyberOrderingTestRepo) UpdateLogEmailSent(ctx context.Context, id int64, sent bool) error {
r.mu.Lock()
defer r.mu.Unlock()
r.calls = append(r.calls, "update_email_sent")
return nil
}
func (r *cyberOrderingTestRepo) ListLogs(ctx context.Context, filter ContentModerationLogFilter) ([]ContentModerationLog, *pagination.PaginationResult, error) {
return nil, nil, nil
}
func (r *cyberOrderingTestRepo) CountFlaggedByUserSince(ctx context.Context, userID int64, since time.Time, excludeCyberPolicy bool) (int, error) {
return 0, nil
}
func (r *cyberOrderingTestRepo) CleanupExpiredLogs(ctx context.Context, hitBefore time.Time, nonHitBefore time.Time) (*ContentModerationCleanupResult, error) {
return &ContentModerationCleanupResult{}, nil
}
func (r *cyberOrderingTestRepo) snapshot() []string {
r.mu.Lock()
defer r.mu.Unlock()
out := make([]string, len(r.calls))
copy(out, r.calls)
return out
}
func (r *cyberOrderingTestRepo) snapshotEmailSents() []bool {
r.mu.Lock()
defer r.mu.Unlock()
out := make([]bool, len(r.emailSents))
copy(out, r.emailSents)
return out
}
func TestRecordCyberPolicyEvent_DisabledWhenRiskControlOff(t *testing.T) {
repo := &contentModerationTestRepo{}
svc := NewContentModerationService(
&contentModerationTestSettingRepo{values: map[string]string{
SettingKeyRiskControlEnabled: "false",
}},
repo,
nil,
nil,
nil,
nil,
nil,
nil,
)
svc.RecordCyberPolicyEvent(context.Background(), CyberPolicyRecordInput{
UserID: 1,
UserEmail: "u@x.com",
Model: "gpt-5",
Endpoint: "/v1/responses",
UpstreamMessage: "flagged",
UpstreamBody: `{"error":{"code":"cyber_policy"}}`,
UpstreamStatus: 400,
})
require.Empty(t, repo.snapshotLogs(), "CreateLog must NOT be called when risk_control_enabled is off")
}
func TestRecordCyberPolicyEvent_WritesLogWhenEnabled(t *testing.T) {
repo := &contentModerationTestRepo{}
svc := NewContentModerationService(
&contentModerationTestSettingRepo{values: map[string]string{
SettingKeyRiskControlEnabled: "true",
}},
repo,
nil,
nil,
nil,
nil,
nil,
nil, // emailService=nil: email path safely skipped
)
svc.RecordCyberPolicyEvent(context.Background(), CyberPolicyRecordInput{
UserID: 1,
UserEmail: "u@x.com",
Model: "gpt-5",
Endpoint: "/v1/responses",
UpstreamMessage: "flagged",
UpstreamBody: `{"error":{"code":"cyber_policy"}}`,
UpstreamStatus: 400,
})
logs := repo.snapshotLogs()
require.Len(t, logs, 1)
log := logs[0]
require.Equal(t, "cyber_policy", log.Action)
require.True(t, log.Flagged)
require.Equal(t, "cyber_policy", log.HighestCategory)
require.Contains(t, log.Error, "flagged")
require.False(t, log.AutoBanned)
// emailService is nil, so EmailSent must be false
require.False(t, log.EmailSent)
// UserID pointer must be set
require.NotNil(t, log.UserID)
require.Equal(t, int64(1), *log.UserID)
// score for cyber_policy is always 1.0
require.Equal(t, 1.0, log.HighestScore)
// mode must be post_upstream
require.Equal(t, "post_upstream", log.Mode)
// provider
require.Equal(t, "openai", log.Provider)
// model
require.Equal(t, "gpt-5", log.Model)
// endpoint
require.Equal(t, "/v1/responses", log.Endpoint)
// violation count >= 1 (side-effects ran)
require.GreaterOrEqual(t, log.ViolationCount, 1)
// Error field should also contain the upstream body JSON
require.True(t, strings.Contains(log.Error, "cyber_policy") || strings.Contains(log.Error, "flagged"),
"Error should mention flagged or cyber_policy")
}
func TestRecordCyberPolicyEvent_RespectsContentModerationScope(t *testing.T) {
groupID := int64(7)
tests := []struct {
name string
config string
groupID *int64
model string
wantCalls []bool
wantLogs int
wantBanned bool
}{
{
name: "excluded group",
config: `{"all_groups":false,"group_ids":[8],"ban_threshold":1}`,
groupID: &groupID,
model: "gpt-5",
wantLogs: 0,
},
{
name: "ungrouped excluded by selected groups",
config: `{"all_groups":false,"group_ids":[7],"ban_threshold":1}`,
groupID: nil,
model: "gpt-5",
wantLogs: 0,
},
{
name: "excluded model",
config: `{"all_groups":true,"model_filter":{"type":"include","models":["gpt-4o"]},"ban_threshold":1}`,
groupID: &groupID,
model: "gpt-5",
wantLogs: 0,
},
{
name: "included group and model",
config: `{"enabled":false,"mode":"off","sample_rate":0,"all_groups":false,"group_ids":[7],"model_filter":{"type":"include","models":["gpt-5"]},"ban_threshold":1}`,
groupID: &groupID,
model: "gpt-5",
wantCalls: []bool{false},
wantLogs: 1,
wantBanned: true,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
repo := &banCountArgsTestRepo{}
userRepo := &contentModerationTestUserRepo{user: &User{ID: 1, Role: RoleUser, Status: StatusActive}}
svc := NewContentModerationService(
&contentModerationTestSettingRepo{values: map[string]string{
SettingKeyRiskControlEnabled: "true",
SettingKeyContentModerationConfig: tt.config,
}},
repo, nil, nil, userRepo, nil, nil, nil,
)
svc.RecordCyberPolicyEvent(context.Background(), CyberPolicyRecordInput{
UserID: 1,
GroupID: tt.groupID,
Model: tt.model,
})
if tt.wantCalls == nil {
require.Empty(t, repo.snapshotCountCalls())
} else {
require.Equal(t, tt.wantCalls, repo.snapshotCountCalls())
}
require.Len(t, repo.snapshotLogs(), tt.wantLogs)
require.Equal(t, tt.wantBanned, userRepo.user.Status == StatusDisabled)
if tt.wantBanned {
require.Len(t, userRepo.updated, 1)
} else {
require.Empty(t, userRepo.updated)
}
})
}
}
func TestRecordCyberPolicyEvent_InitialRuntimeSnapshotLoadFailureSkipsEvent(t *testing.T) {
repo := &banCountArgsTestRepo{}
settingRepo := &contentModerationRuntimeSettingRepo{values: map[string]string{
SettingKeyRiskControlEnabled: "true",
SettingKeyContentModerationConfig: `{invalid`,
}}
svc := NewContentModerationService(settingRepo, repo, nil, nil, nil, nil, nil, nil)
svc.RecordCyberPolicyEvent(context.Background(), CyberPolicyRecordInput{
UserID: 1,
Model: "gpt-5",
})
require.Empty(t, repo.snapshotCountCalls())
require.Empty(t, repo.snapshotLogs())
getValue, getMultiple := settingRepo.calls()
require.Zero(t, getValue)
require.GreaterOrEqual(t, getMultiple, 1)
}
func TestRecordCyberPolicyEvent_RuntimeSnapshotRefreshFailureKeepsStaleScope(t *testing.T) {
repo := &banCountArgsTestRepo{}
settingRepo := &contentModerationRuntimeSettingRepo{values: map[string]string{
SettingKeyRiskControlEnabled: "true",
SettingKeyContentModerationConfig: `{"all_groups":true,"model_filter":{"type":"include","models":["gpt-5"]}}`,
}}
svc := NewContentModerationService(settingRepo, repo, nil, nil, nil, nil, nil, nil)
svc.runtimeCacheTTL = time.Minute
_, err := svc.loadRuntimeSnapshot(context.Background())
require.NoError(t, err)
current := svc.runtimeSnapshot.Load()
require.NotNil(t, current)
expired := *current
expired.loadedAt = time.Now().Add(-2 * time.Minute)
svc.runtimeSnapshot.Store(&expired)
settingRepo.failMultiple(errors.New("database unavailable"))
svc.RecordCyberPolicyEvent(context.Background(), CyberPolicyRecordInput{
UserID: 1,
Model: "gpt-5",
})
require.Len(t, repo.snapshotLogs(), 1)
require.Eventually(t, func() bool {
_, calls := settingRepo.calls()
return calls == 2
}, time.Second, time.Millisecond)
getValue, getMultiple := settingRepo.calls()
require.Zero(t, getValue)
require.Equal(t, 2, getMultiple)
}
// TestRecordCyberPolicyEvent_CreateLogBeforeEmail verifies F7: the moderation
// log is persisted BEFORE email delivery, and EmailSent is patched afterwards —
// SMTP hangs can no longer swallow the audit record.
//
// Note on email ordering: EmailService is a concrete type with no injectable
// send interface, so SMTP-success cannot be simulated in unit tests.
// With emailService=nil the email block is skipped and UpdateLogEmailSent is not
// called (correct: logPersisted && emailSent guard). The test therefore asserts
// the two invariants that ARE observable without real SMTP:
// 1. CreateLog runs first (calls[0]=="create").
// 2. The log is stored with EmailSent=false (not pre-set to true).
//
// The update_email_sent path is covered by integration/e2e tests where a real
// (or test-double) SMTP endpoint is available.
func TestRecordCyberPolicyEvent_CreateLogBeforeEmail(t *testing.T) {
repo := &cyberOrderingTestRepo{}
svc := NewContentModerationService(
&contentModerationTestSettingRepo{values: map[string]string{
SettingKeyRiskControlEnabled: "true",
}},
repo,
nil,
nil,
nil,
nil,
nil,
nil, // emailService=nil: email path safely skipped; see doc comment above
)
svc.RecordCyberPolicyEvent(context.Background(), CyberPolicyRecordInput{
RequestID: "req-1",
UserID: 7,
UserEmail: "u@example.com",
Model: "gpt-5",
UpstreamMessage: "blocked",
})
calls := repo.snapshot()
require.GreaterOrEqual(t, len(calls), 1, "CreateLog must be called")
require.Equal(t, "create", calls[0], "CreateLog must run first (F7: log-before-email)")
// EmailSent must be false when the log is first persisted (new code sets it
// false before CreateLog; email result is patched via UpdateLogEmailSent).
emailSents := repo.snapshotEmailSents()
require.NotEmpty(t, emailSents, "CreateLog must have captured EmailSent value")
require.False(t, emailSents[0], "log must be stored with EmailSent=false initially (F7)")
// With emailService=nil, no email is sent, so UpdateLogEmailSent must NOT
// be called (logPersisted && emailSent guard correctly suppresses the patch).
require.NotContains(t, calls, "update_email_sent",
"UpdateLogEmailSent must not be called when no email was sent")
}
// banCountArgsTestRepo 在 contentModerationTestRepo 基础上记录
// CountFlaggedByUserSince 收到的 excludeCyberPolicy 参数,供透传断言。
type banCountArgsTestRepo struct {
contentModerationTestRepo
argsMu sync.Mutex
countCalls []bool
}
func (r *banCountArgsTestRepo) CountFlaggedByUserSince(ctx context.Context, userID int64, since time.Time, excludeCyberPolicy bool) (int, error) {
r.argsMu.Lock()
r.countCalls = append(r.countCalls, excludeCyberPolicy)
r.argsMu.Unlock()
return r.contentModerationTestRepo.CountFlaggedByUserSince(ctx, userID, since, excludeCyberPolicy)
}
func (r *banCountArgsTestRepo) snapshotCountCalls() []bool {
r.argsMu.Lock()
defer r.argsMu.Unlock()
out := make([]bool, len(r.countCalls))
copy(out, r.countCalls)
return out
}
func TestApplyFlaggedAccountSideEffects_PassesExcludeCyberFlag(t *testing.T) {
repo := &banCountArgsTestRepo{}
svc := NewContentModerationService(
&contentModerationTestSettingRepo{values: map[string]string{}},
repo, nil, nil, nil, nil, nil, nil,
)
userID := int64(42)
cfgExclude := defaultContentModerationConfig()
cfgExclude.CyberPolicyExcludeFromBanCount = true
svc.applyFlaggedAccountSideEffects(context.Background(), cfgExclude, &ContentModerationLog{Flagged: true, UserID: &userID})
cfgDefault := defaultContentModerationConfig() // 默认 false
svc.applyFlaggedAccountSideEffects(context.Background(), cfgDefault, &ContentModerationLog{Flagged: true, UserID: &userID})
require.Equal(t, []bool{true, false}, repo.snapshotCountCalls(),
"applyFlaggedAccountSideEffects 必须把 cfg.CyberPolicyExcludeFromBanCount 透传给 COUNT 查询")
}
func TestRecordCyberPolicyEvent_ExcludeFromBanCount_SkipsBanJudgment(t *testing.T) {
repo := &banCountArgsTestRepo{}
svc := NewContentModerationService(
&contentModerationTestSettingRepo{values: map[string]string{
SettingKeyRiskControlEnabled: "true",
SettingKeyContentModerationConfig: `{"cyber_policy_exclude_from_ban_count":true}`,
}},
repo, nil, nil, nil, nil, nil, nil,
)
svc.RecordCyberPolicyEvent(context.Background(), CyberPolicyRecordInput{
UserID: 1,
UserEmail: "u@x.com",
Model: "gpt-5",
Endpoint: "/v1/responses",
UpstreamMessage: "flagged",
UpstreamStatus: 400,
})
require.Empty(t, repo.snapshotCountCalls(), "开关开时不得执行封号计数查询")
logs := repo.snapshotLogs()
require.Len(t, logs, 1, "风控日志必须照记")
require.True(t, logs[0].Flagged, "日志仍标记 Flagged=true(列表可见可筛)")
require.Equal(t, "cyber_policy", logs[0].Action)
require.Equal(t, 0, logs[0].ViolationCount, "不参与计数时 ViolationCount 保持 0")
require.False(t, logs[0].AutoBanned)
}
func TestRecordCyberPolicyEvent_DefaultCountsTowardBan(t *testing.T) {
repo := &banCountArgsTestRepo{}
svc := NewContentModerationService(
&contentModerationTestSettingRepo{values: map[string]string{
SettingKeyRiskControlEnabled: "true",
}},
repo, nil, nil, nil, nil, nil, nil,
)
svc.RecordCyberPolicyEvent(context.Background(), CyberPolicyRecordInput{
UserID: 1,
UserEmail: "u@x.com",
Model: "gpt-5",
Endpoint: "/v1/responses",
UpstreamMessage: "flagged",
UpstreamStatus: 400,
})
require.Equal(t, []bool{false}, repo.snapshotCountCalls(),
"默认配置必须执行计数查询且不排除 cyber 行")
logs := repo.snapshotLogs()
require.Len(t, logs, 1)
require.GreaterOrEqual(t, logs[0].ViolationCount, 1, "默认路径行为不变(现状回归)")
}