Files
sub2api/backend/internal/securityaudit/prompt_service_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

146 lines
5.7 KiB
Go

package securityaudit
import (
"context"
"encoding/json"
"strings"
"testing"
"time"
"github.com/Wei-Shaw/sub2api/internal/service"
"github.com/stretchr/testify/require"
)
type staticSettingRepository struct {
values map[string]string
}
func (r staticSettingRepository) Get(context.Context, string) (*service.Setting, error) {
return nil, service.ErrSettingNotFound
}
func (r staticSettingRepository) GetValue(context.Context, string) (string, error) {
return "", service.ErrSettingNotFound
}
func (r staticSettingRepository) Set(context.Context, string, string) error { return nil }
func (r staticSettingRepository) GetMultiple(_ context.Context, keys []string) (map[string]string, error) {
result := make(map[string]string, len(keys))
for _, key := range keys {
result[key] = r.values[key]
}
return result, nil
}
func (r staticSettingRepository) SetMultiple(context.Context, map[string]string) error { return nil }
func (r staticSettingRepository) GetAll(context.Context) (map[string]string, error) {
return r.values, nil
}
func (r staticSettingRepository) Delete(context.Context, string) error { return nil }
func TestPromptServiceHasExplicitIdempotentLifecycle(t *testing.T) {
config := NewConfigManager(nil, staticSettingRepository{values: map[string]string{
SettingKeyPromptAuditConfig: "",
SettingKeyRiskControl: "false",
}}, nil, prefixEncryptor{}, testTotpKeyConfig())
service := NewPromptService(
config,
NewPostgreSQLRepository(nil),
NewRedisPayloadStore(nil),
NewOpenAICompatibleScanner(),
NewAtomicMetrics(),
)
require.Nil(t, service.cancel, "construction must not start background work")
require.NoError(t, service.Start(context.Background()))
require.NotNil(t, service.cancel)
require.NoError(t, service.Start(context.Background()), "Start must be idempotent")
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
defer cancel()
require.NoError(t, service.Shutdown(ctx))
require.Nil(t, service.cancel)
require.NoError(t, service.Shutdown(ctx), "Shutdown must be idempotent")
}
func TestPromptServiceStartReportsDependencyFailureWithoutPanic(t *testing.T) {
service := &PromptService{}
require.Error(t, service.Start(context.Background()))
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
require.NoError(t, service.Shutdown(ctx))
}
func TestPromptServiceBlockingLatestTurnOnlyUsesNarrowSnapshot(t *testing.T) {
seen := make([]string, 0, 2)
evaluator := newGuardEvaluator(PromptScannerFunc(func(_ context.Context, _ ActiveEndpoint, chunk string, _ []string) (*NormalizedResult, error) {
seen = append(seen, chunk)
return &NormalizedResult{Decision: EventPass, RiskLevel: RiskLow, Action: ActionAllow, ScannerScores: map[string]float64{}, ScannerEvidence: map[string]string{}}, nil
}), nil, NewAtomicMetrics(), 2, 2)
service := &PromptService{
config: &fakeConfigStore{active: true, cfg: ActiveConfig{
RiskControlEnabled: true, Enabled: true, BlockingEnabled: true, BlockingLatestTurnOnly: true, AllGroups: true,
Scanners: AllScannerIDs, Endpoints: []ActiveEndpoint{{ID: "guard-1", Enabled: true, TimeoutMS: 1000, InputLimit: 4096}},
}},
evaluator: evaluator,
}
decision, err := service.Evaluate(context.Background(), Request{Protocol: "openai_chat_completions", Body: []byte(`{"messages":[{"role":"system","content":"system instruction"},{"role":"user","content":"older user input"},{"role":"assistant","content":"previous output"},{"role":"user","content":"latest user input"}]}`)})
require.NoError(t, err)
require.Equal(t, DecisionAllow, decision.Kind)
require.Equal(t, []string{"latest user input", "previous output"}, seen)
}
func TestPromptServiceRejectsInvalidDeleteConfirmationClaims(t *testing.T) {
now := time.Date(2026, 7, 16, 10, 0, 0, 0, time.UTC)
start, end := now.Add(-time.Hour), now.Add(time.Hour)
filter := EventFilter{Decision: string(EventCritical), StartAt: &start, EndAt: &end}
const snapshotMaxID int64 = 10
filterHash := FilterHash(filter, snapshotMaxID)
validClaims := deleteClaims{
FilterHash: filterHash, SnapshotMaxID: snapshotMaxID, AdminID: 7,
IssuedAt: now, ExpiresAt: now.Add(5 * time.Minute),
}
claimsToken := func(claims deleteClaims) string {
raw, err := json.Marshal(claims)
require.NoError(t, err)
return string(raw)
}
validRequest := DeleteByFilterRequest{
Filter: filter, SnapshotMaxID: snapshotMaxID, FilterHash: filterHash,
ConfirmationToken: claimsToken(validClaims), Confirm: true,
}
tests := []struct {
name string
request DeleteByFilterRequest
adminID int64
}{
{name: "confirm false", request: func() DeleteByFilterRequest { value := validRequest; value.Confirm = false; return value }(), adminID: 7},
{name: "malformed token", request: func() DeleteByFilterRequest {
value := validRequest
value.ConfirmationToken = "not-json"
return value
}(), adminID: 7},
{name: "different administrator", request: validRequest, adminID: 8},
{name: "filter hash mismatch", request: func() DeleteByFilterRequest {
value := validRequest
value.FilterHash = strings.Repeat("b", 64)
return value
}(), adminID: 7},
{name: "snapshot mismatch", request: func() DeleteByFilterRequest { value := validRequest; value.SnapshotMaxID++; return value }(), adminID: 7},
{name: "expired", request: func() DeleteByFilterRequest {
value := validRequest
claims := validClaims
claims.ExpiresAt = now
value.ConfirmationToken = claimsToken(claims)
return value
}(), adminID: 7},
}
service := &PromptService{config: &fakeConfigStore{}, clock: fixedClock{now: now}}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
result, err := service.DeleteByFilter(context.Background(), test.request, test.adminID)
require.Error(t, err)
require.Nil(t, result)
})
}
}