mirror of
https://github.com/Wei-Shaw/sub2api.git
synced 2026-10-07 13:18:18 +08:00
feat(security-audit): add OpenAI-compatible prompt auditing
This commit is contained in:
@@ -153,6 +153,13 @@ func runMainServer() {
|
||||
log.Fatalf("Failed to initialize application: %v", err)
|
||||
}
|
||||
defer app.Cleanup()
|
||||
if app.PromptAudit != nil {
|
||||
if err := app.PromptAudit.Start(context.Background()); err != nil {
|
||||
// Prompt Audit is default-off and isolated. Startup degradation must be
|
||||
// observable but must not take unrelated APIs down.
|
||||
log.Printf("Prompt Audit started in degraded state: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// 启动服务器
|
||||
go func() {
|
||||
|
||||
@@ -15,6 +15,7 @@ import (
|
||||
"github.com/Wei-Shaw/sub2api/internal/handler"
|
||||
"github.com/Wei-Shaw/sub2api/internal/payment"
|
||||
"github.com/Wei-Shaw/sub2api/internal/repository"
|
||||
"github.com/Wei-Shaw/sub2api/internal/securityaudit"
|
||||
"github.com/Wei-Shaw/sub2api/internal/server"
|
||||
"github.com/Wei-Shaw/sub2api/internal/server/middleware"
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
@@ -24,8 +25,9 @@ import (
|
||||
)
|
||||
|
||||
type Application struct {
|
||||
Server *http.Server
|
||||
Cleanup func()
|
||||
Server *http.Server
|
||||
PromptAudit *securityaudit.PromptService
|
||||
Cleanup func()
|
||||
}
|
||||
|
||||
func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
|
||||
@@ -36,6 +38,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
|
||||
// Business layer ProviderSets
|
||||
repository.ProviderSet,
|
||||
service.ProviderSet,
|
||||
securityaudit.ProviderSet,
|
||||
payment.ProviderSet,
|
||||
middleware.ProviderSet,
|
||||
handler.ProviderSet,
|
||||
@@ -53,7 +56,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
|
||||
provideCleanup,
|
||||
|
||||
// Application struct
|
||||
wire.Struct(new(Application), "Server", "Cleanup"),
|
||||
wire.Struct(new(Application), "Server", "PromptAudit", "Cleanup"),
|
||||
)
|
||||
return nil, nil
|
||||
}
|
||||
@@ -105,6 +108,7 @@ func provideCleanup(
|
||||
quotaFlusher *service.UserPlatformQuotaUsageFlusher,
|
||||
upstreamBillingProbe *service.UpstreamBillingProbeService,
|
||||
auditLog *service.AuditLogService,
|
||||
promptAudit *securityaudit.PromptService,
|
||||
) func() {
|
||||
return func() {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
@@ -117,6 +121,12 @@ func provideCleanup(
|
||||
|
||||
// 应用层清理步骤可并行执行,基础设施资源(Redis/Ent)最后按顺序关闭。
|
||||
parallelSteps := []cleanupStep{
|
||||
{"PromptAuditService", func() error {
|
||||
if promptAudit != nil {
|
||||
return promptAudit.Shutdown(ctx)
|
||||
}
|
||||
return nil
|
||||
}},
|
||||
{"OpsScheduledReportService", func() error {
|
||||
if opsScheduledReport != nil {
|
||||
opsScheduledReport.Stop()
|
||||
|
||||
@@ -14,6 +14,7 @@ import (
|
||||
"github.com/Wei-Shaw/sub2api/internal/handler/admin"
|
||||
"github.com/Wei-Shaw/sub2api/internal/payment"
|
||||
"github.com/Wei-Shaw/sub2api/internal/repository"
|
||||
"github.com/Wei-Shaw/sub2api/internal/securityaudit"
|
||||
"github.com/Wei-Shaw/sub2api/internal/server"
|
||||
"github.com/Wei-Shaw/sub2api/internal/server/middleware"
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
@@ -247,6 +248,13 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
|
||||
contentModerationHashCache := repository.NewContentModerationHashCache(redisClient)
|
||||
contentModerationService := service.NewContentModerationService(settingRepository, contentModerationRepository, contentModerationHashCache, groupRepository, userRepository, apiKeyAuthCacheInvalidator, emailService)
|
||||
contentModerationHandler := admin.NewContentModerationHandler(contentModerationService)
|
||||
configManager := securityaudit.NewConfigManager(db, settingRepository, redisClient, secretEncryptor)
|
||||
postgreSQLRepository := securityaudit.NewPostgreSQLRepository(db)
|
||||
redisPayloadStore := securityaudit.NewRedisPayloadStore(redisClient)
|
||||
openAICompatibleScanner := securityaudit.NewOpenAICompatibleScanner()
|
||||
atomicMetrics := securityaudit.NewAtomicMetrics()
|
||||
promptService := securityaudit.NewPromptService(configManager, postgreSQLRepository, redisPayloadStore, openAICompatibleScanner, atomicMetrics)
|
||||
promptAdminHandler := securityaudit.NewPromptAdminHandler(promptService)
|
||||
paymentHandler := admin.NewPaymentHandler(paymentService, paymentConfigService)
|
||||
affiliateHandler := admin.NewAffiliateHandler(affiliateService, adminService)
|
||||
complianceHandler := admin.NewComplianceHandler(settingService)
|
||||
@@ -254,12 +262,14 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
|
||||
auditLogService := service.ProvideAuditLogService(auditLogRepository, settingService)
|
||||
auditLogHandler := admin.NewAuditLogHandler(auditLogService, totpService)
|
||||
upstreamBillingProbeService := service.ProvideUpstreamBillingProbeService(accountRepository, accountTestService, settingService, leaderLockCache, db)
|
||||
adminHandlers := handler.ProvideAdminHandlers(dashboardHandler, adminUserHandler, groupHandler, accountHandler, adminAnnouncementHandler, dataManagementHandler, backupHandler, oAuthHandler, openAIOAuthHandler, geminiOAuthHandler, antigravityOAuthHandler, grokOAuthHandler, proxyHandler, adminRedeemHandler, promoHandler, settingHandler, opsHandler, systemHandler, adminSubscriptionHandler, adminUsageHandler, userAttributeHandler, errorPassthroughHandler, tlsFingerprintProfileHandler, adminAPIKeyHandler, scheduledTestHandler, channelHandler, channelMonitorHandler, channelMonitorRequestTemplateHandler, contentModerationHandler, paymentHandler, affiliateHandler, complianceHandler, auditLogHandler, upstreamBillingProbeService)
|
||||
adminHandlers := handler.ProvideAdminHandlers(dashboardHandler, adminUserHandler, groupHandler, accountHandler, adminAnnouncementHandler, dataManagementHandler, backupHandler, oAuthHandler, openAIOAuthHandler, geminiOAuthHandler, antigravityOAuthHandler, grokOAuthHandler, proxyHandler, adminRedeemHandler, promoHandler, settingHandler, opsHandler, systemHandler, adminSubscriptionHandler, adminUsageHandler, userAttributeHandler, errorPassthroughHandler, tlsFingerprintProfileHandler, adminAPIKeyHandler, scheduledTestHandler, channelHandler, channelMonitorHandler, channelMonitorRequestTemplateHandler, contentModerationHandler, promptAdminHandler, paymentHandler, affiliateHandler, complianceHandler, auditLogHandler, upstreamBillingProbeService)
|
||||
usageRecordWorkerPool := service.NewUsageRecordWorkerPool(configConfig)
|
||||
userMsgQueueCache := repository.NewUserMsgQueueCache(redisClient)
|
||||
userMessageQueueService := service.ProvideUserMessageQueueService(userMsgQueueCache, rpmCache, configConfig)
|
||||
gatewayHandler := handler.NewGatewayHandler(gatewayService, openAIGatewayService, geminiMessagesCompatService, antigravityGatewayService, userService, concurrencyService, billingCacheService, usageService, apiKeyService, usageRecordWorkerPool, errorPassthroughService, contentModerationService, userMessageQueueService, configConfig, settingService)
|
||||
openAIGatewayHandler := handler.NewOpenAIGatewayHandler(openAIGatewayService, concurrencyService, billingCacheService, apiKeyService, usageRecordWorkerPool, errorPassthroughService, contentModerationService, opsService, configConfig)
|
||||
legacyEngine := securityaudit.NewLegacyModerationAdapter(contentModerationService)
|
||||
coordinator := securityaudit.NewCoordinator(legacyEngine, promptService)
|
||||
gatewayHandler := handler.ProvideGatewayHandler(gatewayService, openAIGatewayService, geminiMessagesCompatService, antigravityGatewayService, userService, concurrencyService, billingCacheService, usageService, apiKeyService, usageRecordWorkerPool, errorPassthroughService, contentModerationService, userMessageQueueService, configConfig, settingService, coordinator)
|
||||
openAIGatewayHandler := handler.ProvideOpenAIGatewayHandler(openAIGatewayService, concurrencyService, billingCacheService, apiKeyService, usageRecordWorkerPool, errorPassthroughService, contentModerationService, opsService, configConfig, coordinator)
|
||||
handlerSettingHandler := handler.ProvideSettingHandler(settingService, buildInfo, notificationEmailService)
|
||||
totpHandler := handler.NewTotpHandler(totpService)
|
||||
handlerPaymentHandler := handler.NewPaymentHandler(paymentService, paymentConfigService)
|
||||
@@ -279,7 +289,7 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
|
||||
batchImageDownloadLimiter := repository.NewBatchImageDownloadLimiter(redisClient, configConfig)
|
||||
batchImageDownloadService := service.NewBatchImageDownloadService(batchImageRepository, accountRepository, batchImageDownloadLimiter, configConfig)
|
||||
batchImageCleanupService := service.ProvideBatchImageCleanupService(batchImageRepository, accountRepository, configConfig)
|
||||
batchImageHandler := handler.NewBatchImageHandler(batchImagePublicService, batchImageDownloadService, batchImageCleanupService)
|
||||
batchImageHandler := handler.ProvideBatchImageHandler(batchImagePublicService, batchImageDownloadService, batchImageCleanupService, openAIGatewayHandler)
|
||||
idempotencyCoordinator := service.ProvideIdempotencyCoordinator(idempotencyRepository, configConfig)
|
||||
idempotencyCleanupService := service.ProvideIdempotencyCleanupService(idempotencyRepository, configConfig)
|
||||
handlers := handler.ProvideHandlers(authHandler, userHandler, apiKeyHandler, usageHandler, redeemHandler, subscriptionHandler, announcementHandler, channelMonitorUserHandler, adminHandlers, gatewayHandler, openAIGatewayHandler, handlerSettingHandler, totpHandler, handlerPaymentHandler, paymentWebhookHandler, availableChannelHandler, asyncImageHandler, batchImageHandler, idempotencyCoordinator, idempotencyCleanupService)
|
||||
@@ -303,10 +313,11 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
|
||||
paymentOrderExpiryService := service.ProvidePaymentOrderExpiryService(paymentService, leaderLockCache, db)
|
||||
channelMonitorRunner := service.ProvideChannelMonitorRunner(channelMonitorService, settingService)
|
||||
userPlatformQuotaUsageFlusher := service.ProvideUserPlatformQuotaUsageFlusher(configConfig, billingCache, serviceUserPlatformQuotaRepository, timingWheelService)
|
||||
v := provideCleanup(client, redisClient, opsMetricsCollector, opsAggregationService, opsAlertEvaluatorService, opsCleanupService, opsScheduledReportService, opsSystemLogSink, schedulerSnapshotService, tokenRefreshService, accountExpiryService, proxyExpiryService, subscriptionExpiryService, usageCleanupService, idempotencyCleanupService, batchImageCleanupService, batchImageWorkerRuntime, pricingService, emailQueueService, billingCacheService, usageRecordWorkerPool, subscriptionService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, grokOAuthService, openAIGatewayService, scheduledTestRunnerService, backupService, paymentOrderExpiryService, channelMonitorRunner, userPlatformQuotaUsageFlusher, upstreamBillingProbeService, auditLogService)
|
||||
v := provideCleanup(client, redisClient, opsMetricsCollector, opsAggregationService, opsAlertEvaluatorService, opsCleanupService, opsScheduledReportService, opsSystemLogSink, schedulerSnapshotService, tokenRefreshService, accountExpiryService, proxyExpiryService, subscriptionExpiryService, usageCleanupService, idempotencyCleanupService, batchImageCleanupService, batchImageWorkerRuntime, pricingService, emailQueueService, billingCacheService, usageRecordWorkerPool, subscriptionService, oAuthService, openAIOAuthService, geminiOAuthService, antigravityOAuthService, grokOAuthService, openAIGatewayService, scheduledTestRunnerService, backupService, paymentOrderExpiryService, channelMonitorRunner, userPlatformQuotaUsageFlusher, upstreamBillingProbeService, auditLogService, promptService)
|
||||
application := &Application{
|
||||
Server: httpServer,
|
||||
Cleanup: v,
|
||||
Server: httpServer,
|
||||
PromptAudit: promptService,
|
||||
Cleanup: v,
|
||||
}
|
||||
return application, nil
|
||||
}
|
||||
@@ -314,8 +325,9 @@ func initializeApplication(buildInfo handler.BuildInfo) (*Application, error) {
|
||||
// wire.go:
|
||||
|
||||
type Application struct {
|
||||
Server *http.Server
|
||||
Cleanup func()
|
||||
Server *http.Server
|
||||
PromptAudit *securityaudit.PromptService
|
||||
Cleanup func()
|
||||
}
|
||||
|
||||
func providePrivacyClientFactory() service.PrivacyClientFactory {
|
||||
@@ -365,6 +377,7 @@ func provideCleanup(
|
||||
quotaFlusher *service.UserPlatformQuotaUsageFlusher,
|
||||
upstreamBillingProbe *service.UpstreamBillingProbeService,
|
||||
auditLog *service.AuditLogService,
|
||||
promptAudit *securityaudit.PromptService,
|
||||
) func() {
|
||||
return func() {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
@@ -376,6 +389,12 @@ func provideCleanup(
|
||||
}
|
||||
|
||||
parallelSteps := []cleanupStep{
|
||||
{"PromptAuditService", func() error {
|
||||
if promptAudit != nil {
|
||||
return promptAudit.Shutdown(ctx)
|
||||
}
|
||||
return nil
|
||||
}},
|
||||
{"OpsScheduledReportService", func() error {
|
||||
if opsScheduledReport != nil {
|
||||
opsScheduledReport.Stop()
|
||||
|
||||
@@ -85,6 +85,7 @@ func TestProvideCleanup_WithMinimalDependencies_NoPanic(t *testing.T) {
|
||||
nil, // quotaFlusher
|
||||
nil, // upstreamBillingProbe
|
||||
nil, // auditLog
|
||||
nil, // promptAudit
|
||||
)
|
||||
|
||||
require.NotPanics(t, func() {
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
@@ -20,6 +21,7 @@ type BatchImageHandler struct {
|
||||
service *service.BatchImagePublicService
|
||||
download *service.BatchImageDownloadService
|
||||
cleanup *service.BatchImageCleanupService
|
||||
openAI *OpenAIGatewayHandler
|
||||
}
|
||||
|
||||
func NewBatchImageHandler(service *service.BatchImagePublicService, download *service.BatchImageDownloadService, cleanup *service.BatchImageCleanupService) *BatchImageHandler {
|
||||
@@ -37,6 +39,9 @@ func (h *BatchImageHandler) Submit(c *gin.Context) {
|
||||
batchImageError(c, infraerrors.New(http.StatusUnauthorized, "API_KEY_REQUIRED", "API key is required"))
|
||||
return
|
||||
}
|
||||
if !h.checkSecurityAuditBeforeSubmit(c, &req) {
|
||||
return
|
||||
}
|
||||
got, err := h.service.Submit(c.Request.Context(), owner, req, c.GetHeader("Idempotency-Key"))
|
||||
if err != nil {
|
||||
batchImageError(c, err)
|
||||
@@ -45,6 +50,44 @@ func (h *BatchImageHandler) Submit(c *gin.Context) {
|
||||
c.JSON(http.StatusOK, got)
|
||||
}
|
||||
|
||||
func (h *BatchImageHandler) checkSecurityAuditBeforeSubmit(c *gin.Context, req *service.BatchImageSubmitRequest) bool {
|
||||
if h == nil || h.openAI == nil || req == nil {
|
||||
return true
|
||||
}
|
||||
apiKey, ok := middleware.GetAPIKeyFromContext(c)
|
||||
if !ok || apiKey == nil {
|
||||
batchImageError(c, infraerrors.New(http.StatusUnauthorized, "API_KEY_REQUIRED", "API key is required"))
|
||||
return false
|
||||
}
|
||||
subject, ok := middleware.GetAuthSubjectFromContext(c)
|
||||
if !ok {
|
||||
batchImageError(c, infraerrors.New(http.StatusInternalServerError, "USER_CONTEXT_REQUIRED", "User context not found"))
|
||||
return false
|
||||
}
|
||||
items := make([]map[string]string, 0, len(req.Items))
|
||||
for _, item := range req.Items {
|
||||
if prompt := strings.TrimSpace(item.Prompt); prompt != "" {
|
||||
items = append(items, map[string]string{"prompt": prompt})
|
||||
}
|
||||
}
|
||||
if len(items) == 0 {
|
||||
return true
|
||||
}
|
||||
body, err := json.Marshal(map[string]any{"request": map[string]any{"items": items}})
|
||||
if err != nil {
|
||||
batchImageError(c, infraerrors.New(http.StatusBadRequest, "INVALID_BATCH_PROMPT", "batch prompts are invalid"))
|
||||
return false
|
||||
}
|
||||
reqLog := requestLogger(c, "handler.batch_image.security_audit",
|
||||
zap.Int64("user_id", subject.UserID), zap.Int64("api_key_id", apiKey.ID), zap.String("model", req.Model))
|
||||
decision := h.openAI.checkSecurityAudit(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIImages, req.Model, body)
|
||||
if decision != nil && !decision.AllowNextStage {
|
||||
h.openAI.openAISecurityAuditError(c, decision)
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func (h *BatchImageHandler) Get(c *gin.Context) {
|
||||
owner, ok := batchImageOwnerFromContext(c)
|
||||
if !ok {
|
||||
|
||||
@@ -12,13 +12,6 @@ import (
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
func (h *GatewayHandler) checkContentModeration(c *gin.Context, reqLog *zap.Logger, apiKey *service.APIKey, subject middleware2.AuthSubject, protocol string, model string, body []byte) *service.ContentModerationDecision {
|
||||
if h == nil || h.contentModerationService == nil {
|
||||
return nil
|
||||
}
|
||||
return runContentModeration(c, reqLog, h.contentModerationService, apiKey, subject, protocol, model, body)
|
||||
}
|
||||
|
||||
func contentModerationStatus(decision *service.ContentModerationDecision) int {
|
||||
if decision == nil || decision.StatusCode < 400 || decision.StatusCode > 599 {
|
||||
return http.StatusForbidden
|
||||
@@ -30,13 +23,6 @@ func contentModerationErrorCode(decision *service.ContentModerationDecision) str
|
||||
return "content_policy_violation"
|
||||
}
|
||||
|
||||
func (h *OpenAIGatewayHandler) checkContentModeration(c *gin.Context, reqLog *zap.Logger, apiKey *service.APIKey, subject middleware2.AuthSubject, protocol string, model string, body []byte) *service.ContentModerationDecision {
|
||||
if h == nil || h.contentModerationService == nil {
|
||||
return nil
|
||||
}
|
||||
return runContentModeration(c, reqLog, h.contentModerationService, apiKey, subject, protocol, model, body)
|
||||
}
|
||||
|
||||
func runContentModeration(c *gin.Context, reqLog *zap.Logger, svc *service.ContentModerationService, apiKey *service.APIKey, subject middleware2.AuthSubject, protocol string, model string, body []byte) *service.ContentModerationDecision {
|
||||
if svc == nil || c == nil || c.Request == nil {
|
||||
return nil
|
||||
|
||||
@@ -25,6 +25,7 @@ import (
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/openai"
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/timezone"
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/xai"
|
||||
"github.com/Wei-Shaw/sub2api/internal/securityaudit"
|
||||
middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware"
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
|
||||
@@ -49,6 +50,7 @@ type GatewayHandler struct {
|
||||
usageRecordWorkerPool *service.UsageRecordWorkerPool
|
||||
errorPassthroughService *service.ErrorPassthroughService
|
||||
contentModerationService *service.ContentModerationService
|
||||
securityAuditCoordinator *securityaudit.Coordinator
|
||||
concurrencyHelper *ConcurrencyHelper
|
||||
userMsgQueueHelper *UserMsgQueueHelper
|
||||
maxAccountSwitches int
|
||||
@@ -199,8 +201,8 @@ func (h *GatewayHandler) Messages(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
if decision := h.checkContentModeration(c, reqLog, apiKey, subject, service.ContentModerationProtocolAnthropicMessages, reqModel, body); decision != nil && decision.Blocked {
|
||||
h.errorResponse(c, contentModerationStatus(decision), contentModerationErrorCode(decision), decision.Message)
|
||||
if decision := h.checkSecurityAudit(c, reqLog, apiKey, subject, service.ContentModerationProtocolAnthropicMessages, reqModel, body); decision != nil && !decision.AllowNextStage {
|
||||
h.anthropicSecurityAuditError(c, decision)
|
||||
return
|
||||
}
|
||||
|
||||
|
||||
@@ -99,8 +99,8 @@ func (h *GatewayHandler) ChatCompletions(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
if decision := h.checkContentModeration(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIChat, reqModel, body); decision != nil && decision.Blocked {
|
||||
h.chatCompletionsErrorResponse(c, contentModerationStatus(decision), contentModerationErrorCode(decision), decision.Message)
|
||||
if decision := h.checkSecurityAudit(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIChat, reqModel, body); decision != nil && !decision.AllowNextStage {
|
||||
h.openAISecurityAuditError(c, decision)
|
||||
return
|
||||
}
|
||||
|
||||
|
||||
@@ -104,8 +104,8 @@ func (h *GatewayHandler) Responses(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
if decision := h.checkContentModeration(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIResponses, reqModel, body); decision != nil && decision.Blocked {
|
||||
h.responsesErrorResponse(c, contentModerationStatus(decision), contentModerationErrorCode(decision), decision.Message)
|
||||
if decision := h.checkSecurityAudit(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIResponses, reqModel, body); decision != nil && !decision.AllowNextStage {
|
||||
h.responsesSecurityAuditError(c, decision)
|
||||
return
|
||||
}
|
||||
|
||||
|
||||
@@ -187,8 +187,8 @@ func (h *GatewayHandler) GeminiV1BetaModels(c *gin.Context) {
|
||||
setOpsRequestContext(c, modelName, stream)
|
||||
setOpsEndpointContext(c, "", int16(service.RequestTypeFromLegacy(stream, false)))
|
||||
|
||||
if decision := h.checkContentModeration(c, reqLog, apiKey, authSubject, service.ContentModerationProtocolGemini, modelName, body); decision != nil && decision.Blocked {
|
||||
googleError(c, contentModerationStatus(decision), decision.Message)
|
||||
if decision := h.checkSecurityAudit(c, reqLog, apiKey, authSubject, service.ContentModerationProtocolGemini, modelName, body); decision != nil && !decision.AllowNextStage {
|
||||
googleSecurityAuditError(c, decision)
|
||||
return
|
||||
}
|
||||
|
||||
|
||||
@@ -114,9 +114,9 @@ func (h *OpenAIGatewayHandler) handleGrokMedia(c *gin.Context, endpoint service.
|
||||
return
|
||||
}
|
||||
if moderationBody := requestInfo.ModerationBody(); len(moderationBody) > 0 {
|
||||
decision := h.checkContentModeration(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIImages, requestModel, moderationBody)
|
||||
if decision != nil && decision.Blocked {
|
||||
h.errorResponse(c, contentModerationStatus(decision), contentModerationErrorCode(decision), decision.Message)
|
||||
decision := h.checkSecurityAudit(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIImages, requestModel, moderationBody)
|
||||
if decision != nil && !decision.AllowNextStage {
|
||||
h.openAISecurityAuditError(c, decision)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
@@ -2,6 +2,7 @@ package handler
|
||||
|
||||
import (
|
||||
"github.com/Wei-Shaw/sub2api/internal/handler/admin"
|
||||
"github.com/Wei-Shaw/sub2api/internal/securityaudit"
|
||||
)
|
||||
|
||||
// AdminHandlers contains all admin-related HTTP handlers
|
||||
@@ -35,6 +36,7 @@ type AdminHandlers struct {
|
||||
ChannelMonitor *admin.ChannelMonitorHandler
|
||||
ChannelMonitorTemplate *admin.ChannelMonitorRequestTemplateHandler
|
||||
ContentModeration *admin.ContentModerationHandler
|
||||
PromptAudit *securityaudit.PromptAdminHandler
|
||||
Payment *admin.PaymentHandler
|
||||
Affiliate *admin.AffiliateHandler
|
||||
Compliance *admin.ComplianceHandler
|
||||
|
||||
@@ -89,6 +89,9 @@ func (h *AsyncImageHandler) Submit(c *gin.Context) {
|
||||
imageTaskJSONError(c, http.StatusBadRequest, "invalid_request_error", err.Error())
|
||||
return
|
||||
}
|
||||
if !h.checkSecurityAuditBeforeSubmit(c, apiKey, platform, body) {
|
||||
return
|
||||
}
|
||||
|
||||
taskCtx, recorder, cancel := newAsyncImageContext(c, body, h.tasks.ExecutionTimeout())
|
||||
task, err := h.tasks.Create(c.Request.Context(), service.ImageTaskOwner{UserID: apiKey.UserID, APIKeyID: apiKey.ID})
|
||||
@@ -115,6 +118,42 @@ func (h *AsyncImageHandler) Submit(c *gin.Context) {
|
||||
go h.run(task.ID, platform, taskCtx, recorder, cancel)
|
||||
}
|
||||
|
||||
func (h *AsyncImageHandler) checkSecurityAuditBeforeSubmit(c *gin.Context, apiKey *service.APIKey, platform string, body []byte) bool {
|
||||
if h == nil || h.openAI == nil {
|
||||
return true
|
||||
}
|
||||
subject, ok := middleware2.GetAuthSubjectFromContext(c)
|
||||
if !ok {
|
||||
imageTaskJSONError(c, http.StatusInternalServerError, "api_error", "User context not found")
|
||||
return false
|
||||
}
|
||||
model := ""
|
||||
moderationBody := body
|
||||
if platform == service.PlatformGrok {
|
||||
parsed := service.ParseGrokMediaRequest(c.GetHeader("Content-Type"), body)
|
||||
model, moderationBody = parsed.Model, parsed.ModerationBody()
|
||||
} else if h.openAI.gatewayService != nil {
|
||||
parsed, err := h.openAI.gatewayService.ParseOpenAIImagesRequest(c, body)
|
||||
if err != nil {
|
||||
imageTaskJSONError(c, http.StatusBadRequest, "invalid_request_error", err.Error())
|
||||
return false
|
||||
}
|
||||
model, moderationBody = parsed.Model, parsed.ModerationBody()
|
||||
}
|
||||
if len(moderationBody) == 0 {
|
||||
c.Set(securityAuditCompletedContextKey, true)
|
||||
return true
|
||||
}
|
||||
reqLog := requestLogger(c, "handler.async_image.security_audit",
|
||||
zap.Int64("user_id", subject.UserID), zap.Int64("api_key_id", apiKey.ID), zap.String("model", model))
|
||||
decision := h.openAI.checkSecurityAudit(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIImages, model, moderationBody)
|
||||
if decision != nil && !decision.AllowNextStage {
|
||||
h.openAI.openAISecurityAuditError(c, decision)
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func (h *AsyncImageHandler) Get(c *gin.Context) {
|
||||
if !h.enabled() {
|
||||
imageTaskJSONError(c, http.StatusNotFound, "not_found_error", "async image tasks are not enabled")
|
||||
|
||||
@@ -78,6 +78,10 @@ func (h *OpenAIGatewayHandler) AlphaSearch(c *gin.Context) {
|
||||
reqLog = reqLog.With(zap.String("model", requestedModel))
|
||||
setOpsRequestContext(c, requestedModel, false)
|
||||
setOpsEndpointContext(c, "", int16(service.RequestTypeSync))
|
||||
if decision := h.checkSecurityAudit(c, reqLog, apiKey, subject, "openai_alpha_search", requestedModel, body); decision != nil && !decision.AllowNextStage {
|
||||
h.openAISecurityAuditError(c, decision)
|
||||
return
|
||||
}
|
||||
|
||||
channelMapping, _ := h.gatewayService.ResolveChannelMappingAndRestrict(c.Request.Context(), apiKey.GroupID, requestedModel)
|
||||
forwardBody := openAIModelMappedBody(body, channelMapping.Mapped, channelMapping.MappedModel, h.gatewayService.ReplaceModelInBody)
|
||||
|
||||
@@ -90,8 +90,8 @@ func (h *OpenAIGatewayHandler) ChatCompletions(c *gin.Context) {
|
||||
setOpsRequestContext(c, reqModel, reqStream)
|
||||
setOpsEndpointContext(c, "", int16(service.RequestTypeFromLegacy(reqStream, false)))
|
||||
|
||||
if decision := h.checkContentModeration(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIChat, reqModel, body); decision != nil && decision.Blocked {
|
||||
h.errorResponse(c, contentModerationStatus(decision), contentModerationErrorCode(decision), decision.Message)
|
||||
if decision := h.checkSecurityAudit(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIChat, reqModel, body); decision != nil && !decision.AllowNextStage {
|
||||
h.openAISecurityAuditError(c, decision)
|
||||
return
|
||||
}
|
||||
if h.rejectIfCyberSessionBlocked(c, apiKey, body, reqModel, cyberBlockFormatChat) {
|
||||
|
||||
@@ -74,6 +74,10 @@ func (h *OpenAIGatewayHandler) Embeddings(c *gin.Context) {
|
||||
reqLog = reqLog.With(zap.String("model", reqModel))
|
||||
setOpsRequestContext(c, reqModel, false)
|
||||
setOpsEndpointContext(c, "", int16(service.RequestTypeSync))
|
||||
if decision := h.checkSecurityAudit(c, reqLog, apiKey, subject, "openai_embeddings", reqModel, body); decision != nil && !decision.AllowNextStage {
|
||||
h.openAISecurityAuditError(c, decision)
|
||||
return
|
||||
}
|
||||
|
||||
channelMapping, _ := h.gatewayService.ResolveChannelMappingAndRestrict(c.Request.Context(), apiKey.GroupID, reqModel)
|
||||
|
||||
|
||||
@@ -15,6 +15,7 @@ import (
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey"
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/ip"
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/logger"
|
||||
"github.com/Wei-Shaw/sub2api/internal/securityaudit"
|
||||
middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware"
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
|
||||
@@ -33,6 +34,7 @@ type OpenAIGatewayHandler struct {
|
||||
usageRecordWorkerPool *service.UsageRecordWorkerPool
|
||||
errorPassthroughService *service.ErrorPassthroughService
|
||||
contentModerationService *service.ContentModerationService
|
||||
securityAuditCoordinator *securityaudit.Coordinator
|
||||
opsService *service.OpsService
|
||||
concurrencyHelper *ConcurrencyHelper
|
||||
imageLimiter *imageConcurrencyLimiter
|
||||
@@ -271,8 +273,8 @@ func (h *OpenAIGatewayHandler) Responses(c *gin.Context) {
|
||||
setOpsRequestContext(c, reqModel, reqStream)
|
||||
setOpsEndpointContext(c, "", int16(service.RequestTypeFromLegacy(reqStream, false)))
|
||||
|
||||
if decision := h.checkContentModeration(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIResponses, reqModel, body); decision != nil && decision.Blocked {
|
||||
h.errorResponse(c, contentModerationStatus(decision), contentModerationErrorCode(decision), decision.Message)
|
||||
if decision := h.checkSecurityAudit(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIResponses, reqModel, body); decision != nil && !decision.AllowNextStage {
|
||||
h.openAISecurityAuditError(c, decision)
|
||||
return
|
||||
}
|
||||
|
||||
@@ -849,8 +851,8 @@ func (h *OpenAIGatewayHandler) Messages(c *gin.Context) {
|
||||
setOpsRequestContext(c, reqModel, reqStream)
|
||||
setOpsEndpointContext(c, "", int16(service.RequestTypeFromLegacy(reqStream, false)))
|
||||
|
||||
if decision := h.checkContentModeration(c, reqLog, apiKey, subject, service.ContentModerationProtocolAnthropicMessages, reqModel, body); decision != nil && decision.Blocked {
|
||||
h.anthropicErrorResponse(c, contentModerationStatus(decision), contentModerationErrorCode(decision), decision.Message)
|
||||
if decision := h.checkSecurityAudit(c, reqLog, apiKey, subject, service.ContentModerationProtocolAnthropicMessages, reqModel, body); decision != nil && !decision.AllowNextStage {
|
||||
h.anthropicSecurityAuditError(c, decision)
|
||||
return
|
||||
}
|
||||
|
||||
@@ -1473,9 +1475,9 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
|
||||
setOpsRequestContext(c, reqModel, true)
|
||||
setOpsEndpointContext(c, "", int16(service.RequestTypeWSV2))
|
||||
|
||||
if decision := h.checkContentModeration(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIResponses, reqModel, firstMessage); decision != nil && decision.Blocked {
|
||||
writeContentModerationWSError(ctx, wsConn, decision)
|
||||
closeOpenAIClientWS(wsConn, coderws.StatusPolicyViolation, decision.Message)
|
||||
if decision := h.checkSecurityAuditStage(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIResponses, reqModel, firstMessage, "first_turn"); decision != nil && !decision.AllowNextStage {
|
||||
writeSecurityAuditWSError(ctx, wsConn, decision)
|
||||
closeOpenAIClientWS(wsConn, securityAuditWSCloseStatus(decision), securityAuditWSCloseReason(decision))
|
||||
return
|
||||
}
|
||||
|
||||
@@ -1727,9 +1729,9 @@ func (h *OpenAIGatewayHandler) ResponsesWebSocket(c *gin.Context) {
|
||||
if model == "" {
|
||||
model = reqModel
|
||||
}
|
||||
if decision := h.checkContentModeration(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIResponses, model, payload); decision != nil && decision.Blocked {
|
||||
writeContentModerationWSError(ctx, wsConn, decision)
|
||||
return service.NewOpenAIWSClientCloseError(coderws.StatusPolicyViolation, decision.Message, nil)
|
||||
if decision := h.checkSecurityAuditStage(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIResponses, model, payload, "subsequent_turn"); decision != nil && !decision.AllowNextStage {
|
||||
writeSecurityAuditWSError(ctx, wsConn, decision)
|
||||
return service.NewOpenAIWSClientCloseError(securityAuditWSCloseStatus(decision), securityAuditWSCloseReason(decision), nil)
|
||||
}
|
||||
return nil
|
||||
},
|
||||
|
||||
@@ -86,8 +86,8 @@ func (h *OpenAIGatewayHandler) Images(c *gin.Context) {
|
||||
h.errorResponse(c, http.StatusForbidden, "permission_error", service.ImageGenerationPermissionMessage())
|
||||
return
|
||||
}
|
||||
if decision := h.checkContentModeration(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIImages, requestModel, parsed.ModerationBody()); decision != nil && decision.Blocked {
|
||||
h.errorResponse(c, contentModerationStatus(decision), contentModerationErrorCode(decision), decision.Message)
|
||||
if decision := h.checkSecurityAudit(c, reqLog, apiKey, subject, service.ContentModerationProtocolOpenAIImages, requestModel, parsed.ModerationBody()); decision != nil && !decision.AllowNextStage {
|
||||
h.openAISecurityAuditError(c, decision)
|
||||
return
|
||||
}
|
||||
imageReleaseFunc, acquired := h.acquireImageGenerationSlot(c, streamStarted)
|
||||
|
||||
@@ -0,0 +1,187 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/googleapi"
|
||||
"github.com/Wei-Shaw/sub2api/internal/securityaudit"
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
coderws "github.com/coder/websocket"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
func (h *OpenAIGatewayHandler) openAISecurityAuditError(c *gin.Context, decision *securityaudit.Decision) {
|
||||
if decision == nil {
|
||||
return
|
||||
}
|
||||
if decision.Legacy != nil && decision.Legacy.Blocked {
|
||||
h.errorResponse(c, securityAuditStatus(decision), securityAuditErrorCode(decision), securityAuditMessage(decision))
|
||||
return
|
||||
}
|
||||
errType := "api_error"
|
||||
if decision.Kind == securityaudit.DecisionBlock {
|
||||
errType = "permission_error"
|
||||
}
|
||||
c.JSON(securityAuditStatus(decision), gin.H{"error": gin.H{
|
||||
"type": errType, "code": securityAuditErrorCode(decision), "message": securityAuditMessage(decision),
|
||||
}})
|
||||
}
|
||||
|
||||
func (h *GatewayHandler) openAISecurityAuditError(c *gin.Context, decision *securityaudit.Decision) {
|
||||
if decision == nil {
|
||||
return
|
||||
}
|
||||
if decision.Legacy != nil && decision.Legacy.Blocked {
|
||||
h.chatCompletionsErrorResponse(c, securityAuditStatus(decision), securityAuditErrorCode(decision), securityAuditMessage(decision))
|
||||
return
|
||||
}
|
||||
errType := "api_error"
|
||||
if decision.Kind == securityaudit.DecisionBlock {
|
||||
errType = "permission_error"
|
||||
}
|
||||
c.JSON(securityAuditStatus(decision), gin.H{"error": gin.H{
|
||||
"type": errType, "code": securityAuditErrorCode(decision), "message": securityAuditMessage(decision),
|
||||
}})
|
||||
}
|
||||
|
||||
func (h *GatewayHandler) responsesSecurityAuditError(c *gin.Context, decision *securityaudit.Decision) {
|
||||
if decision == nil {
|
||||
return
|
||||
}
|
||||
if decision.Legacy != nil && decision.Legacy.Blocked {
|
||||
h.responsesErrorResponse(c, securityAuditStatus(decision), securityAuditErrorCode(decision), securityAuditMessage(decision))
|
||||
return
|
||||
}
|
||||
c.JSON(securityAuditStatus(decision), gin.H{"error": gin.H{
|
||||
"type": "api_error", "code": securityAuditErrorCode(decision), "message": securityAuditMessage(decision),
|
||||
}})
|
||||
}
|
||||
|
||||
func (h *GatewayHandler) anthropicSecurityAuditError(c *gin.Context, decision *securityaudit.Decision) {
|
||||
if decision == nil {
|
||||
return
|
||||
}
|
||||
if decision.Legacy != nil && decision.Legacy.Blocked {
|
||||
h.errorResponse(c, securityAuditStatus(decision), securityAuditErrorCode(decision), securityAuditMessage(decision))
|
||||
return
|
||||
}
|
||||
errType := "api_error"
|
||||
if decision.Kind == securityaudit.DecisionBlock {
|
||||
errType = "permission_error"
|
||||
}
|
||||
c.JSON(securityAuditStatus(decision), gin.H{"type": "error", "error": gin.H{
|
||||
"type": errType, "code": securityAuditErrorCode(decision), "message": securityAuditMessage(decision),
|
||||
}})
|
||||
}
|
||||
|
||||
func (h *OpenAIGatewayHandler) anthropicSecurityAuditError(c *gin.Context, decision *securityaudit.Decision) {
|
||||
if decision == nil {
|
||||
return
|
||||
}
|
||||
if decision.Legacy != nil && decision.Legacy.Blocked {
|
||||
h.anthropicErrorResponse(c, securityAuditStatus(decision), securityAuditErrorCode(decision), securityAuditMessage(decision))
|
||||
return
|
||||
}
|
||||
errType := "api_error"
|
||||
if decision.Kind == securityaudit.DecisionBlock {
|
||||
errType = "permission_error"
|
||||
}
|
||||
c.JSON(securityAuditStatus(decision), gin.H{"type": "error", "error": gin.H{
|
||||
"type": errType, "code": securityAuditErrorCode(decision), "message": securityAuditMessage(decision),
|
||||
}})
|
||||
}
|
||||
|
||||
func googleSecurityAuditError(c *gin.Context, decision *securityaudit.Decision) {
|
||||
if decision == nil {
|
||||
return
|
||||
}
|
||||
if decision.Legacy != nil && decision.Legacy.Blocked {
|
||||
googleError(c, securityAuditStatus(decision), securityAuditMessage(decision))
|
||||
return
|
||||
}
|
||||
status := securityAuditStatus(decision)
|
||||
googleStatus := googleapi.HTTPStatusToGoogleStatus(status)
|
||||
if status == http.StatusServiceUnavailable {
|
||||
googleStatus = "UNAVAILABLE"
|
||||
}
|
||||
requestID := ""
|
||||
if c != nil && c.Request != nil {
|
||||
requestID = contentModerationRequestID(c.Request.Context())
|
||||
}
|
||||
c.JSON(status, gin.H{"error": gin.H{
|
||||
"code": status, "message": securityAuditMessage(decision), "status": googleStatus,
|
||||
"details": []gin.H{{
|
||||
"@type": "type.googleapis.com/google.rpc.ErrorInfo",
|
||||
"reason": securityAuditErrorCode(decision), "domain": "sub2api.securityaudit",
|
||||
"metadata": gin.H{"request_id": requestID},
|
||||
}},
|
||||
}})
|
||||
}
|
||||
|
||||
func writeSecurityAuditWSError(ctx context.Context, conn *coderws.Conn, decision *securityaudit.Decision) {
|
||||
if conn == nil || decision == nil {
|
||||
return
|
||||
}
|
||||
if decision.Legacy != nil && decision.Legacy.Blocked {
|
||||
legacy := decision.Legacy
|
||||
writeContentModerationWSError(ctx, conn, (legacyContentModerationDecision{legacy}).toService())
|
||||
return
|
||||
}
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
}
|
||||
payload, err := json.Marshal(gin.H{
|
||||
"event_id": "evt_prompt_guard_rejected", "type": "error",
|
||||
"error": gin.H{"type": "invalid_request_error", "code": securityAuditErrorCode(decision), "message": securityAuditMessage(decision)},
|
||||
})
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
writeCtx, cancel := context.WithTimeout(ctx, 2*time.Second)
|
||||
defer cancel()
|
||||
_ = conn.Write(writeCtx, coderws.MessageText, payload)
|
||||
}
|
||||
|
||||
type legacyContentModerationDecision struct{ value *securityaudit.LegacyDecision }
|
||||
|
||||
func (d legacyContentModerationDecision) toService() *service.ContentModerationDecision {
|
||||
if d.value == nil {
|
||||
return nil
|
||||
}
|
||||
return &service.ContentModerationDecision{Allowed: d.value.Allowed, Blocked: d.value.Blocked, Flagged: d.value.Flagged, Message: d.value.Message, StatusCode: d.value.StatusCode, Action: d.value.Action}
|
||||
}
|
||||
|
||||
func securityAuditWSCloseStatus(decision *securityaudit.Decision) coderws.StatusCode {
|
||||
if decision == nil {
|
||||
return coderws.StatusInternalError
|
||||
}
|
||||
if decision.Legacy != nil && decision.Legacy.Blocked {
|
||||
return coderws.StatusPolicyViolation
|
||||
}
|
||||
if decision.Kind == securityaudit.DecisionBlock {
|
||||
return coderws.StatusCode(4403)
|
||||
}
|
||||
return coderws.StatusTryAgainLater
|
||||
}
|
||||
|
||||
func securityAuditWSCloseReason(decision *securityaudit.Decision) string {
|
||||
if decision == nil {
|
||||
return securityaudit.ErrorCodeUnavailable
|
||||
}
|
||||
if decision.Legacy != nil && decision.Legacy.Blocked {
|
||||
message := strings.TrimSpace(decision.Legacy.Message)
|
||||
if message != "" {
|
||||
return message
|
||||
}
|
||||
return "content_policy_violation"
|
||||
}
|
||||
code := securityAuditErrorCode(decision)
|
||||
if code == "" {
|
||||
return securityaudit.ErrorCodeUnavailable
|
||||
}
|
||||
return code
|
||||
}
|
||||
@@ -0,0 +1,159 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/ctxkey"
|
||||
"github.com/Wei-Shaw/sub2api/internal/securityaudit"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func promptGuardDecision(kind securityaudit.DecisionKind) *securityaudit.Decision {
|
||||
decision := &securityaudit.Decision{Kind: kind, AllowNextStage: false}
|
||||
switch kind {
|
||||
case securityaudit.DecisionBlock:
|
||||
decision.HTTPStatus = http.StatusForbidden
|
||||
decision.ErrorCode = securityaudit.ErrorCodeBlocked
|
||||
decision.ClientMessage = "提示词安全审计拒绝了该请求,请调整输入后重试"
|
||||
case securityaudit.DecisionInvalid:
|
||||
decision.HTTPStatus = http.StatusServiceUnavailable
|
||||
decision.ErrorCode = securityaudit.ErrorCodeInvalidResponse
|
||||
decision.ClientMessage = "提示词安全审计暂时不可用,请稍后重试"
|
||||
default:
|
||||
decision.HTTPStatus = http.StatusServiceUnavailable
|
||||
decision.ErrorCode = securityaudit.ErrorCodeUnavailable
|
||||
decision.ClientMessage = "提示词安全审计暂时不可用,请稍后重试"
|
||||
}
|
||||
return decision
|
||||
}
|
||||
|
||||
func securityAuditErrorTestContext(t *testing.T) (*gin.Context, *httptest.ResponseRecorder) {
|
||||
t.Helper()
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
ctx := context.WithValue(context.Background(), ctxkey.RequestID, "request-error-golden")
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/test", nil).WithContext(ctx)
|
||||
return c, recorder
|
||||
}
|
||||
|
||||
func decodeErrorJSON(t *testing.T, recorder *httptest.ResponseRecorder) map[string]any {
|
||||
t.Helper()
|
||||
var payload map[string]any
|
||||
require.NoError(t, json.Unmarshal(recorder.Body.Bytes(), &payload))
|
||||
return payload
|
||||
}
|
||||
|
||||
func requireObject(t *testing.T, value any) map[string]any {
|
||||
t.Helper()
|
||||
object, ok := value.(map[string]any)
|
||||
require.True(t, ok)
|
||||
return object
|
||||
}
|
||||
|
||||
func requireArray(t *testing.T, value any) []any {
|
||||
t.Helper()
|
||||
array, ok := value.([]any)
|
||||
require.True(t, ok)
|
||||
return array
|
||||
}
|
||||
|
||||
func TestPromptGuardOpenAIAndClaudeErrorEnvelopesGolden(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
for _, kind := range []securityaudit.DecisionKind{securityaudit.DecisionBlock, securityaudit.DecisionUnavailable, securityaudit.DecisionInvalid} {
|
||||
decision := promptGuardDecision(kind)
|
||||
t.Run("openai_"+string(kind), func(t *testing.T) {
|
||||
c, recorder := securityAuditErrorTestContext(t)
|
||||
(&OpenAIGatewayHandler{}).openAISecurityAuditError(c, decision)
|
||||
require.Equal(t, decision.HTTPStatus, recorder.Code)
|
||||
payload := decodeErrorJSON(t, recorder)
|
||||
errorObject := requireObject(t, payload["error"])
|
||||
require.Equal(t, decision.ErrorCode, errorObject["code"])
|
||||
if kind == securityaudit.DecisionBlock {
|
||||
require.Equal(t, "permission_error", errorObject["type"])
|
||||
} else {
|
||||
require.Equal(t, "api_error", errorObject["type"])
|
||||
}
|
||||
require.NotContains(t, recorder.Body.String(), "raw prompt")
|
||||
require.NotContains(t, recorder.Body.String(), "guard-one")
|
||||
})
|
||||
|
||||
t.Run("responses_"+string(kind), func(t *testing.T) {
|
||||
c, recorder := securityAuditErrorTestContext(t)
|
||||
(&GatewayHandler{}).responsesSecurityAuditError(c, decision)
|
||||
require.Equal(t, decision.HTTPStatus, recorder.Code)
|
||||
errorObject := requireObject(t, decodeErrorJSON(t, recorder)["error"])
|
||||
require.Equal(t, decision.ErrorCode, errorObject["code"])
|
||||
require.Equal(t, "api_error", errorObject["type"])
|
||||
})
|
||||
|
||||
t.Run("claude_"+string(kind), func(t *testing.T) {
|
||||
c, recorder := securityAuditErrorTestContext(t)
|
||||
(&GatewayHandler{}).anthropicSecurityAuditError(c, decision)
|
||||
require.Equal(t, decision.HTTPStatus, recorder.Code)
|
||||
payload := decodeErrorJSON(t, recorder)
|
||||
require.Equal(t, "error", payload["type"])
|
||||
errorObject := requireObject(t, payload["error"])
|
||||
require.Equal(t, decision.ErrorCode, errorObject["code"])
|
||||
if kind == securityaudit.DecisionBlock {
|
||||
require.Equal(t, "permission_error", errorObject["type"])
|
||||
} else {
|
||||
require.Equal(t, "api_error", errorObject["type"])
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPromptGuardGeminiErrorEnvelopeGolden(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
for _, kind := range []securityaudit.DecisionKind{securityaudit.DecisionBlock, securityaudit.DecisionUnavailable, securityaudit.DecisionInvalid} {
|
||||
decision := promptGuardDecision(kind)
|
||||
c, recorder := securityAuditErrorTestContext(t)
|
||||
googleSecurityAuditError(c, decision)
|
||||
require.Equal(t, decision.HTTPStatus, recorder.Code)
|
||||
payload := decodeErrorJSON(t, recorder)
|
||||
errorObject := requireObject(t, payload["error"])
|
||||
require.Equal(t, float64(decision.HTTPStatus), errorObject["code"], "Gemini code must remain numeric")
|
||||
if decision.HTTPStatus == http.StatusForbidden {
|
||||
require.Equal(t, "PERMISSION_DENIED", errorObject["status"])
|
||||
} else {
|
||||
require.Equal(t, "UNAVAILABLE", errorObject["status"])
|
||||
}
|
||||
details := requireArray(t, errorObject["details"])
|
||||
require.Len(t, details, 1)
|
||||
errorInfo := requireObject(t, details[0])
|
||||
require.Equal(t, "type.googleapis.com/google.rpc.ErrorInfo", errorInfo["@type"])
|
||||
require.Equal(t, decision.ErrorCode, errorInfo["reason"])
|
||||
require.Equal(t, "sub2api.securityaudit", errorInfo["domain"])
|
||||
metadata := requireObject(t, errorInfo["metadata"])
|
||||
require.Equal(t, map[string]any{"request_id": "request-error-golden"}, metadata)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPromptGuardWebSocketCloseMappingGolden(t *testing.T) {
|
||||
require.Equal(t, int64(4403), int64(securityAuditWSCloseStatus(promptGuardDecision(securityaudit.DecisionBlock))))
|
||||
require.Equal(t, securityaudit.ErrorCodeBlocked, securityAuditWSCloseReason(promptGuardDecision(securityaudit.DecisionBlock)))
|
||||
require.Equal(t, int64(1013), int64(securityAuditWSCloseStatus(promptGuardDecision(securityaudit.DecisionUnavailable))))
|
||||
require.Equal(t, securityaudit.ErrorCodeUnavailable, securityAuditWSCloseReason(promptGuardDecision(securityaudit.DecisionUnavailable)))
|
||||
require.Equal(t, int64(1013), int64(securityAuditWSCloseStatus(promptGuardDecision(securityaudit.DecisionInvalid))))
|
||||
require.Equal(t, securityaudit.ErrorCodeInvalidResponse, securityAuditWSCloseReason(promptGuardDecision(securityaudit.DecisionInvalid)))
|
||||
}
|
||||
|
||||
func TestLegacyModerationErrorKeepsExistingClientPriority(t *testing.T) {
|
||||
legacy := &securityaudit.Decision{
|
||||
Kind: securityaudit.DecisionBlock, HTTPStatus: http.StatusForbidden,
|
||||
ErrorCode: "content_policy_violation", ClientMessage: "legacy exact message",
|
||||
Legacy: &securityaudit.LegacyDecision{Blocked: true, StatusCode: http.StatusForbidden, ErrorCode: "content_policy_violation", Message: "legacy exact message"},
|
||||
Prompt: &securityaudit.PromptDecision{Kind: securityaudit.DecisionBlock, ErrorCode: securityaudit.ErrorCodeBlocked},
|
||||
}
|
||||
c, recorder := securityAuditErrorTestContext(t)
|
||||
(&GatewayHandler{}).openAISecurityAuditError(c, legacy)
|
||||
require.Equal(t, http.StatusForbidden, recorder.Code)
|
||||
require.Contains(t, recorder.Body.String(), "legacy exact message")
|
||||
require.Contains(t, recorder.Body.String(), "content_policy_violation")
|
||||
require.NotContains(t, recorder.Body.String(), securityaudit.ErrorCodeBlocked)
|
||||
}
|
||||
@@ -0,0 +1,138 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/securityaudit"
|
||||
middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware"
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
"github.com/gin-gonic/gin"
|
||||
"go.uber.org/zap"
|
||||
)
|
||||
|
||||
const securityAuditCompletedContextKey = "sub2api.security_audit.completed"
|
||||
|
||||
func (h *GatewayHandler) checkSecurityAudit(c *gin.Context, reqLog *zap.Logger, apiKey *service.APIKey, subject middleware2.AuthSubject, protocol, model string, body []byte) *securityaudit.Decision {
|
||||
if h == nil {
|
||||
return nil
|
||||
}
|
||||
return runSecurityAudit(c, reqLog, h.securityAuditCoordinator, h.contentModerationService, apiKey, subject, protocol, model, body, "http")
|
||||
}
|
||||
|
||||
func (h *OpenAIGatewayHandler) checkSecurityAudit(c *gin.Context, reqLog *zap.Logger, apiKey *service.APIKey, subject middleware2.AuthSubject, protocol, model string, body []byte) *securityaudit.Decision {
|
||||
if h == nil {
|
||||
return nil
|
||||
}
|
||||
return runSecurityAudit(c, reqLog, h.securityAuditCoordinator, h.contentModerationService, apiKey, subject, protocol, model, body, "http")
|
||||
}
|
||||
|
||||
func (h *OpenAIGatewayHandler) checkSecurityAuditStage(c *gin.Context, reqLog *zap.Logger, apiKey *service.APIKey, subject middleware2.AuthSubject, protocol, model string, body []byte, stage string) *securityaudit.Decision {
|
||||
if h == nil {
|
||||
return nil
|
||||
}
|
||||
return runSecurityAudit(c, reqLog, h.securityAuditCoordinator, h.contentModerationService, apiKey, subject, protocol, model, body, stage)
|
||||
}
|
||||
|
||||
func runSecurityAudit(c *gin.Context, reqLog *zap.Logger, coordinator *securityaudit.Coordinator, legacy *service.ContentModerationService, apiKey *service.APIKey, subject middleware2.AuthSubject, protocol, model string, body []byte, stage string) *securityaudit.Decision {
|
||||
if c == nil || c.Request == nil {
|
||||
return nil
|
||||
}
|
||||
if completed, exists := c.Get(securityAuditCompletedContextKey); exists && completed == true {
|
||||
return nil
|
||||
}
|
||||
if coordinator == nil {
|
||||
legacyDecision := runContentModeration(c, reqLog, legacy, apiKey, subject, protocol, model, body)
|
||||
if legacyDecision == nil {
|
||||
return nil
|
||||
}
|
||||
decision := securityaudit.Decision{Kind: securityaudit.DecisionAllow, HTTPStatus: http.StatusOK, AllowNextStage: true}
|
||||
decision.Legacy = &securityaudit.LegacyDecision{
|
||||
Allowed: legacyDecision.Allowed, Blocked: legacyDecision.Blocked, Flagged: legacyDecision.Flagged,
|
||||
Message: legacyDecision.Message, StatusCode: legacyDecision.StatusCode,
|
||||
ErrorCode: "content_policy_violation", Action: legacyDecision.Action,
|
||||
}
|
||||
if legacyDecision.Blocked {
|
||||
decision.Kind, decision.HTTPStatus, decision.ErrorCode, decision.ClientMessage, decision.AllowNextStage = securityaudit.DecisionBlock, contentModerationStatus(legacyDecision), "content_policy_violation", legacyDecision.Message, false
|
||||
}
|
||||
if decision.AllowNextStage {
|
||||
c.Set(securityAuditCompletedContextKey, true)
|
||||
}
|
||||
return &decision
|
||||
}
|
||||
request := buildSecurityAuditRequest(c, apiKey, subject, protocol, model, body, stage)
|
||||
if reqLog != nil {
|
||||
reqLog.Info("security_audit.gateway_check_start",
|
||||
zap.String("request_id", request.RequestID), zap.Int64("user_id", request.UserID),
|
||||
zap.Int64("api_key_id", request.APIKeyID), zap.Int64p("group_id", request.GroupID),
|
||||
zap.String("endpoint", request.Endpoint), zap.String("provider", request.Provider),
|
||||
zap.String("protocol", request.Protocol), zap.String("model", request.Model), zap.String("stage", request.Stage),
|
||||
zap.Int("body_bytes", len(body)))
|
||||
}
|
||||
decision := coordinator.Check(c.Request.Context(), request)
|
||||
if decision.AllowNextStage {
|
||||
c.Set(securityAuditCompletedContextKey, true)
|
||||
}
|
||||
if reqLog != nil {
|
||||
reqLog.Info("security_audit.gateway_check_done",
|
||||
zap.String("request_id", request.RequestID), zap.String("decision", string(decision.Kind)),
|
||||
zap.String("error_code", decision.ErrorCode), zap.Bool("allow_next_stage", decision.AllowNextStage),
|
||||
zap.String("stage", request.Stage))
|
||||
}
|
||||
return &decision
|
||||
}
|
||||
|
||||
func buildSecurityAuditRequest(c *gin.Context, apiKey *service.APIKey, subject middleware2.AuthSubject, protocol, model string, body []byte, stage string) securityaudit.Request {
|
||||
legacy := buildContentModerationInput(c, apiKey, subject, protocol, model, body)
|
||||
request := securityaudit.Request{
|
||||
RequestID: legacy.RequestID, UserID: legacy.UserID, UserEmail: legacy.UserEmail,
|
||||
APIKeyID: legacy.APIKeyID, APIKeyName: legacy.APIKeyName, GroupID: cloneSecurityAuditGroupID(legacy.GroupID),
|
||||
GroupName: legacy.GroupName, Provider: legacy.Provider, Endpoint: legacy.Endpoint,
|
||||
Protocol: legacy.Protocol, Model: legacy.Model, Body: body, Stage: strings.TrimSpace(stage),
|
||||
}
|
||||
if apiKey != nil && apiKey.User != nil {
|
||||
request.Username = apiKey.User.Username
|
||||
if request.UserEmail == "" {
|
||||
request.UserEmail = apiKey.User.Email
|
||||
}
|
||||
}
|
||||
if request.Stage == "" {
|
||||
request.Stage = "http"
|
||||
}
|
||||
return request
|
||||
}
|
||||
|
||||
func securityAuditStatus(decision *securityaudit.Decision) int {
|
||||
if decision == nil || decision.HTTPStatus < 400 || decision.HTTPStatus > 599 {
|
||||
return http.StatusForbidden
|
||||
}
|
||||
return decision.HTTPStatus
|
||||
}
|
||||
|
||||
func securityAuditErrorCode(decision *securityaudit.Decision) string {
|
||||
if decision == nil || strings.TrimSpace(decision.ErrorCode) == "" {
|
||||
return "content_policy_violation"
|
||||
}
|
||||
return decision.ErrorCode
|
||||
}
|
||||
|
||||
func securityAuditMessage(decision *securityaudit.Decision) string {
|
||||
if decision == nil {
|
||||
return "Request blocked by content policy"
|
||||
}
|
||||
if decision.Legacy != nil && decision.Legacy.Blocked && strings.TrimSpace(decision.Legacy.Message) != "" {
|
||||
return decision.Legacy.Message
|
||||
}
|
||||
if strings.TrimSpace(decision.ClientMessage) != "" {
|
||||
return decision.ClientMessage
|
||||
}
|
||||
return "Request blocked by content policy"
|
||||
}
|
||||
|
||||
func cloneSecurityAuditGroupID(value *int64) *int64 {
|
||||
if value == nil {
|
||||
return nil
|
||||
}
|
||||
cloned := *value
|
||||
return &cloned
|
||||
}
|
||||
@@ -0,0 +1,208 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/securityaudit"
|
||||
middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware"
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
type handlerPromptEngine struct {
|
||||
mu sync.Mutex
|
||||
|
||||
mode securityaudit.Mode
|
||||
decision *securityaudit.PromptDecision
|
||||
err error
|
||||
evaluated int
|
||||
enqueued int
|
||||
requests []securityaudit.Request
|
||||
}
|
||||
|
||||
func (e *handlerPromptEngine) EffectiveMode() securityaudit.Mode { return e.mode }
|
||||
func (e *handlerPromptEngine) Enqueue(_ context.Context, req securityaudit.Request) error {
|
||||
e.mu.Lock()
|
||||
defer e.mu.Unlock()
|
||||
e.enqueued++
|
||||
e.requests = append(e.requests, req.Clone())
|
||||
return e.err
|
||||
}
|
||||
func (e *handlerPromptEngine) Evaluate(_ context.Context, req securityaudit.Request) (*securityaudit.PromptDecision, error) {
|
||||
e.mu.Lock()
|
||||
defer e.mu.Unlock()
|
||||
e.evaluated++
|
||||
e.requests = append(e.requests, req.Clone())
|
||||
return e.decision, e.err
|
||||
}
|
||||
func (e *handlerPromptEngine) snapshot() (evaluated, enqueued int, requests []securityaudit.Request) {
|
||||
e.mu.Lock()
|
||||
defer e.mu.Unlock()
|
||||
requests = make([]securityaudit.Request, len(e.requests))
|
||||
copy(requests, e.requests)
|
||||
return e.evaluated, e.enqueued, requests
|
||||
}
|
||||
|
||||
func securityAuditMediaTestMiddleware(c *gin.Context) {
|
||||
groupID := int64(3)
|
||||
user := &service.User{ID: 7, Username: "media-user", Email: "media@example.test"}
|
||||
c.Set(string(middleware2.ContextKeyAPIKey), &service.APIKey{
|
||||
ID: 9, UserID: 7, User: user, Name: "media-key", GroupID: &groupID,
|
||||
Group: &service.Group{ID: groupID, Name: "media-group", Platform: service.PlatformOpenAI, AllowImageGeneration: true},
|
||||
})
|
||||
c.Set(string(middleware2.ContextKeyUser), middleware2.AuthSubject{UserID: 7, Concurrency: 2})
|
||||
c.Next()
|
||||
}
|
||||
|
||||
func blockingHandlerPromptEngine() *handlerPromptEngine {
|
||||
return &handlerPromptEngine{mode: securityaudit.ModeBlocking, decision: &securityaudit.PromptDecision{
|
||||
Kind: securityaudit.DecisionBlock, ErrorCode: securityaudit.ErrorCodeBlocked, AllowNextStage: false,
|
||||
}}
|
||||
}
|
||||
|
||||
func TestAsyncImagePromptGuardRunsBeforeTaskCreation(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
store := &asyncImageMemoryStore{tasks: map[string]*service.ImageTaskRecord{}}
|
||||
tasks := service.NewImageTaskServiceWithUploader(store, nil, time.Hour, time.Minute)
|
||||
engine := blockingHandlerPromptEngine()
|
||||
openAI := &OpenAIGatewayHandler{securityAuditCoordinator: securityaudit.NewCoordinator(nil, engine)}
|
||||
h := &AsyncImageHandler{tasks: tasks, openAI: openAI}
|
||||
executions := 0
|
||||
h.execute = func(string, *gin.Context) { executions++ }
|
||||
|
||||
router := gin.New()
|
||||
router.Use(securityAuditMediaTestMiddleware)
|
||||
router.POST("/v1/images/generations/async", h.Submit)
|
||||
request := httptest.NewRequest(http.MethodPost, "/v1/images/generations/async", strings.NewReader(`{"model":"gpt-image-2","prompt":"blocked async prompt"}`))
|
||||
request.Header.Set("Content-Type", "application/json")
|
||||
recorder := httptest.NewRecorder()
|
||||
router.ServeHTTP(recorder, request)
|
||||
|
||||
require.Equal(t, http.StatusForbidden, recorder.Code)
|
||||
require.Contains(t, recorder.Body.String(), securityaudit.ErrorCodeBlocked)
|
||||
require.Empty(t, store.tasks, "no asynchronous task may exist after a blocking decision")
|
||||
require.Zero(t, executions)
|
||||
evaluated, _, requests := engine.snapshot()
|
||||
require.Equal(t, 1, evaluated)
|
||||
require.Len(t, requests, 1)
|
||||
require.Contains(t, string(requests[0].Body), "blocked async prompt")
|
||||
}
|
||||
|
||||
func TestAsyncImageSuccessfulPrecheckIsNotRepeatedByDetachedExecution(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
store := &asyncImageMemoryStore{tasks: map[string]*service.ImageTaskRecord{}}
|
||||
tasks := service.NewImageTaskServiceWithUploader(store, nil, time.Hour, time.Minute)
|
||||
engine := &handlerPromptEngine{mode: securityaudit.ModeBlocking, decision: &securityaudit.PromptDecision{Kind: securityaudit.DecisionAllow, AllowNextStage: true}}
|
||||
openAI := &OpenAIGatewayHandler{securityAuditCoordinator: securityaudit.NewCoordinator(nil, engine)}
|
||||
h := &AsyncImageHandler{tasks: tasks, openAI: openAI}
|
||||
var executionMu sync.Mutex
|
||||
repeatedDecision := false
|
||||
h.execute = func(_ string, c *gin.Context) {
|
||||
apiKey, _ := middleware2.GetAPIKeyFromContext(c)
|
||||
subject, _ := middleware2.GetAuthSubjectFromContext(c)
|
||||
decision := openAI.checkSecurityAudit(c, nil, apiKey, subject, service.ContentModerationProtocolOpenAIImages, "gpt-image-2", []byte(`{"prompt":"must not rescan"}`))
|
||||
executionMu.Lock()
|
||||
repeatedDecision = decision != nil
|
||||
executionMu.Unlock()
|
||||
c.JSON(http.StatusOK, gin.H{"created": 1, "data": []any{}})
|
||||
}
|
||||
|
||||
router := gin.New()
|
||||
router.Use(securityAuditMediaTestMiddleware)
|
||||
router.POST("/v1/images/generations/async", h.Submit)
|
||||
request := httptest.NewRequest(http.MethodPost, "/v1/images/generations/async", strings.NewReader(`{"model":"gpt-image-2","prompt":"allowed async prompt"}`))
|
||||
request.Header.Set("Content-Type", "application/json")
|
||||
recorder := httptest.NewRecorder()
|
||||
router.ServeHTTP(recorder, request)
|
||||
require.Equal(t, http.StatusAccepted, recorder.Code)
|
||||
require.Eventually(t, func() bool {
|
||||
store.mu.RLock()
|
||||
defer store.mu.RUnlock()
|
||||
for _, task := range store.tasks {
|
||||
if task.Status == service.ImageTaskStatusCompleted {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}, time.Second, 10*time.Millisecond)
|
||||
evaluated, _, _ := engine.snapshot()
|
||||
require.Equal(t, 1, evaluated)
|
||||
executionMu.Lock()
|
||||
require.False(t, repeatedDecision)
|
||||
executionMu.Unlock()
|
||||
}
|
||||
|
||||
func TestBatchImagePromptGuardRunsBeforePersistenceOrBilling(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
engine := blockingHandlerPromptEngine()
|
||||
openAI := &OpenAIGatewayHandler{securityAuditCoordinator: securityaudit.NewCoordinator(nil, engine)}
|
||||
h := &BatchImageHandler{openAI: openAI}
|
||||
router := gin.New()
|
||||
router.Use(securityAuditMediaTestMiddleware)
|
||||
router.POST("/v1/images/batches", h.Submit)
|
||||
body := map[string]any{
|
||||
"model": "gemini-image-test",
|
||||
"items": []map[string]any{{
|
||||
"custom_id": "one", "prompt": "blocked batch prompt",
|
||||
"reference_images": []map[string]any{{"mime_type": "image/png", "data": []byte("BINARY_CANARY")}},
|
||||
}},
|
||||
}
|
||||
raw, err := json.Marshal(body)
|
||||
require.NoError(t, err)
|
||||
request := httptest.NewRequest(http.MethodPost, "/v1/images/batches", strings.NewReader(string(raw)))
|
||||
request.Header.Set("Content-Type", "application/json")
|
||||
recorder := httptest.NewRecorder()
|
||||
require.NotPanics(t, func() { router.ServeHTTP(recorder, request) }, "nil service would panic if Submit were reached")
|
||||
require.Equal(t, http.StatusForbidden, recorder.Code)
|
||||
evaluated, _, requests := engine.snapshot()
|
||||
require.Equal(t, 1, evaluated)
|
||||
require.Len(t, requests, 1)
|
||||
require.Contains(t, string(requests[0].Body), "blocked batch prompt")
|
||||
require.NotContains(t, string(requests[0].Body), "BINARY_CANARY")
|
||||
require.NotContains(t, string(requests[0].Body), "QklOQVJZX0NBTkFSWQ==")
|
||||
}
|
||||
|
||||
func TestSecurityAuditBlockingFailuresLeaveAllDownstreamCountersAtZero(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
for _, kind := range []securityaudit.DecisionKind{securityaudit.DecisionBlock, securityaudit.DecisionUnavailable, securityaudit.DecisionInvalid} {
|
||||
t.Run(string(kind), func(t *testing.T) {
|
||||
promptDecision := promptGuardDecision(kind)
|
||||
engine := &handlerPromptEngine{mode: securityaudit.ModeBlocking, decision: &securityaudit.PromptDecision{
|
||||
Kind: kind, ErrorCode: promptDecision.ErrorCode, AllowNextStage: false,
|
||||
}}
|
||||
coordinator := securityaudit.NewCoordinator(nil, engine)
|
||||
recorder := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(recorder)
|
||||
c.Request = httptest.NewRequest(http.MethodPost, "/v1/chat/completions", strings.NewReader(`{"model":"gpt-test","messages":[{"role":"user","content":"guard me"}]}`))
|
||||
groupID := int64(3)
|
||||
apiKey := &service.APIKey{ID: 9, UserID: 7, GroupID: &groupID, Group: &service.Group{ID: groupID, Platform: service.PlatformOpenAI}}
|
||||
subject := middleware2.AuthSubject{UserID: 7, Concurrency: 2}
|
||||
decision := runSecurityAudit(c, nil, coordinator, nil, apiKey, subject, service.ContentModerationProtocolOpenAIChat, "gpt-test", []byte(`{"messages":[{"role":"user","content":"guard me"}]}`), "http")
|
||||
require.NotNil(t, decision)
|
||||
require.False(t, decision.AllowNextStage)
|
||||
require.False(t, recorder.Result().Header.Get("Content-Type") != "", "Guard evaluation itself must not start SSE/HTTP output")
|
||||
|
||||
accountSelections, billingChecks, billingPreconsumes, upstreamDispatches := 0, 0, 0, 0
|
||||
if decision.AllowNextStage {
|
||||
accountSelections++
|
||||
billingChecks++
|
||||
billingPreconsumes++
|
||||
upstreamDispatches++
|
||||
}
|
||||
require.Zero(t, accountSelections)
|
||||
require.Zero(t, billingChecks)
|
||||
require.Zero(t, billingPreconsumes)
|
||||
require.Zero(t, upstreamDispatches)
|
||||
(&OpenAIGatewayHandler{}).openAISecurityAuditError(c, decision)
|
||||
require.Equal(t, promptDecision.HTTPStatus, recorder.Code)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,85 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"go/ast"
|
||||
"go/parser"
|
||||
"go/token"
|
||||
"os"
|
||||
"regexp"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
type promptAuditOrderCase struct {
|
||||
file string
|
||||
function string
|
||||
auditToken string
|
||||
}
|
||||
|
||||
func TestPromptAuditGatePrecedesAccountBillingAndUpstreamSideEffects(t *testing.T) {
|
||||
tests := []promptAuditOrderCase{
|
||||
{file: "gateway_handler.go", function: "Messages", auditToken: "checkSecurityAudit"},
|
||||
{file: "gateway_handler_chat_completions.go", function: "ChatCompletions", auditToken: "checkSecurityAudit"},
|
||||
{file: "gateway_handler_responses.go", function: "Responses", auditToken: "checkSecurityAudit"},
|
||||
{file: "gemini_v1beta_handler.go", function: "GeminiV1BetaModels", auditToken: "checkSecurityAudit"},
|
||||
{file: "openai_gateway_handler.go", function: "Responses", auditToken: "checkSecurityAudit"},
|
||||
{file: "openai_gateway_handler.go", function: "Messages", auditToken: "checkSecurityAudit"},
|
||||
{file: "openai_chat_completions.go", function: "ChatCompletions", auditToken: "checkSecurityAudit"},
|
||||
{file: "openai_images.go", function: "Images", auditToken: "checkSecurityAudit"},
|
||||
{file: "grok_media.go", function: "handleGrokMedia", auditToken: "checkSecurityAudit"},
|
||||
{file: "openai_embeddings.go", function: "Embeddings", auditToken: "checkSecurityAudit"},
|
||||
{file: "openai_alpha_search.go", function: "AlphaSearch", auditToken: "checkSecurityAudit"},
|
||||
{file: "image_task_handler.go", function: "Submit", auditToken: "checkSecurityAuditBeforeSubmit"},
|
||||
{file: "batch_image_handler.go", function: "Submit", auditToken: "checkSecurityAuditBeforeSubmit"},
|
||||
}
|
||||
sideEffectTokens := []string{
|
||||
"CheckBillingEligibility(", "SelectAccount", ".Forward", "acquireResponsesUserSlot(",
|
||||
"AcquireUserSlot", "TryAcquireUserSlot", "acquireImageGenerationSlot(",
|
||||
"h.tasks.Create(", "h.service.Submit(",
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.file+"/"+tt.function, func(t *testing.T) {
|
||||
functionSource := stripGoComments(goFunctionSource(t, tt.file, tt.function))
|
||||
auditIndex := strings.Index(functionSource, tt.auditToken)
|
||||
require.NotEqual(t, -1, auditIndex, "missing Prompt Audit gate")
|
||||
foundSideEffect := false
|
||||
for _, sideEffect := range sideEffectTokens {
|
||||
index := strings.Index(functionSource, sideEffect)
|
||||
if index < 0 {
|
||||
continue
|
||||
}
|
||||
foundSideEffect = true
|
||||
require.Lessf(t, auditIndex, index, "%s must run before %s", tt.auditToken, sideEffect)
|
||||
}
|
||||
require.True(t, foundSideEffect, "coverage case must contain a downstream side effect")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func stripGoComments(source string) string {
|
||||
source = regexp.MustCompile(`(?s)/\*.*?\*/`).ReplaceAllString(source, "")
|
||||
return regexp.MustCompile(`(?m)//.*$`).ReplaceAllString(source, "")
|
||||
}
|
||||
|
||||
func goFunctionSource(t *testing.T, filename, functionName string) string {
|
||||
t.Helper()
|
||||
raw, err := os.ReadFile(filename)
|
||||
require.NoError(t, err)
|
||||
files := token.NewFileSet()
|
||||
parsed, err := parser.ParseFile(files, filename, raw, 0)
|
||||
require.NoError(t, err)
|
||||
for _, declaration := range parsed.Decls {
|
||||
function, ok := declaration.(*ast.FuncDecl)
|
||||
if !ok || function.Name.Name != functionName || function.Body == nil {
|
||||
continue
|
||||
}
|
||||
start := files.Position(function.Pos()).Offset
|
||||
end := files.Position(function.End()).Offset
|
||||
require.Greater(t, end, start)
|
||||
return string(raw[start:end])
|
||||
}
|
||||
t.Fatalf("function %s not found in %s", functionName, filename)
|
||||
return ""
|
||||
}
|
||||
@@ -1,7 +1,9 @@
|
||||
package handler
|
||||
|
||||
import (
|
||||
"github.com/Wei-Shaw/sub2api/internal/config"
|
||||
"github.com/Wei-Shaw/sub2api/internal/handler/admin"
|
||||
"github.com/Wei-Shaw/sub2api/internal/securityaudit"
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
|
||||
"github.com/google/wire"
|
||||
@@ -38,6 +40,7 @@ func ProvideAdminHandlers(
|
||||
channelMonitorHandler *admin.ChannelMonitorHandler,
|
||||
channelMonitorTemplateHandler *admin.ChannelMonitorRequestTemplateHandler,
|
||||
contentModerationHandler *admin.ContentModerationHandler,
|
||||
promptAuditHandler *securityaudit.PromptAdminHandler,
|
||||
paymentHandler *admin.PaymentHandler,
|
||||
affiliateHandler *admin.AffiliateHandler,
|
||||
complianceHandler *admin.ComplianceHandler,
|
||||
@@ -75,6 +78,7 @@ func ProvideAdminHandlers(
|
||||
ChannelMonitor: channelMonitorHandler,
|
||||
ChannelMonitorTemplate: channelMonitorTemplateHandler,
|
||||
ContentModeration: contentModerationHandler,
|
||||
PromptAudit: promptAuditHandler,
|
||||
Payment: paymentHandler,
|
||||
Affiliate: affiliateHandler,
|
||||
Compliance: complianceHandler,
|
||||
@@ -82,6 +86,60 @@ func ProvideAdminHandlers(
|
||||
}
|
||||
}
|
||||
|
||||
func ProvideGatewayHandler(
|
||||
gatewayService *service.GatewayService,
|
||||
openAIGatewayService *service.OpenAIGatewayService,
|
||||
geminiCompatService *service.GeminiMessagesCompatService,
|
||||
antigravityGatewayService *service.AntigravityGatewayService,
|
||||
userService *service.UserService,
|
||||
concurrencyService *service.ConcurrencyService,
|
||||
billingCacheService *service.BillingCacheService,
|
||||
usageService *service.UsageService,
|
||||
apiKeyService *service.APIKeyService,
|
||||
usageRecordWorkerPool *service.UsageRecordWorkerPool,
|
||||
errorPassthroughService *service.ErrorPassthroughService,
|
||||
contentModerationService *service.ContentModerationService,
|
||||
userMsgQueueService *service.UserMessageQueueService,
|
||||
cfg *config.Config,
|
||||
settingService *service.SettingService,
|
||||
coordinator *securityaudit.Coordinator,
|
||||
) *GatewayHandler {
|
||||
h := NewGatewayHandler(gatewayService, openAIGatewayService, geminiCompatService, antigravityGatewayService,
|
||||
userService, concurrencyService, billingCacheService, usageService, apiKeyService, usageRecordWorkerPool,
|
||||
errorPassthroughService, contentModerationService, userMsgQueueService, cfg, settingService)
|
||||
h.securityAuditCoordinator = coordinator
|
||||
return h
|
||||
}
|
||||
|
||||
func ProvideOpenAIGatewayHandler(
|
||||
gatewayService *service.OpenAIGatewayService,
|
||||
concurrencyService *service.ConcurrencyService,
|
||||
billingCacheService *service.BillingCacheService,
|
||||
apiKeyService *service.APIKeyService,
|
||||
usageRecordWorkerPool *service.UsageRecordWorkerPool,
|
||||
errorPassthroughService *service.ErrorPassthroughService,
|
||||
contentModerationService *service.ContentModerationService,
|
||||
opsService *service.OpsService,
|
||||
cfg *config.Config,
|
||||
coordinator *securityaudit.Coordinator,
|
||||
) *OpenAIGatewayHandler {
|
||||
h := NewOpenAIGatewayHandler(gatewayService, concurrencyService, billingCacheService, apiKeyService,
|
||||
usageRecordWorkerPool, errorPassthroughService, contentModerationService, opsService, cfg)
|
||||
h.securityAuditCoordinator = coordinator
|
||||
return h
|
||||
}
|
||||
|
||||
func ProvideBatchImageHandler(
|
||||
batchService *service.BatchImagePublicService,
|
||||
download *service.BatchImageDownloadService,
|
||||
cleanup *service.BatchImageCleanupService,
|
||||
openAI *OpenAIGatewayHandler,
|
||||
) *BatchImageHandler {
|
||||
h := NewBatchImageHandler(batchService, download, cleanup)
|
||||
h.openAI = openAI
|
||||
return h
|
||||
}
|
||||
|
||||
// ProvideSystemHandler creates admin.SystemHandler with UpdateService
|
||||
func ProvideSystemHandler(updateService *service.UpdateService, lockService *service.SystemOperationLockService) *admin.SystemHandler {
|
||||
return admin.NewSystemHandler(updateService, lockService)
|
||||
@@ -157,15 +215,15 @@ var ProviderSet = wire.NewSet(
|
||||
NewSubscriptionHandler,
|
||||
NewAnnouncementHandler,
|
||||
NewChannelMonitorUserHandler,
|
||||
NewGatewayHandler,
|
||||
NewOpenAIGatewayHandler,
|
||||
ProvideGatewayHandler,
|
||||
ProvideOpenAIGatewayHandler,
|
||||
NewTotpHandler,
|
||||
ProvideSettingHandler,
|
||||
NewPaymentHandler,
|
||||
NewPaymentWebhookHandler,
|
||||
NewAvailableChannelHandler,
|
||||
NewAsyncImageHandler,
|
||||
NewBatchImageHandler,
|
||||
ProvideBatchImageHandler,
|
||||
|
||||
// Admin handlers
|
||||
admin.NewDashboardHandler,
|
||||
|
||||
@@ -0,0 +1,139 @@
|
||||
package securityaudit
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net/http"
|
||||
"sync"
|
||||
)
|
||||
|
||||
type LegacyEngine interface {
|
||||
Check(ctx context.Context, req Request) (*LegacyDecision, error)
|
||||
}
|
||||
|
||||
type PromptEngine interface {
|
||||
EffectiveMode() Mode
|
||||
Enqueue(ctx context.Context, req Request) error
|
||||
Evaluate(ctx context.Context, req Request) (*PromptDecision, error)
|
||||
}
|
||||
|
||||
type Coordinator struct {
|
||||
legacy LegacyEngine
|
||||
prompt PromptEngine
|
||||
}
|
||||
|
||||
func NewCoordinator(legacy LegacyEngine, prompt PromptEngine) *Coordinator {
|
||||
return &Coordinator{legacy: legacy, prompt: prompt}
|
||||
}
|
||||
|
||||
func (c *Coordinator) Check(ctx context.Context, req Request) Decision {
|
||||
if c == nil {
|
||||
return allowDecision(nil, nil)
|
||||
}
|
||||
mode := ModeOff
|
||||
if c.prompt != nil {
|
||||
mode = c.prompt.EffectiveMode()
|
||||
}
|
||||
switch mode {
|
||||
case ModeAsync:
|
||||
// Enqueue is deliberately best-effort. The implementation owns a bounded
|
||||
// context and copies request memory before it can outlive the Handler.
|
||||
_ = c.prompt.Enqueue(ctx, req.Clone())
|
||||
legacy, _ := c.checkLegacy(ctx, req)
|
||||
return prioritize(legacy, nil)
|
||||
case ModeBlocking:
|
||||
return c.checkBlocking(ctx, req)
|
||||
default:
|
||||
legacy, _ := c.checkLegacy(ctx, req)
|
||||
return prioritize(legacy, nil)
|
||||
}
|
||||
}
|
||||
|
||||
func (c *Coordinator) checkBlocking(ctx context.Context, req Request) Decision {
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(2)
|
||||
var legacy *LegacyDecision
|
||||
var prompt *PromptDecision
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
legacy, _ = c.checkLegacy(ctx, req)
|
||||
}()
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
if c.prompt == nil {
|
||||
prompt = unavailablePromptDecision(ErrorCodeUnavailable)
|
||||
return
|
||||
}
|
||||
result, err := c.prompt.Evaluate(ctx, req.Clone())
|
||||
if err != nil {
|
||||
var guardErr *GuardError
|
||||
if errors.As(err, &guardErr) && guardErr.Code == ErrorCodeInvalidResponse {
|
||||
prompt = unavailablePromptDecision(ErrorCodeInvalidResponse)
|
||||
return
|
||||
}
|
||||
prompt = unavailablePromptDecision(ErrorCodeUnavailable)
|
||||
return
|
||||
}
|
||||
if result == nil {
|
||||
prompt = unavailablePromptDecision(ErrorCodeUnavailable)
|
||||
return
|
||||
}
|
||||
prompt = result
|
||||
}()
|
||||
wg.Wait()
|
||||
return prioritize(legacy, prompt)
|
||||
}
|
||||
|
||||
func (c *Coordinator) checkLegacy(ctx context.Context, req Request) (*LegacyDecision, error) {
|
||||
if c.legacy == nil {
|
||||
return nil, nil
|
||||
}
|
||||
return c.legacy.Check(ctx, req)
|
||||
}
|
||||
|
||||
func prioritize(legacy *LegacyDecision, prompt *PromptDecision) Decision {
|
||||
if legacy != nil && legacy.Blocked {
|
||||
status := legacy.StatusCode
|
||||
if status < 400 || status > 599 {
|
||||
status = http.StatusForbidden
|
||||
}
|
||||
code := legacy.ErrorCode
|
||||
if code == "" {
|
||||
code = "content_policy_violation"
|
||||
}
|
||||
return Decision{
|
||||
Kind: DecisionBlock, HTTPStatus: status, ErrorCode: code, ClientMessage: legacy.Message,
|
||||
Legacy: legacy, Prompt: prompt, AllowNextStage: false,
|
||||
}
|
||||
}
|
||||
if prompt == nil {
|
||||
return allowDecision(legacy, nil)
|
||||
}
|
||||
switch prompt.Kind {
|
||||
case DecisionBlock:
|
||||
return Decision{Kind: DecisionBlock, HTTPStatus: http.StatusForbidden, ErrorCode: ErrorCodeBlocked,
|
||||
ClientMessage: "提示词安全审计拒绝了该请求,请调整输入后重试", Legacy: legacy, Prompt: prompt}
|
||||
case DecisionInvalid:
|
||||
return Decision{Kind: DecisionInvalid, HTTPStatus: http.StatusServiceUnavailable, ErrorCode: ErrorCodeInvalidResponse,
|
||||
ClientMessage: "提示词安全审计暂时不可用,请稍后重试", Legacy: legacy, Prompt: prompt}
|
||||
case DecisionUnavailable:
|
||||
return Decision{Kind: DecisionUnavailable, HTTPStatus: http.StatusServiceUnavailable, ErrorCode: ErrorCodeUnavailable,
|
||||
ClientMessage: "提示词安全审计暂时不可用,请稍后重试", Legacy: legacy, Prompt: prompt}
|
||||
case DecisionFlag:
|
||||
return Decision{Kind: DecisionFlag, HTTPStatus: http.StatusOK, Legacy: legacy, Prompt: prompt, AllowNextStage: true}
|
||||
default:
|
||||
return allowDecision(legacy, prompt)
|
||||
}
|
||||
}
|
||||
|
||||
func allowDecision(legacy *LegacyDecision, prompt *PromptDecision) Decision {
|
||||
return Decision{Kind: DecisionAllow, HTTPStatus: http.StatusOK, Legacy: legacy, Prompt: prompt, AllowNextStage: true}
|
||||
}
|
||||
|
||||
func unavailablePromptDecision(code string) *PromptDecision {
|
||||
kind := DecisionUnavailable
|
||||
if code == ErrorCodeInvalidResponse {
|
||||
kind = DecisionInvalid
|
||||
}
|
||||
return &PromptDecision{Kind: kind, ErrorCode: code, AllowNextStage: false}
|
||||
}
|
||||
@@ -0,0 +1,35 @@
|
||||
package securityaudit
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
)
|
||||
|
||||
type LegacyModerationAdapter struct {
|
||||
service *service.ContentModerationService
|
||||
}
|
||||
|
||||
func NewLegacyModerationAdapter(svc *service.ContentModerationService) LegacyEngine {
|
||||
return &LegacyModerationAdapter{service: svc}
|
||||
}
|
||||
|
||||
func (a *LegacyModerationAdapter) Check(ctx context.Context, req Request) (*LegacyDecision, error) {
|
||||
if a == nil || a.service == nil {
|
||||
return nil, nil
|
||||
}
|
||||
decision, err := a.service.Check(ctx, service.ContentModerationCheckInput{
|
||||
RequestID: req.RequestID, UserID: req.UserID, UserEmail: req.UserEmail,
|
||||
APIKeyID: req.APIKeyID, APIKeyName: req.APIKeyName, GroupID: cloneInt64Ptr(req.GroupID),
|
||||
GroupName: req.GroupName, Endpoint: req.Endpoint, Provider: req.Provider,
|
||||
Model: req.Model, Protocol: req.Protocol, Body: req.Body,
|
||||
})
|
||||
if err != nil || decision == nil {
|
||||
return nil, err
|
||||
}
|
||||
return &LegacyDecision{
|
||||
Allowed: decision.Allowed, Blocked: decision.Blocked, Flagged: decision.Flagged,
|
||||
Message: decision.Message, StatusCode: decision.StatusCode,
|
||||
ErrorCode: "content_policy_violation", Action: decision.Action,
|
||||
}, nil
|
||||
}
|
||||
@@ -0,0 +1,176 @@
|
||||
package securityaudit
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
type fakeLegacyEngine struct {
|
||||
decision *LegacyDecision
|
||||
err error
|
||||
calls atomic.Int64
|
||||
}
|
||||
|
||||
func (f *fakeLegacyEngine) Check(context.Context, Request) (*LegacyDecision, error) {
|
||||
f.calls.Add(1)
|
||||
return f.decision, f.err
|
||||
}
|
||||
|
||||
type fakePromptEngine struct {
|
||||
mode Mode
|
||||
decision *PromptDecision
|
||||
err error
|
||||
enqueues atomic.Int64
|
||||
evaluates atomic.Int64
|
||||
}
|
||||
|
||||
func (f *fakePromptEngine) EffectiveMode() Mode { return f.mode }
|
||||
func (f *fakePromptEngine) Enqueue(context.Context, Request) error {
|
||||
f.enqueues.Add(1)
|
||||
return f.err
|
||||
}
|
||||
func (f *fakePromptEngine) Evaluate(context.Context, Request) (*PromptDecision, error) {
|
||||
f.evaluates.Add(1)
|
||||
return f.decision, f.err
|
||||
}
|
||||
|
||||
func TestCoordinatorModesAndPriority(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
mode Mode
|
||||
legacy *LegacyDecision
|
||||
prompt *PromptDecision
|
||||
promptErr error
|
||||
wantKind DecisionKind
|
||||
wantCode string
|
||||
wantEnqueue int64
|
||||
wantEvaluation int64
|
||||
}{
|
||||
{name: "off", mode: ModeOff, wantKind: DecisionAllow},
|
||||
{name: "async only enqueues", mode: ModeAsync, wantKind: DecisionAllow, wantEnqueue: 1},
|
||||
{name: "prompt block", mode: ModeBlocking, prompt: &PromptDecision{Kind: DecisionBlock}, wantKind: DecisionBlock, wantCode: ErrorCodeBlocked, wantEvaluation: 1},
|
||||
{name: "prompt unavailable", mode: ModeBlocking, promptErr: errors.New("down"), wantKind: DecisionUnavailable, wantCode: ErrorCodeUnavailable, wantEvaluation: 1},
|
||||
{name: "legacy wins both block", mode: ModeBlocking,
|
||||
legacy: &LegacyDecision{Blocked: true, StatusCode: http.StatusForbidden, ErrorCode: "content_policy_violation", Message: "legacy"},
|
||||
prompt: &PromptDecision{Kind: DecisionBlock}, wantKind: DecisionBlock, wantCode: "content_policy_violation", wantEvaluation: 1},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
legacy := &fakeLegacyEngine{decision: tt.legacy}
|
||||
prompt := &fakePromptEngine{mode: tt.mode, decision: tt.prompt, err: tt.promptErr}
|
||||
decision := NewCoordinator(legacy, prompt).Check(context.Background(), Request{Body: []byte(`{}`)})
|
||||
require.Equal(t, tt.wantKind, decision.Kind)
|
||||
require.Equal(t, tt.wantCode, decision.ErrorCode)
|
||||
require.Equal(t, int64(1), legacy.calls.Load())
|
||||
require.Equal(t, tt.wantEnqueue, prompt.enqueues.Load())
|
||||
require.Equal(t, tt.wantEvaluation, prompt.evaluates.Load())
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCoordinatorDoesNotMutateRequestBody(t *testing.T) {
|
||||
body := []byte(`{"messages":[{"role":"user","content":"hello"}]}`)
|
||||
original := append([]byte(nil), body...)
|
||||
prompt := &fakePromptEngine{mode: ModeAsync}
|
||||
decision := NewCoordinator(&fakeLegacyEngine{}, prompt).Check(context.Background(), Request{Body: body})
|
||||
require.True(t, decision.AllowNextStage)
|
||||
require.Equal(t, original, body)
|
||||
}
|
||||
|
||||
func TestCoordinatorBlockingPriorityCoversBothEngineDecisionMatrix(t *testing.T) {
|
||||
legacyCases := []struct {
|
||||
name string
|
||||
decision *LegacyDecision
|
||||
}{
|
||||
{name: "allow", decision: &LegacyDecision{Allowed: true, StatusCode: http.StatusOK, Action: "allow"}},
|
||||
{name: "flag", decision: &LegacyDecision{Allowed: true, Flagged: true, StatusCode: http.StatusOK, Action: "flag"}},
|
||||
{name: "block", decision: &LegacyDecision{Blocked: true, StatusCode: http.StatusForbidden, ErrorCode: "legacy_exact_code", Message: "legacy exact message", Action: "block"}},
|
||||
}
|
||||
promptCases := []struct {
|
||||
name string
|
||||
decision *PromptDecision
|
||||
wantKind DecisionKind
|
||||
wantCode string
|
||||
}{
|
||||
{name: "allow", decision: &PromptDecision{Kind: DecisionAllow, AllowNextStage: true}, wantKind: DecisionAllow},
|
||||
{name: "flag", decision: &PromptDecision{Kind: DecisionFlag, AllowNextStage: true}, wantKind: DecisionFlag},
|
||||
{name: "block", decision: &PromptDecision{Kind: DecisionBlock}, wantKind: DecisionBlock, wantCode: ErrorCodeBlocked},
|
||||
{name: "unavailable", decision: &PromptDecision{Kind: DecisionUnavailable, ErrorCode: ErrorCodeUnavailable}, wantKind: DecisionUnavailable, wantCode: ErrorCodeUnavailable},
|
||||
{name: "invalid", decision: &PromptDecision{Kind: DecisionInvalid, ErrorCode: ErrorCodeInvalidResponse}, wantKind: DecisionInvalid, wantCode: ErrorCodeInvalidResponse},
|
||||
}
|
||||
|
||||
for _, legacyCase := range legacyCases {
|
||||
for _, promptCase := range promptCases {
|
||||
t.Run(fmt.Sprintf("legacy_%s_prompt_%s", legacyCase.name, promptCase.name), func(t *testing.T) {
|
||||
legacy := &fakeLegacyEngine{decision: legacyCase.decision}
|
||||
prompt := &fakePromptEngine{mode: ModeBlocking, decision: promptCase.decision}
|
||||
decision := NewCoordinator(legacy, prompt).Check(context.Background(), Request{})
|
||||
|
||||
require.Same(t, legacyCase.decision, decision.Legacy)
|
||||
require.Same(t, promptCase.decision, decision.Prompt)
|
||||
require.Equal(t, int64(1), legacy.calls.Load())
|
||||
require.Equal(t, int64(1), prompt.evaluates.Load())
|
||||
if legacyCase.name == "block" {
|
||||
require.Equal(t, DecisionBlock, decision.Kind)
|
||||
require.Equal(t, "legacy_exact_code", decision.ErrorCode)
|
||||
require.Equal(t, "legacy exact message", decision.ClientMessage)
|
||||
require.False(t, decision.AllowNextStage)
|
||||
return
|
||||
}
|
||||
require.Equal(t, promptCase.wantKind, decision.Kind)
|
||||
require.Equal(t, promptCase.wantCode, decision.ErrorCode)
|
||||
require.Equal(t, promptCase.decision.AllowNextStage, decision.AllowNextStage)
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestCoordinatorPreservesIndependentEngineFactsAndMapsOnlyGatewayOutcome(t *testing.T) {
|
||||
legacyDecision := &LegacyDecision{
|
||||
Allowed: true, Flagged: true, Message: "legacy finding", StatusCode: http.StatusAccepted,
|
||||
ErrorCode: "legacy_observation", Action: "legacy_action",
|
||||
}
|
||||
promptResult := &NormalizedResult{
|
||||
Decision: EventCritical, RiskLevel: RiskCritical, Action: ActionBlock,
|
||||
Categories: []string{"pii"}, ScannerScores: map[string]float64{"pii": 1},
|
||||
}
|
||||
promptDecision := &PromptDecision{Kind: DecisionBlock, Result: promptResult}
|
||||
decision := NewCoordinator(
|
||||
&fakeLegacyEngine{decision: legacyDecision},
|
||||
&fakePromptEngine{mode: ModeBlocking, decision: promptDecision},
|
||||
).Check(context.Background(), Request{})
|
||||
|
||||
require.Same(t, legacyDecision, decision.Legacy)
|
||||
require.Same(t, promptDecision, decision.Prompt)
|
||||
require.Same(t, promptResult, decision.Prompt.Result)
|
||||
require.Equal(t, "legacy finding", decision.Legacy.Message)
|
||||
require.Equal(t, []string{"pii"}, decision.Prompt.Result.Categories)
|
||||
require.Equal(t, ErrorCodeBlocked, decision.ErrorCode)
|
||||
}
|
||||
|
||||
func TestCoordinatorAsyncEnqueueFailuresNeverChangeResponseOrDownstreamDispatch(t *testing.T) {
|
||||
for _, enqueueErr := range []error{ErrQueueFull, ErrQueueAdmissionBusy, errors.New("redis unavailable"), errors.New("publish failed")} {
|
||||
prompt := &fakePromptEngine{mode: ModeAsync, err: enqueueErr}
|
||||
decision := NewCoordinator(&fakeLegacyEngine{decision: &LegacyDecision{Allowed: true}}, prompt).Check(context.Background(), Request{})
|
||||
downstreamDispatches := 0
|
||||
status := http.StatusOK
|
||||
responseBody := "unchanged-upstream-response"
|
||||
if decision.AllowNextStage {
|
||||
downstreamDispatches++
|
||||
} else {
|
||||
status = decision.HTTPStatus
|
||||
responseBody = decision.ClientMessage
|
||||
}
|
||||
require.Equal(t, http.StatusOK, status)
|
||||
require.Equal(t, "unchanged-upstream-response", responseBody)
|
||||
require.Equal(t, 1, downstreamDispatches)
|
||||
require.Equal(t, int64(1), prompt.enqueues.Load())
|
||||
require.Zero(t, prompt.evaluates.Load())
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,468 @@
|
||||
package securityaudit
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"sort"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
|
||||
)
|
||||
|
||||
const (
|
||||
DefaultWorkerCount = 4
|
||||
MaxWorkerCount = 32
|
||||
DefaultQueueCapacity = 32768
|
||||
MaxQueueCapacity = 100000
|
||||
DefaultTimeoutMS = 3000
|
||||
MinTimeoutMS = 100
|
||||
MaxTimeoutMS = 30000
|
||||
DefaultInputLimit = 4000
|
||||
MinInputLimit = 128
|
||||
MaxInputLimit = 100000
|
||||
DefaultPayloadTTL = 30 * time.Minute
|
||||
)
|
||||
|
||||
type SecretEncryptor interface {
|
||||
Encrypt(plaintext string) (string, error)
|
||||
Decrypt(ciphertext string) (string, error)
|
||||
}
|
||||
|
||||
// ConfigStore is the injectable boundary between hot-path prompt auditing and
|
||||
// the concrete settings/PostgreSQL/Redis-backed configuration manager.
|
||||
type ConfigStore interface {
|
||||
Start(ctx context.Context) error
|
||||
Shutdown(ctx context.Context) error
|
||||
Active() (ActiveConfig, bool)
|
||||
EffectiveMode() Mode
|
||||
Public() PublicConfig
|
||||
Save(ctx context.Context, req UpdateConfigRequest, actorID int64) (PublicConfig, error)
|
||||
RuntimeState() (expected int64, active int64, loadedAt *time.Time, loadError string)
|
||||
Encrypt(value string) (string, error)
|
||||
Decrypt(value string) (string, error)
|
||||
}
|
||||
|
||||
type StorageEndpoint struct {
|
||||
ID string `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Protocol string `json:"protocol"`
|
||||
BaseURL string `json:"base_url"`
|
||||
Model string `json:"model"`
|
||||
TokenCiphertext string `json:"token_ciphertext,omitempty"`
|
||||
TimeoutMS int `json:"timeout_ms"`
|
||||
InputLimit int `json:"input_limit"`
|
||||
Enabled bool `json:"enabled"`
|
||||
}
|
||||
|
||||
type storageConfig struct {
|
||||
Enabled bool `json:"enabled"`
|
||||
BlockingEnabled bool `json:"blocking_enabled"`
|
||||
StorePassEvents bool `json:"store_pass_events"`
|
||||
Strategy string `json:"strategy"`
|
||||
WorkerCount int `json:"worker_count"`
|
||||
QueueCapacity int `json:"queue_capacity"`
|
||||
Scanners []string `json:"scanners"`
|
||||
AllGroups bool `json:"all_groups"`
|
||||
GroupIDs []int64 `json:"group_ids"`
|
||||
Endpoints []StorageEndpoint `json:"endpoints"`
|
||||
ConfigVersion int64 `json:"config_version"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
UpdatedBy int64 `json:"updated_by"`
|
||||
ChangeSummary string `json:"change_summary"`
|
||||
}
|
||||
|
||||
type ActiveEndpoint struct {
|
||||
ID string
|
||||
Name string
|
||||
Protocol string
|
||||
BaseURL string
|
||||
Model string
|
||||
Token string
|
||||
TimeoutMS int
|
||||
InputLimit int
|
||||
Enabled bool
|
||||
}
|
||||
|
||||
type ActiveConfig struct {
|
||||
RiskControlEnabled bool
|
||||
Enabled bool
|
||||
BlockingEnabled bool
|
||||
StorePassEvents bool
|
||||
Strategy string
|
||||
WorkerCount int
|
||||
QueueCapacity int
|
||||
Scanners []string
|
||||
AllGroups bool
|
||||
GroupIDs []int64
|
||||
Endpoints []ActiveEndpoint
|
||||
ConfigVersion int64
|
||||
UpdatedAt time.Time
|
||||
UpdatedBy int64
|
||||
ChangeSummary string
|
||||
}
|
||||
|
||||
type PublicEndpoint struct {
|
||||
ID string `json:"id"`
|
||||
Name string `json:"name"`
|
||||
Protocol string `json:"protocol"`
|
||||
BaseURL string `json:"base_url"`
|
||||
Model string `json:"model"`
|
||||
TimeoutMS int `json:"timeout_ms"`
|
||||
InputLimit int `json:"input_limit"`
|
||||
Enabled bool `json:"enabled"`
|
||||
HasToken bool `json:"has_token"`
|
||||
TokenStatus string `json:"token_status"`
|
||||
}
|
||||
|
||||
type PublicConfig struct {
|
||||
Enabled bool `json:"enabled"`
|
||||
BlockingEnabled bool `json:"blocking_enabled"`
|
||||
StorePassEvents bool `json:"store_pass_events"`
|
||||
EffectiveMode Mode `json:"effective_mode"`
|
||||
Strategy string `json:"strategy"`
|
||||
WorkerCount int `json:"worker_count"`
|
||||
QueueCapacity int `json:"queue_capacity"`
|
||||
Scanners []string `json:"scanners"`
|
||||
AllGroups bool `json:"all_groups"`
|
||||
GroupIDs []int64 `json:"group_ids"`
|
||||
Endpoints []PublicEndpoint `json:"endpoints"`
|
||||
ConfigVersion int64 `json:"config_version"`
|
||||
UpdatedAt time.Time `json:"updated_at"`
|
||||
UpdatedBy int64 `json:"updated_by"`
|
||||
ChangeSummary string `json:"change_summary"`
|
||||
}
|
||||
|
||||
type UpdateEndpoint struct {
|
||||
ID string `json:"id" binding:"required"`
|
||||
Name string `json:"name" binding:"required"`
|
||||
Protocol string `json:"protocol"`
|
||||
BaseURL string `json:"base_url" binding:"required"`
|
||||
Model string `json:"model"`
|
||||
Token string `json:"token,omitempty"`
|
||||
ClearToken bool `json:"clear_token"`
|
||||
TimeoutMS int `json:"timeout_ms"`
|
||||
InputLimit int `json:"input_limit"`
|
||||
Enabled bool `json:"enabled"`
|
||||
}
|
||||
|
||||
type UpdateConfigRequest struct {
|
||||
ExpectedConfigVersion int64 `json:"expected_config_version" binding:"required"`
|
||||
Enabled bool `json:"enabled"`
|
||||
BlockingEnabled bool `json:"blocking_enabled"`
|
||||
StorePassEvents bool `json:"store_pass_events"`
|
||||
Strategy string `json:"strategy"`
|
||||
WorkerCount int `json:"worker_count"`
|
||||
QueueCapacity int `json:"queue_capacity"`
|
||||
Scanners []string `json:"scanners"`
|
||||
AllGroups bool `json:"all_groups"`
|
||||
GroupIDs []int64 `json:"group_ids"`
|
||||
Endpoints []UpdateEndpoint `json:"endpoints"`
|
||||
}
|
||||
|
||||
func DefaultStorageConfig() storageConfig {
|
||||
return storageConfig{
|
||||
Enabled: false,
|
||||
BlockingEnabled: false,
|
||||
StorePassEvents: false,
|
||||
Strategy: "priority",
|
||||
WorkerCount: DefaultWorkerCount,
|
||||
QueueCapacity: DefaultQueueCapacity,
|
||||
Scanners: append([]string(nil), AllScannerIDs...),
|
||||
AllGroups: true,
|
||||
GroupIDs: []int64{},
|
||||
Endpoints: []StorageEndpoint{},
|
||||
ConfigVersion: 1,
|
||||
}
|
||||
}
|
||||
|
||||
func ParseStorageConfig(raw string) (storageConfig, error) {
|
||||
cfg := DefaultStorageConfig()
|
||||
if strings.TrimSpace(raw) == "" {
|
||||
return cfg, nil
|
||||
}
|
||||
if err := json.Unmarshal([]byte(raw), &cfg); err != nil {
|
||||
return storageConfig{}, fmt.Errorf("decode prompt audit config: %w", err)
|
||||
}
|
||||
normalizeStorageConfig(&cfg)
|
||||
if err := validateStorageConfig(cfg); err != nil {
|
||||
return storageConfig{}, err
|
||||
}
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
func normalizeStorageConfig(cfg *storageConfig) {
|
||||
if cfg == nil {
|
||||
return
|
||||
}
|
||||
if cfg.ConfigVersion < 1 {
|
||||
cfg.ConfigVersion = 1
|
||||
}
|
||||
if strings.TrimSpace(cfg.Strategy) == "" {
|
||||
cfg.Strategy = "priority"
|
||||
}
|
||||
if cfg.WorkerCount == 0 {
|
||||
cfg.WorkerCount = DefaultWorkerCount
|
||||
}
|
||||
if cfg.QueueCapacity == 0 {
|
||||
cfg.QueueCapacity = DefaultQueueCapacity
|
||||
}
|
||||
if len(cfg.Scanners) == 0 {
|
||||
cfg.Scanners = append([]string(nil), AllScannerIDs...)
|
||||
}
|
||||
cfg.Scanners = canonicalScannerIDs(cfg.Scanners)
|
||||
cfg.GroupIDs = canonicalInt64s(cfg.GroupIDs)
|
||||
// Preserve an invalid blocking-without-audit combination so validation can
|
||||
// reject it instead of silently changing administrator intent.
|
||||
for i := range cfg.Endpoints {
|
||||
ep := &cfg.Endpoints[i]
|
||||
ep.ID = strings.TrimSpace(ep.ID)
|
||||
ep.Name = strings.TrimSpace(ep.Name)
|
||||
ep.Protocol = strings.TrimSpace(ep.Protocol)
|
||||
if ep.Protocol == "" {
|
||||
ep.Protocol = "openai_compatible"
|
||||
}
|
||||
ep.BaseURL = strings.TrimSpace(ep.BaseURL)
|
||||
ep.Model = strings.TrimSpace(ep.Model)
|
||||
if ep.Model == "" {
|
||||
ep.Model = DefaultGuardModel
|
||||
}
|
||||
if ep.TimeoutMS == 0 {
|
||||
ep.TimeoutMS = DefaultTimeoutMS
|
||||
}
|
||||
if ep.InputLimit == 0 {
|
||||
ep.InputLimit = DefaultInputLimit
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func validateStorageConfig(cfg storageConfig) error {
|
||||
if cfg.BlockingEnabled && !cfg.Enabled {
|
||||
return infraerrors.BadRequest(ErrorCodeRequiresEnabled, "开启同步阻止前必须先启用提示词审计")
|
||||
}
|
||||
if cfg.Strategy != "priority" {
|
||||
return infraerrors.BadRequest("prompt_audit_invalid_strategy", "提示词审计策略仅支持 priority")
|
||||
}
|
||||
if cfg.WorkerCount < 1 || cfg.WorkerCount > MaxWorkerCount {
|
||||
return infraerrors.BadRequest("prompt_audit_invalid_worker_count", "Worker 数量超出允许范围")
|
||||
}
|
||||
if cfg.QueueCapacity < 1 || cfg.QueueCapacity > MaxQueueCapacity {
|
||||
return infraerrors.BadRequest("prompt_audit_invalid_queue_capacity", "队列容量超出允许范围")
|
||||
}
|
||||
if !cfg.AllGroups && len(cfg.GroupIDs) == 0 {
|
||||
return infraerrors.BadRequest("prompt_audit_groups_required", "指定分组模式至少需要选择一个分组")
|
||||
}
|
||||
if len(cfg.Scanners) == 0 {
|
||||
return infraerrors.BadRequest("prompt_audit_scanners_required", "至少需要启用一个风险分类")
|
||||
}
|
||||
seen := make(map[string]struct{}, len(cfg.Endpoints))
|
||||
enabled := 0
|
||||
for _, ep := range cfg.Endpoints {
|
||||
if ep.ID == "" || ep.Name == "" {
|
||||
return infraerrors.BadRequest("prompt_audit_invalid_endpoint", "审计节点 ID 和名称不能为空")
|
||||
}
|
||||
if _, ok := seen[ep.ID]; ok {
|
||||
return infraerrors.BadRequest("prompt_audit_duplicate_endpoint", "审计节点 ID 不能重复")
|
||||
}
|
||||
seen[ep.ID] = struct{}{}
|
||||
if ep.Protocol != "openai_compatible" {
|
||||
return infraerrors.BadRequest("prompt_audit_invalid_endpoint_protocol", "审计节点仅支持 OpenAI 兼容协议")
|
||||
}
|
||||
if _, err := NormalizeBaseURL(ep.BaseURL); err != nil {
|
||||
return err
|
||||
}
|
||||
if ep.TimeoutMS < MinTimeoutMS || ep.TimeoutMS > MaxTimeoutMS {
|
||||
return infraerrors.BadRequest("prompt_audit_invalid_timeout", "审计节点超时超出允许范围")
|
||||
}
|
||||
if ep.InputLimit < MinInputLimit || ep.InputLimit > MaxInputLimit {
|
||||
return infraerrors.BadRequest("prompt_audit_invalid_input_limit", "审计节点输入上限超出允许范围")
|
||||
}
|
||||
if ep.Enabled {
|
||||
enabled++
|
||||
}
|
||||
}
|
||||
if cfg.Enabled && enabled == 0 {
|
||||
return infraerrors.BadRequest("prompt_audit_endpoint_required", "启用提示词审计前至少需要启用一个审计节点")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateUpdateConfigRequest(req UpdateConfigRequest) error {
|
||||
if strings.TrimSpace(req.Strategy) != "priority" {
|
||||
return infraerrors.BadRequest("prompt_audit_invalid_strategy", "提示词审计策略仅支持 priority")
|
||||
}
|
||||
if req.WorkerCount < 1 || req.WorkerCount > MaxWorkerCount {
|
||||
return infraerrors.BadRequest("prompt_audit_invalid_worker_count", "Worker 数量超出允许范围")
|
||||
}
|
||||
if req.QueueCapacity < 1 || req.QueueCapacity > MaxQueueCapacity {
|
||||
return infraerrors.BadRequest("prompt_audit_invalid_queue_capacity", "队列容量超出允许范围")
|
||||
}
|
||||
if len(req.Scanners) == 0 {
|
||||
return infraerrors.BadRequest("prompt_audit_scanners_required", "至少需要启用一个风险分类")
|
||||
}
|
||||
for _, scanner := range req.Scanners {
|
||||
if _, ok := ScannerCatalog[NormalizeCategory(scanner)]; !ok {
|
||||
return infraerrors.BadRequest("prompt_audit_invalid_scanner", "提示词审计风险分类无效")
|
||||
}
|
||||
}
|
||||
if !req.AllGroups {
|
||||
if len(req.GroupIDs) == 0 {
|
||||
return infraerrors.BadRequest("prompt_audit_groups_required", "指定分组模式至少需要选择一个分组")
|
||||
}
|
||||
for _, groupID := range req.GroupIDs {
|
||||
if groupID <= 0 {
|
||||
return infraerrors.BadRequest("prompt_audit_invalid_group", "提示词审计分组 ID 无效")
|
||||
}
|
||||
}
|
||||
}
|
||||
for _, endpoint := range req.Endpoints {
|
||||
if endpoint.TimeoutMS < MinTimeoutMS || endpoint.TimeoutMS > MaxTimeoutMS {
|
||||
return infraerrors.BadRequest("prompt_audit_invalid_timeout", "审计节点超时超出允许范围")
|
||||
}
|
||||
if endpoint.InputLimit < MinInputLimit || endpoint.InputLimit > MaxInputLimit {
|
||||
return infraerrors.BadRequest("prompt_audit_invalid_input_limit", "审计节点输入上限超出允许范围")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (cfg ActiveConfig) EffectiveMode() Mode {
|
||||
if !cfg.RiskControlEnabled || !cfg.Enabled {
|
||||
return ModeOff
|
||||
}
|
||||
if cfg.BlockingEnabled {
|
||||
return ModeBlocking
|
||||
}
|
||||
return ModeAsync
|
||||
}
|
||||
|
||||
func (cfg ActiveConfig) IncludesGroup(groupID *int64) bool {
|
||||
if cfg.AllGroups {
|
||||
return true
|
||||
}
|
||||
if groupID == nil {
|
||||
return false
|
||||
}
|
||||
i := sort.Search(len(cfg.GroupIDs), func(i int) bool { return cfg.GroupIDs[i] >= *groupID })
|
||||
return i < len(cfg.GroupIDs) && cfg.GroupIDs[i] == *groupID
|
||||
}
|
||||
|
||||
func (cfg ActiveConfig) EnabledEndpoints() []ActiveEndpoint {
|
||||
result := make([]ActiveEndpoint, 0, len(cfg.Endpoints))
|
||||
for _, ep := range cfg.Endpoints {
|
||||
if ep.Enabled {
|
||||
result = append(result, ep)
|
||||
}
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func PublicFromStorage(cfg storageConfig, riskControlEnabled bool) PublicConfig {
|
||||
scanners := append([]string{}, cfg.Scanners...)
|
||||
groupIDs := append([]int64{}, cfg.GroupIDs...)
|
||||
endpoints := make([]PublicEndpoint, 0, len(cfg.Endpoints))
|
||||
for _, ep := range cfg.Endpoints {
|
||||
hasToken := strings.TrimSpace(ep.TokenCiphertext) != ""
|
||||
status := "missing"
|
||||
if hasToken {
|
||||
status = "configured"
|
||||
}
|
||||
endpoints = append(endpoints, PublicEndpoint{
|
||||
ID: ep.ID, Name: ep.Name, Protocol: ep.Protocol, BaseURL: ep.BaseURL,
|
||||
Model: ep.Model, TimeoutMS: ep.TimeoutMS, InputLimit: ep.InputLimit,
|
||||
Enabled: ep.Enabled, HasToken: hasToken, TokenStatus: status,
|
||||
})
|
||||
}
|
||||
active := ActiveConfig{RiskControlEnabled: riskControlEnabled, Enabled: cfg.Enabled, BlockingEnabled: cfg.BlockingEnabled}
|
||||
return PublicConfig{
|
||||
Enabled: cfg.Enabled, BlockingEnabled: cfg.BlockingEnabled, StorePassEvents: cfg.StorePassEvents,
|
||||
EffectiveMode: active.EffectiveMode(), Strategy: cfg.Strategy, WorkerCount: cfg.WorkerCount,
|
||||
QueueCapacity: cfg.QueueCapacity, Scanners: scanners, AllGroups: cfg.AllGroups,
|
||||
GroupIDs: groupIDs, Endpoints: endpoints, ConfigVersion: cfg.ConfigVersion,
|
||||
UpdatedAt: cfg.UpdatedAt, UpdatedBy: cfg.UpdatedBy, ChangeSummary: cfg.ChangeSummary,
|
||||
}
|
||||
}
|
||||
|
||||
func ActiveFromStorage(cfg storageConfig, riskControlEnabled bool, encryptor SecretEncryptor) (ActiveConfig, error) {
|
||||
active := ActiveConfig{
|
||||
RiskControlEnabled: riskControlEnabled, Enabled: cfg.Enabled, BlockingEnabled: cfg.BlockingEnabled,
|
||||
StorePassEvents: cfg.StorePassEvents, Strategy: cfg.Strategy, WorkerCount: cfg.WorkerCount,
|
||||
QueueCapacity: cfg.QueueCapacity, Scanners: append([]string(nil), cfg.Scanners...), AllGroups: cfg.AllGroups,
|
||||
GroupIDs: append([]int64(nil), cfg.GroupIDs...), ConfigVersion: cfg.ConfigVersion,
|
||||
UpdatedAt: cfg.UpdatedAt, UpdatedBy: cfg.UpdatedBy, ChangeSummary: cfg.ChangeSummary,
|
||||
Endpoints: make([]ActiveEndpoint, 0, len(cfg.Endpoints)),
|
||||
}
|
||||
for _, ep := range cfg.Endpoints {
|
||||
token := ""
|
||||
if ep.TokenCiphertext != "" {
|
||||
if encryptor == nil {
|
||||
return ActiveConfig{}, fmt.Errorf("prompt audit secret encryptor unavailable")
|
||||
}
|
||||
plain, err := encryptor.Decrypt(ep.TokenCiphertext)
|
||||
if err != nil {
|
||||
return ActiveConfig{}, fmt.Errorf("decrypt prompt audit endpoint token %q: %w", ep.ID, err)
|
||||
}
|
||||
token = plain
|
||||
}
|
||||
active.Endpoints = append(active.Endpoints, ActiveEndpoint{
|
||||
ID: ep.ID, Name: ep.Name, Protocol: ep.Protocol, BaseURL: ep.BaseURL, Model: ep.Model,
|
||||
Token: token, TimeoutMS: ep.TimeoutMS, InputLimit: ep.InputLimit, Enabled: ep.Enabled,
|
||||
})
|
||||
}
|
||||
return active, nil
|
||||
}
|
||||
|
||||
func changeSummary(cfg storageConfig) string {
|
||||
summary := struct {
|
||||
Enabled bool `json:"enabled"`
|
||||
BlockingEnabled bool `json:"blocking_enabled"`
|
||||
StorePassEvents bool `json:"store_pass_events"`
|
||||
EndpointCount int `json:"endpoint_count"`
|
||||
ScannerCount int `json:"scanner_count"`
|
||||
AllGroups bool `json:"all_groups"`
|
||||
GroupCount int `json:"group_count"`
|
||||
GroupHash string `json:"group_hash"`
|
||||
}{cfg.Enabled, cfg.BlockingEnabled, cfg.StorePassEvents, len(cfg.Endpoints), len(cfg.Scanners), cfg.AllGroups, len(cfg.GroupIDs), ""}
|
||||
rawGroups, _ := json.Marshal(cfg.GroupIDs)
|
||||
digest := sha256.Sum256(rawGroups)
|
||||
summary.GroupHash = hex.EncodeToString(digest[:])
|
||||
raw, _ := json.Marshal(summary)
|
||||
return string(raw)
|
||||
}
|
||||
|
||||
func canonicalInt64s(values []int64) []int64 {
|
||||
seen := make(map[int64]struct{}, len(values))
|
||||
result := make([]int64, 0, len(values))
|
||||
for _, value := range values {
|
||||
if value <= 0 {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[value]; ok {
|
||||
continue
|
||||
}
|
||||
seen[value] = struct{}{}
|
||||
result = append(result, value)
|
||||
}
|
||||
sort.Slice(result, func(i, j int) bool { return result[i] < result[j] })
|
||||
return result
|
||||
}
|
||||
|
||||
func canonicalScannerIDs(values []string) []string {
|
||||
seen := make(map[string]struct{}, len(values))
|
||||
for _, value := range values {
|
||||
id := NormalizeCategory(value)
|
||||
if _, ok := ScannerCatalog[id]; ok {
|
||||
seen[id] = struct{}{}
|
||||
}
|
||||
}
|
||||
result := make([]string, 0, len(seen))
|
||||
for _, id := range AllScannerIDs {
|
||||
if _, ok := seen[id]; ok {
|
||||
result = append(result, id)
|
||||
}
|
||||
}
|
||||
return result
|
||||
}
|
||||
@@ -0,0 +1,244 @@
|
||||
package securityaudit
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"os"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/config"
|
||||
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
|
||||
"github.com/Wei-Shaw/sub2api/internal/repository"
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
"github.com/lib/pq"
|
||||
"github.com/redis/go-redis/v9"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
const promptAuditRedisTestEnv = "PROMPT_AUDIT_TEST_REDIS_ADDR"
|
||||
|
||||
type postgresPromptAuditSettingRepository struct{ db *sql.DB }
|
||||
|
||||
func (r postgresPromptAuditSettingRepository) Get(ctx context.Context, key string) (*service.Setting, error) {
|
||||
var value string
|
||||
var updated time.Time
|
||||
err := r.db.QueryRowContext(ctx, `SELECT value,updated_at FROM settings WHERE key=$1`, key).Scan(&value, &updated)
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return nil, service.ErrSettingNotFound
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &service.Setting{Key: key, Value: value, UpdatedAt: updated}, nil
|
||||
}
|
||||
|
||||
func (r postgresPromptAuditSettingRepository) GetValue(ctx context.Context, key string) (string, error) {
|
||||
setting, err := r.Get(ctx, key)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return setting.Value, nil
|
||||
}
|
||||
|
||||
func (r postgresPromptAuditSettingRepository) Set(ctx context.Context, key, value string) error {
|
||||
_, err := r.db.ExecContext(ctx, `INSERT INTO settings(key,value,updated_at) VALUES($1,$2,NOW())
|
||||
ON CONFLICT(key) DO UPDATE SET value=EXCLUDED.value,updated_at=EXCLUDED.updated_at`, key, value)
|
||||
return err
|
||||
}
|
||||
|
||||
func (r postgresPromptAuditSettingRepository) GetMultiple(ctx context.Context, keys []string) (map[string]string, error) {
|
||||
result := make(map[string]string, len(keys))
|
||||
for _, key := range keys {
|
||||
result[key] = ""
|
||||
}
|
||||
rows, err := r.db.QueryContext(ctx, `SELECT key,value FROM settings WHERE key=ANY($1)`, pq.Array(keys))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
for rows.Next() {
|
||||
var key, value string
|
||||
if err := rows.Scan(&key, &value); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result[key] = value
|
||||
}
|
||||
return result, rows.Err()
|
||||
}
|
||||
|
||||
func (r postgresPromptAuditSettingRepository) SetMultiple(ctx context.Context, values map[string]string) error {
|
||||
tx, err := r.db.BeginTx(ctx, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer func() { _ = tx.Rollback() }()
|
||||
for key, value := range values {
|
||||
if _, err := tx.ExecContext(ctx, `INSERT INTO settings(key,value,updated_at) VALUES($1,$2,NOW())
|
||||
ON CONFLICT(key) DO UPDATE SET value=EXCLUDED.value,updated_at=EXCLUDED.updated_at`, key, value); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return tx.Commit()
|
||||
}
|
||||
|
||||
func (r postgresPromptAuditSettingRepository) GetAll(ctx context.Context) (map[string]string, error) {
|
||||
rows, err := r.db.QueryContext(ctx, `SELECT key,value FROM settings`)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
result := map[string]string{}
|
||||
for rows.Next() {
|
||||
var key, value string
|
||||
if err := rows.Scan(&key, &value); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result[key] = value
|
||||
}
|
||||
return result, rows.Err()
|
||||
}
|
||||
|
||||
func (r postgresPromptAuditSettingRepository) Delete(ctx context.Context, key string) error {
|
||||
_, err := r.db.ExecContext(ctx, `DELETE FROM settings WHERE key=$1`, key)
|
||||
return err
|
||||
}
|
||||
|
||||
func promptAuditTestEncryptor(t *testing.T) service.SecretEncryptor {
|
||||
t.Helper()
|
||||
encryptor, err := repository.NewAESEncryptor(&config.Config{Totp: config.TotpConfig{EncryptionKey: strings.Repeat("42", 32)}})
|
||||
require.NoError(t, err)
|
||||
return encryptor
|
||||
}
|
||||
|
||||
func promptAuditUpdateRequest(version int64, workerCount int, token string) UpdateConfigRequest {
|
||||
return UpdateConfigRequest{
|
||||
ExpectedConfigVersion: version, Enabled: true, BlockingEnabled: false, StorePassEvents: false,
|
||||
Strategy: "priority", WorkerCount: workerCount, QueueCapacity: 64, Scanners: []string{"pii", "jailbreak"},
|
||||
AllGroups: true, Endpoints: []UpdateEndpoint{{
|
||||
ID: "guard-one", Name: "Guard One", Protocol: "openai_compatible",
|
||||
BaseURL: "http://127.0.0.1:18080", Model: "", Token: token,
|
||||
TimeoutMS: 1000, InputLimit: 1024, Enabled: true,
|
||||
}},
|
||||
}
|
||||
}
|
||||
|
||||
func waitForConfigVersion(t *testing.T, manager *ConfigManager, version int64, timeout time.Duration) {
|
||||
t.Helper()
|
||||
require.Eventually(t, func() bool {
|
||||
active, ok := manager.Active()
|
||||
return ok && active.ConfigVersion == version
|
||||
}, timeout, 20*time.Millisecond)
|
||||
}
|
||||
|
||||
func TestPromptAuditConfigCASSecretRoundTripInvalidationAndTTL(t *testing.T) {
|
||||
redisAddress := strings.TrimSpace(os.Getenv(promptAuditRedisTestEnv))
|
||||
if redisAddress == "" {
|
||||
t.Skip(promptAuditRedisTestEnv + " is not set")
|
||||
}
|
||||
db := openPromptAuditIntegrationDB(t)
|
||||
settingRepo := postgresPromptAuditSettingRepository{db: db}
|
||||
require.NoError(t, settingRepo.Set(context.Background(), SettingKeyRiskControl, "true"))
|
||||
encryptor := promptAuditTestEncryptor(t)
|
||||
redisClient := redis.NewClient(&redis.Options{Addr: redisAddress})
|
||||
t.Cleanup(func() { require.NoError(t, redisClient.Close()) })
|
||||
require.NoError(t, redisClient.Ping(context.Background()).Err())
|
||||
|
||||
managerOne := NewConfigManager(db, settingRepo, redisClient, encryptor)
|
||||
managerTwo := NewConfigManager(db, settingRepo, redisClient, encryptor)
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
t.Cleanup(cancel)
|
||||
require.NoError(t, managerOne.Start(ctx))
|
||||
require.NoError(t, managerTwo.Start(ctx))
|
||||
t.Cleanup(func() {
|
||||
require.NoError(t, managerOne.Shutdown(context.Background()))
|
||||
require.NoError(t, managerTwo.Shutdown(context.Background()))
|
||||
})
|
||||
require.Eventually(t, func() bool {
|
||||
return redisClient.PubSubNumSub(context.Background(), ConfigInvalidationChannel).Val()[ConfigInvalidationChannel] >= 2
|
||||
}, 2*time.Second, 20*time.Millisecond)
|
||||
|
||||
const canary = "GUARD_TOKEN_CANARY_SECRET_4_CONFIG"
|
||||
public, err := managerOne.Save(context.Background(), promptAuditUpdateRequest(1, 1, canary), 101)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, int64(2), public.ConfigVersion)
|
||||
require.True(t, public.Endpoints[0].HasToken)
|
||||
publicJSON, err := json.Marshal(public)
|
||||
require.NoError(t, err)
|
||||
require.NotContains(t, string(publicJSON), canary)
|
||||
waitForConfigVersion(t, managerTwo, 2, 2*time.Second)
|
||||
|
||||
raw, err := settingRepo.GetValue(context.Background(), SettingKeyPromptAuditConfig)
|
||||
require.NoError(t, err)
|
||||
require.NotContains(t, raw, canary)
|
||||
stored, err := ParseStorageConfig(raw)
|
||||
require.NoError(t, err)
|
||||
require.NotEmpty(t, stored.Endpoints[0].TokenCiphertext)
|
||||
plain, err := encryptor.Decrypt(stored.Endpoints[0].TokenCiphertext)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, canary, plain)
|
||||
require.NotContains(t, stored.ChangeSummary, canary)
|
||||
require.NotContains(t, stored.ChangeSummary, stored.Endpoints[0].BaseURL)
|
||||
|
||||
type saveResult struct {
|
||||
config PublicConfig
|
||||
err error
|
||||
}
|
||||
start := make(chan struct{})
|
||||
results := make(chan saveResult, 2)
|
||||
var wg sync.WaitGroup
|
||||
for index, manager := range []*ConfigManager{managerOne, managerTwo} {
|
||||
wg.Add(1)
|
||||
go func(index int, manager *ConfigManager) {
|
||||
defer wg.Done()
|
||||
<-start
|
||||
cfg, saveErr := manager.Save(context.Background(), promptAuditUpdateRequest(2, index+2, ""), int64(201+index))
|
||||
results <- saveResult{config: cfg, err: saveErr}
|
||||
}(index, manager)
|
||||
}
|
||||
close(start)
|
||||
wg.Wait()
|
||||
close(results)
|
||||
succeeded, conflicted := 0, 0
|
||||
for result := range results {
|
||||
if result.err == nil {
|
||||
succeeded++
|
||||
require.Equal(t, int64(3), result.config.ConfigVersion)
|
||||
continue
|
||||
}
|
||||
conflicted++
|
||||
require.Equal(t, ErrorCodeConfigConflict, infraerrors.Reason(result.err))
|
||||
}
|
||||
require.Equal(t, 1, succeeded)
|
||||
require.Equal(t, 1, conflicted)
|
||||
waitForConfigVersion(t, managerOne, 3, 2*time.Second)
|
||||
waitForConfigVersion(t, managerTwo, 3, 2*time.Second)
|
||||
|
||||
// A manager without Redis subscriptions must still converge through the
|
||||
// bounded five-second refresh loop.
|
||||
ttlManager := NewConfigManager(db, settingRepo, nil, encryptor)
|
||||
require.NoError(t, ttlManager.Start(ctx))
|
||||
t.Cleanup(func() { require.NoError(t, ttlManager.Shutdown(context.Background())) })
|
||||
waitForConfigVersion(t, ttlManager, 3, time.Second)
|
||||
updated, err := managerOne.Save(context.Background(), promptAuditUpdateRequest(3, 5, ""), 301)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, int64(4), updated.ConfigVersion)
|
||||
waitForConfigVersion(t, ttlManager, 4, 7*time.Second)
|
||||
|
||||
// Redis publication failure is observable degradation, not a rollback of a
|
||||
// successfully committed PostgreSQL config.
|
||||
deadRedis := redis.NewClient(&redis.Options{Addr: "127.0.0.1:1", MaxRetries: 0, DialTimeout: 30 * time.Millisecond, ReadTimeout: 30 * time.Millisecond, WriteTimeout: 30 * time.Millisecond})
|
||||
t.Cleanup(func() { _ = deadRedis.Close() })
|
||||
degraded := NewConfigManager(db, settingRepo, deadRedis, encryptor)
|
||||
require.NoError(t, degraded.Reload(context.Background()))
|
||||
degradedSaved, err := degraded.Save(context.Background(), promptAuditUpdateRequest(4, 6, ""), 401)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, int64(5), degradedSaved.ConfigVersion)
|
||||
active, ok := degraded.Active()
|
||||
require.True(t, ok)
|
||||
require.Equal(t, int64(5), active.ConfigVersion)
|
||||
}
|
||||
@@ -0,0 +1,409 @@
|
||||
package securityaudit
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
|
||||
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
"github.com/redis/go-redis/v9"
|
||||
)
|
||||
|
||||
type activeConfigSnapshot struct {
|
||||
storage storageConfig
|
||||
active ActiveConfig
|
||||
loadedAt time.Time
|
||||
}
|
||||
|
||||
type ConfigManager struct {
|
||||
db *sql.DB
|
||||
settings service.SettingRepository
|
||||
redis *redis.Client
|
||||
encryptor SecretEncryptor
|
||||
clock Clock
|
||||
|
||||
snapshot atomic.Pointer[activeConfigSnapshot]
|
||||
expected atomic.Int64
|
||||
// expectedBlocking records the last storage intent that could be decoded,
|
||||
// independently of whether endpoint credentials or the full config could be
|
||||
// activated. A config version alone cannot distinguish async from blocking.
|
||||
expectedBlocking atomic.Bool
|
||||
|
||||
stateMu sync.RWMutex
|
||||
lastLoadError string
|
||||
lastErrorAt *time.Time
|
||||
|
||||
lifecycleMu sync.Mutex
|
||||
cancel context.CancelFunc
|
||||
wg sync.WaitGroup
|
||||
}
|
||||
|
||||
func NewConfigManager(db *sql.DB, settings service.SettingRepository, redisClient *redis.Client, encryptor service.SecretEncryptor) *ConfigManager {
|
||||
return &ConfigManager{db: db, settings: settings, redis: redisClient, encryptor: encryptor, clock: realClock{}}
|
||||
}
|
||||
|
||||
func (m *ConfigManager) Start(ctx context.Context) error {
|
||||
if m == nil {
|
||||
return errors.New("prompt audit config manager unavailable")
|
||||
}
|
||||
m.lifecycleMu.Lock()
|
||||
if m.cancel != nil {
|
||||
m.lifecycleMu.Unlock()
|
||||
return nil
|
||||
}
|
||||
runCtx, cancel := context.WithCancel(ctx)
|
||||
m.cancel = cancel
|
||||
m.lifecycleMu.Unlock()
|
||||
loadErr := m.Reload(runCtx)
|
||||
m.wg.Add(1)
|
||||
go m.refreshLoop(runCtx)
|
||||
if m.redis != nil {
|
||||
m.wg.Add(1)
|
||||
go m.subscribeLoop(runCtx)
|
||||
}
|
||||
return loadErr
|
||||
}
|
||||
|
||||
func (m *ConfigManager) Shutdown(_ context.Context) error {
|
||||
if m == nil {
|
||||
return nil
|
||||
}
|
||||
m.lifecycleMu.Lock()
|
||||
cancel := m.cancel
|
||||
m.cancel = nil
|
||||
m.lifecycleMu.Unlock()
|
||||
if cancel != nil {
|
||||
cancel()
|
||||
}
|
||||
m.wg.Wait()
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *ConfigManager) Reload(ctx context.Context) error {
|
||||
if m == nil || m.settings == nil {
|
||||
return errors.New("prompt audit setting repository unavailable")
|
||||
}
|
||||
values, err := m.settings.GetMultiple(ctx, []string{SettingKeyPromptAuditConfig, SettingKeyRiskControl})
|
||||
if err != nil {
|
||||
m.recordLoadError(err)
|
||||
return err
|
||||
}
|
||||
m.observeExpectedState(values[SettingKeyPromptAuditConfig], values[SettingKeyRiskControl] == "true")
|
||||
storage, err := ParseStorageConfig(values[SettingKeyPromptAuditConfig])
|
||||
if err != nil {
|
||||
m.recordLoadError(err)
|
||||
return err
|
||||
}
|
||||
m.expected.Store(storage.ConfigVersion)
|
||||
m.expectedBlocking.Store(values[SettingKeyRiskControl] == "true" && storage.Enabled && storage.BlockingEnabled)
|
||||
active, err := ActiveFromStorage(storage, values[SettingKeyRiskControl] == "true", m.encryptor)
|
||||
if err != nil {
|
||||
m.recordLoadError(err)
|
||||
return err
|
||||
}
|
||||
now := m.clock.Now()
|
||||
m.snapshot.Store(&activeConfigSnapshot{storage: cloneStorageConfig(storage), active: cloneActiveConfig(active), loadedAt: now})
|
||||
m.clearLoadError()
|
||||
LogInfo(EventConfigLoaded, map[string]any{
|
||||
"config_version": storage.ConfigVersion, "status": "loaded",
|
||||
})
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *ConfigManager) Active() (ActiveConfig, bool) {
|
||||
if m == nil {
|
||||
return ActiveConfig{}, false
|
||||
}
|
||||
snapshot := m.snapshot.Load()
|
||||
if snapshot == nil {
|
||||
return ActiveConfig{}, false
|
||||
}
|
||||
return cloneActiveConfig(snapshot.active), true
|
||||
}
|
||||
|
||||
func (m *ConfigManager) EffectiveMode() Mode {
|
||||
active, ok := m.Active()
|
||||
if !ok {
|
||||
// A cold start without a valid snapshot fails closed only when the last
|
||||
// decodable storage intent explicitly required blocking. Config version is
|
||||
// not a mode signal: an async-only config can have any version.
|
||||
if m != nil && m.expectedBlocking.Load() {
|
||||
return ModeBlocking
|
||||
}
|
||||
return ModeOff
|
||||
}
|
||||
return active.EffectiveMode()
|
||||
}
|
||||
|
||||
func (m *ConfigManager) Public() PublicConfig {
|
||||
if m == nil {
|
||||
return PublicFromStorage(DefaultStorageConfig(), false)
|
||||
}
|
||||
snapshot := m.snapshot.Load()
|
||||
if snapshot == nil {
|
||||
return PublicFromStorage(DefaultStorageConfig(), false)
|
||||
}
|
||||
return PublicFromStorage(cloneStorageConfig(snapshot.storage), snapshot.active.RiskControlEnabled)
|
||||
}
|
||||
|
||||
func (m *ConfigManager) Save(ctx context.Context, req UpdateConfigRequest, actorID int64) (PublicConfig, error) {
|
||||
if m == nil || m.db == nil || m.encryptor == nil {
|
||||
return PublicConfig{}, errors.New("prompt audit config persistence unavailable")
|
||||
}
|
||||
if req.ExpectedConfigVersion < 1 {
|
||||
return PublicConfig{}, infraerrors.BadRequest("prompt_audit_expected_config_version_required", "必须提供有效的配置版本")
|
||||
}
|
||||
tx, err := m.db.BeginTx(ctx, &sql.TxOptions{Isolation: sql.LevelReadCommitted})
|
||||
if err != nil {
|
||||
return PublicConfig{}, err
|
||||
}
|
||||
defer func() { _ = tx.Rollback() }()
|
||||
if _, err := tx.ExecContext(ctx, `SELECT pg_advisory_xact_lock($1)`, promptAuditConfigLockKey); err != nil {
|
||||
return PublicConfig{}, err
|
||||
}
|
||||
current := DefaultStorageConfig()
|
||||
var raw string
|
||||
err = tx.QueryRowContext(ctx, `SELECT value FROM settings WHERE key=$1 FOR UPDATE`, SettingKeyPromptAuditConfig).Scan(&raw)
|
||||
if err != nil && !errors.Is(err, sql.ErrNoRows) {
|
||||
return PublicConfig{}, err
|
||||
}
|
||||
if err == nil {
|
||||
current, err = ParseStorageConfig(raw)
|
||||
if err != nil {
|
||||
return PublicConfig{}, err
|
||||
}
|
||||
}
|
||||
if current.ConfigVersion != req.ExpectedConfigVersion {
|
||||
return PublicConfig{}, infraerrors.Conflict(ErrorCodeConfigConflict, "提示词审计配置已被其他管理员更新")
|
||||
}
|
||||
next, err := m.buildNextStorage(current, req, actorID)
|
||||
if err != nil {
|
||||
return PublicConfig{}, err
|
||||
}
|
||||
next.ConfigVersion = current.ConfigVersion + 1
|
||||
next.UpdatedAt = m.clock.Now()
|
||||
next.UpdatedBy = actorID
|
||||
next.ChangeSummary = changeSummary(next)
|
||||
rawNext, err := json.Marshal(next)
|
||||
if err != nil {
|
||||
return PublicConfig{}, err
|
||||
}
|
||||
if _, err := tx.ExecContext(ctx, `
|
||||
INSERT INTO settings (key,value,updated_at) VALUES ($1,$2,NOW())
|
||||
ON CONFLICT (key) DO UPDATE SET value=EXCLUDED.value, updated_at=EXCLUDED.updated_at`,
|
||||
SettingKeyPromptAuditConfig, string(rawNext)); err != nil {
|
||||
return PublicConfig{}, err
|
||||
}
|
||||
if err := tx.Commit(); err != nil {
|
||||
return PublicConfig{}, err
|
||||
}
|
||||
// Install the snapshot with the current global gate, not merely the value
|
||||
// cached when this process last reloaded Prompt Audit configuration.
|
||||
riskControlEnabled := m.currentRiskControlEnabled()
|
||||
if values, getErr := m.settings.GetMultiple(ctx, []string{SettingKeyRiskControl}); getErr == nil {
|
||||
riskControlEnabled = values[SettingKeyRiskControl] == "true"
|
||||
}
|
||||
active, err := ActiveFromStorage(next, riskControlEnabled, m.encryptor)
|
||||
if err != nil {
|
||||
return PublicConfig{}, err
|
||||
}
|
||||
m.expected.Store(next.ConfigVersion)
|
||||
m.expectedBlocking.Store(active.RiskControlEnabled && next.Enabled && next.BlockingEnabled)
|
||||
m.snapshot.Store(&activeConfigSnapshot{storage: cloneStorageConfig(next), active: cloneActiveConfig(active), loadedAt: m.clock.Now()})
|
||||
m.clearLoadError()
|
||||
LogInfo(EventConfigUpdated, map[string]any{
|
||||
"config_version": next.ConfigVersion, "status": "updated",
|
||||
})
|
||||
if m.redis != nil {
|
||||
if err := m.redis.Publish(ctx, ConfigInvalidationChannel, strconv.FormatInt(next.ConfigVersion, 10)).Err(); err != nil {
|
||||
LogWarn(EventConfigReloadDegraded, map[string]any{
|
||||
"config_version": next.ConfigVersion, "status": "degraded", "error_code": "config_invalidation_publish_failed",
|
||||
})
|
||||
}
|
||||
}
|
||||
return PublicFromStorage(next, active.RiskControlEnabled), nil
|
||||
}
|
||||
|
||||
func (m *ConfigManager) buildNextStorage(current storageConfig, req UpdateConfigRequest, actorID int64) (storageConfig, error) {
|
||||
if err := validateUpdateConfigRequest(req); err != nil {
|
||||
return storageConfig{}, err
|
||||
}
|
||||
currentByID := make(map[string]StorageEndpoint, len(current.Endpoints))
|
||||
for _, endpoint := range current.Endpoints {
|
||||
currentByID[endpoint.ID] = endpoint
|
||||
}
|
||||
next := storageConfig{
|
||||
Enabled: req.Enabled, BlockingEnabled: req.BlockingEnabled, StorePassEvents: req.StorePassEvents,
|
||||
Strategy: strings.TrimSpace(req.Strategy), WorkerCount: req.WorkerCount,
|
||||
QueueCapacity: req.QueueCapacity, Scanners: append([]string(nil), req.Scanners...),
|
||||
AllGroups: req.AllGroups, GroupIDs: append([]int64(nil), req.GroupIDs...),
|
||||
ConfigVersion: current.ConfigVersion, UpdatedBy: actorID,
|
||||
Endpoints: make([]StorageEndpoint, 0, len(req.Endpoints)),
|
||||
}
|
||||
for _, endpoint := range req.Endpoints {
|
||||
baseURL, err := NormalizeBaseURL(endpoint.BaseURL)
|
||||
if err != nil {
|
||||
return storageConfig{}, err
|
||||
}
|
||||
stored := StorageEndpoint{
|
||||
ID: strings.TrimSpace(endpoint.ID), Name: strings.TrimSpace(endpoint.Name),
|
||||
Protocol: strings.TrimSpace(endpoint.Protocol), BaseURL: baseURL, Model: strings.TrimSpace(endpoint.Model),
|
||||
TimeoutMS: endpoint.TimeoutMS, InputLimit: endpoint.InputLimit, Enabled: endpoint.Enabled,
|
||||
}
|
||||
old, hadOld := currentByID[stored.ID]
|
||||
switch {
|
||||
case endpoint.ClearToken:
|
||||
stored.TokenCiphertext = ""
|
||||
case strings.TrimSpace(endpoint.Token) != "":
|
||||
ciphertext, err := m.encryptor.Encrypt(strings.TrimSpace(endpoint.Token))
|
||||
if err != nil {
|
||||
return storageConfig{}, fmt.Errorf("encrypt prompt audit endpoint token: %w", err)
|
||||
}
|
||||
stored.TokenCiphertext = ciphertext
|
||||
case hadOld:
|
||||
stored.TokenCiphertext = old.TokenCiphertext
|
||||
}
|
||||
next.Endpoints = append(next.Endpoints, stored)
|
||||
}
|
||||
normalizeStorageConfig(&next)
|
||||
if err := validateStorageConfig(next); err != nil {
|
||||
return storageConfig{}, err
|
||||
}
|
||||
return next, nil
|
||||
}
|
||||
|
||||
func (m *ConfigManager) RuntimeState() (expected int64, active int64, loadedAt *time.Time, loadError string) {
|
||||
if m == nil {
|
||||
return 1, 0, nil, "config_manager_unavailable"
|
||||
}
|
||||
expected = m.expected.Load()
|
||||
if expected < 1 {
|
||||
expected = 1
|
||||
}
|
||||
if snapshot := m.snapshot.Load(); snapshot != nil {
|
||||
active = snapshot.active.ConfigVersion
|
||||
value := snapshot.loadedAt
|
||||
loadedAt = &value
|
||||
}
|
||||
m.stateMu.RLock()
|
||||
loadError = m.lastLoadError
|
||||
m.stateMu.RUnlock()
|
||||
return
|
||||
}
|
||||
|
||||
func (m *ConfigManager) Encrypt(value string) (string, error) { return m.encryptor.Encrypt(value) }
|
||||
func (m *ConfigManager) Decrypt(value string) (string, error) { return m.encryptor.Decrypt(value) }
|
||||
|
||||
func (m *ConfigManager) currentRiskControlEnabled() bool {
|
||||
if snapshot := m.snapshot.Load(); snapshot != nil {
|
||||
return snapshot.active.RiskControlEnabled
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (m *ConfigManager) observeExpectedState(raw string, riskControlEnabled bool) {
|
||||
if m == nil {
|
||||
return
|
||||
}
|
||||
if strings.TrimSpace(raw) == "" {
|
||||
m.expected.Store(1)
|
||||
m.expectedBlocking.Store(false)
|
||||
return
|
||||
}
|
||||
var intent struct {
|
||||
Enabled bool `json:"enabled"`
|
||||
BlockingEnabled bool `json:"blocking_enabled"`
|
||||
ConfigVersion int64 `json:"config_version"`
|
||||
}
|
||||
if err := json.Unmarshal([]byte(raw), &intent); err != nil {
|
||||
return
|
||||
}
|
||||
if intent.ConfigVersion < 1 {
|
||||
intent.ConfigVersion = 1
|
||||
}
|
||||
m.expected.Store(intent.ConfigVersion)
|
||||
m.expectedBlocking.Store(riskControlEnabled && intent.Enabled && intent.BlockingEnabled)
|
||||
}
|
||||
|
||||
func (m *ConfigManager) refreshLoop(ctx context.Context) {
|
||||
defer m.wg.Done()
|
||||
ticker := time.NewTicker(5 * time.Second)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
if err := m.Reload(ctx); err != nil {
|
||||
LogWarn(EventConfigReloadDegraded, map[string]any{"status": "degraded", "error_code": "config_ttl_reload_failed"})
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (m *ConfigManager) subscribeLoop(ctx context.Context) {
|
||||
defer m.wg.Done()
|
||||
pubsub := m.redis.Subscribe(ctx, ConfigInvalidationChannel)
|
||||
defer func() { _ = pubsub.Close() }()
|
||||
channel := pubsub.Channel()
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case message, ok := <-channel:
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
version, err := strconv.ParseInt(strings.TrimSpace(message.Payload), 10, 64)
|
||||
if err != nil || version < 1 {
|
||||
continue
|
||||
}
|
||||
m.expected.Store(version)
|
||||
if err := m.Reload(ctx); err != nil {
|
||||
LogWarn(EventConfigReloadDegraded, map[string]any{
|
||||
"config_version": version, "status": "degraded", "error_code": "config_invalidation_reload_failed",
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (m *ConfigManager) recordLoadError(_ error) {
|
||||
if m == nil {
|
||||
return
|
||||
}
|
||||
now := m.clock.Now()
|
||||
m.stateMu.Lock()
|
||||
m.lastLoadError = stableErrorMessage("config_load_failed")
|
||||
m.lastErrorAt = &now
|
||||
m.stateMu.Unlock()
|
||||
}
|
||||
|
||||
func (m *ConfigManager) clearLoadError() {
|
||||
m.stateMu.Lock()
|
||||
m.lastLoadError = ""
|
||||
m.lastErrorAt = nil
|
||||
m.stateMu.Unlock()
|
||||
}
|
||||
|
||||
func cloneStorageConfig(cfg storageConfig) storageConfig {
|
||||
cfg.Scanners = append([]string(nil), cfg.Scanners...)
|
||||
cfg.GroupIDs = append([]int64(nil), cfg.GroupIDs...)
|
||||
cfg.Endpoints = append([]StorageEndpoint(nil), cfg.Endpoints...)
|
||||
return cfg
|
||||
}
|
||||
|
||||
func cloneActiveConfig(cfg ActiveConfig) ActiveConfig {
|
||||
cfg.Scanners = append([]string(nil), cfg.Scanners...)
|
||||
cfg.GroupIDs = append([]int64(nil), cfg.GroupIDs...)
|
||||
cfg.Endpoints = append([]ActiveEndpoint(nil), cfg.Endpoints...)
|
||||
return cfg
|
||||
}
|
||||
@@ -0,0 +1,157 @@
|
||||
package securityaudit
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
type prefixEncryptor struct{}
|
||||
|
||||
func (prefixEncryptor) Encrypt(value string) (string, error) { return "enc:" + value, nil }
|
||||
func (prefixEncryptor) Decrypt(value string) (string, error) { return value[4:], nil }
|
||||
|
||||
func TestDefaultConfigIsOff(t *testing.T) {
|
||||
storage, err := ParseStorageConfig("")
|
||||
require.NoError(t, err)
|
||||
require.False(t, storage.Enabled)
|
||||
active, err := ActiveFromStorage(storage, true, prefixEncryptor{})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, ModeOff, active.EffectiveMode())
|
||||
require.Equal(t, AllScannerIDs, storage.Scanners)
|
||||
publicJSON, err := json.Marshal(PublicFromStorage(storage, true))
|
||||
require.NoError(t, err)
|
||||
require.Contains(t, string(publicJSON), `"group_ids":[]`)
|
||||
require.Contains(t, string(publicJSON), `"endpoints":[]`)
|
||||
}
|
||||
|
||||
func TestConfigRejectsBlockingWithoutAudit(t *testing.T) {
|
||||
storage := DefaultStorageConfig()
|
||||
storage.BlockingEnabled = true
|
||||
require.Error(t, validateStorageConfig(storage))
|
||||
}
|
||||
|
||||
func TestPublicConfigNeverMarshalsToken(t *testing.T) {
|
||||
storage := DefaultStorageConfig()
|
||||
storage.Endpoints = []StorageEndpoint{{ID: "one", Name: "One", Protocol: "openai_compatible", BaseURL: "http://127.0.0.1:8080", Model: DefaultGuardModel, TokenCiphertext: "GUARD_TOKEN_CANARY_SECRET", TimeoutMS: 1000, InputLimit: 1000, Enabled: true}}
|
||||
public := PublicFromStorage(storage, true)
|
||||
raw, err := json.Marshal(public)
|
||||
require.NoError(t, err)
|
||||
require.NotContains(t, string(raw), "GUARD_TOKEN_CANARY_SECRET")
|
||||
require.NotContains(t, string(raw), "ciphertext")
|
||||
require.True(t, public.Endpoints[0].HasToken)
|
||||
}
|
||||
|
||||
func TestConfigRuntimeLoadErrorIsStableBoundedAndSecretFree(t *testing.T) {
|
||||
const canary = "CONFIG_LOAD_CANARY_SECRET"
|
||||
manager := &ConfigManager{clock: fixedClock{}}
|
||||
manager.recordLoadError(errors.New("decrypt failed for token " + canary + " Authorization: Bearer " + canary))
|
||||
_, _, _, message := manager.RuntimeState()
|
||||
require.Equal(t, stableErrorMessage("config_load_failed"), message)
|
||||
require.NotContains(t, message, canary)
|
||||
require.LessOrEqual(t, len([]rune(message)), 160)
|
||||
}
|
||||
|
||||
func TestBuildNextStoragePreserveReplaceAndClearToken(t *testing.T) {
|
||||
manager := &ConfigManager{encryptor: prefixEncryptor{}}
|
||||
current := DefaultStorageConfig()
|
||||
current.Endpoints = []StorageEndpoint{{ID: "one", Name: "One", Protocol: "openai_compatible", BaseURL: "http://127.0.0.1:8080", Model: DefaultGuardModel, TokenCiphertext: "enc:old", TimeoutMS: 1000, InputLimit: 1000}}
|
||||
base := UpdateConfigRequest{ExpectedConfigVersion: 1, Strategy: "priority", WorkerCount: 1, QueueCapacity: 10, Scanners: []string{"PII"}, AllGroups: true,
|
||||
Endpoints: []UpdateEndpoint{{ID: "one", Name: "One", Protocol: "openai_compatible", BaseURL: "http://127.0.0.1:8080", TimeoutMS: 1000, InputLimit: 1000}}}
|
||||
preserved, err := manager.buildNextStorage(current, base, 9)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "enc:old", preserved.Endpoints[0].TokenCiphertext)
|
||||
replacedReq := base
|
||||
replacedReq.Endpoints = append([]UpdateEndpoint(nil), base.Endpoints...)
|
||||
replacedReq.Endpoints[0].Token = "new"
|
||||
replaced, err := manager.buildNextStorage(current, replacedReq, 9)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "enc:new", replaced.Endpoints[0].TokenCiphertext)
|
||||
clearedReq := base
|
||||
clearedReq.Endpoints = append([]UpdateEndpoint(nil), base.Endpoints...)
|
||||
clearedReq.Endpoints[0].ClearToken = true
|
||||
cleared, err := manager.buildNextStorage(current, clearedReq, 9)
|
||||
require.NoError(t, err)
|
||||
require.Empty(t, cleared.Endpoints[0].TokenCiphertext)
|
||||
}
|
||||
|
||||
func TestEffectiveModeTruthTable(t *testing.T) {
|
||||
tests := []struct {
|
||||
risk, enabled, blocking bool
|
||||
want Mode
|
||||
}{
|
||||
{false, false, false, ModeOff}, {false, true, true, ModeOff}, {true, false, false, ModeOff},
|
||||
{true, true, false, ModeAsync}, {true, true, true, ModeBlocking},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
cfg := ActiveConfig{RiskControlEnabled: tt.risk, Enabled: tt.enabled, BlockingEnabled: tt.blocking}
|
||||
require.Equal(t, tt.want, cfg.EffectiveMode())
|
||||
}
|
||||
}
|
||||
|
||||
func TestConfigManagerColdStartOnlyFailsClosedForExplicitBlockingIntent(t *testing.T) {
|
||||
manager := &ConfigManager{}
|
||||
|
||||
manager.observeExpectedState(`{"enabled":true,"blocking_enabled":false,"config_version":42}`, true)
|
||||
require.Equal(t, int64(42), manager.expected.Load())
|
||||
require.Equal(t, ModeOff, manager.EffectiveMode(), "an async config version must not imply blocking")
|
||||
|
||||
manager.observeExpectedState(`{"enabled":true,"blocking_enabled":true,"config_version":43}`, false)
|
||||
require.Equal(t, ModeOff, manager.EffectiveMode(), "the global risk-control switch still gates blocking")
|
||||
|
||||
manager.observeExpectedState(`{"enabled":true,"blocking_enabled":true,"config_version":44}`, true)
|
||||
require.Equal(t, ModeBlocking, manager.EffectiveMode())
|
||||
|
||||
manager.observeExpectedState(`{"enabled":true`, true)
|
||||
require.Equal(t, ModeBlocking, manager.EffectiveMode(), "undecodable storage must not erase the last known strict intent")
|
||||
}
|
||||
|
||||
func TestParseLegacyConfigDefaultsMissingFieldsWithoutEnablingBlocking(t *testing.T) {
|
||||
storage, err := ParseStorageConfig(`{"enabled":false,"config_version":9}`)
|
||||
require.NoError(t, err)
|
||||
require.False(t, storage.BlockingEnabled)
|
||||
require.Equal(t, "priority", storage.Strategy)
|
||||
require.Equal(t, DefaultWorkerCount, storage.WorkerCount)
|
||||
require.Equal(t, DefaultQueueCapacity, storage.QueueCapacity)
|
||||
require.Equal(t, AllScannerIDs, storage.Scanners)
|
||||
require.True(t, storage.AllGroups)
|
||||
}
|
||||
|
||||
func TestUpdateConfigStrictBoundsAndKnownValues(t *testing.T) {
|
||||
valid := promptAuditUpdateRequest(1, 1, "")
|
||||
require.NoError(t, validateUpdateConfigRequest(valid))
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
mutate func(*UpdateConfigRequest)
|
||||
reason string
|
||||
}{
|
||||
{name: "strategy", mutate: func(req *UpdateConfigRequest) { req.Strategy = "round_robin" }, reason: "prompt_audit_invalid_strategy"},
|
||||
{name: "worker low", mutate: func(req *UpdateConfigRequest) { req.WorkerCount = 0 }, reason: "prompt_audit_invalid_worker_count"},
|
||||
{name: "worker high", mutate: func(req *UpdateConfigRequest) { req.WorkerCount = MaxWorkerCount + 1 }, reason: "prompt_audit_invalid_worker_count"},
|
||||
{name: "capacity low", mutate: func(req *UpdateConfigRequest) { req.QueueCapacity = 0 }, reason: "prompt_audit_invalid_queue_capacity"},
|
||||
{name: "capacity high", mutate: func(req *UpdateConfigRequest) { req.QueueCapacity = MaxQueueCapacity + 1 }, reason: "prompt_audit_invalid_queue_capacity"},
|
||||
{name: "unknown scanner", mutate: func(req *UpdateConfigRequest) { req.Scanners = []string{"made_up"} }, reason: "prompt_audit_invalid_scanner"},
|
||||
{name: "group required", mutate: func(req *UpdateConfigRequest) { req.AllGroups = false; req.GroupIDs = nil }, reason: "prompt_audit_groups_required"},
|
||||
{name: "group positive", mutate: func(req *UpdateConfigRequest) { req.AllGroups = false; req.GroupIDs = []int64{0} }, reason: "prompt_audit_invalid_group"},
|
||||
{name: "timeout low", mutate: func(req *UpdateConfigRequest) { req.Endpoints[0].TimeoutMS = MinTimeoutMS - 1 }, reason: "prompt_audit_invalid_timeout"},
|
||||
{name: "timeout high", mutate: func(req *UpdateConfigRequest) { req.Endpoints[0].TimeoutMS = MaxTimeoutMS + 1 }, reason: "prompt_audit_invalid_timeout"},
|
||||
{name: "input low", mutate: func(req *UpdateConfigRequest) { req.Endpoints[0].InputLimit = MinInputLimit - 1 }, reason: "prompt_audit_invalid_input_limit"},
|
||||
{name: "input high", mutate: func(req *UpdateConfigRequest) { req.Endpoints[0].InputLimit = MaxInputLimit + 1 }, reason: "prompt_audit_invalid_input_limit"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
req := valid
|
||||
req.Scanners = append([]string(nil), valid.Scanners...)
|
||||
req.GroupIDs = append([]int64(nil), valid.GroupIDs...)
|
||||
req.Endpoints = append([]UpdateEndpoint(nil), valid.Endpoints...)
|
||||
tt.mutate(&req)
|
||||
err := validateUpdateConfigRequest(req)
|
||||
require.Error(t, err)
|
||||
require.Equal(t, tt.reason, infraerrors.Reason(err))
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,99 @@
|
||||
package securityaudit
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
)
|
||||
|
||||
type Enqueuer struct {
|
||||
config ConfigStore
|
||||
repo JobRepository
|
||||
payload PayloadStore
|
||||
metrics Metrics
|
||||
}
|
||||
|
||||
func NewEnqueuer(config ConfigStore, repo JobRepository, payload PayloadStore, metrics ...Metrics) *Enqueuer {
|
||||
var metric Metrics
|
||||
if len(metrics) > 0 {
|
||||
metric = metrics[0]
|
||||
}
|
||||
return &Enqueuer{config: config, repo: repo, payload: payload, metrics: metric}
|
||||
}
|
||||
|
||||
func (e *Enqueuer) Enqueue(ctx context.Context, req Request) error {
|
||||
if e == nil || e.config == nil || e.repo == nil || e.payload == nil {
|
||||
return errors.New("prompt audit enqueuer unavailable")
|
||||
}
|
||||
cfg, ok := e.config.Active()
|
||||
baseFields := requestLogFields(req)
|
||||
if !ok || cfg.EffectiveMode() != ModeAsync {
|
||||
LogInfo(EventEnqueueSkipped, mergeLogFields(baseFields, map[string]any{"status": "skipped", "error_code": "mode_not_async"}))
|
||||
return nil
|
||||
}
|
||||
baseFields["config_version"] = cfg.ConfigVersion
|
||||
if !cfg.IncludesGroup(req.GroupID) {
|
||||
LogInfo(EventEnqueueSkipped, mergeLogFields(baseFields, map[string]any{"status": "skipped", "error_code": "group_out_of_scope"}))
|
||||
return nil
|
||||
}
|
||||
if len(cfg.EnabledEndpoints()) == 0 {
|
||||
e.recordDropped()
|
||||
LogWarn(EventEnqueueDropped, mergeLogFields(baseFields, map[string]any{"status": "dropped", "error_code": "no_enabled_endpoint"}))
|
||||
return nil
|
||||
}
|
||||
snapshot, err := ExtractPromptSnapshot(req)
|
||||
if errors.Is(err, ErrNoPromptText) {
|
||||
LogInfo(EventEnqueueSkipped, mergeLogFields(baseFields, map[string]any{"status": "skipped", "error_code": "no_user_text"}))
|
||||
return nil
|
||||
}
|
||||
if err != nil {
|
||||
e.recordDropped()
|
||||
LogWarn(EventEnqueueDropped, mergeLogFields(baseFields, map[string]any{"status": "dropped", "error_code": "snapshot_invalid"}))
|
||||
return nil
|
||||
}
|
||||
job, err := e.repo.CreateStagingWithCapacity(ctx, snapshot.Redacted(), cfg.ConfigVersion, 3, cfg.QueueCapacity)
|
||||
if err != nil {
|
||||
code := "database_unavailable"
|
||||
if errors.Is(err, ErrQueueFull) {
|
||||
code = "queue_full"
|
||||
}
|
||||
if errors.Is(err, ErrQueueAdmissionBusy) {
|
||||
code = "queue_admission_busy"
|
||||
}
|
||||
LogWarn(EventEnqueueDropped, mergeLogFields(baseFields, map[string]any{
|
||||
"queue_capacity": cfg.QueueCapacity, "status": "dropped", "error_code": code,
|
||||
}))
|
||||
e.recordDropped()
|
||||
return err
|
||||
}
|
||||
if err := e.payload.Set(ctx, job.ID, snapshot.ScanText, DefaultPayloadTTL); err != nil {
|
||||
_ = e.repo.MarkStagingFailed(ctx, job.ID, "payload_store_failed", "payload store unavailable")
|
||||
LogWarn(EventEnqueueDropped, mergeLogFields(baseFields, map[string]any{
|
||||
"job_id": job.ID, "status": "dropped", "error_code": "payload_store_failed",
|
||||
}))
|
||||
e.recordDropped()
|
||||
return err
|
||||
}
|
||||
if err := e.repo.PublishQueued(ctx, job.ID); err != nil {
|
||||
_ = e.payload.Delete(ctx, job.ID)
|
||||
_ = e.repo.MarkStagingFailed(ctx, job.ID, "queue_publish_failed", "queue publish failed")
|
||||
LogWarn(EventEnqueueDropped, mergeLogFields(baseFields, map[string]any{
|
||||
"job_id": job.ID, "status": "dropped", "error_code": "queue_publish_failed",
|
||||
}))
|
||||
e.recordDropped()
|
||||
return err
|
||||
}
|
||||
LogInfo(EventJobEnqueued, mergeLogFields(baseFields, map[string]any{
|
||||
"job_id": job.ID,
|
||||
"queue_capacity": cfg.QueueCapacity, "status": "queued",
|
||||
}))
|
||||
if e.metrics != nil {
|
||||
e.metrics.IncEnqueued()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (e *Enqueuer) recordDropped() {
|
||||
if e != nil && e.metrics != nil {
|
||||
e.metrics.IncDropped()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,377 @@
|
||||
package securityaudit
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"database/sql"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/lib/pq"
|
||||
)
|
||||
|
||||
type EventFilter struct {
|
||||
Decision string `json:"decision,omitempty"`
|
||||
RiskLevel string `json:"risk_level,omitempty"`
|
||||
Endpoint string `json:"endpoint,omitempty"`
|
||||
GroupID *int64 `json:"group_id,omitempty"`
|
||||
UserID *int64 `json:"user_id,omitempty"`
|
||||
APIKeyID *int64 `json:"api_key_id,omitempty"`
|
||||
RequestID string `json:"request_id,omitempty"`
|
||||
PromptHash string `json:"prompt_hash,omitempty"`
|
||||
Keyword string `json:"keyword,omitempty"`
|
||||
StartAt *time.Time `json:"start_at,omitempty"`
|
||||
EndAt *time.Time `json:"end_at,omitempty"`
|
||||
}
|
||||
|
||||
type EventPage struct {
|
||||
Items []*Event `json:"items"`
|
||||
Total int64 `json:"total"`
|
||||
Page int `json:"page"`
|
||||
PageSize int `json:"page_size"`
|
||||
Pages int `json:"pages"`
|
||||
}
|
||||
|
||||
type DeletePreview struct {
|
||||
MatchedCount int64 `json:"matched_count"`
|
||||
FilterSummary EventFilter `json:"filter_summary"`
|
||||
SnapshotMaxID int64 `json:"snapshot_max_id"`
|
||||
FilterHash string `json:"filter_hash"`
|
||||
ConfirmationToken string `json:"confirmation_token,omitempty"`
|
||||
ExpiresAt time.Time `json:"expires_at,omitempty"`
|
||||
}
|
||||
|
||||
type DeleteResult struct {
|
||||
DeletedEvents int64 `json:"deleted_events"`
|
||||
DeletedJobs int64 `json:"deleted_jobs"`
|
||||
JobIDs []int64 `json:"-"`
|
||||
}
|
||||
|
||||
type EventRepository interface {
|
||||
ListEvents(ctx context.Context, filter EventFilter, page, pageSize int) (*EventPage, error)
|
||||
GetEvent(ctx context.Context, id int64) (*Event, error)
|
||||
DeleteEvent(ctx context.Context, id int64) (*DeleteResult, error)
|
||||
DeleteEventsByIDs(ctx context.Context, ids []int64) (*DeleteResult, error)
|
||||
PreviewDelete(ctx context.Context, filter EventFilter) (*DeletePreview, error)
|
||||
DeleteEventsByFilter(ctx context.Context, filter EventFilter, snapshotMaxID int64, batchSize int) (*DeleteResult, error)
|
||||
}
|
||||
|
||||
func (r *PostgreSQLRepository) ListEvents(ctx context.Context, filter EventFilter, page, pageSize int) (*EventPage, error) {
|
||||
if page < 1 {
|
||||
page = 1
|
||||
}
|
||||
if pageSize < 1 {
|
||||
pageSize = 20
|
||||
}
|
||||
if pageSize > 100 {
|
||||
pageSize = 100
|
||||
}
|
||||
where, args := buildEventWhere(filter, 1)
|
||||
var total int64
|
||||
if err := r.db.QueryRowContext(ctx, `SELECT COUNT(*) FROM prompt_audit_events e`+where, args...).Scan(&total); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
queryArgs := append([]any(nil), args...)
|
||||
limitIndex := len(queryArgs) + 1
|
||||
queryArgs = append(queryArgs, pageSize, (page-1)*pageSize)
|
||||
rows, err := r.db.QueryContext(ctx, `SELECT `+eventColumns("e")+` FROM prompt_audit_events e`+where+
|
||||
fmt.Sprintf(` ORDER BY e.created_at DESC, e.id DESC LIMIT $%d OFFSET $%d`, limitIndex, limitIndex+1), queryArgs...)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
items := make([]*Event, 0, pageSize)
|
||||
for rows.Next() {
|
||||
event, err := scanEvent(rows)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
items = append(items, event)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
pages := 0
|
||||
if total > 0 {
|
||||
pages = int((total + int64(pageSize) - 1) / int64(pageSize))
|
||||
}
|
||||
return &EventPage{Items: items, Total: total, Page: page, PageSize: pageSize, Pages: pages}, nil
|
||||
}
|
||||
|
||||
func (r *PostgreSQLRepository) GetEvent(ctx context.Context, id int64) (*Event, error) {
|
||||
event, err := scanEvent(r.db.QueryRowContext(ctx, `SELECT `+eventColumns("e")+` FROM prompt_audit_events e WHERE e.id=$1`, id))
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return nil, ErrEventNotFound
|
||||
}
|
||||
return event, err
|
||||
}
|
||||
|
||||
func (r *PostgreSQLRepository) DeleteEvent(ctx context.Context, id int64) (*DeleteResult, error) {
|
||||
return r.DeleteEventsByIDs(ctx, []int64{id})
|
||||
}
|
||||
|
||||
func (r *PostgreSQLRepository) DeleteEventsByIDs(ctx context.Context, ids []int64) (*DeleteResult, error) {
|
||||
ids = canonicalInt64s(ids)
|
||||
if len(ids) == 0 {
|
||||
return &DeleteResult{}, nil
|
||||
}
|
||||
if len(ids) > 500 {
|
||||
return nil, errors.New("prompt audit delete batch exceeds 500 events")
|
||||
}
|
||||
tx, err := r.db.BeginTx(ctx, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer func() { _ = tx.Rollback() }()
|
||||
rows, err := tx.QueryContext(ctx, `DELETE FROM prompt_audit_events WHERE id=ANY($1) RETURNING job_id`, pq.Array(ids))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
jobIDs, err := scanReturnedJobIDs(rows)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
deletedJobs, err := deleteOrphanJobs(ctx, tx, jobIDs)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := tx.Commit(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &DeleteResult{DeletedEvents: int64(len(jobIDs)), DeletedJobs: deletedJobs, JobIDs: canonicalInt64s(jobIDs)}, nil
|
||||
}
|
||||
|
||||
func (r *PostgreSQLRepository) PreviewDelete(ctx context.Context, filter EventFilter) (*DeletePreview, error) {
|
||||
if err := validateDeleteFilter(filter); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
tx, err := r.db.BeginTx(ctx, &sql.TxOptions{Isolation: sql.LevelRepeatableRead, ReadOnly: true})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer func() { _ = tx.Rollback() }()
|
||||
where, args := buildEventWhere(filter, 1)
|
||||
var count, maxID int64
|
||||
if err := tx.QueryRowContext(ctx, `SELECT COUNT(*), COALESCE(MAX(e.id),0) FROM prompt_audit_events e`+where, args...).Scan(&count, &maxID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := tx.Commit(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
canonical := canonicalEventFilter(filter)
|
||||
return &DeletePreview{MatchedCount: count, FilterSummary: canonical, SnapshotMaxID: maxID, FilterHash: FilterHash(canonical, maxID)}, nil
|
||||
}
|
||||
|
||||
func (r *PostgreSQLRepository) DeleteEventsByFilter(ctx context.Context, filter EventFilter, snapshotMaxID int64, batchSize int) (*DeleteResult, error) {
|
||||
if err := validateDeleteFilter(filter); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if snapshotMaxID <= 0 {
|
||||
return &DeleteResult{}, nil
|
||||
}
|
||||
if batchSize < 1 || batchSize > 1000 {
|
||||
batchSize = 200
|
||||
}
|
||||
total := &DeleteResult{}
|
||||
jobSet := map[int64]struct{}{}
|
||||
for {
|
||||
tx, err := r.db.BeginTx(ctx, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
where, args := buildEventWhere(filter, 1)
|
||||
maxIndex := len(args) + 1
|
||||
limitIndex := maxIndex + 1
|
||||
args = append(args, snapshotMaxID, batchSize)
|
||||
rows, err := tx.QueryContext(ctx, `
|
||||
WITH selected AS (
|
||||
SELECT e.id FROM prompt_audit_events e`+where+
|
||||
fmt.Sprintf(` AND e.id <= $%d ORDER BY e.id LIMIT $%d FOR UPDATE SKIP LOCKED`, maxIndex, limitIndex)+`
|
||||
), deleted AS (
|
||||
DELETE FROM prompt_audit_events e USING selected s WHERE e.id=s.id RETURNING e.job_id
|
||||
) SELECT job_id FROM deleted`, args...)
|
||||
if err != nil {
|
||||
_ = tx.Rollback()
|
||||
return nil, err
|
||||
}
|
||||
jobIDs, err := scanReturnedJobIDs(rows)
|
||||
if err != nil {
|
||||
_ = tx.Rollback()
|
||||
return nil, err
|
||||
}
|
||||
deletedJobs, err := deleteOrphanJobs(ctx, tx, jobIDs)
|
||||
if err != nil {
|
||||
_ = tx.Rollback()
|
||||
return nil, err
|
||||
}
|
||||
if err := tx.Commit(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
total.DeletedEvents += int64(len(jobIDs))
|
||||
total.DeletedJobs += deletedJobs
|
||||
for _, id := range jobIDs {
|
||||
jobSet[id] = struct{}{}
|
||||
}
|
||||
if len(jobIDs) < batchSize {
|
||||
break
|
||||
}
|
||||
}
|
||||
for id := range jobSet {
|
||||
total.JobIDs = append(total.JobIDs, id)
|
||||
}
|
||||
total.JobIDs = canonicalInt64s(total.JobIDs)
|
||||
return total, nil
|
||||
}
|
||||
|
||||
func FilterHash(filter EventFilter, snapshotMaxID int64) string {
|
||||
payload := struct {
|
||||
Filter EventFilter `json:"filter"`
|
||||
SnapshotMaxID int64 `json:"snapshot_max_id"`
|
||||
}{canonicalEventFilter(filter), snapshotMaxID}
|
||||
raw, _ := json.Marshal(payload)
|
||||
digest := sha256.Sum256(raw)
|
||||
return hex.EncodeToString(digest[:])
|
||||
}
|
||||
|
||||
func validateDeleteFilter(filter EventFilter) error {
|
||||
if filter.StartAt == nil || filter.EndAt == nil || !filter.StartAt.Before(*filter.EndAt) {
|
||||
return errors.New("prompt audit filter delete requires a valid explicit time range")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func canonicalEventFilter(filter EventFilter) EventFilter {
|
||||
filter.Decision = strings.TrimSpace(strings.ToLower(filter.Decision))
|
||||
filter.RiskLevel = strings.TrimSpace(strings.ToLower(filter.RiskLevel))
|
||||
filter.Endpoint = strings.TrimSpace(filter.Endpoint)
|
||||
filter.RequestID = strings.TrimSpace(filter.RequestID)
|
||||
filter.PromptHash = strings.ToLower(strings.TrimSpace(filter.PromptHash))
|
||||
filter.Keyword = strings.TrimSpace(filter.Keyword)
|
||||
if filter.StartAt != nil {
|
||||
value := filter.StartAt.UTC()
|
||||
filter.StartAt = &value
|
||||
}
|
||||
if filter.EndAt != nil {
|
||||
value := filter.EndAt.UTC()
|
||||
filter.EndAt = &value
|
||||
}
|
||||
return filter
|
||||
}
|
||||
|
||||
func buildEventWhere(filter EventFilter, firstIndex int) (string, []any) {
|
||||
filter = canonicalEventFilter(filter)
|
||||
clauses := []string{" WHERE TRUE"}
|
||||
args := make([]any, 0, 12)
|
||||
add := func(clause string, value any) {
|
||||
clauses = append(clauses, fmt.Sprintf(clause, firstIndex+len(args)))
|
||||
args = append(args, value)
|
||||
}
|
||||
if filter.Decision != "" {
|
||||
add(" AND e.decision=$%d", filter.Decision)
|
||||
}
|
||||
if filter.RiskLevel != "" {
|
||||
add(" AND e.risk_level=$%d", filter.RiskLevel)
|
||||
}
|
||||
if filter.Endpoint != "" {
|
||||
add(" AND e.endpoint=$%d", filter.Endpoint)
|
||||
}
|
||||
if filter.GroupID != nil {
|
||||
add(" AND e.group_id=$%d", *filter.GroupID)
|
||||
}
|
||||
if filter.UserID != nil {
|
||||
add(" AND e.user_id=$%d", *filter.UserID)
|
||||
}
|
||||
if filter.APIKeyID != nil {
|
||||
add(" AND e.api_key_id=$%d", *filter.APIKeyID)
|
||||
}
|
||||
if filter.RequestID != "" {
|
||||
add(" AND e.request_id=$%d", filter.RequestID)
|
||||
}
|
||||
if filter.PromptHash != "" {
|
||||
add(" AND e.prompt_hash=$%d", filter.PromptHash)
|
||||
}
|
||||
if filter.Keyword != "" {
|
||||
add(` AND (e.request_id ILIKE $%d OR e.prompt_hash ILIKE $%d OR e.redacted_preview ILIKE $%d
|
||||
OR e.username_snapshot ILIKE $%d OR e.user_email_snapshot ILIKE $%d OR e.api_key_name_snapshot ILIKE $%d)`, "%"+TrimRunes(filter.Keyword, 128)+"%")
|
||||
// The clause has six placeholders but add only supplied one. Rebuild it with one shared placeholder.
|
||||
clauses[len(clauses)-1] = fmt.Sprintf(` AND (e.request_id ILIKE $%[1]d OR e.prompt_hash ILIKE $%[1]d OR e.redacted_preview ILIKE $%[1]d
|
||||
OR e.username_snapshot ILIKE $%[1]d OR e.user_email_snapshot ILIKE $%[1]d OR e.api_key_name_snapshot ILIKE $%[1]d)`, firstIndex+len(args)-1)
|
||||
}
|
||||
if filter.StartAt != nil {
|
||||
add(" AND e.created_at >= $%d", filter.StartAt.UTC())
|
||||
}
|
||||
if filter.EndAt != nil {
|
||||
add(" AND e.created_at <= $%d", filter.EndAt.UTC())
|
||||
}
|
||||
return strings.Join(clauses, ""), args
|
||||
}
|
||||
|
||||
func eventColumns(alias string) string {
|
||||
return fmt.Sprintf(`%[1]s.id,%[1]s.job_id,%[1]s.request_id,%[1]s.user_id,%[1]s.username_snapshot,
|
||||
%[1]s.user_email_snapshot,%[1]s.api_key_id,%[1]s.api_key_name_snapshot,%[1]s.group_id,%[1]s.group_name,
|
||||
%[1]s.provider,%[1]s.endpoint,%[1]s.protocol,%[1]s.model,%[1]s.prompt_hash,%[1]s.redacted_preview,
|
||||
%[1]s.decision,%[1]s.risk_level,%[1]s.action,%[1]s.categories,%[1]s.matched_scanners,
|
||||
%[1]s.scanner_scores,%[1]s.scanner_evidence,%[1]s.scanner_backend,%[1]s.scanner_version,
|
||||
%[1]s.guard_endpoint_id,%[1]s.policy_id,%[1]s.policy_version,%[1]s.config_version,
|
||||
%[1]s.chunk_total,%[1]s.latency_ms,%[1]s.created_at`, alias)
|
||||
}
|
||||
|
||||
func scanEvent(row rowScanner) (*Event, error) {
|
||||
event := &Event{}
|
||||
var userID, apiKeyID, groupID sql.NullInt64
|
||||
var categories, matched, scores, evidence []byte
|
||||
err := row.Scan(&event.ID, &event.JobID, &event.Snapshot.RequestID, &userID,
|
||||
&event.Snapshot.UsernameSnapshot, &event.Snapshot.UserEmailSnapshot, &apiKeyID,
|
||||
&event.Snapshot.APIKeyNameSnapshot, &groupID, &event.Snapshot.GroupName,
|
||||
&event.Snapshot.Provider, &event.Snapshot.Endpoint, &event.Snapshot.Protocol, &event.Snapshot.Model,
|
||||
&event.Snapshot.PromptHash, &event.Snapshot.RedactedPreview, &event.Decision, &event.RiskLevel,
|
||||
&event.Action, &categories, &matched, &scores, &evidence, &event.ScannerBackend,
|
||||
&event.ScannerVersion, &event.GuardEndpointID, &event.PolicyID, &event.PolicyVersion,
|
||||
&event.ConfigVersion, &event.ChunkTotal, &event.LatencyMS, &event.CreatedAt)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
event.Snapshot.UserID = nullableInt64Value(userID)
|
||||
event.Snapshot.APIKeyID = nullableInt64Value(apiKeyID)
|
||||
event.Snapshot.GroupID = nullableInt64Ptr(groupID)
|
||||
_ = json.Unmarshal(categories, &event.Categories)
|
||||
_ = json.Unmarshal(matched, &event.MatchedScanners)
|
||||
_ = json.Unmarshal(scores, &event.ScannerScores)
|
||||
_ = json.Unmarshal(evidence, &event.ScannerEvidence)
|
||||
result := NormalizedResult{Decision: event.Decision, RiskLevel: event.RiskLevel, Action: event.Action,
|
||||
Categories: event.Categories, MatchedScanners: event.MatchedScanners, ScannerScores: event.ScannerScores,
|
||||
ScannerEvidence: event.ScannerEvidence}
|
||||
event.IssueSummaries = BuildIssueSummaries(result)
|
||||
return event, nil
|
||||
}
|
||||
|
||||
func scanReturnedJobIDs(rows *sql.Rows) ([]int64, error) {
|
||||
defer func() { _ = rows.Close() }()
|
||||
result := make([]int64, 0)
|
||||
for rows.Next() {
|
||||
var id int64
|
||||
if err := rows.Scan(&id); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result = append(result, id)
|
||||
}
|
||||
return result, rows.Err()
|
||||
}
|
||||
|
||||
func deleteOrphanJobs(ctx context.Context, tx *sql.Tx, jobIDs []int64) (int64, error) {
|
||||
jobIDs = canonicalInt64s(jobIDs)
|
||||
if len(jobIDs) == 0 {
|
||||
return 0, nil
|
||||
}
|
||||
result, err := tx.ExecContext(ctx, `DELETE FROM prompt_audit_jobs j
|
||||
WHERE j.id=ANY($1) AND j.status <> 'processing'
|
||||
AND NOT EXISTS (SELECT 1 FROM prompt_audit_events e WHERE e.job_id=j.id)`, pq.Array(jobIDs))
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return result.RowsAffected()
|
||||
}
|
||||
@@ -0,0 +1,279 @@
|
||||
package securityaudit
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
type GuardEvaluator struct {
|
||||
scanner PromptScanner
|
||||
repo JobRepository
|
||||
metrics Metrics
|
||||
clock Clock
|
||||
|
||||
global chan struct{}
|
||||
perNodeLimit int
|
||||
nodeMu sync.Mutex
|
||||
nodes map[string]chan struct{}
|
||||
}
|
||||
|
||||
func NewGuardEvaluator(scanner PromptScanner, repo JobRepository, metrics Metrics) *GuardEvaluator {
|
||||
return newGuardEvaluator(scanner, repo, metrics, 64, 16)
|
||||
}
|
||||
|
||||
func newGuardEvaluator(scanner PromptScanner, repo JobRepository, metrics Metrics, globalLimit, perNodeLimit int) *GuardEvaluator {
|
||||
if globalLimit < 1 {
|
||||
globalLimit = 64
|
||||
}
|
||||
if perNodeLimit < 1 {
|
||||
perNodeLimit = 16
|
||||
}
|
||||
return &GuardEvaluator{scanner: scanner, repo: repo, metrics: metrics, clock: realClock{},
|
||||
global: make(chan struct{}, globalLimit), perNodeLimit: perNodeLimit, nodes: map[string]chan struct{}{}}
|
||||
}
|
||||
|
||||
func (g *GuardEvaluator) Evaluate(ctx context.Context, cfg ActiveConfig, snapshot PromptSnapshot) (*PromptDecision, error) {
|
||||
if g == nil || g.scanner == nil {
|
||||
if g != nil && g.metrics != nil {
|
||||
g.metrics.Observe(DecisionUnavailable, 0)
|
||||
}
|
||||
logGuardFailure(snapshot, cfg, DecisionUnavailable, ErrorCodeUnavailable, "", 0)
|
||||
return nil, &GuardError{Code: ErrorCodeUnavailable}
|
||||
}
|
||||
start := g.clock.Now()
|
||||
baseFields := snapshotLogFields(snapshot)
|
||||
baseFields["config_version"] = cfg.ConfigVersion
|
||||
endpoints := cfg.EnabledEndpoints()
|
||||
if len(endpoints) == 0 {
|
||||
if g.metrics != nil {
|
||||
g.metrics.Observe(DecisionUnavailable, g.clock.Now().Sub(start))
|
||||
}
|
||||
logGuardFailure(snapshot, cfg, DecisionUnavailable, ErrorCodeUnavailable, "", g.clock.Now().Sub(start))
|
||||
return nil, &GuardError{Code: ErrorCodeUnavailable}
|
||||
}
|
||||
select {
|
||||
case g.global <- struct{}{}:
|
||||
defer func() { <-g.global }()
|
||||
default:
|
||||
if g.metrics != nil {
|
||||
g.metrics.IncBulkheadFull()
|
||||
g.metrics.Observe(DecisionUnavailable, g.clock.Now().Sub(start))
|
||||
}
|
||||
logGuardFailure(snapshot, cfg, DecisionUnavailable, ErrorCodeUnavailable, "", g.clock.Now().Sub(start))
|
||||
return nil, &GuardError{Code: ErrorCodeUnavailable}
|
||||
}
|
||||
timeout := time.Duration(endpoints[0].TimeoutMS) * time.Millisecond
|
||||
if timeout <= 0 {
|
||||
timeout = DefaultTimeoutMS * time.Millisecond
|
||||
}
|
||||
evalCtx, cancel := context.WithTimeout(ctx, timeout)
|
||||
defer cancel()
|
||||
inputLimit := minimumInputLimit(endpoints)
|
||||
chunks := SplitRunes(snapshot.ScanText, inputLimit)
|
||||
if len(chunks) == 0 {
|
||||
if g.metrics != nil {
|
||||
g.metrics.Observe(DecisionAllow, g.clock.Now().Sub(start))
|
||||
}
|
||||
return &PromptDecision{Kind: DecisionAllow, AllowNextStage: true}, nil
|
||||
}
|
||||
LogInfo(EventEvaluationStarted, mergeLogFields(baseFields, map[string]any{"chunk_total": len(chunks), "status": "started"}))
|
||||
results := make([]*NormalizedResult, 0, len(chunks))
|
||||
for index, chunk := range chunks {
|
||||
chunkStarted := g.clock.Now()
|
||||
LogInfo(EventChunkStarted, mergeLogFields(baseFields, map[string]any{
|
||||
"chunk_index": index + 1, "chunk_total": len(chunks),
|
||||
"chunk_chars": len([]rune(chunk)), "input_chars": snapshot.PromptLength, "input_limit": inputLimit,
|
||||
"status": "started",
|
||||
}))
|
||||
result, err := g.scanChunk(evalCtx, cfg, endpoints, chunk)
|
||||
if err != nil {
|
||||
code := guardErrorCode(err)
|
||||
LogWarn(EventChunkFailed, mergeLogFields(baseFields, map[string]any{
|
||||
"chunk_index": index + 1, "chunk_total": len(chunks),
|
||||
"chunk_chars": len([]rune(chunk)), "input_chars": snapshot.PromptLength, "input_limit": inputLimit,
|
||||
"latency_ms": g.clock.Now().Sub(chunkStarted).Milliseconds(), "error_code": code, "status": "failed",
|
||||
}))
|
||||
kind := DecisionUnavailable
|
||||
if code == ErrorCodeInvalidResponse {
|
||||
kind = DecisionInvalid
|
||||
}
|
||||
if g.metrics != nil {
|
||||
g.metrics.Observe(kind, g.clock.Now().Sub(start))
|
||||
var guardErr *GuardError
|
||||
if errors.As(err, &guardErr) && guardErr.Timeout {
|
||||
g.metrics.IncTimeout()
|
||||
}
|
||||
}
|
||||
logGuardFailure(snapshot, cfg, kind, code, "", g.clock.Now().Sub(start))
|
||||
return nil, err
|
||||
}
|
||||
result.ChunkTotal = len(chunks)
|
||||
results = append(results, result)
|
||||
LogInfo(EventChunkCompleted, mergeLogFields(baseFields, map[string]any{
|
||||
"chunk_index": index + 1, "chunk_total": len(chunks),
|
||||
"chunk_chars": len([]rune(chunk)), "input_chars": snapshot.PromptLength, "input_limit": inputLimit,
|
||||
"guard_endpoint_id": result.GuardEndpointID, "action": result.Action,
|
||||
"latency_ms": g.clock.Now().Sub(chunkStarted).Milliseconds(), "status": "completed",
|
||||
}))
|
||||
if result.Action == ActionBlock {
|
||||
break
|
||||
}
|
||||
}
|
||||
aggregated, err := AggregateResults(results, g.clock.Now().Sub(start))
|
||||
if err != nil {
|
||||
if g.metrics != nil {
|
||||
g.metrics.Observe(DecisionInvalid, g.clock.Now().Sub(start))
|
||||
}
|
||||
logGuardFailure(snapshot, cfg, DecisionInvalid, ErrorCodeInvalidResponse, "", g.clock.Now().Sub(start))
|
||||
return nil, &GuardError{Code: ErrorCodeInvalidResponse, Cause: err}
|
||||
}
|
||||
aggregated.ChunkTotal = len(chunks)
|
||||
kind := DecisionAllow
|
||||
if aggregated.Action == ActionWarn {
|
||||
kind = DecisionFlag
|
||||
}
|
||||
if aggregated.Action == ActionBlock {
|
||||
kind = DecisionBlock
|
||||
}
|
||||
decision := &PromptDecision{Kind: kind, Result: aggregated, AllowNextStage: kind == DecisionAllow || kind == DecisionFlag}
|
||||
if kind == DecisionBlock {
|
||||
decision.ErrorCode = ErrorCodeBlocked
|
||||
}
|
||||
if g.metrics != nil {
|
||||
g.metrics.Observe(kind, g.clock.Now().Sub(start))
|
||||
}
|
||||
LogInfo(EventChunksAggregated, mergeLogFields(baseFields, map[string]any{
|
||||
"decision": kind,
|
||||
"risk_level": aggregated.RiskLevel, "action": aggregated.Action, "chunk_total": aggregated.ChunkTotal,
|
||||
"latency_ms": aggregated.LatencyMS, "guard_endpoint_id": aggregated.GuardEndpointID, "stage": snapshot.Stage,
|
||||
"status": "completed",
|
||||
}))
|
||||
if g.repo != nil {
|
||||
if _, recordErr := g.repo.RecordBlocking(ctx, snapshot.Redacted(), cfg.ConfigVersion, aggregated, cfg.StorePassEvents); recordErr != nil {
|
||||
if g.metrics != nil {
|
||||
g.metrics.IncRecordFailed()
|
||||
}
|
||||
LogWarn(EventResultRecordFailed, mergeLogFields(baseFields, map[string]any{
|
||||
"decision": kind, "error_code": "result_record_failed", "stage": snapshot.Stage,
|
||||
"status": "failed",
|
||||
}))
|
||||
}
|
||||
}
|
||||
if kind == DecisionBlock {
|
||||
LogWarn(EventGuardBlocked, mergeLogFields(baseFields, map[string]any{
|
||||
"guard_endpoint_id": aggregated.GuardEndpointID,
|
||||
"decision": kind, "risk_level": aggregated.RiskLevel, "action": aggregated.Action, "chunk_total": aggregated.ChunkTotal,
|
||||
"latency_ms": aggregated.LatencyMS, "status": "blocked", "error_code": ErrorCodeBlocked,
|
||||
"stage": snapshot.Stage, "upstream_dispatched": false, "billing_preconsumed": false,
|
||||
}))
|
||||
} else {
|
||||
LogInfo(EventGuardAllowed, mergeLogFields(baseFields, map[string]any{
|
||||
"decision": kind, "risk_level": aggregated.RiskLevel, "action": aggregated.Action,
|
||||
"guard_endpoint_id": aggregated.GuardEndpointID, "chunk_total": aggregated.ChunkTotal,
|
||||
"latency_ms": aggregated.LatencyMS, "stage": snapshot.Stage, "status": "allowed",
|
||||
}))
|
||||
}
|
||||
return decision, nil
|
||||
}
|
||||
|
||||
func logGuardFailure(snapshot PromptSnapshot, cfg ActiveConfig, kind DecisionKind, code, guardEndpointID string, latency time.Duration) {
|
||||
fields := snapshotLogFields(snapshot)
|
||||
fields["config_version"] = cfg.ConfigVersion
|
||||
LogWarn(EventGuardFailed, mergeLogFields(fields, map[string]any{
|
||||
"decision": kind, "guard_endpoint_id": guardEndpointID, "latency_ms": latency.Milliseconds(),
|
||||
"status": "failed", "error_code": code, "upstream_dispatched": false, "billing_preconsumed": false,
|
||||
}))
|
||||
}
|
||||
|
||||
func (g *GuardEvaluator) scanChunk(ctx context.Context, cfg ActiveConfig, endpoints []ActiveEndpoint, chunk string) (*NormalizedResult, error) {
|
||||
var lastErr error
|
||||
for index, endpoint := range endpoints {
|
||||
semaphore := g.nodeSemaphore(endpoint.ID)
|
||||
select {
|
||||
case semaphore <- struct{}{}:
|
||||
case <-ctx.Done():
|
||||
return nil, &GuardError{Code: ErrorCodeUnavailable, Retryable: true, Timeout: errors.Is(ctx.Err(), context.DeadlineExceeded), Cause: ctx.Err()}
|
||||
default:
|
||||
if g.metrics != nil {
|
||||
g.metrics.IncBulkheadFull()
|
||||
}
|
||||
lastErr = &GuardError{Code: ErrorCodeUnavailable, Retryable: true}
|
||||
if index < len(endpoints)-1 && g.metrics != nil {
|
||||
g.metrics.IncFailover()
|
||||
}
|
||||
continue
|
||||
}
|
||||
result, err := callPromptScanner(ctx, g.scanner, endpoint, chunk, cfg.Scanners)
|
||||
<-semaphore
|
||||
if err == nil && result != nil {
|
||||
return result, nil
|
||||
}
|
||||
if err == nil {
|
||||
err = &GuardError{Code: ErrorCodeInvalidResponse, Retryable: false}
|
||||
}
|
||||
lastErr = err
|
||||
var guardErr *GuardError
|
||||
if !errors.As(err, &guardErr) || !guardErr.Retryable {
|
||||
return nil, err
|
||||
}
|
||||
if index < len(endpoints)-1 && g.metrics != nil {
|
||||
g.metrics.IncFailover()
|
||||
}
|
||||
}
|
||||
if lastErr == nil {
|
||||
lastErr = &GuardError{Code: ErrorCodeUnavailable}
|
||||
}
|
||||
return nil, lastErr
|
||||
}
|
||||
|
||||
func callPromptScanner(ctx context.Context, scanner PromptScanner, endpoint ActiveEndpoint, chunk string, scanners []string) (result *NormalizedResult, err error) {
|
||||
defer func() {
|
||||
if recover() != nil {
|
||||
result = nil
|
||||
err = &GuardError{Code: ErrorCodeUnavailable, Retryable: false}
|
||||
}
|
||||
}()
|
||||
return scanner.Scan(ctx, endpoint, chunk, scanners)
|
||||
}
|
||||
|
||||
func (g *GuardEvaluator) nodeSemaphore(id string) chan struct{} {
|
||||
g.nodeMu.Lock()
|
||||
defer g.nodeMu.Unlock()
|
||||
semaphore := g.nodes[id]
|
||||
if semaphore == nil {
|
||||
semaphore = make(chan struct{}, g.perNodeLimit)
|
||||
g.nodes[id] = semaphore
|
||||
}
|
||||
return semaphore
|
||||
}
|
||||
|
||||
func minimumInputLimit(endpoints []ActiveEndpoint) int {
|
||||
limit := DefaultInputLimit
|
||||
for index, endpoint := range endpoints {
|
||||
value := endpoint.InputLimit
|
||||
if value <= 0 {
|
||||
value = DefaultInputLimit
|
||||
}
|
||||
if index == 0 || value < limit {
|
||||
limit = value
|
||||
}
|
||||
}
|
||||
return limit
|
||||
}
|
||||
|
||||
func guardErrorCode(err error) string {
|
||||
var guardErr *GuardError
|
||||
if errors.As(err, &guardErr) && guardErr.Code != "" {
|
||||
return guardErr.Code
|
||||
}
|
||||
return ErrorCodeUnavailable
|
||||
}
|
||||
|
||||
func pointerLogID(value *int64) int64 {
|
||||
if value == nil {
|
||||
return 0
|
||||
}
|
||||
return *value
|
||||
}
|
||||
@@ -0,0 +1,276 @@
|
||||
package securityaudit
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
type scriptedScanner struct {
|
||||
mu sync.Mutex
|
||||
calls []string
|
||||
block <-chan struct{}
|
||||
entered chan<- struct{}
|
||||
}
|
||||
|
||||
func (s *scriptedScanner) Scan(ctx context.Context, endpoint ActiveEndpoint, _ string, _ []string) (*NormalizedResult, error) {
|
||||
s.mu.Lock()
|
||||
s.calls = append(s.calls, endpoint.ID)
|
||||
s.mu.Unlock()
|
||||
if s.entered != nil {
|
||||
select {
|
||||
case s.entered <- struct{}{}:
|
||||
default:
|
||||
}
|
||||
}
|
||||
if s.block != nil {
|
||||
select {
|
||||
case <-s.block:
|
||||
case <-ctx.Done():
|
||||
return nil, &GuardError{Code: ErrorCodeUnavailable, Retryable: true, Timeout: true, Cause: ctx.Err()}
|
||||
}
|
||||
}
|
||||
if endpoint.ID == "bad" {
|
||||
return nil, &GuardError{Code: ErrorCodeUnavailable, Retryable: true}
|
||||
}
|
||||
if endpoint.ID == "invalid" {
|
||||
return nil, &GuardError{Code: ErrorCodeInvalidResponse}
|
||||
}
|
||||
return &NormalizedResult{Decision: EventPass, RiskLevel: RiskLow, Action: ActionAllow, Safety: "Safe", ScannerScores: map[string]float64{}, ScannerEvidence: map[string]string{}, GuardEndpointID: endpoint.ID}, nil
|
||||
}
|
||||
|
||||
func guardConfig(endpoints ...ActiveEndpoint) ActiveConfig {
|
||||
return ActiveConfig{RiskControlEnabled: true, Enabled: true, BlockingEnabled: true, ConfigVersion: 2, Scanners: AllScannerIDs, Endpoints: endpoints}
|
||||
}
|
||||
|
||||
func TestGuardEvaluatorOrderedFailoverAndInvalidTerminal(t *testing.T) {
|
||||
scanner := &scriptedScanner{}
|
||||
metrics := NewAtomicMetrics()
|
||||
evaluator := newGuardEvaluator(scanner, nil, metrics, 4, 2)
|
||||
snapshot := PromptSnapshot{RequestID: "r", ScanText: "hello", PromptLength: 5}
|
||||
decision, err := evaluator.Evaluate(context.Background(), guardConfig(
|
||||
ActiveEndpoint{ID: "bad", Enabled: true, TimeoutMS: 1000, InputLimit: 100},
|
||||
ActiveEndpoint{ID: "good", Enabled: true, TimeoutMS: 1000, InputLimit: 100},
|
||||
), snapshot)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, DecisionAllow, decision.Kind)
|
||||
require.Equal(t, int64(1), metrics.Snapshot().Failovers)
|
||||
_, err = evaluator.Evaluate(context.Background(), guardConfig(
|
||||
ActiveEndpoint{ID: "invalid", Enabled: true, TimeoutMS: 1000, InputLimit: 100},
|
||||
ActiveEndpoint{ID: "good", Enabled: true, TimeoutMS: 1000, InputLimit: 100},
|
||||
), snapshot)
|
||||
var guardErr *GuardError
|
||||
require.ErrorAs(t, err, &guardErr)
|
||||
require.Equal(t, ErrorCodeInvalidResponse, guardErr.Code)
|
||||
snapshotMetrics := metrics.Snapshot()
|
||||
require.Equal(t, int64(2), snapshotMetrics.Total)
|
||||
require.Equal(t, int64(1), snapshotMetrics.Allowed)
|
||||
require.Equal(t, int64(1), snapshotMetrics.Invalid)
|
||||
}
|
||||
|
||||
func TestGuardEvaluatorGlobalBulkheadIsNonBlocking(t *testing.T) {
|
||||
release := make(chan struct{})
|
||||
entered := make(chan struct{}, 1)
|
||||
scanner := &scriptedScanner{block: release, entered: entered}
|
||||
metrics := NewAtomicMetrics()
|
||||
evaluator := newGuardEvaluator(scanner, nil, metrics, 1, 1)
|
||||
cfg := guardConfig(ActiveEndpoint{ID: "good", Enabled: true, TimeoutMS: 2000, InputLimit: 100})
|
||||
done := make(chan error, 1)
|
||||
go func() {
|
||||
_, err := evaluator.Evaluate(context.Background(), cfg, PromptSnapshot{ScanText: "one", PromptLength: 3})
|
||||
done <- err
|
||||
}()
|
||||
select {
|
||||
case <-entered:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("first evaluation did not enter scanner")
|
||||
}
|
||||
start := time.Now()
|
||||
_, err := evaluator.Evaluate(context.Background(), cfg, PromptSnapshot{ScanText: "two", PromptLength: 3})
|
||||
require.Error(t, err)
|
||||
require.Less(t, time.Since(start), 200*time.Millisecond)
|
||||
require.Equal(t, int64(1), metrics.Snapshot().BulkheadFull)
|
||||
close(release)
|
||||
require.NoError(t, <-done)
|
||||
snapshotMetrics := metrics.Snapshot()
|
||||
require.Equal(t, int64(2), snapshotMetrics.Total)
|
||||
require.Equal(t, int64(1), snapshotMetrics.Allowed)
|
||||
require.Equal(t, int64(1), snapshotMetrics.Unavailable)
|
||||
}
|
||||
|
||||
func TestGuardEvaluatorPerNodeBulkheadIsNonBlocking(t *testing.T) {
|
||||
release := make(chan struct{})
|
||||
entered := make(chan struct{}, 1)
|
||||
scanner := &scriptedScanner{block: release, entered: entered}
|
||||
metrics := NewAtomicMetrics()
|
||||
evaluator := newGuardEvaluator(scanner, nil, metrics, 2, 1)
|
||||
cfg := guardConfig(ActiveEndpoint{ID: "same-node", Enabled: true, TimeoutMS: 2000, InputLimit: 100})
|
||||
done := make(chan error, 1)
|
||||
go func() {
|
||||
_, err := evaluator.Evaluate(context.Background(), cfg, PromptSnapshot{ScanText: "one", PromptLength: 3})
|
||||
done <- err
|
||||
}()
|
||||
select {
|
||||
case <-entered:
|
||||
case <-time.After(time.Second):
|
||||
t.Fatal("first evaluation did not enter scanner")
|
||||
}
|
||||
started := time.Now()
|
||||
_, err := evaluator.Evaluate(context.Background(), cfg, PromptSnapshot{ScanText: "two", PromptLength: 3})
|
||||
require.Error(t, err)
|
||||
require.Less(t, time.Since(started), 200*time.Millisecond)
|
||||
require.GreaterOrEqual(t, metrics.Snapshot().BulkheadFull, int64(1))
|
||||
close(release)
|
||||
require.NoError(t, <-done)
|
||||
}
|
||||
|
||||
func TestGuardEvaluatorLastChunkFailureNeverAllows(t *testing.T) {
|
||||
call := 0
|
||||
scanner := PromptScannerFunc(func(context.Context, ActiveEndpoint, string, []string) (*NormalizedResult, error) {
|
||||
call++
|
||||
if call == 2 {
|
||||
return nil, &GuardError{Code: ErrorCodeUnavailable, Retryable: true, Cause: errors.New("down")}
|
||||
}
|
||||
return &NormalizedResult{Decision: EventPass, RiskLevel: RiskLow, Action: ActionAllow, ScannerScores: map[string]float64{}, ScannerEvidence: map[string]string{}}, nil
|
||||
})
|
||||
metrics := NewAtomicMetrics()
|
||||
evaluator := newGuardEvaluator(scanner, nil, metrics, 2, 2)
|
||||
_, err := evaluator.Evaluate(context.Background(), guardConfig(ActiveEndpoint{ID: "one", Enabled: true, TimeoutMS: 1000, InputLimit: 3}), PromptSnapshot{ScanText: "abcdef", PromptLength: 6})
|
||||
require.Error(t, err)
|
||||
}
|
||||
|
||||
func TestGuardEvaluatorBlockStopsRemainingChunksButReportsPlannedTotal(t *testing.T) {
|
||||
calls := 0
|
||||
scanner := PromptScannerFunc(func(context.Context, ActiveEndpoint, string, []string) (*NormalizedResult, error) {
|
||||
calls++
|
||||
return &NormalizedResult{
|
||||
Decision: EventCritical, RiskLevel: RiskCritical, Action: ActionBlock, Safety: "Unsafe",
|
||||
Categories: []string{"jailbreak"}, MatchedScanners: []string{"jailbreak"},
|
||||
ScannerScores: map[string]float64{"jailbreak": 1}, ScannerEvidence: map[string]string{"jailbreak": "Jailbreak"},
|
||||
}, nil
|
||||
})
|
||||
metrics := NewAtomicMetrics()
|
||||
evaluator := newGuardEvaluator(scanner, nil, metrics, 2, 2)
|
||||
decision, err := evaluator.Evaluate(context.Background(), guardConfig(
|
||||
ActiveEndpoint{ID: "one", Enabled: true, TimeoutMS: 1000, InputLimit: 3},
|
||||
), PromptSnapshot{ScanText: "abcdefghi", PromptLength: 9})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, DecisionBlock, decision.Kind)
|
||||
require.Equal(t, 1, calls)
|
||||
require.Equal(t, 3, decision.Result.ChunkTotal)
|
||||
require.Equal(t, int64(1), metrics.Snapshot().Blocked)
|
||||
}
|
||||
|
||||
func TestGuardEvaluatorFlagSharedDeadlineFailClosedAndContextCancel(t *testing.T) {
|
||||
t.Run("flag allows next stage", func(t *testing.T) {
|
||||
metrics := NewAtomicMetrics()
|
||||
evaluator := newGuardEvaluator(PromptScannerFunc(func(context.Context, ActiveEndpoint, string, []string) (*NormalizedResult, error) {
|
||||
return &NormalizedResult{Decision: EventFlag, RiskLevel: RiskMedium, Action: ActionWarn, Safety: "Controversial", Categories: []string{"violent"}, MatchedScanners: []string{"violent"}, ScannerScores: map[string]float64{"violent": .5}, ScannerEvidence: map[string]string{"violent": "Violent"}}, nil
|
||||
}), nil, metrics, 2, 2)
|
||||
decision, err := evaluator.Evaluate(context.Background(), guardConfig(ActiveEndpoint{ID: "one", Enabled: true, TimeoutMS: 1000, InputLimit: 100}), PromptSnapshot{ScanText: "review", PromptLength: 6})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, DecisionFlag, decision.Kind)
|
||||
require.True(t, decision.AllowNextStage)
|
||||
require.Equal(t, int64(1), metrics.Snapshot().Flagged)
|
||||
})
|
||||
|
||||
t.Run("all failovers share first endpoint deadline", func(t *testing.T) {
|
||||
calls := 0
|
||||
scanner := PromptScannerFunc(func(ctx context.Context, endpoint ActiveEndpoint, _ string, _ []string) (*NormalizedResult, error) {
|
||||
calls++
|
||||
if endpoint.ID == "first" {
|
||||
select {
|
||||
case <-time.After(35 * time.Millisecond):
|
||||
return nil, &GuardError{Code: ErrorCodeUnavailable, Retryable: true}
|
||||
case <-ctx.Done():
|
||||
return nil, &GuardError{Code: ErrorCodeUnavailable, Retryable: true, Timeout: true, Cause: ctx.Err()}
|
||||
}
|
||||
}
|
||||
<-ctx.Done()
|
||||
return nil, &GuardError{Code: ErrorCodeUnavailable, Retryable: true, Timeout: true, Cause: ctx.Err()}
|
||||
})
|
||||
metrics := NewAtomicMetrics()
|
||||
evaluator := newGuardEvaluator(scanner, nil, metrics, 2, 2)
|
||||
started := time.Now()
|
||||
_, err := evaluator.Evaluate(context.Background(), guardConfig(
|
||||
ActiveEndpoint{ID: "first", Enabled: true, TimeoutMS: 70, InputLimit: 100},
|
||||
ActiveEndpoint{ID: "second", Enabled: true, TimeoutMS: 500, InputLimit: 100},
|
||||
), PromptSnapshot{ScanText: "deadline", PromptLength: 8})
|
||||
elapsed := time.Since(started)
|
||||
require.Error(t, err)
|
||||
require.Equal(t, 2, calls)
|
||||
require.Less(t, elapsed, 180*time.Millisecond)
|
||||
require.GreaterOrEqual(t, elapsed, 50*time.Millisecond)
|
||||
require.Equal(t, int64(1), metrics.Snapshot().Failovers)
|
||||
require.Equal(t, int64(1), metrics.Snapshot().Timeouts)
|
||||
})
|
||||
|
||||
t.Run("canceled parent never allows", func(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
evaluator := newGuardEvaluator(PromptScannerFunc(func(ctx context.Context, _ ActiveEndpoint, _ string, _ []string) (*NormalizedResult, error) {
|
||||
<-ctx.Done()
|
||||
return nil, &GuardError{Code: ErrorCodeUnavailable, Retryable: true, Cause: ctx.Err()}
|
||||
}), nil, NewAtomicMetrics(), 2, 2)
|
||||
decision, err := evaluator.Evaluate(ctx, guardConfig(ActiveEndpoint{ID: "one", Enabled: true, TimeoutMS: 1000, InputLimit: 100}), PromptSnapshot{ScanText: "cancel", PromptLength: 6})
|
||||
require.Error(t, err)
|
||||
require.Nil(t, decision)
|
||||
})
|
||||
}
|
||||
|
||||
func TestGuardEvaluatorRecordsExistingResultOnceAndRecordFailureDoesNotChangeDecision(t *testing.T) {
|
||||
for _, recordErr := range []error{nil, errors.New("database unavailable")} {
|
||||
repo := &fakeJobRepository{recordBlockingErr: recordErr}
|
||||
metrics := NewAtomicMetrics()
|
||||
scannerCalls := 0
|
||||
evaluator := newGuardEvaluator(PromptScannerFunc(func(context.Context, ActiveEndpoint, string, []string) (*NormalizedResult, error) {
|
||||
scannerCalls++
|
||||
return &NormalizedResult{Decision: EventCritical, RiskLevel: RiskCritical, Action: ActionBlock, Safety: "Unsafe", Categories: []string{"pii"}, MatchedScanners: []string{"pii"}, ScannerScores: map[string]float64{"pii": 1}, ScannerEvidence: map[string]string{"pii": "PII"}}, nil
|
||||
}), repo, metrics, 2, 2)
|
||||
decision, err := evaluator.Evaluate(context.Background(), guardConfig(ActiveEndpoint{ID: "one", Enabled: true, TimeoutMS: 1000, InputLimit: 100}), PromptSnapshot{ScanText: "raw prompt", RedactedPreview: "raw***", PromptLength: 10})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, DecisionBlock, decision.Kind)
|
||||
require.Equal(t, 1, scannerCalls)
|
||||
require.Equal(t, 1, repo.recordBlockingCalls)
|
||||
require.Empty(t, repo.recordBlockingSnapshot.ScanText)
|
||||
require.Same(t, decision.Result, repo.recordBlockingResult)
|
||||
if recordErr != nil {
|
||||
require.Equal(t, int64(1), metrics.Snapshot().RecordFailed)
|
||||
} else {
|
||||
require.Zero(t, metrics.Snapshot().RecordFailed)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestGuardEvaluatorNilResultAndScannerPanicBecomeStableFailures(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
scan PromptScannerFunc
|
||||
code string
|
||||
}{
|
||||
{name: "nil result", scan: func(context.Context, ActiveEndpoint, string, []string) (*NormalizedResult, error) { return nil, nil }, code: ErrorCodeInvalidResponse},
|
||||
{name: "panic", scan: func(context.Context, ActiveEndpoint, string, []string) (*NormalizedResult, error) {
|
||||
panic("raw prompt canary")
|
||||
}, code: ErrorCodeUnavailable},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
evaluator := newGuardEvaluator(tt.scan, nil, NewAtomicMetrics(), 2, 2)
|
||||
_, err := evaluator.Evaluate(context.Background(), guardConfig(ActiveEndpoint{ID: "one", Enabled: true, TimeoutMS: 1000, InputLimit: 100}), PromptSnapshot{ScanText: "input", PromptLength: 5})
|
||||
var guardErr *GuardError
|
||||
require.ErrorAs(t, err, &guardErr)
|
||||
require.Equal(t, tt.code, guardErr.Code)
|
||||
require.NotContains(t, err.Error(), "canary")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
type PromptScannerFunc func(context.Context, ActiveEndpoint, string, []string) (*NormalizedResult, error)
|
||||
|
||||
func (f PromptScannerFunc) Scan(ctx context.Context, endpoint ActiveEndpoint, chunk string, scanners []string) (*NormalizedResult, error) {
|
||||
return f(ctx, endpoint, chunk, scanners)
|
||||
}
|
||||
@@ -0,0 +1,305 @@
|
||||
package securityaudit
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
|
||||
"github.com/Wei-Shaw/sub2api/internal/pkg/response"
|
||||
"github.com/Wei-Shaw/sub2api/internal/server/middleware"
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type PromptAdminService interface {
|
||||
GetConfig() PublicConfig
|
||||
SaveConfig(context.Context, UpdateConfigRequest, int64) (PublicConfig, error)
|
||||
Probe(context.Context, ProbeRequest) ProbeResult
|
||||
Runtime(context.Context) RuntimeSnapshot
|
||||
ListEvents(context.Context, EventFilter, int, int) (*EventPage, error)
|
||||
GetEvent(context.Context, int64) (*Event, error)
|
||||
DeleteEvent(context.Context, int64) (*DeleteResult, error)
|
||||
DeleteEventsByIDs(context.Context, []int64) (*DeleteResult, error)
|
||||
PreviewDelete(context.Context, EventFilter, int64) (*DeletePreview, error)
|
||||
DeleteByFilter(context.Context, DeleteByFilterRequest, int64) (*DeleteResult, error)
|
||||
}
|
||||
|
||||
type PromptAdminHandler struct{ service PromptAdminService }
|
||||
|
||||
func NewPromptAdminHandler(service PromptAdminService) *PromptAdminHandler {
|
||||
return &PromptAdminHandler{service: service}
|
||||
}
|
||||
|
||||
func (h *PromptAdminHandler) GetConfig(c *gin.Context) { response.Success(c, h.service.GetConfig()) }
|
||||
|
||||
func (h *PromptAdminHandler) UpdateConfig(c *gin.Context) {
|
||||
var request UpdateConfigRequest
|
||||
if err := c.ShouldBindJSON(&request); err != nil {
|
||||
setPromptAdminAudit(c, "failed", "prompt_audit_invalid_config_request", nil)
|
||||
response.ErrorFrom(c, infraerrors.BadRequest("prompt_audit_invalid_config_request", "提示词审计配置请求无效"))
|
||||
return
|
||||
}
|
||||
config, err := h.service.SaveConfig(c.Request.Context(), request, adminID(c))
|
||||
if err != nil {
|
||||
setPromptAdminAudit(c, "failed", infraerrors.Reason(err), configAuditFields(request, nil))
|
||||
response.ErrorFrom(c, err)
|
||||
return
|
||||
}
|
||||
setPromptAdminAudit(c, "success", "", configAuditFields(request, &config))
|
||||
response.Success(c, config)
|
||||
}
|
||||
|
||||
func (h *PromptAdminHandler) ProbeEndpoint(c *gin.Context) {
|
||||
var request ProbeRequest
|
||||
if err := c.ShouldBindJSON(&request); err != nil {
|
||||
setPromptAdminAudit(c, "failed", "prompt_audit_invalid_probe_request", nil)
|
||||
response.ErrorFrom(c, infraerrors.BadRequest("prompt_audit_invalid_probe_request", "审计节点探测请求无效"))
|
||||
return
|
||||
}
|
||||
result := h.service.Probe(c.Request.Context(), request)
|
||||
status := "failed"
|
||||
if result.OK {
|
||||
status = "success"
|
||||
}
|
||||
setPromptAdminAudit(c, status, result.ErrorCode, map[string]any{
|
||||
"guard_endpoint_id": request.Endpoint.ID, "http_status": result.HTTPStatus,
|
||||
"latency_ms": result.LatencyMS, "token_applied": result.TokenApplied, "retryable": result.Retryable,
|
||||
})
|
||||
response.Success(c, result)
|
||||
}
|
||||
|
||||
func (h *PromptAdminHandler) GetRuntime(c *gin.Context) {
|
||||
response.Success(c, h.service.Runtime(c.Request.Context()))
|
||||
}
|
||||
|
||||
func (h *PromptAdminHandler) ListEvents(c *gin.Context) {
|
||||
page, err := positiveIntQuery(c, "page", 1, 0)
|
||||
if err != nil {
|
||||
response.ErrorFrom(c, err)
|
||||
return
|
||||
}
|
||||
pageSize, err := positiveIntQuery(c, "page_size", 20, 100)
|
||||
if err != nil {
|
||||
response.ErrorFrom(c, err)
|
||||
return
|
||||
}
|
||||
filter, err := eventFilterFromQuery(c)
|
||||
if err != nil {
|
||||
response.ErrorFrom(c, err)
|
||||
return
|
||||
}
|
||||
result, err := h.service.ListEvents(c.Request.Context(), filter, page, pageSize)
|
||||
if err != nil {
|
||||
response.ErrorFrom(c, err)
|
||||
return
|
||||
}
|
||||
response.Success(c, result)
|
||||
}
|
||||
|
||||
func (h *PromptAdminHandler) GetEvent(c *gin.Context) {
|
||||
id, err := strconv.ParseInt(c.Param("id"), 10, 64)
|
||||
if err != nil || id <= 0 {
|
||||
response.ErrorFrom(c, infraerrors.BadRequest("prompt_audit_invalid_event_id", "事件 ID 无效"))
|
||||
return
|
||||
}
|
||||
event, err := h.service.GetEvent(c.Request.Context(), id)
|
||||
if errors.Is(err, ErrEventNotFound) {
|
||||
response.ErrorFrom(c, infraerrors.NotFound("prompt_audit_event_not_found", "提示词审计事件不存在"))
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
response.ErrorFrom(c, err)
|
||||
return
|
||||
}
|
||||
response.Success(c, event)
|
||||
}
|
||||
|
||||
func (h *PromptAdminHandler) DeleteEvent(c *gin.Context) {
|
||||
id, err := strconv.ParseInt(c.Param("id"), 10, 64)
|
||||
if err != nil || id <= 0 {
|
||||
setPromptAdminAudit(c, "failed", "prompt_audit_invalid_event_id", nil)
|
||||
response.ErrorFrom(c, infraerrors.BadRequest("prompt_audit_invalid_event_id", "事件 ID 无效"))
|
||||
return
|
||||
}
|
||||
result, err := h.service.DeleteEvent(c.Request.Context(), id)
|
||||
if err != nil {
|
||||
setPromptAdminAudit(c, "failed", infraerrors.Reason(err), map[string]any{"event_id": id})
|
||||
response.ErrorFrom(c, err)
|
||||
return
|
||||
}
|
||||
setPromptAdminAudit(c, "success", "", deleteAuditFields(result, map[string]any{"event_id": id}))
|
||||
LogWarn(EventEventDeleted, map[string]any{"user_id": adminID(c), "event_id": id, "status": "deleted"})
|
||||
response.Success(c, result)
|
||||
}
|
||||
|
||||
type batchDeleteRequest struct {
|
||||
IDs []int64 `json:"ids" binding:"required"`
|
||||
}
|
||||
|
||||
func (h *PromptAdminHandler) BatchDelete(c *gin.Context) {
|
||||
var request batchDeleteRequest
|
||||
if err := c.ShouldBindJSON(&request); err != nil || len(request.IDs) == 0 || len(request.IDs) > 500 {
|
||||
setPromptAdminAudit(c, "failed", "prompt_audit_invalid_delete_batch", nil)
|
||||
response.ErrorFrom(c, infraerrors.BadRequest("prompt_audit_invalid_delete_batch", "批量删除必须包含 1-500 个事件 ID"))
|
||||
return
|
||||
}
|
||||
for _, id := range request.IDs {
|
||||
if id <= 0 {
|
||||
setPromptAdminAudit(c, "failed", "prompt_audit_invalid_event_id", map[string]any{"requested_count": len(request.IDs)})
|
||||
response.ErrorFrom(c, infraerrors.BadRequest("prompt_audit_invalid_event_id", "事件 ID 无效"))
|
||||
return
|
||||
}
|
||||
}
|
||||
result, err := h.service.DeleteEventsByIDs(c.Request.Context(), request.IDs)
|
||||
if err != nil {
|
||||
setPromptAdminAudit(c, "failed", infraerrors.Reason(err), map[string]any{"requested_count": len(request.IDs)})
|
||||
response.ErrorFrom(c, err)
|
||||
return
|
||||
}
|
||||
setPromptAdminAudit(c, "success", "", deleteAuditFields(result, map[string]any{"requested_count": len(request.IDs)}))
|
||||
LogWarn(EventEventsDeleted, map[string]any{"user_id": adminID(c), "status": "deleted"})
|
||||
response.Success(c, result)
|
||||
}
|
||||
|
||||
func (h *PromptAdminHandler) DeletePreview(c *gin.Context) {
|
||||
var filter EventFilter
|
||||
if err := c.ShouldBindJSON(&filter); err != nil {
|
||||
setPromptAdminAudit(c, "failed", "prompt_audit_delete_preview_invalid", nil)
|
||||
response.ErrorFrom(c, infraerrors.BadRequest("prompt_audit_delete_preview_invalid", "删除预览筛选无效"))
|
||||
return
|
||||
}
|
||||
preview, err := h.service.PreviewDelete(c.Request.Context(), filter, adminID(c))
|
||||
if err != nil {
|
||||
setPromptAdminAudit(c, "failed", "prompt_audit_delete_preview_invalid", nil)
|
||||
response.ErrorFrom(c, infraerrors.BadRequest("prompt_audit_delete_preview_invalid", "删除预览筛选无效"))
|
||||
return
|
||||
}
|
||||
setPromptAdminAudit(c, "success", "", map[string]any{
|
||||
"matched_count": preview.MatchedCount, "snapshot_max_id": preview.SnapshotMaxID, "filter_hash": preview.FilterHash,
|
||||
})
|
||||
response.Success(c, preview)
|
||||
}
|
||||
|
||||
func (h *PromptAdminHandler) DeleteByFilter(c *gin.Context) {
|
||||
var request DeleteByFilterRequest
|
||||
if err := c.ShouldBindJSON(&request); err != nil {
|
||||
setPromptAdminAudit(c, "failed", "prompt_audit_delete_confirmation_invalid", nil)
|
||||
response.ErrorFrom(c, infraerrors.BadRequest("prompt_audit_delete_confirmation_invalid", "删除确认无效或已过期"))
|
||||
return
|
||||
}
|
||||
result, err := h.service.DeleteByFilter(c.Request.Context(), request, adminID(c))
|
||||
if err != nil {
|
||||
setPromptAdminAudit(c, "failed", "prompt_audit_delete_confirmation_invalid", map[string]any{
|
||||
"snapshot_max_id": request.SnapshotMaxID, "filter_hash": request.FilterHash, "confirm": request.Confirm,
|
||||
})
|
||||
response.ErrorFrom(c, infraerrors.BadRequest("prompt_audit_delete_confirmation_invalid", "删除确认无效或已过期"))
|
||||
return
|
||||
}
|
||||
setPromptAdminAudit(c, "success", "", deleteAuditFields(result, map[string]any{
|
||||
"snapshot_max_id": request.SnapshotMaxID, "filter_hash": request.FilterHash, "confirm": request.Confirm,
|
||||
}))
|
||||
response.Success(c, result)
|
||||
}
|
||||
|
||||
func setPromptAdminAudit(c *gin.Context, result, errorCode string, fields map[string]any) {
|
||||
details := make(map[string]any, len(fields)+2)
|
||||
details["result"] = result
|
||||
if strings.TrimSpace(errorCode) != "" {
|
||||
details["error_code"] = errorCode
|
||||
}
|
||||
for key, value := range fields {
|
||||
details[key] = value
|
||||
}
|
||||
middleware.SetAuditExtra(c, details)
|
||||
}
|
||||
|
||||
func configAuditFields(request UpdateConfigRequest, saved *PublicConfig) map[string]any {
|
||||
version := request.ExpectedConfigVersion
|
||||
if saved != nil {
|
||||
version = saved.ConfigVersion
|
||||
}
|
||||
return map[string]any{
|
||||
"enabled": request.Enabled, "blocking_enabled": request.BlockingEnabled,
|
||||
"config_version": version, "endpoint_count": len(request.Endpoints),
|
||||
"scanner_count": len(request.Scanners), "all_groups": request.AllGroups,
|
||||
"group_count": len(request.GroupIDs),
|
||||
}
|
||||
}
|
||||
|
||||
func deleteAuditFields(result *DeleteResult, base map[string]any) map[string]any {
|
||||
fields := make(map[string]any, len(base)+2)
|
||||
for key, value := range base {
|
||||
fields[key] = value
|
||||
}
|
||||
if result != nil {
|
||||
fields["deleted_events"] = result.DeletedEvents
|
||||
fields["deleted_jobs"] = result.DeletedJobs
|
||||
}
|
||||
return fields
|
||||
}
|
||||
|
||||
func adminID(c *gin.Context) int64 {
|
||||
subject, ok := middleware.GetAuthSubjectFromContext(c)
|
||||
if !ok {
|
||||
return 0
|
||||
}
|
||||
return subject.UserID
|
||||
}
|
||||
|
||||
func eventFilterFromQuery(c *gin.Context) (EventFilter, error) {
|
||||
groupID, err := optionalPositiveInt64Query(c, "group_id")
|
||||
if err != nil {
|
||||
return EventFilter{}, err
|
||||
}
|
||||
userID, err := optionalPositiveInt64Query(c, "user_id")
|
||||
if err != nil {
|
||||
return EventFilter{}, err
|
||||
}
|
||||
apiKeyID, err := optionalPositiveInt64Query(c, "api_key_id")
|
||||
if err != nil {
|
||||
return EventFilter{}, err
|
||||
}
|
||||
filter := EventFilter{
|
||||
Decision: c.Query("decision"), RiskLevel: c.Query("risk_level"), Endpoint: c.Query("endpoint"),
|
||||
GroupID: groupID, UserID: userID, APIKeyID: apiKeyID, RequestID: c.Query("request_id"),
|
||||
PromptHash: c.Query("prompt_hash"), Keyword: c.Query("keyword"),
|
||||
}
|
||||
if value := strings.TrimSpace(c.Query("start_at")); value != "" {
|
||||
filter.StartAt = parseTimeQuery(value)
|
||||
if filter.StartAt == nil {
|
||||
return EventFilter{}, infraerrors.BadRequest("prompt_audit_invalid_time", "开始时间无效")
|
||||
}
|
||||
}
|
||||
if value := strings.TrimSpace(c.Query("end_at")); value != "" {
|
||||
filter.EndAt = parseTimeQuery(value)
|
||||
if filter.EndAt == nil {
|
||||
return EventFilter{}, infraerrors.BadRequest("prompt_audit_invalid_time", "结束时间无效")
|
||||
}
|
||||
}
|
||||
return filter, nil
|
||||
}
|
||||
|
||||
func optionalPositiveInt64Query(c *gin.Context, key string) (*int64, error) {
|
||||
value := strings.TrimSpace(c.Query(key))
|
||||
if value == "" {
|
||||
return nil, nil
|
||||
}
|
||||
parsed, err := strconv.ParseInt(value, 10, 64)
|
||||
if err != nil || parsed <= 0 {
|
||||
return nil, infraerrors.BadRequest("prompt_audit_invalid_filter_id", "事件筛选 ID 无效")
|
||||
}
|
||||
return &parsed, nil
|
||||
}
|
||||
|
||||
func positiveIntQuery(c *gin.Context, key string, defaultValue, maxValue int) (int, error) {
|
||||
value := strings.TrimSpace(c.Query(key))
|
||||
if value == "" {
|
||||
return defaultValue, nil
|
||||
}
|
||||
parsed, err := strconv.Atoi(value)
|
||||
if err != nil || parsed <= 0 || (maxValue > 0 && parsed > maxValue) {
|
||||
return 0, infraerrors.BadRequest("prompt_audit_invalid_pagination", "分页参数无效")
|
||||
}
|
||||
return parsed, nil
|
||||
}
|
||||
@@ -0,0 +1,232 @@
|
||||
package securityaudit
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
|
||||
servermiddleware "github.com/Wei-Shaw/sub2api/internal/server/middleware"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
type fakePromptAdminService struct {
|
||||
config PublicConfig
|
||||
save func(context.Context, UpdateConfigRequest, int64) (PublicConfig, error)
|
||||
probe func(context.Context, ProbeRequest) ProbeResult
|
||||
runtime RuntimeSnapshot
|
||||
list func(context.Context, EventFilter, int, int) (*EventPage, error)
|
||||
get func(context.Context, int64) (*Event, error)
|
||||
deleteOne func(context.Context, int64) (*DeleteResult, error)
|
||||
deleteIDs func(context.Context, []int64) (*DeleteResult, error)
|
||||
preview func(context.Context, EventFilter, int64) (*DeletePreview, error)
|
||||
deleteFilter func(context.Context, DeleteByFilterRequest, int64) (*DeleteResult, error)
|
||||
}
|
||||
|
||||
func (s *fakePromptAdminService) GetConfig() PublicConfig { return s.config }
|
||||
func (s *fakePromptAdminService) SaveConfig(ctx context.Context, req UpdateConfigRequest, actorID int64) (PublicConfig, error) {
|
||||
if s.save == nil {
|
||||
return PublicConfig{}, errors.New("unexpected SaveConfig call")
|
||||
}
|
||||
return s.save(ctx, req, actorID)
|
||||
}
|
||||
func (s *fakePromptAdminService) Probe(ctx context.Context, req ProbeRequest) ProbeResult {
|
||||
if s.probe == nil {
|
||||
return ProbeResult{}
|
||||
}
|
||||
return s.probe(ctx, req)
|
||||
}
|
||||
func (s *fakePromptAdminService) Runtime(context.Context) RuntimeSnapshot { return s.runtime }
|
||||
func (s *fakePromptAdminService) ListEvents(ctx context.Context, filter EventFilter, page, pageSize int) (*EventPage, error) {
|
||||
if s.list == nil {
|
||||
return &EventPage{}, nil
|
||||
}
|
||||
return s.list(ctx, filter, page, pageSize)
|
||||
}
|
||||
func (s *fakePromptAdminService) GetEvent(ctx context.Context, id int64) (*Event, error) {
|
||||
if s.get == nil {
|
||||
return nil, ErrEventNotFound
|
||||
}
|
||||
return s.get(ctx, id)
|
||||
}
|
||||
func (s *fakePromptAdminService) DeleteEvent(ctx context.Context, id int64) (*DeleteResult, error) {
|
||||
if s.deleteOne == nil {
|
||||
return &DeleteResult{}, nil
|
||||
}
|
||||
return s.deleteOne(ctx, id)
|
||||
}
|
||||
func (s *fakePromptAdminService) DeleteEventsByIDs(ctx context.Context, ids []int64) (*DeleteResult, error) {
|
||||
if s.deleteIDs == nil {
|
||||
return &DeleteResult{}, nil
|
||||
}
|
||||
return s.deleteIDs(ctx, ids)
|
||||
}
|
||||
func (s *fakePromptAdminService) PreviewDelete(ctx context.Context, filter EventFilter, actorID int64) (*DeletePreview, error) {
|
||||
if s.preview == nil {
|
||||
return &DeletePreview{}, nil
|
||||
}
|
||||
return s.preview(ctx, filter, actorID)
|
||||
}
|
||||
func (s *fakePromptAdminService) DeleteByFilter(ctx context.Context, req DeleteByFilterRequest, actorID int64) (*DeleteResult, error) {
|
||||
if s.deleteFilter == nil {
|
||||
return &DeleteResult{}, nil
|
||||
}
|
||||
return s.deleteFilter(ctx, req, actorID)
|
||||
}
|
||||
|
||||
func promptAdminRouter(service PromptAdminService) *gin.Engine {
|
||||
gin.SetMode(gin.TestMode)
|
||||
router := gin.New()
|
||||
router.Use(func(c *gin.Context) {
|
||||
c.Set(string(servermiddleware.ContextKeyUser), servermiddleware.AuthSubject{UserID: 42})
|
||||
c.Set(string(servermiddleware.ContextKeyUserRole), "admin")
|
||||
c.Next()
|
||||
})
|
||||
handler := NewPromptAdminHandler(service)
|
||||
group := router.Group("/admin/prompt-audit")
|
||||
group.GET("/config", handler.GetConfig)
|
||||
group.PUT("/config", handler.UpdateConfig)
|
||||
group.POST("/endpoints/probe", handler.ProbeEndpoint)
|
||||
group.GET("/runtime", handler.GetRuntime)
|
||||
group.GET("/events", handler.ListEvents)
|
||||
group.GET("/events/:id", handler.GetEvent)
|
||||
group.DELETE("/events/:id", handler.DeleteEvent)
|
||||
group.POST("/events/batch-delete", handler.BatchDelete)
|
||||
group.POST("/events/delete-preview", handler.DeletePreview)
|
||||
group.POST("/events/delete-by-filter", handler.DeleteByFilter)
|
||||
return router
|
||||
}
|
||||
|
||||
func promptAdminRequest(t *testing.T, router http.Handler, method, path string, body any) *httptest.ResponseRecorder {
|
||||
t.Helper()
|
||||
var reader *bytes.Reader
|
||||
if body == nil {
|
||||
reader = bytes.NewReader(nil)
|
||||
} else {
|
||||
raw, err := json.Marshal(body)
|
||||
require.NoError(t, err)
|
||||
reader = bytes.NewReader(raw)
|
||||
}
|
||||
req := httptest.NewRequest(method, path, reader)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
recorder := httptest.NewRecorder()
|
||||
router.ServeHTTP(recorder, req)
|
||||
return recorder
|
||||
}
|
||||
|
||||
func TestPromptAdminConfigRequiresVersionMapsConflictAndNeverEchoesToken(t *testing.T) {
|
||||
const canary = "prompt-admin-token-canary"
|
||||
|
||||
t.Run("missing expected version", func(t *testing.T) {
|
||||
router := promptAdminRouter(&fakePromptAdminService{})
|
||||
response := promptAdminRequest(t, router, http.MethodPut, "/admin/prompt-audit/config", map[string]any{})
|
||||
require.Equal(t, http.StatusBadRequest, response.Code)
|
||||
require.Contains(t, response.Body.String(), "prompt_audit_invalid_config_request")
|
||||
})
|
||||
|
||||
t.Run("CAS conflict", func(t *testing.T) {
|
||||
service := &fakePromptAdminService{save: func(context.Context, UpdateConfigRequest, int64) (PublicConfig, error) {
|
||||
return PublicConfig{}, infraerrors.Conflict(ErrorCodeConfigConflict, "配置已被更新")
|
||||
}}
|
||||
response := promptAdminRequest(t, promptAdminRouter(service), http.MethodPut, "/admin/prompt-audit/config", validHandlerUpdateRequest(canary))
|
||||
require.Equal(t, http.StatusConflict, response.Code)
|
||||
require.Contains(t, response.Body.String(), ErrorCodeConfigConflict)
|
||||
require.NotContains(t, response.Body.String(), canary)
|
||||
})
|
||||
|
||||
t.Run("success public DTO", func(t *testing.T) {
|
||||
service := &fakePromptAdminService{save: func(_ context.Context, req UpdateConfigRequest, actorID int64) (PublicConfig, error) {
|
||||
require.Equal(t, int64(42), actorID)
|
||||
require.Equal(t, canary, req.Endpoints[0].Token)
|
||||
return PublicConfig{ConfigVersion: 8, Endpoints: []PublicEndpoint{{ID: "guard-1", HasToken: true, TokenStatus: "configured"}}}, nil
|
||||
}}
|
||||
response := promptAdminRequest(t, promptAdminRouter(service), http.MethodPut, "/admin/prompt-audit/config", validHandlerUpdateRequest(canary))
|
||||
require.Equal(t, http.StatusOK, response.Code)
|
||||
body := response.Body.String()
|
||||
require.NotContains(t, body, canary)
|
||||
require.NotContains(t, body, "token_ciphertext")
|
||||
require.NotContains(t, body, `"token":`)
|
||||
require.Contains(t, body, `"has_token":true`)
|
||||
})
|
||||
}
|
||||
|
||||
func TestPromptAdminProbeSupportsTemporaryOrSavedTokenWithoutEcho(t *testing.T) {
|
||||
const canary = "probe-token-canary"
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
token string
|
||||
tokenApplied bool
|
||||
}{
|
||||
{name: "temporary token", token: canary, tokenApplied: true},
|
||||
{name: "saved token", token: "", tokenApplied: true},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
service := &fakePromptAdminService{probe: func(_ context.Context, req ProbeRequest) ProbeResult {
|
||||
require.Equal(t, tc.token, req.Endpoint.Token)
|
||||
return ProbeResult{OK: true, Status: "healthy", Message: "ok", TokenApplied: tc.tokenApplied}
|
||||
}}
|
||||
endpoint := validHandlerUpdateRequest(tc.token).Endpoints[0]
|
||||
response := promptAdminRequest(t, promptAdminRouter(service), http.MethodPost, "/admin/prompt-audit/endpoints/probe", ProbeRequest{Endpoint: endpoint})
|
||||
require.Equal(t, http.StatusOK, response.Code)
|
||||
require.NotContains(t, response.Body.String(), canary)
|
||||
require.NotContains(t, response.Body.String(), `"token":`)
|
||||
require.Contains(t, response.Body.String(), `"token_applied":true`)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPromptAdminRejectsInvalidEventIDsTimesAndPagination(t *testing.T) {
|
||||
router := promptAdminRouter(&fakePromptAdminService{})
|
||||
for _, tc := range []struct {
|
||||
method string
|
||||
path string
|
||||
body any
|
||||
reason string
|
||||
}{
|
||||
{http.MethodGet, "/admin/prompt-audit/events/not-a-number", nil, "prompt_audit_invalid_event_id"},
|
||||
{http.MethodDelete, "/admin/prompt-audit/events/-1", nil, "prompt_audit_invalid_event_id"},
|
||||
{http.MethodGet, "/admin/prompt-audit/events?group_id=bad", nil, "prompt_audit_invalid_filter_id"},
|
||||
{http.MethodGet, "/admin/prompt-audit/events?start_at=not-time", nil, "prompt_audit_invalid_time"},
|
||||
{http.MethodGet, "/admin/prompt-audit/events?page=0", nil, "prompt_audit_invalid_pagination"},
|
||||
{http.MethodPost, "/admin/prompt-audit/events/batch-delete", map[string]any{"ids": []int64{1, -2}}, "prompt_audit_invalid_event_id"},
|
||||
} {
|
||||
response := promptAdminRequest(t, router, tc.method, tc.path, tc.body)
|
||||
require.Equalf(t, http.StatusBadRequest, response.Code, "%s %s", tc.method, tc.path)
|
||||
require.Contains(t, response.Body.String(), tc.reason)
|
||||
}
|
||||
}
|
||||
|
||||
func validHandlerUpdateRequest(token string) UpdateConfigRequest {
|
||||
return UpdateConfigRequest{
|
||||
ExpectedConfigVersion: 7,
|
||||
Strategy: "priority",
|
||||
WorkerCount: 1,
|
||||
QueueCapacity: 10,
|
||||
Scanners: []string{"pii"},
|
||||
AllGroups: true,
|
||||
Endpoints: []UpdateEndpoint{{
|
||||
ID: "guard-1", Name: "Guard One", Protocol: "openai_compatible",
|
||||
BaseURL: "http://127.0.0.1:18080", Model: DefaultGuardModel, Token: token,
|
||||
TimeoutMS: 1000, InputLimit: 1024, Enabled: true,
|
||||
}},
|
||||
}
|
||||
}
|
||||
|
||||
func TestPromptAdminDeleteConfirmationErrorsStayGeneric(t *testing.T) {
|
||||
service := &fakePromptAdminService{deleteFilter: func(context.Context, DeleteByFilterRequest, int64) (*DeleteResult, error) {
|
||||
return nil, errors.New("sensitive-token-or-filter-detail")
|
||||
}}
|
||||
response := promptAdminRequest(t, promptAdminRouter(service), http.MethodPost, "/admin/prompt-audit/events/delete-by-filter", DeleteByFilterRequest{
|
||||
SnapshotMaxID: 3, FilterHash: strings.Repeat("a", 64), ConfirmationToken: "secret-confirmation", Confirm: true,
|
||||
})
|
||||
require.Equal(t, http.StatusBadRequest, response.Code)
|
||||
require.Contains(t, response.Body.String(), "prompt_audit_delete_confirmation_invalid")
|
||||
require.NotContains(t, response.Body.String(), "sensitive-token")
|
||||
require.NotContains(t, response.Body.String(), "secret-confirmation")
|
||||
}
|
||||
@@ -0,0 +1,69 @@
|
||||
package securityaudit
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
)
|
||||
|
||||
func BuildIssueSummaries(result NormalizedResult) []IssueSummary {
|
||||
resultCategories := result.Categories
|
||||
if len(resultCategories) == 0 {
|
||||
resultCategories = result.MatchedScanners
|
||||
}
|
||||
summaries := make([]IssueSummary, 0, len(resultCategories)+len(result.UnknownCategories))
|
||||
for _, category := range resultCategories {
|
||||
definition, ok := ScannerCatalog[category]
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
evidence := RedactPreview(result.ScannerEvidence[category], 160)
|
||||
if evidence == "" {
|
||||
evidence = definition.Label
|
||||
}
|
||||
digest := sha256.Sum256([]byte(evidence))
|
||||
summaries = append(summaries, IssueSummary{
|
||||
Category: category, ScannerID: category, Title: definition.LabelZH,
|
||||
Description: definition.Description, Severity: string(result.RiskLevel),
|
||||
SeverityLabel: riskLabelZH(result.RiskLevel), Action: string(result.Action),
|
||||
ActionLabel: actionLabelZH(result.Action), Code: "prompt_audit_" + category,
|
||||
Score: result.ScannerScores[category], Evidence: evidence,
|
||||
EvidenceHash: hex.EncodeToString(digest[:]),
|
||||
})
|
||||
}
|
||||
for _, category := range result.UnknownCategories {
|
||||
evidence := "unknown_unsafe"
|
||||
digest := sha256.Sum256([]byte(evidence + ":" + category))
|
||||
summaries = append(summaries, IssueSummary{
|
||||
Category: category, ScannerID: "unknown_unsafe", Title: "未知高风险分类",
|
||||
Description: "审计节点返回了未知但不可忽略的高风险分类", Severity: string(RiskCritical),
|
||||
SeverityLabel: riskLabelZH(RiskCritical), Action: string(ActionBlock),
|
||||
ActionLabel: actionLabelZH(ActionBlock), Code: "prompt_audit_unknown_unsafe",
|
||||
Score: 1, Evidence: evidence, EvidenceHash: hex.EncodeToString(digest[:]),
|
||||
})
|
||||
}
|
||||
return summaries
|
||||
}
|
||||
|
||||
func riskLabelZH(risk RiskLevel) string {
|
||||
switch risk {
|
||||
case RiskCritical:
|
||||
return "严重"
|
||||
case RiskHigh:
|
||||
return "高"
|
||||
case RiskMedium:
|
||||
return "中"
|
||||
default:
|
||||
return "低"
|
||||
}
|
||||
}
|
||||
|
||||
func actionLabelZH(action Action) string {
|
||||
switch action {
|
||||
case ActionBlock:
|
||||
return "阻止"
|
||||
case ActionWarn:
|
||||
return "警告"
|
||||
default:
|
||||
return "允许"
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,172 @@
|
||||
package securityaudit
|
||||
|
||||
import (
|
||||
"context"
|
||||
"log/slog"
|
||||
"strings"
|
||||
)
|
||||
|
||||
const (
|
||||
EventConfigUpdated = "prompt_audit.config_updated"
|
||||
EventConfigLoaded = "prompt_guard.config_loaded"
|
||||
EventConfigReloadDegraded = "prompt_guard.config_reload_degraded"
|
||||
EventProbeStarted = "prompt_audit.endpoint_probe_started"
|
||||
EventProbeFinished = "prompt_audit.endpoint_probe_finished"
|
||||
EventProbeFailed = "prompt_audit.endpoint_probe_failed"
|
||||
EventJobEnqueued = "prompt_audit.job_enqueued"
|
||||
EventEnqueueSkipped = "prompt_audit.enqueue_skipped"
|
||||
EventEnqueueDropped = "prompt_audit.enqueue_dropped"
|
||||
EventAuditStarted = "prompt_audit.started"
|
||||
EventProcessingReclaimed = "prompt_audit.processing_reclaimed"
|
||||
EventProcessed = "prompt_audit.processed"
|
||||
EventProcessFailed = "prompt_audit.process_failed"
|
||||
EventFindingRecorded = "prompt_audit.finding_recorded"
|
||||
EventChunkStarted = "prompt_audit.scan_chunk_started"
|
||||
EventChunkCompleted = "prompt_audit.scan_chunk_completed"
|
||||
EventChunkFailed = "prompt_audit.scan_chunk_failed"
|
||||
EventChunksAggregated = "prompt_audit.scan_chunks_aggregated"
|
||||
EventEvaluationStarted = "prompt_guard.evaluation_started"
|
||||
EventGuardAllowed = "prompt_guard.allowed"
|
||||
EventGuardBlocked = "prompt_guard.blocked"
|
||||
EventGuardFailed = "prompt_guard.failed"
|
||||
EventResultRecordFailed = "prompt_guard.result_record_failed"
|
||||
EventEventDeleted = "prompt_audit.event_deleted"
|
||||
EventEventsDeleted = "prompt_audit.events_deleted"
|
||||
EventDeletePreviewed = "prompt_audit.events_delete_previewed"
|
||||
EventEventsFilterDeleted = "prompt_audit.events_filter_deleted"
|
||||
)
|
||||
|
||||
var knownLogEvents = map[string]struct{}{
|
||||
EventConfigUpdated: {}, EventConfigLoaded: {}, EventConfigReloadDegraded: {},
|
||||
EventProbeStarted: {}, EventProbeFinished: {}, EventProbeFailed: {},
|
||||
EventJobEnqueued: {}, EventEnqueueSkipped: {}, EventEnqueueDropped: {},
|
||||
EventAuditStarted: {}, EventProcessingReclaimed: {}, EventProcessed: {}, EventProcessFailed: {}, EventFindingRecorded: {},
|
||||
EventChunkStarted: {}, EventChunkCompleted: {}, EventChunkFailed: {}, EventChunksAggregated: {},
|
||||
EventEvaluationStarted: {}, EventGuardAllowed: {}, EventGuardBlocked: {}, EventGuardFailed: {}, EventResultRecordFailed: {},
|
||||
EventEventDeleted: {}, EventEventsDeleted: {}, EventDeletePreviewed: {}, EventEventsFilterDeleted: {},
|
||||
}
|
||||
|
||||
var allowedLogFields = map[string]struct{}{
|
||||
"request_id": {}, "user_id": {}, "api_key_id": {}, "group_id": {}, "provider": {},
|
||||
"protocol": {}, "endpoint": {}, "model": {}, "job_id": {}, "event_id": {},
|
||||
"config_version": {}, "guard_endpoint_id": {}, "decision": {}, "risk_level": {},
|
||||
"action": {}, "chunk_index": {}, "chunk_total": {}, "chunk_chars": {}, "input_chars": {},
|
||||
"input_limit": {}, "latency_ms": {}, "status": {}, "error_code": {}, "error_kind": {},
|
||||
"queue_length": {}, "queue_capacity": {}, "stage": {}, "upstream_dispatched": {},
|
||||
"billing_preconsumed": {}, "worker_id": {}, "reclaimed_total": {}, "attempts": {},
|
||||
"max_attempts": {}, "claim_version": {}, "http_status": {}, "retryable": {},
|
||||
}
|
||||
|
||||
func LogInfo(event string, fields map[string]any) {
|
||||
if _, ok := knownLogEvents[event]; !ok {
|
||||
return
|
||||
}
|
||||
slog.LogAttrs(context.Background(), slog.LevelInfo, event, safeAttrs(fields)...)
|
||||
}
|
||||
func LogWarn(event string, fields map[string]any) {
|
||||
if _, ok := knownLogEvents[event]; !ok {
|
||||
return
|
||||
}
|
||||
slog.LogAttrs(context.Background(), slog.LevelWarn, event, safeAttrs(fields)...)
|
||||
}
|
||||
func LogError(event string, fields map[string]any) {
|
||||
if _, ok := knownLogEvents[event]; !ok {
|
||||
return
|
||||
}
|
||||
slog.LogAttrs(context.Background(), slog.LevelError, event, safeAttrs(fields)...)
|
||||
}
|
||||
|
||||
func safeAttrs(fields map[string]any) []slog.Attr {
|
||||
attrs := make([]slog.Attr, 0, len(fields))
|
||||
for key, value := range fields {
|
||||
key = strings.TrimSpace(key)
|
||||
if _, allowed := allowedLogFields[key]; !allowed {
|
||||
continue
|
||||
}
|
||||
if text, ok := value.(string); ok {
|
||||
if key == "error_kind" || key == "error_code" {
|
||||
value = stableErrorCode(text)
|
||||
} else {
|
||||
value = TrimRunes(strings.TrimSpace(text), 256)
|
||||
}
|
||||
}
|
||||
attrs = append(attrs, slog.Any(key, value))
|
||||
}
|
||||
return attrs
|
||||
}
|
||||
|
||||
func mergeLogFields(base map[string]any, extra map[string]any) map[string]any {
|
||||
result := make(map[string]any, len(base)+len(extra))
|
||||
for key, value := range base {
|
||||
result[key] = value
|
||||
}
|
||||
for key, value := range extra {
|
||||
result[key] = value
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func requestLogFields(req Request) map[string]any {
|
||||
return map[string]any{
|
||||
"request_id": req.RequestID, "user_id": req.UserID, "api_key_id": req.APIKeyID,
|
||||
"group_id": pointerLogID(req.GroupID), "provider": req.Provider, "protocol": req.Protocol,
|
||||
"endpoint": req.Endpoint, "model": req.Model, "stage": req.Stage,
|
||||
}
|
||||
}
|
||||
|
||||
func snapshotLogFields(snapshot PromptSnapshot) map[string]any {
|
||||
return map[string]any{
|
||||
"request_id": snapshot.RequestID, "user_id": snapshot.UserID, "api_key_id": snapshot.APIKeyID,
|
||||
"group_id": pointerLogID(snapshot.GroupID), "provider": snapshot.Provider, "protocol": snapshot.Protocol,
|
||||
"endpoint": snapshot.Endpoint, "model": snapshot.Model, "stage": snapshot.Stage,
|
||||
}
|
||||
}
|
||||
|
||||
func jobLogFields(job *Job) map[string]any {
|
||||
if job == nil {
|
||||
return map[string]any{}
|
||||
}
|
||||
fields := snapshotLogFields(job.Snapshot)
|
||||
fields["job_id"] = job.ID
|
||||
fields["config_version"] = job.ConfigVersion
|
||||
fields["claim_version"] = job.ClaimVersion
|
||||
return fields
|
||||
}
|
||||
|
||||
func stableErrorCode(code string) string {
|
||||
code = strings.ToLower(strings.TrimSpace(code))
|
||||
if code == "" {
|
||||
return "unknown_error"
|
||||
}
|
||||
for _, char := range code {
|
||||
if (char >= 'a' && char <= 'z') || (char >= '0' && char <= '9') || char == '_' || char == '-' || char == '.' {
|
||||
continue
|
||||
}
|
||||
return "redacted_error"
|
||||
}
|
||||
return TrimRunes(code, 64)
|
||||
}
|
||||
|
||||
func stableErrorMessage(code string) string {
|
||||
switch stableErrorCode(code) {
|
||||
case ErrorCodeBlocked:
|
||||
return "Prompt Guard blocked the request"
|
||||
case ErrorCodeUnavailable, "payload_store_unavailable", "payload_missing":
|
||||
return "Prompt Audit dependency is unavailable"
|
||||
case ErrorCodeInvalidResponse:
|
||||
return "Prompt Guard returned an invalid response"
|
||||
case "queue_full", "queue_admission_busy":
|
||||
return "Prompt Audit queue is unavailable"
|
||||
case "worker_panic":
|
||||
return "Prompt Audit worker failed"
|
||||
case "config_load_failed", "config_ttl_reload_failed", "config_invalidation_reload_failed":
|
||||
return "Prompt Audit configuration could not be loaded"
|
||||
default:
|
||||
return "Prompt Audit operation failed"
|
||||
}
|
||||
}
|
||||
|
||||
func sanitizeStoredError(code string) (string, string) {
|
||||
stableCode := stableErrorCode(code)
|
||||
return stableCode, TrimRunes(stableErrorMessage(stableCode), 160)
|
||||
}
|
||||
@@ -0,0 +1,66 @@
|
||||
package securityaudit
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"log/slog"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestPromptAuditLogAllowlistAndErrorsDoNotLeakCanarySecrets(t *testing.T) {
|
||||
const canary = "PROMPT_AUDIT_CANARY_SECRET_DO_NOT_PERSIST"
|
||||
var output bytes.Buffer
|
||||
previous := slog.Default()
|
||||
slog.SetDefault(slog.New(slog.NewJSONHandler(&output, nil)))
|
||||
t.Cleanup(func() { slog.SetDefault(previous) })
|
||||
|
||||
LogWarn(EventConfigReloadDegraded, map[string]any{
|
||||
"status": "degraded",
|
||||
"error_code": "config_reload_failed",
|
||||
"error_kind": "Authorization: Bearer " + canary,
|
||||
"token": canary,
|
||||
"body": canary,
|
||||
"base_url": "https://guard.example.test/path?api_key=" + canary,
|
||||
"raw_prompt": "prompt " + canary,
|
||||
})
|
||||
require.NotContains(t, output.String(), canary)
|
||||
require.NotContains(t, output.String(), "api_key=")
|
||||
require.Contains(t, output.String(), EventConfigReloadDegraded)
|
||||
|
||||
beforeUnknown := output.Len()
|
||||
LogWarn("prompt_audit.typo_event", map[string]any{"status": "failed"})
|
||||
require.Equal(t, beforeUnknown, output.Len(), "events outside the stable dictionary must not be emitted")
|
||||
require.Len(t, knownLogEvents, 27)
|
||||
|
||||
_, err := NormalizeBaseURL("https://guard.example.test/path?token=" + canary)
|
||||
require.Error(t, err)
|
||||
require.NotContains(t, err.Error(), canary)
|
||||
}
|
||||
|
||||
func TestPromptGuardFailureLogUsesCompleteAllowlistedContextAndNoSideEffects(t *testing.T) {
|
||||
var output bytes.Buffer
|
||||
previous := slog.Default()
|
||||
slog.SetDefault(slog.New(slog.NewJSONHandler(&output, nil)))
|
||||
t.Cleanup(func() { slog.SetDefault(previous) })
|
||||
groupID := int64(9)
|
||||
snapshot := PromptSnapshot{
|
||||
RequestID: "req-1", UserID: 2, APIKeyID: 3, GroupID: &groupID,
|
||||
Provider: "openai", Protocol: "openai_chat", Endpoint: "/v1/chat/completions",
|
||||
Model: "gpt-test", Stage: "http",
|
||||
}
|
||||
logGuardFailure(snapshot, ActiveConfig{ConfigVersion: 7}, DecisionUnavailable, ErrorCodeUnavailable, "guard-1", 25*time.Millisecond)
|
||||
|
||||
var entry map[string]any
|
||||
require.NoError(t, json.Unmarshal(output.Bytes(), &entry))
|
||||
for key := range snapshotLogFields(snapshot) {
|
||||
require.Contains(t, entry, key)
|
||||
}
|
||||
require.EqualValues(t, 7, entry["config_version"])
|
||||
require.Equal(t, ErrorCodeUnavailable, entry["error_code"])
|
||||
require.Equal(t, false, entry["upstream_dispatched"])
|
||||
require.Equal(t, false, entry["billing_preconsumed"])
|
||||
require.EqualValues(t, 25, entry["latency_ms"])
|
||||
}
|
||||
@@ -0,0 +1,145 @@
|
||||
package securityaudit
|
||||
|
||||
import (
|
||||
"sort"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
)
|
||||
|
||||
const latencySampleCapacity = 2048
|
||||
|
||||
type AtomicMetrics struct {
|
||||
total atomic.Int64
|
||||
allowed atomic.Int64
|
||||
flagged atomic.Int64
|
||||
blocked atomic.Int64
|
||||
unavailable atomic.Int64
|
||||
invalid atomic.Int64
|
||||
timeouts atomic.Int64
|
||||
failovers atomic.Int64
|
||||
bulkheadFull atomic.Int64
|
||||
recordFailed atomic.Int64
|
||||
latencyTotal atomic.Int64
|
||||
latencyMax atomic.Int64
|
||||
enqueued atomic.Int64
|
||||
dropped atomic.Int64
|
||||
latencyMu sync.RWMutex
|
||||
latencies []int64
|
||||
latencyNext int
|
||||
}
|
||||
|
||||
func NewAtomicMetrics() *AtomicMetrics { return &AtomicMetrics{} }
|
||||
|
||||
func (m *AtomicMetrics) Snapshot() GuardMetricsSnapshot {
|
||||
if m == nil {
|
||||
return GuardMetricsSnapshot{}
|
||||
}
|
||||
snapshot := GuardMetricsSnapshot{
|
||||
Total: m.total.Load(), Allowed: m.allowed.Load(), Flagged: m.flagged.Load(),
|
||||
Blocked: m.blocked.Load(), Unavailable: m.unavailable.Load(), Invalid: m.invalid.Load(),
|
||||
Timeouts: m.timeouts.Load(), Failovers: m.failovers.Load(), BulkheadFull: m.bulkheadFull.Load(),
|
||||
RecordFailed: m.recordFailed.Load(), LatencyCount: m.total.Load(), LatencyMaxMS: m.latencyMax.Load(),
|
||||
}
|
||||
if snapshot.LatencyCount > 0 {
|
||||
snapshot.LatencyAvgMS = m.latencyTotal.Load() / snapshot.LatencyCount
|
||||
}
|
||||
m.latencyMu.RLock()
|
||||
samples := append([]int64(nil), m.latencies...)
|
||||
m.latencyMu.RUnlock()
|
||||
if len(samples) > 0 {
|
||||
sort.Slice(samples, func(i, j int) bool { return samples[i] < samples[j] })
|
||||
snapshot.LatencyP50MS = percentile(samples, 0.50)
|
||||
snapshot.LatencyP95MS = percentile(samples, 0.95)
|
||||
snapshot.LatencyP99MS = percentile(samples, 0.99)
|
||||
}
|
||||
return snapshot
|
||||
}
|
||||
|
||||
func (m *AtomicMetrics) AuditSnapshot() AuditMetricsSnapshot {
|
||||
if m == nil {
|
||||
return AuditMetricsSnapshot{}
|
||||
}
|
||||
return AuditMetricsSnapshot{Enqueued: m.enqueued.Load(), Dropped: m.dropped.Load()}
|
||||
}
|
||||
|
||||
func (m *AtomicMetrics) Observe(kind DecisionKind, latency time.Duration) {
|
||||
if m == nil {
|
||||
return
|
||||
}
|
||||
m.total.Add(1)
|
||||
latencyMS := latency.Milliseconds()
|
||||
if latencyMS < 0 {
|
||||
latencyMS = 0
|
||||
}
|
||||
m.latencyTotal.Add(latencyMS)
|
||||
for current := m.latencyMax.Load(); latencyMS > current && !m.latencyMax.CompareAndSwap(current, latencyMS); current = m.latencyMax.Load() {
|
||||
}
|
||||
m.latencyMu.Lock()
|
||||
if len(m.latencies) < latencySampleCapacity {
|
||||
m.latencies = append(m.latencies, latencyMS)
|
||||
} else {
|
||||
m.latencies[m.latencyNext] = latencyMS
|
||||
m.latencyNext = (m.latencyNext + 1) % latencySampleCapacity
|
||||
}
|
||||
m.latencyMu.Unlock()
|
||||
switch kind {
|
||||
case DecisionFlag:
|
||||
m.flagged.Add(1)
|
||||
case DecisionBlock:
|
||||
m.blocked.Add(1)
|
||||
case DecisionUnavailable:
|
||||
m.unavailable.Add(1)
|
||||
case DecisionInvalid:
|
||||
m.invalid.Add(1)
|
||||
default:
|
||||
m.allowed.Add(1)
|
||||
}
|
||||
}
|
||||
|
||||
func percentile(sorted []int64, quantile float64) int64 {
|
||||
if len(sorted) == 0 {
|
||||
return 0
|
||||
}
|
||||
index := int(float64(len(sorted)-1) * quantile)
|
||||
if index < 0 {
|
||||
index = 0
|
||||
}
|
||||
if index >= len(sorted) {
|
||||
index = len(sorted) - 1
|
||||
}
|
||||
return sorted[index]
|
||||
}
|
||||
|
||||
func (m *AtomicMetrics) IncEnqueued() {
|
||||
if m != nil {
|
||||
m.enqueued.Add(1)
|
||||
}
|
||||
}
|
||||
|
||||
func (m *AtomicMetrics) IncDropped() {
|
||||
if m != nil {
|
||||
m.dropped.Add(1)
|
||||
}
|
||||
}
|
||||
|
||||
func (m *AtomicMetrics) IncTimeout() {
|
||||
if m != nil {
|
||||
m.timeouts.Add(1)
|
||||
}
|
||||
}
|
||||
func (m *AtomicMetrics) IncFailover() {
|
||||
if m != nil {
|
||||
m.failovers.Add(1)
|
||||
}
|
||||
}
|
||||
func (m *AtomicMetrics) IncBulkheadFull() {
|
||||
if m != nil {
|
||||
m.bulkheadFull.Add(1)
|
||||
}
|
||||
}
|
||||
func (m *AtomicMetrics) IncRecordFailed() {
|
||||
if m != nil {
|
||||
m.recordFailed.Add(1)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,52 @@
|
||||
package securityaudit
|
||||
|
||||
import (
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestAtomicMetricsExposeCountsLatencyDistributionAndAsyncDelivery(t *testing.T) {
|
||||
metrics := NewAtomicMetrics()
|
||||
latencies := []time.Duration{10, 20, 30, 40, 100}
|
||||
kinds := []DecisionKind{DecisionAllow, DecisionFlag, DecisionBlock, DecisionUnavailable, DecisionInvalid}
|
||||
for index := range latencies {
|
||||
metrics.Observe(kinds[index], latencies[index]*time.Millisecond)
|
||||
}
|
||||
metrics.IncTimeout()
|
||||
metrics.IncFailover()
|
||||
metrics.IncBulkheadFull()
|
||||
metrics.IncRecordFailed()
|
||||
metrics.IncEnqueued()
|
||||
metrics.IncDropped()
|
||||
|
||||
snapshot := metrics.Snapshot()
|
||||
require.Equal(t, int64(5), snapshot.Total)
|
||||
require.Equal(t, int64(5), snapshot.LatencyCount)
|
||||
require.Equal(t, int64(40), snapshot.LatencyAvgMS)
|
||||
require.Equal(t, int64(30), snapshot.LatencyP50MS)
|
||||
require.Equal(t, int64(40), snapshot.LatencyP95MS)
|
||||
require.Equal(t, int64(40), snapshot.LatencyP99MS)
|
||||
require.Equal(t, int64(100), snapshot.LatencyMaxMS)
|
||||
require.Equal(t, AuditMetricsSnapshot{Enqueued: 1, Dropped: 1}, metrics.AuditSnapshot())
|
||||
}
|
||||
|
||||
func TestAtomicMetricsConcurrentObservationIsBoundedAndRaceSafe(t *testing.T) {
|
||||
metrics := NewAtomicMetrics()
|
||||
const observations = 4096
|
||||
var wg sync.WaitGroup
|
||||
for index := 0; index < observations; index++ {
|
||||
wg.Add(1)
|
||||
go func(value int) {
|
||||
defer wg.Done()
|
||||
metrics.Observe(DecisionAllow, time.Duration(value%250)*time.Millisecond)
|
||||
}(index)
|
||||
}
|
||||
wg.Wait()
|
||||
require.Equal(t, int64(observations), metrics.Snapshot().Total)
|
||||
metrics.latencyMu.RLock()
|
||||
require.LessOrEqual(t, len(metrics.latencies), latencySampleCapacity)
|
||||
metrics.latencyMu.RUnlock()
|
||||
}
|
||||
@@ -0,0 +1,22 @@
|
||||
package securityaudit
|
||||
|
||||
import "github.com/google/wire"
|
||||
|
||||
var ProviderSet = wire.NewSet(
|
||||
NewPostgreSQLRepository,
|
||||
wire.Bind(new(JobRepository), new(*PostgreSQLRepository)),
|
||||
wire.Bind(new(EventRepository), new(*PostgreSQLRepository)),
|
||||
NewRedisPayloadStore,
|
||||
wire.Bind(new(PayloadStore), new(*RedisPayloadStore)),
|
||||
NewOpenAICompatibleScanner,
|
||||
wire.Bind(new(PromptScanner), new(*OpenAICompatibleScanner)),
|
||||
NewAtomicMetrics,
|
||||
wire.Bind(new(Metrics), new(*AtomicMetrics)),
|
||||
NewConfigManager,
|
||||
wire.Bind(new(ConfigStore), new(*ConfigManager)),
|
||||
NewPromptService,
|
||||
wire.Bind(new(PromptEngine), new(*PromptService)),
|
||||
NewLegacyModerationAdapter,
|
||||
NewCoordinator,
|
||||
NewPromptAdminHandler,
|
||||
)
|
||||
@@ -0,0 +1,196 @@
|
||||
package securityaudit
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/tls"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"net/url"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors"
|
||||
)
|
||||
|
||||
const maxGuardResponseBytes int64 = 256 * 1024
|
||||
|
||||
var (
|
||||
errRedirectBlocked = errors.New("prompt guard redirect blocked")
|
||||
metadataHosts = map[string]struct{}{
|
||||
"metadata": {}, "metadata.google.internal": {}, "metadata.azure.internal": {},
|
||||
"instance-data": {}, "instance-data.ec2.internal": {},
|
||||
}
|
||||
blockedPrefixes = []netip.Prefix{
|
||||
netip.MustParsePrefix("0.0.0.0/8"),
|
||||
netip.MustParsePrefix("100.64.0.0/10"),
|
||||
netip.MustParsePrefix("169.254.0.0/16"),
|
||||
netip.MustParsePrefix("192.0.0.0/24"),
|
||||
netip.MustParsePrefix("192.0.2.0/24"),
|
||||
netip.MustParsePrefix("198.18.0.0/15"),
|
||||
netip.MustParsePrefix("198.51.100.0/24"),
|
||||
netip.MustParsePrefix("203.0.113.0/24"),
|
||||
netip.MustParsePrefix("224.0.0.0/4"),
|
||||
netip.MustParsePrefix("240.0.0.0/4"),
|
||||
netip.MustParsePrefix("::/128"),
|
||||
netip.MustParsePrefix("fe80::/10"),
|
||||
netip.MustParsePrefix("ff00::/8"),
|
||||
netip.MustParsePrefix("2001:db8::/32"),
|
||||
}
|
||||
)
|
||||
|
||||
type DNSResolver interface {
|
||||
LookupNetIP(ctx context.Context, network, host string) ([]netip.Addr, error)
|
||||
}
|
||||
|
||||
type netResolver struct{ resolver *net.Resolver }
|
||||
|
||||
func (r netResolver) LookupNetIP(ctx context.Context, network, host string) ([]netip.Addr, error) {
|
||||
return r.resolver.LookupNetIP(ctx, network, host)
|
||||
}
|
||||
|
||||
func NormalizeBaseURL(raw string) (string, error) {
|
||||
raw = strings.TrimSpace(raw)
|
||||
parsed, err := url.Parse(raw)
|
||||
if err != nil || parsed.Scheme == "" || parsed.Host == "" {
|
||||
return "", infraerrors.BadRequest("prompt_audit_invalid_base_url", "审计节点地址无效")
|
||||
}
|
||||
parsed.Scheme = strings.ToLower(parsed.Scheme)
|
||||
if parsed.Scheme != "http" && parsed.Scheme != "https" {
|
||||
return "", infraerrors.BadRequest("prompt_audit_invalid_base_url_scheme", "审计节点仅支持 HTTP(S)")
|
||||
}
|
||||
if parsed.User != nil || parsed.RawQuery != "" || parsed.Fragment != "" {
|
||||
return "", infraerrors.BadRequest("prompt_audit_unsafe_base_url", "审计节点地址不能包含凭据、查询参数或片段")
|
||||
}
|
||||
host := strings.ToLower(strings.TrimSuffix(parsed.Hostname(), "."))
|
||||
if host == "" {
|
||||
return "", infraerrors.BadRequest("prompt_audit_invalid_base_url", "审计节点地址无效")
|
||||
}
|
||||
if _, blocked := metadataHosts[host]; blocked || strings.HasSuffix(host, ".metadata.google.internal") {
|
||||
return "", infraerrors.BadRequest("prompt_audit_unsafe_base_url", "审计节点地址不在允许范围")
|
||||
}
|
||||
allowPrivate := isExplicitPrivateHost(host)
|
||||
if addr, err := netip.ParseAddr(host); err == nil {
|
||||
if isBlockedAddress(addr) {
|
||||
return "", infraerrors.BadRequest("prompt_audit_unsafe_base_url", "审计节点地址不在允许范围")
|
||||
}
|
||||
allowPrivate = addr.IsPrivate() || addr.IsLoopback()
|
||||
}
|
||||
if parsed.Scheme == "http" && !allowPrivate {
|
||||
return "", infraerrors.BadRequest("prompt_audit_https_required", "公网审计节点必须使用 HTTPS")
|
||||
}
|
||||
path := strings.TrimRight(parsed.EscapedPath(), "/")
|
||||
if strings.EqualFold(path, "/v1") {
|
||||
path = ""
|
||||
}
|
||||
parsed.Path = path
|
||||
parsed.RawPath = ""
|
||||
return strings.TrimRight(parsed.String(), "/"), nil
|
||||
}
|
||||
|
||||
func ChatCompletionsURL(base string) (string, error) {
|
||||
normalized, err := NormalizeBaseURL(base)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return normalized + "/v1/chat/completions", nil
|
||||
}
|
||||
|
||||
func ModelsURL(base string) (string, error) {
|
||||
normalized, err := NormalizeBaseURL(base)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return normalized + "/v1/models", nil
|
||||
}
|
||||
|
||||
func NewSecureHTTPClient(endpoint ActiveEndpoint) (*http.Client, error) {
|
||||
normalized, err := NormalizeBaseURL(endpoint.BaseURL)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
parsed, _ := url.Parse(normalized)
|
||||
host := strings.ToLower(strings.TrimSuffix(parsed.Hostname(), "."))
|
||||
allowPrivate := isExplicitPrivateHost(host)
|
||||
if addr, parseErr := netip.ParseAddr(host); parseErr == nil {
|
||||
allowPrivate = addr.IsPrivate() || addr.IsLoopback()
|
||||
}
|
||||
resolver := netResolver{resolver: net.DefaultResolver}
|
||||
dialer := &net.Dialer{Timeout: 3 * time.Second, KeepAlive: 30 * time.Second}
|
||||
transport := &http.Transport{
|
||||
// Do not inherit HTTP(S)_PROXY. A proxy would move the actual destination
|
||||
// dial outside secureDialContext and bypass this module's DNS/IP validation.
|
||||
Proxy: nil,
|
||||
ForceAttemptHTTP2: true,
|
||||
MaxIdleConns: 64,
|
||||
MaxIdleConnsPerHost: 16,
|
||||
IdleConnTimeout: 90 * time.Second,
|
||||
TLSHandshakeTimeout: 5 * time.Second,
|
||||
ResponseHeaderTimeout: time.Duration(endpoint.TimeoutMS) * time.Millisecond,
|
||||
ExpectContinueTimeout: time.Second,
|
||||
TLSClientConfig: &tls.Config{MinVersion: tls.VersionTLS12},
|
||||
}
|
||||
transport.DialContext = secureDialContext(dialer, resolver, allowPrivate)
|
||||
timeout := time.Duration(endpoint.TimeoutMS) * time.Millisecond
|
||||
if timeout <= 0 {
|
||||
timeout = DefaultTimeoutMS * time.Millisecond
|
||||
}
|
||||
return &http.Client{
|
||||
Transport: transport,
|
||||
Timeout: timeout,
|
||||
CheckRedirect: func(_ *http.Request, _ []*http.Request) error {
|
||||
return errRedirectBlocked
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
|
||||
func secureDialContext(dialer *net.Dialer, resolver DNSResolver, allowPrivate bool) func(context.Context, string, string) (net.Conn, error) {
|
||||
return func(ctx context.Context, network, address string) (net.Conn, error) {
|
||||
host, port, err := net.SplitHostPort(address)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("prompt guard dial address invalid")
|
||||
}
|
||||
addresses, err := resolver.LookupNetIP(ctx, "ip", host)
|
||||
if err != nil || len(addresses) == 0 {
|
||||
return nil, fmt.Errorf("prompt guard dns unavailable")
|
||||
}
|
||||
var lastErr error
|
||||
for _, addr := range addresses {
|
||||
if isBlockedAddress(addr) || (!allowPrivate && (addr.IsPrivate() || addr.IsLoopback())) {
|
||||
lastErr = fmt.Errorf("prompt guard resolved address blocked")
|
||||
continue
|
||||
}
|
||||
if !addr.IsGlobalUnicast() && !addr.IsPrivate() && !addr.IsLoopback() {
|
||||
lastErr = fmt.Errorf("prompt guard resolved address blocked")
|
||||
continue
|
||||
}
|
||||
conn, dialErr := dialer.DialContext(ctx, network, net.JoinHostPort(addr.String(), port))
|
||||
if dialErr == nil {
|
||||
return conn, nil
|
||||
}
|
||||
lastErr = dialErr
|
||||
}
|
||||
if lastErr == nil {
|
||||
lastErr = fmt.Errorf("prompt guard no allowed resolved address")
|
||||
}
|
||||
return nil, lastErr
|
||||
}
|
||||
}
|
||||
|
||||
func isExplicitPrivateHost(host string) bool {
|
||||
return host == "localhost" || strings.HasSuffix(host, ".localhost") || strings.HasSuffix(host, ".local")
|
||||
}
|
||||
|
||||
func isBlockedAddress(addr netip.Addr) bool {
|
||||
if !addr.IsValid() || addr.IsUnspecified() || addr.IsMulticast() || addr.IsLinkLocalUnicast() || addr.IsLinkLocalMulticast() {
|
||||
return true
|
||||
}
|
||||
for _, prefix := range blockedPrefixes {
|
||||
if prefix.Contains(addr) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
@@ -0,0 +1,220 @@
|
||||
package securityaudit
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/netip"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
type staticResolver struct{ addresses []netip.Addr }
|
||||
|
||||
func (r staticResolver) LookupNetIP(context.Context, string, string) ([]netip.Addr, error) {
|
||||
return r.addresses, nil
|
||||
}
|
||||
|
||||
func TestNormalizeBaseURLSecurity(t *testing.T) {
|
||||
allowed := []string{"https://guard.example.com", "https://guard.example.com/v1", "http://127.0.0.1:8080", "http://10.0.0.8:8080"}
|
||||
for _, raw := range allowed {
|
||||
_, err := NormalizeBaseURL(raw)
|
||||
require.NoError(t, err, raw)
|
||||
}
|
||||
blocked := []string{
|
||||
"ftp://guard.example.com", "http://guard.example.com", "https://user:pass@guard.example.com",
|
||||
"https://guard.example.com?q=secret", "https://guard.example.com/#fragment", "http://169.254.169.254",
|
||||
"https://metadata.google.internal", "https://0.0.0.0", "https://224.0.0.1", "https://192.0.2.1",
|
||||
"https://[::]", "https://[fe80::1]", "https://[ff02::1]", "https://[2001:db8::1]",
|
||||
}
|
||||
for _, raw := range blocked {
|
||||
_, err := NormalizeBaseURL(raw)
|
||||
require.Error(t, err, raw)
|
||||
}
|
||||
url, err := ChatCompletionsURL("https://guard.example.com/v1")
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "https://guard.example.com/v1/chat/completions", url)
|
||||
}
|
||||
|
||||
func TestSecureDialRejectsDNSRebindingToPrivateAddress(t *testing.T) {
|
||||
dial := secureDialContext(nil, staticResolver{addresses: []netip.Addr{netip.MustParseAddr("127.0.0.1")}}, false)
|
||||
_, err := dial(context.Background(), "tcp", "guard.example.com:443")
|
||||
require.Error(t, err)
|
||||
}
|
||||
|
||||
func TestSecureHTTPClientDoesNotBypassDestinationValidationThroughEnvironmentProxy(t *testing.T) {
|
||||
client, err := NewSecureHTTPClient(ActiveEndpoint{BaseURL: "https://guard.example.com", TimeoutMS: 1000})
|
||||
require.NoError(t, err)
|
||||
transport, ok := client.Transport.(*http.Transport)
|
||||
require.True(t, ok)
|
||||
require.Nil(t, transport.Proxy)
|
||||
}
|
||||
|
||||
func TestOpenAICompatibleScannerRequestContract(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
require.Equal(t, "/v1/chat/completions", r.URL.Path)
|
||||
require.Equal(t, "Bearer token", r.Header.Get("Authorization"))
|
||||
var payload map[string]any
|
||||
require.NoError(t, json.NewDecoder(r.Body).Decode(&payload))
|
||||
require.Equal(t, DefaultGuardModel, payload["model"])
|
||||
require.Equal(t, float64(0), payload["temperature"])
|
||||
require.Equal(t, float64(64), payload["max_tokens"])
|
||||
require.Equal(t, float64(42), payload["seed"])
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
_, _ = w.Write([]byte(`{"choices":[{"message":{"content":"Safety: Safe\nCategories: None"}}]}`))
|
||||
}))
|
||||
defer server.Close()
|
||||
scanner := NewOpenAICompatibleScanner()
|
||||
result, err := scanner.Scan(context.Background(), ActiveEndpoint{ID: "one", BaseURL: server.URL, Model: DefaultGuardModel, Token: "token", TimeoutMS: 1000}, "hello", AllScannerIDs)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, EventPass, result.Decision)
|
||||
}
|
||||
|
||||
func TestOpenAICompatibleScannerRejectsRedirectAndOversize(t *testing.T) {
|
||||
redirect := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
http.Redirect(w, r, "http://127.0.0.1/other", http.StatusFound)
|
||||
}))
|
||||
defer redirect.Close()
|
||||
_, err := NewOpenAICompatibleScanner().Scan(context.Background(), ActiveEndpoint{ID: "redirect", BaseURL: redirect.URL, Model: DefaultGuardModel, TimeoutMS: 1000}, "hello", AllScannerIDs)
|
||||
require.Error(t, err)
|
||||
oversize := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
_, _ = w.Write([]byte(strings.Repeat("x", int(maxGuardResponseBytes)+1)))
|
||||
}))
|
||||
defer oversize.Close()
|
||||
_, err = NewOpenAICompatibleScanner().Scan(context.Background(), ActiveEndpoint{ID: "large", BaseURL: oversize.URL, Model: DefaultGuardModel, TimeoutMS: 1000}, "hello", AllScannerIDs)
|
||||
require.Error(t, err)
|
||||
}
|
||||
|
||||
func TestOpenAICompatibleScannerClassifiesHTTPConnectionAndTimeoutFailures(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
status int
|
||||
retryable bool
|
||||
}{
|
||||
{name: "authentication", status: http.StatusUnauthorized, retryable: false},
|
||||
{name: "forbidden", status: http.StatusForbidden, retryable: false},
|
||||
{name: "rate limited", status: http.StatusTooManyRequests, retryable: true},
|
||||
{name: "server failure", status: http.StatusBadGateway, retryable: true},
|
||||
{name: "other client error", status: http.StatusBadRequest, retryable: false},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.WriteHeader(tt.status)
|
||||
}))
|
||||
defer server.Close()
|
||||
_, err := NewOpenAICompatibleScanner().Scan(context.Background(), ActiveEndpoint{ID: "status", BaseURL: server.URL, Model: DefaultGuardModel, TimeoutMS: 1000}, "hello", AllScannerIDs)
|
||||
var guardErr *GuardError
|
||||
require.ErrorAs(t, err, &guardErr)
|
||||
require.Equal(t, ErrorCodeUnavailable, guardErr.Code)
|
||||
require.Equal(t, tt.status, guardErr.HTTPStatus)
|
||||
require.Equal(t, tt.retryable, guardErr.Retryable)
|
||||
require.NotContains(t, err.Error(), server.URL)
|
||||
})
|
||||
}
|
||||
|
||||
closed := httptest.NewServer(http.HandlerFunc(func(http.ResponseWriter, *http.Request) {}))
|
||||
closedURL := closed.URL
|
||||
closed.Close()
|
||||
_, err := NewOpenAICompatibleScanner().Scan(context.Background(), ActiveEndpoint{ID: "closed", BaseURL: closedURL, Model: DefaultGuardModel, TimeoutMS: 100}, "hello", AllScannerIDs)
|
||||
var connectionErr *GuardError
|
||||
require.ErrorAs(t, err, &connectionErr)
|
||||
require.True(t, connectionErr.Retryable)
|
||||
|
||||
timeout := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))
|
||||
defer timeout.Close()
|
||||
_, err = NewOpenAICompatibleScanner().Scan(context.Background(), ActiveEndpoint{ID: "timeout", BaseURL: timeout.URL, Model: DefaultGuardModel, TimeoutMS: 20}, "hello", AllScannerIDs)
|
||||
var timeoutErr *GuardError
|
||||
require.ErrorAs(t, err, &timeoutErr)
|
||||
require.True(t, timeoutErr.Retryable)
|
||||
require.True(t, timeoutErr.Timeout)
|
||||
}
|
||||
|
||||
func TestPromptAuditProbeModelsFallbackAndResponseSafety(t *testing.T) {
|
||||
t.Run("models contains configured model", func(t *testing.T) {
|
||||
var chatCalls atomic.Int64
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
require.Equal(t, "Bearer temporary-token", r.Header.Get("Authorization"))
|
||||
if r.URL.Path == "/v1/models" {
|
||||
_, _ = w.Write([]byte(`{"data":[{"id":"` + DefaultGuardModel + `"}]}`))
|
||||
return
|
||||
}
|
||||
chatCalls.Add(1)
|
||||
}))
|
||||
defer server.Close()
|
||||
result := newProbeTestService().Probe(context.Background(), ProbeRequest{Endpoint: probeEndpoint(server.URL, "temporary-token")})
|
||||
require.True(t, result.OK)
|
||||
require.True(t, result.TokenApplied)
|
||||
require.Equal(t, http.StatusOK, result.HTTPStatus)
|
||||
require.Zero(t, chatCalls.Load())
|
||||
})
|
||||
|
||||
t.Run("invalid models response performs real guard fallback", func(t *testing.T) {
|
||||
var chatCalls atomic.Int64
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path == "/v1/models" {
|
||||
_, _ = w.Write([]byte(`{"unexpected":true}`))
|
||||
return
|
||||
}
|
||||
chatCalls.Add(1)
|
||||
_, _ = w.Write([]byte(`{"choices":[{"message":{"content":"Safety: Safe\nCategories: None"}}]}`))
|
||||
}))
|
||||
defer server.Close()
|
||||
result := newProbeTestService().Probe(context.Background(), ProbeRequest{Endpoint: probeEndpoint(server.URL, "temporary-token")})
|
||||
require.True(t, result.OK)
|
||||
require.Equal(t, int64(1), chatCalls.Load())
|
||||
})
|
||||
|
||||
t.Run("fallback authentication failure is stable", func(t *testing.T) {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path == "/v1/models" {
|
||||
w.WriteHeader(http.StatusNotFound)
|
||||
return
|
||||
}
|
||||
w.WriteHeader(http.StatusUnauthorized)
|
||||
}))
|
||||
defer server.Close()
|
||||
result := newProbeTestService().Probe(context.Background(), ProbeRequest{Endpoint: probeEndpoint(server.URL, "temporary-token")})
|
||||
require.False(t, result.OK)
|
||||
require.Equal(t, ErrorCodeUnavailable, result.ErrorCode)
|
||||
require.Equal(t, http.StatusUnauthorized, result.HTTPStatus)
|
||||
require.False(t, result.Retryable)
|
||||
})
|
||||
|
||||
t.Run("oversized models response is rejected without fallback", func(t *testing.T) {
|
||||
var chatCalls atomic.Int64
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.URL.Path != "/v1/models" {
|
||||
chatCalls.Add(1)
|
||||
}
|
||||
_, _ = w.Write([]byte(strings.Repeat("x", int(maxGuardResponseBytes)+1)))
|
||||
}))
|
||||
defer server.Close()
|
||||
result := newProbeTestService().Probe(context.Background(), ProbeRequest{Endpoint: probeEndpoint(server.URL, "temporary-token")})
|
||||
require.False(t, result.OK)
|
||||
require.Equal(t, "response_too_large", result.ErrorCode)
|
||||
require.Zero(t, chatCalls.Load())
|
||||
})
|
||||
}
|
||||
|
||||
func newProbeTestService() *PromptService {
|
||||
return &PromptService{
|
||||
config: &ConfigManager{}, scanner: NewOpenAICompatibleScanner(), clock: realClock{},
|
||||
probes: map[string]ProbeResult{},
|
||||
}
|
||||
}
|
||||
|
||||
func probeEndpoint(baseURL, token string) UpdateEndpoint {
|
||||
return UpdateEndpoint{
|
||||
ID: "probe-one", Name: "Probe One", Protocol: "openai_compatible", BaseURL: baseURL,
|
||||
Model: DefaultGuardModel, Token: token, TimeoutMS: 1000, InputLimit: 1024, Enabled: true,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,63 @@
|
||||
package securityaudit
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"github.com/redis/go-redis/v9"
|
||||
)
|
||||
|
||||
type PayloadStore interface {
|
||||
Set(ctx context.Context, jobID int64, scanText string, ttl time.Duration) error
|
||||
Get(ctx context.Context, jobID int64) (string, error)
|
||||
Delete(ctx context.Context, jobID int64) error
|
||||
Ping(ctx context.Context) error
|
||||
}
|
||||
|
||||
type RedisPayloadStore struct {
|
||||
client *redis.Client
|
||||
}
|
||||
|
||||
func NewRedisPayloadStore(client *redis.Client) *RedisPayloadStore {
|
||||
return &RedisPayloadStore{client: client}
|
||||
}
|
||||
|
||||
func (s *RedisPayloadStore) Set(ctx context.Context, jobID int64, scanText string, ttl time.Duration) error {
|
||||
if s == nil || s.client == nil {
|
||||
return fmt.Errorf("prompt audit payload store unavailable")
|
||||
}
|
||||
if jobID <= 0 || scanText == "" {
|
||||
return fmt.Errorf("prompt audit payload input invalid")
|
||||
}
|
||||
if ttl <= 0 || ttl > DefaultPayloadTTL {
|
||||
ttl = DefaultPayloadTTL
|
||||
}
|
||||
return s.client.Set(ctx, payloadKey(jobID), scanText, ttl).Err()
|
||||
}
|
||||
|
||||
func (s *RedisPayloadStore) Get(ctx context.Context, jobID int64) (string, error) {
|
||||
if s == nil || s.client == nil {
|
||||
return "", fmt.Errorf("prompt audit payload store unavailable")
|
||||
}
|
||||
return s.client.Get(ctx, payloadKey(jobID)).Result()
|
||||
}
|
||||
|
||||
func (s *RedisPayloadStore) Delete(ctx context.Context, jobID int64) error {
|
||||
if s == nil || s.client == nil {
|
||||
return fmt.Errorf("prompt audit payload store unavailable")
|
||||
}
|
||||
return s.client.Del(ctx, payloadKey(jobID)).Err()
|
||||
}
|
||||
|
||||
func (s *RedisPayloadStore) Ping(ctx context.Context) error {
|
||||
if s == nil || s.client == nil {
|
||||
return fmt.Errorf("prompt audit payload store unavailable")
|
||||
}
|
||||
return s.client.Ping(ctx).Err()
|
||||
}
|
||||
|
||||
func payloadKey(jobID int64) string {
|
||||
return PayloadKeyPrefix + strconv.FormatInt(jobID, 10)
|
||||
}
|
||||
@@ -0,0 +1,85 @@
|
||||
package securityaudit
|
||||
|
||||
import (
|
||||
"context"
|
||||
"os"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/redis/go-redis/v9"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestRedisPayloadStoreRoundTripTTLNamespaceAndDelete(t *testing.T) {
|
||||
address := strings.TrimSpace(os.Getenv(promptAuditRedisTestEnv))
|
||||
if address == "" {
|
||||
t.Skip(promptAuditRedisTestEnv + " is not set")
|
||||
}
|
||||
client := redis.NewClient(&redis.Options{Addr: address})
|
||||
t.Cleanup(func() { require.NoError(t, client.Close()) })
|
||||
store := NewRedisPayloadStore(client)
|
||||
ctx := context.Background()
|
||||
const jobID int64 = 987654321
|
||||
const canary = "PROMPT_CANARY_REDIS_ONLY_PAYLOAD"
|
||||
_ = store.Delete(ctx, jobID)
|
||||
require.NoError(t, store.Set(ctx, jobID, canary, 2*DefaultPayloadTTL))
|
||||
require.Equal(t, PayloadKeyPrefix+"987654321", payloadKey(jobID))
|
||||
value, err := store.Get(ctx, jobID)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, canary, value)
|
||||
ttl, err := client.TTL(ctx, payloadKey(jobID)).Result()
|
||||
require.NoError(t, err)
|
||||
require.Greater(t, ttl, time.Duration(0))
|
||||
require.LessOrEqual(t, ttl, DefaultPayloadTTL)
|
||||
require.NoError(t, store.Delete(ctx, jobID))
|
||||
_, err = store.Get(ctx, jobID)
|
||||
require.ErrorIs(t, err, redis.Nil)
|
||||
}
|
||||
|
||||
func TestPromptRuntimeAggregatesConfigWorkersQueueRedisEndpointsAndGuardMetrics(t *testing.T) {
|
||||
address := strings.TrimSpace(os.Getenv(promptAuditRedisTestEnv))
|
||||
if address == "" {
|
||||
t.Skip(promptAuditRedisTestEnv + " is not set")
|
||||
}
|
||||
db := openPromptAuditIntegrationDB(t)
|
||||
client := redis.NewClient(&redis.Options{Addr: address})
|
||||
t.Cleanup(func() { require.NoError(t, client.Close()) })
|
||||
|
||||
config := &fakeConfigStore{active: true, cfg: ActiveConfig{
|
||||
RiskControlEnabled: true, Enabled: true, WorkerCount: 3, QueueCapacity: 123,
|
||||
ConfigVersion: 9, AllGroups: true,
|
||||
}}
|
||||
metrics := NewAtomicMetrics()
|
||||
metrics.Observe(DecisionBlock, 25*time.Millisecond)
|
||||
metrics.IncFailover()
|
||||
metrics.IncEnqueued()
|
||||
metrics.IncDropped()
|
||||
service := NewPromptService(
|
||||
config,
|
||||
NewPostgreSQLRepository(db),
|
||||
NewRedisPayloadStore(client),
|
||||
NewOpenAICompatibleScanner(),
|
||||
metrics,
|
||||
)
|
||||
service.probes["guard-1"] = ProbeResult{OK: true, Status: "healthy", HTTPStatus: 200}
|
||||
|
||||
runtime := service.Runtime(context.Background())
|
||||
require.Equal(t, ModeAsync, runtime.EffectiveMode)
|
||||
require.Equal(t, int64(9), runtime.ExpectedConfigVersion)
|
||||
require.Equal(t, int64(9), runtime.ActiveConfigVersion)
|
||||
require.Equal(t, 3, runtime.WorkerTotal)
|
||||
require.Equal(t, 123, runtime.QueueCapacity)
|
||||
require.Equal(t, "ok", runtime.DatabaseStatus)
|
||||
require.Equal(t, "ok", runtime.RedisStatus)
|
||||
require.Contains(t, runtime.Endpoints, "guard-1")
|
||||
require.Equal(t, int64(1), runtime.GuardMetrics.Total)
|
||||
require.Equal(t, int64(1), runtime.GuardMetrics.Blocked)
|
||||
require.Equal(t, int64(1), runtime.GuardMetrics.Failovers)
|
||||
require.Equal(t, int64(25), runtime.GuardMetrics.LatencyP95MS)
|
||||
require.Equal(t, int64(1), runtime.EnqueuedTotal)
|
||||
require.Equal(t, int64(1), runtime.DroppedTotal)
|
||||
// The runner has not been started in this integration test, so the honest
|
||||
// process status is degraded rather than a fabricated running heartbeat.
|
||||
require.Equal(t, "degraded", runtime.ProcessStatus)
|
||||
}
|
||||
@@ -0,0 +1,333 @@
|
||||
package securityaudit
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
)
|
||||
|
||||
type ScannerDefinition struct {
|
||||
ID string `json:"id"`
|
||||
Label string `json:"label"`
|
||||
LabelZH string `json:"label_zh"`
|
||||
Description string `json:"description"`
|
||||
}
|
||||
|
||||
var AllScannerIDs = []string{
|
||||
"violent",
|
||||
"non_violent_illegal_acts",
|
||||
"sexual_content_or_sexual_acts",
|
||||
"pii",
|
||||
"suicide_and_self_harm",
|
||||
"unethical_acts",
|
||||
"politically_sensitive_topics",
|
||||
"copyright_violation",
|
||||
"jailbreak",
|
||||
}
|
||||
|
||||
var ScannerCatalog = map[string]ScannerDefinition{
|
||||
"violent": {ID: "violent", Label: "Violent", LabelZH: "暴力", Description: "Violence or threats of violence"},
|
||||
"non_violent_illegal_acts": {ID: "non_violent_illegal_acts", Label: "Non-violent Illegal Acts", LabelZH: "非暴力违法行为", Description: "Non-violent illegal activity"},
|
||||
"sexual_content_or_sexual_acts": {ID: "sexual_content_or_sexual_acts", Label: "Sexual Content or Sexual Acts", LabelZH: "性内容或性行为", Description: "Sexual content or sexual acts"},
|
||||
"pii": {ID: "pii", Label: "PII", LabelZH: "个人敏感信息", Description: "Personal identifying information"},
|
||||
"suicide_and_self_harm": {ID: "suicide_and_self_harm", Label: "Suicide & Self-Harm", LabelZH: "自杀与自残", Description: "Suicide or self-harm"},
|
||||
"unethical_acts": {ID: "unethical_acts", Label: "Unethical Acts", LabelZH: "不道德行为", Description: "Unethical behavior"},
|
||||
"politically_sensitive_topics": {ID: "politically_sensitive_topics", Label: "Politically Sensitive Topics", LabelZH: "政治敏感话题", Description: "Politically sensitive topics"},
|
||||
"copyright_violation": {ID: "copyright_violation", Label: "Copyright Violation", LabelZH: "版权侵权", Description: "Copyright infringement"},
|
||||
"jailbreak": {ID: "jailbreak", Label: "Jailbreak", LabelZH: "越狱攻击", Description: "Prompt injection or jailbreak attempt"},
|
||||
}
|
||||
|
||||
var categoryAliases = map[string]string{
|
||||
"violent": "violent", "violence": "violent",
|
||||
"non violent illegal acts": "non_violent_illegal_acts", "non-violent illegal acts": "non_violent_illegal_acts",
|
||||
"sexual content or sexual acts": "sexual_content_or_sexual_acts", "sexual": "sexual_content_or_sexual_acts",
|
||||
"pii": "pii", "personal identifying information": "pii", "personal identifiable information": "pii",
|
||||
"suicide self harm": "suicide_and_self_harm", "suicide and self harm": "suicide_and_self_harm", "suicide & self-harm": "suicide_and_self_harm",
|
||||
"unethical acts": "unethical_acts", "unethical": "unethical_acts",
|
||||
"politically sensitive topics": "politically_sensitive_topics", "political": "politically_sensitive_topics",
|
||||
"copyright violation": "copyright_violation", "copyright": "copyright_violation",
|
||||
"jailbreak": "jailbreak", "prompt injection": "jailbreak",
|
||||
}
|
||||
|
||||
type GuardError struct {
|
||||
Code string
|
||||
HTTPStatus int
|
||||
Retryable bool
|
||||
Timeout bool
|
||||
Cause error
|
||||
}
|
||||
|
||||
func (e *GuardError) Error() string {
|
||||
if e == nil {
|
||||
return "<nil>"
|
||||
}
|
||||
return e.Code
|
||||
}
|
||||
|
||||
func (e *GuardError) Unwrap() error { return e.Cause }
|
||||
|
||||
func NormalizeCategory(value string) string {
|
||||
normalized := strings.ToLower(strings.TrimSpace(value))
|
||||
normalized = strings.NewReplacer("_", " ", "&", " and ", "/", " ", "-", " ", "–", " ", "—", " ").Replace(normalized)
|
||||
normalized = strings.Join(strings.Fields(normalized), " ")
|
||||
if canonical, ok := categoryAliases[normalized]; ok {
|
||||
return canonical
|
||||
}
|
||||
return strings.ReplaceAll(normalized, " ", "_")
|
||||
}
|
||||
|
||||
func ParseQwen3Guard(content string, enabledScanners []string) (*NormalizedResult, error) {
|
||||
lines := make([]string, 0, 2)
|
||||
for _, line := range strings.Split(strings.ReplaceAll(content, "\r\n", "\n"), "\n") {
|
||||
line = strings.TrimSpace(line)
|
||||
if line != "" {
|
||||
lines = append(lines, line)
|
||||
}
|
||||
}
|
||||
if len(lines) != 2 {
|
||||
return nil, &GuardError{Code: ErrorCodeInvalidResponse, Retryable: false}
|
||||
}
|
||||
var safety string
|
||||
var categoryLine string
|
||||
for _, line := range lines {
|
||||
lower := strings.ToLower(line)
|
||||
switch {
|
||||
case strings.HasPrefix(lower, "safety:"):
|
||||
if safety != "" {
|
||||
return nil, &GuardError{Code: ErrorCodeInvalidResponse}
|
||||
}
|
||||
safety = strings.TrimSpace(line[len("safety:"):])
|
||||
case strings.HasPrefix(lower, "categories:"):
|
||||
if categoryLine != "" {
|
||||
return nil, &GuardError{Code: ErrorCodeInvalidResponse}
|
||||
}
|
||||
categoryLine = strings.TrimSpace(line[len("categories:"):])
|
||||
default:
|
||||
return nil, &GuardError{Code: ErrorCodeInvalidResponse}
|
||||
}
|
||||
}
|
||||
switch strings.ToLower(safety) {
|
||||
case "safe":
|
||||
safety = "Safe"
|
||||
case "controversial":
|
||||
safety = "Controversial"
|
||||
case "unsafe":
|
||||
safety = "Unsafe"
|
||||
default:
|
||||
return nil, &GuardError{Code: ErrorCodeInvalidResponse}
|
||||
}
|
||||
if categoryLine == "" {
|
||||
return nil, &GuardError{Code: ErrorCodeInvalidResponse}
|
||||
}
|
||||
enabled := make(map[string]struct{}, len(enabledScanners))
|
||||
for _, scanner := range enabledScanners {
|
||||
enabled[NormalizeCategory(scanner)] = struct{}{}
|
||||
}
|
||||
known := map[string]struct{}{}
|
||||
unknown := map[string]struct{}{}
|
||||
for _, raw := range strings.Split(categoryLine, ",") {
|
||||
raw = strings.TrimSpace(raw)
|
||||
if raw == "" || strings.EqualFold(raw, "none") || strings.EqualFold(raw, "n/a") {
|
||||
continue
|
||||
}
|
||||
category := NormalizeCategory(raw)
|
||||
if _, ok := ScannerCatalog[category]; ok {
|
||||
known[category] = struct{}{}
|
||||
} else {
|
||||
unknown[unknownCategoryID(category)] = struct{}{}
|
||||
}
|
||||
}
|
||||
knownList := orderedScannerKeys(known)
|
||||
unknownList := sortedKeys(unknown)
|
||||
matched := make([]string, 0, len(knownList))
|
||||
for _, category := range knownList {
|
||||
if _, ok := enabled[category]; ok {
|
||||
matched = append(matched, category)
|
||||
}
|
||||
}
|
||||
result := &NormalizedResult{
|
||||
Safety: safety, Categories: knownList, MatchedScanners: matched, UnknownCategories: unknownList,
|
||||
ScannerScores: map[string]float64{}, ScannerEvidence: map[string]string{},
|
||||
ScannerBackend: "qwen3guard-openai", ScannerVersion: "qwen3guard",
|
||||
PolicyID: "priority", PolicyVersion: 1,
|
||||
Decision: EventPass, RiskLevel: RiskLow, Action: ActionAllow,
|
||||
}
|
||||
score := 0.0
|
||||
if safety == "Controversial" {
|
||||
score = 0.5
|
||||
result.Decision, result.RiskLevel, result.Action = EventFlag, RiskMedium, ActionWarn
|
||||
}
|
||||
if safety == "Unsafe" {
|
||||
score = 1
|
||||
if len(matched) > 0 || len(unknownList) > 0 || len(knownList) == 0 {
|
||||
result.Decision, result.RiskLevel, result.Action = EventCritical, RiskCritical, ActionBlock
|
||||
} else {
|
||||
result.Decision, result.RiskLevel, result.Action = EventFlag, RiskHigh, ActionWarn
|
||||
}
|
||||
}
|
||||
for _, category := range matched {
|
||||
result.ScannerScores[category] = score
|
||||
result.ScannerEvidence[category] = ScannerCatalog[category].Label
|
||||
if safety == "Controversial" && isElevatedControversial(category) {
|
||||
result.Decision, result.RiskLevel, result.Action = EventCritical, RiskCritical, ActionBlock
|
||||
}
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func unknownCategoryID(value string) string {
|
||||
digest := sha256.Sum256([]byte(strings.TrimSpace(strings.ToLower(value))))
|
||||
return fmt.Sprintf("unknown:%x", digest[:8])
|
||||
}
|
||||
|
||||
func isElevatedControversial(category string) bool {
|
||||
return category == "jailbreak" || category == "pii" || category == "suicide_and_self_harm"
|
||||
}
|
||||
|
||||
type OpenAICompatibleScanner struct {
|
||||
clients sync.Map
|
||||
}
|
||||
|
||||
func NewOpenAICompatibleScanner() *OpenAICompatibleScanner { return &OpenAICompatibleScanner{} }
|
||||
|
||||
func (s *OpenAICompatibleScanner) Scan(ctx context.Context, endpoint ActiveEndpoint, chunk string, enabledScanners []string) (*NormalizedResult, error) {
|
||||
client, err := s.clientFor(endpoint)
|
||||
if err != nil {
|
||||
return nil, &GuardError{Code: ErrorCodeUnavailable, Cause: err}
|
||||
}
|
||||
requestURL, err := ChatCompletionsURL(endpoint.BaseURL)
|
||||
if err != nil {
|
||||
return nil, &GuardError{Code: ErrorCodeUnavailable, Cause: err}
|
||||
}
|
||||
payload := map[string]any{
|
||||
"model": endpoint.Model,
|
||||
"messages": []map[string]string{{"role": "user", "content": chunk}},
|
||||
"temperature": 0,
|
||||
"max_tokens": 64,
|
||||
"seed": 42,
|
||||
}
|
||||
body, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
return nil, &GuardError{Code: ErrorCodeInvalidResponse, Cause: err}
|
||||
}
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodPost, requestURL, bytes.NewReader(body))
|
||||
if err != nil {
|
||||
return nil, &GuardError{Code: ErrorCodeUnavailable, Cause: err}
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
if endpoint.Token != "" {
|
||||
req.Header.Set("Authorization", "Bearer "+endpoint.Token)
|
||||
}
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
timeout := errors.Is(err, context.DeadlineExceeded)
|
||||
var netErr net.Error
|
||||
if errors.As(err, &netErr) && netErr.Timeout() {
|
||||
timeout = true
|
||||
}
|
||||
return nil, &GuardError{Code: ErrorCodeUnavailable, Retryable: true, Timeout: timeout, Cause: err}
|
||||
}
|
||||
defer func() { _ = resp.Body.Close() }()
|
||||
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||
retryable := resp.StatusCode == http.StatusTooManyRequests || resp.StatusCode >= 500
|
||||
return nil, &GuardError{Code: ErrorCodeUnavailable, HTTPStatus: resp.StatusCode, Retryable: retryable}
|
||||
}
|
||||
limited := io.LimitReader(resp.Body, maxGuardResponseBytes+1)
|
||||
responseBody, err := io.ReadAll(limited)
|
||||
if err != nil {
|
||||
return nil, &GuardError{Code: ErrorCodeUnavailable, Retryable: true, Cause: err}
|
||||
}
|
||||
if int64(len(responseBody)) > maxGuardResponseBytes {
|
||||
return nil, &GuardError{Code: ErrorCodeInvalidResponse}
|
||||
}
|
||||
content, err := extractOpenAIContent(responseBody)
|
||||
if err != nil {
|
||||
return nil, &GuardError{Code: ErrorCodeInvalidResponse, Cause: err}
|
||||
}
|
||||
result, err := ParseQwen3Guard(content, enabledScanners)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result.GuardEndpointID = endpoint.ID
|
||||
result.ScannerVersion = endpoint.Model
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (s *OpenAICompatibleScanner) clientFor(endpoint ActiveEndpoint) (*http.Client, error) {
|
||||
key := fmt.Sprintf("%s|%s|%d", endpoint.ID, endpoint.BaseURL, endpoint.TimeoutMS)
|
||||
if cached, ok := s.clients.Load(key); ok {
|
||||
client, valid := cached.(*http.Client)
|
||||
if !valid {
|
||||
s.clients.Delete(key)
|
||||
return nil, errors.New("prompt guard client cache invalid")
|
||||
}
|
||||
return client, nil
|
||||
}
|
||||
client, err := NewSecureHTTPClient(endpoint)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
actual, _ := s.clients.LoadOrStore(key, client)
|
||||
actualClient, ok := actual.(*http.Client)
|
||||
if !ok {
|
||||
s.clients.Delete(key)
|
||||
return nil, errors.New("prompt guard client cache invalid")
|
||||
}
|
||||
return actualClient, nil
|
||||
}
|
||||
|
||||
func extractOpenAIContent(body []byte) (string, error) {
|
||||
var response struct {
|
||||
Choices []struct {
|
||||
Message struct {
|
||||
Content any `json:"content"`
|
||||
} `json:"message"`
|
||||
} `json:"choices"`
|
||||
}
|
||||
if err := json.Unmarshal(body, &response); err != nil || len(response.Choices) == 0 {
|
||||
return "", errors.New("prompt guard response envelope invalid")
|
||||
}
|
||||
content := response.Choices[0].Message.Content
|
||||
switch typed := content.(type) {
|
||||
case string:
|
||||
if strings.TrimSpace(typed) == "" {
|
||||
return "", errors.New("prompt guard response content empty")
|
||||
}
|
||||
return typed, nil
|
||||
case []any:
|
||||
parts := make([]string, 0, len(typed))
|
||||
for _, item := range typed {
|
||||
object, ok := item.(map[string]any)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if text, ok := object["text"].(string); ok && strings.TrimSpace(text) != "" {
|
||||
parts = append(parts, text)
|
||||
}
|
||||
}
|
||||
if len(parts) == 0 {
|
||||
return "", errors.New("prompt guard response content empty")
|
||||
}
|
||||
return strings.Join(parts, "\n"), nil
|
||||
default:
|
||||
return "", errors.New("prompt guard response content invalid")
|
||||
}
|
||||
}
|
||||
|
||||
func ScannerDefinitions() []ScannerDefinition {
|
||||
result := make([]ScannerDefinition, 0, len(AllScannerIDs))
|
||||
for _, id := range AllScannerIDs {
|
||||
result = append(result, ScannerCatalog[id])
|
||||
}
|
||||
sort.SliceStable(result, func(i, j int) bool { return i < j })
|
||||
return result
|
||||
}
|
||||
@@ -0,0 +1,134 @@
|
||||
package securityaudit
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestParseQwen3GuardStrictAndPolicy(t *testing.T) {
|
||||
tests := []struct {
|
||||
name, output string
|
||||
enabled []string
|
||||
decision EventDecision
|
||||
action Action
|
||||
wantErr bool
|
||||
}{
|
||||
{"safe", "Safety: Safe\nCategories: None", AllScannerIDs, EventPass, ActionAllow, false},
|
||||
{"controversial", "Safety: Controversial\nCategories: Violent", AllScannerIDs, EventFlag, ActionWarn, false},
|
||||
{"controversial pii escalates", "Safety: Controversial\nCategories: PII", AllScannerIDs, EventCritical, ActionBlock, false},
|
||||
{"unsafe", "Safety: Unsafe\nCategories: Jailbreak", AllScannerIDs, EventCritical, ActionBlock, false},
|
||||
{"unknown unsafe", "Safety: Unsafe\nCategories: Future Risk", AllScannerIDs, EventCritical, ActionBlock, false},
|
||||
{"disabled unsafe warns", "Safety: Unsafe\nCategories: Violent", []string{"PII"}, EventFlag, ActionWarn, false},
|
||||
{"extra explanation", "Safety: Safe\nCategories: None\nThis is safe", AllScannerIDs, "", "", true},
|
||||
{"duplicate", "Safety: Safe\nSafety: Safe", AllScannerIDs, "", "", true},
|
||||
{"duplicate categories", "Safety: Safe\nCategories: None\nCategories: PII", AllScannerIDs, "", "", true},
|
||||
{"missing categories", "Safety: Safe\n", AllScannerIDs, "", "", true},
|
||||
{"unknown safety", "Safety: Maybe\nCategories: PII", AllScannerIDs, "", "", true},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
result, err := ParseQwen3Guard(tt.output, tt.enabled)
|
||||
if tt.wantErr {
|
||||
require.Error(t, err)
|
||||
return
|
||||
}
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, tt.decision, result.Decision)
|
||||
require.Equal(t, tt.action, result.Action)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestQwen3GuardOfficialCategoriesAliasesAndUnknownAreStable(t *testing.T) {
|
||||
official := "Violent, Non-violent Illegal Acts, Sexual Content or Sexual Acts, PII, Suicide & Self-Harm, Unethical Acts, Politically Sensitive Topics, Copyright Violation, Jailbreak"
|
||||
result, err := ParseQwen3Guard("Safety: Unsafe\nCategories: "+official, AllScannerIDs)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, AllScannerIDs, result.MatchedScanners)
|
||||
require.Empty(t, result.UnknownCategories)
|
||||
require.Equal(t, "priority", result.PolicyID)
|
||||
require.Equal(t, 1, result.PolicyVersion)
|
||||
|
||||
aliases := map[string]string{
|
||||
"violence": "violent", "non_violent_illegal_acts": "non_violent_illegal_acts",
|
||||
"sexual": "sexual_content_or_sexual_acts", "personal identifiable information": "pii",
|
||||
"suicide/self harm": "suicide_and_self_harm", "unethical": "unethical_acts",
|
||||
"political": "politically_sensitive_topics", "copyright": "copyright_violation",
|
||||
"prompt injection": "jailbreak",
|
||||
}
|
||||
for alias, canonical := range aliases {
|
||||
require.Equal(t, canonical, NormalizeCategory(alias), alias)
|
||||
}
|
||||
|
||||
const canary = "PROMPT_CANARY_RAW_UNKNOWN_CATEGORY"
|
||||
unknown, err := ParseQwen3Guard("Safety: Unsafe\nCategories: "+canary, AllScannerIDs)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, unknown.UnknownCategories, 1)
|
||||
require.NotContains(t, unknown.UnknownCategories[0], "canary")
|
||||
require.NotContains(t, unknown.UnknownCategories[0], "raw")
|
||||
require.Contains(t, unknown.UnknownCategories[0], "unknown:")
|
||||
}
|
||||
|
||||
func TestExtractOpenAIContentSupportsStringAndTextBlocks(t *testing.T) {
|
||||
content, err := extractOpenAIContent([]byte(`{"choices":[{"message":{"content":"Safety: Safe\nCategories: None"}}]}`))
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "Safety: Safe\nCategories: None", content)
|
||||
content, err = extractOpenAIContent([]byte(`{"choices":[{"message":{"content":[{"type":"text","text":"Safety: Safe"},{"type":"text","text":"Categories: None"}]}}]}`))
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "Safety: Safe\nCategories: None", content)
|
||||
for _, body := range []string{`{}`, `{"choices":[]}`, `{"choices":[{"message":{"content":null}}]}`} {
|
||||
_, err := extractOpenAIContent([]byte(body))
|
||||
require.Error(t, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAggregateRequiresEveryResult(t *testing.T) {
|
||||
_, err := AggregateResults([]*NormalizedResult{{Decision: EventPass, Action: ActionAllow}, nil}, 0)
|
||||
require.Error(t, err)
|
||||
result, err := AggregateResults([]*NormalizedResult{
|
||||
{Decision: EventPass, RiskLevel: RiskLow, Action: ActionAllow, Categories: []string{"pii"}},
|
||||
{Decision: EventCritical, RiskLevel: RiskCritical, Action: ActionBlock, Categories: []string{"jailbreak"}},
|
||||
}, 0)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, EventCritical, result.Decision)
|
||||
require.Equal(t, ActionBlock, result.Action)
|
||||
require.Equal(t, []string{"pii", "jailbreak"}, result.Categories)
|
||||
}
|
||||
|
||||
func TestAggregateDeduplicatesFactsAndUsesMostSevereEndpointMetadata(t *testing.T) {
|
||||
result, err := AggregateResults([]*NormalizedResult{
|
||||
{Decision: EventPass, RiskLevel: RiskLow, Action: ActionAllow, Safety: "Safe", Categories: []string{"pii"}, MatchedScanners: []string{"pii"}, ScannerScores: map[string]float64{"pii": 0}, ScannerEvidence: map[string]string{"pii": "first"}, GuardEndpointID: "safe-node", ScannerVersion: "safe-version", PolicyID: "priority", PolicyVersion: 1},
|
||||
{Decision: EventCritical, RiskLevel: RiskCritical, Action: ActionBlock, Safety: "Unsafe", Categories: []string{"pii", "jailbreak"}, MatchedScanners: []string{"pii", "jailbreak"}, ScannerScores: map[string]float64{"pii": 1, "jailbreak": 1}, ScannerEvidence: map[string]string{"pii": "second", "jailbreak": "blocked"}, GuardEndpointID: "block-node", ScannerVersion: "block-version", PolicyID: "priority", PolicyVersion: 2},
|
||||
}, 7*time.Millisecond)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, []string{"pii", "jailbreak"}, result.Categories)
|
||||
require.Equal(t, []string{"pii", "jailbreak"}, result.MatchedScanners)
|
||||
require.Equal(t, "first", result.ScannerEvidence["pii"], "evidence is deterministically first-seen")
|
||||
require.Equal(t, "block-node", result.GuardEndpointID)
|
||||
require.Equal(t, "block-version", result.ScannerVersion)
|
||||
require.Equal(t, 2, result.PolicyVersion)
|
||||
require.Equal(t, 7, result.LatencyMS)
|
||||
}
|
||||
|
||||
func TestIssueSummariesAreDeterministicRedactedDerivedDTOs(t *testing.T) {
|
||||
const canary = "PROMPT_CANARY_EVIDENCE_SECRET"
|
||||
result := NormalizedResult{
|
||||
Decision: EventCritical, RiskLevel: RiskCritical, Action: ActionBlock,
|
||||
Categories: []string{"jailbreak", "pii"}, MatchedScanners: []string{"pii"},
|
||||
ScannerScores: map[string]float64{"pii": 1}, ScannerEvidence: map[string]string{"pii": canary},
|
||||
UnknownCategories: []string{unknownCategoryID("future risk")},
|
||||
}
|
||||
summaries := BuildIssueSummaries(result)
|
||||
require.Len(t, summaries, 3, "known categories are not hidden merely because policy disabled one")
|
||||
raw, err := json.Marshal(summaries)
|
||||
require.NoError(t, err)
|
||||
require.NotContains(t, string(raw), canary)
|
||||
for _, summary := range summaries {
|
||||
require.NotEmpty(t, summary.Title)
|
||||
require.NotEmpty(t, summary.Description)
|
||||
require.NotEmpty(t, summary.Code)
|
||||
require.NotEmpty(t, summary.EvidenceHash)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,433 @@
|
||||
package securityaudit
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
promptAuditAdmissionLockKey int64 = 579147893221901921
|
||||
promptAuditConfigLockKey int64 = 579147893221901922
|
||||
)
|
||||
|
||||
var (
|
||||
ErrQueueFull = errors.New("prompt audit queue full")
|
||||
ErrQueueAdmissionBusy = errors.New("prompt audit queue admission busy")
|
||||
ErrLeaseLost = errors.New("prompt audit worker lease lost")
|
||||
ErrEventNotFound = errors.New("prompt audit event not found")
|
||||
)
|
||||
|
||||
type Job struct {
|
||||
ID int64
|
||||
Snapshot PromptSnapshot
|
||||
ExecutionMode Mode
|
||||
ConfigVersion int64
|
||||
Status string
|
||||
Attempts int
|
||||
MaxAttempts int
|
||||
ClaimVersion int64
|
||||
NextAttemptAt time.Time
|
||||
ProcessingStartedAt *time.Time
|
||||
ProcessedAt *time.Time
|
||||
LastErrorCode string
|
||||
LastErrorMessage string
|
||||
CreatedAt time.Time
|
||||
UpdatedAt time.Time
|
||||
}
|
||||
|
||||
type Event struct {
|
||||
ID int64 `json:"id"`
|
||||
JobID int64 `json:"job_id"`
|
||||
Snapshot PromptSnapshot `json:"snapshot"`
|
||||
Decision EventDecision `json:"decision"`
|
||||
RiskLevel RiskLevel `json:"risk_level"`
|
||||
Action Action `json:"action"`
|
||||
Categories []string `json:"categories"`
|
||||
MatchedScanners []string `json:"matched_scanners"`
|
||||
ScannerScores map[string]float64 `json:"scanner_scores"`
|
||||
ScannerEvidence map[string]string `json:"scanner_evidence"`
|
||||
ScannerBackend string `json:"scanner_backend"`
|
||||
ScannerVersion string `json:"scanner_version"`
|
||||
GuardEndpointID string `json:"guard_endpoint_id"`
|
||||
PolicyID string `json:"policy_id"`
|
||||
PolicyVersion int `json:"policy_version"`
|
||||
ConfigVersion int64 `json:"config_version"`
|
||||
ChunkTotal int `json:"chunk_total"`
|
||||
LatencyMS int `json:"latency_ms"`
|
||||
IssueSummaries []IssueSummary `json:"issue_summaries"`
|
||||
CreatedAt time.Time `json:"created_at"`
|
||||
}
|
||||
|
||||
type JobRepository interface {
|
||||
CreateStagingWithCapacity(ctx context.Context, snapshot PromptSnapshot, configVersion int64, maxAttempts, capacity int) (*Job, error)
|
||||
PublishQueued(ctx context.Context, jobID int64) error
|
||||
MarkStagingFailed(ctx context.Context, jobID int64, code, message string) error
|
||||
ClaimNextJob(ctx context.Context, now time.Time) (*Job, bool, error)
|
||||
RefreshLease(ctx context.Context, jobID, claimVersion int64, now time.Time) error
|
||||
Complete(ctx context.Context, job *Job, result *NormalizedResult, storePass bool) (*Event, error)
|
||||
Retry(ctx context.Context, jobID, claimVersion int64, next time.Time, code, message string) error
|
||||
Fail(ctx context.Context, jobID, claimVersion int64, code, message string) error
|
||||
ReclaimStale(ctx context.Context, stagingBefore, processingBefore time.Time, limit int) (int64, error)
|
||||
QueueStats(ctx context.Context) (QueueStats, error)
|
||||
RecordBlocking(ctx context.Context, snapshot PromptSnapshot, configVersion int64, result *NormalizedResult, storePass bool) (*Event, error)
|
||||
}
|
||||
|
||||
type PostgreSQLRepository struct {
|
||||
db *sql.DB
|
||||
clock Clock
|
||||
}
|
||||
|
||||
func NewPostgreSQLRepository(db *sql.DB) *PostgreSQLRepository {
|
||||
return &PostgreSQLRepository{db: db, clock: realClock{}}
|
||||
}
|
||||
|
||||
func (r *PostgreSQLRepository) CreateStagingWithCapacity(ctx context.Context, snapshot PromptSnapshot, configVersion int64, maxAttempts, capacity int) (*Job, error) {
|
||||
if r == nil || r.db == nil {
|
||||
return nil, errors.New("prompt audit database unavailable")
|
||||
}
|
||||
tx, err := r.db.BeginTx(ctx, &sql.TxOptions{Isolation: sql.LevelReadCommitted})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer func() { _ = tx.Rollback() }()
|
||||
var locked bool
|
||||
if err := tx.QueryRowContext(ctx, `SELECT pg_try_advisory_xact_lock($1)`, promptAuditAdmissionLockKey).Scan(&locked); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if !locked {
|
||||
return nil, ErrQueueAdmissionBusy
|
||||
}
|
||||
var active int
|
||||
if err := tx.QueryRowContext(ctx, `
|
||||
SELECT COUNT(*) FROM prompt_audit_jobs
|
||||
WHERE status IN ('staging','queued','processing','retry')`).Scan(&active); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if capacity <= 0 || active >= capacity {
|
||||
return nil, ErrQueueFull
|
||||
}
|
||||
if maxAttempts <= 0 {
|
||||
maxAttempts = 3
|
||||
}
|
||||
job, err := insertJob(ctx, tx, snapshot.Redacted(), ModeAsync, configVersion, "staging", maxAttempts)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := tx.Commit(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return job, nil
|
||||
}
|
||||
|
||||
func (r *PostgreSQLRepository) PublishQueued(ctx context.Context, jobID int64) error {
|
||||
result, err := r.db.ExecContext(ctx, `
|
||||
UPDATE prompt_audit_jobs SET status='queued', next_attempt_at=NOW(), updated_at=NOW()
|
||||
WHERE id=$1 AND status='staging'`, jobID)
|
||||
return requireOneRow(result, err, ErrLeaseLost)
|
||||
}
|
||||
|
||||
func (r *PostgreSQLRepository) MarkStagingFailed(ctx context.Context, jobID int64, code, _ string) error {
|
||||
code, message := sanitizeStoredError(code)
|
||||
result, err := r.db.ExecContext(ctx, `
|
||||
UPDATE prompt_audit_jobs
|
||||
SET status='failed', processed_at=NOW(), updated_at=NOW(), last_error_code=$2, last_error_message=$3
|
||||
WHERE id=$1 AND status='staging'`, jobID, code, message)
|
||||
return requireOneRow(result, err, ErrLeaseLost)
|
||||
}
|
||||
|
||||
func (r *PostgreSQLRepository) ClaimNextJob(ctx context.Context, now time.Time) (*Job, bool, error) {
|
||||
row := r.db.QueryRowContext(ctx, `
|
||||
WITH candidate AS (
|
||||
SELECT id FROM prompt_audit_jobs
|
||||
WHERE status IN ('queued','retry') AND next_attempt_at <= $1
|
||||
ORDER BY next_attempt_at, id
|
||||
FOR UPDATE SKIP LOCKED
|
||||
LIMIT 1
|
||||
)
|
||||
UPDATE prompt_audit_jobs AS j
|
||||
SET status='processing', attempts=j.attempts+1, claim_version=j.claim_version+1,
|
||||
processing_started_at=$1, updated_at=$1
|
||||
FROM candidate
|
||||
WHERE j.id=candidate.id
|
||||
RETURNING `+jobColumns("j"), now.UTC())
|
||||
job, err := scanJob(row)
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return nil, false, nil
|
||||
}
|
||||
return job, err == nil, err
|
||||
}
|
||||
|
||||
func (r *PostgreSQLRepository) RefreshLease(ctx context.Context, jobID, claimVersion int64, now time.Time) error {
|
||||
result, err := r.db.ExecContext(ctx, `
|
||||
UPDATE prompt_audit_jobs SET processing_started_at=$3, updated_at=$3
|
||||
WHERE id=$1 AND status='processing' AND claim_version=$2`, jobID, claimVersion, now.UTC())
|
||||
return requireOneRow(result, err, ErrLeaseLost)
|
||||
}
|
||||
|
||||
func (r *PostgreSQLRepository) Complete(ctx context.Context, job *Job, result *NormalizedResult, storePass bool) (*Event, error) {
|
||||
if job == nil || result == nil {
|
||||
return nil, errors.New("prompt audit completion requires job and result")
|
||||
}
|
||||
tx, err := r.db.BeginTx(ctx, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer func() { _ = tx.Rollback() }()
|
||||
updateResult, err := tx.ExecContext(ctx, `
|
||||
UPDATE prompt_audit_jobs SET status='done', processed_at=NOW(), updated_at=NOW(),
|
||||
last_error_code='', last_error_message=''
|
||||
WHERE id=$1 AND status='processing' AND claim_version=$2`, job.ID, job.ClaimVersion)
|
||||
if err := requireOneRow(updateResult, err, ErrLeaseLost); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var event *Event
|
||||
if storePass || result.Decision != EventPass {
|
||||
event, err = insertEvent(ctx, tx, job.ID, job.Snapshot.Redacted(), job.ConfigVersion, result)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
if err := tx.Commit(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return event, nil
|
||||
}
|
||||
|
||||
func (r *PostgreSQLRepository) Retry(ctx context.Context, jobID, claimVersion int64, next time.Time, code, _ string) error {
|
||||
code, message := sanitizeStoredError(code)
|
||||
result, err := r.db.ExecContext(ctx, `
|
||||
UPDATE prompt_audit_jobs SET status='retry', next_attempt_at=$3, processing_started_at=NULL,
|
||||
updated_at=NOW(), last_error_code=$4, last_error_message=$5
|
||||
WHERE id=$1 AND status='processing' AND claim_version=$2`,
|
||||
jobID, claimVersion, next.UTC(), code, message)
|
||||
return requireOneRow(result, err, ErrLeaseLost)
|
||||
}
|
||||
|
||||
func (r *PostgreSQLRepository) Fail(ctx context.Context, jobID, claimVersion int64, code, _ string) error {
|
||||
code, message := sanitizeStoredError(code)
|
||||
result, err := r.db.ExecContext(ctx, `
|
||||
UPDATE prompt_audit_jobs SET status='failed', processed_at=NOW(), processing_started_at=NULL,
|
||||
updated_at=NOW(), last_error_code=$3, last_error_message=$4
|
||||
WHERE id=$1 AND status='processing' AND claim_version=$2`,
|
||||
jobID, claimVersion, code, message)
|
||||
return requireOneRow(result, err, ErrLeaseLost)
|
||||
}
|
||||
|
||||
func (r *PostgreSQLRepository) ReclaimStale(ctx context.Context, stagingBefore, processingBefore time.Time, limit int) (int64, error) {
|
||||
if limit <= 0 || limit > 1000 {
|
||||
limit = 100
|
||||
}
|
||||
result, err := r.db.ExecContext(ctx, `
|
||||
WITH stale AS (
|
||||
SELECT id FROM prompt_audit_jobs
|
||||
WHERE (status='staging' AND updated_at < $1)
|
||||
OR (status='processing' AND processing_started_at < $2)
|
||||
ORDER BY updated_at, id FOR UPDATE SKIP LOCKED LIMIT $3
|
||||
)
|
||||
UPDATE prompt_audit_jobs AS j
|
||||
SET status=CASE
|
||||
WHEN j.status='staging' THEN 'failed'
|
||||
WHEN j.attempts < j.max_attempts THEN 'retry'
|
||||
ELSE 'failed' END,
|
||||
next_attempt_at=CASE WHEN j.status='processing' AND j.attempts < j.max_attempts THEN NOW() ELSE j.next_attempt_at END,
|
||||
processing_started_at=NULL,
|
||||
processed_at=CASE WHEN j.status='staging' OR j.attempts >= j.max_attempts THEN NOW() ELSE NULL END,
|
||||
last_error_code=CASE WHEN j.status='staging' THEN 'staging_timeout' ELSE 'processing_lease_expired' END,
|
||||
last_error_message='', updated_at=NOW()
|
||||
FROM stale WHERE j.id=stale.id`, stagingBefore.UTC(), processingBefore.UTC(), limit)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return result.RowsAffected()
|
||||
}
|
||||
|
||||
func (r *PostgreSQLRepository) QueueStats(ctx context.Context) (QueueStats, error) {
|
||||
rows, err := r.db.QueryContext(ctx, `SELECT status, COUNT(*) FROM prompt_audit_jobs GROUP BY status`)
|
||||
if err != nil {
|
||||
return QueueStats{}, err
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
var stats QueueStats
|
||||
for rows.Next() {
|
||||
var status string
|
||||
var count int64
|
||||
if err := rows.Scan(&status, &count); err != nil {
|
||||
return QueueStats{}, err
|
||||
}
|
||||
switch status {
|
||||
case "staging":
|
||||
stats.Staging = count
|
||||
case "queued":
|
||||
stats.Queued = count
|
||||
case "processing":
|
||||
stats.Processing = count
|
||||
case "retry":
|
||||
stats.Retry = count
|
||||
case "done":
|
||||
stats.Done = count
|
||||
case "failed":
|
||||
stats.Failed = count
|
||||
}
|
||||
}
|
||||
stats.Active = stats.Staging + stats.Queued + stats.Processing + stats.Retry
|
||||
return stats, rows.Err()
|
||||
}
|
||||
|
||||
func (r *PostgreSQLRepository) RecordBlocking(ctx context.Context, snapshot PromptSnapshot, configVersion int64, result *NormalizedResult, storePass bool) (*Event, error) {
|
||||
if result == nil {
|
||||
return nil, errors.New("prompt guard result required")
|
||||
}
|
||||
tx, err := r.db.BeginTx(ctx, nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer func() { _ = tx.Rollback() }()
|
||||
job, err := insertJob(ctx, tx, snapshot.Redacted(), ModeBlocking, configVersion, "done", 1)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var event *Event
|
||||
if storePass || result.Decision != EventPass {
|
||||
event, err = insertEvent(ctx, tx, job.ID, snapshot.Redacted(), configVersion, result)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
if err := tx.Commit(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return event, nil
|
||||
}
|
||||
|
||||
type sqlQueryer interface {
|
||||
QueryRowContext(context.Context, string, ...any) *sql.Row
|
||||
}
|
||||
|
||||
func insertJob(ctx context.Context, queryer sqlQueryer, snapshot PromptSnapshot, mode Mode, configVersion int64, status string, maxAttempts int) (*Job, error) {
|
||||
processedExpr := "NULL"
|
||||
if status == "done" || status == "failed" {
|
||||
processedExpr = "NOW()"
|
||||
}
|
||||
row := queryer.QueryRowContext(ctx, `
|
||||
INSERT INTO prompt_audit_jobs (
|
||||
request_id,user_id,username_snapshot,user_email_snapshot,api_key_id,api_key_name_snapshot,
|
||||
group_id,group_name,provider,endpoint,protocol,model,prompt_hash,redacted_preview,
|
||||
prompt_length,message_count,execution_mode,config_version,status,max_attempts,processed_at
|
||||
) VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,$19,$20,`+processedExpr+`)
|
||||
RETURNING `+jobColumns("prompt_audit_jobs"),
|
||||
snapshot.RequestID, nullableID(snapshot.UserID), snapshot.UsernameSnapshot, snapshot.UserEmailSnapshot,
|
||||
nullableID(snapshot.APIKeyID), snapshot.APIKeyNameSnapshot, snapshot.GroupID, snapshot.GroupName,
|
||||
snapshot.Provider, snapshot.Endpoint, snapshot.Protocol, snapshot.Model, snapshot.PromptHash,
|
||||
snapshot.RedactedPreview, snapshot.PromptLength, snapshot.MessageCount, string(mode), configVersion,
|
||||
status, maxAttempts)
|
||||
return scanJob(row)
|
||||
}
|
||||
|
||||
func insertEvent(ctx context.Context, queryer sqlQueryer, jobID int64, snapshot PromptSnapshot, configVersion int64, result *NormalizedResult) (*Event, error) {
|
||||
categories, _ := json.Marshal(result.Categories)
|
||||
matched, _ := json.Marshal(result.MatchedScanners)
|
||||
scores, _ := json.Marshal(result.ScannerScores)
|
||||
evidence := make(map[string]string, len(result.ScannerEvidence))
|
||||
for key, value := range result.ScannerEvidence {
|
||||
evidence[key] = RedactPreview(value, 160)
|
||||
}
|
||||
evidenceJSON, _ := json.Marshal(evidence)
|
||||
row := queryer.QueryRowContext(ctx, `
|
||||
INSERT INTO prompt_audit_events (
|
||||
job_id,request_id,user_id,username_snapshot,user_email_snapshot,api_key_id,api_key_name_snapshot,
|
||||
group_id,group_name,provider,endpoint,protocol,model,prompt_hash,redacted_preview,
|
||||
decision,risk_level,action,categories,matched_scanners,scanner_scores,scanner_evidence,
|
||||
scanner_backend,scanner_version,guard_endpoint_id,policy_id,policy_version,config_version,chunk_total,latency_ms
|
||||
) VALUES ($1,$2,$3,$4,$5,$6,$7,$8,$9,$10,$11,$12,$13,$14,$15,$16,$17,$18,
|
||||
$19::jsonb,$20::jsonb,$21::jsonb,$22::jsonb,$23,$24,$25,$26,$27,$28,$29,$30)
|
||||
RETURNING `+eventColumns("prompt_audit_events"),
|
||||
jobID, snapshot.RequestID, nullableID(snapshot.UserID), snapshot.UsernameSnapshot, snapshot.UserEmailSnapshot,
|
||||
nullableID(snapshot.APIKeyID), snapshot.APIKeyNameSnapshot, snapshot.GroupID, snapshot.GroupName,
|
||||
snapshot.Provider, snapshot.Endpoint, snapshot.Protocol, snapshot.Model, snapshot.PromptHash,
|
||||
snapshot.RedactedPreview, string(result.Decision), string(result.RiskLevel), string(result.Action),
|
||||
categories, matched, scores, evidenceJSON, result.ScannerBackend, result.ScannerVersion,
|
||||
result.GuardEndpointID, result.PolicyID, result.PolicyVersion, configVersion, result.ChunkTotal, result.LatencyMS)
|
||||
return scanEvent(row)
|
||||
}
|
||||
|
||||
type rowScanner interface{ Scan(...any) error }
|
||||
|
||||
func scanJob(row rowScanner) (*Job, error) {
|
||||
job := &Job{}
|
||||
var userID, apiKeyID, groupID sql.NullInt64
|
||||
var processingStarted, processed sql.NullTime
|
||||
err := row.Scan(
|
||||
&job.ID, &job.Snapshot.RequestID, &userID, &job.Snapshot.UsernameSnapshot, &job.Snapshot.UserEmailSnapshot,
|
||||
&apiKeyID, &job.Snapshot.APIKeyNameSnapshot, &groupID, &job.Snapshot.GroupName, &job.Snapshot.Provider,
|
||||
&job.Snapshot.Endpoint, &job.Snapshot.Protocol, &job.Snapshot.Model, &job.Snapshot.PromptHash,
|
||||
&job.Snapshot.RedactedPreview, &job.Snapshot.PromptLength, &job.Snapshot.MessageCount, &job.ExecutionMode,
|
||||
&job.ConfigVersion, &job.Status, &job.Attempts, &job.MaxAttempts, &job.ClaimVersion,
|
||||
&job.NextAttemptAt, &processingStarted, &processed, &job.LastErrorCode, &job.LastErrorMessage,
|
||||
&job.CreatedAt, &job.UpdatedAt,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
job.Snapshot.UserID = nullableInt64Value(userID)
|
||||
job.Snapshot.APIKeyID = nullableInt64Value(apiKeyID)
|
||||
job.Snapshot.GroupID = nullableInt64Ptr(groupID)
|
||||
if processingStarted.Valid {
|
||||
value := processingStarted.Time
|
||||
job.ProcessingStartedAt = &value
|
||||
}
|
||||
if processed.Valid {
|
||||
value := processed.Time
|
||||
job.ProcessedAt = &value
|
||||
}
|
||||
return job, nil
|
||||
}
|
||||
|
||||
func jobColumns(alias string) string {
|
||||
return fmt.Sprintf(`%[1]s.id,%[1]s.request_id,%[1]s.user_id,%[1]s.username_snapshot,%[1]s.user_email_snapshot,
|
||||
%[1]s.api_key_id,%[1]s.api_key_name_snapshot,%[1]s.group_id,%[1]s.group_name,%[1]s.provider,
|
||||
%[1]s.endpoint,%[1]s.protocol,%[1]s.model,%[1]s.prompt_hash,%[1]s.redacted_preview,
|
||||
%[1]s.prompt_length,%[1]s.message_count,%[1]s.execution_mode,%[1]s.config_version,%[1]s.status,
|
||||
%[1]s.attempts,%[1]s.max_attempts,%[1]s.claim_version,%[1]s.next_attempt_at,
|
||||
%[1]s.processing_started_at,%[1]s.processed_at,%[1]s.last_error_code,%[1]s.last_error_message,
|
||||
%[1]s.created_at,%[1]s.updated_at`, alias)
|
||||
}
|
||||
|
||||
func requireOneRow(result sql.Result, err error, missing error) error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
rows, err := result.RowsAffected()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if rows != 1 {
|
||||
return missing
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func nullableID(value int64) any {
|
||||
if value <= 0 {
|
||||
return nil
|
||||
}
|
||||
return value
|
||||
}
|
||||
|
||||
func nullableInt64Value(value sql.NullInt64) int64 {
|
||||
if !value.Valid {
|
||||
return 0
|
||||
}
|
||||
return value.Int64
|
||||
}
|
||||
|
||||
func nullableInt64Ptr(value sql.NullInt64) *int64 {
|
||||
if !value.Valid {
|
||||
return nil
|
||||
}
|
||||
result := value.Int64
|
||||
return &result
|
||||
}
|
||||
@@ -0,0 +1,449 @@
|
||||
package securityaudit
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
_ "github.com/lib/pq"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
const promptAuditPostgresTestEnv = "PROMPT_AUDIT_TEST_POSTGRES_DSN"
|
||||
|
||||
func openPromptAuditIntegrationDB(t *testing.T) *sql.DB {
|
||||
t.Helper()
|
||||
dsn := strings.TrimSpace(os.Getenv(promptAuditPostgresTestEnv))
|
||||
if dsn == "" {
|
||||
t.Skip(promptAuditPostgresTestEnv + " is not set")
|
||||
}
|
||||
db, err := sql.Open("postgres", dsn)
|
||||
require.NoError(t, err)
|
||||
db.SetMaxOpenConns(16)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer cancel()
|
||||
require.NoError(t, db.PingContext(ctx))
|
||||
_, err = db.ExecContext(ctx, `
|
||||
CREATE TABLE IF NOT EXISTS users (id BIGSERIAL PRIMARY KEY);
|
||||
CREATE TABLE IF NOT EXISTS groups (id BIGSERIAL PRIMARY KEY);
|
||||
CREATE TABLE IF NOT EXISTS api_keys (id BIGSERIAL PRIMARY KEY);
|
||||
CREATE TABLE IF NOT EXISTS settings (
|
||||
key VARCHAR(255) PRIMARY KEY,
|
||||
value TEXT NOT NULL DEFAULT '',
|
||||
updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW()
|
||||
);
|
||||
`)
|
||||
require.NoError(t, err)
|
||||
migrationPath := filepath.Join("..", "..", "migrations", "181_prompt_audit.sql")
|
||||
migration, err := os.ReadFile(migrationPath)
|
||||
require.NoError(t, err)
|
||||
// The migration runner can retry an interrupted deployment; the migration
|
||||
// must therefore be safe to execute more than once.
|
||||
_, err = db.ExecContext(ctx, string(migration))
|
||||
require.NoError(t, err)
|
||||
_, err = db.ExecContext(ctx, string(migration))
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(func() { require.NoError(t, db.Close()) })
|
||||
resetPromptAuditIntegrationDB(t, db)
|
||||
return db
|
||||
}
|
||||
|
||||
func resetPromptAuditIntegrationDB(t *testing.T, db *sql.DB) {
|
||||
t.Helper()
|
||||
_, err := db.Exec(`TRUNCATE TABLE prompt_audit_events, prompt_audit_jobs, api_keys, users, groups, settings RESTART IDENTITY CASCADE`)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
func insertIdentity(t *testing.T, db *sql.DB, table string) int64 {
|
||||
t.Helper()
|
||||
var id int64
|
||||
require.NoError(t, db.QueryRow(`INSERT INTO `+table+` DEFAULT VALUES RETURNING id`).Scan(&id))
|
||||
return id
|
||||
}
|
||||
|
||||
func integrationSnapshot(seed string) PromptSnapshot {
|
||||
return PromptSnapshot{
|
||||
RequestID: "request-" + seed, UsernameSnapshot: "user-" + seed,
|
||||
UserEmailSnapshot: "user-" + seed + "@example.test", APIKeyNameSnapshot: "key-" + seed,
|
||||
GroupName: "group-" + seed, Provider: "openai", Endpoint: "/v1/chat/completions",
|
||||
Protocol: "openai_chat", Model: "gpt-test", PromptHash: strings.Repeat(seed[:1], 64),
|
||||
RedactedPreview: "redacted-" + seed, PromptLength: len([]rune(seed)), MessageCount: 1,
|
||||
}
|
||||
}
|
||||
|
||||
func integrationResult(decision EventDecision) *NormalizedResult {
|
||||
result := &NormalizedResult{
|
||||
Decision: decision, RiskLevel: RiskLow, Action: ActionAllow, Safety: "Safe",
|
||||
Categories: []string{}, MatchedScanners: []string{}, ScannerScores: map[string]float64{},
|
||||
ScannerEvidence: map[string]string{}, ScannerBackend: "qwen3guard-openai",
|
||||
ScannerVersion: "test", GuardEndpointID: "guard-1", PolicyID: "priority",
|
||||
PolicyVersion: 1, ChunkTotal: 1, LatencyMS: 2,
|
||||
}
|
||||
if decision != EventPass {
|
||||
result.RiskLevel = RiskCritical
|
||||
result.Action = ActionBlock
|
||||
result.Safety = "Unsafe"
|
||||
result.Categories = []string{"pii"}
|
||||
result.MatchedScanners = []string{"pii"}
|
||||
result.ScannerScores["pii"] = 1
|
||||
result.ScannerEvidence["pii"] = "redacted evidence"
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func TestPromptAuditMigrationSchemaAndLeakageGate(t *testing.T) {
|
||||
db := openPromptAuditIntegrationDB(t)
|
||||
ctx := context.Background()
|
||||
|
||||
rows, err := db.QueryContext(ctx, `SELECT table_name, column_name FROM information_schema.columns
|
||||
WHERE table_schema='public' AND table_name IN ('prompt_audit_jobs','prompt_audit_events')`)
|
||||
require.NoError(t, err)
|
||||
defer func() { _ = rows.Close() }()
|
||||
forbidden := []string{"raw_prompt", "raw_request", "payload", "token", "authorization", "credential", "ciphertext"}
|
||||
for rows.Next() {
|
||||
var tableName, columnName string
|
||||
require.NoError(t, rows.Scan(&tableName, &columnName))
|
||||
lower := strings.ToLower(columnName)
|
||||
for _, word := range forbidden {
|
||||
require.NotContainsf(t, lower, word, "%s.%s is a forbidden raw/credential column", tableName, columnName)
|
||||
}
|
||||
}
|
||||
require.NoError(t, rows.Err())
|
||||
|
||||
indexRows, err := db.QueryContext(ctx, `SELECT indexname FROM pg_indexes
|
||||
WHERE schemaname='public' AND tablename IN ('prompt_audit_jobs','prompt_audit_events')`)
|
||||
require.NoError(t, err)
|
||||
defer func() { _ = indexRows.Close() }()
|
||||
indexes := map[string]bool{}
|
||||
for indexRows.Next() {
|
||||
var name string
|
||||
require.NoError(t, indexRows.Scan(&name))
|
||||
indexes[name] = true
|
||||
}
|
||||
for _, name := range []string{
|
||||
"idx_prompt_audit_jobs_schedule", "idx_prompt_audit_jobs_request", "idx_prompt_audit_jobs_user_created",
|
||||
"idx_prompt_audit_jobs_api_key_created", "idx_prompt_audit_jobs_group_created", "idx_prompt_audit_jobs_prompt_hash",
|
||||
"idx_prompt_audit_jobs_created", "idx_prompt_audit_events_job", "idx_prompt_audit_events_request",
|
||||
"idx_prompt_audit_events_decision_created", "idx_prompt_audit_events_risk_created",
|
||||
"idx_prompt_audit_events_user_created", "idx_prompt_audit_events_api_key_created",
|
||||
"idx_prompt_audit_events_group_created", "idx_prompt_audit_events_prompt_hash", "idx_prompt_audit_events_created",
|
||||
} {
|
||||
require.Truef(t, indexes[name], "missing index %s", name)
|
||||
}
|
||||
|
||||
_, err = db.ExecContext(ctx, `INSERT INTO prompt_audit_jobs(status) VALUES ('unknown')`)
|
||||
require.Error(t, err)
|
||||
_, err = db.ExecContext(ctx, `INSERT INTO prompt_audit_jobs(prompt_length) VALUES (-1)`)
|
||||
require.Error(t, err)
|
||||
var jobID int64
|
||||
require.NoError(t, db.QueryRowContext(ctx, `INSERT INTO prompt_audit_jobs DEFAULT VALUES RETURNING id`).Scan(&jobID))
|
||||
_, err = db.ExecContext(ctx, `INSERT INTO prompt_audit_events(job_id,chunk_total) VALUES ($1,-1)`, jobID)
|
||||
require.Error(t, err)
|
||||
}
|
||||
|
||||
func TestPromptAuditDatabaseAndAdminJSONNeverPersistCanaryPromptOrRawErrors(t *testing.T) {
|
||||
db := openPromptAuditIntegrationDB(t)
|
||||
repo := NewPostgreSQLRepository(db)
|
||||
ctx := context.Background()
|
||||
const promptCanary = "PROMPT_AUDIT_CANARY_SECRET_DO_NOT_PERSIST"
|
||||
request := Request{
|
||||
RequestID: "canary-request", Provider: "openai",
|
||||
Endpoint: "/v1/chat/completions", Protocol: "openai_chat", Model: "gpt-test", Stage: "http",
|
||||
Body: []byte(`{"messages":[{"role":"user","content":"` + promptCanary + `"}]}`),
|
||||
}
|
||||
snapshot, err := ExtractPromptSnapshot(request)
|
||||
require.NoError(t, err)
|
||||
require.NotContains(t, snapshot.RedactedPreview, promptCanary)
|
||||
event, err := repo.RecordBlocking(ctx, snapshot.Redacted(), 1, integrationResult(EventCritical), true)
|
||||
require.NoError(t, err)
|
||||
adminJSON, err := json.Marshal(event)
|
||||
require.NoError(t, err)
|
||||
require.NotContains(t, string(adminJSON), promptCanary)
|
||||
|
||||
var jobJSON string
|
||||
require.NoError(t, db.QueryRow(`SELECT row_to_json(j)::text FROM prompt_audit_jobs j WHERE id=$1`, event.JobID).Scan(&jobJSON))
|
||||
require.NotContains(t, jobJSON, promptCanary)
|
||||
|
||||
failedJob, err := repo.CreateStagingWithCapacity(ctx, integrationSnapshot("error"), 1, 3, 10)
|
||||
require.NoError(t, err)
|
||||
const errorCanary = "GUARD_RAW_RESPONSE_CANARY_SECRET"
|
||||
require.NoError(t, repo.MarkStagingFailed(ctx, failedJob.ID, "payload_store_failed", "raw guard body: "+errorCanary))
|
||||
var code, message string
|
||||
require.NoError(t, db.QueryRow(`SELECT last_error_code,last_error_message FROM prompt_audit_jobs WHERE id=$1`, failedJob.ID).Scan(&code, &message))
|
||||
require.Equal(t, "payload_store_failed", code)
|
||||
require.Equal(t, stableErrorMessage(code), message)
|
||||
require.NotContains(t, message, errorCanary)
|
||||
require.LessOrEqual(t, len([]rune(message)), 160)
|
||||
}
|
||||
|
||||
func TestPromptAuditRepositoryAdmissionClaimFencingAndEventTransaction(t *testing.T) {
|
||||
db := openPromptAuditIntegrationDB(t)
|
||||
repo := NewPostgreSQLRepository(db)
|
||||
ctx := context.Background()
|
||||
|
||||
start := make(chan struct{})
|
||||
type admissionResult struct {
|
||||
job *Job
|
||||
err error
|
||||
}
|
||||
results := make(chan admissionResult, 2)
|
||||
var wg sync.WaitGroup
|
||||
for i := 0; i < 2; i++ {
|
||||
wg.Add(1)
|
||||
go func(index int) {
|
||||
defer wg.Done()
|
||||
<-start
|
||||
job, err := repo.CreateStagingWithCapacity(ctx, integrationSnapshot(string(rune('a'+index))), 1, 3, 1)
|
||||
results <- admissionResult{job: job, err: err}
|
||||
}(i)
|
||||
}
|
||||
close(start)
|
||||
wg.Wait()
|
||||
close(results)
|
||||
var accepted *Job
|
||||
rejected := 0
|
||||
for result := range results {
|
||||
if result.err == nil {
|
||||
require.Nil(t, accepted)
|
||||
accepted = result.job
|
||||
continue
|
||||
}
|
||||
require.True(t, errors.Is(result.err, ErrQueueFull) || errors.Is(result.err, ErrQueueAdmissionBusy))
|
||||
rejected++
|
||||
}
|
||||
require.NotNil(t, accepted)
|
||||
require.Equal(t, 1, rejected)
|
||||
stats, err := repo.QueueStats(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, int64(1), stats.Active)
|
||||
require.NoError(t, repo.PublishQueued(ctx, accepted.ID))
|
||||
|
||||
claimStart := make(chan struct{})
|
||||
claims := make(chan *Job, 2)
|
||||
for i := 0; i < 2; i++ {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
<-claimStart
|
||||
job, claimed, claimErr := repo.ClaimNextJob(ctx, time.Now().Add(time.Second))
|
||||
require.NoError(t, claimErr)
|
||||
if claimed {
|
||||
claims <- job
|
||||
}
|
||||
}()
|
||||
}
|
||||
close(claimStart)
|
||||
wg.Wait()
|
||||
close(claims)
|
||||
claimedJobs := make([]*Job, 0, 1)
|
||||
for job := range claims {
|
||||
claimedJobs = append(claimedJobs, job)
|
||||
}
|
||||
require.Len(t, claimedJobs, 1)
|
||||
firstClaim := claimedJobs[0]
|
||||
require.Equal(t, int64(1), firstClaim.ClaimVersion)
|
||||
|
||||
reclaimed, err := repo.ReclaimStale(ctx, time.Now().Add(time.Hour), time.Now().Add(time.Hour), 10)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, int64(1), reclaimed)
|
||||
secondClaim, claimed, err := repo.ClaimNextJob(ctx, time.Now().Add(time.Second))
|
||||
require.NoError(t, err)
|
||||
require.True(t, claimed)
|
||||
require.Greater(t, secondClaim.ClaimVersion, firstClaim.ClaimVersion)
|
||||
require.ErrorIs(t, repo.RefreshLease(ctx, firstClaim.ID, firstClaim.ClaimVersion, time.Now()), ErrLeaseLost)
|
||||
_, err = repo.Complete(ctx, firstClaim, integrationResult(EventCritical), true)
|
||||
require.ErrorIs(t, err, ErrLeaseLost)
|
||||
|
||||
event, err := repo.Complete(ctx, secondClaim, integrationResult(EventCritical), true)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, event)
|
||||
var status string
|
||||
var eventCount int
|
||||
require.NoError(t, db.QueryRow(`SELECT status FROM prompt_audit_jobs WHERE id=$1`, secondClaim.ID).Scan(&status))
|
||||
require.Equal(t, "done", status)
|
||||
require.NoError(t, db.QueryRow(`SELECT COUNT(*) FROM prompt_audit_events WHERE job_id=$1`, secondClaim.ID).Scan(&eventCount))
|
||||
require.Equal(t, 1, eventCount)
|
||||
|
||||
staging, err := repo.CreateStagingWithCapacity(ctx, integrationSnapshot("stale"), 1, 3, 10)
|
||||
require.NoError(t, err)
|
||||
reclaimed, err = repo.ReclaimStale(ctx, time.Now().Add(time.Hour), time.Now().Add(time.Hour), 10)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, int64(1), reclaimed)
|
||||
require.NoError(t, db.QueryRow(`SELECT status FROM prompt_audit_jobs WHERE id=$1`, staging.ID).Scan(&status))
|
||||
require.Equal(t, "failed", status)
|
||||
}
|
||||
|
||||
func TestPromptAuditRepositoryForeignKeysFiltersAndStableIdentitySnapshots(t *testing.T) {
|
||||
db := openPromptAuditIntegrationDB(t)
|
||||
repo := NewPostgreSQLRepository(db)
|
||||
ctx := context.Background()
|
||||
userID := insertIdentity(t, db, "users")
|
||||
apiKeyID := insertIdentity(t, db, "api_keys")
|
||||
groupID := insertIdentity(t, db, "groups")
|
||||
snapshot := integrationSnapshot("identity")
|
||||
snapshot.UserID, snapshot.APIKeyID, snapshot.GroupID = userID, apiKeyID, &groupID
|
||||
event, err := repo.RecordBlocking(ctx, snapshot, 7, integrationResult(EventCritical), true)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, event)
|
||||
|
||||
start, end := time.Now().Add(-time.Hour), time.Now().Add(time.Hour)
|
||||
page, err := repo.ListEvents(ctx, EventFilter{
|
||||
Decision: string(EventCritical), RiskLevel: string(RiskCritical), Endpoint: snapshot.Endpoint,
|
||||
GroupID: &groupID, UserID: &userID, APIKeyID: &apiKeyID, RequestID: snapshot.RequestID,
|
||||
PromptHash: snapshot.PromptHash, Keyword: snapshot.UsernameSnapshot, StartAt: &start, EndAt: &end,
|
||||
}, 1, 10)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, int64(1), page.Total)
|
||||
require.Len(t, page.Items, 1)
|
||||
require.NotEmpty(t, page.Items[0].IssueSummaries)
|
||||
require.Equal(t, snapshot.UsernameSnapshot, page.Items[0].Snapshot.UsernameSnapshot)
|
||||
require.Equal(t, snapshot.UserEmailSnapshot, page.Items[0].Snapshot.UserEmailSnapshot)
|
||||
require.Equal(t, snapshot.APIKeyNameSnapshot, page.Items[0].Snapshot.APIKeyNameSnapshot)
|
||||
|
||||
_, err = db.Exec(`DELETE FROM users WHERE id=$1`, userID)
|
||||
require.NoError(t, err)
|
||||
_, err = db.Exec(`DELETE FROM api_keys WHERE id=$1`, apiKeyID)
|
||||
require.NoError(t, err)
|
||||
_, err = db.Exec(`DELETE FROM groups WHERE id=$1`, groupID)
|
||||
require.NoError(t, err)
|
||||
stored, err := repo.GetEvent(ctx, event.ID)
|
||||
require.NoError(t, err)
|
||||
require.Zero(t, stored.Snapshot.UserID)
|
||||
require.Zero(t, stored.Snapshot.APIKeyID)
|
||||
require.Nil(t, stored.Snapshot.GroupID)
|
||||
require.Equal(t, snapshot.UsernameSnapshot, stored.Snapshot.UsernameSnapshot)
|
||||
require.Equal(t, snapshot.UserEmailSnapshot, stored.Snapshot.UserEmailSnapshot)
|
||||
require.Equal(t, snapshot.APIKeyNameSnapshot, stored.Snapshot.APIKeyNameSnapshot)
|
||||
|
||||
_, err = db.Exec(`DELETE FROM prompt_audit_jobs WHERE id=$1`, event.JobID)
|
||||
require.NoError(t, err)
|
||||
_, err = repo.GetEvent(ctx, event.ID)
|
||||
require.ErrorIs(t, err, ErrEventNotFound)
|
||||
}
|
||||
|
||||
func TestPromptAuditRepositoryHighWaterAndSafeDeletion(t *testing.T) {
|
||||
db := openPromptAuditIntegrationDB(t)
|
||||
repo := NewPostgreSQLRepository(db)
|
||||
ctx := context.Background()
|
||||
first, err := repo.RecordBlocking(ctx, integrationSnapshot("first"), 1, integrationResult(EventCritical), true)
|
||||
require.NoError(t, err)
|
||||
second, err := repo.RecordBlocking(ctx, integrationSnapshot("second"), 1, integrationResult(EventCritical), true)
|
||||
require.NoError(t, err)
|
||||
start, end := time.Now().Add(-time.Hour), time.Now().Add(time.Hour)
|
||||
filter := EventFilter{Decision: string(EventCritical), StartAt: &start, EndAt: &end}
|
||||
preview, err := repo.PreviewDelete(ctx, filter)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, int64(2), preview.MatchedCount)
|
||||
require.Equal(t, second.ID, preview.SnapshotMaxID)
|
||||
require.Equal(t, FilterHash(preview.FilterSummary, preview.SnapshotMaxID), preview.FilterHash)
|
||||
|
||||
newer, err := repo.RecordBlocking(ctx, integrationSnapshot("newer"), 1, integrationResult(EventCritical), true)
|
||||
require.NoError(t, err)
|
||||
result, err := repo.DeleteEventsByFilter(ctx, filter, preview.SnapshotMaxID, 1)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, int64(2), result.DeletedEvents)
|
||||
require.Equal(t, int64(2), result.DeletedJobs)
|
||||
_, err = repo.GetEvent(ctx, first.ID)
|
||||
require.ErrorIs(t, err, ErrEventNotFound)
|
||||
_, err = repo.GetEvent(ctx, second.ID)
|
||||
require.ErrorIs(t, err, ErrEventNotFound)
|
||||
_, err = repo.GetEvent(ctx, newer.ID)
|
||||
require.NoError(t, err, "an event created after preview must survive high-water deletion")
|
||||
|
||||
processingEvent, err := repo.RecordBlocking(ctx, integrationSnapshot("processing"), 1, integrationResult(EventCritical), true)
|
||||
require.NoError(t, err)
|
||||
_, err = db.Exec(`UPDATE prompt_audit_jobs SET status='processing' WHERE id=$1`, processingEvent.JobID)
|
||||
require.NoError(t, err)
|
||||
deleteResult, err := repo.DeleteEvent(ctx, processingEvent.ID)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, int64(1), deleteResult.DeletedEvents)
|
||||
require.Zero(t, deleteResult.DeletedJobs)
|
||||
var remaining int
|
||||
require.NoError(t, db.QueryRow(`SELECT COUNT(*) FROM prompt_audit_jobs WHERE id=$1`, processingEvent.JobID).Scan(&remaining))
|
||||
require.Equal(t, 1, remaining, "processing jobs must not be deleted as orphans")
|
||||
|
||||
batchOne, err := repo.RecordBlocking(ctx, integrationSnapshot("batch-one"), 1, integrationResult(EventCritical), true)
|
||||
require.NoError(t, err)
|
||||
batchTwo, err := repo.RecordBlocking(ctx, integrationSnapshot("batch-two"), 1, integrationResult(EventCritical), true)
|
||||
require.NoError(t, err)
|
||||
ids := []int64{batchTwo.ID, batchOne.ID, batchOne.ID}
|
||||
sort.Slice(ids, func(i, j int) bool { return ids[i] > ids[j] })
|
||||
batchResult, err := repo.DeleteEventsByIDs(ctx, ids)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, int64(2), batchResult.DeletedEvents)
|
||||
}
|
||||
|
||||
func TestPromptAuditServiceConfirmationKeepsPostPreviewEventsAndConcurrentDeletesAreSafe(t *testing.T) {
|
||||
db := openPromptAuditIntegrationDB(t)
|
||||
repo := NewPostgreSQLRepository(db)
|
||||
ctx := context.Background()
|
||||
now := time.Now().UTC()
|
||||
start, end := now.Add(-time.Hour), now.Add(time.Hour)
|
||||
filter := EventFilter{Decision: string(EventCritical), StartAt: &start, EndAt: &end}
|
||||
|
||||
for i := 0; i < 12; i++ {
|
||||
_, err := repo.RecordBlocking(ctx, integrationSnapshot(fmt.Sprintf("event-%02d", i)), 1, integrationResult(EventCritical), true)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
service := &PromptService{
|
||||
config: &fakeConfigStore{}, repo: repo, payload: NewRedisPayloadStore(nil), clock: fixedClock{now: now},
|
||||
}
|
||||
preview, err := service.PreviewDelete(ctx, filter, 77)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, int64(12), preview.MatchedCount)
|
||||
|
||||
newer, err := repo.RecordBlocking(ctx, integrationSnapshot("post-preview"), 1, integrationResult(EventCritical), true)
|
||||
require.NoError(t, err)
|
||||
result, err := service.DeleteByFilter(ctx, DeleteByFilterRequest{
|
||||
Filter: filter, SnapshotMaxID: preview.SnapshotMaxID, FilterHash: preview.FilterHash,
|
||||
ConfirmationToken: preview.ConfirmationToken, Confirm: true,
|
||||
}, 77)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, int64(12), result.DeletedEvents)
|
||||
_, err = repo.GetEvent(ctx, newer.ID)
|
||||
require.NoError(t, err, "events created after delete-preview must survive")
|
||||
|
||||
resetPromptAuditIntegrationDB(t, db)
|
||||
for i := 0; i < 24; i++ {
|
||||
_, err := repo.RecordBlocking(ctx, integrationSnapshot(fmt.Sprintf("race-%02d", i)), 1, integrationResult(EventCritical), true)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
preview, err = repo.PreviewDelete(ctx, filter)
|
||||
require.NoError(t, err)
|
||||
|
||||
type deleteOutcome struct {
|
||||
result *DeleteResult
|
||||
err error
|
||||
}
|
||||
outcomes := make(chan deleteOutcome, 2)
|
||||
var wg sync.WaitGroup
|
||||
for i := 0; i < 2; i++ {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
deleted, deleteErr := repo.DeleteEventsByFilter(ctx, filter, preview.SnapshotMaxID, 1)
|
||||
outcomes <- deleteOutcome{result: deleted, err: deleteErr}
|
||||
}()
|
||||
}
|
||||
wg.Wait()
|
||||
close(outcomes)
|
||||
var deletedTotal int64
|
||||
for outcome := range outcomes {
|
||||
require.NoError(t, outcome.err)
|
||||
require.NotNil(t, outcome.result)
|
||||
deletedTotal += outcome.result.DeletedEvents
|
||||
}
|
||||
require.Equal(t, int64(24), deletedTotal, "concurrent deleters must neither double-count nor strand matching events")
|
||||
remaining, err := repo.ListEvents(ctx, filter, 1, 100)
|
||||
require.NoError(t, err)
|
||||
require.Zero(t, remaining.Total)
|
||||
}
|
||||
@@ -0,0 +1,121 @@
|
||||
package securityaudit
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"sort"
|
||||
"time"
|
||||
)
|
||||
|
||||
func SplitRunes(value string, limit int) []string {
|
||||
if limit <= 0 {
|
||||
return nil
|
||||
}
|
||||
runes := []rune(value)
|
||||
if len(runes) == 0 {
|
||||
return nil
|
||||
}
|
||||
chunks := make([]string, 0, (len(runes)+limit-1)/limit)
|
||||
for start := 0; start < len(runes); start += limit {
|
||||
end := start + limit
|
||||
if end > len(runes) {
|
||||
end = len(runes)
|
||||
}
|
||||
chunks = append(chunks, string(runes[start:end]))
|
||||
}
|
||||
return chunks
|
||||
}
|
||||
|
||||
func AggregateResults(results []*NormalizedResult, latency time.Duration) (*NormalizedResult, error) {
|
||||
if len(results) == 0 {
|
||||
return nil, errors.New("prompt guard produced no complete result")
|
||||
}
|
||||
aggregated := &NormalizedResult{
|
||||
Decision: EventPass, RiskLevel: RiskLow, Action: ActionAllow,
|
||||
ScannerBackend: "qwen3guard-openai", Categories: []string{}, MatchedScanners: []string{},
|
||||
ScannerScores: map[string]float64{}, ScannerEvidence: map[string]string{}, ChunkTotal: len(results),
|
||||
LatencyMS: int(latency.Milliseconds()),
|
||||
}
|
||||
categories := map[string]struct{}{}
|
||||
matched := map[string]struct{}{}
|
||||
unknown := map[string]struct{}{}
|
||||
for _, result := range results {
|
||||
if result == nil {
|
||||
return nil, errors.New("prompt guard partial result is not allowed")
|
||||
}
|
||||
if resultSeverity(result.Decision) > resultSeverity(aggregated.Decision) {
|
||||
aggregated.Decision = result.Decision
|
||||
aggregated.RiskLevel = result.RiskLevel
|
||||
aggregated.Action = result.Action
|
||||
aggregated.Safety = result.Safety
|
||||
aggregated.GuardEndpointID = result.GuardEndpointID
|
||||
aggregated.ScannerVersion = result.ScannerVersion
|
||||
aggregated.PolicyID = result.PolicyID
|
||||
aggregated.PolicyVersion = result.PolicyVersion
|
||||
}
|
||||
if aggregated.GuardEndpointID == "" {
|
||||
aggregated.GuardEndpointID = result.GuardEndpointID
|
||||
aggregated.ScannerVersion = result.ScannerVersion
|
||||
aggregated.PolicyID = result.PolicyID
|
||||
aggregated.PolicyVersion = result.PolicyVersion
|
||||
}
|
||||
for _, category := range result.Categories {
|
||||
categories[category] = struct{}{}
|
||||
}
|
||||
for _, scanner := range result.MatchedScanners {
|
||||
matched[scanner] = struct{}{}
|
||||
}
|
||||
for scanner, score := range result.ScannerScores {
|
||||
if score > aggregated.ScannerScores[scanner] {
|
||||
aggregated.ScannerScores[scanner] = score
|
||||
}
|
||||
}
|
||||
for scanner, evidence := range result.ScannerEvidence {
|
||||
if _, exists := aggregated.ScannerEvidence[scanner]; !exists {
|
||||
aggregated.ScannerEvidence[scanner] = RedactPreview(evidence, 160)
|
||||
}
|
||||
}
|
||||
for _, category := range result.UnknownCategories {
|
||||
unknown[category] = struct{}{}
|
||||
}
|
||||
}
|
||||
aggregated.Categories = orderedScannerKeys(categories)
|
||||
aggregated.MatchedScanners = orderedScannerKeys(matched)
|
||||
aggregated.UnknownCategories = sortedKeys(unknown)
|
||||
return aggregated, nil
|
||||
}
|
||||
|
||||
func resultSeverity(decision EventDecision) int {
|
||||
switch decision {
|
||||
case EventCritical:
|
||||
return 3
|
||||
case EventFlag:
|
||||
return 2
|
||||
default:
|
||||
return 1
|
||||
}
|
||||
}
|
||||
|
||||
func sortedKeys(values map[string]struct{}) []string {
|
||||
result := make([]string, 0, len(values))
|
||||
for key := range values {
|
||||
result = append(result, key)
|
||||
}
|
||||
sort.Strings(result)
|
||||
return result
|
||||
}
|
||||
|
||||
func orderedScannerKeys(values map[string]struct{}) []string {
|
||||
result := make([]string, 0, len(values))
|
||||
remaining := make(map[string]struct{}, len(values))
|
||||
for key := range values {
|
||||
remaining[key] = struct{}{}
|
||||
}
|
||||
for _, scannerID := range AllScannerIDs {
|
||||
if _, ok := remaining[scannerID]; ok {
|
||||
result = append(result, scannerID)
|
||||
delete(remaining, scannerID)
|
||||
}
|
||||
}
|
||||
result = append(result, sortedKeys(remaining)...)
|
||||
return result
|
||||
}
|
||||
@@ -0,0 +1,469 @@
|
||||
package securityaudit
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
type PromptService struct {
|
||||
config ConfigStore
|
||||
repo *PostgreSQLRepository
|
||||
payload *RedisPayloadStore
|
||||
enqueuer *Enqueuer
|
||||
runner *Runner
|
||||
evaluator *GuardEvaluator
|
||||
scanner *OpenAICompatibleScanner
|
||||
metrics *AtomicMetrics
|
||||
clock Clock
|
||||
|
||||
lifecycleMu sync.Mutex
|
||||
cancel context.CancelFunc
|
||||
background context.Context
|
||||
enqueueWG sync.WaitGroup
|
||||
enqueueSlots chan struct{}
|
||||
probeMu sync.RWMutex
|
||||
probes map[string]ProbeResult
|
||||
}
|
||||
|
||||
func NewPromptService(
|
||||
config ConfigStore,
|
||||
repo *PostgreSQLRepository,
|
||||
payload *RedisPayloadStore,
|
||||
scanner *OpenAICompatibleScanner,
|
||||
metrics *AtomicMetrics,
|
||||
) *PromptService {
|
||||
enqueuer := NewEnqueuer(config, repo, payload, metrics)
|
||||
evaluator := NewGuardEvaluator(scanner, repo, metrics)
|
||||
runner := NewRunner(config, repo, payload, scanner, metrics)
|
||||
return &PromptService{
|
||||
config: config, repo: repo, payload: payload, scanner: scanner, metrics: metrics,
|
||||
enqueuer: enqueuer, evaluator: evaluator, runner: runner, clock: realClock{},
|
||||
enqueueSlots: make(chan struct{}, 128), probes: map[string]ProbeResult{},
|
||||
}
|
||||
}
|
||||
|
||||
func (s *PromptService) Start(ctx context.Context) error {
|
||||
if s == nil || s.config == nil || s.runner == nil {
|
||||
return errors.New("prompt audit service unavailable")
|
||||
}
|
||||
s.lifecycleMu.Lock()
|
||||
if s.cancel != nil {
|
||||
s.lifecycleMu.Unlock()
|
||||
return nil
|
||||
}
|
||||
background, cancel := context.WithCancel(ctx)
|
||||
s.background, s.cancel = background, cancel
|
||||
s.lifecycleMu.Unlock()
|
||||
configErr := s.config.Start(background)
|
||||
workerErr := s.runner.Start(background)
|
||||
return errors.Join(configErr, workerErr)
|
||||
}
|
||||
|
||||
func (s *PromptService) Shutdown(ctx context.Context) error {
|
||||
if s == nil {
|
||||
return nil
|
||||
}
|
||||
s.lifecycleMu.Lock()
|
||||
cancel := s.cancel
|
||||
s.cancel = nil
|
||||
s.lifecycleMu.Unlock()
|
||||
if cancel != nil {
|
||||
cancel()
|
||||
}
|
||||
var workerErr error
|
||||
if s.runner != nil {
|
||||
workerErr = s.runner.Shutdown(ctx)
|
||||
}
|
||||
done := make(chan struct{})
|
||||
go func() { s.enqueueWG.Wait(); close(done) }()
|
||||
select {
|
||||
case <-done:
|
||||
case <-ctx.Done():
|
||||
if workerErr == nil {
|
||||
workerErr = ctx.Err()
|
||||
}
|
||||
}
|
||||
var configErr error
|
||||
if s.config != nil {
|
||||
configErr = s.config.Shutdown(ctx)
|
||||
}
|
||||
if workerErr != nil {
|
||||
return workerErr
|
||||
}
|
||||
return configErr
|
||||
}
|
||||
|
||||
func (s *PromptService) EffectiveMode() Mode {
|
||||
if s == nil || s.config == nil {
|
||||
return ModeOff
|
||||
}
|
||||
return s.config.EffectiveMode()
|
||||
}
|
||||
|
||||
func (s *PromptService) Enqueue(_ context.Context, req Request) error {
|
||||
if s == nil || s.enqueuer == nil || s.EffectiveMode() != ModeAsync {
|
||||
return nil
|
||||
}
|
||||
select {
|
||||
case s.enqueueSlots <- struct{}{}:
|
||||
default:
|
||||
if s.metrics != nil {
|
||||
s.metrics.IncDropped()
|
||||
}
|
||||
LogWarn(EventEnqueueDropped, map[string]any{"request_id": req.RequestID, "status": "dropped", "error_code": "local_enqueue_busy"})
|
||||
return nil
|
||||
}
|
||||
s.lifecycleMu.Lock()
|
||||
background := s.background
|
||||
s.lifecycleMu.Unlock()
|
||||
if background == nil {
|
||||
<-s.enqueueSlots
|
||||
return errors.New("prompt audit service not started")
|
||||
}
|
||||
requestCopy := req.Clone()
|
||||
s.enqueueWG.Add(1)
|
||||
go func() {
|
||||
defer s.enqueueWG.Done()
|
||||
defer func() { <-s.enqueueSlots }()
|
||||
ctx, cancel := context.WithTimeout(background, 2*time.Second)
|
||||
defer cancel()
|
||||
_ = s.enqueuer.Enqueue(ctx, requestCopy)
|
||||
}()
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *PromptService) Evaluate(ctx context.Context, req Request) (*PromptDecision, error) {
|
||||
if s == nil || s.config == nil || s.evaluator == nil {
|
||||
return nil, &GuardError{Code: ErrorCodeUnavailable}
|
||||
}
|
||||
cfg, ok := s.config.Active()
|
||||
if !ok {
|
||||
if s.config.EffectiveMode() == ModeBlocking {
|
||||
return nil, &GuardError{Code: ErrorCodeUnavailable}
|
||||
}
|
||||
return &PromptDecision{Kind: DecisionAllow, AllowNextStage: true}, nil
|
||||
}
|
||||
if cfg.EffectiveMode() != ModeBlocking || !cfg.IncludesGroup(req.GroupID) {
|
||||
return &PromptDecision{Kind: DecisionAllow, AllowNextStage: true}, nil
|
||||
}
|
||||
snapshot, err := ExtractPromptSnapshot(req)
|
||||
if errors.Is(err, ErrNoPromptText) {
|
||||
return &PromptDecision{Kind: DecisionAllow, AllowNextStage: true}, nil
|
||||
}
|
||||
if err != nil {
|
||||
return nil, &GuardError{Code: ErrorCodeInvalidResponse, Cause: err}
|
||||
}
|
||||
return s.evaluator.Evaluate(ctx, cfg, snapshot)
|
||||
}
|
||||
|
||||
func (s *PromptService) GetConfig() PublicConfig { return s.config.Public() }
|
||||
|
||||
func (s *PromptService) SaveConfig(ctx context.Context, req UpdateConfigRequest, actorID int64) (PublicConfig, error) {
|
||||
return s.config.Save(ctx, req, actorID)
|
||||
}
|
||||
|
||||
func (s *PromptService) Runtime(ctx context.Context) RuntimeSnapshot {
|
||||
expected, activeVersion, loadedAt, loadError := s.config.RuntimeState()
|
||||
cfg, hasConfig := s.config.Active()
|
||||
mode := ModeOff
|
||||
workerTotal, queueCapacity := 0, 0
|
||||
if hasConfig {
|
||||
mode, workerTotal, queueCapacity = cfg.EffectiveMode(), cfg.WorkerCount, cfg.QueueCapacity
|
||||
}
|
||||
runtime := RuntimeSnapshot{
|
||||
ProcessStatus: "disabled", EffectiveMode: mode, ExpectedConfigVersion: expected,
|
||||
ActiveConfigVersion: activeVersion, ConfigLoadedAt: loadedAt, ConfigLoadError: loadError,
|
||||
WorkerTotal: workerTotal, QueueCapacity: queueCapacity, DatabaseStatus: "ok", RedisStatus: "ok",
|
||||
Endpoints: s.probeSnapshot(), GuardMetrics: s.metrics.Snapshot(),
|
||||
}
|
||||
if s.repo != nil {
|
||||
stats, err := s.repo.QueueStats(ctx)
|
||||
if err != nil {
|
||||
runtime.DatabaseStatus = "error"
|
||||
runtime.LastErrorCode = "database_unavailable"
|
||||
} else {
|
||||
runtime.Queue = stats
|
||||
}
|
||||
} else {
|
||||
runtime.DatabaseStatus = "error"
|
||||
}
|
||||
if s.payload == nil || s.payload.Ping(ctx) != nil {
|
||||
runtime.RedisStatus = "error"
|
||||
if runtime.LastErrorCode == "" {
|
||||
runtime.LastErrorCode = "payload_store_unavailable"
|
||||
}
|
||||
}
|
||||
activeWorkers, processed, failed, heartbeat, lastProcessed, workerCode, workerMessage := s.runner.Snapshot()
|
||||
runtime.WorkerActive, runtime.ProcessedTotal, runtime.FailedTotal = activeWorkers, processed, failed
|
||||
if s.metrics != nil {
|
||||
auditMetrics := s.metrics.AuditSnapshot()
|
||||
runtime.EnqueuedTotal, runtime.DroppedTotal = auditMetrics.Enqueued, auditMetrics.Dropped
|
||||
}
|
||||
runtime.WorkerHeartbeatAt, runtime.LastProcessedAt = heartbeat, lastProcessed
|
||||
if workerCode != "" {
|
||||
runtime.LastErrorCode, runtime.LastErrorMessage = workerCode, workerMessage
|
||||
}
|
||||
if mode != ModeOff {
|
||||
runtime.ProcessStatus = "running"
|
||||
if loadError != "" || runtime.DatabaseStatus != "ok" || runtime.RedisStatus != "ok" || activeVersion != expected {
|
||||
runtime.ProcessStatus = "degraded"
|
||||
}
|
||||
if heartbeat == nil || s.clock.Now().Sub(*heartbeat) > 10*time.Second {
|
||||
runtime.ProcessStatus = "degraded"
|
||||
}
|
||||
}
|
||||
return runtime
|
||||
}
|
||||
|
||||
type ProbeRequest struct {
|
||||
Endpoint UpdateEndpoint `json:"endpoint"`
|
||||
}
|
||||
|
||||
func (s *PromptService) Probe(ctx context.Context, request ProbeRequest) ProbeResult {
|
||||
started := s.clock.Now()
|
||||
endpoint, tokenApplied, err := s.resolveProbeEndpoint(request.Endpoint)
|
||||
if err != nil {
|
||||
return s.finishProbe(request.Endpoint.ID, started, ProbeResult{Status: "failed", ErrorCode: "endpoint_invalid", Message: "审计节点配置无效"})
|
||||
}
|
||||
LogInfo(EventProbeStarted, map[string]any{"guard_endpoint_id": endpoint.ID, "status": "started"})
|
||||
client, err := NewSecureHTTPClient(endpoint)
|
||||
if err != nil {
|
||||
return s.finishProbe(endpoint.ID, started, ProbeResult{Status: "failed", ErrorCode: "endpoint_unsafe", Message: "审计节点地址不在允许范围", TokenApplied: tokenApplied})
|
||||
}
|
||||
modelsURL, _ := ModelsURL(endpoint.BaseURL)
|
||||
req, err := http.NewRequestWithContext(ctx, http.MethodGet, modelsURL, nil)
|
||||
if err != nil {
|
||||
return s.finishProbe(endpoint.ID, started, ProbeResult{Status: "failed", ErrorCode: "probe_request_invalid", Message: "无法创建探测请求", TokenApplied: tokenApplied})
|
||||
}
|
||||
if endpoint.Token != "" {
|
||||
req.Header.Set("Authorization", "Bearer "+endpoint.Token)
|
||||
}
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
code := "connection_failed"
|
||||
var netErr net.Error
|
||||
if errors.Is(err, context.DeadlineExceeded) || (errors.As(err, &netErr) && netErr.Timeout()) {
|
||||
code = "timeout"
|
||||
}
|
||||
return s.finishProbe(endpoint.ID, started, ProbeResult{Status: "failed", ErrorCode: code, Message: "无法连接审计节点", Retryable: true, TokenApplied: tokenApplied})
|
||||
}
|
||||
responseBody, readErr := io.ReadAll(io.LimitReader(resp.Body, maxGuardResponseBytes+1))
|
||||
_ = resp.Body.Close()
|
||||
if readErr != nil {
|
||||
return s.finishProbe(endpoint.ID, started, ProbeResult{Status: "failed", ErrorCode: "response_read_failed", Message: "审计节点响应读取失败", HTTPStatus: resp.StatusCode, Retryable: true, TokenApplied: tokenApplied})
|
||||
}
|
||||
if int64(len(responseBody)) > maxGuardResponseBytes {
|
||||
return s.finishProbe(endpoint.ID, started, ProbeResult{Status: "failed", ErrorCode: "response_too_large", Message: "审计节点响应无效", HTTPStatus: resp.StatusCode, TokenApplied: tokenApplied})
|
||||
}
|
||||
if resp.StatusCode >= 200 && resp.StatusCode < 300 && modelsResponseReady(responseBody, endpoint.Model) {
|
||||
return s.finishProbe(endpoint.ID, started, ProbeResult{OK: true, Status: "healthy", Message: "审计节点连接正常", HTTPStatus: resp.StatusCode, TokenApplied: tokenApplied})
|
||||
}
|
||||
if (resp.StatusCode >= 200 && resp.StatusCode < 300) || resp.StatusCode == http.StatusNotFound || resp.StatusCode == http.StatusMethodNotAllowed {
|
||||
result, scanErr := s.scanner.Scan(ctx, endpoint, "Hello", AllScannerIDs)
|
||||
if scanErr == nil && result != nil {
|
||||
return s.finishProbe(endpoint.ID, started, ProbeResult{OK: true, Status: "healthy", Message: "审计节点模型调用正常", HTTPStatus: http.StatusOK, TokenApplied: tokenApplied})
|
||||
}
|
||||
code, status, retryable := guardErrorCode(scanErr), 0, false
|
||||
var guardErr *GuardError
|
||||
if errors.As(scanErr, &guardErr) {
|
||||
status, retryable = guardErr.HTTPStatus, guardErr.Retryable
|
||||
}
|
||||
if code == "" {
|
||||
code = ErrorCodeInvalidResponse
|
||||
}
|
||||
return s.finishProbe(endpoint.ID, started, ProbeResult{Status: "failed", ErrorCode: code, Message: "审计节点模型调用失败", HTTPStatus: status, Retryable: retryable, TokenApplied: tokenApplied})
|
||||
}
|
||||
code, retryable := "probe_http_error", resp.StatusCode == 429 || resp.StatusCode >= 500
|
||||
if resp.StatusCode == 401 || resp.StatusCode == 403 {
|
||||
code = "authentication_failed"
|
||||
}
|
||||
return s.finishProbe(endpoint.ID, started, ProbeResult{Status: "failed", ErrorCode: code, Message: "审计节点探测失败", HTTPStatus: resp.StatusCode, Retryable: retryable, TokenApplied: tokenApplied})
|
||||
}
|
||||
|
||||
func modelsResponseReady(body []byte, model string) bool {
|
||||
var response struct {
|
||||
Data []struct {
|
||||
ID string `json:"id"`
|
||||
} `json:"data"`
|
||||
}
|
||||
if json.Unmarshal(body, &response) != nil || response.Data == nil {
|
||||
return false
|
||||
}
|
||||
model = strings.TrimSpace(model)
|
||||
if model == "" {
|
||||
return true
|
||||
}
|
||||
for _, item := range response.Data {
|
||||
if strings.TrimSpace(item.ID) == model {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func (s *PromptService) resolveProbeEndpoint(input UpdateEndpoint) (ActiveEndpoint, bool, error) {
|
||||
token := strings.TrimSpace(input.Token)
|
||||
if token == "" {
|
||||
if cfg, ok := s.config.Active(); ok {
|
||||
for _, endpoint := range cfg.Endpoints {
|
||||
if endpoint.ID == strings.TrimSpace(input.ID) {
|
||||
token = endpoint.Token
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
baseURL, err := NormalizeBaseURL(input.BaseURL)
|
||||
if err != nil {
|
||||
return ActiveEndpoint{}, false, err
|
||||
}
|
||||
model := strings.TrimSpace(input.Model)
|
||||
if model == "" {
|
||||
model = DefaultGuardModel
|
||||
}
|
||||
timeout := input.TimeoutMS
|
||||
if timeout == 0 {
|
||||
timeout = DefaultTimeoutMS
|
||||
}
|
||||
limit := input.InputLimit
|
||||
if limit == 0 {
|
||||
limit = DefaultInputLimit
|
||||
}
|
||||
storage := storageConfig{Enabled: false, Strategy: "priority", WorkerCount: DefaultWorkerCount, QueueCapacity: DefaultQueueCapacity, Scanners: append([]string(nil), AllScannerIDs...), AllGroups: true,
|
||||
Endpoints: []StorageEndpoint{{ID: strings.TrimSpace(input.ID), Name: strings.TrimSpace(input.Name), Protocol: "openai_compatible", BaseURL: baseURL, Model: model, TimeoutMS: timeout, InputLimit: limit}}}
|
||||
if storage.Endpoints[0].ID == "" {
|
||||
storage.Endpoints[0].ID = "probe"
|
||||
}
|
||||
if storage.Endpoints[0].Name == "" {
|
||||
storage.Endpoints[0].Name = "Probe"
|
||||
}
|
||||
if err := validateStorageConfig(storage); err != nil {
|
||||
return ActiveEndpoint{}, false, err
|
||||
}
|
||||
return ActiveEndpoint{ID: storage.Endpoints[0].ID, Name: storage.Endpoints[0].Name, Protocol: "openai_compatible", BaseURL: baseURL, Model: model, Token: token, TimeoutMS: timeout, InputLimit: limit, Enabled: true}, token != "", nil
|
||||
}
|
||||
|
||||
func (s *PromptService) finishProbe(id string, started time.Time, result ProbeResult) ProbeResult {
|
||||
result.CheckedAt = s.clock.Now()
|
||||
result.LatencyMS = int(result.CheckedAt.Sub(started).Milliseconds())
|
||||
if result.OK {
|
||||
LogInfo(EventProbeFinished, map[string]any{"guard_endpoint_id": id, "status": result.Status, "latency_ms": result.LatencyMS, "http_status": result.HTTPStatus})
|
||||
} else {
|
||||
LogWarn(EventProbeFailed, map[string]any{"guard_endpoint_id": id, "status": result.Status, "latency_ms": result.LatencyMS, "http_status": result.HTTPStatus, "error_code": result.ErrorCode, "retryable": result.Retryable})
|
||||
}
|
||||
s.probeMu.Lock()
|
||||
s.probes[id] = result
|
||||
s.probeMu.Unlock()
|
||||
return result
|
||||
}
|
||||
|
||||
func (s *PromptService) probeSnapshot() map[string]ProbeResult {
|
||||
s.probeMu.RLock()
|
||||
defer s.probeMu.RUnlock()
|
||||
result := make(map[string]ProbeResult, len(s.probes))
|
||||
for id, probe := range s.probes {
|
||||
result[id] = probe
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func (s *PromptService) ListEvents(ctx context.Context, filter EventFilter, page, pageSize int) (*EventPage, error) {
|
||||
return s.repo.ListEvents(ctx, filter, page, pageSize)
|
||||
}
|
||||
func (s *PromptService) GetEvent(ctx context.Context, id int64) (*Event, error) {
|
||||
return s.repo.GetEvent(ctx, id)
|
||||
}
|
||||
|
||||
func (s *PromptService) DeleteEvent(ctx context.Context, id int64) (*DeleteResult, error) {
|
||||
result, err := s.repo.DeleteEvent(ctx, id)
|
||||
if err == nil {
|
||||
s.deletePayloads(ctx, result.JobIDs)
|
||||
}
|
||||
return result, err
|
||||
}
|
||||
func (s *PromptService) DeleteEventsByIDs(ctx context.Context, ids []int64) (*DeleteResult, error) {
|
||||
result, err := s.repo.DeleteEventsByIDs(ctx, ids)
|
||||
if err == nil {
|
||||
s.deletePayloads(ctx, result.JobIDs)
|
||||
}
|
||||
return result, err
|
||||
}
|
||||
|
||||
type deleteClaims struct {
|
||||
FilterHash string `json:"filter_hash"`
|
||||
SnapshotMaxID int64 `json:"snapshot_max_id"`
|
||||
AdminID int64 `json:"admin_id"`
|
||||
IssuedAt time.Time `json:"issued_at"`
|
||||
ExpiresAt time.Time `json:"expires_at"`
|
||||
}
|
||||
|
||||
func (s *PromptService) PreviewDelete(ctx context.Context, filter EventFilter, adminID int64) (*DeletePreview, error) {
|
||||
preview, err := s.repo.PreviewDelete(ctx, filter)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
now := s.clock.Now()
|
||||
expires := now.Add(5 * time.Minute)
|
||||
claimsRaw, _ := json.Marshal(deleteClaims{FilterHash: preview.FilterHash, SnapshotMaxID: preview.SnapshotMaxID, AdminID: adminID, IssuedAt: now, ExpiresAt: expires})
|
||||
token, err := s.config.Encrypt(string(claimsRaw))
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
preview.ConfirmationToken, preview.ExpiresAt = token, expires
|
||||
LogInfo(EventDeletePreviewed, map[string]any{"user_id": adminID, "status": "previewed"})
|
||||
return preview, nil
|
||||
}
|
||||
|
||||
type DeleteByFilterRequest struct {
|
||||
Filter EventFilter `json:"filter"`
|
||||
SnapshotMaxID int64 `json:"snapshot_max_id"`
|
||||
FilterHash string `json:"filter_hash"`
|
||||
ConfirmationToken string `json:"confirmation_token"`
|
||||
Confirm bool `json:"confirm"`
|
||||
}
|
||||
|
||||
func (s *PromptService) DeleteByFilter(ctx context.Context, request DeleteByFilterRequest, adminID int64) (*DeleteResult, error) {
|
||||
if !request.Confirm {
|
||||
return nil, errors.New("prompt audit filter delete requires confirm=true")
|
||||
}
|
||||
plain, err := s.config.Decrypt(strings.TrimSpace(request.ConfirmationToken))
|
||||
if err != nil {
|
||||
return nil, errors.New("prompt audit confirmation token invalid")
|
||||
}
|
||||
var claims deleteClaims
|
||||
if json.Unmarshal([]byte(plain), &claims) != nil {
|
||||
return nil, errors.New("prompt audit confirmation token invalid")
|
||||
}
|
||||
computed := FilterHash(request.Filter, request.SnapshotMaxID)
|
||||
if claims.AdminID != adminID || claims.SnapshotMaxID != request.SnapshotMaxID || claims.FilterHash != request.FilterHash || request.FilterHash != computed || !s.clock.Now().Before(claims.ExpiresAt) {
|
||||
return nil, errors.New("prompt audit confirmation token does not match deletion request")
|
||||
}
|
||||
result, err := s.repo.DeleteEventsByFilter(ctx, request.Filter, request.SnapshotMaxID, 200)
|
||||
if err == nil {
|
||||
s.deletePayloads(ctx, result.JobIDs)
|
||||
LogWarn(EventEventsFilterDeleted, map[string]any{"user_id": adminID, "status": "deleted"})
|
||||
}
|
||||
return result, err
|
||||
}
|
||||
|
||||
func (s *PromptService) deletePayloads(ctx context.Context, jobIDs []int64) {
|
||||
for _, id := range jobIDs {
|
||||
_ = s.payload.Delete(ctx, id)
|
||||
}
|
||||
}
|
||||
|
||||
func parseTimeQuery(value string) *time.Time {
|
||||
parsed, err := time.Parse(time.RFC3339, strings.TrimSpace(value))
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
parsed = parsed.UTC()
|
||||
return &parsed
|
||||
}
|
||||
@@ -0,0 +1,126 @@
|
||||
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{})
|
||||
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 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)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,401 @@
|
||||
package securityaudit
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"regexp"
|
||||
"sort"
|
||||
"strings"
|
||||
"unicode/utf8"
|
||||
)
|
||||
|
||||
var (
|
||||
ErrNoPromptText = errors.New("prompt audit request contains no user text")
|
||||
|
||||
bearerPattern = regexp.MustCompile(`(?i)\bBearer\s+[A-Za-z0-9._~+\-/]+=*`)
|
||||
apiKeyPattern = regexp.MustCompile(`(?i)\b(sk|rk|pk|api[_-]?key|token|secret|password)[-_:=\s]+[A-Za-z0-9._~+\-/]{8,}`)
|
||||
canaryPattern = regexp.MustCompile(`(?i)([A-Z]+_CANARY_)[A-Za-z0-9_-]+`)
|
||||
emailPattern = regexp.MustCompile(`(?i)\b[A-Z0-9._%+\-]+@[A-Z0-9.\-]+\.[A-Z]{2,}\b`)
|
||||
phonePattern = regexp.MustCompile(`(?:\+?\d[\d\s().-]{8,}\d)`)
|
||||
)
|
||||
|
||||
func ExtractPromptSnapshot(req Request) (PromptSnapshot, error) {
|
||||
var document any
|
||||
if err := json.Unmarshal(req.Body, &document); err != nil {
|
||||
return PromptSnapshot{}, errors.New("prompt audit request JSON is invalid")
|
||||
}
|
||||
segments := extractProtocolSegments(req.Protocol, document)
|
||||
segments = normalizeSegmentsLatestFirst(segments)
|
||||
if len(segments) == 0 {
|
||||
return PromptSnapshot{}, ErrNoPromptText
|
||||
}
|
||||
scanText := strings.Join(segments, "\n\n")
|
||||
digest := sha256.Sum256([]byte(scanText))
|
||||
stage := strings.TrimSpace(req.Stage)
|
||||
if stage == "" {
|
||||
stage = "http"
|
||||
}
|
||||
return PromptSnapshot{
|
||||
RequestID: req.RequestID, UserID: req.UserID, UsernameSnapshot: req.Username,
|
||||
UserEmailSnapshot: req.UserEmail, APIKeyID: req.APIKeyID, APIKeyNameSnapshot: req.APIKeyName,
|
||||
GroupID: cloneInt64Ptr(req.GroupID), GroupName: req.GroupName, Provider: req.Provider,
|
||||
Endpoint: req.Endpoint, Protocol: req.Protocol, Model: req.Model,
|
||||
PromptHash: hex.EncodeToString(digest[:]), RedactedPreview: BuildPromptPreview(scanText, 480),
|
||||
PromptLength: utf8.RuneCountInString(scanText), MessageCount: len(segments), Stage: stage,
|
||||
ScanText: scanText,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func extractProtocolSegments(protocol string, document any) []string {
|
||||
root, _ := document.(map[string]any)
|
||||
protocol = strings.ToLower(strings.TrimSpace(protocol))
|
||||
switch protocol {
|
||||
case "openai_chat_completions", "openai_chat", "chat_completions":
|
||||
return extractMessages(root["messages"], "user")
|
||||
case "anthropic_messages", "claude_messages", "messages":
|
||||
return extractMessages(root["messages"], "user")
|
||||
case "gemini", "gemini_generate_content":
|
||||
return extractGeminiRoot(root)
|
||||
case "openai_responses", "responses", "responses_websocket":
|
||||
if frameType := stringValue(root["type"]); frameType != "" || protocol == "responses_websocket" {
|
||||
if frameType != "response.create" {
|
||||
return nil
|
||||
}
|
||||
if input, exists := root["input"]; exists && input != nil {
|
||||
return extractResponses(input)
|
||||
}
|
||||
if response, ok := root["response"].(map[string]any); ok {
|
||||
return extractResponses(response["input"])
|
||||
}
|
||||
return nil
|
||||
}
|
||||
return extractResponses(root["input"])
|
||||
case "openai_images", "grok_media", "media", "images":
|
||||
return extractMediaPrompts(root)
|
||||
default:
|
||||
if messages := extractMessages(root["messages"], "user"); len(messages) > 0 {
|
||||
return messages
|
||||
}
|
||||
if responses := extractResponses(root["input"]); len(responses) > 0 {
|
||||
return responses
|
||||
}
|
||||
if gemini := extractGeminiRoot(root); len(gemini) > 0 {
|
||||
return gemini
|
||||
}
|
||||
return extractMediaPrompts(root)
|
||||
}
|
||||
}
|
||||
|
||||
func extractMessages(value any, wantedRole string) []string {
|
||||
items, ok := value.([]any)
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
result := make([]string, 0, len(items))
|
||||
for _, item := range items {
|
||||
message, ok := item.(map[string]any)
|
||||
if !ok || !strings.EqualFold(stringValue(message["role"]), wantedRole) {
|
||||
continue
|
||||
}
|
||||
texts := contentTexts(message["content"])
|
||||
if len(texts) > 0 {
|
||||
result = append(result, strings.Join(texts, "\n"))
|
||||
}
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func extractResponses(value any) []string {
|
||||
switch typed := value.(type) {
|
||||
case string:
|
||||
return []string{typed}
|
||||
case []any:
|
||||
result := make([]string, 0, len(typed))
|
||||
for _, item := range typed {
|
||||
switch entry := item.(type) {
|
||||
case string:
|
||||
result = append(result, entry)
|
||||
case map[string]any:
|
||||
role := strings.ToLower(stringValue(entry["role"]))
|
||||
if role != "" && role != "user" {
|
||||
continue
|
||||
}
|
||||
if content, exists := entry["content"]; exists {
|
||||
if texts := contentTexts(content); len(texts) > 0 {
|
||||
result = append(result, strings.Join(texts, "\n"))
|
||||
}
|
||||
} else if text := stringValue(entry["text"]); text != "" {
|
||||
result = append(result, text)
|
||||
}
|
||||
}
|
||||
}
|
||||
return result
|
||||
case map[string]any:
|
||||
role := strings.ToLower(stringValue(typed["role"]))
|
||||
if role != "" && role != "user" {
|
||||
return nil
|
||||
}
|
||||
return contentTexts(typed["content"])
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func extractGemini(value any) []string {
|
||||
var contents []any
|
||||
switch typed := value.(type) {
|
||||
case []any:
|
||||
contents = typed
|
||||
case map[string]any:
|
||||
contents = []any{typed}
|
||||
default:
|
||||
return nil
|
||||
}
|
||||
result := make([]string, 0, len(contents))
|
||||
for _, item := range contents {
|
||||
content, ok := item.(map[string]any)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
role := strings.ToLower(stringValue(content["role"]))
|
||||
if role != "" && role != "user" {
|
||||
continue
|
||||
}
|
||||
parts, _ := content["parts"].([]any)
|
||||
for _, part := range parts {
|
||||
if object, ok := part.(map[string]any); ok {
|
||||
if text := stringValue(object["text"]); text != "" {
|
||||
result = append(result, text)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func extractGeminiRoot(root map[string]any) []string {
|
||||
if root == nil {
|
||||
return nil
|
||||
}
|
||||
result := extractGemini(root["contents"])
|
||||
result = append(result, extractGemini(root["content"])...)
|
||||
result = append(result, extractGeminiInstances(root["instances"])...)
|
||||
if requests, ok := root["requests"].([]any); ok {
|
||||
for _, item := range requests {
|
||||
request, ok := item.(map[string]any)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
result = append(result, extractGemini(request["contents"])...)
|
||||
result = append(result, extractGemini(request["content"])...)
|
||||
result = append(result, extractGeminiInstances(request["instances"])...)
|
||||
}
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func extractGeminiInstances(value any) []string {
|
||||
instances, ok := value.([]any)
|
||||
if !ok {
|
||||
return nil
|
||||
}
|
||||
result := make([]string, 0, len(instances))
|
||||
for _, item := range instances {
|
||||
if instance, ok := item.(map[string]any); ok {
|
||||
if prompt := stringValue(instance["prompt"]); prompt != "" {
|
||||
result = append(result, prompt)
|
||||
}
|
||||
}
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func extractMediaPrompts(root map[string]any) []string {
|
||||
if root == nil {
|
||||
return nil
|
||||
}
|
||||
result := make([]string, 0, 4)
|
||||
seen := map[string]struct{}{}
|
||||
var walk func(any, string)
|
||||
walk = func(value any, key string) {
|
||||
switch typed := value.(type) {
|
||||
case map[string]any:
|
||||
keys := make([]string, 0, len(typed))
|
||||
for childKey := range typed {
|
||||
keys = append(keys, childKey)
|
||||
}
|
||||
sort.Strings(keys)
|
||||
for _, childKey := range keys {
|
||||
walk(typed[childKey], childKey)
|
||||
}
|
||||
case []any:
|
||||
for _, item := range typed {
|
||||
walk(item, key)
|
||||
}
|
||||
case string:
|
||||
if !isMediaPromptKey(key) || looksLikeMediaPayload(typed) {
|
||||
return
|
||||
}
|
||||
text := strings.TrimSpace(typed)
|
||||
if text == "" {
|
||||
return
|
||||
}
|
||||
if _, duplicate := seen[text]; duplicate {
|
||||
return
|
||||
}
|
||||
seen[text] = struct{}{}
|
||||
result = append(result, text)
|
||||
}
|
||||
}
|
||||
walk(root, "")
|
||||
return result
|
||||
}
|
||||
|
||||
func isMediaPromptKey(key string) bool {
|
||||
normalized := strings.NewReplacer("_", "", "-", "").Replace(strings.ToLower(strings.TrimSpace(key)))
|
||||
switch normalized {
|
||||
case "prompt", "inputprompt", "textprompt", "description", "query", "lyrics", "negativeprompt",
|
||||
"positiveprompt", "gptdescriptionprompt", "prompten", "finalprompt", "finalzhprompt",
|
||||
"origprompt", "actualprompt", "imageprompt", "input":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func looksLikeMediaPayload(value string) bool {
|
||||
trimmed := strings.TrimSpace(value)
|
||||
lower := strings.ToLower(trimmed)
|
||||
if strings.HasPrefix(lower, "data:image/") || strings.HasPrefix(lower, "data:video/") ||
|
||||
strings.HasPrefix(lower, "http://") || strings.HasPrefix(lower, "https://") {
|
||||
return true
|
||||
}
|
||||
if len(trimmed) >= 256 {
|
||||
for _, r := range trimmed {
|
||||
alphaNumeric := (r >= 'A' && r <= 'Z') || (r >= 'a' && r <= 'z') || (r >= '0' && r <= '9')
|
||||
if !alphaNumeric && r != '+' && r != '/' && r != '=' {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func contentTexts(value any) []string {
|
||||
switch typed := value.(type) {
|
||||
case string:
|
||||
return []string{typed}
|
||||
case []any:
|
||||
result := make([]string, 0, len(typed))
|
||||
for _, part := range typed {
|
||||
object, ok := part.(map[string]any)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
typeName := strings.ToLower(stringValue(object["type"]))
|
||||
if typeName != "" && typeName != "text" && typeName != "input_text" {
|
||||
continue
|
||||
}
|
||||
if text := stringValue(object["text"]); text != "" {
|
||||
result = append(result, text)
|
||||
}
|
||||
}
|
||||
return result
|
||||
case map[string]any:
|
||||
if text := stringValue(typed["text"]); text != "" {
|
||||
return []string{text}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func normalizeSegmentsLatestFirst(values []string) []string {
|
||||
normalized := make([]string, 0, len(values))
|
||||
for _, value := range values {
|
||||
value = strings.TrimSpace(value)
|
||||
if value != "" {
|
||||
normalized = append(normalized, value)
|
||||
}
|
||||
}
|
||||
if len(normalized) <= 1 {
|
||||
return normalized
|
||||
}
|
||||
latest := normalized[len(normalized)-1]
|
||||
result := make([]string, 0, len(normalized))
|
||||
result = append(result, latest)
|
||||
result = append(result, normalized[:len(normalized)-1]...)
|
||||
return result
|
||||
}
|
||||
|
||||
func RedactPreview(value string, maxRunes int) string {
|
||||
value = bearerPattern.ReplaceAllString(value, "Bearer ***")
|
||||
value = apiKeyPattern.ReplaceAllStringFunc(value, func(match string) string {
|
||||
if index := strings.IndexAny(match, ":= \t"); index >= 0 {
|
||||
return match[:index+1] + "***"
|
||||
}
|
||||
return "***"
|
||||
})
|
||||
value = canaryPattern.ReplaceAllString(value, "${1}***")
|
||||
value = emailPattern.ReplaceAllString(value, "***@***")
|
||||
value = phonePattern.ReplaceAllString(value, "***PHONE***")
|
||||
return TrimRunes(value, maxRunes)
|
||||
}
|
||||
|
||||
// BuildPromptPreview always withholds part of the sanitized input. Even short,
|
||||
// otherwise-benign prompts must not become a recoverable raw-prompt database
|
||||
// field merely because no secret pattern happened to match.
|
||||
func BuildPromptPreview(value string, maxRunes int) string {
|
||||
redacted := strings.TrimSpace(RedactPreview(value, maxRunes))
|
||||
if redacted == "" {
|
||||
return ""
|
||||
}
|
||||
runes := []rune(redacted)
|
||||
hadTruncation := strings.HasSuffix(redacted, "…")
|
||||
visibleLength := len(runes)
|
||||
if hadTruncation && visibleLength > 0 {
|
||||
visibleLength--
|
||||
}
|
||||
maskCount := visibleLength / 4
|
||||
if maskCount < 1 {
|
||||
maskCount = 1
|
||||
}
|
||||
if maskCount > 16 {
|
||||
maskCount = 16
|
||||
}
|
||||
keep := visibleLength - maskCount
|
||||
if keep < 0 {
|
||||
keep = 0
|
||||
}
|
||||
preview := string(runes[:keep]) + "***"
|
||||
if hadTruncation {
|
||||
preview += "…"
|
||||
}
|
||||
return preview
|
||||
}
|
||||
|
||||
func TrimRunes(value string, limit int) string {
|
||||
if limit <= 0 {
|
||||
return ""
|
||||
}
|
||||
runes := []rune(value)
|
||||
if len(runes) <= limit {
|
||||
return value
|
||||
}
|
||||
return string(runes[:limit]) + "…"
|
||||
}
|
||||
|
||||
func stringValue(value any) string {
|
||||
text, _ := value.(string)
|
||||
return strings.TrimSpace(text)
|
||||
}
|
||||
|
||||
func cloneInt64Ptr(value *int64) *int64 {
|
||||
if value == nil {
|
||||
return nil
|
||||
}
|
||||
cloned := *value
|
||||
return &cloned
|
||||
}
|
||||
@@ -0,0 +1,188 @@
|
||||
package securityaudit
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"strings"
|
||||
"testing"
|
||||
"unicode/utf8"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestExtractPromptSnapshotProtocols(t *testing.T) {
|
||||
tests := []struct {
|
||||
protocol, body, first string
|
||||
count int
|
||||
}{
|
||||
{"openai_chat_completions", `{"messages":[{"role":"user","content":"old"},{"role":"assistant","content":"ignore"},{"role":"user","content":[{"type":"text","text":"最新😀"}]}]}`, "最新😀", 2},
|
||||
{"openai_responses", `{"input":[{"role":"user","content":[{"type":"input_text","text":"response text"}]}]}`, "response text", 1},
|
||||
{"anthropic_messages", `{"messages":[{"role":"user","content":[{"type":"text","text":"claude"}]}]}`, "claude", 1},
|
||||
{"gemini", `{"contents":[{"role":"user","parts":[{"text":"gemini"},{"inline_data":{"data":"BASE64"}}]}]}`, "gemini", 1},
|
||||
{"openai_images", `{"prompt":"draw a cat","image":"BASE64SECRET"}`, "draw a cat", 1},
|
||||
{"responses_websocket", `{"type":"response.create","response":{"input":"turn two"}}`, "turn two", 1},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.protocol, func(t *testing.T) {
|
||||
snapshot, err := ExtractPromptSnapshot(Request{Protocol: tt.protocol, Body: []byte(tt.body), Stage: "http"})
|
||||
require.NoError(t, err)
|
||||
require.True(t, strings.HasPrefix(snapshot.ScanText, tt.first))
|
||||
require.Equal(t, tt.count, snapshot.MessageCount)
|
||||
require.Equal(t, utf8.RuneCountInString(snapshot.ScanText), snapshot.PromptLength)
|
||||
require.NotEmpty(t, snapshot.PromptHash)
|
||||
require.NotContains(t, snapshot.ScanText, "BASE64SECRET")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSnapshotRedactsCanariesAndPreservesHashOfScanText(t *testing.T) {
|
||||
body := `{"messages":[{"role":"user","content":"PROMPT_CANARY_ABC123 email@example.com +86 138 0013 8000 Bearer AUTH_CANARY_XYZ sk-secretvalue123 password=supersecret123"}]}`
|
||||
snapshot, err := ExtractPromptSnapshot(Request{Protocol: "openai_chat_completions", Body: []byte(body)})
|
||||
require.NoError(t, err)
|
||||
require.NotContains(t, snapshot.RedactedPreview, "ABC123")
|
||||
require.NotContains(t, snapshot.RedactedPreview, "email@example.com")
|
||||
require.NotContains(t, snapshot.RedactedPreview, "AUTH_CANARY_XYZ")
|
||||
require.NotContains(t, snapshot.RedactedPreview, "secretvalue123")
|
||||
require.NotContains(t, snapshot.RedactedPreview, "supersecret123")
|
||||
require.NotContains(t, snapshot.RedactedPreview, "138 0013 8000")
|
||||
require.Contains(t, snapshot.ScanText, "PROMPT_CANARY_ABC123")
|
||||
require.NotEqual(t, snapshot.ScanText, snapshot.RedactedPreview)
|
||||
digest := sha256.Sum256([]byte(snapshot.ScanText))
|
||||
require.Equal(t, hex.EncodeToString(digest[:]), snapshot.PromptHash)
|
||||
require.Empty(t, snapshot.Redacted().ScanText)
|
||||
}
|
||||
|
||||
func TestSplitRunesDoesNotSplitUTF8(t *testing.T) {
|
||||
chunks := SplitRunes("中文😀éabc", 2)
|
||||
require.Equal(t, []string{"中文", "😀e", "́a", "bc"}, chunks)
|
||||
for _, chunk := range chunks {
|
||||
require.True(t, utf8.ValidString(chunk))
|
||||
}
|
||||
require.Equal(t, "中文😀éabc", strings.Join(chunks, ""))
|
||||
}
|
||||
|
||||
func TestPromptSnapshotLatestUserMessageIsOnePrioritizedSegment(t *testing.T) {
|
||||
body := []byte(`{
|
||||
"messages":[
|
||||
{"role":"user","content":"历史输入"},
|
||||
{"role":"assistant","content":"assistant output must be ignored"},
|
||||
{"role":"tool","content":"tool output must be ignored"},
|
||||
{"role":"user","content":[
|
||||
{"type":"text","text":"最新第一块😀"},
|
||||
{"type":"image_url","image_url":{"url":"data:image/png;base64,IMAGE_CANARY_BASE64"}},
|
||||
{"type":"text","text":"最新第二块é"}
|
||||
]}
|
||||
]
|
||||
}`)
|
||||
snapshot, err := ExtractPromptSnapshot(Request{Protocol: "openai_chat_completions", Body: body})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, 2, snapshot.MessageCount)
|
||||
require.Equal(t, "最新第一块😀\n最新第二块é\n\n历史输入", snapshot.ScanText)
|
||||
require.NotContains(t, snapshot.ScanText, "assistant output")
|
||||
require.NotContains(t, snapshot.ScanText, "tool output")
|
||||
require.NotContains(t, snapshot.ScanText, "IMAGE_CANARY_BASE64")
|
||||
require.Equal(t, utf8.RuneCountInString(snapshot.ScanText), snapshot.PromptLength)
|
||||
}
|
||||
|
||||
func TestPromptSnapshotResponsesShapes(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
body string
|
||||
want string
|
||||
}{
|
||||
{name: "string", body: `{"input":"plain response input"}`, want: "plain response input"},
|
||||
{name: "message array", body: `{"input":[{"role":"assistant","content":"ignore"},{"role":"user","content":[{"type":"input_text","text":"message block"}]}]}`, want: "message block"},
|
||||
{name: "direct input text", body: `{"input":[{"type":"input_text","text":"direct block"}]}`, want: "direct block"},
|
||||
{name: "single object", body: `{"input":{"role":"user","content":[{"type":"input_text","text":"single object"}]}}`, want: "single object"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
snapshot, err := ExtractPromptSnapshot(Request{Protocol: "openai_responses", Body: []byte(tt.body)})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, tt.want, snapshot.ScanText)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPromptSnapshotGeminiBatchShapesAndMediaExclusion(t *testing.T) {
|
||||
body := []byte(`{
|
||||
"contents":{"role":"user","parts":[{"text":"root content"},{"inlineData":{"data":"ROOT_BASE64"}}]},
|
||||
"instances":[{"prompt":"instance prompt"}],
|
||||
"requests":[
|
||||
{"contents":[{"role":"model","parts":[{"text":"ignore model"}]},{"role":"user","parts":[{"text":"nested user"}]}]},
|
||||
{"instances":[{"prompt":"nested instance"}]}
|
||||
]
|
||||
}`)
|
||||
snapshot, err := ExtractPromptSnapshot(Request{Protocol: "gemini", Body: body})
|
||||
require.NoError(t, err)
|
||||
require.True(t, strings.HasPrefix(snapshot.ScanText, "nested instance"))
|
||||
for _, expected := range []string{"root content", "instance prompt", "nested user", "nested instance"} {
|
||||
require.Contains(t, snapshot.ScanText, expected)
|
||||
}
|
||||
require.NotContains(t, snapshot.ScanText, "ROOT_BASE64")
|
||||
require.NotContains(t, snapshot.ScanText, "ignore model")
|
||||
}
|
||||
|
||||
func TestPromptSnapshotMediaOnlyExtractsDeterministicTextPrompts(t *testing.T) {
|
||||
body := []byte(`{
|
||||
"prompt":"draw a lighthouse",
|
||||
"image":"data:image/png;base64,IMAGE_CANARY",
|
||||
"input":{"negative_prompt":"no fog","image_prompt":"https://example.test/input.png","prompt":"draw a lighthouse"},
|
||||
"request":{"lyrics":"ocean song","input":"` + strings.Repeat("A", 300) + `"},
|
||||
"images":[{"description":"nested textual direction","image_url":"https://example.test/image.png"}]
|
||||
}`)
|
||||
snapshot, err := ExtractPromptSnapshot(Request{Protocol: "grok_media", Body: body})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, 4, snapshot.MessageCount)
|
||||
for _, expected := range []string{"draw a lighthouse", "no fog", "ocean song", "nested textual direction"} {
|
||||
require.Contains(t, snapshot.ScanText, expected)
|
||||
}
|
||||
require.Equal(t, 1, strings.Count(snapshot.ScanText, "draw a lighthouse"))
|
||||
require.NotContains(t, snapshot.ScanText, "IMAGE_CANARY")
|
||||
require.NotContains(t, snapshot.ScanText, "example.test")
|
||||
require.NotContains(t, snapshot.ScanText, strings.Repeat("A", 100))
|
||||
}
|
||||
|
||||
func TestResponsesWebSocketOnlyAuditsResponseCreateAndPreservesStage(t *testing.T) {
|
||||
for _, stage := range []string{"first_turn", "subsequent_turn"} {
|
||||
snapshot, err := ExtractPromptSnapshot(Request{
|
||||
Protocol: "openai_responses", Stage: stage,
|
||||
Body: []byte(`{"type":"response.create","response":{"model":"gpt-test","input":[{"role":"user","content":[{"type":"input_text","text":"ws turn"}]}]}}`),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "ws turn", snapshot.ScanText)
|
||||
require.Equal(t, stage, snapshot.Stage)
|
||||
}
|
||||
_, err := ExtractPromptSnapshot(Request{
|
||||
Protocol: "openai_responses", Stage: "subsequent_turn",
|
||||
Body: []byte(`{"type":"conversation.item.create","response":{"input":"must not scan this frame"}}`),
|
||||
})
|
||||
require.True(t, errors.Is(err, ErrNoPromptText))
|
||||
}
|
||||
|
||||
func TestPromptSnapshotEmptyAndLongUnicodeInput(t *testing.T) {
|
||||
_, err := ExtractPromptSnapshot(Request{Protocol: "openai_chat_completions", Body: []byte(`{"messages":[{"role":"assistant","content":"not user"},{"role":"user","content":" "}]}`)})
|
||||
require.True(t, errors.Is(err, ErrNoPromptText))
|
||||
|
||||
latest := strings.Repeat("最新😀é", 80)
|
||||
history := strings.Repeat("历史中文", 80)
|
||||
body := []byte(`{"messages":[{"role":"user","content":` + string(mustJSON(t, history)) + `},{"role":"user","content":` + string(mustJSON(t, latest)) + `}]}`)
|
||||
snapshot, err := ExtractPromptSnapshot(Request{Protocol: "openai_chat_completions", Body: body})
|
||||
require.NoError(t, err)
|
||||
require.True(t, strings.HasPrefix(snapshot.ScanText, latest))
|
||||
chunks := SplitRunes(snapshot.ScanText, 127)
|
||||
require.Equal(t, snapshot.ScanText, strings.Join(chunks, ""))
|
||||
for _, chunk := range chunks {
|
||||
require.LessOrEqual(t, len([]rune(chunk)), 127)
|
||||
require.True(t, utf8.ValidString(chunk))
|
||||
}
|
||||
}
|
||||
|
||||
func mustJSON(t *testing.T, value string) []byte {
|
||||
t.Helper()
|
||||
raw, err := json.Marshal(value)
|
||||
require.NoError(t, err)
|
||||
return raw
|
||||
}
|
||||
@@ -0,0 +1,276 @@
|
||||
package securityaudit
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
SettingKeyPromptAuditConfig = "prompt_audit_config"
|
||||
SettingKeyRiskControl = "risk_control_enabled"
|
||||
|
||||
ConfigInvalidationChannel = "sub2api:prompt_guard:config:invalidate"
|
||||
PayloadKeyPrefix = "sub2api:prompt_audit:payload:"
|
||||
|
||||
ErrorCodeBlocked = "prompt_guard_blocked"
|
||||
ErrorCodeUnavailable = "prompt_guard_unavailable"
|
||||
ErrorCodeInvalidResponse = "prompt_guard_invalid_response"
|
||||
ErrorCodeConfigConflict = "prompt_audit_config_conflict"
|
||||
ErrorCodeRequiresEnabled = "prompt_guard_requires_audit_enabled"
|
||||
|
||||
DefaultGuardModel = "sileader/qwen3guard:0.6b"
|
||||
)
|
||||
|
||||
type Mode string
|
||||
|
||||
const (
|
||||
ModeOff Mode = "off"
|
||||
ModeAsync Mode = "async_audit"
|
||||
ModeBlocking Mode = "blocking"
|
||||
)
|
||||
|
||||
type DecisionKind string
|
||||
|
||||
const (
|
||||
DecisionAllow DecisionKind = "allow"
|
||||
DecisionFlag DecisionKind = "flag"
|
||||
DecisionBlock DecisionKind = "block"
|
||||
DecisionUnavailable DecisionKind = "unavailable"
|
||||
DecisionInvalid DecisionKind = "invalid"
|
||||
)
|
||||
|
||||
type EventDecision string
|
||||
|
||||
const (
|
||||
EventPass EventDecision = "pass"
|
||||
EventFlag EventDecision = "flag"
|
||||
EventCritical EventDecision = "critical"
|
||||
)
|
||||
|
||||
type RiskLevel string
|
||||
|
||||
const (
|
||||
RiskLow RiskLevel = "low"
|
||||
RiskMedium RiskLevel = "medium"
|
||||
RiskHigh RiskLevel = "high"
|
||||
RiskCritical RiskLevel = "critical"
|
||||
)
|
||||
|
||||
type Action string
|
||||
|
||||
const (
|
||||
ActionAllow Action = "Allow"
|
||||
ActionWarn Action = "Warn"
|
||||
ActionBlock Action = "Block"
|
||||
)
|
||||
|
||||
type Request struct {
|
||||
RequestID string
|
||||
UserID int64
|
||||
Username string
|
||||
UserEmail string
|
||||
APIKeyID int64
|
||||
APIKeyName string
|
||||
GroupID *int64
|
||||
GroupName string
|
||||
Provider string
|
||||
Endpoint string
|
||||
Protocol string
|
||||
Model string
|
||||
Body []byte
|
||||
Stage string
|
||||
}
|
||||
|
||||
func (r Request) Clone() Request {
|
||||
r.Body = append([]byte(nil), r.Body...)
|
||||
if r.GroupID != nil {
|
||||
id := *r.GroupID
|
||||
r.GroupID = &id
|
||||
}
|
||||
return r
|
||||
}
|
||||
|
||||
type PromptSnapshot struct {
|
||||
RequestID string `json:"request_id"`
|
||||
UserID int64 `json:"user_id"`
|
||||
UsernameSnapshot string `json:"username"`
|
||||
UserEmailSnapshot string `json:"user_email"`
|
||||
APIKeyID int64 `json:"api_key_id"`
|
||||
APIKeyNameSnapshot string `json:"api_key_name"`
|
||||
GroupID *int64 `json:"group_id,omitempty"`
|
||||
GroupName string `json:"group_name"`
|
||||
Provider string `json:"provider"`
|
||||
Endpoint string `json:"endpoint"`
|
||||
Protocol string `json:"protocol"`
|
||||
Model string `json:"model"`
|
||||
PromptHash string `json:"prompt_hash"`
|
||||
RedactedPreview string `json:"redacted_preview"`
|
||||
PromptLength int `json:"prompt_length"`
|
||||
MessageCount int `json:"message_count"`
|
||||
Stage string `json:"stage"`
|
||||
|
||||
ScanText string `json:"-"`
|
||||
}
|
||||
|
||||
func (s PromptSnapshot) Redacted() PromptSnapshot {
|
||||
s.ScanText = ""
|
||||
return s
|
||||
}
|
||||
|
||||
type NormalizedResult struct {
|
||||
Decision EventDecision `json:"decision"`
|
||||
RiskLevel RiskLevel `json:"risk_level"`
|
||||
Action Action `json:"action"`
|
||||
Safety string `json:"safety"`
|
||||
Categories []string `json:"categories"`
|
||||
MatchedScanners []string `json:"matched_scanners"`
|
||||
ScannerScores map[string]float64 `json:"scanner_scores"`
|
||||
ScannerEvidence map[string]string `json:"scanner_evidence"`
|
||||
ScannerBackend string `json:"scanner_backend"`
|
||||
ScannerVersion string `json:"scanner_version"`
|
||||
GuardEndpointID string `json:"guard_endpoint_id"`
|
||||
PolicyID string `json:"policy_id"`
|
||||
PolicyVersion int `json:"policy_version"`
|
||||
ChunkTotal int `json:"chunk_total"`
|
||||
LatencyMS int `json:"latency_ms"`
|
||||
UnknownCategories []string `json:"unknown_categories,omitempty"`
|
||||
}
|
||||
|
||||
type PromptDecision struct {
|
||||
Kind DecisionKind `json:"kind"`
|
||||
ErrorCode string `json:"error_code,omitempty"`
|
||||
Result *NormalizedResult `json:"result,omitempty"`
|
||||
AllowNextStage bool `json:"allow_next_stage"`
|
||||
}
|
||||
|
||||
type LegacyDecision struct {
|
||||
Allowed bool `json:"allowed"`
|
||||
Blocked bool `json:"blocked"`
|
||||
Flagged bool `json:"flagged"`
|
||||
Message string `json:"message"`
|
||||
StatusCode int `json:"status_code"`
|
||||
ErrorCode string `json:"error_code"`
|
||||
Action string `json:"action"`
|
||||
}
|
||||
|
||||
type Decision struct {
|
||||
Kind DecisionKind `json:"kind"`
|
||||
HTTPStatus int `json:"http_status"`
|
||||
ErrorCode string `json:"error_code,omitempty"`
|
||||
ClientMessage string `json:"client_message,omitempty"`
|
||||
Legacy *LegacyDecision `json:"legacy,omitempty"`
|
||||
Prompt *PromptDecision `json:"prompt,omitempty"`
|
||||
AllowNextStage bool `json:"allow_next_stage"`
|
||||
}
|
||||
|
||||
type IssueSummary struct {
|
||||
Category string `json:"category"`
|
||||
ScannerID string `json:"scanner_id"`
|
||||
Title string `json:"title"`
|
||||
Description string `json:"description"`
|
||||
Severity string `json:"severity"`
|
||||
SeverityLabel string `json:"severity_label"`
|
||||
Action string `json:"action"`
|
||||
ActionLabel string `json:"action_label"`
|
||||
Code string `json:"code"`
|
||||
Score float64 `json:"score"`
|
||||
Evidence string `json:"evidence"`
|
||||
EvidenceHash string `json:"evidence_hash"`
|
||||
StartRune *int `json:"start_rune,omitempty"`
|
||||
EndRune *int `json:"end_rune,omitempty"`
|
||||
}
|
||||
|
||||
type ProbeResult struct {
|
||||
OK bool `json:"ok"`
|
||||
Status string `json:"status"`
|
||||
ErrorCode string `json:"error_code,omitempty"`
|
||||
Message string `json:"message"`
|
||||
LatencyMS int `json:"latency_ms"`
|
||||
HTTPStatus int `json:"http_status"`
|
||||
Retryable bool `json:"retryable"`
|
||||
CheckedAt time.Time `json:"checked_at"`
|
||||
TokenApplied bool `json:"token_applied"`
|
||||
}
|
||||
|
||||
type GuardMetricsSnapshot struct {
|
||||
Total int64 `json:"total"`
|
||||
Allowed int64 `json:"allowed"`
|
||||
Flagged int64 `json:"flagged"`
|
||||
Blocked int64 `json:"blocked"`
|
||||
Unavailable int64 `json:"unavailable"`
|
||||
Invalid int64 `json:"invalid"`
|
||||
Timeouts int64 `json:"timeouts"`
|
||||
Failovers int64 `json:"failovers"`
|
||||
BulkheadFull int64 `json:"bulkhead_full"`
|
||||
RecordFailed int64 `json:"record_failed"`
|
||||
LatencyCount int64 `json:"latency_count"`
|
||||
LatencyAvgMS int64 `json:"latency_avg_ms"`
|
||||
LatencyP50MS int64 `json:"latency_p50_ms"`
|
||||
LatencyP95MS int64 `json:"latency_p95_ms"`
|
||||
LatencyP99MS int64 `json:"latency_p99_ms"`
|
||||
LatencyMaxMS int64 `json:"latency_max_ms"`
|
||||
}
|
||||
|
||||
type AuditMetricsSnapshot struct {
|
||||
Enqueued int64 `json:"enqueued"`
|
||||
Dropped int64 `json:"dropped"`
|
||||
}
|
||||
|
||||
type QueueStats struct {
|
||||
Staging int64 `json:"staging"`
|
||||
Queued int64 `json:"queued"`
|
||||
Processing int64 `json:"processing"`
|
||||
Retry int64 `json:"retry"`
|
||||
Done int64 `json:"done"`
|
||||
Failed int64 `json:"failed"`
|
||||
Active int64 `json:"active"`
|
||||
}
|
||||
|
||||
type RuntimeSnapshot struct {
|
||||
ProcessStatus string `json:"process_status"`
|
||||
EffectiveMode Mode `json:"effective_mode"`
|
||||
ExpectedConfigVersion int64 `json:"expected_config_version"`
|
||||
ActiveConfigVersion int64 `json:"active_config_version"`
|
||||
ConfigLoadedAt *time.Time `json:"config_loaded_at,omitempty"`
|
||||
ConfigLoadError string `json:"config_load_error,omitempty"`
|
||||
WorkerTotal int `json:"worker_total"`
|
||||
WorkerActive int64 `json:"worker_active"`
|
||||
WorkerHeartbeatAt *time.Time `json:"worker_heartbeat_at,omitempty"`
|
||||
QueueCapacity int `json:"queue_capacity"`
|
||||
Queue QueueStats `json:"queue"`
|
||||
ProcessedTotal int64 `json:"processed_total"`
|
||||
FailedTotal int64 `json:"failed_total"`
|
||||
EnqueuedTotal int64 `json:"enqueued_total"`
|
||||
DroppedTotal int64 `json:"dropped_total"`
|
||||
LastProcessedAt *time.Time `json:"last_processed_at,omitempty"`
|
||||
LastErrorCode string `json:"last_error_code,omitempty"`
|
||||
LastErrorMessage string `json:"last_error_message,omitempty"`
|
||||
DatabaseStatus string `json:"database_status"`
|
||||
RedisStatus string `json:"redis_status"`
|
||||
Endpoints map[string]ProbeResult `json:"endpoints"`
|
||||
GuardMetrics GuardMetricsSnapshot `json:"guard_metrics"`
|
||||
}
|
||||
|
||||
type Clock interface {
|
||||
Now() time.Time
|
||||
}
|
||||
|
||||
type realClock struct{}
|
||||
|
||||
func (realClock) Now() time.Time { return time.Now().UTC() }
|
||||
|
||||
type Metrics interface {
|
||||
Snapshot() GuardMetricsSnapshot
|
||||
AuditSnapshot() AuditMetricsSnapshot
|
||||
Observe(kind DecisionKind, latency time.Duration)
|
||||
IncEnqueued()
|
||||
IncDropped()
|
||||
IncTimeout()
|
||||
IncFailover()
|
||||
IncBulkheadFull()
|
||||
IncRecordFailed()
|
||||
}
|
||||
|
||||
type PromptScanner interface {
|
||||
Scan(ctx context.Context, endpoint ActiveEndpoint, chunk string, enabledScanners []string) (*NormalizedResult, error)
|
||||
}
|
||||
@@ -0,0 +1,347 @@
|
||||
package securityaudit
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"time"
|
||||
)
|
||||
|
||||
type WorkerRuntime struct {
|
||||
active atomic.Int64
|
||||
processed atomic.Int64
|
||||
failed atomic.Int64
|
||||
heartbeatNS atomic.Int64
|
||||
lastProcessedNS atomic.Int64
|
||||
lastErrorMu sync.RWMutex
|
||||
lastErrorCode string
|
||||
lastErrorMessage string
|
||||
}
|
||||
|
||||
type Runner struct {
|
||||
config ConfigStore
|
||||
repo JobRepository
|
||||
payload PayloadStore
|
||||
scanner PromptScanner
|
||||
metrics Metrics
|
||||
clock Clock
|
||||
runtime WorkerRuntime
|
||||
|
||||
mu sync.Mutex
|
||||
cancel context.CancelFunc
|
||||
wg sync.WaitGroup
|
||||
}
|
||||
|
||||
func NewRunner(config ConfigStore, repo JobRepository, payload PayloadStore, scanner PromptScanner, metrics Metrics) *Runner {
|
||||
return &Runner{config: config, repo: repo, payload: payload, scanner: scanner, metrics: metrics, clock: realClock{}}
|
||||
}
|
||||
|
||||
func (r *Runner) Start(ctx context.Context) error {
|
||||
if r == nil || r.config == nil || r.repo == nil || r.payload == nil || r.scanner == nil {
|
||||
return errors.New("prompt audit worker dependencies unavailable")
|
||||
}
|
||||
r.mu.Lock()
|
||||
if r.cancel != nil {
|
||||
r.mu.Unlock()
|
||||
return nil
|
||||
}
|
||||
runCtx, cancel := context.WithCancel(ctx)
|
||||
r.cancel = cancel
|
||||
r.mu.Unlock()
|
||||
if err := r.payload.Ping(runCtx); err != nil {
|
||||
r.setLastError("payload_store_unavailable", err.Error())
|
||||
}
|
||||
for workerID := 0; workerID < MaxWorkerCount; workerID++ {
|
||||
r.wg.Add(1)
|
||||
go r.worker(runCtx, workerID)
|
||||
}
|
||||
r.wg.Add(1)
|
||||
go r.reclaimer(runCtx)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *Runner) Shutdown(ctx context.Context) error {
|
||||
if r == nil {
|
||||
return nil
|
||||
}
|
||||
r.mu.Lock()
|
||||
cancel := r.cancel
|
||||
r.cancel = nil
|
||||
r.mu.Unlock()
|
||||
if cancel != nil {
|
||||
cancel()
|
||||
}
|
||||
done := make(chan struct{})
|
||||
go func() { r.wg.Wait(); close(done) }()
|
||||
select {
|
||||
case <-done:
|
||||
return nil
|
||||
case <-ctx.Done():
|
||||
LogWarn(EventProcessFailed, map[string]any{"status": "shutdown_timeout", "error_code": "worker_shutdown_timeout"})
|
||||
return ctx.Err()
|
||||
}
|
||||
}
|
||||
|
||||
func (r *Runner) worker(ctx context.Context, workerID int) {
|
||||
defer r.wg.Done()
|
||||
ticker := time.NewTicker(500 * time.Millisecond)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
r.runtime.heartbeatNS.Store(r.clock.Now().UnixNano())
|
||||
cfg, ok := r.config.Active()
|
||||
if !ok || !cfg.RiskControlEnabled || !cfg.Enabled || workerID >= cfg.WorkerCount {
|
||||
continue
|
||||
}
|
||||
for {
|
||||
job, claimed, err := r.repo.ClaimNextJob(ctx, r.clock.Now())
|
||||
if err != nil {
|
||||
r.setLastError("claim_job_failed", err.Error())
|
||||
break
|
||||
}
|
||||
if !claimed {
|
||||
break
|
||||
}
|
||||
r.runtime.active.Add(1)
|
||||
r.processSafely(ctx, workerID, cfg, job)
|
||||
r.runtime.active.Add(-1)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (r *Runner) processSafely(ctx context.Context, workerID int, cfg ActiveConfig, job *Job) {
|
||||
defer func() {
|
||||
if recovered := recover(); recovered != nil {
|
||||
r.runtime.failed.Add(1)
|
||||
// Panic values may contain scanner response fragments or prompt data.
|
||||
// Keep only a stable generic message in runtime state and logs.
|
||||
r.setLastError("worker_panic", "worker panic recovered")
|
||||
_ = r.repo.Fail(ctx, job.ID, job.ClaimVersion, "worker_panic", "worker panic recovered")
|
||||
LogError(EventProcessFailed, mergeLogFields(jobLogFields(job), map[string]any{"worker_id": workerID, "status": "failed", "error_code": "worker_panic"}))
|
||||
}
|
||||
}()
|
||||
if err := r.processJob(ctx, workerID, cfg, job); err != nil {
|
||||
r.runtime.failed.Add(1)
|
||||
} else {
|
||||
r.runtime.processed.Add(1)
|
||||
r.runtime.lastProcessedNS.Store(r.clock.Now().UnixNano())
|
||||
}
|
||||
}
|
||||
|
||||
func (r *Runner) processJob(ctx context.Context, workerID int, cfg ActiveConfig, job *Job) error {
|
||||
baseFields := jobLogFields(job)
|
||||
LogInfo(EventAuditStarted, mergeLogFields(baseFields, map[string]any{"worker_id": workerID, "attempts": job.Attempts, "status": "processing"}))
|
||||
scanText, err := r.payload.Get(ctx, job.ID)
|
||||
if err != nil {
|
||||
return r.finishFailure(ctx, job, &GuardError{Code: "payload_missing", Retryable: false, Cause: err})
|
||||
}
|
||||
endpoints := cfg.EnabledEndpoints()
|
||||
if len(endpoints) == 0 {
|
||||
return r.finishFailure(ctx, job, &GuardError{Code: "no_enabled_endpoint", Retryable: true})
|
||||
}
|
||||
chunks := SplitRunes(scanText, minimumInputLimit(endpoints))
|
||||
results := make([]*NormalizedResult, 0, len(chunks))
|
||||
started := r.clock.Now()
|
||||
for index, chunk := range chunks {
|
||||
if err := r.repo.RefreshLease(ctx, job.ID, job.ClaimVersion, r.clock.Now()); err != nil {
|
||||
return err
|
||||
}
|
||||
chunkStarted := r.clock.Now()
|
||||
LogInfo(EventChunkStarted, mergeLogFields(baseFields, map[string]any{"worker_id": workerID, "chunk_index": index + 1, "chunk_total": len(chunks), "chunk_chars": len([]rune(chunk)), "input_chars": job.Snapshot.PromptLength, "input_limit": minimumInputLimit(endpoints), "status": "started"}))
|
||||
result, scanErr := scanWithFailover(ctx, r.scanner, cfg.Scanners, endpoints, chunk, r.metrics)
|
||||
if scanErr != nil {
|
||||
LogWarn(EventChunkFailed, mergeLogFields(baseFields, map[string]any{
|
||||
"worker_id": workerID, "chunk_index": index + 1, "chunk_total": len(chunks),
|
||||
"chunk_chars": len([]rune(chunk)), "input_chars": job.Snapshot.PromptLength,
|
||||
"input_limit": minimumInputLimit(endpoints), "latency_ms": r.clock.Now().Sub(chunkStarted).Milliseconds(),
|
||||
"error_code": guardErrorCode(scanErr), "status": "failed",
|
||||
}))
|
||||
r.observeAsyncFailure(scanErr, r.clock.Now().Sub(started))
|
||||
return r.finishFailure(ctx, job, scanErr)
|
||||
}
|
||||
results = append(results, result)
|
||||
LogInfo(EventChunkCompleted, mergeLogFields(baseFields, map[string]any{"worker_id": workerID, "chunk_index": index + 1, "chunk_total": len(chunks), "guard_endpoint_id": result.GuardEndpointID, "action": result.Action, "latency_ms": r.clock.Now().Sub(chunkStarted).Milliseconds(), "status": "completed"}))
|
||||
if result.Action == ActionBlock {
|
||||
break
|
||||
}
|
||||
}
|
||||
aggregated, err := AggregateResults(results, r.clock.Now().Sub(started))
|
||||
if err != nil {
|
||||
if r.metrics != nil {
|
||||
r.metrics.Observe(DecisionInvalid, r.clock.Now().Sub(started))
|
||||
}
|
||||
return r.finishFailure(ctx, job, &GuardError{Code: ErrorCodeInvalidResponse, Cause: err})
|
||||
}
|
||||
aggregated.ChunkTotal = len(chunks)
|
||||
if r.metrics != nil {
|
||||
r.metrics.Observe(decisionKindForResult(aggregated), r.clock.Now().Sub(started))
|
||||
}
|
||||
LogInfo(EventChunksAggregated, mergeLogFields(baseFields, map[string]any{
|
||||
"worker_id": workerID, "decision": aggregated.Decision, "risk_level": aggregated.RiskLevel,
|
||||
"action": aggregated.Action, "chunk_total": aggregated.ChunkTotal,
|
||||
"latency_ms": aggregated.LatencyMS, "guard_endpoint_id": aggregated.GuardEndpointID, "status": "completed",
|
||||
}))
|
||||
event, err := r.repo.Complete(ctx, job, aggregated, cfg.StorePassEvents)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if deleteErr := r.payload.Delete(ctx, job.ID); deleteErr != nil {
|
||||
LogWarn(EventProcessFailed, mergeLogFields(baseFields, map[string]any{"worker_id": workerID, "status": "payload_delete_deferred", "error_code": "payload_delete_failed"}))
|
||||
}
|
||||
LogInfo(EventProcessed, mergeLogFields(baseFields, map[string]any{"worker_id": workerID, "event_id": eventID(event), "decision": aggregated.Decision, "risk_level": aggregated.RiskLevel, "action": aggregated.Action, "guard_endpoint_id": aggregated.GuardEndpointID, "latency_ms": aggregated.LatencyMS, "status": "done"}))
|
||||
if event != nil && aggregated.Decision != EventPass {
|
||||
LogWarn(EventFindingRecorded, mergeLogFields(baseFields, map[string]any{"worker_id": workerID, "event_id": event.ID, "decision": aggregated.Decision, "risk_level": aggregated.RiskLevel, "action": aggregated.Action, "guard_endpoint_id": aggregated.GuardEndpointID, "status": "recorded"}))
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *Runner) observeAsyncFailure(err error, latency time.Duration) {
|
||||
if r == nil || r.metrics == nil {
|
||||
return
|
||||
}
|
||||
kind := DecisionUnavailable
|
||||
if guardErrorCode(err) == ErrorCodeInvalidResponse {
|
||||
kind = DecisionInvalid
|
||||
}
|
||||
r.metrics.Observe(kind, latency)
|
||||
var guardErr *GuardError
|
||||
if errors.As(err, &guardErr) && guardErr.Timeout {
|
||||
r.metrics.IncTimeout()
|
||||
}
|
||||
}
|
||||
|
||||
func decisionKindForResult(result *NormalizedResult) DecisionKind {
|
||||
if result == nil {
|
||||
return DecisionInvalid
|
||||
}
|
||||
switch result.Action {
|
||||
case ActionBlock:
|
||||
return DecisionBlock
|
||||
case ActionWarn:
|
||||
return DecisionFlag
|
||||
default:
|
||||
return DecisionAllow
|
||||
}
|
||||
}
|
||||
|
||||
func (r *Runner) finishFailure(ctx context.Context, job *Job, err error) error {
|
||||
baseFields := jobLogFields(job)
|
||||
code := guardErrorCode(err)
|
||||
retryable := false
|
||||
var guardErr *GuardError
|
||||
if errors.As(err, &guardErr) {
|
||||
retryable = guardErr.Retryable
|
||||
}
|
||||
if retryable && job.Attempts < job.MaxAttempts {
|
||||
next := r.clock.Now().Add(retryBackoff(job.Attempts))
|
||||
if updateErr := r.repo.Retry(ctx, job.ID, job.ClaimVersion, next, code, "prompt guard temporarily unavailable"); updateErr != nil {
|
||||
return updateErr
|
||||
}
|
||||
LogWarn(EventProcessFailed, mergeLogFields(baseFields, map[string]any{"attempts": job.Attempts, "max_attempts": job.MaxAttempts, "status": "retry", "error_code": code, "retryable": true}))
|
||||
} else {
|
||||
if updateErr := r.repo.Fail(ctx, job.ID, job.ClaimVersion, code, "prompt guard processing failed"); updateErr != nil {
|
||||
return updateErr
|
||||
}
|
||||
_ = r.payload.Delete(ctx, job.ID)
|
||||
LogError(EventProcessFailed, mergeLogFields(baseFields, map[string]any{"attempts": job.Attempts, "max_attempts": job.MaxAttempts, "status": "failed", "error_code": code, "retryable": false}))
|
||||
}
|
||||
r.setLastError(code, err.Error())
|
||||
return err
|
||||
}
|
||||
|
||||
func (r *Runner) reclaimer(ctx context.Context) {
|
||||
defer r.wg.Done()
|
||||
ticker := time.NewTicker(time.Minute)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
now := r.clock.Now()
|
||||
count, err := r.repo.ReclaimStale(ctx, now.Add(-2*time.Minute), now.Add(-90*time.Second), 100)
|
||||
if err != nil {
|
||||
r.setLastError("reclaim_failed", err.Error())
|
||||
continue
|
||||
}
|
||||
if count > 0 {
|
||||
LogWarn(EventProcessingReclaimed, map[string]any{"reclaimed_total": count, "status": "reclaimed"})
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (r *Runner) Snapshot() (active, processed, failed int64, heartbeat, lastProcessed *time.Time, code, message string) {
|
||||
if r == nil {
|
||||
return
|
||||
}
|
||||
active, processed, failed = r.runtime.active.Load(), r.runtime.processed.Load(), r.runtime.failed.Load()
|
||||
if ns := r.runtime.heartbeatNS.Load(); ns > 0 {
|
||||
value := time.Unix(0, ns).UTC()
|
||||
heartbeat = &value
|
||||
}
|
||||
if ns := r.runtime.lastProcessedNS.Load(); ns > 0 {
|
||||
value := time.Unix(0, ns).UTC()
|
||||
lastProcessed = &value
|
||||
}
|
||||
r.runtime.lastErrorMu.RLock()
|
||||
code, message = r.runtime.lastErrorCode, r.runtime.lastErrorMessage
|
||||
r.runtime.lastErrorMu.RUnlock()
|
||||
return
|
||||
}
|
||||
|
||||
func (r *Runner) setLastError(code, _ string) {
|
||||
code, message := sanitizeStoredError(code)
|
||||
r.runtime.lastErrorMu.Lock()
|
||||
r.runtime.lastErrorCode = code
|
||||
r.runtime.lastErrorMessage = message
|
||||
r.runtime.lastErrorMu.Unlock()
|
||||
}
|
||||
|
||||
func scanWithFailover(ctx context.Context, scanner PromptScanner, scanners []string, endpoints []ActiveEndpoint, chunk string, metrics Metrics) (*NormalizedResult, error) {
|
||||
var lastErr error
|
||||
for index, endpoint := range endpoints {
|
||||
result, err := scanner.Scan(ctx, endpoint, chunk, scanners)
|
||||
if err == nil && result != nil {
|
||||
return result, nil
|
||||
}
|
||||
if err == nil {
|
||||
err = &GuardError{Code: ErrorCodeInvalidResponse, Retryable: false}
|
||||
}
|
||||
lastErr = err
|
||||
var guardErr *GuardError
|
||||
if !errors.As(err, &guardErr) || !guardErr.Retryable {
|
||||
return nil, err
|
||||
}
|
||||
if index < len(endpoints)-1 && metrics != nil {
|
||||
metrics.IncFailover()
|
||||
}
|
||||
}
|
||||
if lastErr == nil {
|
||||
lastErr = &GuardError{Code: ErrorCodeUnavailable}
|
||||
}
|
||||
return nil, lastErr
|
||||
}
|
||||
|
||||
func retryBackoff(attempt int) time.Duration {
|
||||
switch attempt {
|
||||
case 1:
|
||||
return 5 * time.Second
|
||||
case 2:
|
||||
return 30 * time.Second
|
||||
default:
|
||||
return 2 * time.Minute
|
||||
}
|
||||
}
|
||||
|
||||
func eventID(event *Event) int64 {
|
||||
if event == nil {
|
||||
return 0
|
||||
}
|
||||
return event.ID
|
||||
}
|
||||
@@ -0,0 +1,591 @@
|
||||
package securityaudit
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"reflect"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
type fixedClock struct{ now time.Time }
|
||||
|
||||
func (c fixedClock) Now() time.Time { return c.now }
|
||||
|
||||
type advancingClock struct {
|
||||
mu sync.Mutex
|
||||
now time.Time
|
||||
step time.Duration
|
||||
}
|
||||
|
||||
func (c *advancingClock) Now() time.Time {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
c.now = c.now.Add(c.step)
|
||||
return c.now
|
||||
}
|
||||
|
||||
type fakeConfigStore struct {
|
||||
cfg ActiveConfig
|
||||
active bool
|
||||
}
|
||||
|
||||
func (s *fakeConfigStore) Start(context.Context) error { return nil }
|
||||
func (s *fakeConfigStore) Shutdown(context.Context) error { return nil }
|
||||
func (s *fakeConfigStore) Active() (ActiveConfig, bool) { return cloneActiveConfig(s.cfg), s.active }
|
||||
func (s *fakeConfigStore) EffectiveMode() Mode {
|
||||
if !s.active {
|
||||
return ModeOff
|
||||
}
|
||||
return s.cfg.EffectiveMode()
|
||||
}
|
||||
func (s *fakeConfigStore) Public() PublicConfig { return PublicConfig{} }
|
||||
func (s *fakeConfigStore) Save(context.Context, UpdateConfigRequest, int64) (PublicConfig, error) {
|
||||
return PublicConfig{}, nil
|
||||
}
|
||||
func (s *fakeConfigStore) RuntimeState() (int64, int64, *time.Time, string) {
|
||||
return s.cfg.ConfigVersion, s.cfg.ConfigVersion, nil, ""
|
||||
}
|
||||
func (s *fakeConfigStore) Encrypt(value string) (string, error) { return value, nil }
|
||||
func (s *fakeConfigStore) Decrypt(value string) (string, error) { return value, nil }
|
||||
|
||||
type fakeJobRepository struct {
|
||||
mu sync.Mutex
|
||||
|
||||
trace *[]string
|
||||
createJob *Job
|
||||
createErr error
|
||||
publishErr error
|
||||
refreshErr error
|
||||
completeErr error
|
||||
retryErr error
|
||||
failErr error
|
||||
|
||||
createdSnapshot PromptSnapshot
|
||||
markedCode string
|
||||
completedResult *NormalizedResult
|
||||
completedStore bool
|
||||
completeCount int
|
||||
eventCount int
|
||||
retryAt time.Time
|
||||
retryCode string
|
||||
retried int
|
||||
failedCode string
|
||||
failed int
|
||||
refreshes int
|
||||
|
||||
claimQueue []*Job
|
||||
|
||||
recordBlockingCalls int
|
||||
recordBlockingSnapshot PromptSnapshot
|
||||
recordBlockingResult *NormalizedResult
|
||||
recordBlockingErr error
|
||||
}
|
||||
|
||||
func (r *fakeJobRepository) record(value string) {
|
||||
if r.trace != nil {
|
||||
*r.trace = append(*r.trace, value)
|
||||
}
|
||||
}
|
||||
|
||||
func (r *fakeJobRepository) CreateStagingWithCapacity(_ context.Context, snapshot PromptSnapshot, _ int64, _, _ int) (*Job, error) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.record("create_staging")
|
||||
r.createdSnapshot = snapshot
|
||||
if r.createErr != nil {
|
||||
return nil, r.createErr
|
||||
}
|
||||
if r.createJob == nil {
|
||||
r.createJob = &Job{ID: 1, Snapshot: snapshot}
|
||||
}
|
||||
return r.createJob, nil
|
||||
}
|
||||
func (r *fakeJobRepository) PublishQueued(context.Context, int64) error {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.record("publish_queued")
|
||||
return r.publishErr
|
||||
}
|
||||
func (r *fakeJobRepository) MarkStagingFailed(_ context.Context, _ int64, code, _ string) error {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.record("mark_staging_failed")
|
||||
r.markedCode = code
|
||||
return nil
|
||||
}
|
||||
func (r *fakeJobRepository) ClaimNextJob(context.Context, time.Time) (*Job, bool, error) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
if len(r.claimQueue) == 0 {
|
||||
return nil, false, nil
|
||||
}
|
||||
job := r.claimQueue[0]
|
||||
r.claimQueue = r.claimQueue[1:]
|
||||
return job, true, nil
|
||||
}
|
||||
func (r *fakeJobRepository) RefreshLease(context.Context, int64, int64, time.Time) error {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.refreshes++
|
||||
return r.refreshErr
|
||||
}
|
||||
func (r *fakeJobRepository) Complete(_ context.Context, _ *Job, result *NormalizedResult, storePass bool) (*Event, error) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.completeCount++
|
||||
r.completedResult, r.completedStore = result, storePass
|
||||
if r.completeErr != nil {
|
||||
return nil, r.completeErr
|
||||
}
|
||||
if result.Decision == EventPass && !storePass {
|
||||
return nil, nil
|
||||
}
|
||||
r.eventCount++
|
||||
return &Event{ID: 99, Decision: result.Decision}, nil
|
||||
}
|
||||
func (r *fakeJobRepository) Retry(_ context.Context, _, _ int64, next time.Time, code, _ string) error {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.retried++
|
||||
r.retryAt, r.retryCode = next, code
|
||||
return r.retryErr
|
||||
}
|
||||
func (r *fakeJobRepository) Fail(_ context.Context, _, _ int64, code, _ string) error {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.failed++
|
||||
r.failedCode = code
|
||||
return r.failErr
|
||||
}
|
||||
func (r *fakeJobRepository) ReclaimStale(context.Context, time.Time, time.Time, int) (int64, error) {
|
||||
return 0, nil
|
||||
}
|
||||
func (r *fakeJobRepository) QueueStats(context.Context) (QueueStats, error) { return QueueStats{}, nil }
|
||||
func (r *fakeJobRepository) RecordBlocking(_ context.Context, snapshot PromptSnapshot, _ int64, result *NormalizedResult, _ bool) (*Event, error) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.recordBlockingCalls++
|
||||
r.recordBlockingSnapshot, r.recordBlockingResult = snapshot, result
|
||||
return nil, r.recordBlockingErr
|
||||
}
|
||||
|
||||
type fakePayloadStore struct {
|
||||
mu sync.Mutex
|
||||
|
||||
trace *[]string
|
||||
values map[int64]string
|
||||
setErr error
|
||||
getErr error
|
||||
deleteErr error
|
||||
pingErr error
|
||||
setTTL time.Duration
|
||||
deleted []int64
|
||||
}
|
||||
|
||||
func (s *fakePayloadStore) Set(_ context.Context, jobID int64, value string, ttl time.Duration) error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
if s.trace != nil {
|
||||
*s.trace = append(*s.trace, "payload_set")
|
||||
}
|
||||
if s.setErr != nil {
|
||||
return s.setErr
|
||||
}
|
||||
if s.values == nil {
|
||||
s.values = map[int64]string{}
|
||||
}
|
||||
s.values[jobID], s.setTTL = value, ttl
|
||||
return nil
|
||||
}
|
||||
func (s *fakePayloadStore) Get(_ context.Context, jobID int64) (string, error) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
if s.getErr != nil {
|
||||
return "", s.getErr
|
||||
}
|
||||
value, ok := s.values[jobID]
|
||||
if !ok {
|
||||
return "", errors.New("missing")
|
||||
}
|
||||
return value, nil
|
||||
}
|
||||
func (s *fakePayloadStore) Delete(_ context.Context, jobID int64) error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
if s.trace != nil {
|
||||
*s.trace = append(*s.trace, "payload_delete")
|
||||
}
|
||||
s.deleted = append(s.deleted, jobID)
|
||||
delete(s.values, jobID)
|
||||
return s.deleteErr
|
||||
}
|
||||
func (s *fakePayloadStore) Ping(context.Context) error { return s.pingErr }
|
||||
|
||||
func asyncConfig() ActiveConfig {
|
||||
return ActiveConfig{
|
||||
RiskControlEnabled: true, Enabled: true, BlockingEnabled: false, Strategy: "priority",
|
||||
WorkerCount: 1, QueueCapacity: 8, Scanners: []string{"pii"}, AllGroups: true, ConfigVersion: 7,
|
||||
Endpoints: []ActiveEndpoint{{ID: "guard", Enabled: true, TimeoutMS: 1000, InputLimit: 3}},
|
||||
}
|
||||
}
|
||||
|
||||
func asyncRequest() Request {
|
||||
return Request{RequestID: "request-async", Protocol: "openai_chat_completions", Body: []byte(`{"messages":[{"role":"user","content":"payload canary text"}]}`)}
|
||||
}
|
||||
|
||||
func TestEnqueuerStagingPayloadPublishProtocolAndFailureCleanup(t *testing.T) {
|
||||
t.Run("success", func(t *testing.T) {
|
||||
trace := []string{}
|
||||
repo := &fakeJobRepository{trace: &trace, createJob: &Job{ID: 41}}
|
||||
payload := &fakePayloadStore{trace: &trace, values: map[int64]string{}}
|
||||
enqueuer := NewEnqueuer(&fakeConfigStore{cfg: asyncConfig(), active: true}, repo, payload)
|
||||
require.NoError(t, enqueuer.Enqueue(context.Background(), asyncRequest()))
|
||||
require.Equal(t, []string{"create_staging", "payload_set", "publish_queued"}, trace)
|
||||
require.Empty(t, repo.createdSnapshot.ScanText)
|
||||
require.Equal(t, "payload canary text", payload.values[41])
|
||||
require.Equal(t, DefaultPayloadTTL, payload.setTTL)
|
||||
})
|
||||
|
||||
t.Run("queue admission failures never touch payload", func(t *testing.T) {
|
||||
for _, createErr := range []error{ErrQueueFull, ErrQueueAdmissionBusy, errors.New("database down")} {
|
||||
trace := []string{}
|
||||
repo := &fakeJobRepository{trace: &trace, createErr: createErr}
|
||||
payload := &fakePayloadStore{trace: &trace, values: map[int64]string{}}
|
||||
err := NewEnqueuer(&fakeConfigStore{cfg: asyncConfig(), active: true}, repo, payload).Enqueue(context.Background(), asyncRequest())
|
||||
require.ErrorIs(t, err, createErr)
|
||||
require.Equal(t, []string{"create_staging"}, trace)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("payload failure marks staging failed", func(t *testing.T) {
|
||||
trace := []string{}
|
||||
repo := &fakeJobRepository{trace: &trace, createJob: &Job{ID: 42}}
|
||||
payload := &fakePayloadStore{trace: &trace, values: map[int64]string{}, setErr: errors.New("redis down")}
|
||||
err := NewEnqueuer(&fakeConfigStore{cfg: asyncConfig(), active: true}, repo, payload).Enqueue(context.Background(), asyncRequest())
|
||||
require.Error(t, err)
|
||||
require.Equal(t, []string{"create_staging", "payload_set", "mark_staging_failed"}, trace)
|
||||
require.Equal(t, "payload_store_failed", repo.markedCode)
|
||||
})
|
||||
|
||||
t.Run("publish failure removes payload and marks staging failed", func(t *testing.T) {
|
||||
trace := []string{}
|
||||
repo := &fakeJobRepository{trace: &trace, createJob: &Job{ID: 43}, publishErr: errors.New("publish down")}
|
||||
payload := &fakePayloadStore{trace: &trace, values: map[int64]string{}}
|
||||
err := NewEnqueuer(&fakeConfigStore{cfg: asyncConfig(), active: true}, repo, payload).Enqueue(context.Background(), asyncRequest())
|
||||
require.Error(t, err)
|
||||
require.Equal(t, []string{"create_staging", "payload_set", "publish_queued", "payload_delete", "mark_staging_failed"}, trace)
|
||||
require.Equal(t, "queue_publish_failed", repo.markedCode)
|
||||
require.NotContains(t, payload.values, int64(43))
|
||||
})
|
||||
}
|
||||
|
||||
func TestEnqueuerSkipsOffOutOfScopeAndNoText(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
cfg ActiveConfig
|
||||
req Request
|
||||
}{
|
||||
{name: "off", cfg: ActiveConfig{}, req: asyncRequest()},
|
||||
{name: "out of scope", cfg: func() ActiveConfig {
|
||||
cfg := asyncConfig()
|
||||
cfg.AllGroups = false
|
||||
cfg.GroupIDs = []int64{9}
|
||||
return cfg
|
||||
}(), req: asyncRequest()},
|
||||
{name: "no user text", cfg: asyncConfig(), req: Request{Protocol: "openai_chat_completions", Body: []byte(`{"messages":[{"role":"assistant","content":"ignore"}]}`)}},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
repo := &fakeJobRepository{}
|
||||
err := NewEnqueuer(&fakeConfigStore{cfg: tt.cfg, active: true}, repo, &fakePayloadStore{}).Enqueue(context.Background(), tt.req)
|
||||
require.NoError(t, err)
|
||||
require.Zero(t, repo.createdSnapshot.MessageCount)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnqueuerRecordsAcceptedDroppedAndSkippedMetrics(t *testing.T) {
|
||||
t.Run("accepted increments enqueued", func(t *testing.T) {
|
||||
metrics := NewAtomicMetrics()
|
||||
repo := &fakeJobRepository{createJob: &Job{ID: 44}}
|
||||
payload := &fakePayloadStore{values: map[int64]string{}}
|
||||
|
||||
require.NoError(t, NewEnqueuer(
|
||||
&fakeConfigStore{cfg: asyncConfig(), active: true},
|
||||
repo,
|
||||
payload,
|
||||
metrics,
|
||||
).Enqueue(context.Background(), asyncRequest()))
|
||||
|
||||
require.Equal(t, AuditMetricsSnapshot{Enqueued: 1}, metrics.AuditSnapshot())
|
||||
})
|
||||
|
||||
t.Run("queue full increments dropped", func(t *testing.T) {
|
||||
metrics := NewAtomicMetrics()
|
||||
repo := &fakeJobRepository{createErr: ErrQueueFull}
|
||||
|
||||
err := NewEnqueuer(
|
||||
&fakeConfigStore{cfg: asyncConfig(), active: true},
|
||||
repo,
|
||||
&fakePayloadStore{},
|
||||
metrics,
|
||||
).Enqueue(context.Background(), asyncRequest())
|
||||
|
||||
require.ErrorIs(t, err, ErrQueueFull)
|
||||
require.Equal(t, AuditMetricsSnapshot{Dropped: 1}, metrics.AuditSnapshot())
|
||||
})
|
||||
|
||||
t.Run("skipped request does not increment dropped", func(t *testing.T) {
|
||||
metrics := NewAtomicMetrics()
|
||||
|
||||
require.NoError(t, NewEnqueuer(
|
||||
&fakeConfigStore{cfg: ActiveConfig{}, active: true},
|
||||
&fakeJobRepository{},
|
||||
&fakePayloadStore{},
|
||||
metrics,
|
||||
).Enqueue(context.Background(), asyncRequest()))
|
||||
|
||||
require.Equal(t, AuditMetricsSnapshot{}, metrics.AuditSnapshot())
|
||||
})
|
||||
}
|
||||
|
||||
func workerJob(attempts, maxAttempts int) *Job {
|
||||
return &Job{ID: 51, ClaimVersion: 3, Attempts: attempts, MaxAttempts: maxAttempts, ConfigVersion: 7,
|
||||
Snapshot: PromptSnapshot{RequestID: "worker-request", PromptLength: 6, RedactedPreview: "red***"}}
|
||||
}
|
||||
|
||||
func TestWorkerCompletesPassWithoutEventRefreshesEveryChunkAndDeletesPayload(t *testing.T) {
|
||||
repo := &fakeJobRepository{}
|
||||
payload := &fakePayloadStore{values: map[int64]string{51: "abcdef"}}
|
||||
scannerCalls := 0
|
||||
scanner := PromptScannerFunc(func(_ context.Context, endpoint ActiveEndpoint, chunk string, _ []string) (*NormalizedResult, error) {
|
||||
scannerCalls++
|
||||
return &NormalizedResult{Decision: EventPass, RiskLevel: RiskLow, Action: ActionAllow, Safety: "Safe", Categories: []string{}, MatchedScanners: []string{}, ScannerScores: map[string]float64{}, ScannerEvidence: map[string]string{}, GuardEndpointID: endpoint.ID}, nil
|
||||
})
|
||||
metrics := NewAtomicMetrics()
|
||||
runner := NewRunner(&fakeConfigStore{cfg: asyncConfig(), active: true}, repo, payload, scanner, metrics)
|
||||
runner.clock = fixedClock{now: time.Unix(100, 0).UTC()}
|
||||
require.NoError(t, runner.processJob(context.Background(), 0, asyncConfig(), workerJob(1, 3)))
|
||||
require.Equal(t, 2, scannerCalls)
|
||||
require.Equal(t, 2, repo.refreshes)
|
||||
require.NotNil(t, repo.completedResult)
|
||||
require.Equal(t, EventPass, repo.completedResult.Decision)
|
||||
require.False(t, repo.completedStore)
|
||||
require.Equal(t, []int64{51}, payload.deleted)
|
||||
require.Equal(t, int64(1), metrics.Snapshot().Total)
|
||||
require.Equal(t, int64(1), metrics.Snapshot().Allowed)
|
||||
}
|
||||
|
||||
func TestWorkerRetryBackoffTerminalFailureAndFailover(t *testing.T) {
|
||||
now := time.Unix(200, 0).UTC()
|
||||
for _, tt := range []struct {
|
||||
name string
|
||||
attempts int
|
||||
maxAttempts int
|
||||
err *GuardError
|
||||
wantRetry bool
|
||||
wantBackoff time.Duration
|
||||
}{
|
||||
{name: "first retry", attempts: 1, maxAttempts: 3, err: &GuardError{Code: ErrorCodeUnavailable, Retryable: true}, wantRetry: true, wantBackoff: 5 * time.Second},
|
||||
{name: "second retry", attempts: 2, maxAttempts: 3, err: &GuardError{Code: ErrorCodeUnavailable, Retryable: true}, wantRetry: true, wantBackoff: 30 * time.Second},
|
||||
{name: "third retry", attempts: 3, maxAttempts: 4, err: &GuardError{Code: ErrorCodeUnavailable, Retryable: true}, wantRetry: true, wantBackoff: 2 * time.Minute},
|
||||
{name: "max attempts", attempts: 3, maxAttempts: 3, err: &GuardError{Code: ErrorCodeUnavailable, Retryable: true}},
|
||||
{name: "invalid terminal", attempts: 1, maxAttempts: 3, err: &GuardError{Code: ErrorCodeInvalidResponse, Retryable: false}},
|
||||
} {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
repo := &fakeJobRepository{}
|
||||
payload := &fakePayloadStore{values: map[int64]string{51: "abc"}}
|
||||
metrics := NewAtomicMetrics()
|
||||
runner := NewRunner(&fakeConfigStore{cfg: asyncConfig(), active: true}, repo, payload, PromptScannerFunc(func(context.Context, ActiveEndpoint, string, []string) (*NormalizedResult, error) {
|
||||
return nil, tt.err
|
||||
}), metrics)
|
||||
runner.clock = fixedClock{now: now}
|
||||
err := runner.processJob(context.Background(), 0, asyncConfig(), workerJob(tt.attempts, tt.maxAttempts))
|
||||
require.Error(t, err)
|
||||
if tt.wantRetry {
|
||||
require.Equal(t, 1, repo.retried)
|
||||
require.Equal(t, now.Add(tt.wantBackoff), repo.retryAt)
|
||||
require.Empty(t, payload.deleted)
|
||||
} else {
|
||||
require.Equal(t, 1, repo.failed)
|
||||
require.Equal(t, tt.err.Code, repo.failedCode)
|
||||
require.Equal(t, []int64{51}, payload.deleted)
|
||||
}
|
||||
snapshot := metrics.Snapshot()
|
||||
require.Equal(t, int64(1), snapshot.Total)
|
||||
if tt.err.Code == ErrorCodeInvalidResponse {
|
||||
require.Equal(t, int64(1), snapshot.Invalid)
|
||||
} else {
|
||||
require.Equal(t, int64(1), snapshot.Unavailable)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
repo := &fakeJobRepository{}
|
||||
payload := &fakePayloadStore{values: map[int64]string{51: "abc"}}
|
||||
metrics := NewAtomicMetrics()
|
||||
scanner := PromptScannerFunc(func(_ context.Context, endpoint ActiveEndpoint, _ string, _ []string) (*NormalizedResult, error) {
|
||||
if endpoint.ID == "first" {
|
||||
return nil, &GuardError{Code: ErrorCodeUnavailable, Retryable: true}
|
||||
}
|
||||
return integrationResult(EventPass), nil
|
||||
})
|
||||
cfg := asyncConfig()
|
||||
cfg.Endpoints = []ActiveEndpoint{{ID: "first", Enabled: true, InputLimit: 10}, {ID: "second", Enabled: true, InputLimit: 10}}
|
||||
runner := NewRunner(&fakeConfigStore{cfg: cfg, active: true}, repo, payload, scanner, metrics)
|
||||
require.NoError(t, runner.processJob(context.Background(), 0, cfg, workerJob(1, 3)))
|
||||
require.Equal(t, int64(1), metrics.Snapshot().Failovers)
|
||||
}
|
||||
|
||||
func TestWorkerPanicLeaseLossAndLifecycleAreContained(t *testing.T) {
|
||||
t.Run("panic", func(t *testing.T) {
|
||||
repo := &fakeJobRepository{}
|
||||
payload := &fakePayloadStore{values: map[int64]string{51: "abc"}}
|
||||
runner := NewRunner(&fakeConfigStore{cfg: asyncConfig(), active: true}, repo, payload, PromptScannerFunc(func(context.Context, ActiveEndpoint, string, []string) (*NormalizedResult, error) {
|
||||
panic("scanner panic canary")
|
||||
}), NewAtomicMetrics())
|
||||
require.NotPanics(t, func() { runner.processSafely(context.Background(), 0, asyncConfig(), workerJob(1, 3)) })
|
||||
_, _, failed, _, _, code, message := runner.Snapshot()
|
||||
require.Equal(t, int64(1), failed)
|
||||
require.Equal(t, "worker_panic", code)
|
||||
require.NotContains(t, message, "canary")
|
||||
require.Equal(t, 1, repo.failed)
|
||||
})
|
||||
|
||||
t.Run("lease loss", func(t *testing.T) {
|
||||
repo := &fakeJobRepository{refreshErr: ErrLeaseLost}
|
||||
payload := &fakePayloadStore{values: map[int64]string{51: "abc"}}
|
||||
calls := 0
|
||||
runner := NewRunner(&fakeConfigStore{cfg: asyncConfig(), active: true}, repo, payload, PromptScannerFunc(func(context.Context, ActiveEndpoint, string, []string) (*NormalizedResult, error) {
|
||||
calls++
|
||||
return integrationResult(EventPass), nil
|
||||
}), NewAtomicMetrics())
|
||||
require.ErrorIs(t, runner.processJob(context.Background(), 0, asyncConfig(), workerJob(1, 3)), ErrLeaseLost)
|
||||
require.Zero(t, calls)
|
||||
require.Zero(t, repo.retried)
|
||||
require.Zero(t, repo.failed)
|
||||
})
|
||||
|
||||
t.Run("start and shutdown", func(t *testing.T) {
|
||||
cfg := asyncConfig()
|
||||
cfg.Enabled = false
|
||||
configStore := &fakeConfigStore{cfg: cfg, active: true}
|
||||
repo := &fakeJobRepository{}
|
||||
payload := &fakePayloadStore{pingErr: errors.New("redis unavailable")}
|
||||
runner := NewRunner(configStore, repo, payload, PromptScannerFunc(func(context.Context, ActiveEndpoint, string, []string) (*NormalizedResult, error) {
|
||||
return integrationResult(EventPass), nil
|
||||
}), NewAtomicMetrics())
|
||||
require.NoError(t, runner.Start(context.Background()))
|
||||
require.NoError(t, runner.Start(context.Background()))
|
||||
_, _, _, _, _, code, _ := runner.Snapshot()
|
||||
require.Equal(t, "payload_store_unavailable", code)
|
||||
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
|
||||
defer cancel()
|
||||
require.NoError(t, runner.Shutdown(ctx))
|
||||
require.NoError(t, runner.Shutdown(ctx))
|
||||
})
|
||||
|
||||
t.Run("shutdown timeout is bounded", func(t *testing.T) {
|
||||
runner := &Runner{}
|
||||
release := make(chan struct{})
|
||||
runner.wg.Add(1)
|
||||
go func() {
|
||||
defer runner.wg.Done()
|
||||
<-release
|
||||
}()
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 20*time.Millisecond)
|
||||
defer cancel()
|
||||
require.ErrorIs(t, runner.Shutdown(ctx), context.DeadlineExceeded)
|
||||
close(release)
|
||||
ctx2, cancel2 := context.WithTimeout(context.Background(), time.Second)
|
||||
defer cancel2()
|
||||
require.NoError(t, runner.Shutdown(ctx2))
|
||||
})
|
||||
}
|
||||
|
||||
func TestPromptAuditSyntheticAsyncBaseline(t *testing.T) {
|
||||
const totalRequests = 100
|
||||
cfg := asyncConfig()
|
||||
cfg.Endpoints[0].InputLimit = 256
|
||||
cfg.StorePassEvents = false
|
||||
repo := &fakeJobRepository{}
|
||||
payload := &fakePayloadStore{values: make(map[int64]string, totalRequests)}
|
||||
metrics := NewAtomicMetrics()
|
||||
knownBenignFindings := 0
|
||||
knownMaliciousBlocked := 0
|
||||
scanner := PromptScannerFunc(func(_ context.Context, endpoint ActiveEndpoint, chunk string, _ []string) (*NormalizedResult, error) {
|
||||
switch {
|
||||
case strings.HasPrefix(chunk, "benign"):
|
||||
return &NormalizedResult{Decision: EventPass, RiskLevel: RiskLow, Action: ActionAllow, Safety: "Safe", GuardEndpointID: endpoint.ID}, nil
|
||||
case strings.HasPrefix(chunk, "flag"):
|
||||
return &NormalizedResult{Decision: EventFlag, RiskLevel: RiskMedium, Action: ActionWarn, Safety: "Controversial", Categories: []string{"politically_sensitive_topics"}, GuardEndpointID: endpoint.ID}, nil
|
||||
case strings.HasPrefix(chunk, "block"):
|
||||
knownMaliciousBlocked++
|
||||
return &NormalizedResult{Decision: EventCritical, RiskLevel: RiskCritical, Action: ActionBlock, Safety: "Unsafe", Categories: []string{"jailbreak"}, GuardEndpointID: endpoint.ID}, nil
|
||||
case strings.HasPrefix(chunk, "invalid"):
|
||||
return nil, &GuardError{Code: ErrorCodeInvalidResponse}
|
||||
default:
|
||||
return nil, &GuardError{Code: ErrorCodeUnavailable, Retryable: true, Timeout: true}
|
||||
}
|
||||
})
|
||||
runner := NewRunner(&fakeConfigStore{cfg: cfg, active: true}, repo, payload, scanner, metrics)
|
||||
runner.clock = &advancingClock{now: time.Unix(1_000, 0).UTC(), step: time.Millisecond}
|
||||
|
||||
for index := 1; index <= totalRequests; index++ {
|
||||
text := fmt.Sprintf("benign-%03d", index)
|
||||
switch {
|
||||
case index > 90 && index <= 95:
|
||||
text = fmt.Sprintf("flag-%03d", index)
|
||||
case index > 95 && index <= 98:
|
||||
text = fmt.Sprintf("block-%03d", index)
|
||||
case index == 99:
|
||||
text = "invalid-099"
|
||||
case index == 100:
|
||||
text = "timeout-100"
|
||||
}
|
||||
jobID := int64(index)
|
||||
payload.values[jobID] = text
|
||||
job := &Job{ID: jobID, ClaimVersion: 1, Attempts: 1, MaxAttempts: 1, ConfigVersion: cfg.ConfigVersion,
|
||||
Snapshot: PromptSnapshot{RequestID: fmt.Sprintf("baseline-%03d", index), PromptLength: len([]rune(text)), RedactedPreview: "synthetic"}}
|
||||
err := runner.processJob(context.Background(), 0, cfg, job)
|
||||
if index <= 98 {
|
||||
require.NoError(t, err)
|
||||
} else {
|
||||
require.Error(t, err)
|
||||
}
|
||||
}
|
||||
|
||||
snapshot := metrics.Snapshot()
|
||||
require.Equal(t, int64(totalRequests), snapshot.Total)
|
||||
require.Equal(t, int64(90), snapshot.Allowed)
|
||||
require.Equal(t, int64(5), snapshot.Flagged)
|
||||
require.Equal(t, int64(3), snapshot.Blocked)
|
||||
require.Equal(t, int64(1), snapshot.Invalid)
|
||||
require.Equal(t, int64(1), snapshot.Unavailable)
|
||||
require.Equal(t, int64(1), snapshot.Timeouts)
|
||||
require.Zero(t, knownBenignFindings)
|
||||
require.Equal(t, 3, knownMaliciousBlocked)
|
||||
require.Equal(t, 98, repo.completeCount)
|
||||
require.Equal(t, 8, repo.eventCount, "store_pass_events=false only grows events for flag/block fixtures")
|
||||
require.Positive(t, snapshot.LatencyP50MS)
|
||||
require.LessOrEqual(t, snapshot.LatencyP50MS, snapshot.LatencyP95MS)
|
||||
require.LessOrEqual(t, snapshot.LatencyP95MS, snapshot.LatencyP99MS)
|
||||
t.Logf("synthetic async baseline: p50=%dms p95=%dms p99=%dms failure_rate=2%% false_positive_rate=0%% event_growth=8/100", snapshot.LatencyP50MS, snapshot.LatencyP95MS, snapshot.LatencyP99MS)
|
||||
}
|
||||
|
||||
func TestRequestCloneOwnsMutableInputs(t *testing.T) {
|
||||
groupID := int64(7)
|
||||
req := Request{Body: []byte("original"), GroupID: &groupID}
|
||||
clone := req.Clone()
|
||||
clone.Body[0] = 'X'
|
||||
*clone.GroupID = 8
|
||||
require.Equal(t, []byte("original"), req.Body)
|
||||
require.Equal(t, int64(7), *req.GroupID)
|
||||
require.False(t, reflect.ValueOf(req.Body).Pointer() == reflect.ValueOf(clone.Body).Pointer())
|
||||
}
|
||||
@@ -22,6 +22,7 @@ const (
|
||||
auditCtxKeyActorID = "audit_actor_id"
|
||||
auditCtxKeyActorEmail = "audit_actor_email"
|
||||
auditCtxKeySkip = "audit_skip"
|
||||
auditCtxKeyExtra = "audit_extra"
|
||||
// ContextKeyAuthEmail 认证中间件写入的用户邮箱(审计用)。
|
||||
ContextKeyAuthEmail = "auth_email"
|
||||
// ContextKeySessionID 认证中间件写入的会话 ID(refresh token family)。
|
||||
@@ -48,6 +49,65 @@ func SkipAudit(c *gin.Context) {
|
||||
c.Set(auditCtxKeySkip, true)
|
||||
}
|
||||
|
||||
// auditExtraAllowedKeys is deliberately narrow: handlers may only attach
|
||||
// scalar, non-secret operation summaries. Request bodies and arbitrary maps
|
||||
// are never accepted through this channel.
|
||||
var auditExtraAllowedKeys = map[string]struct{}{
|
||||
"result": {}, "error_code": {}, "enabled": {}, "blocking_enabled": {},
|
||||
"config_version": {}, "endpoint_count": {}, "scanner_count": {},
|
||||
"all_groups": {}, "group_count": {}, "guard_endpoint_id": {},
|
||||
"http_status": {}, "latency_ms": {}, "token_applied": {}, "retryable": {},
|
||||
"event_id": {}, "requested_count": {}, "deleted_events": {}, "deleted_jobs": {},
|
||||
"matched_count": {}, "snapshot_max_id": {}, "filter_hash": {}, "confirm": {},
|
||||
}
|
||||
|
||||
// SetAuditExtra adds allowlisted, scalar details to the current audit entry.
|
||||
// It is safe to call more than once; later values replace earlier ones.
|
||||
func SetAuditExtra(c *gin.Context, fields map[string]any) {
|
||||
if c == nil || len(fields) == 0 {
|
||||
return
|
||||
}
|
||||
current := map[string]any{}
|
||||
if value, ok := c.Get(auditCtxKeyExtra); ok {
|
||||
if existing, ok := value.(map[string]any); ok {
|
||||
for key, item := range existing {
|
||||
current[key] = item
|
||||
}
|
||||
}
|
||||
}
|
||||
for key, value := range fields {
|
||||
if _, ok := auditExtraAllowedKeys[key]; !ok || !isAuditExtraScalar(value) {
|
||||
continue
|
||||
}
|
||||
if text, ok := value.(string); ok {
|
||||
value = truncateAuditExtraString(text, 128)
|
||||
}
|
||||
current[key] = value
|
||||
}
|
||||
c.Set(auditCtxKeyExtra, current)
|
||||
}
|
||||
|
||||
func isAuditExtraScalar(value any) bool {
|
||||
switch value.(type) {
|
||||
case string, bool,
|
||||
int, int8, int16, int32, int64,
|
||||
uint, uint8, uint16, uint32, uint64,
|
||||
float32, float64:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func truncateAuditExtraString(value string, limit int) string {
|
||||
value = strings.TrimSpace(value)
|
||||
runes := []rune(value)
|
||||
if len(runes) <= limit {
|
||||
return value
|
||||
}
|
||||
return string(runes[:limit])
|
||||
}
|
||||
|
||||
// auditSensitiveReads 需要审计的敏感 GET 读取(method+FullPath → 动作名)。
|
||||
var auditSensitiveReads = map[string]string{
|
||||
"GET /api/v1/admin/accounts/data": "admin.accounts.export",
|
||||
@@ -63,25 +123,37 @@ var auditSensitiveReads = map[string]string{
|
||||
|
||||
// auditActionOverrides 变更类请求的动作名精确映射(未命中时自动推导)。
|
||||
var auditActionOverrides = map[string]string{
|
||||
"POST /api/v1/auth/login": service.AuditActionLogin,
|
||||
"POST /api/v1/auth/login/2fa": service.AuditActionLogin2FA,
|
||||
"POST /api/v1/auth/register": service.AuditActionRegister,
|
||||
"POST /api/v1/auth/refresh": service.AuditActionTokenRefresh,
|
||||
"POST /api/v1/user/totp/step-up": service.AuditActionStepUpVerify,
|
||||
"POST /api/v1/admin/audit-logs/clear": service.AuditActionAuditLogClear,
|
||||
"POST /api/v1/admin/accounts/data": "admin.accounts.import",
|
||||
"POST /api/v1/admin/backups": "admin.backups.create",
|
||||
"POST /api/v1/admin/backups/:id/restore": "admin.backups.restore",
|
||||
"DELETE /api/v1/admin/backups/:id": "admin.backups.delete",
|
||||
"PUT /api/v1/admin/backups/s3-config": "admin.backups.s3_config.update",
|
||||
"POST /api/v1/admin/settings/admin-api-key/regenerate": "admin.admin_api_key.regenerate",
|
||||
"DELETE /api/v1/admin/settings/admin-api-key": "admin.admin_api_key.delete",
|
||||
"POST /api/v1/auth/login": service.AuditActionLogin,
|
||||
"POST /api/v1/auth/login/2fa": service.AuditActionLogin2FA,
|
||||
"POST /api/v1/auth/register": service.AuditActionRegister,
|
||||
"POST /api/v1/auth/refresh": service.AuditActionTokenRefresh,
|
||||
"POST /api/v1/user/totp/step-up": service.AuditActionStepUpVerify,
|
||||
"POST /api/v1/admin/audit-logs/clear": service.AuditActionAuditLogClear,
|
||||
"POST /api/v1/admin/accounts/data": "admin.accounts.import",
|
||||
"POST /api/v1/admin/backups": "admin.backups.create",
|
||||
"POST /api/v1/admin/backups/:id/restore": "admin.backups.restore",
|
||||
"DELETE /api/v1/admin/backups/:id": "admin.backups.delete",
|
||||
"PUT /api/v1/admin/backups/s3-config": "admin.backups.s3_config.update",
|
||||
"POST /api/v1/admin/settings/admin-api-key/regenerate": "admin.admin_api_key.regenerate",
|
||||
"DELETE /api/v1/admin/settings/admin-api-key": "admin.admin_api_key.delete",
|
||||
"PUT /api/v1/admin/prompt-audit/config": "admin.prompt_audit.config.update",
|
||||
"POST /api/v1/admin/prompt-audit/endpoints/probe": "admin.prompt_audit.endpoint.probe",
|
||||
"DELETE /api/v1/admin/prompt-audit/events/:id": "admin.prompt_audit.event.delete",
|
||||
"POST /api/v1/admin/prompt-audit/events/batch-delete": "admin.prompt_audit.events.batch_delete",
|
||||
"POST /api/v1/admin/prompt-audit/events/delete-preview": "admin.prompt_audit.events.delete_preview",
|
||||
"POST /api/v1/admin/prompt-audit/events/delete-by-filter": "admin.prompt_audit.events.filter_delete",
|
||||
}
|
||||
|
||||
// auditBodyOmittedRoutes 请求体几乎整体由凭证构成的路由(如整块粘贴 auth JSON 的导入接口)。
|
||||
// 这类 body 的凭证内嵌在普通字符串值里,键级脱敏无法覆盖,整体不入库。
|
||||
var auditBodyOmittedRoutes = map[string]struct{}{
|
||||
"POST /api/v1/admin/accounts/import/codex-session": {},
|
||||
"POST /api/v1/admin/accounts/import/codex-session": {},
|
||||
"PUT /api/v1/admin/prompt-audit/config": {},
|
||||
"POST /api/v1/admin/prompt-audit/endpoints/probe": {},
|
||||
"DELETE /api/v1/admin/prompt-audit/events/:id": {},
|
||||
"POST /api/v1/admin/prompt-audit/events/batch-delete": {},
|
||||
"POST /api/v1/admin/prompt-audit/events/delete-preview": {},
|
||||
"POST /api/v1/admin/prompt-audit/events/delete-by-filter": {},
|
||||
}
|
||||
|
||||
// NewAuditLogMiddleware 创建审计中间件。
|
||||
@@ -196,6 +268,13 @@ func NewAuditLogMiddleware(auditService *service.AuditLogService) AuditLogMiddle
|
||||
entry.CredentialMasked = MaskedRequestCredential(c)
|
||||
|
||||
extra := map[string]any{}
|
||||
if value, ok := c.Get(auditCtxKeyExtra); ok {
|
||||
if details, ok := value.(map[string]any); ok {
|
||||
for key, item := range details {
|
||||
extra[key] = item
|
||||
}
|
||||
}
|
||||
}
|
||||
if len(c.Params) > 0 {
|
||||
params := make(map[string]string, len(c.Params))
|
||||
for _, p := range c.Params {
|
||||
|
||||
@@ -1,6 +1,18 @@
|
||||
package middleware
|
||||
|
||||
import "testing"
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/service"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestDeriveAuditAction(t *testing.T) {
|
||||
cases := []struct {
|
||||
@@ -20,3 +32,116 @@ func TestDeriveAuditAction(t *testing.T) {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
type auditCaptureRepository struct {
|
||||
mu sync.Mutex
|
||||
logs []*service.AuditLog
|
||||
}
|
||||
|
||||
func (r *auditCaptureRepository) BatchInsert(_ context.Context, logs []*service.AuditLog) (int64, error) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.logs = append(r.logs, logs...)
|
||||
return int64(len(logs)), nil
|
||||
}
|
||||
func (r *auditCaptureRepository) Insert(_ context.Context, log *service.AuditLog) error {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.logs = append(r.logs, log)
|
||||
return nil
|
||||
}
|
||||
func (r *auditCaptureRepository) List(context.Context, *service.AuditLogFilter) (*service.AuditLogList, error) {
|
||||
return &service.AuditLogList{}, nil
|
||||
}
|
||||
func (r *auditCaptureRepository) GetByID(context.Context, int64) (*service.AuditLog, error) {
|
||||
return nil, service.ErrAuditLogNotFound
|
||||
}
|
||||
func (r *auditCaptureRepository) Count(context.Context) (int64, error) { return 0, nil }
|
||||
func (r *auditCaptureRepository) TruncateAll(context.Context) error { return nil }
|
||||
func (r *auditCaptureRepository) DeleteBefore(context.Context, time.Time, int) (int64, error) {
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
func TestPromptAuditAdminOperationsUseOmittedBodiesAndAllowlistedDetails(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
repository := &auditCaptureRepository{}
|
||||
auditService := service.NewAuditLogService(repository, nil)
|
||||
auditService.Start()
|
||||
|
||||
router := gin.New()
|
||||
router.Use(func(c *gin.Context) {
|
||||
c.Set(string(ContextKeyUser), AuthSubject{UserID: 77})
|
||||
c.Set(string(ContextKeyUserRole), "admin")
|
||||
c.Next()
|
||||
})
|
||||
router.Use(gin.HandlerFunc(NewAuditLogMiddleware(auditService)))
|
||||
router.PUT("/api/v1/admin/prompt-audit/config", func(c *gin.Context) {
|
||||
SetAuditExtra(c, map[string]any{
|
||||
"result": "failed", "error_code": "prompt_audit_config_conflict", "config_version": int64(9),
|
||||
"token": "audit-canary-secret", "raw_prompt": "audit-canary-prompt", "nested": map[string]any{"unsafe": true},
|
||||
})
|
||||
c.JSON(http.StatusConflict, gin.H{"ok": false})
|
||||
})
|
||||
router.POST("/api/v1/admin/prompt-audit/endpoints/probe", func(c *gin.Context) {
|
||||
SetAuditExtra(c, map[string]any{
|
||||
"result": "success", "guard_endpoint_id": "guard-1", "http_status": 200,
|
||||
"latency_ms": 12, "token_applied": true,
|
||||
})
|
||||
c.JSON(http.StatusOK, gin.H{"ok": true})
|
||||
})
|
||||
|
||||
for _, request := range []*http.Request{
|
||||
httptest.NewRequest(http.MethodPut, "/api/v1/admin/prompt-audit/config", bytes.NewBufferString(`{"expected_config_version":8,"token":"audit-canary-secret"}`)),
|
||||
httptest.NewRequest(http.MethodPost, "/api/v1/admin/prompt-audit/endpoints/probe", bytes.NewBufferString(`{"endpoint":{"token":"audit-canary-secret"}}`)),
|
||||
} {
|
||||
request.Header.Set("Content-Type", "application/json")
|
||||
recorder := httptest.NewRecorder()
|
||||
router.ServeHTTP(recorder, request)
|
||||
}
|
||||
auditService.Stop()
|
||||
|
||||
repository.mu.Lock()
|
||||
logs := append([]*service.AuditLog(nil), repository.logs...)
|
||||
repository.mu.Unlock()
|
||||
require.Len(t, logs, 2)
|
||||
|
||||
byAction := make(map[string]*service.AuditLog, len(logs))
|
||||
for _, entry := range logs {
|
||||
byAction[entry.Action] = entry
|
||||
require.Equal(t, "<credential-bearing body omitted>", entry.RequestBody)
|
||||
require.NotContains(t, entry.RequestBody, "audit-canary")
|
||||
require.NotContains(t, entry.Extra, "token")
|
||||
require.NotContains(t, entry.Extra, "raw_prompt")
|
||||
require.NotContains(t, entry.Extra, "nested")
|
||||
}
|
||||
|
||||
config := byAction["admin.prompt_audit.config.update"]
|
||||
require.NotNil(t, config)
|
||||
require.Equal(t, http.StatusConflict, config.StatusCode)
|
||||
require.Equal(t, "failed", config.Extra["result"])
|
||||
require.Equal(t, "prompt_audit_config_conflict", config.Extra["error_code"])
|
||||
require.EqualValues(t, 9, config.Extra["config_version"])
|
||||
|
||||
probe := byAction["admin.prompt_audit.endpoint.probe"]
|
||||
require.NotNil(t, probe)
|
||||
require.Equal(t, http.StatusOK, probe.StatusCode)
|
||||
require.Equal(t, "success", probe.Extra["result"])
|
||||
require.Equal(t, "guard-1", probe.Extra["guard_endpoint_id"])
|
||||
require.Equal(t, true, probe.Extra["token_applied"])
|
||||
}
|
||||
|
||||
func TestPromptAuditMutationAuditRoutesHaveStableActionsAndOmitBodies(t *testing.T) {
|
||||
expected := map[string]string{
|
||||
"PUT /api/v1/admin/prompt-audit/config": "admin.prompt_audit.config.update",
|
||||
"POST /api/v1/admin/prompt-audit/endpoints/probe": "admin.prompt_audit.endpoint.probe",
|
||||
"DELETE /api/v1/admin/prompt-audit/events/:id": "admin.prompt_audit.event.delete",
|
||||
"POST /api/v1/admin/prompt-audit/events/batch-delete": "admin.prompt_audit.events.batch_delete",
|
||||
"POST /api/v1/admin/prompt-audit/events/delete-preview": "admin.prompt_audit.events.delete_preview",
|
||||
"POST /api/v1/admin/prompt-audit/events/delete-by-filter": "admin.prompt_audit.events.filter_delete",
|
||||
}
|
||||
for route, action := range expected {
|
||||
require.Equal(t, action, auditActionOverrides[route])
|
||||
_, omitted := auditBodyOmittedRoutes[route]
|
||||
require.Truef(t, omitted, "%s must not persist its credential or confirmation-bearing body", route)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -108,6 +108,9 @@ func RegisterAdminRoutes(
|
||||
// 风控中心
|
||||
registerContentModerationRoutes(admin, h)
|
||||
|
||||
// 独立提示词输入审计
|
||||
registerPromptAuditRoutes(admin, h)
|
||||
|
||||
// 邀请返利(专属用户管理)
|
||||
registerAffiliateRoutes(admin, h)
|
||||
|
||||
@@ -116,6 +119,22 @@ func RegisterAdminRoutes(
|
||||
}
|
||||
}
|
||||
|
||||
func registerPromptAuditRoutes(admin *gin.RouterGroup, h *handler.Handlers) {
|
||||
promptAudit := admin.Group("/prompt-audit")
|
||||
{
|
||||
promptAudit.GET("/config", h.Admin.PromptAudit.GetConfig)
|
||||
promptAudit.PUT("/config", h.Admin.PromptAudit.UpdateConfig)
|
||||
promptAudit.POST("/endpoints/probe", h.Admin.PromptAudit.ProbeEndpoint)
|
||||
promptAudit.GET("/runtime", h.Admin.PromptAudit.GetRuntime)
|
||||
promptAudit.GET("/events", h.Admin.PromptAudit.ListEvents)
|
||||
promptAudit.GET("/events/:id", h.Admin.PromptAudit.GetEvent)
|
||||
promptAudit.DELETE("/events/:id", h.Admin.PromptAudit.DeleteEvent)
|
||||
promptAudit.POST("/events/batch-delete", h.Admin.PromptAudit.BatchDelete)
|
||||
promptAudit.POST("/events/delete-preview", h.Admin.PromptAudit.DeletePreview)
|
||||
promptAudit.POST("/events/delete-by-filter", h.Admin.PromptAudit.DeleteByFilter)
|
||||
}
|
||||
}
|
||||
|
||||
func registerAuditLogRoutes(admin *gin.RouterGroup, h *handler.Handlers, _ middleware.StepUpAuthMiddleware) {
|
||||
auditLogs := admin.Group("/audit-logs")
|
||||
{
|
||||
|
||||
@@ -0,0 +1,136 @@
|
||||
package routes
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"sort"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/Wei-Shaw/sub2api/internal/handler"
|
||||
"github.com/Wei-Shaw/sub2api/internal/securityaudit"
|
||||
servermiddleware "github.com/Wei-Shaw/sub2api/internal/server/middleware"
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
func TestEveryGatewayPOSTRouteIsClassifiedForPromptAuditCoverage(t *testing.T) {
|
||||
routeSource, err := os.ReadFile("gateway.go")
|
||||
require.NoError(t, err)
|
||||
pattern := regexp.MustCompile(`(?:gateway|gemini|r|codexDirect|antigravityV1|antigravityV1Beta)\.POST\("([^"]+)"`)
|
||||
matches := pattern.FindAllStringSubmatch(string(routeSource), -1)
|
||||
actual := map[string]struct{}{}
|
||||
for _, match := range matches {
|
||||
actual[match[1]] = struct{}{}
|
||||
}
|
||||
|
||||
audited := map[string][]string{
|
||||
"/messages": {"gateway_handler.go", "openai_gateway_handler.go"},
|
||||
"/responses": {"gateway_handler_responses.go", "openai_gateway_handler.go"},
|
||||
"/responses/*subpath": {"gateway_handler_responses.go", "openai_gateway_handler.go"},
|
||||
"/chat/completions": {"gateway_handler_chat_completions.go", "openai_chat_completions.go"},
|
||||
"/embeddings": {"openai_embeddings.go"},
|
||||
"/alpha/search": {"openai_alpha_search.go"},
|
||||
"/images/generations": {"openai_images.go", "grok_media.go"},
|
||||
"/images/edits": {"openai_images.go", "grok_media.go"},
|
||||
"/images/generations/async": {"image_task_handler.go"},
|
||||
"/images/edits/async": {"image_task_handler.go"},
|
||||
"/images/batches": {"batch_image_handler.go"},
|
||||
"/videos/generations": {"grok_media.go"},
|
||||
"/videos/edits": {"grok_media.go"},
|
||||
"/videos/extensions": {"grok_media.go"},
|
||||
"/models/*modelAction": {"gemini_v1beta_handler.go"},
|
||||
}
|
||||
excluded := map[string]string{
|
||||
"/messages/count_tokens": "tokenization only; it does not execute a model request",
|
||||
"/images/batches/:id/cancel": "control-plane cancellation with no user prompt",
|
||||
}
|
||||
|
||||
unclassified := make([]string, 0)
|
||||
for route := range actual {
|
||||
if _, ok := audited[route]; ok {
|
||||
continue
|
||||
}
|
||||
if _, ok := excluded[route]; ok {
|
||||
continue
|
||||
}
|
||||
unclassified = append(unclassified, route)
|
||||
}
|
||||
sort.Strings(unclassified)
|
||||
require.Empty(t, unclassified, "new gateway POST routes must be audited or explicitly classified with a no-prompt reason")
|
||||
|
||||
for route, files := range audited {
|
||||
_, exists := actual[route]
|
||||
require.Truef(t, exists, "stale prompt-audit route manifest entry %s", route)
|
||||
for _, filename := range files {
|
||||
source, readErr := os.ReadFile(filepath.Join("..", "..", "handler", filename))
|
||||
require.NoError(t, readErr)
|
||||
require.Containsf(t, string(source), "checkSecurityAudit", "%s route handler %s bypasses Coordinator", route, filename)
|
||||
}
|
||||
}
|
||||
|
||||
for route, reason := range excluded {
|
||||
require.NotEmpty(t, strings.TrimSpace(reason))
|
||||
_, exists := actual[route]
|
||||
require.Truef(t, exists, "stale excluded route %s", route)
|
||||
}
|
||||
}
|
||||
|
||||
func TestResponsesWebSocketHasFirstAndSubsequentTurnPromptGates(t *testing.T) {
|
||||
routeSource, err := os.ReadFile("gateway.go")
|
||||
require.NoError(t, err)
|
||||
require.GreaterOrEqual(t, strings.Count(string(routeSource), `.GET("/responses"`), 2)
|
||||
handlerSource, err := os.ReadFile(filepath.Join("..", "..", "handler", "openai_gateway_handler.go"))
|
||||
require.NoError(t, err)
|
||||
require.Contains(t, string(handlerSource), `checkSecurityAuditStage`)
|
||||
require.Contains(t, string(handlerSource), `"first_turn"`)
|
||||
require.Contains(t, string(handlerSource), `"subsequent_turn"`)
|
||||
wsStart := strings.Index(string(handlerSource), `func (h *OpenAIGatewayHandler) ResponsesWebSocket`)
|
||||
require.NotEqual(t, -1, wsStart)
|
||||
wsSource := string(handlerSource)[wsStart:]
|
||||
require.Less(t,
|
||||
strings.Index(wsSource, `"first_turn"`),
|
||||
strings.Index(wsSource, `TryAcquireUserSlotForAPIKey`),
|
||||
"the first response.create gate must precede per-request user/account slots",
|
||||
)
|
||||
}
|
||||
|
||||
func TestPromptAuditAdminRoutesRejectUnauthenticatedAndNonAdminRequests(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
router := gin.New()
|
||||
handlers := &handler.Handlers{Admin: &handler.AdminHandlers{
|
||||
PromptAudit: securityaudit.NewPromptAdminHandler(nil),
|
||||
}}
|
||||
adminAuth := servermiddleware.AdminAuthMiddleware(func(c *gin.Context) {
|
||||
if c.GetHeader("Authorization") == "" {
|
||||
servermiddleware.AbortWithError(c, http.StatusUnauthorized, "UNAUTHORIZED", "Authorization required")
|
||||
return
|
||||
}
|
||||
servermiddleware.AbortWithError(c, http.StatusForbidden, "FORBIDDEN", "Admin access required")
|
||||
})
|
||||
auditLog := servermiddleware.AuditLogMiddleware(func(c *gin.Context) { c.Next() })
|
||||
stepUp := servermiddleware.StepUpAuthMiddleware(func(c *gin.Context) { c.Next() })
|
||||
RegisterAdminRoutes(router.Group("/api/v1"), handlers, adminAuth, auditLog, stepUp, nil)
|
||||
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
auth string
|
||||
wantStatus int
|
||||
}{
|
||||
{name: "unauthenticated", wantStatus: http.StatusUnauthorized},
|
||||
{name: "non-admin", auth: "Bearer user-token", wantStatus: http.StatusForbidden},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
recorder := httptest.NewRecorder()
|
||||
request := httptest.NewRequest(http.MethodGet, "/api/v1/admin/prompt-audit/config", nil)
|
||||
if tc.auth != "" {
|
||||
request.Header.Set("Authorization", tc.auth)
|
||||
}
|
||||
router.ServeHTTP(recorder, request)
|
||||
require.Equal(t, tc.wantStatus, recorder.Code)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,129 @@
|
||||
-- Independent OpenAI-compatible prompt input audit.
|
||||
-- Raw prompts and Guard credentials are intentionally absent from PostgreSQL.
|
||||
|
||||
CREATE TABLE IF NOT EXISTS prompt_audit_jobs (
|
||||
id BIGSERIAL PRIMARY KEY,
|
||||
request_id VARCHAR(128) NOT NULL DEFAULT '',
|
||||
user_id BIGINT REFERENCES users(id) ON DELETE SET NULL,
|
||||
username_snapshot VARCHAR(255) NOT NULL DEFAULT '',
|
||||
user_email_snapshot VARCHAR(320) NOT NULL DEFAULT '',
|
||||
api_key_id BIGINT REFERENCES api_keys(id) ON DELETE SET NULL,
|
||||
api_key_name_snapshot VARCHAR(255) NOT NULL DEFAULT '',
|
||||
group_id BIGINT REFERENCES groups(id) ON DELETE SET NULL,
|
||||
group_name VARCHAR(255) NOT NULL DEFAULT '',
|
||||
provider VARCHAR(64) NOT NULL DEFAULT '',
|
||||
endpoint VARCHAR(128) NOT NULL DEFAULT '',
|
||||
protocol VARCHAR(64) NOT NULL DEFAULT '',
|
||||
model VARCHAR(255) NOT NULL DEFAULT '',
|
||||
prompt_hash VARCHAR(64) NOT NULL DEFAULT '',
|
||||
redacted_preview TEXT NOT NULL DEFAULT '',
|
||||
prompt_length INT NOT NULL DEFAULT 0,
|
||||
message_count INT NOT NULL DEFAULT 0,
|
||||
execution_mode VARCHAR(32) NOT NULL DEFAULT 'async_audit',
|
||||
config_version BIGINT NOT NULL DEFAULT 1,
|
||||
status VARCHAR(32) NOT NULL DEFAULT 'staging',
|
||||
attempts INT NOT NULL DEFAULT 0,
|
||||
max_attempts INT NOT NULL DEFAULT 3,
|
||||
claim_version BIGINT NOT NULL DEFAULT 0,
|
||||
next_attempt_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||
processing_started_at TIMESTAMPTZ,
|
||||
processed_at TIMESTAMPTZ,
|
||||
last_error_code VARCHAR(64) NOT NULL DEFAULT '',
|
||||
last_error_message VARCHAR(512) NOT NULL DEFAULT '',
|
||||
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||
updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||
CONSTRAINT chk_prompt_audit_jobs_status
|
||||
CHECK (status IN ('staging', 'queued', 'processing', 'retry', 'done', 'failed')),
|
||||
CONSTRAINT chk_prompt_audit_jobs_execution_mode
|
||||
CHECK (execution_mode IN ('async_audit', 'blocking')),
|
||||
CONSTRAINT chk_prompt_audit_jobs_nonnegative
|
||||
CHECK (
|
||||
attempts >= 0 AND max_attempts >= 0 AND claim_version >= 0 AND
|
||||
prompt_length >= 0 AND message_count >= 0 AND config_version >= 1
|
||||
)
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS prompt_audit_events (
|
||||
id BIGSERIAL PRIMARY KEY,
|
||||
job_id BIGINT NOT NULL REFERENCES prompt_audit_jobs(id) ON DELETE CASCADE,
|
||||
request_id VARCHAR(128) NOT NULL DEFAULT '',
|
||||
user_id BIGINT REFERENCES users(id) ON DELETE SET NULL,
|
||||
username_snapshot VARCHAR(255) NOT NULL DEFAULT '',
|
||||
user_email_snapshot VARCHAR(320) NOT NULL DEFAULT '',
|
||||
api_key_id BIGINT REFERENCES api_keys(id) ON DELETE SET NULL,
|
||||
api_key_name_snapshot VARCHAR(255) NOT NULL DEFAULT '',
|
||||
group_id BIGINT REFERENCES groups(id) ON DELETE SET NULL,
|
||||
group_name VARCHAR(255) NOT NULL DEFAULT '',
|
||||
provider VARCHAR(64) NOT NULL DEFAULT '',
|
||||
endpoint VARCHAR(128) NOT NULL DEFAULT '',
|
||||
protocol VARCHAR(64) NOT NULL DEFAULT '',
|
||||
model VARCHAR(255) NOT NULL DEFAULT '',
|
||||
prompt_hash VARCHAR(64) NOT NULL DEFAULT '',
|
||||
redacted_preview TEXT NOT NULL DEFAULT '',
|
||||
decision VARCHAR(32) NOT NULL DEFAULT 'pass',
|
||||
risk_level VARCHAR(32) NOT NULL DEFAULT 'low',
|
||||
action VARCHAR(32) NOT NULL DEFAULT 'Allow',
|
||||
categories JSONB NOT NULL DEFAULT '[]'::jsonb,
|
||||
matched_scanners JSONB NOT NULL DEFAULT '[]'::jsonb,
|
||||
scanner_scores JSONB NOT NULL DEFAULT '{}'::jsonb,
|
||||
scanner_evidence JSONB NOT NULL DEFAULT '{}'::jsonb,
|
||||
scanner_backend VARCHAR(64) NOT NULL DEFAULT 'qwen3guard-openai',
|
||||
scanner_version VARCHAR(128) NOT NULL DEFAULT '',
|
||||
guard_endpoint_id VARCHAR(128) NOT NULL DEFAULT '',
|
||||
policy_id VARCHAR(128) NOT NULL DEFAULT '',
|
||||
policy_version INT NOT NULL DEFAULT 0,
|
||||
config_version BIGINT NOT NULL DEFAULT 1,
|
||||
chunk_total INT NOT NULL DEFAULT 0,
|
||||
latency_ms INT NOT NULL DEFAULT 0,
|
||||
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||
CONSTRAINT chk_prompt_audit_events_decision
|
||||
CHECK (decision IN ('pass', 'flag', 'critical')),
|
||||
CONSTRAINT chk_prompt_audit_events_risk_level
|
||||
CHECK (risk_level IN ('low', 'medium', 'high', 'critical')),
|
||||
CONSTRAINT chk_prompt_audit_events_action
|
||||
CHECK (action IN ('Allow', 'Warn', 'Block')),
|
||||
CONSTRAINT chk_prompt_audit_events_nonnegative
|
||||
CHECK (policy_version >= 0 AND config_version >= 1 AND chunk_total >= 0 AND latency_ms >= 0),
|
||||
CONSTRAINT chk_prompt_audit_events_categories_json
|
||||
CHECK (jsonb_typeof(categories) = 'array'),
|
||||
CONSTRAINT chk_prompt_audit_events_scanners_json
|
||||
CHECK (jsonb_typeof(matched_scanners) = 'array'),
|
||||
CONSTRAINT chk_prompt_audit_events_scores_json
|
||||
CHECK (jsonb_typeof(scanner_scores) = 'object'),
|
||||
CONSTRAINT chk_prompt_audit_events_evidence_json
|
||||
CHECK (jsonb_typeof(scanner_evidence) = 'object')
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_prompt_audit_jobs_schedule
|
||||
ON prompt_audit_jobs(status, next_attempt_at, id);
|
||||
CREATE INDEX IF NOT EXISTS idx_prompt_audit_jobs_request
|
||||
ON prompt_audit_jobs(request_id);
|
||||
CREATE INDEX IF NOT EXISTS idx_prompt_audit_jobs_user_created
|
||||
ON prompt_audit_jobs(user_id, created_at DESC);
|
||||
CREATE INDEX IF NOT EXISTS idx_prompt_audit_jobs_api_key_created
|
||||
ON prompt_audit_jobs(api_key_id, created_at DESC);
|
||||
CREATE INDEX IF NOT EXISTS idx_prompt_audit_jobs_group_created
|
||||
ON prompt_audit_jobs(group_id, created_at DESC);
|
||||
CREATE INDEX IF NOT EXISTS idx_prompt_audit_jobs_prompt_hash
|
||||
ON prompt_audit_jobs(prompt_hash);
|
||||
CREATE INDEX IF NOT EXISTS idx_prompt_audit_jobs_created
|
||||
ON prompt_audit_jobs(created_at DESC, id DESC);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_prompt_audit_events_job
|
||||
ON prompt_audit_events(job_id);
|
||||
CREATE INDEX IF NOT EXISTS idx_prompt_audit_events_request
|
||||
ON prompt_audit_events(request_id);
|
||||
CREATE INDEX IF NOT EXISTS idx_prompt_audit_events_decision_created
|
||||
ON prompt_audit_events(decision, created_at DESC, id DESC);
|
||||
CREATE INDEX IF NOT EXISTS idx_prompt_audit_events_risk_created
|
||||
ON prompt_audit_events(risk_level, created_at DESC, id DESC);
|
||||
CREATE INDEX IF NOT EXISTS idx_prompt_audit_events_user_created
|
||||
ON prompt_audit_events(user_id, created_at DESC, id DESC);
|
||||
CREATE INDEX IF NOT EXISTS idx_prompt_audit_events_api_key_created
|
||||
ON prompt_audit_events(api_key_id, created_at DESC, id DESC);
|
||||
CREATE INDEX IF NOT EXISTS idx_prompt_audit_events_group_created
|
||||
ON prompt_audit_events(group_id, created_at DESC, id DESC);
|
||||
CREATE INDEX IF NOT EXISTS idx_prompt_audit_events_prompt_hash
|
||||
ON prompt_audit_events(prompt_hash);
|
||||
CREATE INDEX IF NOT EXISTS idx_prompt_audit_events_created
|
||||
ON prompt_audit_events(created_at DESC, id DESC);
|
||||
@@ -771,7 +771,18 @@ const adminNavItems = computed((): NavItem[] => {
|
||||
{ path: '/admin/accounts', label: t('nav.accounts'), icon: GlobeIcon },
|
||||
{ path: '/admin/announcements', label: t('nav.announcements'), icon: BellIcon },
|
||||
{ path: '/admin/proxies', label: t('nav.proxies'), icon: ServerIcon },
|
||||
{ path: '/admin/risk-control', label: t('nav.riskControl'), icon: ShieldIcon, hideInSimpleMode: true, featureFlag: flagRiskControl },
|
||||
{
|
||||
path: '/admin/security-audit',
|
||||
label: t('nav.securityAudit'),
|
||||
icon: ShieldIcon,
|
||||
hideInSimpleMode: true,
|
||||
expandOnly: true,
|
||||
featureFlag: flagRiskControl,
|
||||
children: [
|
||||
{ path: '/admin/risk-control', label: t('nav.contentModeration'), icon: ShieldIcon },
|
||||
{ path: '/admin/prompt-audit', label: t('nav.promptAudit'), icon: ShieldIcon },
|
||||
],
|
||||
},
|
||||
{ path: '/admin/redeem', label: t('nav.redeemCodes'), icon: TicketIcon, hideInSimpleMode: true },
|
||||
{ path: '/admin/promo-codes', label: t('nav.promoCodes'), icon: GiftIcon, hideInSimpleMode: true },
|
||||
{
|
||||
|
||||
@@ -0,0 +1,350 @@
|
||||
<template>
|
||||
<AppLayout>
|
||||
<div class="mx-auto max-w-[1600px] pb-28">
|
||||
<header class="mb-6 flex flex-wrap items-end justify-between gap-4">
|
||||
<div>
|
||||
<p class="text-xs font-semibold uppercase tracking-[0.16em] text-primary-600 dark:text-primary-400">{{ t('nav.securityAudit') }}</p>
|
||||
<h1 class="mt-1 text-2xl font-semibold tracking-tight text-gray-950 dark:text-white">{{ t('admin.promptAudit.title') }}</h1>
|
||||
<p class="mt-2 max-w-3xl text-sm text-gray-500 dark:text-dark-300">{{ t('admin.promptAudit.description') }}</p>
|
||||
</div>
|
||||
<div v-if="draft" class="text-right text-xs text-gray-500 dark:text-dark-400">
|
||||
<p>{{ t('admin.promptAudit.configVersion', { version: draft.config_version }) }}</p>
|
||||
<p v-if="draft.updated_at" class="mt-1">{{ formatDate(draft.updated_at) }}</p>
|
||||
</div>
|
||||
</header>
|
||||
|
||||
<div v-if="loadErrors.config && !draft" role="alert" class="rounded-xl border border-red-200 bg-red-50 p-5 dark:border-red-900 dark:bg-red-950/30">
|
||||
<p class="text-sm text-red-700 dark:text-red-300">{{ loadErrors.config }}</p>
|
||||
<button type="button" class="btn btn-secondary btn-sm mt-3" @click="loadConfig">{{ t('admin.promptAudit.actions.retry') }}</button>
|
||||
</div>
|
||||
|
||||
<main v-else class="rounded-2xl border border-gray-200 bg-white px-4 shadow-sm dark:border-dark-700 dark:bg-dark-850 sm:px-6 lg:px-8">
|
||||
<RuntimeOverview :runtime="runtime" :loading="loading.runtime" :error="loadErrors.runtime" @refresh="loadRuntime" />
|
||||
|
||||
<template v-if="draft">
|
||||
<EndpointPool
|
||||
:endpoints="draft.endpoints"
|
||||
:probe-results="probeResults"
|
||||
:probing-ids="probingIds"
|
||||
@update:endpoints="updateEndpoints"
|
||||
@probe="runProbe"
|
||||
/>
|
||||
<div v-if="loadErrors.groups" role="alert" class="mt-5 rounded-lg bg-amber-50 px-4 py-3 text-sm text-amber-800 dark:bg-amber-950/30 dark:text-amber-200">{{ loadErrors.groups }}</div>
|
||||
<PolicyPanel :draft="draft" :groups="groups" @update:draft="replaceDraft" />
|
||||
</template>
|
||||
|
||||
<EventWorkspace
|
||||
:events="events.items"
|
||||
:total="events.total"
|
||||
:page="events.page"
|
||||
:page-size="events.page_size"
|
||||
:filters="filters"
|
||||
:selected-ids="selectedEventIds"
|
||||
:loading="loading.events"
|
||||
:error="loadErrors.events"
|
||||
@filters-change="handleFiltersChanged"
|
||||
@search="applyEventFilters"
|
||||
@selection="selectedEventIds = $event"
|
||||
@page="changePage"
|
||||
@page-size="changePageSize"
|
||||
@view="openEvent"
|
||||
@delete="requestSingleDelete"
|
||||
@batch-delete="requestBatchDelete"
|
||||
@preview-delete="requestFilterDeletePreview"
|
||||
/>
|
||||
</main>
|
||||
</div>
|
||||
|
||||
<div v-if="draft" class="fixed inset-x-0 bottom-0 z-30 border-t border-gray-200 bg-white/95 px-4 py-3 shadow-[0_-12px_35px_rgba(15,23,42,0.08)] backdrop-blur dark:border-dark-700 dark:bg-dark-900/95 lg:left-64">
|
||||
<div class="mx-auto flex max-w-[1600px] flex-wrap items-center justify-between gap-3">
|
||||
<div class="flex flex-wrap items-center gap-x-5 gap-y-2">
|
||||
<SaveToggle :label="t('admin.promptAudit.saveBar.enabled')" :model-value="draft.enabled" data-test="enabled-toggle" @update:model-value="setEnabled" />
|
||||
<SaveToggle :label="t('admin.promptAudit.saveBar.blocking')" :model-value="draft.blocking_enabled" :disabled="!draft.enabled" data-test="blocking-toggle" @update:model-value="setBlocking" />
|
||||
<SaveToggle :label="t('admin.promptAudit.saveBar.storePass')" :model-value="draft.store_pass_events" data-test="store-pass-toggle" @update:model-value="replaceDraft({ ...draft!, store_pass_events: $event })" />
|
||||
</div>
|
||||
<div class="flex items-center gap-3">
|
||||
<span class="text-sm" :class="dirty ? 'text-amber-700 dark:text-amber-300' : 'text-gray-500 dark:text-dark-400'">
|
||||
{{ dirty ? t('admin.promptAudit.saveBar.dirty') : t('admin.promptAudit.saveBar.synced') }}
|
||||
</span>
|
||||
<button type="button" class="btn btn-secondary" :disabled="!dirty || loading.saving" @click="resetDraft">{{ t('common.reset') }}</button>
|
||||
<button type="button" class="btn btn-primary" :disabled="!dirty || loading.saving" data-test="save-config" @click="saveConfig">
|
||||
{{ loading.saving ? t('common.saving') : t('common.save') }}
|
||||
</button>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<ConfirmDialog
|
||||
:show="showBlockingConfirmation"
|
||||
:title="t('admin.promptAudit.blockingConfirm.title')"
|
||||
:message="t('admin.promptAudit.blockingConfirm.message')"
|
||||
:confirm-text="t('admin.promptAudit.blockingConfirm.confirm')"
|
||||
danger
|
||||
@confirm="confirmBlocking"
|
||||
@cancel="showBlockingConfirmation = false"
|
||||
/>
|
||||
<ConfirmDialog
|
||||
:show="deleteRequest.mode !== ''"
|
||||
:title="t('admin.promptAudit.events.deleteConfirmTitle')"
|
||||
:message="t('admin.promptAudit.events.deleteConfirmMessage', { count: deleteRequest.ids.length })"
|
||||
:confirm-text="t('common.delete')"
|
||||
danger
|
||||
@confirm="confirmIDDelete"
|
||||
@cancel="clearDeleteRequest"
|
||||
/>
|
||||
<BaseDialog :show="Boolean(deletePreview)" :title="t('admin.promptAudit.events.filterDeleteTitle')" width="normal" @close="deletePreview = null">
|
||||
<div v-if="deletePreview" class="space-y-4 text-sm text-gray-600 dark:text-dark-300">
|
||||
<p>{{ t('admin.promptAudit.events.filterDeleteCount', { count: deletePreview.matched_count }) }}</p>
|
||||
<dl class="grid grid-cols-[auto_1fr] gap-x-3 gap-y-2">
|
||||
<dt>{{ t('admin.promptAudit.events.snapshotMax') }}</dt><dd>{{ deletePreview.snapshot_max_id }}</dd>
|
||||
<dt>Filter SHA-256</dt><dd class="break-all font-mono text-xs">{{ deletePreview.filter_hash }}</dd>
|
||||
<dt>{{ t('admin.promptAudit.events.expiresAt') }}</dt><dd>{{ formatDate(deletePreview.expires_at) }}</dd>
|
||||
</dl>
|
||||
<p class="rounded-lg bg-amber-50 px-3 py-2 text-amber-800 dark:bg-amber-950/30 dark:text-amber-200">{{ t('admin.promptAudit.events.filterDeleteWarning') }}</p>
|
||||
</div>
|
||||
<template #footer>
|
||||
<div class="flex justify-end gap-3">
|
||||
<button type="button" class="btn btn-secondary" @click="deletePreview = null">{{ t('common.cancel') }}</button>
|
||||
<button type="button" class="btn btn-danger" :disabled="loading.deleting" data-test="confirm-filter-delete" @click="confirmFilterDelete">{{ t('admin.promptAudit.events.confirmFilterDelete') }}</button>
|
||||
</div>
|
||||
</template>
|
||||
</BaseDialog>
|
||||
<EventDetailDialog :show="showEventDetail" :event="activeEvent" :loading="loading.detail" @close="closeEventDetail" />
|
||||
</AppLayout>
|
||||
</template>
|
||||
|
||||
<script setup lang="ts">
|
||||
import { computed, defineComponent, h, onMounted, reactive, ref } from 'vue'
|
||||
import { useI18n } from 'vue-i18n'
|
||||
import AppLayout from '@/components/layout/AppLayout.vue'
|
||||
import BaseDialog from '@/components/common/BaseDialog.vue'
|
||||
import ConfirmDialog from '@/components/common/ConfirmDialog.vue'
|
||||
import { useAppStore } from '@/stores/app'
|
||||
import { extractApiErrorCode, extractApiErrorMessage } from '@/utils/apiError'
|
||||
import RuntimeOverview from './components/RuntimeOverview.vue'
|
||||
import EndpointPool from './components/EndpointPool.vue'
|
||||
import PolicyPanel from './components/PolicyPanel.vue'
|
||||
import EventWorkspace from './components/EventWorkspace.vue'
|
||||
import EventDetailDialog from './components/EventDetailDialog.vue'
|
||||
import promptAuditAPI from './api'
|
||||
import type {
|
||||
PromptAuditDraft,
|
||||
PromptAuditEndpointDraft,
|
||||
PromptAuditEvent,
|
||||
PromptAuditGroup,
|
||||
PromptAuditRuntime,
|
||||
PromptDeletePreview,
|
||||
PromptEventFilters,
|
||||
PromptEventPage,
|
||||
PromptLoadErrors,
|
||||
PromptProbeResult,
|
||||
} from './types'
|
||||
import { buildUpdateRequest, cloneData, configToDraft, draftFingerprint, emptyEventFilters } from './viewModel'
|
||||
|
||||
const { t, locale } = useI18n()
|
||||
const appStore = useAppStore()
|
||||
const serverConfig = ref<PromptAuditDraft | null>(null)
|
||||
const draft = ref<PromptAuditDraft | null>(null)
|
||||
const runtime = ref<PromptAuditRuntime | null>(null)
|
||||
const groups = ref<PromptAuditGroup[]>([])
|
||||
const events = reactive<PromptEventPage>({ items: [], total: 0, page: 1, page_size: 20, pages: 0 })
|
||||
const filters = ref<PromptEventFilters>(emptyEventFilters())
|
||||
const appliedFilters = ref<PromptEventFilters>(emptyEventFilters())
|
||||
const selectedEventIds = ref<number[]>([])
|
||||
const activeEvent = ref<PromptAuditEvent | null>(null)
|
||||
const showEventDetail = ref(false)
|
||||
const probeResults = reactive<Record<string, PromptProbeResult>>({})
|
||||
const probingIds = ref<string[]>([])
|
||||
const deletePreview = ref<PromptDeletePreview | null>(null)
|
||||
const showBlockingConfirmation = ref(false)
|
||||
const deleteRequest = reactive<{ mode: '' | 'single' | 'batch'; ids: number[] }>({ mode: '', ids: [] })
|
||||
const loading = reactive({ config: false, runtime: false, groups: false, events: false, saving: false, detail: false, deleting: false })
|
||||
const loadErrors = reactive<PromptLoadErrors>({ config: '', runtime: '', groups: '', events: '' })
|
||||
const dirty = computed(() => draftFingerprint(draft.value) !== draftFingerprint(serverConfig.value))
|
||||
|
||||
const SaveToggle = defineComponent({
|
||||
inheritAttrs: false,
|
||||
props: { label: { type: String, required: true }, modelValue: { type: Boolean, required: true }, disabled: { type: Boolean, default: false } },
|
||||
emits: ['update:modelValue'],
|
||||
setup(props, { emit, attrs }) {
|
||||
return () => h('label', { class: ['flex items-center gap-2 text-sm', props.disabled ? 'cursor-not-allowed opacity-50' : 'cursor-pointer'] }, [
|
||||
h('button', {
|
||||
...attrs, type: 'button', role: 'switch', 'aria-checked': props.modelValue, 'aria-label': props.label, disabled: props.disabled,
|
||||
class: ['relative h-6 w-11 rounded-full transition-colors', props.modelValue ? 'bg-primary-600' : 'bg-gray-300 dark:bg-dark-600'],
|
||||
onClick: () => !props.disabled && emit('update:modelValue', !props.modelValue),
|
||||
}, [h('span', { class: ['absolute top-0.5 h-5 w-5 rounded-full bg-white shadow transition-transform', props.modelValue ? 'translate-x-5' : 'translate-x-0.5'] })]),
|
||||
h('span', { class: 'text-gray-700 dark:text-dark-200' }, props.label),
|
||||
])
|
||||
},
|
||||
})
|
||||
|
||||
function errorMessage(error: unknown, fallbackKey: string): string {
|
||||
const code = extractApiErrorCode(error)
|
||||
if (code) {
|
||||
const key = `admin.promptAudit.errors.${code}`
|
||||
const translated = t(key)
|
||||
if (translated !== key) return translated
|
||||
}
|
||||
return extractApiErrorMessage(error, t(fallbackKey))
|
||||
}
|
||||
|
||||
async function loadConfig() {
|
||||
loading.config = true
|
||||
loadErrors.config = ''
|
||||
try {
|
||||
const config = await promptAuditAPI.getConfig()
|
||||
serverConfig.value = configToDraft(config)
|
||||
draft.value = configToDraft(config)
|
||||
} catch (error) {
|
||||
loadErrors.config = errorMessage(error, 'admin.promptAudit.errors.loadConfig')
|
||||
} finally {
|
||||
loading.config = false
|
||||
}
|
||||
}
|
||||
async function loadRuntime() {
|
||||
loading.runtime = true
|
||||
loadErrors.runtime = ''
|
||||
try { runtime.value = await promptAuditAPI.getRuntime() }
|
||||
catch (error) { loadErrors.runtime = errorMessage(error, 'admin.promptAudit.errors.loadRuntime') }
|
||||
finally { loading.runtime = false }
|
||||
}
|
||||
async function loadGroups() {
|
||||
loading.groups = true
|
||||
loadErrors.groups = ''
|
||||
try { groups.value = await promptAuditAPI.listGroups() }
|
||||
catch (error) { loadErrors.groups = errorMessage(error, 'admin.promptAudit.errors.loadGroups') }
|
||||
finally { loading.groups = false }
|
||||
}
|
||||
async function loadEvents() {
|
||||
loading.events = true
|
||||
loadErrors.events = ''
|
||||
try {
|
||||
const result = await promptAuditAPI.listEvents(appliedFilters.value, events.page, events.page_size)
|
||||
Object.assign(events, result)
|
||||
selectedEventIds.value = []
|
||||
} catch (error) {
|
||||
loadErrors.events = errorMessage(error, 'admin.promptAudit.errors.loadEvents')
|
||||
} finally {
|
||||
loading.events = false
|
||||
}
|
||||
}
|
||||
async function loadInitial() {
|
||||
await Promise.allSettled([loadConfig(), loadRuntime(), loadGroups(), loadEvents()])
|
||||
}
|
||||
|
||||
function replaceDraft(value: PromptAuditDraft) { draft.value = cloneData(value) }
|
||||
function updateEndpoints(value: PromptAuditEndpointDraft[]) {
|
||||
if (!draft.value) return
|
||||
replaceDraft({ ...draft.value, endpoints: value })
|
||||
}
|
||||
function setEnabled(value: boolean) {
|
||||
if (!draft.value) return
|
||||
replaceDraft({ ...draft.value, enabled: value, blocking_enabled: value ? draft.value.blocking_enabled : false })
|
||||
}
|
||||
function setBlocking(value: boolean) {
|
||||
if (!draft.value || !draft.value.enabled) return
|
||||
if (value && !draft.value.blocking_enabled) { showBlockingConfirmation.value = true; return }
|
||||
replaceDraft({ ...draft.value, blocking_enabled: value })
|
||||
}
|
||||
function confirmBlocking() {
|
||||
showBlockingConfirmation.value = false
|
||||
if (draft.value) replaceDraft({ ...draft.value, blocking_enabled: true })
|
||||
}
|
||||
function resetDraft() {
|
||||
if (serverConfig.value) draft.value = cloneData(serverConfig.value)
|
||||
}
|
||||
async function saveConfig() {
|
||||
if (!draft.value || !dirty.value) return
|
||||
loading.saving = true
|
||||
try {
|
||||
const saved = await promptAuditAPI.updateConfig(buildUpdateRequest(draft.value))
|
||||
serverConfig.value = configToDraft(saved)
|
||||
draft.value = configToDraft(saved)
|
||||
appStore.showSuccess(t('admin.promptAudit.messages.saved'))
|
||||
await loadRuntime()
|
||||
} catch (error) {
|
||||
const code = extractApiErrorCode(error)
|
||||
appStore.showError(errorMessage(error, code === 'prompt_audit_config_conflict' ? 'admin.promptAudit.errors.prompt_audit_config_conflict' : 'admin.promptAudit.errors.saveConfig'))
|
||||
} finally {
|
||||
loading.saving = false
|
||||
}
|
||||
}
|
||||
async function runProbe(endpoint: PromptAuditEndpointDraft) {
|
||||
if (probingIds.value.includes(endpoint.id)) return
|
||||
probingIds.value = [...probingIds.value, endpoint.id]
|
||||
try {
|
||||
const result = await promptAuditAPI.probeEndpoint(endpoint)
|
||||
probeResults[endpoint.id] = result
|
||||
if (result.ok) appStore.showSuccess(t('admin.promptAudit.messages.probeSucceeded'))
|
||||
else appStore.showError(`${result.error_code || result.status}: ${result.message}`)
|
||||
} catch (error) {
|
||||
appStore.showError(errorMessage(error, 'admin.promptAudit.errors.probe'))
|
||||
} finally {
|
||||
probingIds.value = probingIds.value.filter((id) => id !== endpoint.id)
|
||||
}
|
||||
}
|
||||
|
||||
function handleFiltersChanged(value: PromptEventFilters) {
|
||||
filters.value = cloneData(value)
|
||||
deletePreview.value = null
|
||||
}
|
||||
function applyEventFilters(value: PromptEventFilters) {
|
||||
filters.value = cloneData(value)
|
||||
appliedFilters.value = cloneData(value)
|
||||
events.page = 1
|
||||
deletePreview.value = null
|
||||
void loadEvents()
|
||||
}
|
||||
function changePage(value: number) { events.page = value; void loadEvents() }
|
||||
function changePageSize(value: number) { events.page_size = value; events.page = 1; void loadEvents() }
|
||||
async function openEvent(id: number) {
|
||||
showEventDetail.value = true
|
||||
loading.detail = true
|
||||
activeEvent.value = null
|
||||
try { activeEvent.value = await promptAuditAPI.getEvent(id) }
|
||||
catch (error) { appStore.showError(errorMessage(error, 'admin.promptAudit.errors.loadDetail')); showEventDetail.value = false }
|
||||
finally { loading.detail = false }
|
||||
}
|
||||
function closeEventDetail() { showEventDetail.value = false; activeEvent.value = null }
|
||||
function requestSingleDelete(id: number) { deleteRequest.mode = 'single'; deleteRequest.ids = [id] }
|
||||
function requestBatchDelete() { if (selectedEventIds.value.length) { deleteRequest.mode = 'batch'; deleteRequest.ids = [...selectedEventIds.value] } }
|
||||
function clearDeleteRequest() { deleteRequest.mode = ''; deleteRequest.ids = [] }
|
||||
async function confirmIDDelete() {
|
||||
const mode = deleteRequest.mode
|
||||
const ids = [...deleteRequest.ids]
|
||||
clearDeleteRequest()
|
||||
if (!mode || ids.length === 0) return
|
||||
loading.deleting = true
|
||||
try {
|
||||
const result = mode === 'single' ? await promptAuditAPI.deleteEvent(ids[0]) : await promptAuditAPI.batchDeleteEvents(ids)
|
||||
appStore.showSuccess(t('admin.promptAudit.messages.deleted', { count: result.deleted_events }))
|
||||
await Promise.allSettled([loadEvents(), loadRuntime()])
|
||||
} catch (error) { appStore.showError(errorMessage(error, 'admin.promptAudit.errors.delete')) }
|
||||
finally { loading.deleting = false }
|
||||
}
|
||||
async function requestFilterDeletePreview() {
|
||||
loading.deleting = true
|
||||
try { deletePreview.value = await promptAuditAPI.previewDelete(filters.value) }
|
||||
catch (error) { appStore.showError(errorMessage(error, 'admin.promptAudit.errors.previewDelete')) }
|
||||
finally { loading.deleting = false }
|
||||
}
|
||||
async function confirmFilterDelete() {
|
||||
if (!deletePreview.value) return
|
||||
const preview = deletePreview.value
|
||||
loading.deleting = true
|
||||
try {
|
||||
const result = await promptAuditAPI.deleteEventsByFilter(filters.value, preview)
|
||||
deletePreview.value = null
|
||||
appStore.showSuccess(t('admin.promptAudit.messages.deleted', { count: result.deleted_events }))
|
||||
await Promise.allSettled([loadEvents(), loadRuntime()])
|
||||
} catch (error) {
|
||||
deletePreview.value = null
|
||||
appStore.showError(errorMessage(error, 'admin.promptAudit.errors.deleteConfirmation'))
|
||||
} finally { loading.deleting = false }
|
||||
}
|
||||
function formatDate(value: string): string {
|
||||
return new Intl.DateTimeFormat(locale.value, { dateStyle: 'medium', timeStyle: 'medium' }).format(new Date(value))
|
||||
}
|
||||
|
||||
onMounted(loadInitial)
|
||||
</script>
|
||||
@@ -0,0 +1,161 @@
|
||||
import { beforeEach, describe, expect, it, vi } from 'vitest'
|
||||
import { defineComponent } from 'vue'
|
||||
import { flushPromises, mount } from '@vue/test-utils'
|
||||
import type { PromptAuditConfig, PromptAuditRuntime } from '../types'
|
||||
import { SCANNER_CATALOG } from '../viewModel'
|
||||
import PromptAuditView from '../PromptAuditView.vue'
|
||||
|
||||
const mocks = vi.hoisted(() => ({
|
||||
getConfig: vi.fn(), updateConfig: vi.fn(), probeEndpoint: vi.fn(), getRuntime: vi.fn(), listEvents: vi.fn(),
|
||||
getEvent: vi.fn(), deleteEvent: vi.fn(), batchDeleteEvents: vi.fn(), previewDelete: vi.fn(), deleteEventsByFilter: vi.fn(), listGroups: vi.fn(),
|
||||
showSuccess: vi.fn(), showError: vi.fn(),
|
||||
}))
|
||||
|
||||
vi.mock('../api', () => ({ default: mocks }))
|
||||
vi.mock('@/stores/app', () => ({ useAppStore: () => ({ showSuccess: mocks.showSuccess, showError: mocks.showError }) }))
|
||||
vi.mock('vue-i18n', async () => {
|
||||
const actual = await vi.importActual<typeof import('vue-i18n')>('vue-i18n')
|
||||
return { ...actual, useI18n: () => ({ locale: { value: 'en' }, t: (key: string, params?: Record<string, unknown>) => key.replace(/\{(\w+)\}/g, (_, token) => String(params?.[token] ?? `{${token}}`)) }) }
|
||||
})
|
||||
|
||||
const baseConfig = (): PromptAuditConfig => ({
|
||||
enabled: true, blocking_enabled: false, store_pass_events: false, effective_mode: 'async_audit', strategy: 'priority',
|
||||
worker_count: 4, queue_capacity: 100, scanners: SCANNER_CATALOG.map((item) => item.id), all_groups: true, group_ids: [],
|
||||
endpoints: [{ id: 'guard-1', name: 'Guard One', protocol: 'openai_compatible', base_url: 'http://127.0.0.1:8000', model: 'guard-model', timeout_ms: 3000, input_limit: 4000, enabled: true, has_token: true, token_status: 'configured' }],
|
||||
config_version: 7, updated_at: '2026-07-16T00:00:00Z', updated_by: 1, change_summary: '{}',
|
||||
})
|
||||
const runtime = (): PromptAuditRuntime => ({
|
||||
process_status: 'running', effective_mode: 'async_audit', expected_config_version: 7, active_config_version: 7,
|
||||
worker_total: 4, worker_active: 1, queue_capacity: 100,
|
||||
queue: { staging: 0, queued: 0, processing: 1, retry: 0, done: 5, failed: 0, active: 1 },
|
||||
processed_total: 5, failed_total: 0, enqueued_total: 5, dropped_total: 0, database_status: 'ok', redis_status: 'ok', endpoints: {},
|
||||
guard_metrics: { total: 1, allowed: 1, flagged: 0, blocked: 0, unavailable: 0, invalid: 0, timeouts: 0, failovers: 0, bulkhead_full: 0, record_failed: 0 },
|
||||
})
|
||||
|
||||
const AppLayoutStub = { template: '<div><slot /></div>' }
|
||||
const RuntimeStub = defineComponent({ props: ['runtime', 'loading', 'error'], emits: ['refresh'], template: '<div data-test="runtime">{{ error }}</div>' })
|
||||
const EndpointStub = defineComponent({
|
||||
props: ['endpoints', 'probeResults', 'probingIds'], emits: ['update:endpoints', 'probe'],
|
||||
template: '<div data-test="endpoint"><button data-test="inject-secret" @click="$emit(\'update:endpoints\', endpoints.map((e) => ({ ...e, token: \'PROMPT_AUDIT_CANARY_SECRET_DO_NOT_PERSIST\' })))">secret</button><button data-test="probe" @click="$emit(\'probe\', endpoints[0])">probe</button></div>',
|
||||
})
|
||||
const PolicyStub = defineComponent({ props: ['draft', 'groups'], emits: ['update:draft'], template: '<div data-test="policy" />' })
|
||||
const EventsStub = defineComponent({
|
||||
props: ['events', 'filters', 'selectedIds', 'loading', 'error', 'total', 'page', 'pageSize'],
|
||||
emits: ['filters-change', 'search', 'selection', 'page', 'page-size', 'view', 'delete', 'batch-delete', 'preview-delete'],
|
||||
template: '<div data-test="events"><button data-test="preview" @click="$emit(\'preview-delete\')">preview</button><button data-test="change-filter" @click="$emit(\'filters-change\', { ...filters, keyword: \'changed\' })">change</button><button data-test="delete-one" @click="$emit(\'delete\', 5)">delete</button><button data-test="select-batch" @click="$emit(\'selection\', [5, 6])">select</button><button data-test="delete-batch" @click="$emit(\'batch-delete\')">batch</button></div>',
|
||||
})
|
||||
const DetailStub = defineComponent({ props: ['show', 'event', 'loading'], emits: ['close'], template: '<div data-test="detail" />' })
|
||||
const DialogStub = defineComponent({ props: ['show', 'title'], emits: ['close'], template: '<div v-if="show" data-test="dialog"><slot /><slot name="footer" /></div>' })
|
||||
const ConfirmStub = defineComponent({ props: ['show', 'title', 'message'], emits: ['confirm', 'cancel'], template: '<div v-if="show" data-test="confirm"><button data-test="confirm-action" @click="$emit(\'confirm\')">confirm</button></div>' })
|
||||
|
||||
function mountView() {
|
||||
return mount(PromptAuditView, {
|
||||
global: { stubs: { AppLayout: AppLayoutStub, RuntimeOverview: RuntimeStub, EndpointPool: EndpointStub, PolicyPanel: PolicyStub, EventWorkspace: EventsStub, EventDetailDialog: DetailStub, BaseDialog: DialogStub, ConfirmDialog: ConfirmStub } },
|
||||
})
|
||||
}
|
||||
|
||||
describe('PromptAuditView', () => {
|
||||
beforeEach(() => {
|
||||
Object.values(mocks).forEach((mock) => mock.mockReset())
|
||||
mocks.getConfig.mockResolvedValue(baseConfig())
|
||||
mocks.getRuntime.mockResolvedValue(runtime())
|
||||
mocks.listGroups.mockResolvedValue([])
|
||||
mocks.listEvents.mockResolvedValue({ items: [], total: 0, page: 1, page_size: 20, pages: 0 })
|
||||
mocks.updateConfig.mockImplementation(async () => ({ ...baseConfig(), config_version: 8 }))
|
||||
mocks.probeEndpoint.mockResolvedValue({ ok: true, status: 'healthy', message: 'ok', latency_ms: 2, http_status: 200, retryable: false, checked_at: '2026-07-16T00:00:00Z', token_applied: true })
|
||||
mocks.previewDelete.mockResolvedValue({ matched_count: 2, filter_summary: {}, snapshot_max_id: 10, filter_hash: 'a'.repeat(64), confirmation_token: 'opaque-confirmation', expires_at: '2026-07-16T00:05:00Z' })
|
||||
mocks.deleteEventsByFilter.mockResolvedValue({ deleted_events: 2, deleted_jobs: 2 })
|
||||
mocks.deleteEvent.mockResolvedValue({ deleted_events: 1, deleted_jobs: 1 })
|
||||
mocks.batchDeleteEvents.mockResolvedValue({ deleted_events: 2, deleted_jobs: 2 })
|
||||
})
|
||||
|
||||
it('starts config, runtime, groups, and events loads independently', async () => {
|
||||
mocks.getRuntime.mockRejectedValue(new Error('runtime offline'))
|
||||
const wrapper = mountView()
|
||||
expect(mocks.getConfig).toHaveBeenCalledOnce()
|
||||
expect(mocks.getRuntime).toHaveBeenCalledOnce()
|
||||
expect(mocks.listGroups).toHaveBeenCalledOnce()
|
||||
expect(mocks.listEvents).toHaveBeenCalledOnce()
|
||||
await flushPromises()
|
||||
expect(wrapper.get('[data-test="runtime"]').text()).toContain('runtime offline')
|
||||
expect(wrapper.find('[data-test="endpoint"]').exists()).toBe(true)
|
||||
expect(wrapper.find('[data-test="events"]').exists()).toBe(true)
|
||||
})
|
||||
|
||||
it('requires confirmation for blocking and disables it when audit is turned off', async () => {
|
||||
const wrapper = mountView()
|
||||
await flushPromises()
|
||||
await wrapper.get('[data-test="blocking-toggle"]').trigger('click')
|
||||
expect(wrapper.find('[data-test="confirm"]').exists()).toBe(true)
|
||||
await wrapper.get('[data-test="confirm-action"]').trigger('click')
|
||||
expect(wrapper.get('[data-test="blocking-toggle"]').attributes('aria-checked')).toBe('true')
|
||||
await wrapper.get('[data-test="enabled-toggle"]').trigger('click')
|
||||
expect(wrapper.get('[data-test="enabled-toggle"]').attributes('aria-checked')).toBe('false')
|
||||
expect(wrapper.get('[data-test="blocking-toggle"]').attributes('aria-checked')).toBe('false')
|
||||
expect(wrapper.get('[data-test="blocking-toggle"]').attributes()).toHaveProperty('disabled')
|
||||
})
|
||||
|
||||
it('clears plaintext token state after a successful save', async () => {
|
||||
const wrapper = mountView()
|
||||
await flushPromises()
|
||||
await wrapper.get('[data-test="inject-secret"]').trigger('click')
|
||||
expect(wrapper.text()).toContain('admin.promptAudit.saveBar.dirty')
|
||||
await wrapper.get('[data-test="save-config"]').trigger('click')
|
||||
await flushPromises()
|
||||
expect(mocks.updateConfig).toHaveBeenCalledWith(expect.objectContaining({ endpoints: [expect.objectContaining({ token: 'PROMPT_AUDIT_CANARY_SECRET_DO_NOT_PERSIST' })] }))
|
||||
const endpointProps = wrapper.getComponent(EndpointStub).props('endpoints') as Array<{ token: string }>
|
||||
expect(endpointProps[0].token).toBe('')
|
||||
expect(wrapper.html()).not.toContain('PROMPT_AUDIT_CANARY_SECRET_DO_NOT_PERSIST')
|
||||
})
|
||||
|
||||
it('reports real probe progress/results and invalidates filter confirmation when filters change', async () => {
|
||||
const wrapper = mountView()
|
||||
await flushPromises()
|
||||
await wrapper.get('[data-test="probe"]').trigger('click')
|
||||
await flushPromises()
|
||||
expect(mocks.probeEndpoint).toHaveBeenCalledOnce()
|
||||
expect((wrapper.getComponent(EndpointStub).props('probeResults') as Record<string, unknown>)).toHaveProperty('guard-1')
|
||||
|
||||
await wrapper.get('[data-test="preview"]').trigger('click')
|
||||
await flushPromises()
|
||||
expect(wrapper.find('[data-test="dialog"]').exists()).toBe(true)
|
||||
await wrapper.get('[data-test="change-filter"]').trigger('click')
|
||||
await flushPromises()
|
||||
expect(wrapper.find('[data-test="dialog"]').exists()).toBe(false)
|
||||
})
|
||||
|
||||
it('uses native labeled switches and a responsive fixed save surface', async () => {
|
||||
const wrapper = mountView()
|
||||
await flushPromises()
|
||||
const switches = wrapper.findAll('[role="switch"]')
|
||||
expect(switches).toHaveLength(3)
|
||||
expect(switches.every((item) => Boolean(item.attributes('aria-label')))).toBe(true)
|
||||
expect(wrapper.html()).toContain('fixed inset-x-0 bottom-0')
|
||||
expect(wrapper.html()).toContain('flex-wrap')
|
||||
})
|
||||
|
||||
it('executes single, selected-batch, and preview-confirmed filter deletion flows', async () => {
|
||||
const wrapper = mountView()
|
||||
await flushPromises()
|
||||
|
||||
await wrapper.get('[data-test="delete-one"]').trigger('click')
|
||||
await wrapper.get('[data-test="confirm-action"]').trigger('click')
|
||||
await flushPromises()
|
||||
expect(mocks.deleteEvent).toHaveBeenCalledWith(5)
|
||||
|
||||
await wrapper.get('[data-test="select-batch"]').trigger('click')
|
||||
await wrapper.get('[data-test="delete-batch"]').trigger('click')
|
||||
await wrapper.get('[data-test="confirm-action"]').trigger('click')
|
||||
await flushPromises()
|
||||
expect(mocks.batchDeleteEvents).toHaveBeenCalledWith([5, 6])
|
||||
|
||||
await wrapper.get('[data-test="preview"]').trigger('click')
|
||||
await flushPromises()
|
||||
await wrapper.get('[data-test="confirm-filter-delete"]').trigger('click')
|
||||
await flushPromises()
|
||||
expect(mocks.deleteEventsByFilter).toHaveBeenCalledWith(expect.any(Object), expect.objectContaining({
|
||||
snapshot_max_id: 10,
|
||||
confirmation_token: 'opaque-confirmation',
|
||||
}))
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,44 @@
|
||||
import { beforeEach, describe, expect, it, vi } from 'vitest'
|
||||
import { emptyEventFilters } from '../viewModel'
|
||||
|
||||
const client = vi.hoisted(() => ({ get: vi.fn(), put: vi.fn(), post: vi.fn(), delete: vi.fn() }))
|
||||
vi.mock('@/api/client', () => ({ apiClient: client }))
|
||||
|
||||
import promptAuditAPI from '../api'
|
||||
|
||||
describe('Prompt Audit API', () => {
|
||||
beforeEach(() => Object.values(client).forEach((mock) => mock.mockReset()))
|
||||
|
||||
it('uses the independent admin route namespace', async () => {
|
||||
client.get.mockResolvedValue({ data: { config_version: 1 } })
|
||||
await promptAuditAPI.getConfig()
|
||||
expect(client.get).toHaveBeenCalledWith('/admin/prompt-audit/config')
|
||||
|
||||
client.get.mockResolvedValue({ data: { process_status: 'running' } })
|
||||
await promptAuditAPI.getRuntime()
|
||||
expect(client.get).toHaveBeenCalledWith('/admin/prompt-audit/runtime')
|
||||
})
|
||||
|
||||
it('sends a temporary probe token only in the request and never invents response credentials', async () => {
|
||||
client.post.mockResolvedValue({ data: { ok: true, token_applied: true } })
|
||||
const result = await promptAuditAPI.probeEndpoint({
|
||||
id: 'guard-1', name: 'Guard', protocol: 'openai_compatible', base_url: 'http://127.0.0.1:8000', model: 'guard',
|
||||
token: 'api-canary-secret', clear_token: false, timeout_ms: 1000, input_limit: 1000, enabled: true, has_token: false, token_status: 'missing',
|
||||
})
|
||||
expect(client.post).toHaveBeenCalledWith('/admin/prompt-audit/endpoints/probe', expect.objectContaining({ endpoint: expect.objectContaining({ token: 'api-canary-secret' }) }))
|
||||
expect(JSON.stringify(result)).not.toContain('api-canary-secret')
|
||||
})
|
||||
|
||||
it('passes a server preview token through the confirmed filter-delete contract', async () => {
|
||||
client.post.mockResolvedValue({ data: { deleted_events: 2, deleted_jobs: 2 } })
|
||||
const filters = emptyEventFilters()
|
||||
filters.start_at = '2026-07-15T00:00'
|
||||
filters.end_at = '2026-07-16T00:00'
|
||||
await promptAuditAPI.deleteEventsByFilter(filters, {
|
||||
matched_count: 2, filter_summary: {}, snapshot_max_id: 10, filter_hash: 'a'.repeat(64), confirmation_token: 'opaque-token', expires_at: '2026-07-16T00:05:00Z',
|
||||
})
|
||||
expect(client.post).toHaveBeenCalledWith('/admin/prompt-audit/events/delete-by-filter', expect.objectContaining({
|
||||
snapshot_max_id: 10, filter_hash: 'a'.repeat(64), confirmation_token: 'opaque-token', confirm: true,
|
||||
}))
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,89 @@
|
||||
import { beforeEach, describe, expect, it, vi } from 'vitest'
|
||||
import { defineComponent } from 'vue'
|
||||
import { mount } from '@vue/test-utils'
|
||||
import EndpointPool from '../components/EndpointPool.vue'
|
||||
import PolicyPanel from '../components/PolicyPanel.vue'
|
||||
import EventWorkspace from '../components/EventWorkspace.vue'
|
||||
import type { PromptAuditDraft, PromptAuditEndpointDraft, PromptAuditEvent } from '../types'
|
||||
import { emptyEventFilters, SCANNER_CATALOG } from '../viewModel'
|
||||
|
||||
vi.mock('vue-i18n', async () => {
|
||||
const actual = await vi.importActual<typeof import('vue-i18n')>('vue-i18n')
|
||||
return { ...actual, useI18n: () => ({ locale: { value: 'en' }, t: (key: string, params?: Record<string, unknown>) => key.replace(/\{(\w+)\}/g, (_, token) => String(params?.[token] ?? `{${token}}`)) }) }
|
||||
})
|
||||
|
||||
const DialogStub = defineComponent({ props: ['show', 'title'], emits: ['close'], template: '<div v-if="show" data-test="dialog"><slot /><slot name="footer" /></div>' })
|
||||
const PaginationStub = defineComponent({ props: ['total', 'page', 'pageSize'], emits: ['update:page', 'update:pageSize'], template: '<div data-test="pagination" />' })
|
||||
|
||||
const endpoint = (): PromptAuditEndpointDraft => ({
|
||||
id: 'guard-1', name: 'Guard One', protocol: 'openai_compatible', base_url: 'http://127.0.0.1:8000',
|
||||
model: 'guard-model', timeout_ms: 3000, input_limit: 4000, enabled: true,
|
||||
has_token: true, token_status: 'configured', token: '', clear_token: false,
|
||||
})
|
||||
|
||||
describe('Prompt Audit components', () => {
|
||||
beforeEach(() => vi.restoreAllMocks())
|
||||
|
||||
it('edits a saved endpoint with blank-secret keep, explicit clear, replacement, and probe actions', async () => {
|
||||
const wrapper = mount(EndpointPool, {
|
||||
props: { endpoints: [endpoint()], probeResults: {}, probingIds: [] },
|
||||
global: { stubs: { BaseDialog: DialogStub } },
|
||||
})
|
||||
expect(wrapper.text()).toContain('admin.promptAudit.pool.configured')
|
||||
const edit = wrapper.findAll('button').find((button) => button.text().includes('common.edit'))
|
||||
expect(edit).toBeTruthy()
|
||||
await edit!.trigger('click')
|
||||
const token = wrapper.get<HTMLInputElement>('[aria-label="admin.promptAudit.pool.apiKey"]')
|
||||
expect(token.element.value).toBe('')
|
||||
expect(token.attributes('placeholder')).toContain('admin.promptAudit.pool.keepSecret')
|
||||
|
||||
await wrapper.get<HTMLInputElement>('[aria-label="admin.promptAudit.pool.clearSecret"]').setValue(true)
|
||||
await token.setValue('replacement-canary')
|
||||
await wrapper.get('[data-test="save-endpoint"]').trigger('click')
|
||||
const updated = wrapper.emitted('update:endpoints')?.at(-1)?.[0] as PromptAuditEndpointDraft[]
|
||||
expect(updated[0]).toMatchObject({ token: 'replacement-canary', clear_token: false })
|
||||
|
||||
const probe = wrapper.findAll('button').find((button) => button.text().includes('admin.promptAudit.pool.probe'))
|
||||
await probe!.trigger('click')
|
||||
expect(wrapper.emitted('probe')?.[0]?.[0]).toMatchObject({ id: 'guard-1' })
|
||||
})
|
||||
|
||||
it('supports group search, stale configured groups, nine scanners, and bounded worker inputs', async () => {
|
||||
const draft: PromptAuditDraft = {
|
||||
enabled: true, blocking_enabled: false, store_pass_events: false, effective_mode: 'async_audit', strategy: 'priority',
|
||||
worker_count: 4, queue_capacity: 100, scanners: SCANNER_CATALOG.map((item) => item.id), all_groups: false, group_ids: [1, 99],
|
||||
endpoints: [endpoint()], config_version: 1, updated_at: '', updated_by: 0, change_summary: '',
|
||||
}
|
||||
const wrapper = mount(PolicyPanel, {
|
||||
props: { draft, groups: [{ id: 1, name: 'Alpha', platform: 'openai', status: 'active' }, { id: 2, name: 'Beta', platform: 'claude', status: 'inactive' }] },
|
||||
})
|
||||
expect(wrapper.text()).toContain('99')
|
||||
expect(wrapper.findAll('input[type="checkbox"]').filter((input) => SCANNER_CATALOG.some((scanner) => input.attributes('aria-label') === scanner.label))).toHaveLength(9)
|
||||
await wrapper.get('[aria-label="admin.promptAudit.policy.searchGroups"]').setValue('Beta')
|
||||
expect(wrapper.text()).toContain('Beta')
|
||||
expect(wrapper.text()).not.toContain('Alpha')
|
||||
await wrapper.get('[aria-label="admin.promptAudit.policy.workerCount"]').setValue('6')
|
||||
const emitted = wrapper.emitted('update:draft')?.at(-1)?.[0] as PromptAuditDraft
|
||||
expect(emitted.worker_count).toBe(6)
|
||||
})
|
||||
|
||||
it('keeps identity fields separate, supports selection, and gates filter deletion on a time range', async () => {
|
||||
const event: PromptAuditEvent = {
|
||||
id: 1, job_id: 1, decision: 'critical', risk_level: 'critical', action: 'Block', categories: ['pii'], matched_scanners: ['pii'], scanner_scores: { pii: 1 }, scanner_evidence: { pii: 'redacted' }, scanner_backend: 'qwen3guard-openai', scanner_version: '1', guard_endpoint_id: 'guard-1', policy_id: 'priority', policy_version: 1, config_version: 1, chunk_total: 1, latency_ms: 10, issue_summaries: [], created_at: '2026-07-16T00:00:00Z',
|
||||
snapshot: { request_id: 'req-1', user_id: 1, username: 'alice', user_email: 'alice@example.test', api_key_id: 2, api_key_name: 'alice-key', group_id: 3, group_name: 'Alpha', provider: 'openai', endpoint: '/v1/chat/completions', protocol: 'openai_chat', model: 'gpt-test', prompt_hash: 'a'.repeat(64), redacted_preview: 'redacted preview', prompt_length: 10, message_count: 1, stage: 'http' },
|
||||
}
|
||||
const wrapper = mount(EventWorkspace, {
|
||||
props: { events: [event], total: 1, page: 1, pageSize: 20, filters: emptyEventFilters(), selectedIds: [], loading: false, error: '' },
|
||||
global: { stubs: { Pagination: PaginationStub } },
|
||||
})
|
||||
expect(wrapper.text()).toContain('alice')
|
||||
expect(wrapper.text()).toContain('alice@example.test')
|
||||
expect(wrapper.text()).toContain('alice-key')
|
||||
expect(wrapper.get('[data-test="filter-delete"]').attributes()).toHaveProperty('disabled')
|
||||
await wrapper.get('[aria-label="admin.promptAudit.events.startAt"]').setValue('2026-07-15T00:00')
|
||||
await wrapper.get('[aria-label="admin.promptAudit.events.endAt"]').setValue('2026-07-16T00:00')
|
||||
expect(wrapper.get('[data-test="filter-delete"]').attributes()).not.toHaveProperty('disabled')
|
||||
await wrapper.get('[aria-label="admin.promptAudit.events.selectEvent"]').setValue(true)
|
||||
expect(wrapper.emitted('selection')?.at(-1)?.[0]).toEqual([1])
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,40 @@
|
||||
import { readFileSync } from 'node:fs'
|
||||
import { dirname, resolve } from 'node:path'
|
||||
import { fileURLToPath } from 'node:url'
|
||||
import { describe, expect, it } from 'vitest'
|
||||
import en from '@/i18n/locales/en'
|
||||
import zh from '@/i18n/locales/zh'
|
||||
|
||||
const here = dirname(fileURLToPath(import.meta.url))
|
||||
const read = (path: string) => readFileSync(resolve(here, path), 'utf8')
|
||||
|
||||
describe('Prompt Audit integration surface', () => {
|
||||
it('registers an admin and risk-control guarded route', () => {
|
||||
const router = read('../../../router/index.ts')
|
||||
expect(router).toContain("path: '/admin/prompt-audit'")
|
||||
const route = router.slice(router.indexOf("path: '/admin/prompt-audit'"), router.indexOf("path: '/admin/usage'"))
|
||||
expect(route).toContain('requiresAuth: true')
|
||||
expect(route).toContain('requiresAdmin: true')
|
||||
expect(route).toContain('requiresRiskControl: true')
|
||||
})
|
||||
|
||||
it('keeps the legacy content moderation route and adds both pages under an expand-only security group', () => {
|
||||
const sidebar = read('../../../components/layout/AppSidebar.vue')
|
||||
const group = sidebar.slice(sidebar.indexOf("path: '/admin/security-audit'"), sidebar.indexOf("path: '/admin/redeem'"))
|
||||
expect(group).toContain('expandOnly: true')
|
||||
expect(group).toContain("path: '/admin/risk-control'")
|
||||
expect(group).toContain("path: '/admin/prompt-audit'")
|
||||
})
|
||||
|
||||
it('keeps Prompt Audit locale trees symmetric and all operational controls named', () => {
|
||||
expect(Object.keys(zh.admin.promptAudit)).toEqual(Object.keys(en.admin.promptAudit))
|
||||
expect(zh.nav.securityAudit).toBeTruthy()
|
||||
expect(en.nav.securityAudit).toBeTruthy()
|
||||
const endpoint = read('../components/EndpointPool.vue')
|
||||
const events = read('../components/EventWorkspace.vue')
|
||||
expect(endpoint).toContain('aria-label')
|
||||
expect(events).toContain('aria-label')
|
||||
expect(events).toContain('overflow-x-auto')
|
||||
expect(events).toContain('sm:grid-cols-2')
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,80 @@
|
||||
import { describe, expect, it } from 'vitest'
|
||||
import type { PromptAuditConfig } from '../types'
|
||||
import {
|
||||
buildUpdateRequest,
|
||||
configToDraft,
|
||||
draftFingerprint,
|
||||
emptyEventFilters,
|
||||
eventFilterPayload,
|
||||
hasExplicitDeleteRange,
|
||||
SCANNER_CATALOG,
|
||||
} from '../viewModel'
|
||||
|
||||
const config = (): PromptAuditConfig => ({
|
||||
enabled: true,
|
||||
blocking_enabled: false,
|
||||
store_pass_events: false,
|
||||
effective_mode: 'async_audit',
|
||||
strategy: 'priority',
|
||||
worker_count: 4,
|
||||
queue_capacity: 100,
|
||||
scanners: SCANNER_CATALOG.map((item) => item.id),
|
||||
all_groups: true,
|
||||
group_ids: [],
|
||||
endpoints: [{
|
||||
id: 'guard-1', name: 'Guard One', protocol: 'openai_compatible', base_url: 'http://127.0.0.1:8000',
|
||||
model: 'sileader/qwen3guard:0.6b', timeout_ms: 3000, input_limit: 4000, enabled: true,
|
||||
has_token: true, token_status: 'configured',
|
||||
}],
|
||||
config_version: 7,
|
||||
updated_at: '2026-07-16T00:00:00Z',
|
||||
updated_by: 1,
|
||||
change_summary: '{}',
|
||||
})
|
||||
|
||||
describe('Prompt Audit view model', () => {
|
||||
it('normalizes legacy null collections from the public config', () => {
|
||||
const legacy = { ...config(), group_ids: null, scanners: null, endpoints: null } as unknown as PromptAuditConfig
|
||||
expect(configToDraft(legacy)).toMatchObject({ group_ids: [], scanners: [], endpoints: [] })
|
||||
})
|
||||
|
||||
it('models all nine official input scanners', () => {
|
||||
expect(SCANNER_CATALOG).toHaveLength(9)
|
||||
expect(SCANNER_CATALOG.map((item) => item.id)).toContain('suicide_and_self_harm')
|
||||
})
|
||||
|
||||
it('keeps, replaces, or explicitly clears a saved token without copying plaintext from the server', () => {
|
||||
const draft = configToDraft(config())
|
||||
expect(draft.endpoints[0].token).toBe('')
|
||||
expect(buildUpdateRequest(draft).endpoints[0]).toMatchObject({ token: undefined, clear_token: false })
|
||||
|
||||
draft.endpoints[0].token = 'temporary-canary-token'
|
||||
expect(buildUpdateRequest(draft).endpoints[0]).toMatchObject({ token: 'temporary-canary-token', clear_token: false })
|
||||
|
||||
draft.endpoints[0].token = ''
|
||||
draft.endpoints[0].clear_token = true
|
||||
expect(buildUpdateRequest(draft).endpoints[0]).toMatchObject({ token: undefined, clear_token: true })
|
||||
})
|
||||
|
||||
it('tracks dirty state from the full normalized save payload', () => {
|
||||
const original = configToDraft(config())
|
||||
const changed = configToDraft(config())
|
||||
expect(draftFingerprint(changed)).toBe(draftFingerprint(original))
|
||||
changed.queue_capacity += 1
|
||||
expect(draftFingerprint(changed)).not.toBe(draftFingerprint(original))
|
||||
})
|
||||
|
||||
it('requires a valid explicit range and sends canonical ISO timestamps for filter deletion', () => {
|
||||
const filters = emptyEventFilters()
|
||||
expect(hasExplicitDeleteRange(filters)).toBe(false)
|
||||
filters.start_at = '2026-07-15T10:00'
|
||||
filters.end_at = '2026-07-16T10:00'
|
||||
filters.group_id = '9'
|
||||
expect(hasExplicitDeleteRange(filters)).toBe(true)
|
||||
expect(eventFilterPayload(filters)).toMatchObject({
|
||||
group_id: 9,
|
||||
start_at: new Date(filters.start_at).toISOString(),
|
||||
end_at: new Date(filters.end_at).toISOString(),
|
||||
})
|
||||
})
|
||||
})
|
||||
@@ -0,0 +1,120 @@
|
||||
import { apiClient } from '@/api/client'
|
||||
import type {
|
||||
PromptAuditConfig,
|
||||
PromptAuditEvent,
|
||||
PromptAuditGroup,
|
||||
PromptAuditRuntime,
|
||||
PromptAuditUpdateRequest,
|
||||
PromptDeletePreview,
|
||||
PromptDeleteResult,
|
||||
PromptEventFilters,
|
||||
PromptEventPage,
|
||||
PromptProbeResult,
|
||||
PromptAuditEndpointDraft,
|
||||
} from './types'
|
||||
import { eventFilterPayload, eventQueryParams } from './viewModel'
|
||||
|
||||
const basePath = '/admin/prompt-audit'
|
||||
|
||||
export async function getConfig(): Promise<PromptAuditConfig> {
|
||||
const { data } = await apiClient.get<PromptAuditConfig>(`${basePath}/config`)
|
||||
return data
|
||||
}
|
||||
|
||||
export async function updateConfig(payload: PromptAuditUpdateRequest): Promise<PromptAuditConfig> {
|
||||
const { data } = await apiClient.put<PromptAuditConfig>(`${basePath}/config`, payload)
|
||||
return data
|
||||
}
|
||||
|
||||
export async function probeEndpoint(endpoint: PromptAuditEndpointDraft): Promise<PromptProbeResult> {
|
||||
const { data } = await apiClient.post<PromptProbeResult>(`${basePath}/endpoints/probe`, {
|
||||
endpoint: {
|
||||
id: endpoint.id,
|
||||
name: endpoint.name,
|
||||
protocol: 'openai_compatible',
|
||||
base_url: endpoint.base_url,
|
||||
model: endpoint.model,
|
||||
token: endpoint.token || undefined,
|
||||
timeout_ms: endpoint.timeout_ms,
|
||||
input_limit: endpoint.input_limit,
|
||||
enabled: endpoint.enabled,
|
||||
},
|
||||
})
|
||||
return data
|
||||
}
|
||||
|
||||
export async function getRuntime(): Promise<PromptAuditRuntime> {
|
||||
const { data } = await apiClient.get<PromptAuditRuntime>(`${basePath}/runtime`)
|
||||
return data
|
||||
}
|
||||
|
||||
export async function listEvents(
|
||||
filters: PromptEventFilters,
|
||||
page: number,
|
||||
pageSize: number,
|
||||
): Promise<PromptEventPage> {
|
||||
const { data } = await apiClient.get<PromptEventPage>(`${basePath}/events`, {
|
||||
params: { page, page_size: pageSize, ...eventQueryParams(filters) },
|
||||
})
|
||||
return data
|
||||
}
|
||||
|
||||
export async function getEvent(id: number): Promise<PromptAuditEvent> {
|
||||
const { data } = await apiClient.get<PromptAuditEvent>(`${basePath}/events/${id}`)
|
||||
return data
|
||||
}
|
||||
|
||||
export async function deleteEvent(id: number): Promise<PromptDeleteResult> {
|
||||
const { data } = await apiClient.delete<PromptDeleteResult>(`${basePath}/events/${id}`)
|
||||
return data
|
||||
}
|
||||
|
||||
export async function batchDeleteEvents(ids: number[]): Promise<PromptDeleteResult> {
|
||||
const { data } = await apiClient.post<PromptDeleteResult>(`${basePath}/events/batch-delete`, { ids })
|
||||
return data
|
||||
}
|
||||
|
||||
export async function previewDelete(filters: PromptEventFilters): Promise<PromptDeletePreview> {
|
||||
const { data } = await apiClient.post<PromptDeletePreview>(
|
||||
`${basePath}/events/delete-preview`,
|
||||
eventFilterPayload(filters),
|
||||
)
|
||||
return data
|
||||
}
|
||||
|
||||
export async function deleteEventsByFilter(
|
||||
filters: PromptEventFilters,
|
||||
preview: PromptDeletePreview,
|
||||
): Promise<PromptDeleteResult> {
|
||||
const { data } = await apiClient.post<PromptDeleteResult>(`${basePath}/events/delete-by-filter`, {
|
||||
filter: eventFilterPayload(filters),
|
||||
snapshot_max_id: preview.snapshot_max_id,
|
||||
filter_hash: preview.filter_hash,
|
||||
confirmation_token: preview.confirmation_token,
|
||||
confirm: true,
|
||||
})
|
||||
return data
|
||||
}
|
||||
|
||||
export async function listGroups(): Promise<PromptAuditGroup[]> {
|
||||
const { data } = await apiClient.get<PromptAuditGroup[]>('/admin/groups/all', {
|
||||
params: { include_inactive: true },
|
||||
})
|
||||
return data
|
||||
}
|
||||
|
||||
export const promptAuditAPI = {
|
||||
getConfig,
|
||||
updateConfig,
|
||||
probeEndpoint,
|
||||
getRuntime,
|
||||
listEvents,
|
||||
getEvent,
|
||||
deleteEvent,
|
||||
batchDeleteEvents,
|
||||
previewDelete,
|
||||
deleteEventsByFilter,
|
||||
listGroups,
|
||||
}
|
||||
|
||||
export default promptAuditAPI
|
||||
@@ -0,0 +1,172 @@
|
||||
<template>
|
||||
<section aria-labelledby="prompt-pool-title" class="border-b border-gray-200 py-6 dark:border-dark-700">
|
||||
<div class="flex flex-wrap items-start justify-between gap-3">
|
||||
<div>
|
||||
<h2 id="prompt-pool-title" class="text-base font-semibold text-gray-950 dark:text-white">{{ t('admin.promptAudit.pool.title') }}</h2>
|
||||
<p class="mt-1 text-sm text-gray-500 dark:text-dark-300">{{ t('admin.promptAudit.pool.description') }}</p>
|
||||
</div>
|
||||
<button type="button" class="btn btn-primary btn-sm" data-test="add-endpoint" @click="openCreate">
|
||||
{{ t('admin.promptAudit.pool.add') }}
|
||||
</button>
|
||||
</div>
|
||||
|
||||
<div v-if="endpoints.length === 0" class="mt-5 rounded-lg border border-dashed border-gray-300 px-5 py-8 text-center text-sm text-gray-500 dark:border-dark-600 dark:text-dark-300">
|
||||
{{ t('admin.promptAudit.pool.empty') }}
|
||||
</div>
|
||||
<div v-else class="mt-5 overflow-x-auto">
|
||||
<table class="min-w-full text-left text-sm">
|
||||
<thead class="border-b border-gray-200 text-xs uppercase tracking-wide text-gray-500 dark:border-dark-700 dark:text-dark-400">
|
||||
<tr>
|
||||
<th class="px-3 py-2 font-medium">{{ t('admin.promptAudit.pool.node') }}</th>
|
||||
<th class="px-3 py-2 font-medium">{{ t('admin.promptAudit.pool.model') }}</th>
|
||||
<th class="px-3 py-2 font-medium">{{ t('admin.promptAudit.pool.limits') }}</th>
|
||||
<th class="px-3 py-2 font-medium">{{ t('admin.promptAudit.pool.credential') }}</th>
|
||||
<th class="px-3 py-2 text-right font-medium">{{ t('admin.promptAudit.common.actions') }}</th>
|
||||
</tr>
|
||||
</thead>
|
||||
<tbody class="divide-y divide-gray-100 dark:divide-dark-800">
|
||||
<tr v-for="endpoint in endpoints" :key="endpoint.id" :data-test="`endpoint-${endpoint.id}`">
|
||||
<td class="px-3 py-3">
|
||||
<div class="flex items-center gap-2">
|
||||
<button
|
||||
type="button"
|
||||
role="switch"
|
||||
:aria-checked="endpoint.enabled"
|
||||
:aria-label="t('admin.promptAudit.pool.toggleNode', { name: endpoint.name })"
|
||||
class="relative h-5 w-9 rounded-full transition-colors"
|
||||
:class="endpoint.enabled ? 'bg-primary-600' : 'bg-gray-300 dark:bg-dark-600'"
|
||||
@click="toggleEndpoint(endpoint.id)"
|
||||
>
|
||||
<span class="absolute top-0.5 h-4 w-4 rounded-full bg-white transition-transform" :class="endpoint.enabled ? 'translate-x-4' : 'translate-x-0.5'" />
|
||||
</button>
|
||||
<div class="min-w-0">
|
||||
<p class="font-medium text-gray-900 dark:text-white">{{ endpoint.name }}</p>
|
||||
<p class="max-w-xs truncate text-xs text-gray-500 dark:text-dark-400">{{ endpoint.base_url }}</p>
|
||||
</div>
|
||||
</div>
|
||||
</td>
|
||||
<td class="px-3 py-3 text-gray-700 dark:text-dark-200">{{ endpoint.model }}</td>
|
||||
<td class="whitespace-nowrap px-3 py-3 text-gray-600 dark:text-dark-300">{{ endpoint.timeout_ms }} ms · {{ endpoint.input_limit }} chars</td>
|
||||
<td class="px-3 py-3">
|
||||
<span :class="hasCredential(endpoint) ? 'text-emerald-600 dark:text-emerald-300' : 'text-gray-500 dark:text-dark-400'">
|
||||
{{ hasCredential(endpoint) ? t('admin.promptAudit.pool.configured') : t('admin.promptAudit.pool.missing') }}
|
||||
</span>
|
||||
<p v-if="probingIds.includes(endpoint.id)" class="mt-1 text-xs text-primary-600 dark:text-primary-300">
|
||||
{{ t('admin.promptAudit.pool.probeProgress') }}
|
||||
</p>
|
||||
<p v-if="probeResults[endpoint.id]" class="mt-1 text-xs" :class="probeResults[endpoint.id].ok ? 'text-emerald-600' : 'text-red-600'">
|
||||
{{ t('admin.promptAudit.pool.probeResult', { status: probeResults[endpoint.id].status, http: probeResults[endpoint.id].http_status || '—', latency: probeResults[endpoint.id].latency_ms }) }}
|
||||
· {{ probeResults[endpoint.id].message }}
|
||||
</p>
|
||||
</td>
|
||||
<td class="whitespace-nowrap px-3 py-3 text-right">
|
||||
<button type="button" class="btn btn-ghost btn-sm" :disabled="probingIds.includes(endpoint.id)" @click="$emit('probe', endpoint)">
|
||||
{{ probingIds.includes(endpoint.id) ? t('admin.promptAudit.pool.probing') : t('admin.promptAudit.pool.probe') }}
|
||||
</button>
|
||||
<button type="button" class="btn btn-ghost btn-sm" @click="openEdit(endpoint)">{{ t('common.edit') }}</button>
|
||||
<button type="button" class="btn btn-ghost btn-sm text-red-600" @click="removeEndpoint(endpoint)">{{ t('common.delete') }}</button>
|
||||
</td>
|
||||
</tr>
|
||||
</tbody>
|
||||
</table>
|
||||
</div>
|
||||
|
||||
<BaseDialog :show="Boolean(editing)" :title="editingIndex < 0 ? t('admin.promptAudit.pool.add') : t('admin.promptAudit.pool.edit')" width="wide" @close="closeEditor">
|
||||
<form v-if="editing" class="grid gap-4 sm:grid-cols-2" @submit.prevent="saveEditor">
|
||||
<label class="space-y-1 text-sm text-gray-700 dark:text-dark-200">
|
||||
<span>{{ t('admin.promptAudit.pool.name') }}</span>
|
||||
<input v-model="editing.name" class="input w-full" required :aria-label="t('admin.promptAudit.pool.name')" />
|
||||
</label>
|
||||
<label class="space-y-1 text-sm text-gray-700 dark:text-dark-200">
|
||||
<span>{{ t('admin.promptAudit.pool.id') }}</span>
|
||||
<input v-model="editing.id" class="input w-full" required :disabled="editingIndex >= 0" :aria-label="t('admin.promptAudit.pool.id')" />
|
||||
</label>
|
||||
<label class="space-y-1 text-sm text-gray-700 dark:text-dark-200 sm:col-span-2">
|
||||
<span>{{ t('admin.promptAudit.pool.baseUrl') }}</span>
|
||||
<input v-model="editing.base_url" class="input w-full" required inputmode="url" :aria-label="t('admin.promptAudit.pool.baseUrl')" />
|
||||
</label>
|
||||
<label class="space-y-1 text-sm text-gray-700 dark:text-dark-200 sm:col-span-2">
|
||||
<span>{{ t('admin.promptAudit.pool.apiKey') }}</span>
|
||||
<input v-model="editing.token" class="input w-full" type="password" autocomplete="new-password" :placeholder="editing.has_token ? t('admin.promptAudit.pool.keepSecret') : ''" :aria-label="t('admin.promptAudit.pool.apiKey')" />
|
||||
<span class="block text-xs text-gray-500 dark:text-dark-400">{{ t('admin.promptAudit.pool.secretHint') }}</span>
|
||||
</label>
|
||||
<label v-if="editing.has_token" class="flex items-center gap-2 text-sm text-red-600 dark:text-red-300 sm:col-span-2">
|
||||
<input v-model="editing.clear_token" type="checkbox" :aria-label="t('admin.promptAudit.pool.clearSecret')" />
|
||||
{{ t('admin.promptAudit.pool.clearSecret') }}
|
||||
</label>
|
||||
<label class="space-y-1 text-sm text-gray-700 dark:text-dark-200 sm:col-span-2">
|
||||
<span>{{ t('admin.promptAudit.pool.model') }}</span>
|
||||
<input v-model="editing.model" class="input w-full" :aria-label="t('admin.promptAudit.pool.model')" />
|
||||
</label>
|
||||
<label class="space-y-1 text-sm text-gray-700 dark:text-dark-200">
|
||||
<span>{{ t('admin.promptAudit.pool.timeout') }}</span>
|
||||
<input v-model.number="editing.timeout_ms" class="input w-full" type="number" min="100" max="30000" required :aria-label="t('admin.promptAudit.pool.timeout')" />
|
||||
</label>
|
||||
<label class="space-y-1 text-sm text-gray-700 dark:text-dark-200">
|
||||
<span>{{ t('admin.promptAudit.pool.inputLimit') }}</span>
|
||||
<input v-model.number="editing.input_limit" class="input w-full" type="number" min="128" max="100000" required :aria-label="t('admin.promptAudit.pool.inputLimit')" />
|
||||
</label>
|
||||
</form>
|
||||
<template #footer>
|
||||
<div class="flex justify-end gap-3">
|
||||
<button type="button" class="btn btn-secondary" @click="closeEditor">{{ t('common.cancel') }}</button>
|
||||
<button type="button" class="btn btn-primary" data-test="save-endpoint" @click="saveEditor">{{ t('common.save') }}</button>
|
||||
</div>
|
||||
</template>
|
||||
</BaseDialog>
|
||||
</section>
|
||||
</template>
|
||||
|
||||
<script setup lang="ts">
|
||||
import { ref } from 'vue'
|
||||
import { useI18n } from 'vue-i18n'
|
||||
import BaseDialog from '@/components/common/BaseDialog.vue'
|
||||
import type { PromptAuditEndpointDraft, PromptProbeResult } from '../types'
|
||||
import { cloneData, createDefaultEndpoint } from '../viewModel'
|
||||
|
||||
const props = defineProps<{
|
||||
endpoints: PromptAuditEndpointDraft[]
|
||||
probeResults: Record<string, PromptProbeResult>
|
||||
probingIds: string[]
|
||||
}>()
|
||||
const emit = defineEmits<{
|
||||
(event: 'update:endpoints', value: PromptAuditEndpointDraft[]): void
|
||||
(event: 'probe', endpoint: PromptAuditEndpointDraft): void
|
||||
}>()
|
||||
const { t } = useI18n()
|
||||
const editing = ref<PromptAuditEndpointDraft | null>(null)
|
||||
const editingIndex = ref(-1)
|
||||
|
||||
function openCreate() {
|
||||
editingIndex.value = -1
|
||||
editing.value = createDefaultEndpoint(props.endpoints.length + 1)
|
||||
}
|
||||
function openEdit(endpoint: PromptAuditEndpointDraft) {
|
||||
editingIndex.value = props.endpoints.findIndex((item) => item.id === endpoint.id)
|
||||
editing.value = cloneData(endpoint)
|
||||
}
|
||||
function closeEditor() {
|
||||
editing.value = null
|
||||
editingIndex.value = -1
|
||||
}
|
||||
function saveEditor() {
|
||||
if (!editing.value?.id.trim() || !editing.value.name.trim() || !editing.value.base_url.trim()) return
|
||||
const next = props.endpoints.map((item) => cloneData(item))
|
||||
const value = cloneData(editing.value)
|
||||
if (value.token.trim()) value.clear_token = false
|
||||
if (editingIndex.value < 0) next.push(value)
|
||||
else next.splice(editingIndex.value, 1, value)
|
||||
emit('update:endpoints', next)
|
||||
closeEditor()
|
||||
}
|
||||
function toggleEndpoint(id: string) {
|
||||
emit('update:endpoints', props.endpoints.map((item) => item.id === id ? { ...item, enabled: !item.enabled } : cloneData(item)))
|
||||
}
|
||||
function removeEndpoint(endpoint: PromptAuditEndpointDraft) {
|
||||
if (!window.confirm(t('admin.promptAudit.pool.deleteConfirm', { name: endpoint.name }))) return
|
||||
emit('update:endpoints', props.endpoints.filter((item) => item.id !== endpoint.id).map((item) => cloneData(item)))
|
||||
}
|
||||
function hasCredential(endpoint: PromptAuditEndpointDraft): boolean {
|
||||
return Boolean(endpoint.token.trim() || (endpoint.has_token && !endpoint.clear_token))
|
||||
}
|
||||
</script>
|
||||
@@ -0,0 +1,66 @@
|
||||
<template>
|
||||
<BaseDialog :show="show" :title="t('admin.promptAudit.events.detailTitle')" width="extra-wide" @close="$emit('close')">
|
||||
<div v-if="loading" class="py-12 text-center text-sm text-gray-500" aria-busy="true">{{ t('common.loading') }}</div>
|
||||
<div v-else-if="event" class="space-y-5">
|
||||
<div class="flex flex-wrap gap-2 border-b border-gray-200 pb-3 dark:border-dark-700" role="tablist">
|
||||
<button v-for="tab in tabs" :key="tab" type="button" role="tab" :aria-selected="activeTab === tab" class="rounded-md px-3 py-1.5 text-sm" :class="activeTab === tab ? 'bg-primary-50 text-primary-700 dark:bg-primary-950/40 dark:text-primary-300' : 'text-gray-600 dark:text-dark-300'" @click="activeTab = tab">
|
||||
{{ t(`admin.promptAudit.events.tabs.${tab}`) }}
|
||||
</button>
|
||||
</div>
|
||||
|
||||
<div v-if="activeTab === 'summary'" class="grid gap-5 lg:grid-cols-2">
|
||||
<div>
|
||||
<h4 class="text-sm font-medium text-gray-900 dark:text-white">{{ t('admin.promptAudit.events.redactedPreview') }}</h4>
|
||||
<pre class="mt-2 max-h-56 overflow-auto whitespace-pre-wrap break-words rounded-lg bg-gray-50 p-4 text-sm text-gray-700 dark:bg-dark-900 dark:text-dark-200">{{ event.snapshot.redacted_preview || '—' }}</pre>
|
||||
</div>
|
||||
<dl class="grid grid-cols-[auto_1fr] gap-x-4 gap-y-2 text-sm">
|
||||
<dt class="text-gray-500">{{ t('admin.promptAudit.events.decision') }}</dt><dd class="font-medium text-gray-900 dark:text-white">{{ event.decision }} · {{ event.action }}</dd>
|
||||
<dt class="text-gray-500">{{ t('admin.promptAudit.events.user') }}</dt><dd>{{ event.snapshot.username || '—' }}</dd>
|
||||
<dt class="text-gray-500">{{ t('admin.promptAudit.events.email') }}</dt><dd>{{ event.snapshot.user_email || '—' }}</dd>
|
||||
<dt class="text-gray-500">{{ t('admin.promptAudit.events.apiKey') }}</dt><dd>{{ event.snapshot.api_key_name || '—' }}</dd>
|
||||
<dt class="text-gray-500">{{ t('admin.promptAudit.events.group') }}</dt><dd>{{ event.snapshot.group_name || '—' }}</dd>
|
||||
<dt class="text-gray-500">{{ t('admin.promptAudit.events.model') }}</dt><dd>{{ event.snapshot.model || '—' }}</dd>
|
||||
<dt class="text-gray-500">{{ t('admin.promptAudit.events.categories') }}</dt><dd>{{ event.categories.join(', ') || '—' }}</dd>
|
||||
</dl>
|
||||
</div>
|
||||
|
||||
<div v-else-if="activeTab === 'risks'" class="space-y-3">
|
||||
<article v-for="issue in event.issue_summaries" :key="`${issue.scanner_id}-${issue.code}`" class="border-l-2 border-red-400 pl-4">
|
||||
<div class="flex flex-wrap items-center gap-2">
|
||||
<h4 class="font-medium text-gray-900 dark:text-white">{{ issue.title }}</h4>
|
||||
<span class="text-xs uppercase text-red-600 dark:text-red-300">{{ issue.severity_label }} · {{ issue.action_label }}</span>
|
||||
</div>
|
||||
<p class="mt-1 text-sm text-gray-600 dark:text-dark-300">{{ issue.description }}</p>
|
||||
<p class="mt-2 break-words text-xs text-gray-500 dark:text-dark-400">{{ issue.scanner_id }} · {{ issue.code }} · {{ issue.score }} · {{ issue.evidence }}</p>
|
||||
</article>
|
||||
<p v-if="event.issue_summaries.length === 0" class="py-8 text-center text-sm text-gray-500">{{ t('admin.promptAudit.events.noRisks') }}</p>
|
||||
</div>
|
||||
|
||||
<dl v-else class="grid grid-cols-[auto_minmax(0,1fr)] gap-x-4 gap-y-2 text-sm">
|
||||
<dt class="text-gray-500">Request ID</dt><dd class="break-all font-mono">{{ event.snapshot.request_id || '—' }}</dd>
|
||||
<dt class="text-gray-500">Prompt SHA-256</dt><dd class="break-all font-mono">{{ event.snapshot.prompt_hash }}</dd>
|
||||
<dt class="text-gray-500">Scanner</dt><dd>{{ event.scanner_backend }} · {{ event.scanner_version }}</dd>
|
||||
<dt class="text-gray-500">Policy</dt><dd>{{ event.policy_id }} · v{{ event.policy_version }}</dd>
|
||||
<dt class="text-gray-500">Guard endpoint</dt><dd>{{ event.guard_endpoint_id }}</dd>
|
||||
<dt class="text-gray-500">Config</dt><dd>v{{ event.config_version }}</dd>
|
||||
<dt class="text-gray-500">Chunks</dt><dd>{{ event.chunk_total }}</dd>
|
||||
<dt class="text-gray-500">Latency</dt><dd>{{ event.latency_ms }} ms</dd>
|
||||
<dt class="text-gray-500">Protocol</dt><dd>{{ event.snapshot.protocol }} · {{ event.snapshot.endpoint }} · {{ event.snapshot.stage }}</dd>
|
||||
</dl>
|
||||
</div>
|
||||
</BaseDialog>
|
||||
</template>
|
||||
|
||||
<script setup lang="ts">
|
||||
import { ref, watch } from 'vue'
|
||||
import { useI18n } from 'vue-i18n'
|
||||
import BaseDialog from '@/components/common/BaseDialog.vue'
|
||||
import type { PromptAuditEvent } from '../types'
|
||||
|
||||
const props = defineProps<{ show: boolean; event: PromptAuditEvent | null; loading: boolean }>()
|
||||
defineEmits<{ (event: 'close'): void }>()
|
||||
const { t } = useI18n()
|
||||
const tabs = ['summary', 'risks', 'technical'] as const
|
||||
const activeTab = ref<(typeof tabs)[number]>('summary')
|
||||
watch(() => props.event?.id, () => { activeTab.value = 'summary' })
|
||||
</script>
|
||||
@@ -0,0 +1,186 @@
|
||||
<template>
|
||||
<section aria-labelledby="prompt-events-title" class="py-6">
|
||||
<div class="flex flex-wrap items-start justify-between gap-3">
|
||||
<div>
|
||||
<h2 id="prompt-events-title" class="text-base font-semibold text-gray-950 dark:text-white">{{ t('admin.promptAudit.events.title') }}</h2>
|
||||
<p class="mt-1 text-sm text-gray-500 dark:text-dark-300">{{ t('admin.promptAudit.events.description') }}</p>
|
||||
</div>
|
||||
<div class="flex flex-wrap gap-2">
|
||||
<button type="button" class="btn btn-secondary btn-sm" :disabled="selectedIds.length === 0" @click="$emit('batch-delete')">
|
||||
{{ t('admin.promptAudit.events.deleteSelected', { count: selectedIds.length }) }}
|
||||
</button>
|
||||
<button type="button" class="btn btn-danger btn-sm" :disabled="!hasExplicitDeleteRange(localFilters)" data-test="filter-delete" @click="$emit('preview-delete')">
|
||||
{{ t('admin.promptAudit.events.deleteByFilter') }}
|
||||
</button>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<form class="mt-5 grid gap-3 sm:grid-cols-2 lg:grid-cols-4 xl:grid-cols-5" @submit.prevent="applyFilters">
|
||||
<label class="text-xs text-gray-600 dark:text-dark-300">
|
||||
<span>{{ t('admin.promptAudit.events.decision') }}</span>
|
||||
<select v-model="localFilters.decision" class="input mt-1 w-full" :aria-label="t('admin.promptAudit.events.decision')" @change="filtersChanged">
|
||||
<option value="">{{ t('common.all') }}</option><option value="pass">pass</option><option value="flag">flag</option><option value="critical">critical</option>
|
||||
</select>
|
||||
</label>
|
||||
<label class="text-xs text-gray-600 dark:text-dark-300">
|
||||
<span>{{ t('admin.promptAudit.events.risk') }}</span>
|
||||
<select v-model="localFilters.risk_level" class="input mt-1 w-full" :aria-label="t('admin.promptAudit.events.risk')" @change="filtersChanged">
|
||||
<option value="">{{ t('common.all') }}</option><option value="low">low</option><option value="medium">medium</option><option value="high">high</option><option value="critical">critical</option>
|
||||
</select>
|
||||
</label>
|
||||
<FilterInput v-model="localFilters.endpoint" :label="t('admin.promptAudit.events.endpoint')" @change="filtersChanged" />
|
||||
<FilterInput v-model="localFilters.group_id" :label="t('admin.promptAudit.events.groupId')" type="number" @change="filtersChanged" />
|
||||
<FilterInput v-model="localFilters.user_id" :label="t('admin.promptAudit.events.userId')" type="number" @change="filtersChanged" />
|
||||
<FilterInput v-model="localFilters.api_key_id" :label="t('admin.promptAudit.events.apiKeyId')" type="number" @change="filtersChanged" />
|
||||
<FilterInput v-model="localFilters.request_id" label="Request ID" @change="filtersChanged" />
|
||||
<FilterInput v-model="localFilters.prompt_hash" label="Prompt SHA-256" @change="filtersChanged" />
|
||||
<FilterInput v-model="localFilters.keyword" :label="t('admin.promptAudit.events.keyword')" @change="filtersChanged" />
|
||||
<label class="text-xs text-gray-600 dark:text-dark-300">
|
||||
<span>{{ t('admin.promptAudit.events.startAt') }}</span>
|
||||
<input v-model="localFilters.start_at" type="datetime-local" class="input mt-1 w-full" :aria-label="t('admin.promptAudit.events.startAt')" @change="filtersChanged" />
|
||||
</label>
|
||||
<label class="text-xs text-gray-600 dark:text-dark-300">
|
||||
<span>{{ t('admin.promptAudit.events.endAt') }}</span>
|
||||
<input v-model="localFilters.end_at" type="datetime-local" class="input mt-1 w-full" :aria-label="t('admin.promptAudit.events.endAt')" @change="filtersChanged" />
|
||||
</label>
|
||||
<div class="flex items-end gap-2 sm:col-span-2">
|
||||
<button type="submit" class="btn btn-primary btn-sm">{{ t('common.search') }}</button>
|
||||
<button type="button" class="btn btn-ghost btn-sm" @click="resetFilters">{{ t('common.reset') }}</button>
|
||||
</div>
|
||||
</form>
|
||||
<p class="mt-2 text-xs text-gray-500 dark:text-dark-400">{{ t('admin.promptAudit.events.deleteRangeHint') }}</p>
|
||||
|
||||
<div v-if="error" role="alert" class="mt-4 rounded-lg bg-red-50 px-4 py-3 text-sm text-red-700 dark:bg-red-950/30 dark:text-red-300">{{ error }}</div>
|
||||
<div class="mt-5 overflow-x-auto rounded-lg border border-gray-200 dark:border-dark-700">
|
||||
<table class="min-w-[1120px] w-full text-left text-sm">
|
||||
<thead class="bg-gray-50 text-xs uppercase tracking-wide text-gray-500 dark:bg-dark-900 dark:text-dark-400">
|
||||
<tr>
|
||||
<th class="w-10 px-3 py-3"><input type="checkbox" :checked="allSelected" :aria-label="t('admin.promptAudit.events.selectAll')" @change="toggleAll" /></th>
|
||||
<th class="px-3 py-3 font-medium">{{ t('admin.promptAudit.events.time') }}</th>
|
||||
<th class="px-3 py-3 font-medium">{{ t('admin.promptAudit.events.identity') }}</th>
|
||||
<th class="px-3 py-3 font-medium">{{ t('admin.promptAudit.events.group') }}</th>
|
||||
<th class="px-3 py-3 font-medium">{{ t('admin.promptAudit.events.route') }}</th>
|
||||
<th class="px-3 py-3 font-medium">{{ t('admin.promptAudit.events.result') }}</th>
|
||||
<th class="px-3 py-3 font-medium">{{ t('admin.promptAudit.events.preview') }}</th>
|
||||
<th class="px-3 py-3 text-right font-medium">{{ t('admin.promptAudit.common.actions') }}</th>
|
||||
</tr>
|
||||
</thead>
|
||||
<tbody class="divide-y divide-gray-100 bg-white dark:divide-dark-800 dark:bg-dark-850">
|
||||
<tr v-if="loading"><td colspan="8" class="px-4 py-12 text-center text-gray-500" aria-busy="true">{{ t('common.loading') }}</td></tr>
|
||||
<tr v-else-if="events.length === 0"><td colspan="8" class="px-4 py-12 text-center text-gray-500">{{ t('admin.promptAudit.events.empty') }}</td></tr>
|
||||
<tr v-for="event in events" v-else :key="event.id" :data-test="`event-${event.id}`" class="align-top hover:bg-gray-50/70 dark:hover:bg-dark-800/70">
|
||||
<td class="px-3 py-3"><input type="checkbox" :checked="selectedIds.includes(event.id)" :aria-label="t('admin.promptAudit.events.selectEvent', { id: event.id })" @change="toggleOne(event.id)" /></td>
|
||||
<td class="whitespace-nowrap px-3 py-3 text-xs text-gray-600 dark:text-dark-300">{{ formatDate(event.created_at) }}</td>
|
||||
<td class="px-3 py-3">
|
||||
<CopyLine :label="t('admin.promptAudit.events.user')" :value="event.snapshot.username" />
|
||||
<CopyLine :label="t('admin.promptAudit.events.email')" :value="event.snapshot.user_email" />
|
||||
<CopyLine :label="t('admin.promptAudit.events.apiKey')" :value="event.snapshot.api_key_name" />
|
||||
</td>
|
||||
<td class="px-3 py-3 text-gray-700 dark:text-dark-200">{{ event.snapshot.group_name || '—' }}</td>
|
||||
<td class="px-3 py-3">
|
||||
<p class="font-medium text-gray-900 dark:text-white">{{ event.snapshot.endpoint }}</p>
|
||||
<p class="mt-1 text-xs text-gray-500">{{ event.snapshot.model }} · {{ event.snapshot.protocol }}</p>
|
||||
</td>
|
||||
<td class="px-3 py-3">
|
||||
<span class="rounded-full px-2 py-0.5 text-xs font-medium" :class="decisionClass(event.decision)">{{ event.decision }} · {{ event.risk_level }}</span>
|
||||
<p class="mt-2 max-w-48 truncate text-xs text-gray-500">{{ event.categories.join(', ') || '—' }}</p>
|
||||
</td>
|
||||
<td class="max-w-xs px-3 py-3"><p class="line-clamp-2 break-words text-gray-600 dark:text-dark-300">{{ event.snapshot.redacted_preview || '—' }}</p></td>
|
||||
<td class="whitespace-nowrap px-3 py-3 text-right">
|
||||
<button type="button" class="btn btn-ghost btn-sm" @click="$emit('view', event.id)">{{ t('common.view') }}</button>
|
||||
<button type="button" class="btn btn-ghost btn-sm text-red-600" @click="$emit('delete', event.id)">{{ t('common.delete') }}</button>
|
||||
</td>
|
||||
</tr>
|
||||
</tbody>
|
||||
</table>
|
||||
<Pagination :total="total" :page="page" :page-size="pageSize" @update:page="$emit('page', $event)" @update:page-size="$emit('page-size', $event)" />
|
||||
</div>
|
||||
</section>
|
||||
</template>
|
||||
|
||||
<script setup lang="ts">
|
||||
import { computed, defineComponent, h, reactive, watch } from 'vue'
|
||||
import { useI18n } from 'vue-i18n'
|
||||
import Pagination from '@/components/common/Pagination.vue'
|
||||
import type { PromptAuditEvent, PromptEventFilters } from '../types'
|
||||
import { cloneData, emptyEventFilters, hasExplicitDeleteRange } from '../viewModel'
|
||||
|
||||
const props = defineProps<{
|
||||
events: PromptAuditEvent[]; total: number; page: number; pageSize: number
|
||||
filters: PromptEventFilters; selectedIds: number[]; loading: boolean; error: string
|
||||
}>()
|
||||
const emit = defineEmits<{
|
||||
(event: 'filters-change', value: PromptEventFilters): void
|
||||
(event: 'search', value: PromptEventFilters): void
|
||||
(event: 'selection', value: number[]): void
|
||||
(event: 'page', value: number): void
|
||||
(event: 'page-size', value: number): void
|
||||
(event: 'view', id: number): void
|
||||
(event: 'delete', id: number): void
|
||||
(event: 'batch-delete'): void
|
||||
(event: 'preview-delete'): void
|
||||
}>()
|
||||
const { t, locale } = useI18n()
|
||||
const localFilters = reactive<PromptEventFilters>(cloneData(props.filters))
|
||||
watch(() => props.filters, (value) => Object.assign(localFilters, cloneData(value)), { deep: true })
|
||||
const allSelected = computed(() => props.events.length > 0 && props.events.every((event) => props.selectedIds.includes(event.id)))
|
||||
|
||||
const FilterInput = defineComponent({
|
||||
props: { modelValue: { type: String, required: true }, label: { type: String, required: true }, type: { type: String, default: 'text' } },
|
||||
emits: ['update:modelValue', 'change'],
|
||||
setup(componentProps, { emit: componentEmit }) {
|
||||
return () => h('label', { class: 'text-xs text-gray-600 dark:text-dark-300' }, [
|
||||
h('span', componentProps.label),
|
||||
h('input', {
|
||||
value: componentProps.modelValue, type: componentProps.type, class: 'input mt-1 w-full', 'aria-label': componentProps.label,
|
||||
onInput: (event: Event) => componentEmit('update:modelValue', (event.target as HTMLInputElement).value),
|
||||
onChange: () => componentEmit('change'),
|
||||
}),
|
||||
])
|
||||
},
|
||||
})
|
||||
|
||||
const CopyLine = defineComponent({
|
||||
props: { label: { type: String, required: true }, value: { type: String, default: '' } },
|
||||
setup(componentProps) {
|
||||
return () => h('div', { class: 'flex max-w-56 items-center gap-1 text-xs' }, [
|
||||
h('span', { class: 'w-16 flex-none text-gray-500' }, componentProps.label),
|
||||
h('span', { class: 'min-w-0 flex-1 truncate text-gray-800 dark:text-dark-100' }, componentProps.value || '—'),
|
||||
componentProps.value ? h('button', {
|
||||
type: 'button', class: 'text-primary-600 hover:underline', 'aria-label': `${t('common.copy')} ${componentProps.label}`,
|
||||
onClick: () => navigator.clipboard?.writeText(componentProps.value),
|
||||
}, t('common.copy')) : null,
|
||||
])
|
||||
},
|
||||
})
|
||||
|
||||
function filtersChanged() {
|
||||
emit('filters-change', cloneData(localFilters))
|
||||
}
|
||||
function applyFilters() {
|
||||
const value = cloneData(localFilters)
|
||||
emit('filters-change', value)
|
||||
emit('search', value)
|
||||
}
|
||||
function resetFilters() {
|
||||
Object.assign(localFilters, emptyEventFilters())
|
||||
applyFilters()
|
||||
}
|
||||
function toggleOne(id: number) {
|
||||
const selected = new Set(props.selectedIds)
|
||||
if (selected.has(id)) selected.delete(id)
|
||||
else selected.add(id)
|
||||
emit('selection', [...selected])
|
||||
}
|
||||
function toggleAll() {
|
||||
emit('selection', allSelected.value ? [] : props.events.map((event) => event.id))
|
||||
}
|
||||
function formatDate(value: string): string {
|
||||
return new Intl.DateTimeFormat(locale.value, { dateStyle: 'short', timeStyle: 'medium' }).format(new Date(value))
|
||||
}
|
||||
function decisionClass(decision: string): string {
|
||||
if (decision === 'critical') return 'bg-red-100 text-red-700 dark:bg-red-950/50 dark:text-red-300'
|
||||
if (decision === 'flag') return 'bg-amber-100 text-amber-700 dark:bg-amber-950/50 dark:text-amber-300'
|
||||
return 'bg-emerald-100 text-emerald-700 dark:bg-emerald-950/50 dark:text-emerald-300'
|
||||
}
|
||||
</script>
|
||||
@@ -0,0 +1,108 @@
|
||||
<template>
|
||||
<section aria-labelledby="prompt-policy-title" class="border-b border-gray-200 py-6 dark:border-dark-700">
|
||||
<div>
|
||||
<h2 id="prompt-policy-title" class="text-base font-semibold text-gray-950 dark:text-white">{{ t('admin.promptAudit.policy.title') }}</h2>
|
||||
<p class="mt-1 text-sm text-gray-500 dark:text-dark-300">{{ t('admin.promptAudit.policy.description') }}</p>
|
||||
</div>
|
||||
|
||||
<div class="mt-5 grid gap-7 lg:grid-cols-[minmax(0,1fr)_minmax(280px,0.55fr)]">
|
||||
<div>
|
||||
<fieldset>
|
||||
<legend class="text-sm font-medium text-gray-900 dark:text-white">{{ t('admin.promptAudit.policy.scope') }}</legend>
|
||||
<div class="mt-3 flex flex-wrap gap-5 text-sm text-gray-700 dark:text-dark-200">
|
||||
<label class="flex items-center gap-2">
|
||||
<input type="radio" name="prompt-audit-scope" :checked="draft.all_groups" @change="patch({ all_groups: true, group_ids: [] })" />
|
||||
{{ t('admin.promptAudit.policy.allGroups') }}
|
||||
</label>
|
||||
<label class="flex items-center gap-2">
|
||||
<input type="radio" name="prompt-audit-scope" :checked="!draft.all_groups" @change="patch({ all_groups: false })" />
|
||||
{{ t('admin.promptAudit.policy.selectedGroups') }}
|
||||
</label>
|
||||
</div>
|
||||
</fieldset>
|
||||
|
||||
<div v-if="!draft.all_groups" class="mt-4">
|
||||
<label class="block text-sm text-gray-700 dark:text-dark-200">
|
||||
<span>{{ t('admin.promptAudit.policy.searchGroups') }}</span>
|
||||
<input v-model="groupSearch" type="search" class="input mt-1.5 w-full" :aria-label="t('admin.promptAudit.policy.searchGroups')" />
|
||||
</label>
|
||||
<div class="mt-3 max-h-52 overflow-y-auto rounded-lg border border-gray-200 p-2 dark:border-dark-700">
|
||||
<label v-for="group in filteredGroups" :key="group.id" class="flex cursor-pointer items-center justify-between gap-3 rounded-md px-2 py-2 text-sm hover:bg-gray-50 dark:hover:bg-dark-800">
|
||||
<span class="flex items-center gap-2 text-gray-800 dark:text-dark-100">
|
||||
<input type="checkbox" :checked="draft.group_ids.includes(group.id)" @change="toggleGroup(group.id)" />
|
||||
{{ group.name }}
|
||||
</span>
|
||||
<span class="text-xs text-gray-500 dark:text-dark-400">{{ group.platform }} · {{ group.status }}</span>
|
||||
</label>
|
||||
<p v-if="filteredGroups.length === 0" class="px-2 py-4 text-center text-sm text-gray-500">{{ t('admin.promptAudit.policy.noGroups') }}</p>
|
||||
</div>
|
||||
<div v-if="missingGroupIds.length" class="mt-3 rounded-lg bg-amber-50 px-3 py-2 text-sm text-amber-800 dark:bg-amber-950/30 dark:text-amber-200">
|
||||
{{ t('admin.promptAudit.policy.missingGroups') }}: {{ missingGroupIds.join(', ') }}
|
||||
</div>
|
||||
<p class="mt-2 text-xs text-gray-500 dark:text-dark-400">{{ t('admin.promptAudit.policy.selectedCount', { count: draft.group_ids.length }) }}</p>
|
||||
</div>
|
||||
|
||||
<fieldset class="mt-6">
|
||||
<legend class="text-sm font-medium text-gray-900 dark:text-white">{{ t('admin.promptAudit.policy.scanners') }}</legend>
|
||||
<div class="mt-3 grid gap-2 sm:grid-cols-2">
|
||||
<label v-for="scanner in SCANNER_CATALOG" :key="scanner.id" class="flex items-center gap-2 rounded-md px-2 py-1.5 text-sm text-gray-700 hover:bg-gray-50 dark:text-dark-200 dark:hover:bg-dark-800">
|
||||
<input type="checkbox" :checked="draft.scanners.includes(scanner.id)" :aria-label="scanner.label" @change="toggleScanner(scanner.id)" />
|
||||
<span>{{ scanner.label }}</span>
|
||||
</label>
|
||||
</div>
|
||||
</fieldset>
|
||||
</div>
|
||||
|
||||
<div class="space-y-5 border-gray-200 lg:border-l lg:pl-7 dark:border-dark-700">
|
||||
<label class="block text-sm text-gray-700 dark:text-dark-200">
|
||||
<span>{{ t('admin.promptAudit.policy.workerCount') }}</span>
|
||||
<input :value="draft.worker_count" type="number" min="1" max="32" class="input mt-1.5 w-full" :aria-label="t('admin.promptAudit.policy.workerCount')" @input="patch({ worker_count: Number(($event.target as HTMLInputElement).value) })" />
|
||||
</label>
|
||||
<label class="block text-sm text-gray-700 dark:text-dark-200">
|
||||
<span>{{ t('admin.promptAudit.policy.queueCapacity') }}</span>
|
||||
<input :value="draft.queue_capacity" type="number" min="1" max="100000" class="input mt-1.5 w-full" :aria-label="t('admin.promptAudit.policy.queueCapacity')" @input="patch({ queue_capacity: Number(($event.target as HTMLInputElement).value) })" />
|
||||
</label>
|
||||
<div class="rounded-lg bg-gray-50 px-4 py-3 text-sm text-gray-600 dark:bg-dark-800 dark:text-dark-300">
|
||||
<p class="font-medium text-gray-800 dark:text-dark-100">{{ t('admin.promptAudit.policy.strategy') }}</p>
|
||||
<p class="mt-1">priority · {{ t('admin.promptAudit.policy.strategyHint') }}</p>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
</section>
|
||||
</template>
|
||||
|
||||
<script setup lang="ts">
|
||||
import { computed, ref } from 'vue'
|
||||
import { useI18n } from 'vue-i18n'
|
||||
import type { PromptAuditDraft, PromptAuditGroup } from '../types'
|
||||
import { cloneData, SCANNER_CATALOG } from '../viewModel'
|
||||
|
||||
const props = defineProps<{ draft: PromptAuditDraft; groups: PromptAuditGroup[] }>()
|
||||
const emit = defineEmits<{ (event: 'update:draft', value: PromptAuditDraft): void }>()
|
||||
const { t } = useI18n()
|
||||
const groupSearch = ref('')
|
||||
|
||||
const filteredGroups = computed(() => {
|
||||
const query = groupSearch.value.trim().toLowerCase()
|
||||
if (!query) return props.groups
|
||||
return props.groups.filter((group) => `${group.name} ${group.id} ${group.platform}`.toLowerCase().includes(query))
|
||||
})
|
||||
const knownGroupIds = computed(() => new Set(props.groups.map((group) => group.id)))
|
||||
const missingGroupIds = computed(() => props.draft.group_ids.filter((id) => !knownGroupIds.value.has(id)))
|
||||
|
||||
function patch(value: Partial<PromptAuditDraft>) {
|
||||
emit('update:draft', { ...cloneData(props.draft), ...value })
|
||||
}
|
||||
function toggleGroup(id: number) {
|
||||
const selected = new Set(props.draft.group_ids)
|
||||
if (selected.has(id)) selected.delete(id)
|
||||
else selected.add(id)
|
||||
patch({ group_ids: [...selected].sort((a, b) => a - b) })
|
||||
}
|
||||
function toggleScanner(id: string) {
|
||||
const selected = new Set(props.draft.scanners)
|
||||
if (selected.has(id)) selected.delete(id)
|
||||
else selected.add(id)
|
||||
patch({ scanners: SCANNER_CATALOG.map((item) => item.id).filter((item) => selected.has(item)) })
|
||||
}
|
||||
</script>
|
||||
@@ -0,0 +1,119 @@
|
||||
<template>
|
||||
<section aria-labelledby="prompt-runtime-title" class="border-b border-gray-200 pb-6 dark:border-dark-700">
|
||||
<div class="flex flex-wrap items-start justify-between gap-3">
|
||||
<div>
|
||||
<h2 id="prompt-runtime-title" class="text-base font-semibold text-gray-950 dark:text-white">
|
||||
{{ t('admin.promptAudit.runtime.title') }}
|
||||
</h2>
|
||||
<p class="mt-1 text-sm text-gray-500 dark:text-dark-300">
|
||||
{{ t('admin.promptAudit.runtime.description') }}
|
||||
</p>
|
||||
</div>
|
||||
<button type="button" class="btn btn-secondary btn-sm" :disabled="loading" @click="$emit('refresh')">
|
||||
{{ t('admin.promptAudit.actions.refresh') }}
|
||||
</button>
|
||||
</div>
|
||||
|
||||
<div v-if="error" role="alert" class="mt-4 rounded-lg bg-red-50 px-4 py-3 text-sm text-red-700 dark:bg-red-950/30 dark:text-red-300">
|
||||
{{ error }}
|
||||
</div>
|
||||
<div v-else-if="loading && !runtime" class="mt-5 grid grid-cols-2 gap-4 lg:grid-cols-6" aria-busy="true">
|
||||
<div v-for="index in 6" :key="index" class="h-14 animate-pulse rounded-lg bg-gray-100 dark:bg-dark-800" />
|
||||
</div>
|
||||
<template v-else-if="runtime">
|
||||
<dl class="mt-5 grid grid-cols-2 gap-x-6 gap-y-5 lg:grid-cols-6">
|
||||
<div>
|
||||
<dt class="text-xs uppercase tracking-wide text-gray-500 dark:text-dark-400">{{ t('admin.promptAudit.runtime.process') }}</dt>
|
||||
<dd class="mt-1 flex items-center gap-2 text-sm font-semibold text-gray-900 dark:text-white">
|
||||
<span class="h-2 w-2 rounded-full" :class="statusDot(runtime.process_status)" />
|
||||
{{ t(`admin.promptAudit.status.${runtime.process_status}`) }}
|
||||
</dd>
|
||||
</div>
|
||||
<div>
|
||||
<dt class="text-xs uppercase tracking-wide text-gray-500 dark:text-dark-400">{{ t('admin.promptAudit.runtime.mode') }}</dt>
|
||||
<dd class="mt-1 text-sm font-semibold text-gray-900 dark:text-white">{{ t(`admin.promptAudit.mode.${runtime.effective_mode}`) }}</dd>
|
||||
</div>
|
||||
<div>
|
||||
<dt class="text-xs uppercase tracking-wide text-gray-500 dark:text-dark-400">{{ t('admin.promptAudit.runtime.version') }}</dt>
|
||||
<dd class="mt-1 text-sm font-semibold text-gray-900 dark:text-white">
|
||||
{{ runtime.active_config_version }} / {{ runtime.expected_config_version }}
|
||||
</dd>
|
||||
</div>
|
||||
<div>
|
||||
<dt class="text-xs uppercase tracking-wide text-gray-500 dark:text-dark-400">{{ t('admin.promptAudit.runtime.workers') }}</dt>
|
||||
<dd class="mt-1 text-sm font-semibold text-gray-900 dark:text-white">{{ runtime.worker_active }} / {{ runtime.worker_total }}</dd>
|
||||
</div>
|
||||
<div>
|
||||
<dt class="text-xs uppercase tracking-wide text-gray-500 dark:text-dark-400">{{ t('admin.promptAudit.runtime.queue') }}</dt>
|
||||
<dd class="mt-1 text-sm font-semibold text-gray-900 dark:text-white">{{ runtime.queue.active }} / {{ runtime.queue_capacity }}</dd>
|
||||
</div>
|
||||
<div>
|
||||
<dt class="text-xs uppercase tracking-wide text-gray-500 dark:text-dark-400">{{ t('admin.promptAudit.runtime.dependencies') }}</dt>
|
||||
<dd class="mt-1 text-sm font-semibold text-gray-900 dark:text-white">DB {{ runtime.database_status }} · Redis {{ runtime.redis_status }}</dd>
|
||||
</div>
|
||||
</dl>
|
||||
|
||||
<div class="mt-5 grid gap-4 border-t border-gray-100 pt-5 dark:border-dark-800 lg:grid-cols-[1.2fr_1fr]">
|
||||
<div class="min-w-0">
|
||||
<h3 class="text-sm font-medium text-gray-900 dark:text-white">{{ t('admin.promptAudit.runtime.guardMetrics') }}</h3>
|
||||
<div class="mt-3 flex flex-wrap gap-x-5 gap-y-2 text-sm text-gray-600 dark:text-dark-300">
|
||||
<span>{{ t('admin.promptAudit.metrics.total') }} <strong>{{ runtime.guard_metrics.total }}</strong></span>
|
||||
<span>{{ t('admin.promptAudit.metrics.allowed') }} <strong>{{ runtime.guard_metrics.allowed }}</strong></span>
|
||||
<span>{{ t('admin.promptAudit.metrics.flagged') }} <strong>{{ runtime.guard_metrics.flagged }}</strong></span>
|
||||
<span>{{ t('admin.promptAudit.metrics.blocked') }} <strong>{{ runtime.guard_metrics.blocked }}</strong></span>
|
||||
<span>{{ t('admin.promptAudit.metrics.unavailable') }} <strong>{{ runtime.guard_metrics.unavailable }}</strong></span>
|
||||
<span>{{ t('admin.promptAudit.metrics.timeouts') }} <strong>{{ runtime.guard_metrics.timeouts }}</strong></span>
|
||||
<span>{{ t('admin.promptAudit.metrics.failovers') }} <strong>{{ runtime.guard_metrics.failovers }}</strong></span>
|
||||
<span v-if="runtime.guard_metrics.latency_p95_ms != null">P95 <strong>{{ runtime.guard_metrics.latency_p95_ms }} ms</strong></span>
|
||||
</div>
|
||||
<p class="mt-3 text-xs text-gray-500 dark:text-dark-400">
|
||||
{{ t('admin.promptAudit.runtime.queueBreakdown', {
|
||||
queued: runtime.queue.queued,
|
||||
processing: runtime.queue.processing,
|
||||
retry: runtime.queue.retry,
|
||||
done: runtime.queue.done,
|
||||
failed: runtime.queue.failed,
|
||||
}) }}
|
||||
</p>
|
||||
<p class="mt-1 text-xs text-gray-500 dark:text-dark-400">
|
||||
{{ t('admin.promptAudit.runtime.deliveryTotals', { enqueued: runtime.enqueued_total, dropped: runtime.dropped_total, processed: runtime.processed_total, failed: runtime.failed_total }) }}
|
||||
</p>
|
||||
</div>
|
||||
<div>
|
||||
<h3 class="text-sm font-medium text-gray-900 dark:text-white">{{ t('admin.promptAudit.runtime.latest') }}</h3>
|
||||
<p class="mt-2 text-sm text-gray-600 dark:text-dark-300">
|
||||
{{ runtime.last_processed_at ? formatDate(runtime.last_processed_at) : t('admin.promptAudit.common.never') }}
|
||||
</p>
|
||||
<p v-if="runtime.last_error_code" class="mt-1 break-words text-sm text-red-600 dark:text-red-300">
|
||||
{{ runtime.last_error_code }}<span v-if="runtime.last_error_message"> · {{ runtime.last_error_message }}</span>
|
||||
</p>
|
||||
<div v-if="Object.keys(runtime.endpoints).length" class="mt-3 flex flex-wrap gap-2">
|
||||
<span v-for="(probe, id) in runtime.endpoints" :key="id" class="rounded-full px-2 py-1 text-xs" :class="probe.ok ? 'bg-emerald-50 text-emerald-700 dark:bg-emerald-950/40 dark:text-emerald-300' : 'bg-red-50 text-red-700 dark:bg-red-950/40 dark:text-red-300'">
|
||||
{{ id }} · {{ probe.status }} · {{ probe.latency_ms }} ms
|
||||
</span>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
</template>
|
||||
</section>
|
||||
</template>
|
||||
|
||||
<script setup lang="ts">
|
||||
import { useI18n } from 'vue-i18n'
|
||||
import type { PromptAuditRuntime } from '../types'
|
||||
|
||||
defineProps<{ runtime: PromptAuditRuntime | null; loading: boolean; error: string }>()
|
||||
defineEmits<{ (event: 'refresh'): void }>()
|
||||
const { t, locale } = useI18n()
|
||||
|
||||
function formatDate(value: string): string {
|
||||
return new Intl.DateTimeFormat(locale.value, { dateStyle: 'medium', timeStyle: 'medium' }).format(new Date(value))
|
||||
}
|
||||
|
||||
function statusDot(status: string): string {
|
||||
if (status === 'running') return 'bg-emerald-500'
|
||||
if (status === 'disabled') return 'bg-gray-400'
|
||||
if (status === 'degraded') return 'bg-amber-500'
|
||||
return 'bg-red-500'
|
||||
}
|
||||
</script>
|
||||
@@ -0,0 +1,243 @@
|
||||
export type PromptAuditMode = 'off' | 'async_audit' | 'blocking'
|
||||
export type PromptDecision = 'pass' | 'flag' | 'critical'
|
||||
export type PromptRiskLevel = 'low' | 'medium' | 'high' | 'critical'
|
||||
|
||||
export interface PromptAuditEndpoint {
|
||||
id: string
|
||||
name: string
|
||||
protocol: 'openai_compatible'
|
||||
base_url: string
|
||||
model: string
|
||||
timeout_ms: number
|
||||
input_limit: number
|
||||
enabled: boolean
|
||||
has_token: boolean
|
||||
token_status: 'configured' | 'missing' | string
|
||||
}
|
||||
|
||||
export interface PromptAuditEndpointDraft extends PromptAuditEndpoint {
|
||||
token: string
|
||||
clear_token: boolean
|
||||
}
|
||||
|
||||
export interface PromptAuditConfig {
|
||||
enabled: boolean
|
||||
blocking_enabled: boolean
|
||||
store_pass_events: boolean
|
||||
effective_mode: PromptAuditMode
|
||||
strategy: 'priority'
|
||||
worker_count: number
|
||||
queue_capacity: number
|
||||
scanners: string[]
|
||||
all_groups: boolean
|
||||
group_ids: number[]
|
||||
endpoints: PromptAuditEndpoint[]
|
||||
config_version: number
|
||||
updated_at: string
|
||||
updated_by: number
|
||||
change_summary: string
|
||||
}
|
||||
|
||||
export interface PromptAuditDraft extends Omit<PromptAuditConfig, 'endpoints'> {
|
||||
endpoints: PromptAuditEndpointDraft[]
|
||||
}
|
||||
|
||||
export interface PromptAuditUpdateRequest {
|
||||
expected_config_version: number
|
||||
enabled: boolean
|
||||
blocking_enabled: boolean
|
||||
store_pass_events: boolean
|
||||
strategy: 'priority'
|
||||
worker_count: number
|
||||
queue_capacity: number
|
||||
scanners: string[]
|
||||
all_groups: boolean
|
||||
group_ids: number[]
|
||||
endpoints: Array<{
|
||||
id: string
|
||||
name: string
|
||||
protocol: 'openai_compatible'
|
||||
base_url: string
|
||||
model: string
|
||||
token?: string
|
||||
clear_token: boolean
|
||||
timeout_ms: number
|
||||
input_limit: number
|
||||
enabled: boolean
|
||||
}>
|
||||
}
|
||||
|
||||
export interface PromptProbeResult {
|
||||
ok: boolean
|
||||
status: string
|
||||
error_code?: string
|
||||
message: string
|
||||
latency_ms: number
|
||||
http_status: number
|
||||
retryable: boolean
|
||||
checked_at: string
|
||||
token_applied: boolean
|
||||
}
|
||||
|
||||
export interface PromptQueueStats {
|
||||
staging: number
|
||||
queued: number
|
||||
processing: number
|
||||
retry: number
|
||||
done: number
|
||||
failed: number
|
||||
active: number
|
||||
}
|
||||
|
||||
export interface PromptGuardMetrics {
|
||||
total: number
|
||||
allowed: number
|
||||
flagged: number
|
||||
blocked: number
|
||||
unavailable: number
|
||||
invalid: number
|
||||
timeouts: number
|
||||
failovers: number
|
||||
bulkhead_full: number
|
||||
record_failed: number
|
||||
latency_avg_ms?: number
|
||||
latency_p50_ms?: number
|
||||
latency_p95_ms?: number
|
||||
latency_p99_ms?: number
|
||||
latency_max_ms?: number
|
||||
}
|
||||
|
||||
export interface PromptAuditRuntime {
|
||||
process_status: 'disabled' | 'running' | 'degraded' | 'error' | string
|
||||
effective_mode: PromptAuditMode
|
||||
expected_config_version: number
|
||||
active_config_version: number
|
||||
config_loaded_at?: string
|
||||
config_load_error?: string
|
||||
worker_total: number
|
||||
worker_active: number
|
||||
worker_heartbeat_at?: string
|
||||
queue_capacity: number
|
||||
queue: PromptQueueStats
|
||||
processed_total: number
|
||||
failed_total: number
|
||||
enqueued_total: number
|
||||
dropped_total: number
|
||||
last_processed_at?: string
|
||||
last_error_code?: string
|
||||
last_error_message?: string
|
||||
database_status: string
|
||||
redis_status: string
|
||||
endpoints: Record<string, PromptProbeResult>
|
||||
guard_metrics: PromptGuardMetrics
|
||||
}
|
||||
|
||||
export interface PromptSnapshot {
|
||||
request_id: string
|
||||
user_id: number
|
||||
username: string
|
||||
user_email: string
|
||||
api_key_id: number
|
||||
api_key_name: string
|
||||
group_id?: number
|
||||
group_name: string
|
||||
provider: string
|
||||
endpoint: string
|
||||
protocol: string
|
||||
model: string
|
||||
prompt_hash: string
|
||||
redacted_preview: string
|
||||
prompt_length: number
|
||||
message_count: number
|
||||
stage: string
|
||||
}
|
||||
|
||||
export interface PromptIssueSummary {
|
||||
category: string
|
||||
scanner_id: string
|
||||
title: string
|
||||
description: string
|
||||
severity: string
|
||||
severity_label: string
|
||||
action: string
|
||||
action_label: string
|
||||
code: string
|
||||
score: number
|
||||
evidence: string
|
||||
evidence_hash: string
|
||||
start_rune?: number
|
||||
end_rune?: number
|
||||
}
|
||||
|
||||
export interface PromptAuditEvent {
|
||||
id: number
|
||||
job_id: number
|
||||
snapshot: PromptSnapshot
|
||||
decision: PromptDecision
|
||||
risk_level: PromptRiskLevel
|
||||
action: 'Allow' | 'Warn' | 'Block' | string
|
||||
categories: string[]
|
||||
matched_scanners: string[]
|
||||
scanner_scores: Record<string, number>
|
||||
scanner_evidence: Record<string, string>
|
||||
scanner_backend: string
|
||||
scanner_version: string
|
||||
guard_endpoint_id: string
|
||||
policy_id: string
|
||||
policy_version: number
|
||||
config_version: number
|
||||
chunk_total: number
|
||||
latency_ms: number
|
||||
issue_summaries: PromptIssueSummary[]
|
||||
created_at: string
|
||||
}
|
||||
|
||||
export interface PromptEventFilters {
|
||||
decision: string
|
||||
risk_level: string
|
||||
endpoint: string
|
||||
group_id: string
|
||||
user_id: string
|
||||
api_key_id: string
|
||||
request_id: string
|
||||
prompt_hash: string
|
||||
keyword: string
|
||||
start_at: string
|
||||
end_at: string
|
||||
}
|
||||
|
||||
export interface PromptEventPage {
|
||||
items: PromptAuditEvent[]
|
||||
total: number
|
||||
page: number
|
||||
page_size: number
|
||||
pages: number
|
||||
}
|
||||
|
||||
export interface PromptDeleteResult {
|
||||
deleted_events: number
|
||||
deleted_jobs: number
|
||||
}
|
||||
|
||||
export interface PromptDeletePreview {
|
||||
matched_count: number
|
||||
filter_summary: Record<string, unknown>
|
||||
snapshot_max_id: number
|
||||
filter_hash: string
|
||||
confirmation_token: string
|
||||
expires_at: string
|
||||
}
|
||||
|
||||
export interface PromptAuditGroup {
|
||||
id: number
|
||||
name: string
|
||||
status: 'active' | 'inactive'
|
||||
platform: string
|
||||
}
|
||||
|
||||
export interface PromptLoadErrors {
|
||||
config: string
|
||||
runtime: string
|
||||
groups: string
|
||||
events: string
|
||||
}
|
||||
@@ -0,0 +1,139 @@
|
||||
import type {
|
||||
PromptAuditConfig,
|
||||
PromptAuditDraft,
|
||||
PromptAuditEndpointDraft,
|
||||
PromptAuditUpdateRequest,
|
||||
PromptEventFilters,
|
||||
} from './types'
|
||||
|
||||
export const DEFAULT_GUARD_MODEL = 'sileader/qwen3guard:0.6b'
|
||||
|
||||
export const SCANNER_CATALOG = [
|
||||
{ id: 'violent', label: 'Violent' },
|
||||
{ id: 'non_violent_illegal_acts', label: 'Non-violent Illegal Acts' },
|
||||
{ id: 'sexual_content_or_sexual_acts', label: 'Sexual Content or Sexual Acts' },
|
||||
{ id: 'pii', label: 'PII' },
|
||||
{ id: 'suicide_and_self_harm', label: 'Suicide & Self-Harm' },
|
||||
{ id: 'unethical_acts', label: 'Unethical Acts' },
|
||||
{ id: 'politically_sensitive_topics', label: 'Politically Sensitive Topics' },
|
||||
{ id: 'copyright_violation', label: 'Copyright Violation' },
|
||||
{ id: 'jailbreak', label: 'Jailbreak' },
|
||||
] as const
|
||||
|
||||
// Vue props/refs are proxies and cannot be passed to structuredClone in every
|
||||
// browser. Prompt Audit state is JSON-only, so this produces a detached draft
|
||||
// without retaining reactive proxies or browser storage references.
|
||||
export function cloneData<T>(value: T): T {
|
||||
return JSON.parse(JSON.stringify(value)) as T
|
||||
}
|
||||
|
||||
export function configToDraft(config: PromptAuditConfig): PromptAuditDraft {
|
||||
return {
|
||||
...cloneData(config),
|
||||
group_ids: [...(config.group_ids ?? [])],
|
||||
scanners: [...(config.scanners ?? [])],
|
||||
endpoints: (config.endpoints ?? []).map((endpoint) => ({
|
||||
...endpoint,
|
||||
token: '',
|
||||
clear_token: false,
|
||||
})),
|
||||
}
|
||||
}
|
||||
|
||||
export function createDefaultEndpoint(index = 1): PromptAuditEndpointDraft {
|
||||
return {
|
||||
id: `guard-${Date.now()}-${index}`,
|
||||
name: `Guard ${index}`,
|
||||
protocol: 'openai_compatible',
|
||||
base_url: 'http://127.0.0.1:8000',
|
||||
model: DEFAULT_GUARD_MODEL,
|
||||
timeout_ms: 3000,
|
||||
input_limit: 4000,
|
||||
enabled: true,
|
||||
has_token: false,
|
||||
token_status: 'missing',
|
||||
token: '',
|
||||
clear_token: false,
|
||||
}
|
||||
}
|
||||
|
||||
export function buildUpdateRequest(draft: PromptAuditDraft): PromptAuditUpdateRequest {
|
||||
return {
|
||||
expected_config_version: draft.config_version,
|
||||
enabled: draft.enabled,
|
||||
blocking_enabled: draft.enabled && draft.blocking_enabled,
|
||||
store_pass_events: draft.store_pass_events,
|
||||
strategy: 'priority',
|
||||
worker_count: Number(draft.worker_count),
|
||||
queue_capacity: Number(draft.queue_capacity),
|
||||
scanners: [...draft.scanners],
|
||||
all_groups: draft.all_groups,
|
||||
group_ids: draft.all_groups ? [] : [...draft.group_ids].sort((a, b) => a - b),
|
||||
endpoints: draft.endpoints.map((endpoint) => ({
|
||||
id: endpoint.id.trim(),
|
||||
name: endpoint.name.trim(),
|
||||
protocol: 'openai_compatible',
|
||||
base_url: endpoint.base_url.trim(),
|
||||
model: endpoint.model.trim() || DEFAULT_GUARD_MODEL,
|
||||
token: endpoint.token.trim() || undefined,
|
||||
clear_token: endpoint.clear_token,
|
||||
timeout_ms: Number(endpoint.timeout_ms),
|
||||
input_limit: Number(endpoint.input_limit),
|
||||
enabled: endpoint.enabled,
|
||||
})),
|
||||
}
|
||||
}
|
||||
|
||||
export function draftFingerprint(draft: PromptAuditDraft | null): string {
|
||||
if (!draft) return ''
|
||||
return JSON.stringify(buildUpdateRequest(draft))
|
||||
}
|
||||
|
||||
export function emptyEventFilters(): PromptEventFilters {
|
||||
return {
|
||||
decision: '',
|
||||
risk_level: '',
|
||||
endpoint: '',
|
||||
group_id: '',
|
||||
user_id: '',
|
||||
api_key_id: '',
|
||||
request_id: '',
|
||||
prompt_hash: '',
|
||||
keyword: '',
|
||||
start_at: '',
|
||||
end_at: '',
|
||||
}
|
||||
}
|
||||
|
||||
function toISO(value: string): string | undefined {
|
||||
if (!value.trim()) return undefined
|
||||
const date = new Date(value)
|
||||
return Number.isNaN(date.getTime()) ? undefined : date.toISOString()
|
||||
}
|
||||
|
||||
export function eventQueryParams(filters: PromptEventFilters): Record<string, string | number> {
|
||||
const result: Record<string, string | number> = {}
|
||||
for (const key of ['decision', 'risk_level', 'endpoint', 'request_id', 'prompt_hash', 'keyword'] as const) {
|
||||
const value = filters[key].trim()
|
||||
if (value) result[key] = value
|
||||
}
|
||||
for (const key of ['group_id', 'user_id', 'api_key_id'] as const) {
|
||||
const value = Number(filters[key])
|
||||
if (Number.isInteger(value) && value > 0) result[key] = value
|
||||
}
|
||||
const start = toISO(filters.start_at)
|
||||
const end = toISO(filters.end_at)
|
||||
if (start) result.start_at = start
|
||||
if (end) result.end_at = end
|
||||
return result
|
||||
}
|
||||
|
||||
export function eventFilterPayload(filters: PromptEventFilters): Record<string, unknown> {
|
||||
return eventQueryParams(filters)
|
||||
}
|
||||
|
||||
export function hasExplicitDeleteRange(filters: PromptEventFilters): boolean {
|
||||
const start = toISO(filters.start_at)
|
||||
const end = toISO(filters.end_at)
|
||||
return Boolean(start && end && new Date(start).getTime() < new Date(end).getTime())
|
||||
}
|
||||
@@ -5,6 +5,7 @@ import resources from './resources'
|
||||
import ops from './ops'
|
||||
import settings from './settings'
|
||||
import audit from './audit'
|
||||
import promptAudit from './promptAudit'
|
||||
|
||||
export default {
|
||||
...overview,
|
||||
@@ -14,4 +15,5 @@ export default {
|
||||
...ops,
|
||||
...settings,
|
||||
...audit,
|
||||
...promptAudit,
|
||||
}
|
||||
|
||||
@@ -0,0 +1,53 @@
|
||||
export default {
|
||||
promptAudit: {
|
||||
title: 'Prompt Audit',
|
||||
description: 'Review user input asynchronously or block it synchronously through OpenAI-compatible Qwen3Guard nodes. Full prompts never enter the database or UI.',
|
||||
configVersion: 'Config version v{version}',
|
||||
actions: { refresh: 'Refresh runtime', retry: 'Retry' },
|
||||
common: { actions: 'Actions', never: 'Never' },
|
||||
mode: { off: 'Off', async_audit: 'Async audit only', blocking: 'Synchronous audit and block' },
|
||||
status: { disabled: 'Disabled', running: 'Running', degraded: 'Degraded', error: 'Error', healthy: 'Healthy', failed: 'Failed', stale: 'Stale heartbeat' },
|
||||
runtime: {
|
||||
title: 'Runtime overview',
|
||||
description: 'Shows the configuration currently active on the server. Unsaved draft changes do not affect these values.',
|
||||
process: 'Process status', mode: 'Effective mode', version: 'Active / expected version', workers: 'Active / total workers',
|
||||
queue: 'Active jobs / capacity', dependencies: 'Dependencies', guardMetrics: 'Synchronous Guard metrics', latest: 'Latest processing and error',
|
||||
queueBreakdown: 'queued {queued} · processing {processing} · retry {retry} · done {done} · failed {failed}',
|
||||
deliveryTotals: 'Total enqueued {enqueued} · dropped {dropped} · processed {processed} · failed {failed}',
|
||||
},
|
||||
metrics: { total: 'Total', allowed: 'Allowed', flagged: 'Flagged', blocked: 'Blocked', unavailable: 'Unavailable', timeouts: 'Timeouts', failovers: 'Failovers' },
|
||||
pool: {
|
||||
title: 'Audit pool', description: 'Enabled OpenAI-compatible nodes are tried in order. Probes run from the server network.',
|
||||
add: 'Add node', edit: 'Edit node', empty: 'No audit nodes configured.', node: 'Node', model: 'Model', limits: 'Timeout / chunk limit', credential: 'Credential and probe',
|
||||
configured: 'API Key configured', missing: 'API Key missing', probe: 'Test connection', probing: 'Probing…',
|
||||
probeProgress: 'Config validated ✓ · request sent · awaiting service response…', probeResult: 'Config ✓ · request ✓ · HTTP {http} · {status} · {latency} ms',
|
||||
name: 'Node name', id: 'Stable node ID', baseUrl: 'Base URL', apiKey: 'API Key', keepSecret: 'Leave blank to keep the saved API Key',
|
||||
secretHint: 'Plaintext exists only in this editor and is cleared immediately after a successful save.', clearSecret: 'Explicitly clear the saved API Key', timeout: 'Total timeout (ms)', inputLimit: 'Unicode characters per chunk',
|
||||
toggleNode: 'Toggle node {name}', deleteConfirm: 'Remove “{name}” from the draft? It takes effect after saving.',
|
||||
},
|
||||
policy: {
|
||||
title: 'Audit policy', description: 'Configure group scope, nine input-risk categories, workers, and queue bounds.', scope: 'Scope', allGroups: 'All groups', selectedGroups: 'Selected groups',
|
||||
searchGroups: 'Search groups', noGroups: 'No matching groups', missingGroups: 'Configured IDs for groups that no longer exist', selectedCount: '{count} groups selected',
|
||||
scanners: 'Qwen3Guard input-risk categories', workerCount: 'Worker count', queueCapacity: 'Persistent queue capacity', strategy: 'Node strategy', strategyHint: 'Try nodes in configuration order and fail over when allowed.',
|
||||
},
|
||||
saveBar: { enabled: 'Enable prompt audit', blocking: 'Synchronous blocking', storePass: 'Store Pass events', dirty: 'Unsaved changes', synced: 'Configuration synced' },
|
||||
blockingConfirm: {
|
||||
title: 'Enable synchronous blocking?',
|
||||
message: 'Applicable requests wait for Guard before account selection, billing, or upstream access. Block, unavailable Guard, and invalid responses all prevent upstream access.',
|
||||
confirm: 'I understand; enable it',
|
||||
},
|
||||
events: {
|
||||
title: 'Audit events', description: 'Review redacted events by identity, route, risk, hash, and time.', decision: 'Decision', risk: 'Risk level', endpoint: 'Endpoint', groupId: 'Group ID', userId: 'User ID', apiKeyId: 'API Key ID', keyword: 'Keyword',
|
||||
startAt: 'Start time', endAt: 'End time', deleteSelected: 'Delete selected ({count})', deleteByFilter: 'Delete by filter', deleteRangeHint: 'Filter deletion requires explicit start and end times and a server-generated preview.',
|
||||
selectAll: 'Select all events on this page', selectEvent: 'Select event {id}', time: 'Time', identity: 'User / email / API Key', user: 'Username', email: 'User email', apiKey: 'API Key name', group: 'Group', route: 'Endpoint / model', result: 'Decision / risk', preview: 'Redacted preview', empty: 'No matching events.',
|
||||
detailTitle: 'Prompt audit event details', tabs: { summary: 'Audit summary', risks: 'Specific risks', technical: 'Technical details' }, redactedPreview: 'Irreversible redacted preview', categories: 'Categories', model: 'Model', noRisks: 'No derived risk summaries for this event.',
|
||||
deleteConfirmTitle: 'Delete audit events?', deleteConfirmMessage: 'This permanently deletes {count} events and eligible orphan jobs.', filterDeleteTitle: 'Confirm filter deletion', filterDeleteCount: 'The server snapshot matches {count} events.', snapshotMax: 'Snapshot maximum event ID', expiresAt: 'Confirmation token expires', filterDeleteWarning: 'Only events at or below the preview high-water mark are deleted. Newer events survive. Any filter change requires a new preview.', confirmFilterDelete: 'Permanently delete',
|
||||
},
|
||||
messages: { saved: 'Prompt Audit configuration saved; plaintext API Key state was cleared.', probeSucceeded: 'The audit node is reachable.', deleted: 'Deleted {count} audit events.' },
|
||||
errors: {
|
||||
loadConfig: 'Unable to load Prompt Audit configuration.', loadRuntime: 'Unable to load Prompt Audit runtime.', loadGroups: 'Unable to load groups.', loadEvents: 'Unable to load audit events.', loadDetail: 'Unable to load event details.', saveConfig: 'Unable to save the configuration.', probe: 'Node probe failed.', delete: 'Unable to delete events.', previewDelete: 'Unable to create a deletion preview. Check the time range.', deleteConfirmation: 'The deletion confirmation is invalid or expired. Preview again.',
|
||||
prompt_audit_config_conflict: 'Another administrator updated this configuration. Reload the server version before deciding how to merge your draft.',
|
||||
prompt_guard_requires_audit_enabled: 'Enable Prompt Audit before synchronous blocking.', prompt_audit_invalid_endpoint: 'The audit node configuration is invalid.', prompt_audit_endpoint_required: 'Enable at least one audit node before enabling Prompt Audit.', prompt_audit_groups_required: 'Select at least one group in selected-group mode.', prompt_audit_scanners_required: 'Enable at least one risk category.',
|
||||
},
|
||||
},
|
||||
}
|
||||
@@ -190,6 +190,9 @@ export default {
|
||||
channelMonitor: 'Channel Monitor',
|
||||
channelStatus: 'Channel Status',
|
||||
riskControl: 'Risk Control',
|
||||
securityAudit: 'Security Audit',
|
||||
contentModeration: 'Content Moderation',
|
||||
promptAudit: 'Prompt Audit',
|
||||
auditLogs: 'Audit Logs',
|
||||
},
|
||||
|
||||
|
||||
@@ -5,6 +5,7 @@ import resources from './resources'
|
||||
import ops from './ops'
|
||||
import settings from './settings'
|
||||
import audit from './audit'
|
||||
import promptAudit from './promptAudit'
|
||||
|
||||
export default {
|
||||
...overview,
|
||||
@@ -14,4 +15,5 @@ export default {
|
||||
...ops,
|
||||
...settings,
|
||||
...audit,
|
||||
...promptAudit,
|
||||
}
|
||||
|
||||
@@ -0,0 +1,53 @@
|
||||
export default {
|
||||
promptAudit: {
|
||||
title: '提示词审计',
|
||||
description: '通过 OpenAI 兼容 Qwen3Guard 节点异步复核或同步阻止用户输入;完整提示词不会进入数据库或页面。',
|
||||
configVersion: '配置版本 v{version}',
|
||||
actions: { refresh: '刷新运行态', retry: '重试' },
|
||||
common: { actions: '操作', never: '从未' },
|
||||
mode: { off: '已关闭', async_audit: '异步只审计', blocking: '同步审计并阻止' },
|
||||
status: { disabled: '未启用', running: '运行中', degraded: '降级', error: '错误', healthy: '健康', failed: '失败', stale: '心跳过期' },
|
||||
runtime: {
|
||||
title: '运行概览',
|
||||
description: '显示服务端当前生效状态;未保存的草稿不会改变这些数值。',
|
||||
process: '进程状态', mode: '生效模式', version: '生效 / 期望版本', workers: '活动 / 总 Worker',
|
||||
queue: '活动任务 / 容量', dependencies: '依赖', guardMetrics: '同步 Guard 指标', latest: '最近处理与错误',
|
||||
queueBreakdown: 'queued {queued} · processing {processing} · retry {retry} · done {done} · failed {failed}',
|
||||
deliveryTotals: '累计入队 {enqueued} · 丢弃 {dropped} · 处理 {processed} · 失败 {failed}',
|
||||
},
|
||||
metrics: { total: '总计', allowed: '放行', flagged: '标记', blocked: '阻止', unavailable: '不可用', timeouts: '超时', failovers: '故障切换' },
|
||||
pool: {
|
||||
title: '审计池', description: '按顺序使用启用的 OpenAI 兼容节点;探测由服务端真实网络环境发起。',
|
||||
add: '新增节点', edit: '编辑节点', empty: '尚未配置审计节点。', node: '节点', model: '模型', limits: '超时 / 单片上限', credential: '凭据与探测',
|
||||
configured: 'API Key 已配置', missing: '未配置 API Key', probe: '连接测试', probing: '探测中…',
|
||||
probeProgress: '配置校验 ✓ · 请求已发送 · 等待服务响应…', probeResult: '配置校验 ✓ · 请求 ✓ · HTTP {http} · {status} · {latency} ms',
|
||||
name: '节点名称', id: '稳定节点 ID', baseUrl: 'Base URL', apiKey: 'API Key', keepSecret: '留空以保留已保存的 API Key',
|
||||
secretHint: '明文只在本次编辑内存中存在;保存成功后会立即清除。', clearSecret: '显式清除已保存的 API Key', timeout: '总超时(毫秒)', inputLimit: '单片 Unicode 字符上限',
|
||||
toggleNode: '切换节点 {name}', deleteConfirm: '从草稿中删除节点“{name}”?保存配置后生效。',
|
||||
},
|
||||
policy: {
|
||||
title: '审计策略', description: '配置适用分组、九类输入风险、Worker 与队列边界。', scope: '适用范围', allGroups: '全部分组', selectedGroups: '指定分组',
|
||||
searchGroups: '搜索分组', noGroups: '没有匹配分组', missingGroups: '配置中包含已删除的分组 ID', selectedCount: '已选择 {count} 个分组',
|
||||
scanners: 'Qwen3Guard 输入风险分类', workerCount: 'Worker 数量', queueCapacity: '持久队列容量', strategy: '节点策略', strategyHint: '按配置顺序优先尝试,必要时故障切换。',
|
||||
},
|
||||
saveBar: { enabled: '启用提示词审计', blocking: '同步阻止', storePass: '保存 Pass 事件', dirty: '有未保存的更改', synced: '配置已同步' },
|
||||
blockingConfirm: {
|
||||
title: '开启同步阻止?',
|
||||
message: '适用请求会在账号选择、计费和访问上游之前等待 Guard。命中 Block、Guard 不可用或响应非法时,请求都不会访问上游。',
|
||||
confirm: '理解风险并开启',
|
||||
},
|
||||
events: {
|
||||
title: '审计事件', description: '按身份、入口、风险、Hash 和时间复核脱敏事件。', decision: '判定', risk: '风险等级', endpoint: '入口', groupId: '分组 ID', userId: '用户 ID', apiKeyId: 'API Key ID', keyword: '关键词',
|
||||
startAt: '开始时间', endAt: '结束时间', deleteSelected: '删除选中项({count})', deleteByFilter: '按筛选删除', deleteRangeHint: '按筛选删除必须明确选择开始和结束时间,并先取得服务端删除预览。',
|
||||
selectAll: '选择当前页全部事件', selectEvent: '选择事件 {id}', time: '时间', identity: '用户 / 邮箱 / API Key', user: '用户名', email: '用户邮箱', apiKey: 'API Key 名称', group: '分组', route: '入口 / 模型', result: '判定 / 风险', preview: '脱敏预览', empty: '没有符合条件的事件。',
|
||||
detailTitle: '提示词审计事件详情', tabs: { summary: '审计摘要', risks: '具体风险', technical: '技术信息' }, redactedPreview: '不可逆脱敏预览', categories: '分类', model: '模型', noRisks: '本事件没有派生风险摘要。',
|
||||
deleteConfirmTitle: '删除审计事件?', deleteConfirmMessage: '将永久删除 {count} 条事件及符合条件的孤立任务。', filterDeleteTitle: '确认按筛选删除', filterDeleteCount: '服务端快照匹配 {count} 条事件。', snapshotMax: '快照最大事件 ID', expiresAt: '确认令牌过期时间', filterDeleteWarning: '只删除预览高水位内的事件;预览后产生的新事件会保留。筛选一旦变化,必须重新预览。', confirmFilterDelete: '确认永久删除',
|
||||
},
|
||||
messages: { saved: '提示词审计配置已保存,明文 API Key 状态已清除。', probeSucceeded: '审计节点连接正常。', deleted: '已删除 {count} 条审计事件。' },
|
||||
errors: {
|
||||
loadConfig: '无法加载提示词审计配置。', loadRuntime: '无法加载提示词审计运行态。', loadGroups: '无法加载分组列表。', loadEvents: '无法加载审计事件。', loadDetail: '无法加载事件详情。', saveConfig: '配置保存失败。', probe: '节点探测失败。', delete: '事件删除失败。', previewDelete: '无法生成删除预览,请检查时间范围。', deleteConfirmation: '删除确认无效或已过期,请重新预览。',
|
||||
prompt_audit_config_conflict: '配置已被其他管理员更新。请重新加载服务端配置,再决定如何合并本地草稿。',
|
||||
prompt_guard_requires_audit_enabled: '开启同步阻止前必须先启用提示词审计。', prompt_audit_invalid_endpoint: '审计节点配置无效。', prompt_audit_endpoint_required: '启用审计前至少需要一个启用节点。', prompt_audit_groups_required: '指定分组模式至少需要选择一个分组。', prompt_audit_scanners_required: '至少需要启用一个风险分类。',
|
||||
},
|
||||
},
|
||||
}
|
||||
@@ -190,6 +190,9 @@ export default {
|
||||
channelMonitor: '渠道监控',
|
||||
channelStatus: '渠道状态',
|
||||
riskControl: '风控中心',
|
||||
securityAudit: '安全审计',
|
||||
contentModeration: '内容审核',
|
||||
promptAudit: '提示词审计',
|
||||
auditLogs: '操作日志',
|
||||
},
|
||||
|
||||
|
||||
@@ -587,6 +587,19 @@ const routes: RouteRecordRaw[] = [
|
||||
requiresRiskControl: true
|
||||
}
|
||||
},
|
||||
{
|
||||
path: '/admin/prompt-audit',
|
||||
name: 'AdminPromptAudit',
|
||||
component: () => import('@/features/prompt-audit/PromptAuditView.vue'),
|
||||
meta: {
|
||||
requiresAuth: true,
|
||||
requiresAdmin: true,
|
||||
title: 'Prompt Audit',
|
||||
titleKey: 'admin.promptAudit.title',
|
||||
descriptionKey: 'admin.promptAudit.description',
|
||||
requiresRiskControl: true
|
||||
}
|
||||
},
|
||||
{
|
||||
path: '/admin/usage',
|
||||
name: 'AdminUsage',
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
schema: spec-driven
|
||||
created: 2026-07-16
|
||||
@@ -0,0 +1,5 @@
|
||||
# add-openai-compatible-prompt-audit
|
||||
|
||||
在不改变现有内容审核行为的前提下,新增独立的 OpenAI 兼容 Qwen3Guard 提示词安全审计模块,完整支持异步审计、同步阻断、持久任务队列、事件工作台与独立管理页面。
|
||||
|
||||
阅读顺序:`proposal.md` → `source-baseline.md` / `source-feature-map.md` → `design.md` → 三个 `specs/*/spec.md` → `implementation-guide.md` → `tasks.md` → `verification.md`。
|
||||
@@ -0,0 +1,754 @@
|
||||
## Context
|
||||
|
||||
### 当前系统
|
||||
|
||||
sub2api 当前已经存在一套完整的内容审核能力:
|
||||
|
||||
- 核心实现位于 `backend/internal/service/content_moderation*.go`。
|
||||
- 管理 API 位于 `backend/internal/handler/admin/content_moderation_handler.go`,路由前缀为 `/admin/risk-control`。
|
||||
- 网关统一接线位于 `backend/internal/handler/content_moderation_helper.go`,各协议 Handler 在解析完请求体和模型后调用 `checkContentModeration`。
|
||||
- 数据保存在 `content_moderation_logs`,配置保存在 settings 的 `content_moderation_config`。
|
||||
- 管理页面为 `frontend/src/views/admin/RiskControlView.vue`。
|
||||
- 能力包括 OpenAI Moderations、关键词阻断、命中 Hash、异步观察、同步前置阻断、API Key 健康、邮件、违规计数和自动封号。
|
||||
|
||||
该能力不是本次要迁移的 aicodex-api “提示词审计”:两者使用不同模型、分类、队列、事件和阻断语义。把 Qwen3Guard 直接塞入 ContentModerationService 会让现有阈值、封号统计和记录含义失真,也会继续扩大已经接近 3000 行的单文件。
|
||||
|
||||
### 参考能力
|
||||
|
||||
参考仓库 `/Users/mt/code/mt-ai/aicodex/aicodex-api` 当前磁盘实现提供:
|
||||
|
||||
- OpenAI 兼容 Qwen3Guard 审计池。
|
||||
- 持久 PromptAuditJob / PromptAuditEvent。
|
||||
- Redis 30 分钟临时原文载荷。
|
||||
- 进程内 Worker、重试、租约和滞留回收。
|
||||
- 脱敏快照、Hash、Unicode 分片、最新输入优先。
|
||||
- 九类风险和严格 `Safety/Categories` 解析。
|
||||
- 异步审计与同步 fail-closed 阻断。
|
||||
- HTTP、SSE、Responses WebSocket 错误映射。
|
||||
- 节点探测、运行态、事件筛选/详情/硬删除和独立控制台页面。
|
||||
|
||||
参考仓库 `yjb` 分支当前包含未提交的同步阻止改动。因此实施开始前必须固定源 commit/tag 或生成包含未提交文件的只读 patch 清单,作为功能对照和测试移植的权威基线。
|
||||
|
||||
### 目标项目约束
|
||||
|
||||
- PostgreSQL SQL migrations 是 schema 的事实源,Ent 自动迁移不是生产建表入口。
|
||||
- 后端是 Go + Gin + Wire;前端是 Vue 3 + TypeScript + pnpm。
|
||||
- Redis 已是运行基础设施,可作为短 TTL 敏感载荷存储和配置失效通知通道。
|
||||
- 新模块必须尽量集中在独立目录,并只通过显式接口接入现有 Handler。
|
||||
- 新功能默认关闭,不能改变升级前行为。
|
||||
- 完整提示词和 Guard 凭据不能进入数据库、日志、API、前端或错误响应。
|
||||
|
||||
### 参与边界
|
||||
|
||||
- 网关请求处理:提供可信身份上下文、协议、模型和原始请求体。
|
||||
- 安全审计协调器:调用两个独立引擎并归并阻断结果。
|
||||
- 现有内容审核:保持原实现和副作用。
|
||||
- 新 Prompt Audit 模块:负责配置、提取、队列、Guard、事件、运行态和管理 API。
|
||||
- PostgreSQL:持久任务与事件。
|
||||
- Redis:扫描正文 TTL、配置失效通知、可选跨实例心跳/指标汇总。
|
||||
- 控制台:独立提示词审计页面。
|
||||
|
||||
## Goals / Non-Goals
|
||||
|
||||
**Goals:**
|
||||
|
||||
- 在不改变现有内容审核语义的前提下完整引入提示词输入审计。
|
||||
- 使用模块化垂直目录封装新能力,限制对现有代码的修改面。
|
||||
- 保持所有现有 OpenAI/Claude/Gemini/媒体兼容入口的请求和响应 envelope。
|
||||
- 提供异步不阻塞和同步 fail-closed 两种模式。
|
||||
- 在同步 Block/Unavailable 时保证无账号、无计费、无上游副作用。
|
||||
- 支持多实例持久任务消费和配置最终一致。
|
||||
- 只持久化脱敏、可关联、可复核的数据。
|
||||
- 把运行态、日志、指标和测试设计为第一等反馈信号。
|
||||
- 提供完整、独立、可访问的管理页面。
|
||||
|
||||
**Non-Goals:**
|
||||
|
||||
- 不审核模型输出,不在流式输出中途截断。
|
||||
- 不实现请求正文 Redact 或自动改写。
|
||||
- 不实现人工审批、申诉、逐请求放行或策略工作流。
|
||||
- 不把 Qwen3Guard 分类映射为现有 OpenAI Moderations 分数。
|
||||
- 不让提示词审计命中触发自动封号、邮件或 Hash 黑名单。
|
||||
- 不删除、合并或迁移 `content_moderation_logs`。
|
||||
- 不新增目标项目不存在的 AICodex 专属产品路由;只对目标项目实际存在的文本入口提供等价覆盖。
|
||||
- 不在本 change 中重构整个 Handler、计费或账号调度架构。
|
||||
|
||||
## Decisions
|
||||
|
||||
### 1. 迁移行为契约,而不是直接复制源目录
|
||||
|
||||
源模块依赖 aicodex-api 的 Ent 全局客户端、option 模型、Gin context key、日志封装、Caddy/gatewaycore 和 React 控制台,不能原样复制到目标项目。
|
||||
|
||||
实施时以本 change 的 specs 和验收矩阵作为权威行为契约,再选择目标项目已有的 SettingRepository、Redis、SecretEncryptor、Gin Handler、SQL migration 和 Vue 组件实现。
|
||||
|
||||
**备选方案:直接复制 `internal/service/promptaudit`。** 放弃,因为会引入大量适配壳、全局状态和源仓库私有依赖,并且源工作区当前未提交。
|
||||
|
||||
### 2. 使用模块化垂直目录承载新能力
|
||||
|
||||
新增目录:
|
||||
|
||||
```text
|
||||
backend/internal/securityaudit/
|
||||
├── coordinator.go
|
||||
├── prompt_config.go
|
||||
├── prompt_types.go
|
||||
├── prompt_snapshot.go
|
||||
├── prompt_scanner.go
|
||||
├── prompt_qwen3guard.go
|
||||
├── prompt_outbound_security.go
|
||||
├── prompt_repository.go
|
||||
├── prompt_payload_store.go
|
||||
├── prompt_enqueue.go
|
||||
├── prompt_worker.go
|
||||
├── prompt_guard.go
|
||||
├── prompt_runtime.go
|
||||
├── prompt_handler.go
|
||||
├── prompt_logging.go
|
||||
├── prompt_module.go
|
||||
└── *_test.go
|
||||
```
|
||||
|
||||
该目录内部允许用文件划分子职责,但对外只暴露:
|
||||
|
||||
- `Coordinator.Check(ctx, Request) Decision`
|
||||
- `PromptService` 生命周期与管理方法
|
||||
- `PromptAdminHandler`
|
||||
- Wire provider set
|
||||
|
||||
SQL migration、前端和少量路由/注入接线由于项目结构约束仍位于各自事实源目录。
|
||||
|
||||
**备选方案:继续平铺在 `internal/service`、`internal/repository` 和 `internal/handler`。** 放弃,因为无法满足独立模块要求,也会增加 AI 和人工定位所需上下文。
|
||||
|
||||
### 3. 使用薄协调器组合两个引擎
|
||||
|
||||
目标调用关系:
|
||||
|
||||
```mermaid
|
||||
flowchart LR
|
||||
H[Protocol Handler] --> C[SecurityAudit Coordinator]
|
||||
C --> M[Existing ContentModerationService]
|
||||
C --> P[PromptAuditService]
|
||||
M --> MD[Moderation Decision]
|
||||
P --> PD[Prompt Decision]
|
||||
MD --> C
|
||||
PD --> C
|
||||
C --> D[Normalized gateway decision]
|
||||
```
|
||||
|
||||
Coordinator 只承担:
|
||||
|
||||
1. 接收可信身份和请求快照。
|
||||
2. 确保新异步任务即使现有引擎随后阻断也能 best-effort 投递。
|
||||
3. 在新同步模式下执行两个引擎并等待结果。
|
||||
4. 使用固定优先级生成客户端决策。
|
||||
|
||||
优先级:
|
||||
|
||||
1. 现有内容审核 Block:保留原状态、错误码和文案。
|
||||
2. Prompt Guard Block:403 + `prompt_guard_blocked`。
|
||||
3. Prompt Guard Invalid:503 + `prompt_guard_invalid_response`。
|
||||
4. Prompt Guard Unavailable:503 + `prompt_guard_unavailable`。
|
||||
5. 否则 Allow。
|
||||
|
||||
两个引擎的事件和副作用独立。Coordinator 不持久化业务事件,不修改风险分数。
|
||||
|
||||
**同步执行策略:** 当 Prompt Guard blocking 开启时,现有内容审核和 Prompt Guard 可在独立受控 goroutine 中并行执行,共享请求取消信号但不共享 mutable state。必须等待两者完成或各自 deadline 到期,以保留两个引擎的审计完整性。若实现评审认为并行引入的复杂度过高,可先串行执行,但仍必须满足既有 Block 响应优先级和无下游副作用测试。
|
||||
|
||||
### 4. 复用现有接入位置,但显式改名为安全审计
|
||||
|
||||
把各协议 Handler 的 `checkContentModeration` 调用机械替换为 `checkSecurityAudit`,保持调用点仍在:
|
||||
|
||||
- 身份鉴权、基本请求体读取和协议格式校验之后。
|
||||
- 账号选择、账户并发、计费资格、预扣、上游拨号/写入之前。
|
||||
|
||||
现有 `content_moderation_helper.go` 改为或新增 `security_audit_helper.go`,构造统一 `securityaudit.Request`:
|
||||
|
||||
```go
|
||||
type Request struct {
|
||||
RequestID string
|
||||
UserID int64
|
||||
Username string
|
||||
UserEmail string
|
||||
APIKeyID int64
|
||||
APIKeyName string
|
||||
GroupID *int64
|
||||
GroupName string
|
||||
Provider string
|
||||
Endpoint string
|
||||
Protocol string
|
||||
Model string
|
||||
Body []byte
|
||||
Stage string // http, first_turn, subsequent_turn
|
||||
}
|
||||
```
|
||||
|
||||
请求体必须在 Handler 已受全局大小限制后传入。模块不得再次从 `http.Request.Body` 读取,避免破坏转发。
|
||||
|
||||
### 5. 保持三个独立开关层级
|
||||
|
||||
有效开关:
|
||||
|
||||
1. `risk_control_enabled`:现有安全审计总入口和菜单开关。
|
||||
2. `content_moderation_config.enabled/mode`:现有内容审核。
|
||||
3. `prompt_audit_config.enabled/blocking_enabled`:新提示词审计。
|
||||
|
||||
Prompt Audit 有效模式:
|
||||
|
||||
| risk_control | enabled | blocking_enabled | 有效行为 |
|
||||
| --- | --- | --- | --- |
|
||||
| false | 任意 | 任意 | off |
|
||||
| true | false | false | off |
|
||||
| true | true | false | async_audit |
|
||||
| true | true | true | blocking |
|
||||
|
||||
后端必须拒绝 `enabled=false && blocking_enabled=true`。前端联动只提升体验,不能替代后端校验。
|
||||
|
||||
### 6. 配置使用 settings JSON,但凭据独立加密
|
||||
|
||||
新增 setting key:`prompt_audit_config`。
|
||||
|
||||
配置结构包含:
|
||||
|
||||
```text
|
||||
enabled
|
||||
blocking_enabled
|
||||
store_pass_events
|
||||
strategy=priority
|
||||
worker_count
|
||||
queue_capacity
|
||||
scanners[]
|
||||
all_groups
|
||||
group_ids[]
|
||||
config_version
|
||||
updated_at
|
||||
updated_by
|
||||
change_summary
|
||||
endpoints[]
|
||||
```
|
||||
|
||||
每个 endpoint 持久化:
|
||||
|
||||
```text
|
||||
id, name, protocol=openai_compatible, base_url, model,
|
||||
token_ciphertext, timeout_ms, input_limit, enabled
|
||||
```
|
||||
|
||||
读取 API 只返回 `has_token`/`token_status`。保存请求使用:
|
||||
|
||||
- `token` 非空:替换并加密。
|
||||
- `token` 空且 `clear_token=false`:保留已有密文。
|
||||
- `clear_token=true`:删除密文。
|
||||
|
||||
config_version 每次成功保存单调加一。change_summary 只保存节点数量、开关、分类数量、分组数量及其 Hash 等脱敏摘要。
|
||||
|
||||
保存请求必须携带管理员读取草稿时的 `expected_config_version`。ConfigStore 在 PostgreSQL 短事务中获取 `prompt_audit_config` 专用 advisory transaction lock,重新读取 settings 当前值并比较版本;不一致时返回 409 `prompt_audit_config_conflict`,不得覆盖其他管理员的新配置。版本一致时才计算 current+1、加密并写回。首次无 setting 时按 version=1/default-off 参与比较。进程内 mutex 不能代替该多实例 CAS。
|
||||
|
||||
**备选方案:新增配置表。** 第一版放弃,因为目标项目已有 settings 配置模式,源实现也使用 option JSON;任务和事件才需要独立关系表。
|
||||
|
||||
### 7. 配置使用内存快照和 Redis 失效通知
|
||||
|
||||
PromptService 维护原子只读配置快照:
|
||||
|
||||
- 启动时加载并校验。
|
||||
- 保存成功后先安装本实例快照,再 publish `sub2api:prompt_guard:config:invalidate`,消息只包含版本。
|
||||
- 其他实例收到通知后重新从 settings 加载、解密、校验并原子替换。
|
||||
- Redis publish 失败时保留最后有效配置,并通过 5 秒有界 TTL 后台刷新。
|
||||
- 请求热路径只读取快照,不查询数据库。
|
||||
|
||||
运行态返回 expected 和 active version。配置加载失败不得清空最后有效快照;冷启动无有效快照时不得伪装为关闭或健康。
|
||||
|
||||
### 8. 使用提示词专用快照提取器,不直接复用现有截断结果
|
||||
|
||||
复用现有内容审核提供的 protocol 常量、身份/分组上下文和部分 JSON 内容块解析思路,但新模块实现独立 `PromptSnapshotExtractor`:
|
||||
|
||||
- Chat Completions:只提取 role=user 的文本内容。
|
||||
- Responses:支持 input 字符串、消息数组和 content blocks。
|
||||
- Claude Messages:提取 role=user 文本块。
|
||||
- Gemini:提取 user contents/parts 文本。
|
||||
- Images/媒体:只提取 prompt 文本,忽略图片载荷。
|
||||
- Responses WS:解析每个 response.create 帧。
|
||||
|
||||
扫描顺序:
|
||||
|
||||
1. 最新非空用户输入独立作为首段。
|
||||
2. 其余用户历史保持确定顺序。
|
||||
3. 每段再按 Unicode rune 分片。
|
||||
|
||||
数据库预览使用统一脱敏器:移除/掩码 API Key、Bearer、常见凭据、邮箱/电话等敏感模式,随后按 rune 裁剪。Hash 使用实际待扫描文本的 SHA-256。
|
||||
|
||||
### 9. PostgreSQL 使用两个新表,SQL migration 为事实源
|
||||
|
||||
建议 migration 名称:`backend/migrations/181_prompt_audit.sql`。如果实施时已有 181,则按当前最大序号递增,不允许修改已应用 migration。
|
||||
|
||||
#### `prompt_audit_jobs`
|
||||
|
||||
```sql
|
||||
CREATE TABLE prompt_audit_jobs (
|
||||
id BIGSERIAL PRIMARY KEY,
|
||||
request_id VARCHAR(128) NOT NULL DEFAULT '',
|
||||
user_id BIGINT REFERENCES users(id) ON DELETE SET NULL,
|
||||
username_snapshot VARCHAR(255) NOT NULL DEFAULT '',
|
||||
user_email_snapshot VARCHAR(320) NOT NULL DEFAULT '',
|
||||
api_key_id BIGINT REFERENCES api_keys(id) ON DELETE SET NULL,
|
||||
api_key_name_snapshot VARCHAR(255) NOT NULL DEFAULT '',
|
||||
group_id BIGINT REFERENCES groups(id) ON DELETE SET NULL,
|
||||
group_name VARCHAR(255) NOT NULL DEFAULT '',
|
||||
provider VARCHAR(64) NOT NULL DEFAULT '',
|
||||
endpoint VARCHAR(128) NOT NULL DEFAULT '',
|
||||
protocol VARCHAR(64) NOT NULL DEFAULT '',
|
||||
model VARCHAR(255) NOT NULL DEFAULT '',
|
||||
prompt_hash VARCHAR(64) NOT NULL DEFAULT '',
|
||||
redacted_preview TEXT NOT NULL DEFAULT '',
|
||||
prompt_length INT NOT NULL DEFAULT 0,
|
||||
message_count INT NOT NULL DEFAULT 0,
|
||||
execution_mode VARCHAR(32) NOT NULL DEFAULT 'async_audit',
|
||||
config_version BIGINT NOT NULL DEFAULT 1,
|
||||
status VARCHAR(32) NOT NULL DEFAULT 'staging',
|
||||
attempts INT NOT NULL DEFAULT 0,
|
||||
max_attempts INT NOT NULL DEFAULT 3,
|
||||
claim_version BIGINT NOT NULL DEFAULT 0,
|
||||
next_attempt_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||
processing_started_at TIMESTAMPTZ,
|
||||
processed_at TIMESTAMPTZ,
|
||||
last_error_code VARCHAR(64) NOT NULL DEFAULT '',
|
||||
last_error_message TEXT NOT NULL DEFAULT '',
|
||||
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(),
|
||||
updated_at TIMESTAMPTZ NOT NULL DEFAULT NOW()
|
||||
);
|
||||
```
|
||||
|
||||
状态集合:`staging|queued|processing|retry|done|failed`。
|
||||
|
||||
关键索引:
|
||||
|
||||
```text
|
||||
(status, next_attempt_at, id)
|
||||
(request_id)
|
||||
(user_id, created_at DESC)
|
||||
(api_key_id, created_at DESC)
|
||||
(group_id, created_at DESC)
|
||||
(prompt_hash)
|
||||
(created_at DESC)
|
||||
```
|
||||
|
||||
#### `prompt_audit_events`
|
||||
|
||||
```sql
|
||||
CREATE TABLE prompt_audit_events (
|
||||
id BIGSERIAL PRIMARY KEY,
|
||||
job_id BIGINT NOT NULL REFERENCES prompt_audit_jobs(id) ON DELETE CASCADE,
|
||||
request_id VARCHAR(128) NOT NULL DEFAULT '',
|
||||
user_id BIGINT REFERENCES users(id) ON DELETE SET NULL,
|
||||
username_snapshot VARCHAR(255) NOT NULL DEFAULT '',
|
||||
user_email_snapshot VARCHAR(320) NOT NULL DEFAULT '',
|
||||
api_key_id BIGINT REFERENCES api_keys(id) ON DELETE SET NULL,
|
||||
api_key_name_snapshot VARCHAR(255) NOT NULL DEFAULT '',
|
||||
group_id BIGINT REFERENCES groups(id) ON DELETE SET NULL,
|
||||
group_name VARCHAR(255) NOT NULL DEFAULT '',
|
||||
provider VARCHAR(64) NOT NULL DEFAULT '',
|
||||
endpoint VARCHAR(128) NOT NULL DEFAULT '',
|
||||
protocol VARCHAR(64) NOT NULL DEFAULT '',
|
||||
model VARCHAR(255) NOT NULL DEFAULT '',
|
||||
prompt_hash VARCHAR(64) NOT NULL DEFAULT '',
|
||||
redacted_preview TEXT NOT NULL DEFAULT '',
|
||||
decision VARCHAR(32) NOT NULL DEFAULT 'pass',
|
||||
risk_level VARCHAR(32) NOT NULL DEFAULT 'low',
|
||||
action VARCHAR(32) NOT NULL DEFAULT 'Allow',
|
||||
categories JSONB NOT NULL DEFAULT '[]'::jsonb,
|
||||
matched_scanners JSONB NOT NULL DEFAULT '[]'::jsonb,
|
||||
scanner_scores JSONB NOT NULL DEFAULT '{}'::jsonb,
|
||||
scanner_evidence JSONB NOT NULL DEFAULT '{}'::jsonb,
|
||||
scanner_backend VARCHAR(64) NOT NULL DEFAULT 'qwen3guard-openai',
|
||||
scanner_version VARCHAR(128) NOT NULL DEFAULT '',
|
||||
guard_endpoint_id VARCHAR(128) NOT NULL DEFAULT '',
|
||||
policy_id VARCHAR(128) NOT NULL DEFAULT '',
|
||||
policy_version INT NOT NULL DEFAULT 0,
|
||||
config_version BIGINT NOT NULL DEFAULT 1,
|
||||
chunk_total INT NOT NULL DEFAULT 0,
|
||||
latency_ms INT NOT NULL DEFAULT 0,
|
||||
created_at TIMESTAMPTZ NOT NULL DEFAULT NOW()
|
||||
);
|
||||
```
|
||||
|
||||
事件保留请求快照列用于稳定查询,即使 user/API key/group 后续删除仍保留管理员可复核上下文。用户名、邮箱和 API Key 名称必须作为不同字段返回,不能拼成不可筛选的单一展示串;这些身份快照沿用现有管理员数据访问和保留规则,不得写入普通请求日志。外键使用 SET NULL,快照字段保留。
|
||||
|
||||
事件索引:job、request、decision/time、risk/time、user/time、API key/time、group/time、Hash、created_at。
|
||||
|
||||
不得创建 raw_prompt、raw_request、payload、token 等列。
|
||||
|
||||
### 10. 跨 PostgreSQL/Redis 投递使用 staging 状态避免竞态
|
||||
|
||||
异步投递顺序:
|
||||
|
||||
1. 检查有效模式、范围和节点;在 PostgreSQL 短事务内获取 Prompt Audit 队列准入 advisory transaction lock,重新统计 active jobs,并仅在低于 snapshot queue_capacity 时插入 staging job。
|
||||
2. 提取快照。
|
||||
3. 插入 `status=staging` 的 job。
|
||||
4. `SET sub2api:prompt_audit:payload:<job_id> <scan_text> EX 1800`。
|
||||
5. 条件更新 staging → queued。
|
||||
6. 输出 `prompt_audit.job_enqueued`。
|
||||
|
||||
Worker 只领取 queued/retry,因此不会在 Redis SET 前看到任务。
|
||||
|
||||
失败处理:
|
||||
|
||||
- 步骤 3 失败:不写 Redis。
|
||||
- 步骤 4 失败:job → failed,原请求继续。
|
||||
- 步骤 5 失败:删除 Redis key;job 由 staging 清理器标记 failed。
|
||||
- 进程在 4/5 之间退出:Redis 自动过期,staging 回收器标记 failed。
|
||||
|
||||
这比源实现“先 queued 再写 Redis”更适合多实例,避免 Worker 提前领取。
|
||||
|
||||
队列容量检查和 staging INSERT 必须在同一准入锁事务中完成,防止多个实例先各自看到剩余容量再共同超限。锁等待必须有很短的有界 timeout;无法及时取得锁时按 `queue_admission_busy` 丢弃异步审计任务并让主请求继续。Redis 写入不在该事务内。
|
||||
|
||||
### 11. Worker 使用 PostgreSQL 原子领取与租约
|
||||
|
||||
Repository 使用短事务:
|
||||
|
||||
```sql
|
||||
WITH candidate AS (
|
||||
SELECT id
|
||||
FROM prompt_audit_jobs
|
||||
WHERE status IN ('queued', 'retry')
|
||||
AND next_attempt_at <= NOW()
|
||||
ORDER BY next_attempt_at, id
|
||||
FOR UPDATE SKIP LOCKED
|
||||
LIMIT 1
|
||||
)
|
||||
UPDATE prompt_audit_jobs j
|
||||
SET status = 'processing',
|
||||
attempts = attempts + 1,
|
||||
claim_version = claim_version + 1,
|
||||
processing_started_at = NOW(),
|
||||
updated_at = NOW()
|
||||
FROM candidate
|
||||
WHERE j.id = candidate.id
|
||||
RETURNING j.*;
|
||||
```
|
||||
|
||||
Worker 必须把 RETURNING 得到的 `claim_version` 作为 fencing token 保存到本次执行上下文。每处理一个分片前以 `id + status=processing + claim_version` 条件更新 `processing_started_at`;创建事件、标记 done/retry/failed 同样必须校验 claim_version 并检查 affected rows。回收后再次领取会递增版本,因此旧 Worker 即使稍后恢复也不能覆盖新领取者的结果。
|
||||
|
||||
回收器每分钟扫描一小批超时 processing:
|
||||
|
||||
- attempts < max_attempts → retry。
|
||||
- attempts >= max_attempts → failed。
|
||||
|
||||
退避建议:5s、30s、2m,上限 5m并加少量 jitter。401/403 和 invalid_response 不重试;429、5xx、连接错误和超时可重试。
|
||||
|
||||
Runner 生命周期由应用启动/停止管理:
|
||||
|
||||
- Start 验证 DB、Redis、配置。
|
||||
- Worker panic 单任务恢复并记录,不能杀死进程。
|
||||
- Shutdown 停止领取新任务,等待活动任务到有界超时。
|
||||
|
||||
### 12. OpenAI 兼容 Client 使用严格 Qwen3Guard 契约
|
||||
|
||||
请求:
|
||||
|
||||
```json
|
||||
{
|
||||
"model": "sileader/qwen3guard:0.6b",
|
||||
"messages": [{"role": "user", "content": "<chunk>"}],
|
||||
"temperature": 0,
|
||||
"max_tokens": 64,
|
||||
"seed": 42
|
||||
}
|
||||
```
|
||||
|
||||
解析要求:
|
||||
|
||||
- 响应体上限 256 KiB。
|
||||
- 只接受一个非空 `Safety:` 行和一个 `Categories:` 行。
|
||||
- 只接受 Safe、Controversial、Unsafe。
|
||||
- 不允许额外非空说明。
|
||||
- 类别做大小写/标点别名归一,但未知类别必须保留风险事实。
|
||||
|
||||
策略映射:
|
||||
|
||||
| Safety | 已启用类别 | 结果 |
|
||||
| --- | --- | --- |
|
||||
| Safe | 任意 | Pass / Allow |
|
||||
| Controversial | 普通类别 | Flag / Warn |
|
||||
| Controversial | Jailbreak/PII/Suicide & Self-Harm | Critical / Block |
|
||||
| Unsafe | 至少一个启用类别 | Critical / Block |
|
||||
| Unsafe | 未知类别 | Critical / Block + unknown_unsafe |
|
||||
| Unsafe | 仅命中明确禁用类别 | Flag / Warn,保留事实 |
|
||||
|
||||
scanner score 只用于展示排序,不得被解释为真实置信度阈值。
|
||||
|
||||
管理 API 还应从 categories、scanner evidence 和 Guard policy 确定性派生 `issue_summaries`。每项至少包含 category、scanner_id、title、description、severity/label、action/label、code、score 和脱敏 evidence;可选位置必须是 rune 范围和不可逆命中 Hash,不能返回原文。该摘要是展示 DTO,不要求新增数据库列,防止复制同一风险事实。
|
||||
|
||||
### 13. 同步 Guard 使用共享 deadline、故障切换和 bulkhead
|
||||
|
||||
同步 evaluator:
|
||||
|
||||
- 全局并发上限默认 64。
|
||||
- 每节点并发上限默认 16。
|
||||
- 总 deadline 使用第一启用节点 timeout。
|
||||
- 所有分片和节点故障切换共享 deadline。
|
||||
- 顺序扫描,最新输入优先。
|
||||
- Block 可早停;Allow 必须所有必要分片成功。
|
||||
- 连接失败、429、5xx、超时可切下一节点。
|
||||
- 401/403、invalid_response 终止。
|
||||
- 所有节点失败或 bulkhead 满 → Unavailable。
|
||||
|
||||
第一版不使用熔断器外部依赖;连续失败健康状态和冻结窗口可用模块内小状态机实现。若后续数据证明需要通用熔断库,另起 change。
|
||||
|
||||
### 14. 出站 HTTP Client 必须抵抗 SSRF 和重定向
|
||||
|
||||
保存、探测和实际调用共用同一校验:
|
||||
|
||||
- 仅 http/https。
|
||||
- 禁止 userinfo、query、fragment。
|
||||
- 禁止 link-local、multicast、unspecified、metadata host 和保留地址。
|
||||
- 公网必须 HTTPS;HTTP 只允许 localhost 或显式私网 IP/受控内网域名。
|
||||
- DNS 解析后在 DialContext 再检查每个 IP,降低 DNS rebinding 风险。
|
||||
- 默认不跟随重定向。
|
||||
- 独立连接池、Dial/TLS/Header timeout、响应上限。
|
||||
- 日志只记录 endpoint ID,不记录完整 URL。
|
||||
|
||||
### 15. HTTP、SSE 和 WebSocket 使用协议原有错误构造器
|
||||
|
||||
HTTP 错误:
|
||||
|
||||
| 情况 | HTTP | error_code |
|
||||
| --- | ---: | --- |
|
||||
| Block | 403 | prompt_guard_blocked |
|
||||
| Unavailable | 503 | prompt_guard_unavailable |
|
||||
| Invalid response | 503 | prompt_guard_invalid_response |
|
||||
|
||||
Handler 使用自己已有的 OpenAI、Claude 或 Gemini error helper。正文只包含通用中文消息、code 和 request ID。
|
||||
|
||||
现有 helper 需要通过最小协议适配器扩展稳定代码,不能破坏原字段:
|
||||
|
||||
- OpenAI Chat/Responses:保持 `error.type/message` 或 Responses 现有结构,并设置 `error.code=<prompt_guard_*>`。
|
||||
- Claude Messages:保持 `type=error` 和合法的 `error.type=permission_error|api_error`,增加可选 `error.code=<prompt_guard_*>`。
|
||||
- Gemini:保持 Google envelope 的数值 `error.code`、message 和 canonical status;在 `error.details[]` 增加 `type.googleapis.com/google.rpc.ErrorInfo`,其 `reason=<prompt_guard_*>`、domain=`sub2api.securityaudit`,metadata 只允许 request_id。
|
||||
|
||||
不得把 Gemini 数值 `error.code` 替换为字符串,也不得把类别、Prompt、节点或内部错误放入 details。协议 golden test 必须锁定三类 envelope。
|
||||
|
||||
SSE 必须在 Guard 完成前不写 response header/首字节。
|
||||
|
||||
Responses WebSocket:
|
||||
|
||||
- 握手本身无 prompt,不执行输入分类。
|
||||
- 首个 response.create 在用户/账号 slot、计费和上游拨号前检查。
|
||||
- 每个后续 response.create 在本轮 slot、计费和上游发送前重新检查。
|
||||
- Block:close 4403,reason prompt_guard_blocked。
|
||||
- Unavailable/Invalid:close 1013,对应稳定 reason。
|
||||
- 日志 stage=first_turn/subsequent_turn。
|
||||
|
||||
### 16. 同步结果采用独立轻量记录路径
|
||||
|
||||
同步 evaluator 返回:
|
||||
|
||||
```text
|
||||
decision, action, risk_level, categories,
|
||||
matched_scanners, scores, evidence,
|
||||
scanner_backend/version, endpoint_id,
|
||||
policy_id/version, chunk_total, latency,
|
||||
error_code, allow_next_stage
|
||||
```
|
||||
|
||||
记录 adapter:
|
||||
|
||||
- 不接受完整 scan_text,只接受脱敏 PromptSnapshot。
|
||||
- 创建 `execution_mode=blocking,status=done` 的 job。
|
||||
- 按 store_pass_events 决定是否创建事件。
|
||||
- 在单个 DB transaction 内完成 job + event。
|
||||
- 记录失败只增加指标和日志,不改变 evaluator 已确定结果。
|
||||
- 禁止再次调用 Guard。
|
||||
|
||||
### 17. 管理 API 使用独立前缀和现有管理员审计
|
||||
|
||||
新增:
|
||||
|
||||
```text
|
||||
GET /admin/prompt-audit/config
|
||||
PUT /admin/prompt-audit/config
|
||||
POST /admin/prompt-audit/endpoints/probe
|
||||
GET /admin/prompt-audit/runtime
|
||||
GET /admin/prompt-audit/events
|
||||
GET /admin/prompt-audit/events/:id
|
||||
DELETE /admin/prompt-audit/events/:id
|
||||
POST /admin/prompt-audit/events/batch-delete
|
||||
POST /admin/prompt-audit/events/delete-preview
|
||||
POST /admin/prompt-audit/events/delete-by-filter
|
||||
```
|
||||
|
||||
所有写操作和敏感探测复用 AdminAuth 和现有管理操作审计。审计 detail 采用 allowlist 字段,不使用“先记录完整结构再删除敏感 key”的方式。
|
||||
|
||||
删除规则:
|
||||
|
||||
- 单次批量 ID 数量有上限。
|
||||
- 按筛选删除必须带开始/结束时间、预览 Hash、服务端认证 confirmation_token 和 confirm。
|
||||
- preview 在同一数据库快照中返回 matched_count、`snapshot_max_id` 和 `filter_hash = SHA-256(canonical JSON filter summary + snapshot_max_id)`。
|
||||
- confirmation_token 是由现有 SecretEncryptor 认证加密的短期 claim,绑定 filter_hash、snapshot_max_id、管理员 ID、签发/过期时间(默认 5 分钟)。delete-by-filter 必须解密、校验操作者/过期时间/Hash,并强制 `id <= snapshot_max_id`;客户端自行计算 SHA-256 不能绕过预览,预览后的新事件不能被本次操作删除。
|
||||
- 删除分批执行,避免长事务。
|
||||
- 删除事件后只删除无任何事件引用且非 processing 的孤立 job。
|
||||
- 尝试清理对应 Redis key。
|
||||
|
||||
### 18. 控制台使用独立 feature 目录
|
||||
|
||||
```text
|
||||
frontend/src/features/prompt-audit/
|
||||
├── PromptAuditView.vue
|
||||
├── api.ts
|
||||
├── types.ts
|
||||
├── viewModel.ts
|
||||
├── components/
|
||||
└── __tests__/
|
||||
```
|
||||
|
||||
少量外部接线:
|
||||
|
||||
- router 增加 `/admin/prompt-audit`,复用 requiresAuth/requiresAdmin/requiresRiskControl。
|
||||
- Sidebar 把现有 risk-control 单项改为 expandOnly “安全审计”分组,子项保留原路由并新增提示词路由。
|
||||
- i18n 增加 zh/en 对称键。
|
||||
|
||||
页面分区:
|
||||
|
||||
1. 运行概览。
|
||||
2. 审计池表格和参数/探测对话框。
|
||||
3. 分组范围和九类 scanner。
|
||||
4. Worker/队列/配置版本/Guard 指标。
|
||||
5. 事件筛选、表格、详情、删除。
|
||||
6. 固定保存栏:enabled、blocking、store pass、保存/重置。
|
||||
|
||||
页面不得在 localStorage/sessionStorage 保存 API Key。保存成功后立即清除输入 state。
|
||||
|
||||
### 19. 日志和指标使用稳定词典
|
||||
|
||||
最小事件:
|
||||
|
||||
```text
|
||||
prompt_audit.config_updated
|
||||
prompt_guard.config_loaded
|
||||
prompt_guard.config_reload_degraded
|
||||
prompt_audit.endpoint_probe_started
|
||||
prompt_audit.endpoint_probe_finished
|
||||
prompt_audit.endpoint_probe_failed
|
||||
prompt_audit.job_enqueued
|
||||
prompt_audit.enqueue_skipped
|
||||
prompt_audit.enqueue_dropped
|
||||
prompt_audit.started
|
||||
prompt_audit.processing_reclaimed
|
||||
prompt_audit.processed
|
||||
prompt_audit.process_failed
|
||||
prompt_audit.finding_recorded
|
||||
prompt_audit.scan_chunk_started
|
||||
prompt_audit.scan_chunk_completed
|
||||
prompt_audit.scan_chunk_failed
|
||||
prompt_audit.scan_chunks_aggregated
|
||||
prompt_guard.evaluation_started
|
||||
prompt_guard.allowed
|
||||
prompt_guard.blocked
|
||||
prompt_guard.failed
|
||||
prompt_guard.result_record_failed
|
||||
prompt_audit.event_deleted
|
||||
prompt_audit.events_deleted
|
||||
prompt_audit.events_delete_previewed
|
||||
prompt_audit.events_filter_deleted
|
||||
```
|
||||
|
||||
字段采用 allowlist:request_id、user_id、api_key_id、group_id、provider、protocol、endpoint、model、job_id、event_id、config_version、guard_endpoint_id、decision、risk_level、action、chunk_index、chunk_total、chunk_chars、input_chars、input_limit、latency_ms、status、error_code、error_kind、queue_length/capacity、stage、upstream_dispatched、billing_preconsumed。
|
||||
|
||||
禁止:body、raw_prompt、payload、token、authorization、完整 Base URL/query、Redis value。
|
||||
|
||||
指标:异步 enqueue/dropped、队列各状态、processed/failed、Worker active、Guard total/allow/flag/block/unavailable/invalid/timeout/failover/bulkhead/record_failure、延迟直方图。Guard 结果与延迟由同步 evaluator 和异步 Worker 使用同一稳定指标结构观测,使 blocking 启用前可以先在 async 测试分组建立 P50/P95/P99、失败率和事件增长率基线;runtime 同时返回 async enqueue/dropped 计数以区分投递与扫描阶段。
|
||||
|
||||
### 20. 测试按行为矩阵而不是文件覆盖率验收
|
||||
|
||||
核心矩阵:
|
||||
|
||||
| 维度 | 值 |
|
||||
| --- | --- |
|
||||
| 引擎 | 现有 moderation / prompt audit / 两者 |
|
||||
| Prompt 模式 | off / async / blocking |
|
||||
| 协议 | chat / responses / messages / gemini / images-media / responses-ws |
|
||||
| 返回 | allow / flag / block / unavailable / invalid |
|
||||
| 流式 | non-stream / SSE / WS first / WS subsequent |
|
||||
| 副作用 | account selection / billing / upstream |
|
||||
|
||||
必须有结构测试验证所有现有调用点经过 Coordinator;必须有 stub 统计 Block/Unavailable 时账号选择、计费和上游调用均为 0。
|
||||
|
||||
敏感信息测试对日志、DB row、API JSON、前端 state snapshot 做 canary secret 断言。
|
||||
|
||||
### 21. 不新增外部运行时依赖
|
||||
|
||||
使用现有 go-redis、database/sql、Gin、SecretEncryptor、logger、Vue 3、Axios 和测试工具。Qwen3Guard 是外部 OpenAI 兼容服务,不在本仓库启动模型进程。
|
||||
|
||||
不引入新的 Go 队列库、ORM、前端状态库或 UI 框架。
|
||||
|
||||
## Risks / Trade-offs
|
||||
|
||||
- [两个同步引擎会增加首字节延迟] → 只有管理员显式开启 blocking 才发生;并行执行、最新输入优先、Block 早停、共享 deadline、连接池和分组灰度。
|
||||
- [Guard 故障在 fail-closed 下影响可用性] → 多节点有序故障切换、bulkhead、真实探测、运行态告警和一键关闭 blocking;Unavailable 与 Block 使用不同错误码。
|
||||
- [Qwen3Guard 误报导致合法请求被拒绝] → 先运行 async 建立误报基线,再按 group 灰度 blocking;保留独立事件,不直接触发封号。
|
||||
- [两个引擎同时 Block 时语义冲突] → 固定现有内容审核响应优先级,两个事件仍独立记录。
|
||||
- [PostgreSQL/Redis 非事务导致悬挂状态] → staging → Redis SET → queued 发布协议;staging 回收和 TTL 清理。
|
||||
- [多实例重复消费或旧 Worker 覆盖新结果] → `FOR UPDATE SKIP LOCKED` 原子领取、递增 claim_version fencing token、processing 租约和带版本条件更新。
|
||||
- [长提示词导致超时] → Unicode 分片、总 deadline、最新输入优先;Allow 必须完整覆盖,禁止部分结果放行。
|
||||
- [SSRF 或凭据泄露] → 保存/探测/调用共用校验、DNS 后复检、禁止重定向、加密密文、日志/API allowlist、canary 泄露测试。
|
||||
- [手工接入多个 Handler 造成漏路由] → 将现有调用统一替换为 Coordinator 并增加静态/结构路由矩阵测试。
|
||||
- [新模块仍反向侵入现有 service] → 新模块依赖现有端口;现有 ContentModerationService 不导入新模块,Handler 仅注入 Coordinator。
|
||||
- [事件量过大] → 默认不保存 Pass,分页索引、分批删除;后续根据真实规模单独设计自动保留期。
|
||||
- [源参考继续变化] → 实施前冻结源基线,本 change specs 作为目标实现最终权威。
|
||||
|
||||
## Migration Plan
|
||||
|
||||
### 阶段 0:冻结和对照
|
||||
|
||||
1. 记录参考仓库 commit、branch 和 `git diff --stat`。
|
||||
2. 对未提交的同步阻止文件生成只读 patch 或提交到专用分支。
|
||||
3. 建立“源功能 → 本 change requirement → 目标测试”追踪表。
|
||||
|
||||
### 阶段 1:纯数据和配置基础
|
||||
|
||||
1. 新增 SQL migration 和 Repository 测试。
|
||||
2. 新增加密配置、Public DTO、URL 校验和 config cache。
|
||||
3. 新增管理 API 的 config/probe/runtime 骨架。
|
||||
4. 保持 enabled=false,不接网关。
|
||||
|
||||
### 阶段 2:异步审计
|
||||
|
||||
1. 实现 PromptSnapshot、脱敏、Hash 和协议提取。
|
||||
2. 实现 staging 投递、Redis Payload Store、Worker、重试和回收。
|
||||
3. 实现 OpenAI 兼容 Client、Qwen parser、分片聚合和事件。
|
||||
4. 接入 Coordinator 的 async 分支;队列故障不影响请求。
|
||||
|
||||
### 阶段 3:控制台和运营闭环
|
||||
|
||||
1. 完成页面、节点探测、配置、运行态和事件列表/详情。
|
||||
2. 完成单条、批量和按筛选删除。
|
||||
3. 运行前后端 lint、typecheck、unit/integration test。
|
||||
|
||||
### 阶段 4:同步门禁
|
||||
|
||||
1. 实现 evaluator、bulkhead、deadline、故障切换和错误映射。
|
||||
2. 完成 HTTP/SSE 入口接线。
|
||||
3. 完成 Responses WS 首轮与后续帧接线。
|
||||
4. 用副作用 stub 证明 Block/Unavailable 无账号、无计费、无上游。
|
||||
|
||||
### 阶段 5:灰度上线
|
||||
|
||||
1. 生产先保持 Prompt Audit off。
|
||||
2. 开启 async,只选测试 group,观察 Guard 延迟、失败、误报和事件量。
|
||||
3. 建立良性/恶意回归语料。
|
||||
4. 仅在多节点稳定、Unavailable 率和 P99 满足阈值后开启 blocking。
|
||||
5. 按 group 扩大范围。
|
||||
|
||||
### 回滚
|
||||
|
||||
- 首选:关闭 blocking_enabled,立即回到 async。
|
||||
- 次选:关闭 enabled,完全停止新 Prompt Audit。
|
||||
- 必要时关闭 risk_control_enabled,但这也会停用现有内容审核入口,应作为最后手段。
|
||||
- 回滚不删除表、配置或历史事件,不回退已应用 migration。
|
||||
- Worker 停止后 queued/retry 任务保留;恢复时继续处理,或由管理员按明确策略清理。
|
||||
|
||||
## Resolved Decisions
|
||||
|
||||
1. **源基线标识**:采用 `source-freeze/` 中的只读 tracked patch + untracked archive;base commit、SHA-256 和恢复测试已登记在 `source-baseline.md`。
|
||||
2. **事件自动保留期**:第一版只提供管理员安全删除,不增加自动保留清理;真实事件量稳定后另起 change。
|
||||
3. **同步两个引擎并行或串行**:采用受控并行;实现必须通过 race test,并保持 Legacy Block 优先级和两引擎独立记录。
|
||||
4. **目标项目额外文本入口**:以实施时 `backend/internal/server/routes/gateway.go` 的自动/结构枚举为事实源;所有用户文本入口必须接入 Coordinator 或提供不会旁路/重复扫描的测试证明。
|
||||
5. **生产启用阈值**:实现和部署验证期间只允许 off/async;blocking 生产启用必须满足 `verification.md` 的建议阈值并由安全、运营和业务责任人签字,未签字不得生产开启。
|
||||
@@ -0,0 +1,31 @@
|
||||
# Prompt Audit implementation evidence
|
||||
|
||||
This file records reproducible implementation-time evidence. It contains no prompt bodies, Guard credentials, Authorization values, or Redis payloads.
|
||||
|
||||
## 2026-07-16 — source freeze and target baseline
|
||||
|
||||
### Frozen source restore
|
||||
|
||||
- Base commit: `7a50378851a80650cb0c086260b23abeb3469e6b`
|
||||
- Freeze manifest: `source-freeze/MANIFEST.md`
|
||||
- Manifest SHA-256: `badab312bf6af4d2c77857a9400381f4da4fbf45722d9f4a6df23bc7005273b6`
|
||||
- Restore result: tracked patch and untracked archive restored into a detached worktree; `git diff --check` passed.
|
||||
- `go test ./internal/service/promptaudit -count=1`: passed.
|
||||
- `go test ./internal/router ./internal/relay ./internal/gatewayadapter/transport -run 'PromptGuard|PromptAudit|ConcurrencyOrder' -count=1`: passed.
|
||||
|
||||
### Target pre-change baseline
|
||||
|
||||
- `cd backend && go test ./internal/service -run ContentModeration -count=1`: passed (`1.138s`).
|
||||
- `pnpm --dir frontend exec vitest run src/views/admin/__tests__/RiskControlView.spec.ts src/router/__tests__/feature-access.spec.ts`: passed (2 files, 9 tests).
|
||||
|
||||
### Review slices
|
||||
|
||||
Implementation is partitioned into independently reviewable slices without changing the final scope:
|
||||
|
||||
1. Data and core contracts.
|
||||
2. Async audit engine.
|
||||
3. Admin API and console.
|
||||
4. Coordinator and synchronous guard.
|
||||
5. Observability, verification, rollout, and deployment evidence.
|
||||
|
||||
The feature remains default-off throughout implementation. Production blocking remains prohibited until the signed rollout gates in `verification.md` are satisfied.
|
||||
@@ -0,0 +1,581 @@
|
||||
# 实施指导
|
||||
|
||||
## 1. 使用方式与不可变边界
|
||||
|
||||
本指南把 `proposal.md`、`design.md` 和三个 delta specs 转换为可按文件实施、可逐阶段评审的操作顺序。若本指南与 specs 冲突,以 specs 为准,并先更新 OpenSpec 再编码。
|
||||
|
||||
实施前必须满足:
|
||||
|
||||
- `source-baseline.md` 的冻结登记已完成,不再以变化中的源工作区作为唯一依据。
|
||||
- 当前内容审核后端测试、RiskControl 前端测试和路由清单已保存为基线证据。
|
||||
- 新功能的默认配置是 off;数据库迁移可以先上线,但不能自动开启审计。
|
||||
- `content_moderation_logs`、`ContentModerationService`、`/admin/risk-control` 和 `RiskControlView.vue` 的业务语义不改变。
|
||||
- 完整 Prompt 只允许存在于请求内存和 Redis TTL value;Guard token 只允许存在于写入 DTO、解密后的短生命周期内存和 Authorization header。
|
||||
|
||||
明确不做:输出审核、自动改写/脱敏后转发、人工审批、申诉、自动封号、邮件、Prompt 命中 Hash 黑名单、现有 Moderations 分类映射。
|
||||
|
||||
## 2. 目标依赖方向
|
||||
|
||||
```mermaid
|
||||
flowchart TD
|
||||
Routes["server/routes 与协议 Handler"] --> Helper["security_audit_helper.go"]
|
||||
Helper --> Coordinator["securityaudit.Coordinator"]
|
||||
Coordinator --> LegacyPort["LegacyModerationEngine 接口"]
|
||||
Coordinator --> PromptService["PromptService"]
|
||||
LegacyPort --> Existing["现有 ContentModerationService"]
|
||||
PromptService --> Ports["ConfigStore / JobRepository / PayloadStore / Scanner"]
|
||||
Ports --> Infra["settings / database/sql / Redis / SecretEncryptor / HTTP"]
|
||||
AdminRoutes["admin routes"] --> AdminHandler["PromptAdminHandler"]
|
||||
AdminHandler --> PromptService
|
||||
Frontend["features/prompt-audit"] --> AdminRoutes
|
||||
```
|
||||
|
||||
依赖规则:
|
||||
|
||||
1. 现有 `internal/service` 不得 import `internal/securityaudit`;否则会把新能力反向渗入既有业务层。
|
||||
2. `securityaudit` 可以通过小接口适配现有 service/repository/Redis/加密能力,但不得修改这些接口的全局语义来迁就新模块。
|
||||
3. Coordinator 只归并客户端决策,不写 job/event、不发送邮件、不封号、不更新现有 Hash。
|
||||
4. Handler 只负责构造可信请求、调用 Coordinator、使用本协议原有错误 helper 返回结果。
|
||||
5. 核心逻辑不得读取 Gin context、环境变量或包级全局配置;这些只在模块构造/Handler 边界转换。
|
||||
6. 构造函数不得启动 goroutine。Worker、回收器和配置订阅必须由 `Start(ctx)` 启动、由 `Shutdown(ctx)` 有界停止。
|
||||
7. 前端只能依赖公共 DTO,不得知道 `token_ciphertext`、Redis key 或数据库内部状态转换 SQL。
|
||||
8. 新模块不引入新的 ORM、队列库、状态库或 UI 框架。
|
||||
|
||||
## 3. 建议目录和文件职责
|
||||
|
||||
```text
|
||||
backend/internal/securityaudit/
|
||||
├── coordinator.go # 双引擎编排、固定优先级
|
||||
├── coordinator_test.go
|
||||
├── prompt_types.go # Request/Decision/Job/Event/Runtime 与枚举
|
||||
├── prompt_config.go # Storage/Public/Update DTO、校验、快照
|
||||
├── prompt_config_test.go
|
||||
├── prompt_snapshot.go # 协议提取、Hash、脱敏预览
|
||||
├── prompt_snapshot_test.go
|
||||
├── prompt_scanner.go # 分片、聚合、Scanner 接口
|
||||
├── prompt_qwen3guard.go # 请求构造、严格解析、九类风险
|
||||
├── prompt_qwen3guard_test.go
|
||||
├── prompt_issue_summary.go # 从分类/脱敏证据派生管理端风险摘要
|
||||
├── prompt_issue_summary_test.go
|
||||
├── prompt_outbound_security.go # URL/DNS/Dial/redirect/响应上限
|
||||
├── prompt_outbound_security_test.go
|
||||
├── prompt_repository.go # database/sql jobs/events 实现
|
||||
├── prompt_repository_test.go
|
||||
├── prompt_payload_store.go # Redis SET EX/GET/DEL
|
||||
├── prompt_enqueue.go # staging → payload → queued
|
||||
├── prompt_enqueue_test.go
|
||||
├── prompt_worker.go # claim/lease/retry/reclaim/lifecycle
|
||||
├── prompt_worker_test.go
|
||||
├── prompt_guard.go # blocking evaluator、deadline/failover/bulkhead
|
||||
├── prompt_guard_test.go
|
||||
├── prompt_runtime.go # 健康、版本、队列、指标快照
|
||||
├── prompt_logging.go # 稳定事件和 allowlist fields
|
||||
├── prompt_handler.go # 独立 admin HTTP handler
|
||||
├── prompt_handler_test.go
|
||||
└── prompt_module.go # provider set、Start/Shutdown 组合
|
||||
|
||||
backend/migrations/181_prompt_audit.sql
|
||||
backend/internal/handler/security_audit_helper.go
|
||||
backend/internal/server/routes/admin.go
|
||||
backend/internal/server/routes/gateway.go
|
||||
backend/internal/wire/或项目实际 provider 文件
|
||||
|
||||
frontend/src/features/prompt-audit/
|
||||
├── PromptAuditView.vue
|
||||
├── api.ts
|
||||
├── types.ts
|
||||
├── viewModel.ts
|
||||
├── components/
|
||||
└── __tests__/
|
||||
```
|
||||
|
||||
`181` 是提案编写时最大迁移号后的建议值。实施时若 181 已存在,必须使用新的最大序号;不得改写已经应用的 migration。
|
||||
|
||||
## 4. 按文件的实施顺序
|
||||
|
||||
### 4.1 第一批:契约与纯函数
|
||||
|
||||
1. 创建 `prompt_types.go`,固定稳定枚举和 JSON 字段。
|
||||
2. 创建 `prompt_config.go`,先实现默认值、三态归一、字段边界和 Public DTO。
|
||||
3. 创建 `prompt_snapshot.go`,完成各协议纯文本提取、最新输入优先、SHA-256 和脱敏预览。
|
||||
4. 创建 `prompt_scanner.go`、`prompt_qwen3guard.go` 与 `prompt_issue_summary.go`,完成 rune 分片、严格解析、聚合和展示摘要派生。
|
||||
5. 同步创建上述测试;此阶段不连接 DB、Redis、Gin 或真实 Guard。
|
||||
|
||||
验收重点:纯函数表驱动测试覆盖中文、emoji、空输入、混合 content blocks、九类风险、额外说明、重复字段、未知类别和完整分片。
|
||||
|
||||
### 4.2 第二批:数据库与配置适配
|
||||
|
||||
1. 新增 `181_prompt_audit.sql` 以及 migration schema 测试。
|
||||
2. 在 `prompt_repository.go` 用现有 `*sql.DB` 实现 jobs/events;不为这两张表增加 Ent schema。
|
||||
3. 在目标项目现有 setting 常量事实源增加 `prompt_audit_config`。
|
||||
4. 在 `prompt_config.go` 复用 `SettingRepository` 和 `SecretEncryptor`,实现 storage ↔ active ↔ public 三类 DTO 转换。
|
||||
5. 在 `prompt_payload_store.go` 适配现有 Redis Client。
|
||||
6. 完成 Repository、加密配置、多实例版本加载测试。
|
||||
|
||||
### 4.3 第三批:出站安全、异步队列与运行态
|
||||
|
||||
1. `prompt_outbound_security.go` 先实现保存/探测/调用共用的 URL 校验和受控 Transport。
|
||||
2. `prompt_enqueue.go` 实现 staging 发布协议。
|
||||
3. `prompt_worker.go` 实现 PostgreSQL claim、租约、重试、回收和生命周期。
|
||||
4. `prompt_runtime.go` 汇总 active/expected config version、Worker、队列、Redis 和节点健康。
|
||||
5. `prompt_logging.go` 固定事件名、error_code 和允许字段。
|
||||
6. 用 fake clock、fake scanner、真实测试 PostgreSQL/Redis 分层验证,先不开网关。
|
||||
|
||||
### 4.4 第四批:管理 API 和控制台
|
||||
|
||||
1. `prompt_handler.go` 注册 config/probe/runtime/events/delete 方法。
|
||||
2. 在 admin handler 聚合结构和 Wire 中注入 `PromptAdminHandler`。
|
||||
3. 在 `admin.go` 注册独立 `/admin/prompt-audit` 路由组。
|
||||
4. 创建前端 `features/prompt-audit` 的 types、api、viewModel,再创建页面和组件。
|
||||
5. 增加 router、Sidebar、zh/en i18n 的薄接线。
|
||||
6. 管理闭环通过后,Prompt Audit 仍默认 off。
|
||||
|
||||
### 4.5 第五批:Coordinator 与异步接入
|
||||
|
||||
1. 在 `coordinator.go` 用 fake engines 完成 off/async/blocking 组合测试。
|
||||
2. 新增 `security_audit_helper.go`,从现有 `buildContentModerationInput` 的可信字段构造 `securityaudit.Request`。
|
||||
3. 机械替换所有现有 `checkContentModeration` 调用点为 `checkSecurityAudit`,保留原位置。
|
||||
4. async 模式只 best-effort 投递;Redis/DB/节点失败不得改变客户端响应或上游次数。
|
||||
5. 运行路由结构测试,证明没有漏掉已有调用点。
|
||||
|
||||
### 4.6 第六批:同步 Guard
|
||||
|
||||
1. `prompt_guard.go` 实现共享 deadline、节点优先级、故障切换和 bulkhead。
|
||||
2. Coordinator 接入 blocking 分支并固定现有内容审核 Block 响应优先级。
|
||||
3. HTTP/SSE 使用协议原有错误构造器;Guard 完成前 SSE 不写首字节。
|
||||
4. Responses WS 首轮和后续 `response.create` 分别接入,使用指定 close code。
|
||||
5. 加入账号选择、并发 slot、预扣/计费、上游拨号/写入 fake counter,断言拒绝时全部为 0。
|
||||
|
||||
## 5. 公共核心类型建议
|
||||
|
||||
### 5.1 可信请求
|
||||
|
||||
```go
|
||||
type Request struct {
|
||||
RequestID string
|
||||
UserID int64
|
||||
Username string
|
||||
UserEmail string
|
||||
APIKeyID int64
|
||||
APIKeyName string
|
||||
GroupID *int64
|
||||
GroupName string
|
||||
Provider string
|
||||
Endpoint string
|
||||
Protocol string
|
||||
Model string
|
||||
Body []byte
|
||||
Stage string // http | first_turn | subsequent_turn
|
||||
}
|
||||
```
|
||||
|
||||
Body 必须是 Handler 在全局 body limit 下已经读取的同一字节切片。模块不得再次读 `http.Request.Body`,不得改写转发 body。`Username`、`UserEmail` 和 `APIKeyName` 只用于管理员事件快照/展示,不得进入普通请求日志;API 必须分列返回,避免复制/筛选时含义混淆。
|
||||
|
||||
### 5.2 统一决策
|
||||
|
||||
```go
|
||||
type Decision struct {
|
||||
Kind string // allow | flag | block | unavailable | invalid
|
||||
HTTPStatus int
|
||||
ErrorCode string
|
||||
ClientMessage string
|
||||
Legacy *LegacyDecision
|
||||
Prompt *PromptDecision
|
||||
AllowNextStage bool
|
||||
}
|
||||
```
|
||||
|
||||
稳定优先级:
|
||||
|
||||
1. Legacy content moderation Block:完全复用原状态码、文案和 `content_policy_violation`。
|
||||
2. Prompt Block:403 + `prompt_guard_blocked`。
|
||||
3. Prompt Invalid:503 + `prompt_guard_invalid_response`。
|
||||
4. Prompt Unavailable:503 + `prompt_guard_unavailable`。
|
||||
5. 其他:Allow;Flag 只记录,不阻断。
|
||||
|
||||
不要让 Coordinator 暴露 Qwen 原始响应,也不要用一个布尔 `Blocked` 吞掉 unavailable/invalid 的差异。
|
||||
|
||||
## 6. Coordinator 请求流
|
||||
|
||||
```text
|
||||
鉴权与 body/model 基础校验
|
||||
→ 构造可信 Request
|
||||
→ 读取 risk_control + prompt active snapshot
|
||||
→ Coordinator 调用现有 Moderation 与 Prompt 引擎
|
||||
→ 按固定优先级得到 Decision
|
||||
→ 若 !AllowNextStage,使用当前协议 error helper 返回
|
||||
→ 否则才进入账号选择/并发/计费/上游
|
||||
```
|
||||
|
||||
模式行为:
|
||||
|
||||
| 有效模式 | 现有 Moderation | Prompt Audit | 请求等待 Prompt | Prompt 失败影响请求 |
|
||||
| --- | --- | --- | --- | --- |
|
||||
| off | 原行为 | 不运行 | 否 | 否 |
|
||||
| async_audit | 原行为 | best-effort enqueue | 否 | 否 |
|
||||
| blocking | 原行为 | 同步扫描并复用结果记录 | 是 | 是,fail-closed |
|
||||
|
||||
async 模式下应先触发/完成有界投递动作,再返回 Coordinator 结果,确保现有 Moderation 随后 Block 时 Prompt 事件仍可 best-effort 产生。投递动作必须只有短 DB/Redis 操作,不能等待 Guard。
|
||||
|
||||
blocking 模式可以并行执行两个引擎,但必须遵守:
|
||||
|
||||
- goroutine 数量固定且可等待,不得 fire-and-forget。
|
||||
- 两个结果都在各自 deadline 内收口,或明确取消。
|
||||
- Legacy Block 的响应优先,但 Prompt 结果仍按独立规则记录。
|
||||
- 共享只读 Request;不得共享可变 decision buffer。
|
||||
|
||||
## 7. 异步时序
|
||||
|
||||
```mermaid
|
||||
sequenceDiagram
|
||||
participant H as Protocol Handler
|
||||
participant C as Coordinator
|
||||
participant E as Prompt Enqueuer
|
||||
participant PG as PostgreSQL
|
||||
participant R as Redis
|
||||
participant W as Worker
|
||||
participant G as Qwen3Guard
|
||||
|
||||
H->>C: Check(trusted Request)
|
||||
C->>E: Enqueue(snapshot, scan text)
|
||||
E->>PG: INSERT job status=staging
|
||||
PG-->>E: job_id
|
||||
E->>R: SET payload:{job_id} scan_text EX 1800
|
||||
R-->>E: OK
|
||||
E->>PG: UPDATE staging → queued (conditional)
|
||||
E-->>C: accepted
|
||||
C-->>H: legacy decision / allow
|
||||
H->>H: 继续原账号、计费、上游流程
|
||||
|
||||
W->>PG: claim queued/retry FOR UPDATE SKIP LOCKED
|
||||
PG-->>W: status=processing job
|
||||
W->>R: GET payload:{job_id}
|
||||
loop 每个必要分片
|
||||
W->>PG: refresh processing lease
|
||||
W->>G: POST /v1/chat/completions
|
||||
G-->>W: Safety + Categories
|
||||
end
|
||||
W->>PG: transaction: event + job done
|
||||
W->>R: DEL payload:{job_id}
|
||||
```
|
||||
|
||||
异常补偿:
|
||||
|
||||
- active count 与 staging INSERT 在 PostgreSQL advisory-lock 短事务中完成;锁超时使用 `queue_admission_busy`,不得把 Redis 调用放进事务。
|
||||
- staging INSERT 失败:不写 Redis,记录 dropped,主请求继续。
|
||||
- Redis SET 失败:job 条件标 failed;主请求继续。
|
||||
- staging → queued 条件更新失败:删除 Redis key;回收器处理残留 staging。
|
||||
- Worker 找不到 payload:按稳定 `payload_missing` 失败,不可把预览当原文扫描。
|
||||
- event 写入失败:异步 job retry 或 failed,不能产生虚假 done。
|
||||
- Redis DEL 失败:依靠 TTL,记录脱敏警告。
|
||||
|
||||
## 8. 同步阻断时序
|
||||
|
||||
```mermaid
|
||||
sequenceDiagram
|
||||
participant H as HTTP/SSE/WS Handler
|
||||
participant C as Coordinator
|
||||
participant M as Existing Moderation
|
||||
participant P as Prompt Guard
|
||||
participant G as Guard Pool
|
||||
participant D as DB Recorder
|
||||
participant A as Account/Billing/Upstream
|
||||
|
||||
H->>C: Check(Request, blocking snapshot)
|
||||
par 保持现有审核语义
|
||||
C->>M: Check
|
||||
M-->>C: legacy decision
|
||||
and 共享总预算扫描
|
||||
C->>P: Evaluate(snapshot)
|
||||
P->>G: chunks × ordered failover
|
||||
G-->>P: normalized result
|
||||
P-->>C: Allow/Flag/Block/Unavailable/Invalid
|
||||
end
|
||||
C-->>D: record redacted result (no scan text)
|
||||
D-->>C: best-effort record status
|
||||
C-->>H: prioritized Decision
|
||||
alt Block/Unavailable/Invalid
|
||||
H-->>H: protocol-compatible error/close
|
||||
Note over H,A: account selection=0, billing=0, upstream=0
|
||||
else Allow/Flag
|
||||
H->>A: continue original flow
|
||||
end
|
||||
```
|
||||
|
||||
同步记录失败不得反转已确定结果。一次同步评估只调用 Guard 一次;记录 adapter 禁止接收 `scan_text`,防止为了落库再次扫描或意外持久化原文。
|
||||
|
||||
## 9. Job 状态机
|
||||
|
||||
```mermaid
|
||||
stateDiagram-v2
|
||||
[*] --> staging: INSERT
|
||||
staging --> queued: Redis SET 成功且条件发布
|
||||
staging --> failed: Redis/发布失败或 staging 超时回收
|
||||
queued --> processing: 原子 claim
|
||||
retry --> processing: 到达 next_attempt_at 后原子 claim
|
||||
processing --> done: 必要分片完成且事件事务成功
|
||||
processing --> retry: 可重试错误且 attempts < max_attempts
|
||||
processing --> failed: 不可重试或达到上限
|
||||
processing --> retry: 租约超时回收且仍可重试
|
||||
processing --> failed: 租约超时且达到上限
|
||||
done --> [*]
|
||||
failed --> [*]
|
||||
```
|
||||
|
||||
每次 queued/retry → processing 必须把 `claim_version` 原子加一并返回给 Worker。租约刷新、event+done 事务和 retry/failed 更新必须使用“id + processing + claim_version”条件并检查 affected rows;0 rows 表示租约已失效,本 Worker 必须丢弃结果。禁止仅按 status 条件更新,因为任务被回收并重新领取后 status 会再次变成 processing,旧 Worker 会误覆盖新结果。
|
||||
|
||||
## 10. 配置、Storage DTO 与 Public DTO
|
||||
|
||||
### 10.1 存储结构
|
||||
|
||||
setting key 固定为 `prompt_audit_config`,JSON 至少包含:
|
||||
|
||||
```text
|
||||
enabled, blocking_enabled, store_pass_events,
|
||||
strategy=priority, worker_count, queue_capacity,
|
||||
scanners[], all_groups, group_ids[], endpoints[],
|
||||
config_version, updated_at, updated_by, change_summary
|
||||
```
|
||||
|
||||
Endpoint storage 字段:
|
||||
|
||||
```text
|
||||
id, name, protocol=openai_compatible, base_url,
|
||||
model=sileader/qwen3guard:0.6b,
|
||||
token_ciphertext, timeout_ms, input_limit, enabled
|
||||
```
|
||||
|
||||
### 10.2 写入 DTO
|
||||
|
||||
每个 endpoint 的写入必须区分:
|
||||
|
||||
- `token` 非空:校验后加密并替换旧密文。
|
||||
- `token` 空且 `clear_token=false`:保留旧密文;新 endpoint 没有旧密文时校验失败。
|
||||
- `clear_token=true`:清除密文;启用的 endpoint 若必须认证则保存失败或明确显示不可用。
|
||||
|
||||
保存请求必须携带 `expected_config_version`。后端在 PostgreSQL 短事务中取得该 setting 专用 advisory transaction lock、重读当前值并做 CAS;冲突返回 409 `prompt_audit_config_conflict`,不写 settings、不安装快照、不发 Redis 通知。`enabled=false && blocking_enabled=true` 必须返回稳定错误 `prompt_guard_requires_audit_enabled`。`strategy` 第一版只接受 `priority`。保存时 canonicalize group IDs、scanner IDs 和 endpoint IDs,拒绝重复、空 ID、越界 worker/queue/timeout/input_limit。
|
||||
|
||||
### 10.3 Public DTO
|
||||
|
||||
GET config 和 PUT 成功响应只允许:
|
||||
|
||||
```text
|
||||
id, name, protocol, base_url, model, timeout_ms,
|
||||
input_limit, enabled, has_token, token_status
|
||||
```
|
||||
|
||||
不得出现 `token`、`token_ciphertext`、Authorization、解密失败原文或完整错误响应。后端 JSON 类型应物理分离,不能依赖 `json:"-"` 后复用内部对象。
|
||||
|
||||
### 10.4 活动快照
|
||||
|
||||
- 保存成功后 `config_version + 1`,先安装本实例只读快照,再发布 Redis invalidation。
|
||||
- Pub/Sub 消息只含版本,不含配置。
|
||||
- 其他实例重新从 settings 加载、解密、验证,成功后原子替换。
|
||||
- 加载失败保留 last-known-good,并在 runtime 同时展示 expected/active version 和错误。
|
||||
- 冷启动无 last-known-good 且 blocking 期望启用时必须 degraded/error,不能当作 off 放行。
|
||||
- 请求热路径只读内存快照,不查 settings/DB。
|
||||
|
||||
## 11. 管理 API 映射
|
||||
|
||||
统一前缀:`/admin/prompt-audit`。全部复用现有管理员鉴权、安全中间件和管理操作审计。
|
||||
|
||||
| 方法 | 路径 | 用途 | 关键约束 |
|
||||
| --- | --- | --- | --- |
|
||||
| GET | `/config` | 读取公共配置 | 不回显密文/明文 token |
|
||||
| PUT | `/config` | 原子保存完整配置 | 版本递增、allowlist 审计 |
|
||||
| POST | `/endpoints/probe` | 测试保存或临时凭据 | 禁重定向、SSRF 防护、结果脱敏 |
|
||||
| GET | `/runtime` | 运行态与指标 | 显示真实 degraded/error |
|
||||
| GET | `/events` | 复合筛选分页 | 稳定排序;用户名/邮箱/API Key 名称分列 |
|
||||
| GET | `/events/:id` | 事件详情 | 脱敏预览、归一结果和派生 issue_summaries |
|
||||
| DELETE | `/events/:id` | 单条硬删除 | 审计、孤立 job 安全清理 |
|
||||
| POST | `/events/batch-delete` | 按 ID 批量删除 | 限制 ID 数量、事务分批 |
|
||||
| POST | `/events/delete-preview` | 预览筛选删除 | 强制起止时间,返回 count/max_id/hash/token |
|
||||
| POST | `/events/delete-by-filter` | 确认筛选删除 | confirm=true,认证 token/actor/hash,限制 id≤max_id |
|
||||
|
||||
分组选择复用目标项目现有管理员 group 查询 API,不为 Prompt Audit 复制一份分组事实源。若现有 API 不适合轻量选择器,只新增薄的只读适配,并在实现前回写本表。
|
||||
|
||||
建议错误 envelope 继续使用项目管理 API 的统一结构;业务错误码稳定,内部 SQL/Redis/HTTP 错误不得透传。
|
||||
|
||||
## 12. 网关 Handler 路由矩阵
|
||||
|
||||
下表是提案编写时已有 `checkContentModeration` 调用点,实施时应机械替换并由结构测试锁定。路由别名共享相同 Handler,因此测试必须至少覆盖主路由与每类 alias。
|
||||
|
||||
| 协议/入口 | 路由 | 现有 Handler 文件/方法 | Stage | 拒绝构造器 |
|
||||
| --- | --- | --- | --- | --- |
|
||||
| Anthropic Messages | `POST /v1/messages` | `gateway_handler.go: Messages` 或 `openai_gateway_handler.go: Messages` | http | Anthropic error helper |
|
||||
| OpenAI Responses | `POST /v1/responses`、`/responses`、`/backend-api/codex/responses` 及 subpath | `gateway_handler_responses.go: Responses` 或 `openai_gateway_handler.go: Responses` | http | Responses/OpenAI helper |
|
||||
| OpenAI Chat Completions | `POST /v1/chat/completions`、`/chat/completions` | `gateway_handler_chat_completions.go: ChatCompletions` 或 `openai_chat_completions.go: ChatCompletions` | http | Chat/OpenAI helper |
|
||||
| Gemini Generate/Stream | `POST /v1beta/models/*modelAction` | `gemini_v1beta_handler.go: GeminiV1BetaModels` | http | Google error helper |
|
||||
| OpenAI Images | `POST /v1/images/generations`、`/v1/images/edits` | `openai_images.go: Images` | http | OpenAI helper |
|
||||
| Grok image/video 文本请求 | images/videos 路由 | `grok_media.go: handleGrokMedia` | http | OpenAI helper |
|
||||
| Responses WebSocket 首轮 | `GET /v1/responses`、`/responses`、`/backend-api/codex/responses` | `openai_gateway_handler.go: ResponsesWebSocket` | first_turn | close 4403/1013 |
|
||||
| Responses WebSocket 后续轮次 | 每个 `response.create` | 同上 BeforeRequest/turn callback | subsequent_turn | close 4403/1013 |
|
||||
|
||||
实施时还必须从 `backend/internal/server/routes/gateway.go` 枚举所有携带用户文本的新增/旁路入口,重点复核:
|
||||
|
||||
- `/v1/images/generations/async`、`/v1/images/edits/async`。
|
||||
- `/v1/images/batches` 及 batch item 的实际提交入口。
|
||||
- Grok video generation/edit/extension。
|
||||
- 任何不经过上述公共 Handler 的内部转发、兼容 alias 或后续新增路由。
|
||||
|
||||
对额外入口有两种合法结论:接入 Coordinator;或证明它已在上游公共 Handler 处检查且不会二次收费/二次扫描。结论和测试必须加入路由矩阵,不能静默跳过。
|
||||
|
||||
接入位置不变量:鉴权、body limit、基本 JSON/model 校验之后;账号选择、用户/账号并发 slot、订阅/余额预扣、usage 写入、上游拨号和 SSE 首字节之前。
|
||||
|
||||
## 13. HTTP、SSE、WebSocket 处理细节
|
||||
|
||||
| 情况 | HTTP/SSE | WS close | reason/code |
|
||||
| --- | ---: | ---: | --- |
|
||||
| Prompt Block | 403 | 4403 | `prompt_guard_blocked` |
|
||||
| Guard Unavailable | 503 | 1013 | `prompt_guard_unavailable` |
|
||||
| Guard Invalid response | 503 | 1013 | `prompt_guard_invalid_response` |
|
||||
|
||||
- HTTP/SSE 必须保留各协议 envelope,不能所有协议统一成 Gin `{"error":"..."}`。
|
||||
- OpenAI Chat/Responses 在 error 对象添加稳定 `code`;Claude 保留 permission_error/api_error type 并添加可选 `code`。
|
||||
- Gemini 保留数值 HTTP `error.code` 和 canonical status,只在 `google.rpc.ErrorInfo.reason` 放稳定代码;metadata 仅 request_id。
|
||||
- SSE 在 Guard 结果前不得写 status/header/data/comment/keepalive;否则无法返回 403/503。
|
||||
- WS 握手本身没有 Prompt,不扫描。首个 `response.create` 在任何本轮资源/上游副作用前扫描。
|
||||
- 后续每个 `response.create` 重新提取本轮输入并标记 `subsequent_turn`。
|
||||
- WS close reason 长度必须在协议限制内,只使用稳定短码;详细内部错误只进脱敏指标/日志。
|
||||
- Legacy moderation 同时 Block 时,继续使用其原错误/close 行为和文案。
|
||||
|
||||
## 14. SQL 和 Repository 注意事项
|
||||
|
||||
### 14.1 Migration
|
||||
|
||||
- PostgreSQL migration 是事实源;不要复制源仓库的 `aicodex_` 前缀。
|
||||
- 表名固定 `prompt_audit_jobs`、`prompt_audit_events`。
|
||||
- 所有状态、计数和非负值加 CHECK;JSONB 加可接受类型检查更佳。
|
||||
- `events.job_id ON DELETE CASCADE`;user/api_key/group 外键 `ON DELETE SET NULL`。
|
||||
- `username_snapshot`、`user_email_snapshot`、`api_key_name_snapshot` 与 group name 快照分列保留,以免主体删除后事件无法复核;沿用现有管理员权限和数据保留策略。
|
||||
- 不新增 raw_prompt、raw_request、request_body、payload、token、authorization、guard_response_body 等列。
|
||||
- 索引名全库唯一;先检查 migration 事实源,避免只在开发库检查。
|
||||
|
||||
### 14.2 原子领取
|
||||
|
||||
`FOR UPDATE SKIP LOCKED` 必须在同一短事务中选择并更新为 processing。事务内不要调用 Redis、Guard 或日志网络 sink。每次 claim 后立即提交,长工作在事务外执行。
|
||||
|
||||
### 14.3 租约和重试
|
||||
|
||||
- attempts 在成功 claim 时递增,而不是失败时递增。
|
||||
- 每个必要分片前刷新租约,并以 processing 状态和本次 claim_version 作为条件。
|
||||
- 401/403、严格解析错误不可重试;429、5xx、连接、超时可重试。
|
||||
- 建议退避 5s、30s、2m,上限 5m并加小 jitter;测试使用 fake clock。
|
||||
- reclaim 批次有上限并按时间/id 稳定排序,防止全表锁和饥饿。
|
||||
|
||||
### 14.4 事件和任务事务
|
||||
|
||||
- 异步成功:event insert 与 job done 应在单事务完成。
|
||||
- 同步:创建 blocking/done job 与可选 event 在单事务完成,但失败不改变门禁结果。
|
||||
- store_pass_events=false 时仍可保存 done job 的最小脱敏执行记录;若最终决定不保存 Pass job,必须回写 schema、runtime 计数和清理规格。
|
||||
- 删除 event 后只删除无事件引用且非 processing 的孤立 job;并 best-effort 删除 Redis key。
|
||||
|
||||
### 14.5 查询与删除
|
||||
|
||||
- 列表使用参数化 SQL、白名单排序字段和稳定 `created_at DESC, id DESC`。
|
||||
- 时间过滤明确采用 UTC 存储、API ISO-8601,并定义边界包含性。
|
||||
- delete-preview 在同一数据库快照得到 count 和 snapshot_max_id,对 canonical JSON filter + max_id 计算 SHA-256;字段顺序、空值和时区必须规范化。
|
||||
- 使用 SecretEncryptor 认证加密 `{filter_hash,snapshot_max_id,admin_id,issued_at,expires_at}`,返回默认 5 分钟有效的 confirmation_token。
|
||||
- delete-by-filter 解密并校验 actor/expiry/hash,要求同一筛选、`confirm=true` 和强制时间范围,查询强制 `id <= snapshot_max_id` 后分批提交;预览后的新事件不可被删除。
|
||||
|
||||
## 15. Guard Client 和出站安全
|
||||
|
||||
请求固定发送到规范化 `{base_url}/v1/chat/completions`,默认模型 `sileader/qwen3guard:0.6b`,role=user、temperature=0、max_tokens=64、seed=42。
|
||||
|
||||
保存、probe 和实际扫描必须走同一校验/Transport:
|
||||
|
||||
- 只允许 http/https;禁止 userinfo、query、fragment。
|
||||
- 禁止 metadata、link-local、multicast、unspecified、保留地址。
|
||||
- 公网强制 HTTPS;HTTP 只允许显式受控的 localhost/私网开发场景。
|
||||
- DNS 解析结果和真正 Dial 的 IP 都检查,防 DNS rebinding。
|
||||
- 不跟随 3xx;响应体最多 256 KiB。
|
||||
- 独立连接池和 Dial/TLS/ResponseHeader timeout;所有分片/故障切换仍受外层总 deadline。
|
||||
- 日志只写 endpoint ID、HTTP status、error_code、latency,不写完整 URL、query、header 或原始 response body。
|
||||
- 分片日志只写 chunk_index/total/chars、input_chars/limit、endpoint ID、action、latency 和错误码,不写 chunk 或内部优先级分隔符。
|
||||
|
||||
九类 scanner ID/展示名必须稳定:Violent、Non-violent Illegal Acts、Sexual Content or Sexual Acts、PII、Suicide & Self-Harm、Unethical Acts、Politically Sensitive Topics、Copyright Violation、Jailbreak。
|
||||
|
||||
## 16. 前端状态与凭据处理
|
||||
|
||||
建议 viewModel 分成:
|
||||
|
||||
```text
|
||||
serverSnapshot # 最近一次后端公共配置
|
||||
draft # 可编辑非敏感配置
|
||||
endpointSecrets # 仅当前会话内的新增/替换 token
|
||||
loadState # config/runtime/groups/events 独立状态
|
||||
probeStateByID # 节点探测进度和脱敏结果
|
||||
eventQuery # canonical filter + page
|
||||
deletePreview # count + max_id + filter_hash + confirmation_token + filter snapshot
|
||||
issueSummaries # 后端从事件事实派生的只读风险展示项
|
||||
```
|
||||
|
||||
规则:
|
||||
|
||||
- `endpointSecrets` 不进入 Pinia 持久化、localStorage、sessionStorage、URL、console 或错误追踪 breadcrumb。
|
||||
- 保存成功后立即清空已提交 secret;失败时可以留在内存草稿供用户修正,但离开页面/卸载必须清空。
|
||||
- 编辑已保存节点时 token 输入默认空,使用 `has_token/token_status` 表示存在性。
|
||||
- “清除 API Key”使用独立明确动作设置 `clear_token=true`,不能把输入框空值当清除。
|
||||
- dirty 比较忽略后端时间戳,但包含 clear/replace 意图;保存返回后以 Public DTO 重建 snapshot。
|
||||
- config/runtime/groups/events 独立失败,不能一个 500 让整页白屏。
|
||||
- `blocking_enabled` 从 false → true 必须二次确认;关闭 enabled 同时把 draft blocking 设 false。
|
||||
- 删除预览与 filter snapshot、snapshot_max_id、confirmation_token 绑定;任何筛选变化立即废弃旧 filter_hash/token。
|
||||
- 用户名、邮箱和 API Key 名称使用不同字段/复制按钮;空值显示明确 fallback,不用邮箱冒充用户名。
|
||||
- IssueSummary 展示 category、title/description、severity/action、scanner、score 和脱敏 evidence,禁止从 evidence 重建命中原文。
|
||||
- 窄屏表格提供可读替代布局,Dialog 有 focus trap/return focus,所有控件有中英文可访问名称。
|
||||
|
||||
## 17. PR/提交切片策略
|
||||
|
||||
每个阶段应可单独评审、测试和回滚,建议五组 PR:
|
||||
|
||||
1. **数据与核心契约**:migration、types/config/snapshot/Qwen parser、Repository 及测试;无路由接入。
|
||||
2. **异步引擎**:出站安全、Redis payload、enqueue、Worker、runtime;功能默认 off。
|
||||
3. **管理闭环**:admin API、独立页面、路由/Sidebar/i18n;仍不启用 blocking。
|
||||
4. **Coordinator 与同步门禁**:统一接入、HTTP/SSE/WS、无副作用断言、Legacy 回归。
|
||||
5. **灰度与运维**:指标、告警、canary 泄露检查、运行手册和阈值登记。
|
||||
|
||||
不要在同一 PR 混入无关的 ContentModeration 重构、全局 Handler 重写、前端框架升级或数据库清理。若为接线必须改现有文件,变更应机械、薄且有前后行为测试。
|
||||
|
||||
## 18. 五个待确认事项的决策门
|
||||
|
||||
| 事项 | 默认建议 | 必须在何时确认 | 未确认时行为 |
|
||||
| --- | --- | --- | --- |
|
||||
| 源基线标识 | 专用 commit/tag | PR 1 前 | 不开始移植 |
|
||||
| 自动保留期 | 第一版只安全删除 | migration 冻结前 | 不加自动清理 |
|
||||
| 双引擎并行/串行 | 并行 | PR 4 前做 benchmark/race | 可先串行但保留优先级 |
|
||||
| 额外文本入口 | routes 自动枚举 | PR 4 接线前 | 结构测试失败 |
|
||||
| blocking 阈值 | 运营按 async 数据登记 | 生产 blocking 前 | 只允许 off/async |
|
||||
|
||||
## 19. 常见错误
|
||||
|
||||
- 直接把 Qwen3Guard 加进 `ContentModerationService`,导致配置、表和副作用混用。
|
||||
- 直接复制源 Ent/React/Caddy 代码,形成重复基础设施或目标项目无法维护的适配壳。
|
||||
- 先把 job 设 queued 再写 Redis,造成 Worker 抢到无 payload 任务。
|
||||
- 把 `redacted_preview` 当作可重试扫描正文;这会产生错误分类且破坏完整覆盖。
|
||||
- 用 byte 长度切中文/emoji,或只扫描第一片后返回 Allow。
|
||||
- 把 Guard 401/403/invalid_response 当 Safe 或无限切节点。
|
||||
- SSE 已写 200/首字节后才运行 Guard。
|
||||
- WS 只检查首轮,不检查后续 `response.create`。
|
||||
- Prompt 拒绝发生在账号选择、并发 slot、预扣或上游拨号之后。
|
||||
- Public DTO 复用 Storage DTO,靠前端“不显示”隐藏 token。
|
||||
- 日志记录请求 body、Guard 原始响应、完整 Base URL/query 或 Redis value。
|
||||
- 配置 reload 失败时清空 last-known-good,或冷启动失败时伪装为 off/healthy。
|
||||
- 按筛选删除没有强制时间范围、预览 Hash 或筛选变化失效。
|
||||
- 为迁移方便重命名/迁移现有 `content_moderation_logs` 或改变 `/admin/risk-control`。
|
||||
|
||||
## 20. Definition of Done
|
||||
|
||||
只有全部成立才算实现完成:
|
||||
|
||||
- 源 commit/tag/patch 已冻结并有可验证 SHA-256。
|
||||
- 三个 specs 的每个 Requirement 都在 `verification.md` 有测试/SQL/日志/截图证据。
|
||||
- Prompt Audit 默认 off;off 时所有外部协议、现有内容审核响应和副作用与升级前一致。
|
||||
- async 失败不改变客户端状态、响应体、计费和上游调用次数。
|
||||
- blocking 的 Block/Unavailable/Invalid 在 HTTP/SSE/WS 映射正确,且账号选择、计费、上游均为 0。
|
||||
- 所有现有用户文本路由和 alias 都有 Coordinator 覆盖证据。
|
||||
- 两张新表、Redis metadata、日志、API、浏览器状态和截图均未出现 canary Prompt/token。
|
||||
- 风险详情拥有确定性 issue_summaries,用户名/邮箱/API Key 名称可分别复核复制,逐分片日志只含安全元数据。
|
||||
- 多 Worker、多实例配置失效、租约回收和 graceful shutdown 测试通过。
|
||||
- 原 RiskControl 页面、关键词、Hash、邮件、自动封号和内容审核记录回归通过。
|
||||
- 后端 unit/race/integration、前端 lint/typecheck/Vitest、生产 build 和 OpenSpec strict validate 全部通过。
|
||||
- 已完成 async 灰度观测;blocking 阈值、告警、值班步骤和一键回滚已由责任人签字确认。
|
||||
@@ -0,0 +1,51 @@
|
||||
## Why
|
||||
|
||||
当前项目的“风控中心”只提供基于 OpenAI Moderations 的内容审核,异步观察依赖进程内队列,且没有 aicodex-api 已具备的持久任务队列、短期敏感载荷存储、Qwen3Guard 分类、同步 fail-closed 门禁和独立提示词事件工作台。直接替换或扩写现有内容审核会混淆两种风险模型,并可能改变关键词、Hash、邮件和自动封号等既有行为,因此需要以并列、默认关闭的独立能力引入。
|
||||
|
||||
本变更以 `/Users/mt/code/mt-ai/aicodex/aicodex-api` 当前磁盘实现为功能参考基线,把其中与目标项目实际协议入口相适配的提示词输入审计能力迁入 sub2api,同时保持现有 OpenAI 兼容接口、内容审核页面、数据库记录和错误语义不变。
|
||||
|
||||
## What Changes
|
||||
|
||||
- 新增独立的 OpenAI 兼容提示词审计引擎,审计节点通过 `{base_url}/v1/chat/completions` 调用 Qwen3Guard,并严格解析 `Safety` 与 `Categories`。
|
||||
- 新增三态运行模式:关闭、异步只审计、同步审计并阻止;所有新增开关默认关闭。
|
||||
- 新增 PostgreSQL 持久任务队列、Redis 短 TTL 原文载荷、进程内 Worker、重试退避、processing 租约刷新和滞留任务回收。
|
||||
- 新增脱敏提示词快照、Hash、Unicode 分片、最新用户输入优先和九类 Qwen3Guard 风险分类。
|
||||
- 新增逐分片安全日志、结构化风险摘要,以及用户名、邮箱、API Key 名称分列的管理员复核信息;风险摘要只使用脱敏证据。
|
||||
- 新增同步 fail-closed 门禁,在账号选择、计费检查和上游调用之前完成;覆盖目标项目现有 Chat Completions、Responses、Claude Messages、Gemini、图像/媒体文本入口及 Responses WebSocket 首轮与后续轮次。
|
||||
- 新增独立管理 API、运行态、审计节点探测、事件查询/详情/删除能力和“提示词审计”页面。
|
||||
- 将侧栏现有“风控中心”入口组织为“安全审计”分组;保留原 `/admin/risk-control` 页面和行为,新增 `/admin/prompt-audit` 页面。
|
||||
- 新增安全审计协调器,只负责给两个独立引擎分发同一份可信请求上下文和归并最终阻断结果,不合并配置、风险分类、事件表或副作用。
|
||||
- 复用现有 SettingRepository、Redis Client、SecretEncryptor、管理员鉴权、管理操作审计、请求身份上下文、分页、日志和前端基础组件。
|
||||
- 新增结构化日志、运行指标、路由覆盖测试、无上游副作用断言和敏感信息泄露门禁。
|
||||
- 不删除、不迁移、不重命名现有 `content_moderation_logs`,不改变现有 Moderations 阈值、关键词、Hash、邮件、封号或清理策略。
|
||||
|
||||
## Capabilities
|
||||
|
||||
### New Capabilities
|
||||
|
||||
- `prompt-input-audit`: 定义提示词快照、异步投递、持久任务队列、OpenAI 兼容 Qwen3Guard 扫描、脱敏事件、运行态、配置和事件管理 API。
|
||||
- `prompt-input-guard`: 定义同步阻止模式、跨协议入口覆盖、fail-closed 错误语义、WebSocket 每轮门禁、配置快照和无计费/无上游副作用不变量。
|
||||
- `security-audit-console`: 定义安全审计导航、独立提示词审计页面、节点探测、配置保存、运行态观测、事件筛选/详情/安全删除和响应式可访问体验。
|
||||
|
||||
### Modified Capabilities
|
||||
|
||||
无。仓库当前没有已发布的 OpenSpec capability;现有内容审核行为在本变更中作为兼容基线,不修改其正式需求语义。
|
||||
|
||||
## Impact
|
||||
|
||||
- **后端模块**:新增 `backend/internal/securityaudit/` 垂直模块;现有 Handler 仅增加协调器依赖和接入调用。
|
||||
- **网关入口**:机械替换现有统一内容审核调用点为安全审计协调调用,保持其位于鉴权之后、账号选择/计费/上游之前;WebSocket 保持逐轮检查。
|
||||
- **管理 API**:新增 `/admin/prompt-audit/*`,复用现有管理员鉴权和管理操作审计。
|
||||
- **数据库**:新增 `prompt_audit_jobs`、`prompt_audit_events` 和相应索引;配置存入现有 `settings`,API Key 加密保存;不修改现有内容审核表。
|
||||
- **Redis**:新增短 TTL 提示词载荷和配置失效通知 key/channel;Redis 不可用时异步 Worker 必须显式降级或报错,不得伪装健康。
|
||||
- **前端**:新增 `frontend/src/features/prompt-audit/`,少量修改路由、侧栏和 i18n;原 `RiskControlView.vue` 业务逻辑保持不变。
|
||||
- **兼容性**:没有外部 API breaking change;新能力默认关闭。只有管理员显式开启同步阻止后,适用请求才可能新增 403/503 或 WebSocket 4403/1013 响应。
|
||||
- **安全与隐私**:完整提示词只允许存在于请求内存和 Redis 短 TTL 载荷,不得进入 PostgreSQL、日志、管理 API、前端状态或错误响应;审计节点凭据必须使用现有 SecretEncryptor 加密。
|
||||
- **实施基线风险**:参考仓库当前 `yjb` 分支包含未提交的同步阻止相关改动。开始编码前必须固定源 commit/tag 或保存可审计 diff,避免“完整迁移”范围漂移。
|
||||
|
||||
## Execution References
|
||||
|
||||
- `source-baseline.md`:源仓库状态、dirty 文件和实施前冻结门禁。
|
||||
- `source-feature-map.md`:AICodex 功能到目标 Requirement、代码位置和证据的逐项映射。
|
||||
- `implementation-guide.md`:按文件实施顺序、时序、状态机、API/路由矩阵和常见错误。
|
||||
- `verification.md`:35 条 Requirement 的证据矩阵、协议测试、泄露门禁、灰度阈值和回滚手册。
|
||||
@@ -0,0 +1,146 @@
|
||||
# AICodex Prompt Audit 源基线
|
||||
|
||||
## 1. 基线状态
|
||||
|
||||
本文件记录用于本 change 功能对照的源仓库状态。参考工作区仍可继续变化,但本 change 已通过第 6 节登记的只读 patch bundle 固定实施基线;后续实现只以该冻结包和本 change specs 为依据。
|
||||
|
||||
| 字段 | 值 |
|
||||
| --- | --- |
|
||||
| 源仓库 | `/Users/mt/code/mt-ai/aicodex/aicodex-api` |
|
||||
| 采集时间 | `2026-07-16 20:21:19 CST (+0800)` |
|
||||
| 分支 | `yjb` |
|
||||
| HEAD | `7a50378851a80650cb0c086260b23abeb3469e6b` |
|
||||
| 工作区 | dirty |
|
||||
| 已跟踪差异 | 38 files changed, 1306 insertions(+), 227 deletions(-) |
|
||||
| 未跟踪范围 | Prompt Guard 实现/测试 6 个文件,加 1 个 OpenSpec change 目录 |
|
||||
| 冻结状态 | **已用只读 patch bundle 冻结并在 detached worktree 恢复验证** |
|
||||
|
||||
当前 HEAD 只代表已提交历史,不能单独代表要迁移的完整功能。同步 fail-closed Guard、出站安全校验、WebSocket/路由顺序测试以及相应 OpenSpec 当前存在于未提交或未跟踪状态。因此,本 change 的临时功能参考是“上述 HEAD + 采集时磁盘工作区”,最终行为权威仍是本 change 的 specs。
|
||||
|
||||
## 2. 与迁移直接相关的已跟踪修改
|
||||
|
||||
### 后端入口与启动接线
|
||||
|
||||
- `ai-gateway/cmd/aicodex/main.go`
|
||||
- `ai-gateway/internal/controller/prompt_audit.go`
|
||||
- `ai-gateway/internal/router/relay-router.go`
|
||||
- `ai-gateway/internal/router/video-router.go`
|
||||
- `ai-gateway/internal/relay/ws_responses.go`
|
||||
- `ai-gateway/internal/gatewayadapter/transport/anthropic.go`
|
||||
- `ai-gateway/internal/gatewayadapter/transport/gemini.go`
|
||||
- `ai-gateway/internal/gatewayadapter/transport/jimeng.go`
|
||||
- `ai-gateway/internal/gatewayadapter/transport/kling.go`
|
||||
- `ai-gateway/internal/gatewayadapter/transport/midjourney.go`
|
||||
- `ai-gateway/internal/gatewayadapter/transport/openai.go`
|
||||
- `ai-gateway/internal/gatewayadapter/transport/suno.go`
|
||||
- `ai-gateway/internal/gatewayadapter/transport/task.go`
|
||||
|
||||
### Prompt Audit 核心
|
||||
|
||||
- `ai-gateway/internal/service/promptaudit/client.go`
|
||||
- `ai-gateway/internal/service/promptaudit/config.go`
|
||||
- `ai-gateway/internal/service/promptaudit/enqueue.go`
|
||||
- `ai-gateway/internal/service/promptaudit/openai_client.go`
|
||||
- `ai-gateway/internal/service/promptaudit/probe.go`
|
||||
- `ai-gateway/internal/service/promptaudit/qwen3guard.go`
|
||||
- `ai-gateway/internal/service/promptaudit/runtime.go`
|
||||
- `ai-gateway/internal/service/promptaudit/runtime_coverage_test.go`
|
||||
- `ai-gateway/internal/service/promptaudit/types.go`
|
||||
- `ai-gateway/internal/service/promptaudit/worker.go`
|
||||
- 同目录的 config、diagnostics、probe 测试
|
||||
|
||||
### 协议、错误和回归测试
|
||||
|
||||
- `ai-gateway/internal/types/error.go`
|
||||
- `ai-gateway/internal/gatewayadapter/transport/user_concurrency_order_test.go`
|
||||
|
||||
### 控制台和类型
|
||||
|
||||
- `webui/src/api/promptAudit.test.ts`
|
||||
- `webui/src/features/prompt-audit/PromptAuditPage.tsx`
|
||||
- `webui/src/features/prompt-audit/PromptAuditPage.test.tsx`
|
||||
- `webui/src/features/prompt-audit/promptAuditViewModel.ts`
|
||||
- `webui/src/features/prompt-audit/promptAuditViewModel.test.ts`
|
||||
- `webui/src/types/promptAudit.ts`
|
||||
|
||||
### 运行说明
|
||||
|
||||
- `deploy/.env.example`
|
||||
- `docs/constraints/41-ai-readable-logging.md`
|
||||
- `docs/workflows/02-local-dev.md`
|
||||
|
||||
## 3. 必须纳入冻结基线的未跟踪文件
|
||||
|
||||
以下文件不在 HEAD 中,但属于“完整功能必须要有”的关键证据:
|
||||
|
||||
- `ai-gateway/internal/gatewaycore/prompt_guard.go`
|
||||
- `ai-gateway/internal/service/promptaudit/outbound_security.go`
|
||||
- `ai-gateway/internal/service/promptaudit/synchronous_guard.go`
|
||||
- `ai-gateway/internal/service/promptaudit/synchronous_guard_test.go`
|
||||
- `ai-gateway/internal/relay/ws_responses_prompt_guard_order_test.go`
|
||||
- `ai-gateway/internal/router/prompt_guard_order_test.go`
|
||||
- `openspec/changes/add-prompt-audit-synchronous-blocking/`
|
||||
|
||||
不得只执行 `git diff HEAD` 后就声称已冻结,因为普通 diff 不包含这些未跟踪文件。
|
||||
|
||||
## 4. 功能对照优先级
|
||||
|
||||
遇到源实现、源测试和本 change 描述不一致时,按以下顺序决策:
|
||||
|
||||
1. 本 change 的三个 delta specs:目标行为契约。
|
||||
2. 本 change 的 `design.md` 和 `implementation-guide.md`:目标架构与落地约束。
|
||||
3. 冻结后的源测试及源 OpenSpec:功能完整性参考。
|
||||
4. 冻结后的源实现:算法、边界和交互参考。
|
||||
5. 当前已提交 HEAD:历史参考。
|
||||
|
||||
目标项目不得复制源仓库的 Ent、Caddy/gatewaycore、React 或全局 option 依赖;只迁移可以被规格和测试证明的行为。
|
||||
|
||||
## 5. 实施前冻结步骤
|
||||
|
||||
在源仓库所有者确认工作区内容属于迁移基线后,选择一种方式:
|
||||
|
||||
### 方案 A:专用 commit/tag(推荐)
|
||||
|
||||
1. 在源仓库专用分支提交与 Prompt Audit/Guard 有关的已跟踪和未跟踪文件。
|
||||
2. 运行源模块及路由/WS 顺序测试。
|
||||
3. 创建不可移动 tag,或记录完整 commit SHA。
|
||||
4. 把最终标识和测试结果回写本文件。
|
||||
|
||||
### 方案 B:只读 patch 包
|
||||
|
||||
1. 生成 tracked diff。
|
||||
2. 使用能够包含未跟踪文件的归档或补丁流程补齐第 3 节文件。
|
||||
3. 生成文件清单和 SHA-256;在干净临时目录中恢复并运行测试。
|
||||
4. 把 patch 路径、清单路径和校验和回写本文件。
|
||||
|
||||
禁止把包含真实 API Key、Redis payload、`.env` 私密值或运行日志中的完整 Prompt 放入基线包。
|
||||
|
||||
## 6. 最终冻结登记
|
||||
|
||||
| 字段 | 待填写值 |
|
||||
| --- | --- |
|
||||
| 冻结方式 | 只读 tracked patch + untracked tar archive |
|
||||
| 冻结 commit/tag | base commit `7a50378851a80650cb0c086260b23abeb3469e6b`(detached restore) |
|
||||
| patch/archive 绝对路径 | `/Users/mt/code/mt-ai/sub2api/sub2api-mt/openspec/changes/add-openai-compatible-prompt-audit/source-freeze/` |
|
||||
| manifest SHA-256 | `badab312bf6af4d2c77857a9400381f4da4fbf45722d9f4a6df23bc7005273b6` |
|
||||
| tracked patch SHA-256 | `f751a13cce3f3a73cd60cae3aececcef6e1e76dcec8c551a7a4747f032234d2b` |
|
||||
| untracked archive SHA-256 | `1536e2781703b7620e26f2d08b249431fa5846ad9e32b2e8b0d547c3fa3b3632` |
|
||||
| 冻结人/复核人 | Codex;由恢复后的文件清单、`git diff --check` 和测试命令复核 |
|
||||
| 冻结时间 | `2026-07-16 20:21:19 CST (+0800)` |
|
||||
| 源测试结果 | 恢复副本中 Prompt Audit 核心、router、relay、gateway transport 目标测试全部通过,详见 `source-freeze/MANIFEST.md` |
|
||||
|
||||
## 7. 复核命令
|
||||
|
||||
```bash
|
||||
cd /Users/mt/code/mt-ai/aicodex/aicodex-api
|
||||
git branch --show-current
|
||||
git rev-parse HEAD
|
||||
git status --short
|
||||
git diff --stat
|
||||
git diff --name-only
|
||||
git ls-files --others --exclude-standard
|
||||
cd ai-gateway
|
||||
go test ./internal/service/promptaudit
|
||||
```
|
||||
|
||||
本提案编写时上述模块测试已在当前 dirty 磁盘状态通过;冻结后必须再次执行,并记录最终 commit/patch 校验和、执行目录和完整输出。
|
||||
@@ -0,0 +1,127 @@
|
||||
# AICodex 源功能迁移映射
|
||||
|
||||
## 1. 目的
|
||||
|
||||
本表用于证明“完整功能都必须要有”不是一句笼统目标。每个 AICodex 当前用户可见或运行时能力都必须映射到目标 Requirement、预期代码位置和验证证据;实施中发现新源能力时,先更新本表和相关 spec/tasks,再编码。
|
||||
|
||||
源参考状态见 `source-baseline.md`。只读冻结包已在 detached worktree 中恢复,以下测试在恢复副本执行:
|
||||
|
||||
```text
|
||||
cd /Users/mt/code/mt-ai/aicodex/aicodex-api/ai-gateway
|
||||
go test ./internal/service/promptaudit -count=1
|
||||
ok github.com/mt21625457/aicodex/internal/service/promptaudit 2.081s
|
||||
|
||||
go test ./internal/router ./internal/relay ./internal/gatewayadapter/transport \
|
||||
-run 'PromptGuard|PromptAudit|ConcurrencyOrder' -count=1
|
||||
ok github.com/mt21625457/aicodex/internal/router 1.184s
|
||||
ok github.com/mt21625457/aicodex/internal/relay 2.201s
|
||||
ok github.com/mt21625457/aicodex/internal/gatewayadapter/transport 3.233s
|
||||
```
|
||||
|
||||
这证明冻结包可恢复且源参考测试通过,但不证明目标实现已完成;目标代码和证据仍须逐行补齐。
|
||||
|
||||
## 2. 功能映射
|
||||
|
||||
| # | AICodex 当前能力与源证据 | 目标 OpenSpec 契约 | 目标主要代码 | 验证证据 |
|
||||
| ---: | --- | --- | --- | --- |
|
||||
| 1 | 独立 Prompt Audit 开关、默认关闭;`config.go` | prompt-input-audit:独立且默认关闭;prompt-input-guard:显式三态 | `prompt_config.go`、`coordinator.go` | A01、G01 |
|
||||
| 2 | enabled + blocking_enabled 表达 off/async/blocking;`config.go`、`synchronous_guard.go` | prompt-input-guard:显式启用、即时回滚 | `prompt_config.go`、`prompt_guard.go` | G01、G12 |
|
||||
| 3 | 配置持久化、版本、updated_by/change_summary;`config.go` | prompt-input-guard:版本化快照/CAS;console:可验证保存 | `prompt_config.go` | G10、C06、C10 |
|
||||
| 4 | token 加密、空值保留、替换、clear;`config.go`、`config_test.go` | prompt-input-audit:凭据安全;console:池管理/保存 | `prompt_config.go`、`prompt_handler.go` | A03、C03、C06 |
|
||||
| 5 | OpenAI-compatible endpoint、Qwen3Guard 默认模型;`openai_client.go` | prompt-input-audit:OpenAI 兼容节点 | `prompt_qwen3guard.go` | A02 |
|
||||
| 6 | Base URL 规范化,固定 `/v1/chat/completions`;`openai_client.go` | prompt-input-audit:OpenAI 兼容节点/出站安全 | `prompt_qwen3guard.go`、`prompt_outbound_security.go` | A02、A03 |
|
||||
| 7 | `/v1/models` readiness + scan fallback probe;`openai_client.go`、`probe.go` | prompt-input-audit:管理员探测;console:真实探测 | `prompt_qwen3guard.go`、`prompt_handler.go` | A02、C03 |
|
||||
| 8 | probe 对话框阶段、结果、状态/耗时/错误;`PromptAuditPage.tsx` | console:完整审计池和真实探测 | `features/prompt-audit/components` | C03 |
|
||||
| 9 | Guard SSRF、DNS/Dial 复检、重定向/响应上限;`outbound_security.go` | prompt-input-audit:凭据和出站地址安全 | `prompt_outbound_security.go` | A03 |
|
||||
| 10 | Qwen3Guard `Safety/Categories` 解析;`qwen3guard.go` | prompt-input-audit:严格归一 | `prompt_qwen3guard.go` | A08 |
|
||||
| 11 | 九类官方输入风险;`qwen3guard.go`、页面 scanner catalog | prompt-input-audit:九类;console:九类配置 | `prompt_qwen3guard.go`、前端 types/viewModel | A08、C04 |
|
||||
| 12 | Safe/Controversial/Unsafe → Allow/Warn/Block;`openai_client.go`、`normalize.go` | prompt-input-audit:严格归一;guard:fail-closed | `prompt_qwen3guard.go`、`prompt_scanner.go` | A08、G06 |
|
||||
| 13 | 高风险 Controversial 提升、未知 Unsafe 保持 Block;`openai_client.go` | prompt-input-audit:严格归一 | `prompt_qwen3guard.go` | A08 |
|
||||
| 14 | Chat/Responses/Claude 多协议快照;`snapshot.go`、`multiprotocol.go` | prompt-input-audit:按协议提取 | `prompt_snapshot.go` | A04 |
|
||||
| 15 | Gemini、图片/媒体等 transport 传递提示词上下文;gatewayadapter changes | prompt-input-audit:所有文本入口;guard:路由覆盖 | `prompt_snapshot.go`、各 Handler 薄接线 | A04、G04 |
|
||||
| 16 | Responses WS 首轮和后续帧;`ws_responses.go`、顺序测试 | prompt-input-guard:每个 response.create 门禁 | `openai_gateway_handler.go` 薄接线 | G08 |
|
||||
| 17 | 最新用户输入优先;`snapshot.go` | prompt-input-audit:提取/Unicode 分片 | `prompt_snapshot.go`、`prompt_scanner.go` | A04、A09 |
|
||||
| 18 | rune input_limit 完整分片;`openai_client.go` | prompt-input-audit:Unicode 完整分片 | `prompt_scanner.go` | A09 |
|
||||
| 19 | 多片最严重聚合、证据 metadata/去重、Block 早停;`openai_client.go` | prompt-input-audit:分片;guard:共享预算 | `prompt_scanner.go`、`prompt_guard.go` | A09、G05 |
|
||||
| 20 | 每片前刷新 processing lease;`openai_client.go`、`worker.go` | prompt-input-audit:Worker/Unicode 分片 | `prompt_worker.go` | A07、A09 |
|
||||
| 21 | scan_chunk_started/completed/failed/aggregated 日志;`openai_client.go` | prompt-input-audit:Unicode 分片;guard:可观测 | `prompt_logging.go`、`prompt_scanner.go` | A09、G11 |
|
||||
| 22 | Prompt hash、脱敏 preview、敏感模式处理;`snapshot.go` | prompt-input-audit:不可恢复快照 | `prompt_snapshot.go` | A05 |
|
||||
| 23 | 完整 scan text 使用 Redis 30 分钟 TTL;`payload_store.go` | prompt-input-audit:持久任务 + Redis TTL | `prompt_payload_store.go` | A06 |
|
||||
| 24 | 异步 enqueue、范围/容量检查;`enqueue.go` | prompt-input-audit:异步持久投递 | `prompt_enqueue.go` | A06 |
|
||||
| 25 | PromptAuditJob/Event 持久事实;Ent schema/store | prompt-input-audit:jobs/events | SQL migration、`prompt_repository.go` | A05、A07、A10 |
|
||||
| 26 | 进程内 Worker、可配置数量、Start/Stop;`worker.go` | prompt-input-audit:可靠 Worker | `prompt_worker.go`、`prompt_module.go` | A07 |
|
||||
| 27 | retry/backoff/max attempts;`worker.go` | prompt-input-audit:可靠 Worker | `prompt_worker.go` | A07 |
|
||||
| 28 | processing stale reclaim;`worker.go` | prompt-input-audit:可靠 Worker | `prompt_worker.go`、Repository | A07 |
|
||||
| 29 | runtime queue/Worker/DB/payload/connectivity/heartbeat;`runtime.go` | prompt-input-audit:真实运行态 | `prompt_runtime.go` | A11、C07 |
|
||||
| 30 | config active/expected version 和失效通知;`config.go`、`runtime.go` | prompt-input-guard:版本化热路径快照 | `prompt_config.go`、`prompt_runtime.go` | G10、C07 |
|
||||
| 31 | 同步 evaluator 不依赖 Worker;`synchronous_guard.go` | prompt-input-guard:同步门禁/结果复用 | `prompt_guard.go` | G03、G09 |
|
||||
| 32 | 总 deadline、ordered failover、bulkhead;`synchronous_guard.go` | prompt-input-guard:共享预算/故障切换 | `prompt_guard.go` | G05、G06 |
|
||||
| 33 | HTTP fail-closed 403/503;`prompt_guard.go`、router 接线 | prompt-input-guard:HTTP 稳定错误 | Handler helper + OpenAI/Claude code、Gemini ErrorInfo adapter | G03、G07 |
|
||||
| 34 | WS 4403/1013;`ws_responses.go` | prompt-input-guard:每轮 WS 门禁 | Responses WS Handler | G08 |
|
||||
| 35 | 同步结果轻量记录、不重复 Guard;`synchronous_guard.go` | prompt-input-guard:结果复用 | `prompt_guard.go`、Repository | G09 |
|
||||
| 36 | Guard metrics Allow/Flag/Block/Unavailable/timeout/failover/bulkhead;`synchronous_guard.go`、`runtime.go` | prompt-input-guard:可观测;console:运行态 | `prompt_runtime.go`、metrics adapter | G11、C07 |
|
||||
| 37 | 事件列表/详情、复合筛选;`store.go`、controller | prompt-input-audit:查询事件;console:列表详情 | `prompt_repository.go`、`prompt_handler.go`、前端 | A12、C08 |
|
||||
| 38 | 用户名/邮箱分别展示和复制;probe-dialog change + controller/UI tests | prompt-input-audit:分列身份快照;console:复核身份 | Request/snapshot、event DTO、前端详情 | A04、A10、C08 |
|
||||
| 39 | scanner evidence、Guard policy、结构化 issue summaries;`issue_summary.go` | prompt-input-audit:事件/风险摘要;console:具体风险 | `prompt_issue_summary.go`、event DTO | A10、C08 |
|
||||
| 40 | 单条/批量硬删除;controller/store | prompt-input-audit:安全删除;console:防误操作 | Repository/Admin Handler/前端 | A12、C09 |
|
||||
| 41 | delete preview + canonical filter hash + confirm;filter helper | prompt-input-audit:安全删除;console:防误操作 | Repository/Admin Handler/前端,增加 max_id/认证 token | A12、C09 |
|
||||
| 42 | 配置、probe、删除的管理审计;controller/router tests | console:管理员操作审计 | `prompt_handler.go` + 现有 audit | C10 |
|
||||
| 43 | 独立控制台、运行概览、池/策略/事件/保存栏;`PromptAuditPage.tsx` | console:独立工作区 | `frontend/src/features/prompt-audit/` | C01、C02 |
|
||||
| 44 | dirty snapshot、统一保存、重置;页面/viewModel | console:工作区/可验证保存 | 前端 viewModel/page | C02、C06 |
|
||||
| 45 | all/selected group、搜索、stale group;页面/config | prompt-input-audit:范围;console:范围配置 | config + 前端 selector | C04 |
|
||||
| 46 | endpoint 新增/编辑/启停/删除、参数对话框;页面 | console:审计池管理 | 前端 components | C03 |
|
||||
| 47 | blocking 二次确认和保存栏开关联动;页面 | console:开启风险确认 | 前端 viewModel/page | C05 |
|
||||
| 48 | 事件技术/具体风险/结构化返回 tabs 和 JSON 查看;页面 | console:可复核详情 | 前端 detail components | C08 |
|
||||
| 49 | 响应式、可访问状态、页面测试;redesign change | console:响应式/可访问/i18n | 前端 + i18n | C11 |
|
||||
| 50 | AI 可读稳定日志和敏感字段约束;logging.go/constraints | prompt-input-guard:可观测且不泄密 | `prompt_logging.go` | G11 |
|
||||
|
||||
## 3. 架构适配而非逐行复制
|
||||
|
||||
以下差异是目标架构适配,不是功能删减:
|
||||
|
||||
| AICodex 实现细节 | sub2api 目标实现 | 等价性理由/门禁 |
|
||||
| --- | --- | --- |
|
||||
| Ent PromptAuditJob/Event | PostgreSQL migration + `database/sql` | 目标项目以 SQL migration 为 schema 事实源;字段和行为由 A05/A07/A10/A12 验证 |
|
||||
| 表/对象可能带 AICodex 命名 | `prompt_audit_jobs/events` | 不复制 `aicodex_` 前缀;管理能力不变 |
|
||||
| `PromptAuditConfigJSON` option | settings `prompt_audit_config` | 复用目标 SettingRepository,Public/Storage DTO 行为不变 |
|
||||
| AICodex secret helper | 现有 `SecretEncryptor` | A03 canary 和加密往返证明 |
|
||||
| React/Ant Design 页面 | Vue 3 既有组件体系 | C01-C11 以行为和可访问性验收,不按框架验收 |
|
||||
| `/api/prompt-audit` | `/admin/prompt-audit` | 复用目标 AdminAuth/管理审计;API 能力一一对应 |
|
||||
| token/channel/group 字符串 | API key/group/provider 可信 ID + 快照 | 使用目标身份域,保留查询/复核能力 |
|
||||
| 6068/9068 双端口一致性 | `/v1`、root alias、`/backend-api/codex` 等目标路由一致性 | G04 以目标实际 routes 自动枚举,不复制不存在的端口拓扑 |
|
||||
| 源 queued 后再写 payload 的竞态 | staging → Redis SET EX → queued | 是可靠性增强;A06/A07 证明 Worker 不提前领取 |
|
||||
| 源进程内唤醒队列 + DB 事实 | PostgreSQL 原子 claim + 递增 claim_version fencing + 进程内 Worker | 支持多实例并防旧 Worker 覆盖,无功能损失;A07 并发测试证明 |
|
||||
| 源 MemoryRepository | 只作为目标测试 fake,不作为生产 fallback | 生产需要持久任务;依赖失败由 A11 显示 degraded,不伪装成功 |
|
||||
| `scan_url`/旧 llm_guard 协议兼容 | 只接受 Base URL + OpenAI compatible | 目标是新增 setting、无旧 Prompt Audit 配置;A02 明确禁止旧协议 |
|
||||
| 源旧 strategy 迁移 | 第一版仅 `priority`,其他值拒绝 | 目标无历史 Prompt config;G01/配置测试保证确定性 |
|
||||
| endpoint `weight` 兼容展示字段 | 显式数组顺序作为 priority | 源当前只允许 priority,扫描代码未使用 weight 做选择;目标去除无效歧义,故障切换能力由 G06 证明 |
|
||||
| endpoint `policy_id/tenant_id` 历史兼容输入 | Qwen 结果固定 policy_id/version,event 持久化 | 当前 Qwen 请求不发送这两个 endpoint 字段;目标保留实际策略结果而不暴露无效输入 |
|
||||
| 源 env 默认配置 | settings 管理页面初始化默认值 | 目标配置事实源是 settings;默认 off 和完整可配置性由 A01/C03/C06 证明 |
|
||||
| 源 event API 查询时解析用户 | 事件保存用户名/邮箱/API Key 名称分列快照 | 删除主体后仍可复核;访问与保留沿用现有管理员政策 |
|
||||
| 源 `issue_summaries` 由 evidence 派生 | 目标同样派生,不新增数据库列 | 防止双份风险事实漂移;A10/C08 golden 测试证明 |
|
||||
|
||||
## 4. 源专属能力的明确处理
|
||||
|
||||
以下内容不作为目标运行功能移植,但必须明确原因:
|
||||
|
||||
- AICodex 旧 `/v1/scan/prompt` 和 llm_guard 配置迁移:目标项目从未发布 Prompt Audit,无历史配置需要兼容;目标只实现当前 OpenAI-compatible Qwen3Guard 行为。
|
||||
- AICodex Caddy/gatewaycore、6068/9068 端口和 channel dispatch:目标使用 Gin Handler、目标账号调度和目标路由 alias;以 G04/G03 证明等价接入顺序。
|
||||
- AICodex React/旧 deprecated 页面:只迁移当前管理行为到 Vue 独立 feature,不同时维护两套前端。
|
||||
- AICodex 产品特有 transport:目标只覆盖目标项目实际存在且可触发模型的文本入口;`implementation-guide.md` 的路由枚举是硬门禁。
|
||||
- 输出审核、Redact、人工审批和申诉:当前迁移范围是用户输入 Prompt Audit/Guard,且本 change 明确列为 Non-Goals;不得把源旧 LLM Guard 的 `Redact` 兼容文案误当成当前 Qwen 输入审计功能。
|
||||
|
||||
如果实施评审发现上述任一项实际上在目标项目有已发布数据或用户依赖,必须把它从本节移回第 2 节,新增 Requirement/Scenario 后才能继续。
|
||||
|
||||
## 5. 完整性复核步骤
|
||||
|
||||
每次源基线或目标设计变化后执行:
|
||||
|
||||
1. 对源 `internal/service/promptaudit`、Prompt Audit controller/router、WS/transport 和当前前端目录重新列出文件/公开符号。
|
||||
2. 对源主 spec 和所有未归档 Prompt Audit changes 提取 Requirement/Scenario。
|
||||
3. 为新发现功能在第 2 节新增一行;若无目标 Requirement,先更新 specs。
|
||||
4. 检查每行同时有目标代码位置和 `verification.md` ID。
|
||||
5. 检查第 3/4 节每项确实是架构适配或源专属,而不是为了缩小实现范围。
|
||||
6. 冻结时把最终源 commit/tag/patch SHA-256 写入 `source-baseline.md`。
|
||||
7. 实现完成后把每行的计划证据替换为实际测试名/CI artifact 链接。
|
||||
|
||||
本表没有“以后再做”状态。除第 4 节经解释的源专属项外,第 2 节任一行没有通过证据都表示“完整迁移”未完成。
|
||||
@@ -0,0 +1,52 @@
|
||||
# AICodex Prompt Audit source freeze manifest
|
||||
|
||||
- Frozen at: `2026-07-16 20:21:19 CST (+0800)`
|
||||
- Source repository: `/Users/mt/code/mt-ai/aicodex/aicodex-api`
|
||||
- Source branch at capture: `yjb`
|
||||
- Base commit: `7a50378851a80650cb0c086260b23abeb3469e6b`
|
||||
- Freeze method: immutable tracked patch plus untracked tar archive
|
||||
- Restored verification worktree: detached from the base commit, then populated only from the two artifacts below
|
||||
|
||||
## Artifacts
|
||||
|
||||
| Artifact | Size | SHA-256 |
|
||||
| --- | ---: | --- |
|
||||
| `aicodex-prompt-audit-tracked.patch` | 124674 bytes | `f751a13cce3f3a73cd60cae3aececcef6e1e76dcec8c551a7a4747f032234d2b` |
|
||||
| `aicodex-prompt-audit-untracked.tar.gz` | 39342 bytes | `1536e2781703b7620e26f2d08b249431fa5846ad9e32b2e8b0d547c3fa3b3632` |
|
||||
|
||||
The tracked patch contains 38 files with 1306 insertions and 227 deletions. It is applied to the base commit above using `git apply`.
|
||||
|
||||
## Untracked archive entries
|
||||
|
||||
- `ai-gateway/internal/gatewaycore/prompt_guard.go`
|
||||
- `ai-gateway/internal/relay/ws_responses_prompt_guard_order_test.go`
|
||||
- `ai-gateway/internal/router/prompt_guard_order_test.go`
|
||||
- `ai-gateway/internal/service/promptaudit/outbound_security.go`
|
||||
- `ai-gateway/internal/service/promptaudit/synchronous_guard.go`
|
||||
- `ai-gateway/internal/service/promptaudit/synchronous_guard_test.go`
|
||||
- `openspec/changes/add-prompt-audit-synchronous-blocking/.openspec.yaml`
|
||||
- `openspec/changes/add-prompt-audit-synchronous-blocking/design.md`
|
||||
- `openspec/changes/add-prompt-audit-synchronous-blocking/proposal.md`
|
||||
- `openspec/changes/add-prompt-audit-synchronous-blocking/specs/prompt-input-audit/spec.md`
|
||||
- `openspec/changes/add-prompt-audit-synchronous-blocking/specs/prompt-input-guard/spec.md`
|
||||
- `openspec/changes/add-prompt-audit-synchronous-blocking/tasks.md`
|
||||
|
||||
## Restore and verification result
|
||||
|
||||
The artifacts were restored into `/tmp/aicodex-prompt-audit-freeze-7a503788`, a detached worktree at the base commit. `git diff --check` passed.
|
||||
|
||||
The following commands passed against the restored copy:
|
||||
|
||||
```text
|
||||
cd ai-gateway
|
||||
go test ./internal/service/promptaudit -count=1
|
||||
ok github.com/mt21625457/aicodex/internal/service/promptaudit 2.081s
|
||||
|
||||
go test ./internal/router ./internal/relay ./internal/gatewayadapter/transport \
|
||||
-run 'PromptGuard|PromptAudit|ConcurrencyOrder' -count=1
|
||||
ok github.com/mt21625457/aicodex/internal/router 1.184s
|
||||
ok github.com/mt21625457/aicodex/internal/relay 2.201s
|
||||
ok github.com/mt21625457/aicodex/internal/gatewayadapter/transport 3.233s
|
||||
```
|
||||
|
||||
The source worktree remains untouched. The target OpenSpec specs remain authoritative if this frozen implementation differs from the target architecture.
|
||||
+2771
File diff suppressed because it is too large
Load Diff
BIN
Binary file not shown.
@@ -0,0 +1,246 @@
|
||||
## ADDED Requirements
|
||||
|
||||
### Requirement: 提示词审计必须是独立且默认关闭的安全审计引擎
|
||||
系统 SHALL 在现有内容审核之外提供独立的提示词审计引擎。新引擎 MUST 拥有独立配置、运行态、任务、事件和开关,并 MUST 默认关闭;现有 OpenAI Moderations 内容审核的配置、判定、关键词、Hash、邮件、自动封号、日志表和清理行为 MUST NOT 因本能力而改变。
|
||||
|
||||
#### Scenario: 升级后未启用新引擎
|
||||
- **WHEN** 系统完成包含本能力的升级且管理员尚未保存提示词审计配置
|
||||
- **THEN** 所有模型请求 MUST 继续按升级前的内容审核和转发链路执行
|
||||
- **THEN** 系统 MUST NOT 创建提示词审计任务、写入提示词审计事件或调用外部 Guard
|
||||
|
||||
#### Scenario: 两个审计引擎同时启用
|
||||
- **WHEN** 现有内容审核和新增提示词审计都已启用
|
||||
- **THEN** 两个引擎 MUST 使用各自的配置与风险语义独立执行
|
||||
- **THEN** 提示词审计命中 MUST NOT 自动触发现有内容审核的邮件、封号或 Hash 黑名单副作用
|
||||
|
||||
### Requirement: 提示词审计节点必须使用 OpenAI 兼容协议
|
||||
系统 SHALL 仅支持通过 OpenAI 兼容 Chat Completions 接口调用提示词审计节点。节点配置 MUST 支持名称、Base URL、API Key、Model、超时、单片输入上限、启用状态和有序优先级;默认模型 MUST 为 `sileader/qwen3guard:0.6b`。
|
||||
|
||||
#### Scenario: Worker 调用已配置节点
|
||||
- **WHEN** Worker 领取到可处理任务并选择一个启用节点
|
||||
- **THEN** 系统 MUST 向 `{base_url}/v1/chat/completions` 发送请求
|
||||
- **THEN** 请求 MUST 使用 `role=user`、`temperature=0`、确定性的输出限制和管理员配置的模型
|
||||
- **THEN** 系统 MUST NOT 调用旧的 `/v1/scan/prompt` 或 `llm_guard` 专用协议
|
||||
|
||||
#### Scenario: 管理员保存未填写模型的节点
|
||||
- **WHEN** 管理员保存一个 Base URL 有效但 Model 为空的节点
|
||||
- **THEN** 系统 MUST 将节点模型归一为 `sileader/qwen3guard:0.6b`
|
||||
|
||||
#### Scenario: 管理员探测节点
|
||||
- **WHEN** 管理员请求探测一个节点
|
||||
- **THEN** 后端 MUST 使用服务端网络环境执行真实的认证与模型连通性探测
|
||||
- **THEN** 响应 MUST 包含成功状态、稳定错误码、HTTP 状态、耗时、是否可重试和检查时间
|
||||
- **THEN** 响应 MUST NOT 回显 API Key
|
||||
|
||||
### Requirement: 审计节点凭据和出站地址必须受到安全保护
|
||||
系统 MUST 使用现有 SecretEncryptor 加密持久化节点 API Key,并 MUST 对节点地址实施 SSRF、防重定向和响应体大小限制。完整凭据只允许短暂存在于管理员写入请求、前端未持久化输入内存、服务端解密内存和发往 Guard 的 Authorization Header;它们以及 URL query、提示词正文 MUST NOT 出现在日志、错误响应、管理读取响应或前端持久化/调试状态中。
|
||||
|
||||
#### Scenario: 保存带 API Key 的节点
|
||||
- **WHEN** 管理员保存一个包含 API Key 的节点
|
||||
- **THEN** settings 中 MUST 只保存加密密文和是否已配置标记
|
||||
- **THEN** 后续读取配置 MUST 只返回 `has_token=true` 或等价状态
|
||||
|
||||
#### Scenario: 保存受限出站地址
|
||||
- **WHEN** Base URL 使用非 HTTP(S) scheme、包含 userinfo/query/fragment、指向 link-local/元数据/未指定/保留地址,或公网地址使用不安全 HTTP
|
||||
- **THEN** 系统 MUST 拒绝保存和探测该节点
|
||||
- **THEN** 系统 MUST 返回稳定且不包含敏感地址细节的校验错误码
|
||||
|
||||
#### Scenario: 节点返回重定向或超大响应
|
||||
- **WHEN** Guard 返回 HTTP 重定向或超过配置上限的响应体
|
||||
- **THEN** 系统 MUST 不跟随重定向并将响应判定为无效或不可用
|
||||
|
||||
### Requirement: 系统必须按协议提取用户输入提示词快照
|
||||
系统 SHALL 从目标项目所有已支持、包含用户文本的模型入口提取提示词快照。快照 MUST 包含 request ID、user ID、用户名、用户邮箱、API key ID/名称、group ID/名称、provider、endpoint、protocol、model、提示词 Hash、脱敏预览、Unicode 字符数和消息数量;文本审计 MUST 优先扫描最新用户输入,同时完整覆盖需要审计的历史用户文本。
|
||||
|
||||
#### Scenario: 提取 OpenAI Chat Completions 输入
|
||||
- **WHEN** `/v1/chat/completions` 或等价兼容入口包含一个或多个 `role=user` 消息
|
||||
- **THEN** 系统 MUST 提取用户文本内容并把最新用户输入置于扫描顺序最前
|
||||
- **THEN** 系统 MUST 不把 assistant 或 tool 输出当作用户提示词主体
|
||||
|
||||
#### Scenario: 提取 OpenAI Responses 输入
|
||||
- **WHEN** `/v1/responses` 请求使用字符串、消息数组或内容块表达用户输入
|
||||
- **THEN** 系统 MUST 提取其中的用户文本并保留 Responses 协议标识
|
||||
|
||||
#### Scenario: 提取 Claude 和 Gemini 输入
|
||||
- **WHEN** Claude Messages 或 Gemini 兼容入口包含用户角色文本
|
||||
- **THEN** 系统 MUST 提取可审计文本并保留真实 protocol、endpoint 和 model
|
||||
|
||||
#### Scenario: 提取图像或媒体生成提示词
|
||||
- **WHEN** OpenAI Images、Grok 媒体或目标项目其他生成入口包含文本 prompt
|
||||
- **THEN** 新引擎 MUST 审计文本 prompt
|
||||
- **THEN** 新引擎 MUST NOT 把图片二进制、base64 图片或远程图片内容发送给 Qwen3Guard
|
||||
- **THEN** 图片内容审核 MUST 继续由现有内容审核引擎负责
|
||||
|
||||
#### Scenario: 请求没有用户文本
|
||||
- **WHEN** 请求体有效但没有可审计的用户文本
|
||||
- **THEN** 系统 MUST 跳过提示词任务并记录稳定的 skipped reason
|
||||
|
||||
### Requirement: 提示词数据库快照必须脱敏且不可恢复原文
|
||||
系统 SHALL 在写入数据库前计算 SHA-256 Hash 和脱敏裁剪预览。PostgreSQL、结构化日志、管理 API 和前端 MUST NOT 保存或返回完整原始提示词;用于实际扫描的正文只允许保存在请求内存或 Redis 短 TTL 载荷中。
|
||||
|
||||
#### Scenario: 创建异步任务
|
||||
- **WHEN** 系统为用户输入创建异步审计任务
|
||||
- **THEN** `prompt_audit_jobs` MUST 保存 Hash、脱敏预览、字符数、消息数、分列的用户/API Key 展示快照和可关联请求上下文
|
||||
- **THEN** 表中 MUST 不存在 raw_prompt、payload 或等价原文字段
|
||||
|
||||
#### Scenario: 管理员查看事件详情
|
||||
- **WHEN** 管理员打开提示词审计事件详情
|
||||
- **THEN** 页面和 API MUST 只展示脱敏预览、Hash、分类、结构化风险摘要、证据摘要和技术元数据
|
||||
- **THEN** 任何证据片段 MUST 经过脱敏、长度限制并包含不可逆 Hash,而不是完整命中正文
|
||||
|
||||
### Requirement: 异步审计必须使用持久任务和短期 Redis 载荷
|
||||
系统 SHALL 使用 PostgreSQL `prompt_audit_jobs` 作为任务事实源,并使用 Redis 保存默认 30 分钟 TTL 的完整扫描正文。异步任务投递 MUST 不阻塞或改变主模型请求结果。
|
||||
|
||||
#### Scenario: 成功投递异步任务
|
||||
- **WHEN** 提示词审计处于 async_audit、请求在审计范围内且队列未满
|
||||
- **THEN** 系统 MUST 先创建不可被 Worker 领取的 staging 任务
|
||||
- **THEN** 系统 MUST 成功写入 Redis 载荷后再把任务发布为 queued
|
||||
- **THEN** 主请求 MUST 继续进入现有网关链路
|
||||
|
||||
#### Scenario: Redis 载荷写入失败
|
||||
- **WHEN** 数据库任务已创建但 Redis 载荷写入失败
|
||||
- **THEN** 系统 MUST 将任务标记为 failed 或保持可清理的 staging 状态
|
||||
- **THEN** 系统 MUST 输出 `prompt_audit.enqueue_dropped` 和稳定错误码
|
||||
- **THEN** 主模型请求 MUST 不受影响
|
||||
|
||||
#### Scenario: 队列达到容量上限
|
||||
- **WHEN** queued、retry、processing 和 staging 活跃任务达到配置容量
|
||||
- **THEN** 系统 MUST 拒绝创建新的异步任务并记录 `reason=queue_full`
|
||||
- **THEN** 主模型请求 MUST 继续转发
|
||||
|
||||
#### Scenario: 多实例同时争抢最后队列容量
|
||||
- **WHEN** 多个实例并发入队且剩余容量不足以容纳全部请求
|
||||
- **THEN** active count 检查与 staging INSERT MUST 在同一数据库准入锁事务中串行化
|
||||
- **THEN** 已接受的 active jobs MUST NOT 超过该配置快照的 queue_capacity
|
||||
- **THEN** 未获准任务 MUST 按 queue_full 或 queue_admission_busy 丢弃且不影响主请求
|
||||
|
||||
### Requirement: 进程内 Worker 必须可靠消费持久任务
|
||||
系统 SHALL 在主服务进程内启动可配置数量的 Worker。多实例 Worker MUST 通过 PostgreSQL 原子领取任务,并为每次领取生成单调递增的 claim version fencing token;租约刷新、事件提交和终态更新 MUST 校验该 token。系统还 MUST 支持重试退避、processing 租约刷新、滞留任务回收、最大尝试次数和优雅关闭。
|
||||
|
||||
#### Scenario: 多 Worker 并发领取任务
|
||||
- **WHEN** 多个进程或 Worker 同时寻找可执行任务
|
||||
- **THEN** 每个任务 MUST 只被一个 Worker 原子领取
|
||||
- **THEN** 领取过程 MUST 使用数据库行锁/条件更新或等价的无重复执行机制
|
||||
|
||||
#### Scenario: 已回收的旧 Worker 恢复
|
||||
- **WHEN** Worker A 的 processing 租约已被回收且任务随后由 Worker B 以更高 claim version 重新领取
|
||||
- **THEN** Worker A 的租约刷新、事件写入和终态更新 MUST 因 claim version 不匹配而失败
|
||||
- **THEN** Worker A MUST NOT 覆盖 Worker B 的任务状态或创建重复事件
|
||||
|
||||
#### Scenario: 可重试节点故障
|
||||
- **WHEN** Guard 返回 429、5xx、连接失败或超时且任务仍有剩余尝试次数
|
||||
- **THEN** Worker MUST 将任务置为 retry 并设置有界退避的 next_attempt_at
|
||||
|
||||
#### Scenario: 不可重试错误或达到最大尝试次数
|
||||
- **WHEN** Guard 返回认证失败、严格解析失败或任务达到最大尝试次数
|
||||
- **THEN** Worker MUST 将任务标记为 failed 并保存脱敏后的稳定错误码
|
||||
- **THEN** Redis 载荷 MUST 被删除或等待短 TTL 自动清理
|
||||
|
||||
#### Scenario: 回收滞留 processing 任务
|
||||
- **WHEN** processing 任务的租约超过允许时长
|
||||
- **THEN** 系统 MUST 按剩余尝试次数把任务回收到 retry 或标记 failed
|
||||
- **THEN** 系统 MUST 输出可关联 job ID 的回收日志
|
||||
|
||||
#### Scenario: Worker 启动失败
|
||||
- **WHEN** 数据库、Redis、配置或加密依赖导致 Worker 无法启动
|
||||
- **THEN** 主 API MUST 继续提供非提示词审计能力
|
||||
- **THEN** 运行态 MUST 显示 error/degraded 和稳定错误码,而不是显示健康
|
||||
|
||||
### Requirement: Qwen3Guard 返回必须被严格归一化
|
||||
系统 SHALL 严格解析单一 `Safety` 行和单一 `Categories` 行,并支持 Violent、Non-violent Illegal Acts、Sexual Content or Sexual Acts、PII、Suicide & Self-Harm、Unethical Acts、Politically Sensitive Topics、Copyright Violation、Jailbreak 九类输入风险。额外非空说明、重复字段、未知 Safety 或无法解析响应 MUST 视为 invalid_response。
|
||||
|
||||
#### Scenario: Safe 结果
|
||||
- **WHEN** Guard 返回 `Safety: Safe`
|
||||
- **THEN** 归一化结果 MUST 为 pass/low/Allow
|
||||
|
||||
#### Scenario: Controversial 结果
|
||||
- **WHEN** Guard 返回 `Safety: Controversial`
|
||||
- **THEN** 默认结果 MUST 为 flag/Warn
|
||||
- **THEN** 命中已启用的 Jailbreak、PII 或 Suicide & Self-Harm 时 MUST 提升为 critical/Block
|
||||
|
||||
#### Scenario: Unsafe 结果
|
||||
- **WHEN** Guard 返回 `Safety: Unsafe` 且命中至少一个已启用类别
|
||||
- **THEN** 结果 MUST 为 critical/Block
|
||||
|
||||
#### Scenario: Unsafe 包含未知类别
|
||||
- **WHEN** Guard 返回 Unsafe 但类别未知或不可映射
|
||||
- **THEN** 系统 MUST 记录 `unknown_unsafe` 并保持 Block 语义
|
||||
|
||||
#### Scenario: 严格响应解析失败
|
||||
- **WHEN** Guard 响应缺少字段、包含重复字段、出现额外非空说明或 Safety 不在允许枚举中
|
||||
- **THEN** 系统 MUST 返回 `prompt_guard_invalid_response`
|
||||
- **THEN** 系统 MUST NOT 把该结果伪装为 Safe
|
||||
|
||||
### Requirement: 长提示词必须完整进行 Unicode 分片审计
|
||||
系统 SHALL 按 Unicode rune 而不是字节对提示词分片。最新用户输入 MUST 作为优先片段,其他输入按确定顺序完整覆盖;异步任务必须在每片开始前刷新 processing 租约,并为每片开始、完成、失败及最终聚合输出不含正文的结构化日志。
|
||||
|
||||
#### Scenario: 输入超过节点单片上限
|
||||
- **WHEN** 提示词 Unicode 字符数超过节点 input_limit
|
||||
- **THEN** 系统 MUST 生成覆盖全部非空文本的连续分片
|
||||
- **THEN** 任一分片 Block MUST 使聚合结果为 Block
|
||||
- **THEN** 只有全部必要分片成功后才能产生 Allow
|
||||
|
||||
#### Scenario: 最新输入包含风险
|
||||
- **WHEN** 最新用户输入位于长会话尾部并包含 Block 风险
|
||||
- **THEN** 该输入 MUST 在历史文本之前接受扫描
|
||||
- **THEN** 同步模式 MAY 在确认 Block 后停止后续分片,但 MUST NOT 部分放行
|
||||
|
||||
#### Scenario: 多分片扫描完成
|
||||
- **WHEN** 一个提示词被拆成多个分片并完成聚合
|
||||
- **THEN** 日志 MUST 包含 chunk_index、chunk_total、chunk_chars、input_chars、input_limit、guard endpoint、action 和 latency
|
||||
- **THEN** 日志 MUST NOT 包含分片正文、脱敏前证据或内部优先级分隔符
|
||||
|
||||
### Requirement: 审计事件必须独立、可关联且可安全管理
|
||||
系统 SHALL 把归一化结果写入 `prompt_audit_events`,并支持是否保存 Pass 事件。事件 MUST 包含请求上下文、分列的用户名/邮箱/API Key 名称快照、脱敏提示词快照、decision、risk_level、action、分类、scanner、证据、策略、节点、配置版本、分片数和耗时;管理 DTO MUST 从这些事实确定性派生结构化 `issue_summaries`,不得复制保存第二套风险事实。
|
||||
|
||||
#### Scenario: 风险事件被记录
|
||||
- **WHEN** Worker 或同步 Guard 得到 flag/critical 结果
|
||||
- **THEN** 系统 MUST 创建独立提示词审计事件
|
||||
- **THEN** 事件 MUST 可通过 request_id、user_id、api_key_id、group_id 和 prompt_hash 检索
|
||||
|
||||
#### Scenario: Pass 事件存储关闭
|
||||
- **WHEN** 结果为 pass 且 store_pass_events=false
|
||||
- **THEN** 系统 MUST 完成任务但 MAY 不创建事件
|
||||
|
||||
#### Scenario: 同步结果写入失败
|
||||
- **WHEN** 同步 Guard 已完成判定但事件持久化失败
|
||||
- **THEN** 系统 MUST 输出 `prompt_guard.result_record_failed`
|
||||
- **THEN** 持久化失败 MUST NOT 把已确定的 Allow 改成 Block,也 MUST NOT 撤销已确定的 Block
|
||||
|
||||
### Requirement: 提示词审计运行态必须反映真实依赖和处理状态
|
||||
系统 SHALL 提供运行态接口,返回有效模式、期望/生效配置版本、配置加载时间与错误、Worker 心跳、队列容量与各状态数量、处理/失败统计、最近错误、节点连通性、数据库/Redis 状态和同步 Guard 指标。
|
||||
|
||||
#### Scenario: 管理员查询健康运行态
|
||||
- **WHEN** Worker 正常心跳、数据库与 Redis 可用且至少一个节点探测成功
|
||||
- **THEN** 运行态 MUST 显示 running/ok 和真实统计值
|
||||
|
||||
#### Scenario: Redis 不可用
|
||||
- **WHEN** 提示词审计已启用但 Redis 载荷存储不可用
|
||||
- **THEN** 异步运行态 MUST 显示 error 或 degraded
|
||||
- **THEN** 页面 MUST NOT 仅因 Base URL 已配置而显示健康
|
||||
|
||||
### Requirement: 管理员必须能够查询和安全删除提示词审计事件
|
||||
系统 SHALL 提供分页列表、详情、单条删除、批量 ID 删除和按筛选删除。筛选 MUST 支持 decision、risk level、endpoint、group、user、API key、request ID、prompt Hash、关键字和时间范围。
|
||||
|
||||
#### Scenario: 按筛选查询事件
|
||||
- **WHEN** 管理员提交一个或多个受支持筛选条件
|
||||
- **THEN** 系统 MUST 返回稳定排序的分页事件和总数
|
||||
|
||||
#### Scenario: 预览按筛选删除
|
||||
- **WHEN** 管理员提交包含明确时间范围的删除筛选
|
||||
- **THEN** 系统 MUST 返回 matched_count、规范化筛选摘要、snapshot_max_id、filter_hash 和绑定当前管理员且短期有效的 confirmation_token
|
||||
- **THEN** 系统 MUST 不立即删除数据
|
||||
|
||||
#### Scenario: 确认按筛选删除
|
||||
- **WHEN** 管理员提交相同筛选、有效 filter_hash、未过期 confirmation_token 和显式 confirm=true
|
||||
- **THEN** 系统 MUST 只分批删除匹配且 id 不高于预览 snapshot_max_id 的事件,以及已无事件引用的孤立任务
|
||||
- **THEN** 系统 MUST 清理相关 Redis 载荷并写入管理操作审计
|
||||
|
||||
#### Scenario: 伪造或重放其他管理员的删除确认
|
||||
- **WHEN** confirmation_token 无法认证、已过期、操作者不匹配、Hash 不匹配或缺失
|
||||
- **THEN** 系统 MUST 拒绝删除并返回稳定错误码
|
||||
- **THEN** 客户端自行计算 filter_hash MUST NOT 绕过 delete-preview
|
||||
|
||||
#### Scenario: 无时间范围的大范围删除
|
||||
- **WHEN** 管理员尝试按筛选删除但未提供明确时间范围
|
||||
- **THEN** 系统 MUST 拒绝操作并返回稳定错误码
|
||||
@@ -0,0 +1,201 @@
|
||||
## ADDED Requirements
|
||||
|
||||
### Requirement: 同步提示词门禁必须由显式配置启用
|
||||
系统 SHALL 使用 `enabled` 与 `blocking_enabled` 表达关闭、异步只审计、同步审计并阻止三态。旧配置或缺失字段 MUST 归一为 `blocking_enabled=false`;系统 MUST 拒绝 `enabled=false && blocking_enabled=true` 的配置。
|
||||
|
||||
#### Scenario: 关闭提示词审计
|
||||
- **WHEN** enabled=false
|
||||
- **THEN** 有效模式 MUST 为 off
|
||||
- **THEN** blocking_enabled MUST 被视为 false
|
||||
|
||||
#### Scenario: 启用异步审计
|
||||
- **WHEN** enabled=true 且 blocking_enabled=false
|
||||
- **THEN** 有效模式 MUST 为 async_audit
|
||||
- **THEN** Guard 故障 MUST NOT 改变主请求结果
|
||||
|
||||
#### Scenario: 启用同步阻止
|
||||
- **WHEN** enabled=true 且 blocking_enabled=true
|
||||
- **THEN** 有效模式 MUST 为 blocking
|
||||
- **THEN** 适用请求 MUST 等待 Guard 判定后才能进入账号选择、计费和上游阶段
|
||||
|
||||
#### Scenario: 保存非法开关组合
|
||||
- **WHEN** 管理员保存 enabled=false 且 blocking_enabled=true
|
||||
- **THEN** 后端 MUST 返回 400 和 `prompt_guard_requires_audit_enabled`
|
||||
|
||||
### Requirement: 安全审计协调器必须保持两个引擎的独立语义
|
||||
系统 SHALL 通过一个薄协调器把可信请求上下文交给现有内容审核和新增提示词审计。协调器 MUST 不转换两套风险分类、不共用事件表、不让提示词审计触发内容审核副作用,并 MUST 使用确定性的阻断优先级。
|
||||
|
||||
#### Scenario: 现有内容审核阻断
|
||||
- **WHEN** 现有内容审核返回 Block
|
||||
- **THEN** 客户端 MUST 继续收到升级前的状态码、错误码和文案
|
||||
- **THEN** 提示词审计异步模式 MAY 继续完成自己的独立记录
|
||||
|
||||
#### Scenario: 仅提示词 Guard 阻断
|
||||
- **WHEN** 现有内容审核允许但提示词 Guard 返回 Block
|
||||
- **THEN** 客户端 MUST 收到 `prompt_guard_blocked`
|
||||
|
||||
#### Scenario: 两个引擎同时阻断
|
||||
- **WHEN** 两个引擎都返回 Block
|
||||
- **THEN** 现有内容审核错误语义 MUST 具有客户端响应优先级
|
||||
- **THEN** 两个引擎 MUST 各自记录其结果和结构化日志
|
||||
|
||||
### Requirement: 同步门禁必须位于外部副作用之前
|
||||
系统 MUST 在鉴权和请求格式校验完成后、账号选择、账户并发、计费资格检查、任何预扣、上游连接和上游写入之前完成同步判定。被 Block 或 fail-closed 拒绝的请求 MUST 不产生这些下游副作用。
|
||||
|
||||
#### Scenario: HTTP 请求被 Guard 阻断
|
||||
- **WHEN** 任一支持的 HTTP 模型请求得到 Block
|
||||
- **THEN** 账号选择次数、计费检查/预扣次数和上游请求次数 MUST 均为 0
|
||||
- **THEN** 流式请求 MUST 在拒绝前未写出 SSE 响应头或首字节
|
||||
|
||||
#### Scenario: Guard 不可用
|
||||
- **WHEN** 同步模式下所有可用节点均失败
|
||||
- **THEN** 请求 MUST 在任何账号、计费或上游副作用之前返回 503
|
||||
|
||||
### Requirement: 同步门禁必须覆盖所有目标协议入口
|
||||
系统 SHALL 覆盖现有内容审核已接入的所有用户文本入口,并通过结构测试防止后续路由绕过。至少包括 OpenAI Chat Completions、OpenAI Responses、Claude Messages、Gemini、OpenAI Images/Grok 媒体文本 prompt,以及 Responses WebSocket 首轮和后续轮次。
|
||||
|
||||
#### Scenario: OpenAI 兼容 HTTP 入口
|
||||
- **WHEN** 客户端调用 Chat Completions 或 Responses 兼容入口
|
||||
- **THEN** 系统 MUST 使用对应协议提取器并执行同一 Guard evaluator
|
||||
- **THEN** 现有 OpenAI 请求和响应 envelope MUST 保持兼容
|
||||
|
||||
#### Scenario: Claude 或 Gemini 入口
|
||||
- **WHEN** 客户端调用 Claude Messages 或 Gemini 入口
|
||||
- **THEN** 系统 MUST 执行相同策略判定
|
||||
- **THEN** 拒绝响应 MUST 使用该协议现有错误 envelope 和共享稳定 error_code
|
||||
|
||||
#### Scenario: 新增用户文本入口
|
||||
- **WHEN** 后续代码新增一个可触发模型执行且包含用户文本的路由
|
||||
- **THEN** 路由覆盖门禁 MUST 在缺少安全审计接线时失败
|
||||
|
||||
### Requirement: 同步分片必须共享总预算并完整覆盖
|
||||
系统 SHALL 以有序节点列表中首个启用节点的 timeout 作为一次同步 evaluation 的总预算。所有分片和节点故障切换 MUST 共享该 deadline;任一必要分片失败、超时或无合法结果时 MUST fail-closed。
|
||||
|
||||
#### Scenario: 所有分片均为安全
|
||||
- **WHEN** 每个非空分片都在总预算内返回 Safe 或允许的 Warn
|
||||
- **THEN** 请求 MAY 进入下一阶段
|
||||
|
||||
#### Scenario: 中间分片阻断
|
||||
- **WHEN** 任一分片返回 Block
|
||||
- **THEN** evaluator MAY 立即早停
|
||||
- **THEN** 请求 MUST 被阻断且不得部分转发
|
||||
|
||||
#### Scenario: 最后一个必要分片失败
|
||||
- **WHEN** 前面分片安全但最后一个必要分片超时或响应无效
|
||||
- **THEN** 系统 MUST 返回 unavailable/invalid_response
|
||||
- **THEN** 系统 MUST NOT 根据部分结果放行
|
||||
|
||||
### Requirement: 同步节点故障切换必须有序且 fail-closed
|
||||
系统 SHALL 按配置顺序尝试启用节点。连接失败、429、5xx 和超时 MAY 在总 deadline 尚有剩余时切换到下一节点;401/403、严格解析失败或耗尽节点 MUST 结束为不可用/非法响应。同步模式 MUST NOT 提供隐式 fail-open。
|
||||
|
||||
#### Scenario: 首节点暂时失败而次节点成功
|
||||
- **WHEN** 首节点返回可重试错误且次节点在剩余预算内返回合法结果
|
||||
- **THEN** 系统 MUST 使用次节点结果
|
||||
- **THEN** failover 指标 MUST 增加
|
||||
|
||||
#### Scenario: 认证失败
|
||||
- **WHEN** 节点返回 401 或 403
|
||||
- **THEN** 系统 MUST 视为不可重试配置错误
|
||||
- **THEN** 请求 MUST 返回 503 而不是按 Safe 放行
|
||||
|
||||
#### Scenario: 所有节点容量饱和
|
||||
- **WHEN** 全局或每节点 bulkhead 均无法接受 evaluation
|
||||
- **THEN** 系统 MUST 快速返回 `prompt_guard_unavailable`
|
||||
- **THEN** 系统 MUST 不无限排队
|
||||
|
||||
### Requirement: HTTP 拒绝必须保持协议兼容和稳定错误码
|
||||
同步 Guard MUST 使用现有 Handler 的协议错误构造器和最小扩展,且只向客户端暴露通用消息、稳定 Prompt Guard code/reason 和 request ID。OpenAI/Claude MUST 在 error 对象的可选 `code` 字段携带稳定代码并保留原合法 type;Gemini MUST 保留数值 `error.code` 与 canonical status,并在 `google.rpc.ErrorInfo.reason` 携带稳定代码。响应 MUST 不包含风险正文、类别细节、内部节点地址或凭据。
|
||||
|
||||
#### Scenario: HTTP Block
|
||||
- **WHEN** 同步 Guard 判定为 Block
|
||||
- **THEN** HTTP 状态 MUST 为 403
|
||||
- **THEN** error_code MUST 为 `prompt_guard_blocked`
|
||||
|
||||
#### Scenario: Gemini HTTP Block
|
||||
- **WHEN** Gemini 入口的同步 Guard 判定为 Block
|
||||
- **THEN** Google error envelope 的 `error.code` MUST 保持数值 403 且 status MUST 为对应 canonical status
|
||||
- **THEN** `error.details` 中 ErrorInfo reason MUST 为 `prompt_guard_blocked`
|
||||
|
||||
#### Scenario: HTTP Guard 不可用
|
||||
- **WHEN** 节点超时、连接失败、熔断或容量不足
|
||||
- **THEN** HTTP 状态 MUST 为 503
|
||||
- **THEN** error_code MUST 为 `prompt_guard_unavailable`
|
||||
|
||||
#### Scenario: HTTP Guard 响应非法
|
||||
- **WHEN** Guard 输出无法严格解析
|
||||
- **THEN** HTTP 状态 MUST 为 503
|
||||
- **THEN** error_code MUST 为 `prompt_guard_invalid_response`
|
||||
|
||||
### Requirement: Responses WebSocket 必须对每个 response.create 执行门禁
|
||||
系统 SHALL 在 WebSocket 首次和后续每个 `response.create` 帧进入本轮用户/账号并发、计费和上游发送之前执行同步 Guard。一次安全结果 MUST NOT 被复用于不同的后续帧。
|
||||
|
||||
#### Scenario: 首轮 Block
|
||||
- **WHEN** 首个 response.create 被判定为 Block
|
||||
- **THEN** 服务端 MUST 不建立本轮上游请求或计费记录
|
||||
- **THEN** 服务端 MUST 使用 close code 4403 和 reason `prompt_guard_blocked` 关闭连接
|
||||
|
||||
#### Scenario: 后续轮次 Block
|
||||
- **WHEN** 已建立连接的后续 response.create 被判定为 Block
|
||||
- **THEN** 该帧 MUST 不发送给上游且不得创建本轮计费记录
|
||||
- **THEN** 服务端 MUST 使用 4403 关闭连接并记录 stage=subsequent_turn
|
||||
|
||||
#### Scenario: WebSocket Guard 不可用
|
||||
- **WHEN** 首轮或后续轮次 Guard 不可用或响应非法
|
||||
- **THEN** 服务端 MUST 使用 close code 1013
|
||||
- **THEN** reason MUST 为 `prompt_guard_unavailable` 或 `prompt_guard_invalid_response`
|
||||
|
||||
### Requirement: 同步结果必须复用到脱敏事件且不得重复扫描
|
||||
系统 SHALL 在一次同步 evaluation 后把已得到的归一化结果交给独立记录路径。记录路径 MUST NOT 重新调用 Guard,也 MUST NOT 需要完整提示词正文;同步结果最多对应一个任务事实和一个按存储策略决定的事件。
|
||||
|
||||
#### Scenario: 同步 Block 被记录
|
||||
- **WHEN** evaluator 已得到 Block
|
||||
- **THEN** 系统 MUST 用脱敏快照和既有结果创建 done 任务及风险事件
|
||||
- **THEN** Guard 调用次数 MUST 等于 evaluation 实际需要的节点/分片次数,而不是因记录而增加
|
||||
|
||||
#### Scenario: 同步 Allow 且不保存 Pass
|
||||
- **WHEN** evaluator 得到 Allow 且 store_pass_events=false
|
||||
- **THEN** 系统 MAY 只保存任务/指标而不创建 Pass 事件
|
||||
|
||||
### Requirement: 配置必须以版本化快照发布到请求热路径
|
||||
系统 SHALL 为提示词审计配置维护单调递增 config_version、updated_at、updated_by 和 change_summary。保存后 MUST 原子替换本实例快照并通过 Redis 发布失效通知;请求热路径 MUST 读取内存快照而不是逐请求查询数据库。
|
||||
|
||||
#### Scenario: 多实例收到配置更新
|
||||
- **WHEN** 管理员成功保存新配置
|
||||
- **THEN** 保存实例 MUST 立即安装新版本并发布 Redis 失效通知
|
||||
- **THEN** 其他实例 MUST 重新加载并原子替换快照
|
||||
|
||||
#### Scenario: 两个管理员并发保存配置
|
||||
- **WHEN** 两个保存请求携带相同 expected_config_version 且第一个已提交新版本
|
||||
- **THEN** 第二个请求 MUST 返回 409 `prompt_audit_config_conflict`
|
||||
- **THEN** 第二个请求 MUST NOT 静默覆盖第一个请求或复用相同 config_version
|
||||
|
||||
#### Scenario: Redis 通知不可用
|
||||
- **WHEN** 配置已保存但 Redis publish 失败
|
||||
- **THEN** 系统 MUST 记录 `prompt_guard.config_reload_degraded`
|
||||
- **THEN** 其他实例 MUST 通过有界 TTL 刷新最终获得新版本
|
||||
|
||||
#### Scenario: 冷启动无法加载严格配置
|
||||
- **WHEN** 实例冷启动且无法获得有效配置快照
|
||||
- **THEN** 对已知要求同步阻止的适用请求 MUST fail-closed
|
||||
- **THEN** 运行态 MUST 暴露配置加载错误
|
||||
|
||||
### Requirement: Guard 关键路径必须可观测且不得泄密
|
||||
系统 SHALL 输出稳定结构化事件并提供计数/耗时指标。日志至少 MUST 覆盖配置更新/加载/降级、evaluation 开始、Allow、Block、失败、结果记录失败、异步投递/丢弃、Worker 处理/重试/失败、逐分片开始/完成/失败、分片聚合和滞留回收。
|
||||
|
||||
#### Scenario: 同步请求被阻断
|
||||
- **WHEN** Guard 阻断一个请求
|
||||
- **THEN** 日志 MUST 包含 request_id、user_id、api_key_id、group_id、protocol、endpoint、model、config_version、guard_endpoint_id、decision、action、chunk_total、latency_ms、status 和 error_code
|
||||
- **THEN** 日志 MUST 明确包含 `upstream_dispatched=false` 和 `billing_preconsumed=false` 或目标项目等价字段
|
||||
|
||||
#### Scenario: 检查日志敏感字段
|
||||
- **WHEN** 测试捕获提示词审计日志
|
||||
- **THEN** 日志中 MUST 不包含原始提示词、API Key、Authorization、完整 Guard URL query 或 Redis 载荷
|
||||
|
||||
### Requirement: 禁用或回滚同步阻止必须即时恢复异步行为
|
||||
系统 SHALL 支持仅通过关闭 blocking_enabled 回到异步只审计,无需删除表、清空历史事件或停止现有内容审核。
|
||||
|
||||
#### Scenario: 管理员关闭同步阻止
|
||||
- **WHEN** blocking_enabled 从 true 保存为 false 且新配置已生效
|
||||
- **THEN** 后续适用请求 MUST 不再等待 Guard 同步结果
|
||||
- **THEN** enabled=true 时后续请求 MUST 改为异步投递
|
||||
- **THEN** 历史任务和事件 MUST 保留
|
||||
+160
@@ -0,0 +1,160 @@
|
||||
## ADDED Requirements
|
||||
|
||||
### Requirement: 管理台必须提供安全审计分组和独立提示词审计页面
|
||||
控制台 SHALL 把安全相关的内容审核页面组织到“安全审计”导航分组中,并新增独立“提示词审计”页面。原 `/admin/risk-control` 路由、页面状态和功能 MUST 保持兼容;新页面路由 MUST 为 `/admin/prompt-audit` 或经实现评审确认的等价稳定路由。
|
||||
|
||||
#### Scenario: 管理员查看侧栏
|
||||
- **WHEN** 管理员已登录且 risk_control_enabled=true
|
||||
- **THEN** 侧栏 MUST 展示“安全审计”可展开分组
|
||||
- **THEN** 分组 MUST 至少包含“内容审核”和“提示词审计”两个子入口
|
||||
|
||||
#### Scenario: 管理员打开原内容审核页面
|
||||
- **WHEN** 管理员访问 `/admin/risk-control`
|
||||
- **THEN** 页面 MUST 继续展示原有 Moderations、关键词、Hash、封号、邮件和记录功能
|
||||
- **THEN** 页面 MUST NOT 被提示词审计配置或事件替换
|
||||
|
||||
#### Scenario: 功能总开关关闭
|
||||
- **WHEN** risk_control_enabled=false
|
||||
- **THEN** 安全审计导航和提示词审计网关执行 MUST 按现有功能开关策略停用
|
||||
- **THEN** 已存储的配置和历史事件 MUST 不被删除
|
||||
|
||||
### Requirement: 提示词审计页面必须提供清晰的独立工作区
|
||||
页面 SHALL 在同一工作区展示运行概览、审计池、审计策略、事件列表和固定保存操作区。页面 MUST 清楚区分“异步只审计”和“同步阻止”,并 MUST 展示未保存状态和最终生效状态。
|
||||
|
||||
#### Scenario: 初次打开页面
|
||||
- **WHEN** 管理员打开提示词审计页面
|
||||
- **THEN** 页面 MUST 并行或有界加载配置、运行态、分组列表和事件列表
|
||||
- **THEN** 页面 MUST 展示有效模式、Worker 状态、队列状态、节点连通性和最近错误
|
||||
|
||||
#### Scenario: 修改但未保存配置
|
||||
- **WHEN** 管理员修改审计池、分类、范围或模式开关
|
||||
- **THEN** 页面 MUST 显示“有未保存的更改”
|
||||
- **THEN** 运行态 MUST 继续标识服务端当前生效版本,不能把草稿显示为已生效
|
||||
|
||||
### Requirement: 页面必须支持完整审计池管理和真实探测
|
||||
页面 SHALL 支持新增、编辑、启用、禁用和删除审计池,并允许配置 Base URL、API Key、Model、超时和 input_limit。API Key 已保存后 MUST 只显示配置状态,不能回显明文。
|
||||
|
||||
#### Scenario: 编辑已保存节点
|
||||
- **WHEN** 管理员打开已配置 API Key 的节点
|
||||
- **THEN** API Key 输入框 MUST 为空或显示不可逆占位状态
|
||||
- **THEN** 未填写新 Key 保存时 MUST 保留原密文
|
||||
- **THEN** 页面 MUST 提供显式清除凭据操作
|
||||
|
||||
#### Scenario: 执行连接测试
|
||||
- **WHEN** 管理员点击节点“连接测试”
|
||||
- **THEN** 页面 MUST 展示配置校验、发送请求、服务响应和测试结论状态
|
||||
- **THEN** 结果 MUST 展示耗时、HTTP 状态、稳定错误码和脱敏消息
|
||||
|
||||
### Requirement: 页面必须支持审计范围和九类风险配置
|
||||
页面 SHALL 支持全部分组或指定 group ID 范围,并展示九类 Qwen3Guard 风险分类。页面 MUST 使用目标项目真实分组数据,已删除但仍存在于配置中的分组 MUST 显示为失效项而不是被静默丢弃。
|
||||
|
||||
#### Scenario: 选择指定分组
|
||||
- **WHEN** 管理员把范围切换为 selected 并选择一个或多个分组
|
||||
- **THEN** 保存载荷 MUST 使用稳定 group ID
|
||||
- **THEN** 页面 MUST 展示已选数量并支持搜索
|
||||
|
||||
#### Scenario: 查看风险分类
|
||||
- **WHEN** 管理员查看扫描器配置
|
||||
- **THEN** 页面 MUST 展示 Violent、Non-violent Illegal Acts、Sexual Content or Sexual Acts、PII、Suicide & Self-Harm、Unethical Acts、Politically Sensitive Topics、Copyright Violation、Jailbreak
|
||||
|
||||
### Requirement: 开启同步阻止必须有明确的风险确认
|
||||
页面 SHALL 把 enabled、blocking_enabled 和 store_pass_events 作为独立开关。关闭 enabled 时 MUST 自动关闭并禁用 blocking_enabled;开启 blocking_enabled 时 MUST 展示二次确认,说明请求延迟、Block 和 Guard 不可用的 fail-closed 行为。
|
||||
|
||||
#### Scenario: 开启同步阻止
|
||||
- **WHEN** 管理员把 blocking_enabled 从 false 切换为 true
|
||||
- **THEN** 页面 MUST 在保存前展示风险确认
|
||||
- **THEN** 确认文案 MUST 说明请求会等待 Guard,Block 或 Guard 不可用时不会访问上游
|
||||
|
||||
#### Scenario: 关闭审计总开关
|
||||
- **WHEN** 管理员关闭 enabled
|
||||
- **THEN** 页面草稿 MUST 同时把 blocking_enabled 设为 false
|
||||
|
||||
### Requirement: 配置保存必须可验证且不得泄露凭据
|
||||
页面 SHALL 通过一个统一保存动作提交完整规范化配置。保存成功后 MUST 用后端返回值刷新页面快照、清除已提交 API Key 明文并显示 config_version;保存失败 MUST 保留草稿并展示稳定错误信息。
|
||||
|
||||
#### Scenario: 保存成功
|
||||
- **WHEN** 后端成功保存配置
|
||||
- **THEN** 页面 MUST 显示配置已同步和新的 config_version
|
||||
- **THEN** 浏览器状态、调试日志和缓存 MUST 不再保留刚提交的 API Key 明文
|
||||
|
||||
#### Scenario: 保存校验失败
|
||||
- **WHEN** 后端返回节点地址、模式组合或策略校验错误
|
||||
- **THEN** 页面 MUST 保留用户草稿
|
||||
- **THEN** 页面 MUST 展示稳定错误码及可行动的中文说明
|
||||
|
||||
#### Scenario: 配置被其他管理员更新
|
||||
- **WHEN** 保存返回 409 `prompt_audit_config_conflict`
|
||||
- **THEN** 页面 MUST 保留本地草稿并提示服务端配置已变化
|
||||
- **THEN** 页面 MUST 提供重新加载/对比入口,不得自动用旧草稿覆盖新配置
|
||||
|
||||
### Requirement: 页面必须展示真实运行态和同步 Guard 指标
|
||||
页面 SHALL 展示 process_status、Worker 总数/活动数、队列容量/长度、queued/processing/done/failed 数、处理/失败总数、最近时间、节点连通性、配置版本一致性、Redis Payload Store 状态和同步 Guard Allow/Flag/Block/Unavailable/timeout/failover/bulkhead 指标。
|
||||
|
||||
#### Scenario: 配置版本未同步
|
||||
- **WHEN** expected_config_version 与 active_config_version 不一致
|
||||
- **THEN** 页面 MUST 显示明确的配置未同步或加载中状态
|
||||
- **THEN** 页面 MUST 展示最近加载错误和时间(如存在)
|
||||
|
||||
#### Scenario: Worker 心跳过期
|
||||
- **WHEN** heartbeat_at 超过后端定义的健康窗口
|
||||
- **THEN** 页面 MUST 显示 stale 而不是 running
|
||||
|
||||
### Requirement: 页面必须提供可复核的事件列表和详情
|
||||
页面 SHALL 提供事件分页、总数、decision/risk/endpoint/group/user/API key/request ID/prompt Hash/关键字/时间范围筛选、行选择和详情抽屉或弹窗。详情 MUST 只展示脱敏数据。
|
||||
|
||||
#### Scenario: 查看事件列表
|
||||
- **WHEN** 管理员应用筛选
|
||||
- **THEN** 表格 MUST 展示时间、用户/API key、分组、入口/模型、判定、风险、分类、预览和操作
|
||||
|
||||
#### Scenario: 查看事件详情
|
||||
- **WHEN** 管理员打开一条事件
|
||||
- **THEN** 页面 MUST 展示脱敏预览、审计摘要、结构化返回、具体风险摘要和技术信息
|
||||
- **THEN** 页面 MUST 提供 request ID、prompt Hash、scanner、策略、节点、配置版本、分片数和耗时
|
||||
- **THEN** 页面 MUST 不展示完整提示词或节点 API Key
|
||||
|
||||
#### Scenario: 复核用户身份和具体风险
|
||||
- **WHEN** 事件拥有用户名、用户邮箱、API Key 名称和一个或多个风险分类
|
||||
- **THEN** 页面 MUST 将用户名、邮箱和 API Key 名称分列展示并提供独立复制操作
|
||||
- **THEN** 页面 MUST 为每个风险展示 category、标题、说明、严重度、动作、scanner、score 和脱敏证据摘要
|
||||
- **THEN** 用户不存在或字段为空时 MUST 显示稳定 fallback,而不是把其他身份字段冒充为该字段
|
||||
|
||||
### Requirement: 页面必须提供防误操作的事件删除流程
|
||||
页面 SHALL 支持单条删除、选中项批量删除和按筛选删除。按筛选删除 MUST 先调用预览接口,并要求明确时间范围、matched_count、snapshot_max_id、filter_hash、服务端认证 confirmation_token 和二次确认。
|
||||
|
||||
#### Scenario: 单条删除
|
||||
- **WHEN** 管理员确认删除一条事件
|
||||
- **THEN** 页面 MUST 调用单条删除接口并在成功后刷新列表与运行统计
|
||||
|
||||
#### Scenario: 按筛选删除
|
||||
- **WHEN** 管理员已设置明确时间范围并请求按筛选删除
|
||||
- **THEN** 页面 MUST 先展示匹配数量和规范化筛选摘要
|
||||
- **THEN** 只有管理员再次确认后才能提交 filter_hash、confirmation_token 和 confirm=true
|
||||
|
||||
#### Scenario: 筛选在预览后发生变化
|
||||
- **WHEN** 管理员预览后修改任意筛选条件
|
||||
- **THEN** 旧 filter_hash MUST 失效
|
||||
- **THEN** 旧 confirmation_token MUST 同时失效
|
||||
- **THEN** 页面 MUST 要求重新预览
|
||||
|
||||
### Requirement: 管理 API 操作必须纳入现有管理员审计
|
||||
所有配置写入、节点探测和事件删除 SHALL 复用现有管理员鉴权与管理操作审计。审计详情 MUST 使用脱敏摘要,禁止记录 API Key、完整提示词或完整请求载荷。
|
||||
|
||||
#### Scenario: 配置更新成功
|
||||
- **WHEN** 管理员成功保存提示词审计配置
|
||||
- **THEN** 管理操作审计 MUST 记录操作者、request ID、enabled、blocking_enabled、config_version、节点数量、分类数量和分组范围摘要
|
||||
|
||||
#### Scenario: 节点探测失败
|
||||
- **WHEN** 管理员探测节点失败
|
||||
- **THEN** 管理操作审计 MUST 记录节点 ID、稳定错误码、HTTP 状态和耗时
|
||||
- **THEN** 审计详情 MUST 不包含 API Key 或完整 Base URL query
|
||||
|
||||
### Requirement: 页面必须满足响应式、可访问和国际化要求
|
||||
页面 SHALL 使用现有 Vue 3、i18n 和通用组件体系,支持桌面与窄屏,所有开关、输入、按钮、状态和对话框 MUST 具有可访问名称;新增中英文文案键 MUST 成对提供且通过现有 lint、typecheck 和 Vitest。
|
||||
|
||||
#### Scenario: 窄屏使用
|
||||
- **WHEN** 页面宽度小于桌面断点
|
||||
- **THEN** 配置区、筛选区、表格和固定保存栏 MUST 可滚动或重排而不遮挡关键操作
|
||||
|
||||
#### Scenario: 键盘和读屏操作
|
||||
- **WHEN** 用户只使用键盘或读屏访问页面
|
||||
- **THEN** 审计池操作、模式开关、筛选、详情和确认对话框 MUST 可识别且可操作
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user