From d11bdb13f52b0fa1d356482b4f8b14262420ec30 Mon Sep 17 00:00:00 2001 From: mt21625457 Date: Fri, 17 Jul 2026 00:39:39 +0800 Subject: [PATCH 1/7] feat(security-audit): add OpenAI-compatible prompt auditing --- backend/cmd/server/main.go | 7 + backend/cmd/server/wire.go | 16 +- backend/cmd/server/wire_gen.go | 37 +- backend/cmd/server/wire_gen_test.go | 1 + .../internal/handler/batch_image_handler.go | 43 + .../handler/content_moderation_helper.go | 14 - backend/internal/handler/gateway_handler.go | 6 +- .../gateway_handler_chat_completions.go | 4 +- .../handler/gateway_handler_responses.go | 4 +- .../internal/handler/gemini_v1beta_handler.go | 4 +- backend/internal/handler/grok_media.go | 6 +- backend/internal/handler/handler.go | 2 + .../internal/handler/image_task_handler.go | 39 + .../internal/handler/openai_alpha_search.go | 4 + .../handler/openai_chat_completions.go | 4 +- backend/internal/handler/openai_embeddings.go | 4 + .../handler/openai_gateway_handler.go | 22 +- backend/internal/handler/openai_images.go | 4 +- .../internal/handler/security_audit_errors.go | 187 ++ .../handler/security_audit_errors_test.go | 159 + .../internal/handler/security_audit_helper.go | 138 + .../security_audit_media_submit_test.go | 208 ++ .../handler/security_audit_order_test.go | 85 + backend/internal/handler/wire.go | 64 +- backend/internal/securityaudit/coordinator.go | 139 + .../securityaudit/coordinator_legacy.go | 35 + .../securityaudit/coordinator_test.go | 176 ++ .../internal/securityaudit/prompt_config.go | 468 +++ .../prompt_config_integration_test.go | 244 ++ .../securityaudit/prompt_config_store.go | 409 +++ .../securityaudit/prompt_config_test.go | 157 + .../internal/securityaudit/prompt_enqueue.go | 99 + .../securityaudit/prompt_event_repository.go | 377 +++ .../internal/securityaudit/prompt_guard.go | 279 ++ .../securityaudit/prompt_guard_test.go | 276 ++ .../internal/securityaudit/prompt_handler.go | 305 ++ .../securityaudit/prompt_handler_test.go | 232 ++ .../securityaudit/prompt_issue_summary.go | 69 + .../internal/securityaudit/prompt_logging.go | 172 + .../securityaudit/prompt_logging_test.go | 66 + .../internal/securityaudit/prompt_metrics.go | 145 + .../securityaudit/prompt_metrics_test.go | 52 + .../internal/securityaudit/prompt_module.go | 22 + .../securityaudit/prompt_outbound_security.go | 196 ++ .../prompt_outbound_security_test.go | 220 ++ .../securityaudit/prompt_payload_store.go | 63 + .../prompt_payload_store_integration_test.go | 85 + .../securityaudit/prompt_qwen3guard.go | 333 ++ .../securityaudit/prompt_qwen3guard_test.go | 134 + .../securityaudit/prompt_repository.go | 433 +++ .../prompt_repository_integration_test.go | 449 +++ .../internal/securityaudit/prompt_scanner.go | 121 + .../internal/securityaudit/prompt_service.go | 469 +++ .../securityaudit/prompt_service_test.go | 126 + .../internal/securityaudit/prompt_snapshot.go | 401 +++ .../securityaudit/prompt_snapshot_test.go | 188 ++ .../internal/securityaudit/prompt_types.go | 276 ++ .../internal/securityaudit/prompt_worker.go | 347 +++ .../securityaudit/prompt_worker_test.go | 591 ++++ .../internal/server/middleware/audit_log.go | 107 +- .../server/middleware/audit_log_test.go | 127 +- backend/internal/server/routes/admin.go | 19 + .../prompt_audit_route_coverage_test.go | 136 + backend/migrations/181_prompt_audit.sql | 129 + frontend/src/components/layout/AppSidebar.vue | 13 +- .../features/prompt-audit/PromptAuditView.vue | 350 +++ .../__tests__/PromptAuditView.spec.ts | 161 + .../prompt-audit/__tests__/api.spec.ts | 44 + .../prompt-audit/__tests__/components.spec.ts | 89 + .../__tests__/integrationSurface.spec.ts | 40 + .../prompt-audit/__tests__/viewModel.spec.ts | 80 + frontend/src/features/prompt-audit/api.ts | 120 + .../prompt-audit/components/EndpointPool.vue | 172 + .../components/EventDetailDialog.vue | 66 + .../components/EventWorkspace.vue | 186 ++ .../prompt-audit/components/PolicyPanel.vue | 108 + .../components/RuntimeOverview.vue | 119 + frontend/src/features/prompt-audit/types.ts | 243 ++ .../src/features/prompt-audit/viewModel.ts | 139 + frontend/src/i18n/locales/en/admin/index.ts | 2 + .../src/i18n/locales/en/admin/promptAudit.ts | 53 + frontend/src/i18n/locales/en/common.ts | 3 + frontend/src/i18n/locales/zh/admin/index.ts | 2 + .../src/i18n/locales/zh/admin/promptAudit.ts | 53 + frontend/src/i18n/locales/zh/common.ts | 3 + frontend/src/router/index.ts | 13 + .../.openspec.yaml | 2 + .../README.md | 5 + .../design.md | 754 +++++ .../implementation-evidence.md | 31 + .../implementation-guide.md | 581 ++++ .../proposal.md | 51 + .../source-baseline.md | 146 + .../source-feature-map.md | 127 + .../source-freeze/MANIFEST.md | 52 + .../aicodex-prompt-audit-tracked.patch | 2771 +++++++++++++++++ .../aicodex-prompt-audit-untracked.tar.gz | Bin 0 -> 39342 bytes .../specs/prompt-input-audit/spec.md | 246 ++ .../specs/prompt-input-guard/spec.md | 201 ++ .../specs/security-audit-console/spec.md | 160 + .../tasks.md | 187 ++ .../verification.md | 492 +++ openspec/config.yaml | 20 + 103 files changed, 17549 insertions(+), 70 deletions(-) create mode 100644 backend/internal/handler/security_audit_errors.go create mode 100644 backend/internal/handler/security_audit_errors_test.go create mode 100644 backend/internal/handler/security_audit_helper.go create mode 100644 backend/internal/handler/security_audit_media_submit_test.go create mode 100644 backend/internal/handler/security_audit_order_test.go create mode 100644 backend/internal/securityaudit/coordinator.go create mode 100644 backend/internal/securityaudit/coordinator_legacy.go create mode 100644 backend/internal/securityaudit/coordinator_test.go create mode 100644 backend/internal/securityaudit/prompt_config.go create mode 100644 backend/internal/securityaudit/prompt_config_integration_test.go create mode 100644 backend/internal/securityaudit/prompt_config_store.go create mode 100644 backend/internal/securityaudit/prompt_config_test.go create mode 100644 backend/internal/securityaudit/prompt_enqueue.go create mode 100644 backend/internal/securityaudit/prompt_event_repository.go create mode 100644 backend/internal/securityaudit/prompt_guard.go create mode 100644 backend/internal/securityaudit/prompt_guard_test.go create mode 100644 backend/internal/securityaudit/prompt_handler.go create mode 100644 backend/internal/securityaudit/prompt_handler_test.go create mode 100644 backend/internal/securityaudit/prompt_issue_summary.go create mode 100644 backend/internal/securityaudit/prompt_logging.go create mode 100644 backend/internal/securityaudit/prompt_logging_test.go create mode 100644 backend/internal/securityaudit/prompt_metrics.go create mode 100644 backend/internal/securityaudit/prompt_metrics_test.go create mode 100644 backend/internal/securityaudit/prompt_module.go create mode 100644 backend/internal/securityaudit/prompt_outbound_security.go create mode 100644 backend/internal/securityaudit/prompt_outbound_security_test.go create mode 100644 backend/internal/securityaudit/prompt_payload_store.go create mode 100644 backend/internal/securityaudit/prompt_payload_store_integration_test.go create mode 100644 backend/internal/securityaudit/prompt_qwen3guard.go create mode 100644 backend/internal/securityaudit/prompt_qwen3guard_test.go create mode 100644 backend/internal/securityaudit/prompt_repository.go create mode 100644 backend/internal/securityaudit/prompt_repository_integration_test.go create mode 100644 backend/internal/securityaudit/prompt_scanner.go create mode 100644 backend/internal/securityaudit/prompt_service.go create mode 100644 backend/internal/securityaudit/prompt_service_test.go create mode 100644 backend/internal/securityaudit/prompt_snapshot.go create mode 100644 backend/internal/securityaudit/prompt_snapshot_test.go create mode 100644 backend/internal/securityaudit/prompt_types.go create mode 100644 backend/internal/securityaudit/prompt_worker.go create mode 100644 backend/internal/securityaudit/prompt_worker_test.go create mode 100644 backend/internal/server/routes/prompt_audit_route_coverage_test.go create mode 100644 backend/migrations/181_prompt_audit.sql create mode 100644 frontend/src/features/prompt-audit/PromptAuditView.vue create mode 100644 frontend/src/features/prompt-audit/__tests__/PromptAuditView.spec.ts create mode 100644 frontend/src/features/prompt-audit/__tests__/api.spec.ts create mode 100644 frontend/src/features/prompt-audit/__tests__/components.spec.ts create mode 100644 frontend/src/features/prompt-audit/__tests__/integrationSurface.spec.ts create mode 100644 frontend/src/features/prompt-audit/__tests__/viewModel.spec.ts create mode 100644 frontend/src/features/prompt-audit/api.ts create mode 100644 frontend/src/features/prompt-audit/components/EndpointPool.vue create mode 100644 frontend/src/features/prompt-audit/components/EventDetailDialog.vue create mode 100644 frontend/src/features/prompt-audit/components/EventWorkspace.vue create mode 100644 frontend/src/features/prompt-audit/components/PolicyPanel.vue create mode 100644 frontend/src/features/prompt-audit/components/RuntimeOverview.vue create mode 100644 frontend/src/features/prompt-audit/types.ts create mode 100644 frontend/src/features/prompt-audit/viewModel.ts create mode 100644 frontend/src/i18n/locales/en/admin/promptAudit.ts create mode 100644 frontend/src/i18n/locales/zh/admin/promptAudit.ts create mode 100644 openspec/changes/add-openai-compatible-prompt-audit/.openspec.yaml create mode 100644 openspec/changes/add-openai-compatible-prompt-audit/README.md create mode 100644 openspec/changes/add-openai-compatible-prompt-audit/design.md create mode 100644 openspec/changes/add-openai-compatible-prompt-audit/implementation-evidence.md create mode 100644 openspec/changes/add-openai-compatible-prompt-audit/implementation-guide.md create mode 100644 openspec/changes/add-openai-compatible-prompt-audit/proposal.md create mode 100644 openspec/changes/add-openai-compatible-prompt-audit/source-baseline.md create mode 100644 openspec/changes/add-openai-compatible-prompt-audit/source-feature-map.md create mode 100644 openspec/changes/add-openai-compatible-prompt-audit/source-freeze/MANIFEST.md create mode 100644 openspec/changes/add-openai-compatible-prompt-audit/source-freeze/aicodex-prompt-audit-tracked.patch create mode 100644 openspec/changes/add-openai-compatible-prompt-audit/source-freeze/aicodex-prompt-audit-untracked.tar.gz create mode 100644 openspec/changes/add-openai-compatible-prompt-audit/specs/prompt-input-audit/spec.md create mode 100644 openspec/changes/add-openai-compatible-prompt-audit/specs/prompt-input-guard/spec.md create mode 100644 openspec/changes/add-openai-compatible-prompt-audit/specs/security-audit-console/spec.md create mode 100644 openspec/changes/add-openai-compatible-prompt-audit/tasks.md create mode 100644 openspec/changes/add-openai-compatible-prompt-audit/verification.md create mode 100644 openspec/config.yaml diff --git a/backend/cmd/server/main.go b/backend/cmd/server/main.go index 784f309fc..c9b2b41a4 100644 --- a/backend/cmd/server/main.go +++ b/backend/cmd/server/main.go @@ -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() { diff --git a/backend/cmd/server/wire.go b/backend/cmd/server/wire.go index f5baee9da..cd4454f52 100644 --- a/backend/cmd/server/wire.go +++ b/backend/cmd/server/wire.go @@ -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() diff --git a/backend/cmd/server/wire_gen.go b/backend/cmd/server/wire_gen.go index 6f9b9c6d2..ae9aec86d 100644 --- a/backend/cmd/server/wire_gen.go +++ b/backend/cmd/server/wire_gen.go @@ -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() diff --git a/backend/cmd/server/wire_gen_test.go b/backend/cmd/server/wire_gen_test.go index 4343dbb72..f9a479bc6 100644 --- a/backend/cmd/server/wire_gen_test.go +++ b/backend/cmd/server/wire_gen_test.go @@ -85,6 +85,7 @@ func TestProvideCleanup_WithMinimalDependencies_NoPanic(t *testing.T) { nil, // quotaFlusher nil, // upstreamBillingProbe nil, // auditLog + nil, // promptAudit ) require.NotPanics(t, func() { diff --git a/backend/internal/handler/batch_image_handler.go b/backend/internal/handler/batch_image_handler.go index d739cba99..b93c5a65a 100644 --- a/backend/internal/handler/batch_image_handler.go +++ b/backend/internal/handler/batch_image_handler.go @@ -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 { diff --git a/backend/internal/handler/content_moderation_helper.go b/backend/internal/handler/content_moderation_helper.go index af6dbd8ee..f91fd8a8f 100644 --- a/backend/internal/handler/content_moderation_helper.go +++ b/backend/internal/handler/content_moderation_helper.go @@ -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 diff --git a/backend/internal/handler/gateway_handler.go b/backend/internal/handler/gateway_handler.go index 7a23984d9..43139a87e 100644 --- a/backend/internal/handler/gateway_handler.go +++ b/backend/internal/handler/gateway_handler.go @@ -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 } diff --git a/backend/internal/handler/gateway_handler_chat_completions.go b/backend/internal/handler/gateway_handler_chat_completions.go index 84456a904..326383479 100644 --- a/backend/internal/handler/gateway_handler_chat_completions.go +++ b/backend/internal/handler/gateway_handler_chat_completions.go @@ -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 } diff --git a/backend/internal/handler/gateway_handler_responses.go b/backend/internal/handler/gateway_handler_responses.go index ce1ae4081..b708ab5ce 100644 --- a/backend/internal/handler/gateway_handler_responses.go +++ b/backend/internal/handler/gateway_handler_responses.go @@ -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 } diff --git a/backend/internal/handler/gemini_v1beta_handler.go b/backend/internal/handler/gemini_v1beta_handler.go index c595abfe5..f570d8a6c 100644 --- a/backend/internal/handler/gemini_v1beta_handler.go +++ b/backend/internal/handler/gemini_v1beta_handler.go @@ -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 } diff --git a/backend/internal/handler/grok_media.go b/backend/internal/handler/grok_media.go index eea6149e8..61cebbd3e 100644 --- a/backend/internal/handler/grok_media.go +++ b/backend/internal/handler/grok_media.go @@ -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 } } diff --git a/backend/internal/handler/handler.go b/backend/internal/handler/handler.go index ce0228b3f..13d3076b5 100644 --- a/backend/internal/handler/handler.go +++ b/backend/internal/handler/handler.go @@ -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 diff --git a/backend/internal/handler/image_task_handler.go b/backend/internal/handler/image_task_handler.go index 29d4cc003..dfbb870db 100644 --- a/backend/internal/handler/image_task_handler.go +++ b/backend/internal/handler/image_task_handler.go @@ -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") diff --git a/backend/internal/handler/openai_alpha_search.go b/backend/internal/handler/openai_alpha_search.go index 1457c4de8..eca8931e8 100644 --- a/backend/internal/handler/openai_alpha_search.go +++ b/backend/internal/handler/openai_alpha_search.go @@ -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) diff --git a/backend/internal/handler/openai_chat_completions.go b/backend/internal/handler/openai_chat_completions.go index 84db86ce8..974d35ded 100644 --- a/backend/internal/handler/openai_chat_completions.go +++ b/backend/internal/handler/openai_chat_completions.go @@ -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) { diff --git a/backend/internal/handler/openai_embeddings.go b/backend/internal/handler/openai_embeddings.go index 473bfea5e..f2e883513 100644 --- a/backend/internal/handler/openai_embeddings.go +++ b/backend/internal/handler/openai_embeddings.go @@ -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) diff --git a/backend/internal/handler/openai_gateway_handler.go b/backend/internal/handler/openai_gateway_handler.go index ea23c0cf3..f59bb08e0 100644 --- a/backend/internal/handler/openai_gateway_handler.go +++ b/backend/internal/handler/openai_gateway_handler.go @@ -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 }, diff --git a/backend/internal/handler/openai_images.go b/backend/internal/handler/openai_images.go index d585c159a..c70dc55a4 100644 --- a/backend/internal/handler/openai_images.go +++ b/backend/internal/handler/openai_images.go @@ -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) diff --git a/backend/internal/handler/security_audit_errors.go b/backend/internal/handler/security_audit_errors.go new file mode 100644 index 000000000..b6f42c711 --- /dev/null +++ b/backend/internal/handler/security_audit_errors.go @@ -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 +} diff --git a/backend/internal/handler/security_audit_errors_test.go b/backend/internal/handler/security_audit_errors_test.go new file mode 100644 index 000000000..9a4f90be4 --- /dev/null +++ b/backend/internal/handler/security_audit_errors_test.go @@ -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) +} diff --git a/backend/internal/handler/security_audit_helper.go b/backend/internal/handler/security_audit_helper.go new file mode 100644 index 000000000..03b6e1de5 --- /dev/null +++ b/backend/internal/handler/security_audit_helper.go @@ -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 +} diff --git a/backend/internal/handler/security_audit_media_submit_test.go b/backend/internal/handler/security_audit_media_submit_test.go new file mode 100644 index 000000000..014fd3c1b --- /dev/null +++ b/backend/internal/handler/security_audit_media_submit_test.go @@ -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) + }) + } +} diff --git a/backend/internal/handler/security_audit_order_test.go b/backend/internal/handler/security_audit_order_test.go new file mode 100644 index 000000000..e9cb4e1f9 --- /dev/null +++ b/backend/internal/handler/security_audit_order_test.go @@ -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 "" +} diff --git a/backend/internal/handler/wire.go b/backend/internal/handler/wire.go index 9a5ea696f..57bf64d02 100644 --- a/backend/internal/handler/wire.go +++ b/backend/internal/handler/wire.go @@ -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, diff --git a/backend/internal/securityaudit/coordinator.go b/backend/internal/securityaudit/coordinator.go new file mode 100644 index 000000000..e6d9c0d92 --- /dev/null +++ b/backend/internal/securityaudit/coordinator.go @@ -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} +} diff --git a/backend/internal/securityaudit/coordinator_legacy.go b/backend/internal/securityaudit/coordinator_legacy.go new file mode 100644 index 000000000..2a73c923f --- /dev/null +++ b/backend/internal/securityaudit/coordinator_legacy.go @@ -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 +} diff --git a/backend/internal/securityaudit/coordinator_test.go b/backend/internal/securityaudit/coordinator_test.go new file mode 100644 index 000000000..936f2a7c1 --- /dev/null +++ b/backend/internal/securityaudit/coordinator_test.go @@ -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()) + } +} diff --git a/backend/internal/securityaudit/prompt_config.go b/backend/internal/securityaudit/prompt_config.go new file mode 100644 index 000000000..ed05e10bc --- /dev/null +++ b/backend/internal/securityaudit/prompt_config.go @@ -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 +} diff --git a/backend/internal/securityaudit/prompt_config_integration_test.go b/backend/internal/securityaudit/prompt_config_integration_test.go new file mode 100644 index 000000000..3817abe05 --- /dev/null +++ b/backend/internal/securityaudit/prompt_config_integration_test.go @@ -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) +} diff --git a/backend/internal/securityaudit/prompt_config_store.go b/backend/internal/securityaudit/prompt_config_store.go new file mode 100644 index 000000000..b0781c858 --- /dev/null +++ b/backend/internal/securityaudit/prompt_config_store.go @@ -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 +} diff --git a/backend/internal/securityaudit/prompt_config_test.go b/backend/internal/securityaudit/prompt_config_test.go new file mode 100644 index 000000000..1023e96fc --- /dev/null +++ b/backend/internal/securityaudit/prompt_config_test.go @@ -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)) + }) + } +} diff --git a/backend/internal/securityaudit/prompt_enqueue.go b/backend/internal/securityaudit/prompt_enqueue.go new file mode 100644 index 000000000..625636c90 --- /dev/null +++ b/backend/internal/securityaudit/prompt_enqueue.go @@ -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() + } +} diff --git a/backend/internal/securityaudit/prompt_event_repository.go b/backend/internal/securityaudit/prompt_event_repository.go new file mode 100644 index 000000000..3960aead3 --- /dev/null +++ b/backend/internal/securityaudit/prompt_event_repository.go @@ -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() +} diff --git a/backend/internal/securityaudit/prompt_guard.go b/backend/internal/securityaudit/prompt_guard.go new file mode 100644 index 000000000..4bab06c10 --- /dev/null +++ b/backend/internal/securityaudit/prompt_guard.go @@ -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 +} diff --git a/backend/internal/securityaudit/prompt_guard_test.go b/backend/internal/securityaudit/prompt_guard_test.go new file mode 100644 index 000000000..c48993681 --- /dev/null +++ b/backend/internal/securityaudit/prompt_guard_test.go @@ -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) +} diff --git a/backend/internal/securityaudit/prompt_handler.go b/backend/internal/securityaudit/prompt_handler.go new file mode 100644 index 000000000..b2d5a51a6 --- /dev/null +++ b/backend/internal/securityaudit/prompt_handler.go @@ -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 +} diff --git a/backend/internal/securityaudit/prompt_handler_test.go b/backend/internal/securityaudit/prompt_handler_test.go new file mode 100644 index 000000000..51a4305f3 --- /dev/null +++ b/backend/internal/securityaudit/prompt_handler_test.go @@ -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") +} diff --git a/backend/internal/securityaudit/prompt_issue_summary.go b/backend/internal/securityaudit/prompt_issue_summary.go new file mode 100644 index 000000000..6f7aeed93 --- /dev/null +++ b/backend/internal/securityaudit/prompt_issue_summary.go @@ -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 "允许" + } +} diff --git a/backend/internal/securityaudit/prompt_logging.go b/backend/internal/securityaudit/prompt_logging.go new file mode 100644 index 000000000..4085c8512 --- /dev/null +++ b/backend/internal/securityaudit/prompt_logging.go @@ -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) +} diff --git a/backend/internal/securityaudit/prompt_logging_test.go b/backend/internal/securityaudit/prompt_logging_test.go new file mode 100644 index 000000000..ff936a1ce --- /dev/null +++ b/backend/internal/securityaudit/prompt_logging_test.go @@ -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"]) +} diff --git a/backend/internal/securityaudit/prompt_metrics.go b/backend/internal/securityaudit/prompt_metrics.go new file mode 100644 index 000000000..29e1e274a --- /dev/null +++ b/backend/internal/securityaudit/prompt_metrics.go @@ -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) + } +} diff --git a/backend/internal/securityaudit/prompt_metrics_test.go b/backend/internal/securityaudit/prompt_metrics_test.go new file mode 100644 index 000000000..29f61e935 --- /dev/null +++ b/backend/internal/securityaudit/prompt_metrics_test.go @@ -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() +} diff --git a/backend/internal/securityaudit/prompt_module.go b/backend/internal/securityaudit/prompt_module.go new file mode 100644 index 000000000..691cc7dbd --- /dev/null +++ b/backend/internal/securityaudit/prompt_module.go @@ -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, +) diff --git a/backend/internal/securityaudit/prompt_outbound_security.go b/backend/internal/securityaudit/prompt_outbound_security.go new file mode 100644 index 000000000..0e0982800 --- /dev/null +++ b/backend/internal/securityaudit/prompt_outbound_security.go @@ -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 +} diff --git a/backend/internal/securityaudit/prompt_outbound_security_test.go b/backend/internal/securityaudit/prompt_outbound_security_test.go new file mode 100644 index 000000000..bd687a67e --- /dev/null +++ b/backend/internal/securityaudit/prompt_outbound_security_test.go @@ -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, + } +} diff --git a/backend/internal/securityaudit/prompt_payload_store.go b/backend/internal/securityaudit/prompt_payload_store.go new file mode 100644 index 000000000..42e30a122 --- /dev/null +++ b/backend/internal/securityaudit/prompt_payload_store.go @@ -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) +} diff --git a/backend/internal/securityaudit/prompt_payload_store_integration_test.go b/backend/internal/securityaudit/prompt_payload_store_integration_test.go new file mode 100644 index 000000000..0558963fc --- /dev/null +++ b/backend/internal/securityaudit/prompt_payload_store_integration_test.go @@ -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) +} diff --git a/backend/internal/securityaudit/prompt_qwen3guard.go b/backend/internal/securityaudit/prompt_qwen3guard.go new file mode 100644 index 000000000..aa3e9ea9d --- /dev/null +++ b/backend/internal/securityaudit/prompt_qwen3guard.go @@ -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 "" + } + 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 +} diff --git a/backend/internal/securityaudit/prompt_qwen3guard_test.go b/backend/internal/securityaudit/prompt_qwen3guard_test.go new file mode 100644 index 000000000..5db035401 --- /dev/null +++ b/backend/internal/securityaudit/prompt_qwen3guard_test.go @@ -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) + } +} diff --git a/backend/internal/securityaudit/prompt_repository.go b/backend/internal/securityaudit/prompt_repository.go new file mode 100644 index 000000000..2190048b7 --- /dev/null +++ b/backend/internal/securityaudit/prompt_repository.go @@ -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 +} diff --git a/backend/internal/securityaudit/prompt_repository_integration_test.go b/backend/internal/securityaudit/prompt_repository_integration_test.go new file mode 100644 index 000000000..33a940f9b --- /dev/null +++ b/backend/internal/securityaudit/prompt_repository_integration_test.go @@ -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) +} diff --git a/backend/internal/securityaudit/prompt_scanner.go b/backend/internal/securityaudit/prompt_scanner.go new file mode 100644 index 000000000..26fca2f94 --- /dev/null +++ b/backend/internal/securityaudit/prompt_scanner.go @@ -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 +} diff --git a/backend/internal/securityaudit/prompt_service.go b/backend/internal/securityaudit/prompt_service.go new file mode 100644 index 000000000..05607a8c7 --- /dev/null +++ b/backend/internal/securityaudit/prompt_service.go @@ -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 +} diff --git a/backend/internal/securityaudit/prompt_service_test.go b/backend/internal/securityaudit/prompt_service_test.go new file mode 100644 index 000000000..93d82df56 --- /dev/null +++ b/backend/internal/securityaudit/prompt_service_test.go @@ -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) + }) + } +} diff --git a/backend/internal/securityaudit/prompt_snapshot.go b/backend/internal/securityaudit/prompt_snapshot.go new file mode 100644 index 000000000..b39d92179 --- /dev/null +++ b/backend/internal/securityaudit/prompt_snapshot.go @@ -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 +} diff --git a/backend/internal/securityaudit/prompt_snapshot_test.go b/backend/internal/securityaudit/prompt_snapshot_test.go new file mode 100644 index 000000000..5765e1a51 --- /dev/null +++ b/backend/internal/securityaudit/prompt_snapshot_test.go @@ -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 +} diff --git a/backend/internal/securityaudit/prompt_types.go b/backend/internal/securityaudit/prompt_types.go new file mode 100644 index 000000000..bf38b6463 --- /dev/null +++ b/backend/internal/securityaudit/prompt_types.go @@ -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) +} diff --git a/backend/internal/securityaudit/prompt_worker.go b/backend/internal/securityaudit/prompt_worker.go new file mode 100644 index 000000000..5db007f6b --- /dev/null +++ b/backend/internal/securityaudit/prompt_worker.go @@ -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 +} diff --git a/backend/internal/securityaudit/prompt_worker_test.go b/backend/internal/securityaudit/prompt_worker_test.go new file mode 100644 index 000000000..9cbd6c1cb --- /dev/null +++ b/backend/internal/securityaudit/prompt_worker_test.go @@ -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()) +} diff --git a/backend/internal/server/middleware/audit_log.go b/backend/internal/server/middleware/audit_log.go index 1cb510f06..15d9996de 100644 --- a/backend/internal/server/middleware/audit_log.go +++ b/backend/internal/server/middleware/audit_log.go @@ -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 { diff --git a/backend/internal/server/middleware/audit_log_test.go b/backend/internal/server/middleware/audit_log_test.go index 752bed9f8..026e9163b 100644 --- a/backend/internal/server/middleware/audit_log_test.go +++ b/backend/internal/server/middleware/audit_log_test.go @@ -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, "", 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) + } +} diff --git a/backend/internal/server/routes/admin.go b/backend/internal/server/routes/admin.go index 844f75eae..1f0011509 100644 --- a/backend/internal/server/routes/admin.go +++ b/backend/internal/server/routes/admin.go @@ -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") { diff --git a/backend/internal/server/routes/prompt_audit_route_coverage_test.go b/backend/internal/server/routes/prompt_audit_route_coverage_test.go new file mode 100644 index 000000000..ba38910b4 --- /dev/null +++ b/backend/internal/server/routes/prompt_audit_route_coverage_test.go @@ -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) + }) + } +} diff --git a/backend/migrations/181_prompt_audit.sql b/backend/migrations/181_prompt_audit.sql new file mode 100644 index 000000000..6ca18f872 --- /dev/null +++ b/backend/migrations/181_prompt_audit.sql @@ -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); diff --git a/frontend/src/components/layout/AppSidebar.vue b/frontend/src/components/layout/AppSidebar.vue index c66a7bfbf..ea89b8d80 100644 --- a/frontend/src/components/layout/AppSidebar.vue +++ b/frontend/src/components/layout/AppSidebar.vue @@ -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 }, { diff --git a/frontend/src/features/prompt-audit/PromptAuditView.vue b/frontend/src/features/prompt-audit/PromptAuditView.vue new file mode 100644 index 000000000..e0491bbdc --- /dev/null +++ b/frontend/src/features/prompt-audit/PromptAuditView.vue @@ -0,0 +1,350 @@ + + + diff --git a/frontend/src/features/prompt-audit/__tests__/PromptAuditView.spec.ts b/frontend/src/features/prompt-audit/__tests__/PromptAuditView.spec.ts new file mode 100644 index 000000000..1f4a0598f --- /dev/null +++ b/frontend/src/features/prompt-audit/__tests__/PromptAuditView.spec.ts @@ -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('vue-i18n') + return { ...actual, useI18n: () => ({ locale: { value: 'en' }, t: (key: string, params?: Record) => 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: '
' } +const RuntimeStub = defineComponent({ props: ['runtime', 'loading', 'error'], emits: ['refresh'], template: '
{{ error }}
' }) +const EndpointStub = defineComponent({ + props: ['endpoints', 'probeResults', 'probingIds'], emits: ['update:endpoints', 'probe'], + template: '
', +}) +const PolicyStub = defineComponent({ props: ['draft', 'groups'], emits: ['update:draft'], template: '
' }) +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: '
', +}) +const DetailStub = defineComponent({ props: ['show', 'event', 'loading'], emits: ['close'], template: '
' }) +const DialogStub = defineComponent({ props: ['show', 'title'], emits: ['close'], template: '
' }) +const ConfirmStub = defineComponent({ props: ['show', 'title', 'message'], emits: ['confirm', 'cancel'], template: '
' }) + +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)).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', + })) + }) +}) diff --git a/frontend/src/features/prompt-audit/__tests__/api.spec.ts b/frontend/src/features/prompt-audit/__tests__/api.spec.ts new file mode 100644 index 000000000..ead59e2f3 --- /dev/null +++ b/frontend/src/features/prompt-audit/__tests__/api.spec.ts @@ -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, + })) + }) +}) diff --git a/frontend/src/features/prompt-audit/__tests__/components.spec.ts b/frontend/src/features/prompt-audit/__tests__/components.spec.ts new file mode 100644 index 000000000..62c8a0a10 --- /dev/null +++ b/frontend/src/features/prompt-audit/__tests__/components.spec.ts @@ -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('vue-i18n') + return { ...actual, useI18n: () => ({ locale: { value: 'en' }, t: (key: string, params?: Record) => key.replace(/\{(\w+)\}/g, (_, token) => String(params?.[token] ?? `{${token}}`)) }) } +}) + +const DialogStub = defineComponent({ props: ['show', 'title'], emits: ['close'], template: '
' }) +const PaginationStub = defineComponent({ props: ['total', 'page', 'pageSize'], emits: ['update:page', 'update:pageSize'], template: '
' }) + +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('[aria-label="admin.promptAudit.pool.apiKey"]') + expect(token.element.value).toBe('') + expect(token.attributes('placeholder')).toContain('admin.promptAudit.pool.keepSecret') + + await wrapper.get('[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]) + }) +}) diff --git a/frontend/src/features/prompt-audit/__tests__/integrationSurface.spec.ts b/frontend/src/features/prompt-audit/__tests__/integrationSurface.spec.ts new file mode 100644 index 000000000..3dd006ad9 --- /dev/null +++ b/frontend/src/features/prompt-audit/__tests__/integrationSurface.spec.ts @@ -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') + }) +}) diff --git a/frontend/src/features/prompt-audit/__tests__/viewModel.spec.ts b/frontend/src/features/prompt-audit/__tests__/viewModel.spec.ts new file mode 100644 index 000000000..f4390466a --- /dev/null +++ b/frontend/src/features/prompt-audit/__tests__/viewModel.spec.ts @@ -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(), + }) + }) +}) diff --git a/frontend/src/features/prompt-audit/api.ts b/frontend/src/features/prompt-audit/api.ts new file mode 100644 index 000000000..05fe2aa38 --- /dev/null +++ b/frontend/src/features/prompt-audit/api.ts @@ -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 { + const { data } = await apiClient.get(`${basePath}/config`) + return data +} + +export async function updateConfig(payload: PromptAuditUpdateRequest): Promise { + const { data } = await apiClient.put(`${basePath}/config`, payload) + return data +} + +export async function probeEndpoint(endpoint: PromptAuditEndpointDraft): Promise { + const { data } = await apiClient.post(`${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 { + const { data } = await apiClient.get(`${basePath}/runtime`) + return data +} + +export async function listEvents( + filters: PromptEventFilters, + page: number, + pageSize: number, +): Promise { + const { data } = await apiClient.get(`${basePath}/events`, { + params: { page, page_size: pageSize, ...eventQueryParams(filters) }, + }) + return data +} + +export async function getEvent(id: number): Promise { + const { data } = await apiClient.get(`${basePath}/events/${id}`) + return data +} + +export async function deleteEvent(id: number): Promise { + const { data } = await apiClient.delete(`${basePath}/events/${id}`) + return data +} + +export async function batchDeleteEvents(ids: number[]): Promise { + const { data } = await apiClient.post(`${basePath}/events/batch-delete`, { ids }) + return data +} + +export async function previewDelete(filters: PromptEventFilters): Promise { + const { data } = await apiClient.post( + `${basePath}/events/delete-preview`, + eventFilterPayload(filters), + ) + return data +} + +export async function deleteEventsByFilter( + filters: PromptEventFilters, + preview: PromptDeletePreview, +): Promise { + const { data } = await apiClient.post(`${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 { + const { data } = await apiClient.get('/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 diff --git a/frontend/src/features/prompt-audit/components/EndpointPool.vue b/frontend/src/features/prompt-audit/components/EndpointPool.vue new file mode 100644 index 000000000..fc132358c --- /dev/null +++ b/frontend/src/features/prompt-audit/components/EndpointPool.vue @@ -0,0 +1,172 @@ + + + diff --git a/frontend/src/features/prompt-audit/components/EventDetailDialog.vue b/frontend/src/features/prompt-audit/components/EventDetailDialog.vue new file mode 100644 index 000000000..c606a394b --- /dev/null +++ b/frontend/src/features/prompt-audit/components/EventDetailDialog.vue @@ -0,0 +1,66 @@ + + + diff --git a/frontend/src/features/prompt-audit/components/EventWorkspace.vue b/frontend/src/features/prompt-audit/components/EventWorkspace.vue new file mode 100644 index 000000000..1bf166a89 --- /dev/null +++ b/frontend/src/features/prompt-audit/components/EventWorkspace.vue @@ -0,0 +1,186 @@ + + + diff --git a/frontend/src/features/prompt-audit/components/PolicyPanel.vue b/frontend/src/features/prompt-audit/components/PolicyPanel.vue new file mode 100644 index 000000000..9df3408a7 --- /dev/null +++ b/frontend/src/features/prompt-audit/components/PolicyPanel.vue @@ -0,0 +1,108 @@ + + + diff --git a/frontend/src/features/prompt-audit/components/RuntimeOverview.vue b/frontend/src/features/prompt-audit/components/RuntimeOverview.vue new file mode 100644 index 000000000..58dd68dc3 --- /dev/null +++ b/frontend/src/features/prompt-audit/components/RuntimeOverview.vue @@ -0,0 +1,119 @@ + + + diff --git a/frontend/src/features/prompt-audit/types.ts b/frontend/src/features/prompt-audit/types.ts new file mode 100644 index 000000000..57571953d --- /dev/null +++ b/frontend/src/features/prompt-audit/types.ts @@ -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 { + 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 + 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 + scanner_evidence: Record + 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 + 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 +} diff --git a/frontend/src/features/prompt-audit/viewModel.ts b/frontend/src/features/prompt-audit/viewModel.ts new file mode 100644 index 000000000..6308ce73d --- /dev/null +++ b/frontend/src/features/prompt-audit/viewModel.ts @@ -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(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 { + const result: Record = {} + 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 { + 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()) +} diff --git a/frontend/src/i18n/locales/en/admin/index.ts b/frontend/src/i18n/locales/en/admin/index.ts index 52e9231c7..224ab94b3 100644 --- a/frontend/src/i18n/locales/en/admin/index.ts +++ b/frontend/src/i18n/locales/en/admin/index.ts @@ -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, } diff --git a/frontend/src/i18n/locales/en/admin/promptAudit.ts b/frontend/src/i18n/locales/en/admin/promptAudit.ts new file mode 100644 index 000000000..cbe8113b3 --- /dev/null +++ b/frontend/src/i18n/locales/en/admin/promptAudit.ts @@ -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.', + }, + }, +} diff --git a/frontend/src/i18n/locales/en/common.ts b/frontend/src/i18n/locales/en/common.ts index 3f9ca1dde..be13f5aaa 100644 --- a/frontend/src/i18n/locales/en/common.ts +++ b/frontend/src/i18n/locales/en/common.ts @@ -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', }, diff --git a/frontend/src/i18n/locales/zh/admin/index.ts b/frontend/src/i18n/locales/zh/admin/index.ts index 52e9231c7..224ab94b3 100644 --- a/frontend/src/i18n/locales/zh/admin/index.ts +++ b/frontend/src/i18n/locales/zh/admin/index.ts @@ -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, } diff --git a/frontend/src/i18n/locales/zh/admin/promptAudit.ts b/frontend/src/i18n/locales/zh/admin/promptAudit.ts new file mode 100644 index 000000000..c39f33e46 --- /dev/null +++ b/frontend/src/i18n/locales/zh/admin/promptAudit.ts @@ -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: '至少需要启用一个风险分类。', + }, + }, +} diff --git a/frontend/src/i18n/locales/zh/common.ts b/frontend/src/i18n/locales/zh/common.ts index 62a354a67..02de40946 100644 --- a/frontend/src/i18n/locales/zh/common.ts +++ b/frontend/src/i18n/locales/zh/common.ts @@ -190,6 +190,9 @@ export default { channelMonitor: '渠道监控', channelStatus: '渠道状态', riskControl: '风控中心', + securityAudit: '安全审计', + contentModeration: '内容审核', + promptAudit: '提示词审计', auditLogs: '操作日志', }, diff --git a/frontend/src/router/index.ts b/frontend/src/router/index.ts index 657d25f8d..259ea5559 100644 --- a/frontend/src/router/index.ts +++ b/frontend/src/router/index.ts @@ -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', diff --git a/openspec/changes/add-openai-compatible-prompt-audit/.openspec.yaml b/openspec/changes/add-openai-compatible-prompt-audit/.openspec.yaml new file mode 100644 index 000000000..cd2ce7e9f --- /dev/null +++ b/openspec/changes/add-openai-compatible-prompt-audit/.openspec.yaml @@ -0,0 +1,2 @@ +schema: spec-driven +created: 2026-07-16 diff --git a/openspec/changes/add-openai-compatible-prompt-audit/README.md b/openspec/changes/add-openai-compatible-prompt-audit/README.md new file mode 100644 index 000000000..7a34daefc --- /dev/null +++ b/openspec/changes/add-openai-compatible-prompt-audit/README.md @@ -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`。 diff --git a/openspec/changes/add-openai-compatible-prompt-audit/design.md b/openspec/changes/add-openai-compatible-prompt-audit/design.md new file mode 100644 index 000000000..787f697a3 --- /dev/null +++ b/openspec/changes/add-openai-compatible-prompt-audit/design.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: 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": ""}], + "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=`。 +- Claude Messages:保持 `type=error` 和合法的 `error.type=permission_error|api_error`,增加可选 `error.code=`。 +- Gemini:保持 Google envelope 的数值 `error.code`、message 和 canonical status;在 `error.details[]` 增加 `type.googleapis.com/google.rpc.ErrorInfo`,其 `reason=`、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` 的建议阈值并由安全、运营和业务责任人签字,未签字不得生产开启。 diff --git a/openspec/changes/add-openai-compatible-prompt-audit/implementation-evidence.md b/openspec/changes/add-openai-compatible-prompt-audit/implementation-evidence.md new file mode 100644 index 000000000..b31dc3e73 --- /dev/null +++ b/openspec/changes/add-openai-compatible-prompt-audit/implementation-evidence.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. diff --git a/openspec/changes/add-openai-compatible-prompt-audit/implementation-guide.md b/openspec/changes/add-openai-compatible-prompt-audit/implementation-guide.md new file mode 100644 index 000000000..80636caf3 --- /dev/null +++ b/openspec/changes/add-openai-compatible-prompt-audit/implementation-guide.md @@ -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 阈值、告警、值班步骤和一键回滚已由责任人签字确认。 diff --git a/openspec/changes/add-openai-compatible-prompt-audit/proposal.md b/openspec/changes/add-openai-compatible-prompt-audit/proposal.md new file mode 100644 index 000000000..829d6391d --- /dev/null +++ b/openspec/changes/add-openai-compatible-prompt-audit/proposal.md @@ -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 的证据矩阵、协议测试、泄露门禁、灰度阈值和回滚手册。 diff --git a/openspec/changes/add-openai-compatible-prompt-audit/source-baseline.md b/openspec/changes/add-openai-compatible-prompt-audit/source-baseline.md new file mode 100644 index 000000000..0a8036b76 --- /dev/null +++ b/openspec/changes/add-openai-compatible-prompt-audit/source-baseline.md @@ -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 校验和、执行目录和完整输出。 diff --git a/openspec/changes/add-openai-compatible-prompt-audit/source-feature-map.md b/openspec/changes/add-openai-compatible-prompt-audit/source-feature-map.md new file mode 100644 index 000000000..bc2912dcf --- /dev/null +++ b/openspec/changes/add-openai-compatible-prompt-audit/source-feature-map.md @@ -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 节任一行没有通过证据都表示“完整迁移”未完成。 diff --git a/openspec/changes/add-openai-compatible-prompt-audit/source-freeze/MANIFEST.md b/openspec/changes/add-openai-compatible-prompt-audit/source-freeze/MANIFEST.md new file mode 100644 index 000000000..a7db2eb56 --- /dev/null +++ b/openspec/changes/add-openai-compatible-prompt-audit/source-freeze/MANIFEST.md @@ -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. diff --git a/openspec/changes/add-openai-compatible-prompt-audit/source-freeze/aicodex-prompt-audit-tracked.patch b/openspec/changes/add-openai-compatible-prompt-audit/source-freeze/aicodex-prompt-audit-tracked.patch new file mode 100644 index 000000000..d3f47171b --- /dev/null +++ b/openspec/changes/add-openai-compatible-prompt-audit/source-freeze/aicodex-prompt-audit-tracked.patch @@ -0,0 +1,2771 @@ +diff --git a/ai-gateway/cmd/aicodex/main.go b/ai-gateway/cmd/aicodex/main.go +index 59d11dcb90f59c9868ca836b2acf1a827d78eeb0..e57948b376df63083dd1c24dd3a778707a9793f3 100644 +--- a/ai-gateway/cmd/aicodex/main.go ++++ b/ai-gateway/cmd/aicodex/main.go +@@ -85,7 +85,7 @@ var ( + migrationInitLogDBFn = model.InitLogDB + migrationCloseDBFn = model.CloseDB + promptAuditRunnerFactory = func() *promptaudit.Runner { +- return promptaudit.NewRunner(nil, nil, promptaudit.NewOpenAICompatibleClient(service.GetHttpClient()), nil) ++ return promptaudit.NewRunner(nil, nil, promptaudit.NewOpenAICompatibleClient(nil), nil) + } + ) + +diff --git a/ai-gateway/internal/controller/prompt_audit.go b/ai-gateway/internal/controller/prompt_audit.go +index 25a71a76ee20d9f6a950074c6dc83995a67fe75d..dc34813b01f25864c5451ac3f947e1331b2e1672 100644 +--- a/ai-gateway/internal/controller/prompt_audit.go ++++ b/ai-gateway/internal/controller/prompt_audit.go +@@ -1,6 +1,7 @@ + package controller + + import ( ++ "errors" + "net/http" + "strconv" + "strings" +@@ -9,6 +10,7 @@ import ( + "github.com/gin-gonic/gin" + appent "github.com/mt21625457/aicodex/ent" + "github.com/mt21625457/aicodex/internal/common" ++ "github.com/mt21625457/aicodex/internal/constant" + "github.com/mt21625457/aicodex/internal/service/promptaudit" + ) + +@@ -48,12 +50,21 @@ type promptAuditEventFilterRequest struct { + var ( + previewDeletePromptAuditEventsByFilter = promptaudit.PreviewDeleteEventsByFilter + deletePromptAuditEventsByFilter = promptaudit.DeleteEventsByFilter ++ promptAuditConfigServiceFactory = func() *promptaudit.ConfigService { return promptaudit.NewConfigService(nil) } + ) + + func GetPromptAuditConfig(c *gin.Context) { +- cfg, err := promptaudit.NewConfigService(nil).Public(c.Request.Context()) ++ cfg, err := promptAuditConfigServiceFactory().Public(c.Request.Context()) + if err != nil { +- common.ApiError(c, err) ++ promptaudit.LogWarnEvent( ++ "prompt_guard.config_reload_degraded", ++ promptAuditLogFields(c, ++ promptaudit.Field("status", "degraded"), ++ promptaudit.Field("error_code", "config_read_failed"), ++ promptaudit.Field("error_kind", "config_read_failed"), ++ )..., ++ ) ++ common.ApiErrorMsg(c, "读取提示词审计配置失败") + return + } + common.ApiSuccess(c, cfg) +@@ -68,8 +79,17 @@ func UpdatePromptAuditConfig(c *gin.Context) { + common.ApiErrorMsg(c, "invalid request body") + return + } +- cfg, err := promptaudit.NewConfigService(nil).Save(c.Request.Context(), req) ++ req.UpdatedBy = common.GetContextKeyInt(c, constant.ContextKeyUserId) ++ cfg, err := promptAuditConfigServiceFactory().Save(c.Request.Context(), req) + if err != nil { ++ var validationErr *promptaudit.ConfigValidationError ++ if errors.As(err, &validationErr) { ++ recordAdminAuditFailed(c, "prompt_audit.config.update", "prompt_audit_config", "global", validationErr.Code, promptAuditAdminAuditDetail(map[string]any{ ++ "enabled": req.Enabled, "blocking_enabled": req.BlockingEnabled, ++ })) ++ c.JSON(http.StatusBadRequest, gin.H{"success": false, "message": validationErr.Message, "code": validationErr.Code}) ++ return ++ } + recordAdminAuditFailed(c, "prompt_audit.config.update", "prompt_audit_config", "global", "save_config_failed", promptAuditAdminAuditDetail(map[string]any{ + "enabled": req.Enabled, + "endpoint_count": len(req.Endpoints), +@@ -77,7 +97,7 @@ func UpdatePromptAuditConfig(c *gin.Context) { + "audit_group_mode": req.AuditGroupMode, + "audit_group_count": len(req.AuditGroups), + "audit_group_hash": promptaudit.Config{AuditGroups: req.AuditGroups}.AuditGroupHash(), +- "error": err.Error(), ++ "error_code": "save_config_failed", + })) + promptaudit.LogWarnEvent( + "prompt_audit.config_updated", +@@ -91,15 +111,16 @@ func UpdatePromptAuditConfig(c *gin.Context) { + promptaudit.Field("audit_group_count", len(req.AuditGroups)), + promptaudit.Field("audit_group_hash", promptaudit.Config{AuditGroups: req.AuditGroups}.AuditGroupHash()), + promptaudit.Field("error_code", "save_config_failed"), +- promptaudit.Field("error_kind", err.Error()), ++ promptaudit.Field("error_kind", "config_save_failed"), + )..., + ) +- common.ApiErrorMsg(c, err.Error()) ++ common.ApiErrorMsg(c, "保存提示词审计配置失败") + return + } +- promptaudit.ClearConfigCache() + recordAdminAuditSuccess(c, "prompt_audit.config.update", "prompt_audit_config", "global", promptAuditAdminAuditDetail(map[string]any{ + "enabled": cfg.Enabled, ++ "blocking_enabled": cfg.BlockingEnabled, ++ "config_version": cfg.ConfigVersion, + "endpoint_count": len(cfg.Endpoints), + "scanner_count": len(cfg.Scanners), + "worker_count": cfg.WorkerCount, +@@ -123,6 +144,16 @@ func UpdatePromptAuditConfig(c *gin.Context) { + promptaudit.Field("audit_group_hash", cfg.AuditGroupHash()), + )..., + ) ++ promptaudit.LogInfoEvent( ++ "prompt_guard.config_updated", ++ promptAuditLogFields(c, ++ promptaudit.Field("status", "success"), ++ promptaudit.Field("enabled", cfg.Enabled), ++ promptaudit.Field("blocking_enabled", cfg.BlockingEnabled), ++ promptaudit.Field("config_version", cfg.ConfigVersion), ++ promptaudit.Field("updated_by", req.UpdatedBy), ++ )..., ++ ) + common.ApiSuccess(c, cfg) + } + +diff --git a/ai-gateway/internal/controller/prompt_audit_test.go b/ai-gateway/internal/controller/prompt_audit_test.go +index 76ae4f4c27e2eec3220532fc1dadc562dd2ee9ce..f6e51234283e35a63c77d9ca4a83581c3ee693d8 100644 +--- a/ai-gateway/internal/controller/prompt_audit_test.go ++++ b/ai-gateway/internal/controller/prompt_audit_test.go +@@ -21,6 +21,90 @@ import ( + "github.com/mt21625457/aicodex/internal/testutil" + ) + ++type controllerPromptAuditOptionStore struct { ++ values map[string]string ++ getErr error ++ setErr error ++} ++ ++func (s *controllerPromptAuditOptionStore) GetOption(_ context.Context, key string) (string, bool, error) { ++ if s.getErr != nil { ++ return "", false, s.getErr ++ } ++ value, ok := s.values[key] ++ return value, ok, nil ++} ++ ++func (s *controllerPromptAuditOptionStore) SetOption(_ context.Context, key string, value string) error { ++ if s.setErr != nil { ++ return s.setErr ++ } ++ if s.values == nil { ++ s.values = make(map[string]string) ++ } ++ s.values[key] = value ++ return nil ++} ++ ++func withPromptAuditConfigServiceFactory(t *testing.T, store promptaudit.OptionStore) { ++ t.Helper() ++ previous := promptAuditConfigServiceFactory ++ promptAuditConfigServiceFactory = func() *promptaudit.ConfigService { ++ return promptaudit.NewConfigService(store) ++ } ++ t.Cleanup(func() { ++ promptAuditConfigServiceFactory = previous ++ }) ++} ++ ++func TestPromptAuditConfigAPIRejectsInvalidBlockingCombination(t *testing.T) { ++ withPromptAuditConfigServiceFactory(t, &controllerPromptAuditOptionStore{values: map[string]string{}}) ++ ctx, recorder := newPromptAuditRequestContext(t, http.MethodPut, "/api/prompt-audit/config", `{"enabled":false,"blocking_enabled":true,"strategy":"priority"}`) ++ ++ UpdatePromptAuditConfig(ctx) ++ ++ if recorder.Code != http.StatusBadRequest { ++ t.Fatalf("非法同步阻止组合应返回 400,实际 status=%d body=%s", recorder.Code, recorder.Body.String()) ++ } ++ if body := recorder.Body.String(); !strings.Contains(body, promptaudit.PromptGuardRequiresAuditEnabled) { ++ t.Fatalf("非法同步阻止组合应返回稳定错误码,实际 %s", body) ++ } ++} ++ ++func TestPromptAuditConfigAPIMasksInternalReadAndSaveErrors(t *testing.T) { ++ const sensitiveInternalError = "postgres://admin:super-secret@db.internal/aicodex" ++ ++ t.Run("读取失败", func(t *testing.T) { ++ withPromptAuditConfigServiceFactory(t, &controllerPromptAuditOptionStore{getErr: errors.New(sensitiveInternalError)}) ++ ctx, recorder := newPromptAuditRequestContext(t, http.MethodGet, "/api/prompt-audit/config", "") ++ ++ GetPromptAuditConfig(ctx) ++ ++ body := recorder.Body.String() ++ if strings.Contains(body, sensitiveInternalError) || strings.Contains(body, "super-secret") { ++ t.Fatalf("配置读取错误不得回显内部连接信息: %s", body) ++ } ++ if !strings.Contains(body, "读取提示词审计配置失败") { ++ t.Fatalf("配置读取错误应返回通用消息: %s", body) ++ } ++ }) ++ ++ t.Run("保存失败", func(t *testing.T) { ++ withPromptAuditConfigServiceFactory(t, &controllerPromptAuditOptionStore{values: map[string]string{}, setErr: errors.New(sensitiveInternalError)}) ++ ctx, recorder := newPromptAuditRequestContext(t, http.MethodPut, "/api/prompt-audit/config", `{"enabled":false,"blocking_enabled":false,"strategy":"priority"}`) ++ ++ UpdatePromptAuditConfig(ctx) ++ ++ body := recorder.Body.String() ++ if strings.Contains(body, sensitiveInternalError) || strings.Contains(body, "super-secret") { ++ t.Fatalf("配置保存错误不得回显内部连接信息: %s", body) ++ } ++ if !strings.Contains(body, "保存提示词审计配置失败") { ++ t.Fatalf("配置保存错误应返回通用消息: %s", body) ++ } ++ }) ++} ++ + func TestPromptAuditConfigAPIStoresTokenAsSensitiveValue(t *testing.T) { + withPromptAuditControllerTestDB(t, func() { + gin.SetMode(gin.TestMode) +diff --git a/ai-gateway/internal/gatewayadapter/transport/anthropic.go b/ai-gateway/internal/gatewayadapter/transport/anthropic.go +index 5d15c6bd79bc5530094eb43fde9f47645819cb1a..10dce044a23e590534d6d410e8a46929225dae0c 100644 +--- a/ai-gateway/internal/gatewayadapter/transport/anthropic.go ++++ b/ai-gateway/internal/gatewayadapter/transport/anthropic.go +@@ -113,6 +113,7 @@ func NewAnthropicGatewayHandler(deps AnthropicGatewayDeps) http.Handler { + v1.Use(gatewaycore.RegisterHTTPAuditPostHook()) + v1.Use(middleware.UserConcurrencyLimit()) + v1.Use(middleware.ModelRequestRateLimit()) ++ v1.Use(promptaudit.HTTPGuardMiddleware(types.RelayFormatClaude)) + v1.Use(middleware.Distribute()) + v1.Use(middleware.PriorityAdmission()) + v1.Use(promptaudit.HTTPEnqueueMiddleware(types.RelayFormatClaude)) +diff --git a/ai-gateway/internal/gatewayadapter/transport/gemini.go b/ai-gateway/internal/gatewayadapter/transport/gemini.go +index 53e58476a7fc6ec854766cfac995fb5a4254962e..24d0e12eb16b1fff373fc69eb65268718c81e043 100644 +--- a/ai-gateway/internal/gatewayadapter/transport/gemini.go ++++ b/ai-gateway/internal/gatewayadapter/transport/gemini.go +@@ -100,6 +100,7 @@ func NewGeminiGatewayHandler(deps GeminiGatewayDeps) http.Handler { + relayRouter.Use(gatewaycore.RegisterHTTPAuditPostHook()) + relayRouter.Use(middleware.UserConcurrencyLimit()) + relayRouter.Use(middleware.ModelRequestRateLimit()) ++ relayRouter.Use(promptaudit.HTTPGuardMiddleware(types.RelayFormatGemini)) + relayRouter.Use(middleware.Distribute()) + relayRouter.Use(middleware.PriorityAdmission()) + relayRouter.Use(promptaudit.HTTPEnqueueMiddleware(types.RelayFormatGemini)) +diff --git a/ai-gateway/internal/gatewayadapter/transport/jimeng.go b/ai-gateway/internal/gatewayadapter/transport/jimeng.go +index 8e5ef822bda42b100d8376f635c651a606e4611a..7dbc008716c36ac1ba26d022f0156f82c505876c 100644 +--- a/ai-gateway/internal/gatewayadapter/transport/jimeng.go ++++ b/ai-gateway/internal/gatewayadapter/transport/jimeng.go +@@ -54,6 +54,7 @@ func NewJimengGatewayHandler(deps JimengGatewayDeps) http.Handler { + jimeng.Use(middleware.JimengRequestConvert()) + jimeng.Use(middleware.TokenAuth()) + jimeng.Use(middleware.UserConcurrencyLimit()) ++ jimeng.Use(promptaudit.HTTPGuardMiddleware(types.RelayFormatTask)) + jimeng.Use(middleware.Distribute()) + jimeng.Use(middleware.PriorityAdmission()) + jimeng.Use(promptaudit.HTTPEnqueueMiddleware(types.RelayFormatTask)) +diff --git a/ai-gateway/internal/gatewayadapter/transport/kling.go b/ai-gateway/internal/gatewayadapter/transport/kling.go +index 875dd1157b3cc168da4ee0bfb5036bca98285b55..1c5e91c3deeb39e1e2ddfa2c5c0c2a721819bbba 100644 +--- a/ai-gateway/internal/gatewayadapter/transport/kling.go ++++ b/ai-gateway/internal/gatewayadapter/transport/kling.go +@@ -54,6 +54,7 @@ func NewKlingGatewayHandler(deps KlingGatewayDeps) http.Handler { + kling.Use(middleware.KlingRequestConvert()) + kling.Use(middleware.TokenAuth()) + kling.Use(middleware.UserConcurrencyLimit()) ++ kling.Use(promptaudit.HTTPGuardMiddleware(types.RelayFormatTask)) + kling.Use(middleware.Distribute()) + kling.Use(middleware.PriorityAdmission()) + kling.Use(promptaudit.HTTPEnqueueMiddleware(types.RelayFormatTask)) +diff --git a/ai-gateway/internal/gatewayadapter/transport/midjourney.go b/ai-gateway/internal/gatewayadapter/transport/midjourney.go +index 216dbf70eda862db47766944e03de967c67a2ef7..5d4f3079b8d71791e2e0c113120159034ac3165b 100644 +--- a/ai-gateway/internal/gatewayadapter/transport/midjourney.go ++++ b/ai-gateway/internal/gatewayadapter/transport/midjourney.go +@@ -122,6 +122,7 @@ func registerMidjourneyTransportGroup(group *gin.RouterGroup, deps MidjourneyGat + + group.Use(middleware.TokenAuth()) + group.Use(middleware.UserConcurrencyLimit()) ++ group.Use(promptaudit.HTTPGuardMiddleware(types.RelayFormatMjProxy)) + group.Use(middleware.Distribute()) + group.Use(middleware.PriorityAdmission()) + group.Use(promptaudit.HTTPEnqueueMiddleware(types.RelayFormatMjProxy)) +diff --git a/ai-gateway/internal/gatewayadapter/transport/openai.go b/ai-gateway/internal/gatewayadapter/transport/openai.go +index d8fd0fc79ad78036a6540bf0c1c1e6b5a7b8fded..ea729a96b965174996e5b6f9df58519bfc1d1d09 100644 +--- a/ai-gateway/internal/gatewayadapter/transport/openai.go ++++ b/ai-gateway/internal/gatewayadapter/transport/openai.go +@@ -157,6 +157,7 @@ func NewOpenAIGatewayHandler(deps OpenAIGatewayDeps) http.Handler { + httpRouter.Use(gatewaycore.RegisterHTTPAuditPostHook()) + httpRouter.Use(middleware.UserConcurrencyLimit()) + httpRouter.Use(middleware.ModelRequestRateLimit()) ++ httpRouter.Use(promptaudit.HTTPGuardMiddleware()) + httpRouter.Use(middleware.Distribute()) + httpRouter.Use(middleware.PriorityAdmission()) + httpRouter.Use(promptaudit.HTTPEnqueueMiddleware()) +@@ -195,6 +196,7 @@ func NewOpenAIGatewayHandler(deps OpenAIGatewayDeps) http.Handler { + responsesAliasRouter.Use(gatewaycore.RegisterHTTPAuditPostHook()) + responsesAliasRouter.Use(middleware.UserConcurrencyLimit()) + responsesAliasRouter.Use(middleware.ModelRequestRateLimit()) ++ responsesAliasRouter.Use(promptaudit.HTTPGuardMiddleware(types.RelayFormatOpenAIResponses)) + responsesAliasRouter.Use(middleware.Distribute()) + responsesAliasRouter.Use(middleware.PriorityAdmission()) + responsesAliasRouter.Use(promptaudit.HTTPEnqueueMiddleware()) +diff --git a/ai-gateway/internal/gatewayadapter/transport/suno.go b/ai-gateway/internal/gatewayadapter/transport/suno.go +index 714e282598547c5de444e145e44797ac25d7c4e4..e3532cc3c37ad5a73034a27c4ac9c9dfdd1092db 100644 +--- a/ai-gateway/internal/gatewayadapter/transport/suno.go ++++ b/ai-gateway/internal/gatewayadapter/transport/suno.go +@@ -65,6 +65,7 @@ func NewSunoGatewayHandler(deps SunoGatewayDeps) http.Handler { + suno.Use(middleware.SystemPerformanceCheck()) + suno.Use(middleware.TokenAuth()) + suno.Use(middleware.UserConcurrencyLimit()) ++ suno.Use(promptaudit.HTTPGuardMiddleware(types.RelayFormatTask)) + suno.Use(middleware.Distribute()) + suno.Use(middleware.PriorityAdmission()) + suno.Use(promptaudit.HTTPEnqueueMiddleware(types.RelayFormatTask)) +diff --git a/ai-gateway/internal/gatewayadapter/transport/task.go b/ai-gateway/internal/gatewayadapter/transport/task.go +index 18f8c1fbdc6b039217fcd5b159a5a7c6e4a49480..9233198c44b17d864fb6243a2ec48a71e24bb2e0 100644 +--- a/ai-gateway/internal/gatewayadapter/transport/task.go ++++ b/ai-gateway/internal/gatewayadapter/transport/task.go +@@ -74,6 +74,7 @@ func NewTaskGatewayHandler(deps TaskGatewayDeps) http.Handler { + v1 := engine.Group("/v1", relayMiddlewares...) + v1.Use(middleware.TokenAuth()) + v1.Use(middleware.UserConcurrencyLimit()) ++ v1.Use(promptaudit.HTTPGuardMiddleware(types.RelayFormatTask)) + v1.Use(middleware.Distribute()) + v1.Use(middleware.PriorityAdmission()) + v1.Use(promptaudit.HTTPEnqueueMiddleware(types.RelayFormatTask)) +diff --git a/ai-gateway/internal/gatewayadapter/transport/user_concurrency_order_test.go b/ai-gateway/internal/gatewayadapter/transport/user_concurrency_order_test.go +index f2afd6df93ddac878b88af129ac6ba2965fe1418..ec93e96b6052b4e078754c59fb1ac2302e66a74c 100644 +--- a/ai-gateway/internal/gatewayadapter/transport/user_concurrency_order_test.go ++++ b/ai-gateway/internal/gatewayadapter/transport/user_concurrency_order_test.go +@@ -1,6 +1,8 @@ + package transport + + import ( ++ "context" ++ "fmt" + "go/ast" + "go/parser" + "go/token" +@@ -18,9 +20,22 @@ import ( + "github.com/gin-gonic/gin" + "github.com/go-redis/redis/v8" + "github.com/mt21625457/aicodex/internal/common" ++ "github.com/mt21625457/aicodex/internal/gatewaycore" ++ "github.com/mt21625457/aicodex/internal/service/promptaudit" + "github.com/mt21625457/aicodex/internal/setting" ++ "github.com/mt21625457/aicodex/internal/types" + ) + ++type transportPromptGuardEvaluator struct { ++ result gatewaycore.PromptGuardResult ++ calls atomic.Int32 ++} ++ ++func (e *transportPromptGuardEvaluator) Evaluate(context.Context, gatewaycore.PromptGuardInput) gatewaycore.PromptGuardResult { ++ e.calls.Add(1) ++ return e.result ++} ++ + func TestTransportExecutionRoutesApplyUserConcurrencyAfterTokenAuthBeforeDistribute(t *testing.T) { + tests := []struct { + fileName string +@@ -58,6 +73,203 @@ func TestTransportExecutionRoutesApplyUserConcurrencyAfterTokenAuthBeforeDistrib + } + } + ++func TestTransportPromptGuardRunsBeforeDistributionAndPriorityAdmission(t *testing.T) { ++ tests := []struct { ++ fileName string ++ funcName string ++ want []string ++ }{ ++ {fileName: "openai.go", funcName: "NewOpenAIGatewayHandler", want: []string{"ModelRequestRateLimit", "HTTPGuardMiddleware", "Distribute", "PriorityAdmission"}}, ++ {fileName: "anthropic.go", funcName: "NewAnthropicGatewayHandler", want: []string{"ModelRequestRateLimit", "HTTPGuardMiddleware", "Distribute", "PriorityAdmission"}}, ++ {fileName: "gemini.go", funcName: "NewGeminiGatewayHandler", want: []string{"ModelRequestRateLimit", "HTTPGuardMiddleware", "Distribute", "PriorityAdmission"}}, ++ {fileName: "suno.go", funcName: "NewSunoGatewayHandler", want: []string{"UserConcurrencyLimit", "HTTPGuardMiddleware", "Distribute", "PriorityAdmission"}}, ++ {fileName: "midjourney.go", funcName: "registerMidjourneyTransportGroup", want: []string{"UserConcurrencyLimit", "HTTPGuardMiddleware", "Distribute", "PriorityAdmission"}}, ++ {fileName: "kling.go", funcName: "NewKlingGatewayHandler", want: []string{"UserConcurrencyLimit", "HTTPGuardMiddleware", "Distribute", "PriorityAdmission"}}, ++ {fileName: "jimeng.go", funcName: "NewJimengGatewayHandler", want: []string{"UserConcurrencyLimit", "HTTPGuardMiddleware", "Distribute", "PriorityAdmission"}}, ++ {fileName: "task.go", funcName: "NewTaskGatewayHandler", want: []string{"UserConcurrencyLimit", "HTTPGuardMiddleware", "Distribute", "PriorityAdmission"}}, ++ } ++ for _, tt := range tests { ++ t.Run(tt.fileName+"/"+tt.funcName, func(t *testing.T) { ++ order := middlewareCallOrderInFunction(t, tt.fileName, tt.funcName) ++ if !containsOrderedMiddlewareSequence(order, tt.want) { ++ t.Fatalf("同步门禁必须位于分流与优先级准入之前,want=%v order=%v", tt.want, order) ++ } ++ }) ++ } ++} ++ ++func TestTransportPromptGuardBlocksSupported9068ProtocolsBeforeRelay(t *testing.T) { ++ gin.SetMode(gin.TestMode) ++ cleanup := setupDirectGatewayAuthDB(t) ++ defer cleanup() ++ seedAdminTokenAndChannel(t, "rawtransportguard", 9201) ++ installTransportPromptGuardConfig(t, true) ++ ++ evaluator := &transportPromptGuardEvaluator{result: gatewaycore.PromptGuardResult{ ++ Decision: gatewaycore.PromptGuardDecisionBlock, ++ Action: gatewaycore.PromptGuardDecisionBlock, ++ ErrorCode: gatewaycore.PromptGuardErrorBlocked, ++ AllowNextStage: false, ++ }} ++ restoreEvaluator := promptaudit.SetPromptGuardEvaluatorForTesting(evaluator) ++ t.Cleanup(restoreEvaluator) ++ ++ tests := []struct { ++ name string ++ newHandler func(*atomic.Bool) http.Handler ++ target string ++ body string ++ marker string ++ }{ ++ { ++ name: "OpenAI Chat", ++ newHandler: func(entered *atomic.Bool) http.Handler { ++ return NewOpenAIGatewayHandler(OpenAIGatewayDeps{RelayHTTP: func(c *gin.Context, _ interfaceRelayFormat) { entered.Store(true) }}) ++ }, ++ target: "/v1/chat/completions", ++ body: `{"model":"gpt-5","messages":[{"role":"user","content":"guard-secret-chat"}]}`, ++ marker: gatewaycore.PromptGuardErrorBlocked, ++ }, ++ { ++ name: "OpenAI Responses alias", ++ newHandler: func(entered *atomic.Bool) http.Handler { ++ return NewOpenAIGatewayHandler(OpenAIGatewayDeps{RelayHTTP: func(c *gin.Context, _ interfaceRelayFormat) { entered.Store(true) }}) ++ }, ++ target: "/responses", ++ body: `{"model":"gpt-5","input":"guard-secret-responses"}`, ++ marker: gatewaycore.PromptGuardErrorBlocked, ++ }, ++ { ++ name: "Claude Messages", ++ newHandler: func(_ *atomic.Bool) http.Handler { ++ return NewAnthropicGatewayHandler(AnthropicGatewayDeps{}) ++ }, ++ target: "/v1/messages", ++ body: `{"model":"claude-sonnet","messages":[{"role":"user","content":"guard-secret-claude"}]}`, ++ marker: `"type":"prompt_guard_blocked"`, ++ }, ++ { ++ name: "Gemini", ++ newHandler: func(entered *atomic.Bool) http.Handler { ++ return NewGeminiGatewayHandler(GeminiGatewayDeps{RelayHTTP: func(c *gin.Context, _ interfaceRelayFormat) { entered.Store(true) }}) ++ }, ++ target: "/v1beta/models/gemini-2.5:streamGenerateContent", ++ body: `{"contents":[{"role":"user","parts":[{"text":"guard-secret-gemini"}]}]}`, ++ marker: `"reason":"prompt_guard_blocked"`, ++ }, ++ { ++ name: "text task", ++ newHandler: func(entered *atomic.Bool) http.Handler { ++ return NewTaskGatewayHandler(TaskGatewayDeps{RelayTask: func(c *gin.Context) { entered.Store(true) }}) ++ }, ++ target: "/v1/videos", ++ body: `{"model":"sora-2","prompt":"guard-secret-task"}`, ++ marker: gatewaycore.PromptGuardErrorBlocked, ++ }, ++ } ++ ++ for _, tt := range tests { ++ t.Run(tt.name, func(t *testing.T) { ++ var entered atomic.Bool ++ req := httptest.NewRequest(http.MethodPost, tt.target, strings.NewReader(tt.body)) ++ req.Header.Set("Authorization", "Bearer sk-rawtransportguard-1") ++ req.Header.Set("Content-Type", gin.MIMEJSON) ++ req = req.WithContext(common.WithRequestEntrypoint(req.Context(), common.RequestEntrypointAI9068)) ++ rec := httptest.NewRecorder() ++ tt.newHandler(&entered).ServeHTTP(rec, req) ++ if rec.Code != http.StatusForbidden || entered.Load() || !strings.Contains(rec.Body.String(), tt.marker) { ++ t.Fatalf("9068 同步门禁未在 relay 前阻止: status=%d entered=%v body=%s", rec.Code, entered.Load(), rec.Body.String()) ++ } ++ for _, forbidden := range []string{"guard-secret-", "127.0.0.1:18080"} { ++ if strings.Contains(rec.Body.String(), forbidden) { ++ t.Fatalf("9068 错误响应泄露敏感信息 %q: %s", forbidden, rec.Body.String()) ++ } ++ } ++ }) ++ } ++ if got := evaluator.calls.Load(); got != int32(len(tests)) { ++ t.Fatalf("每个协议请求必须且仅调用一次同步 Guard: calls=%d want=%d", got, len(tests)) ++ } ++} ++ ++func TestTransportPromptGuardDisabledKeeps9068AsynchronousAudit(t *testing.T) { ++ gin.SetMode(gin.TestMode) ++ cleanup := setupDirectGatewayAuthDB(t) ++ defer cleanup() ++ seedAdminTokenAndChannel(t, "rawtransportobserve", 9202) ++ installTransportPromptGuardConfig(t, false) ++ ++ evaluator := &transportPromptGuardEvaluator{result: gatewaycore.PromptGuardResult{ ++ Decision: gatewaycore.PromptGuardDecisionUnavailable, ++ ErrorCode: gatewaycore.PromptGuardErrorUnavailable, ++ AllowNextStage: false, ++ }} ++ restoreEvaluator := promptaudit.SetPromptGuardEvaluatorForTesting(evaluator) ++ t.Cleanup(restoreEvaluator) ++ repo := promptaudit.NewMemoryRepository(promptaudit.EntRepository{}) ++ restoreDefaults := promptaudit.SetDefaultsForTesting(repo, promptaudit.NewMemoryPayloadStore()) ++ t.Cleanup(restoreDefaults) ++ ++ var entered atomic.Bool ++ handler := NewOpenAIGatewayHandler(OpenAIGatewayDeps{RelayHTTP: func(c *gin.Context, _ interfaceRelayFormat) { ++ entered.Store(true) ++ c.Status(http.StatusNoContent) ++ }}) ++ req := httptest.NewRequest(http.MethodPost, "/v1/chat/completions", strings.NewReader(`{"model":"gpt-5","messages":[{"role":"user","content":"observe-only"}]}`)) ++ req.Header.Set("Authorization", "Bearer sk-rawtransportobserve-1") ++ req.Header.Set("Content-Type", gin.MIMEJSON) ++ req = req.WithContext(common.WithRequestEntrypoint(req.Context(), common.RequestEntrypointAI9068)) ++ rec := httptest.NewRecorder() ++ handler.ServeHTTP(rec, req) ++ ++ if rec.Code != http.StatusNoContent || !entered.Load() || evaluator.calls.Load() != 0 { ++ t.Fatalf("关闭同步阻止时必须继续原异步链路: status=%d entered=%v guard_calls=%d body=%s", rec.Code, entered.Load(), evaluator.calls.Load(), rec.Body.String()) ++ } ++ deadline := time.Now().Add(time.Second) ++ for { ++ active, err := repo.CountActiveJobs(context.Background()) ++ if err != nil { ++ t.Fatalf("读取异步审计队列失败: %v", err) ++ } ++ if active == 1 { ++ break ++ } ++ if time.Now().After(deadline) { ++ t.Fatalf("异步只审计模式未创建任务: active=%d", active) ++ } ++ time.Sleep(5 * time.Millisecond) ++ } ++} ++ ++// interfaceRelayFormat 是测试回调中 types.RelayFormat 的本地别名,避免与标准库类型混淆。 ++type interfaceRelayFormat = types.RelayFormat ++ ++func installTransportPromptGuardConfig(t *testing.T, blocking bool) { ++ t.Helper() ++ common.OptionMapRWMutex.Lock() ++ hadMap := common.OptionMap != nil ++ if common.OptionMap == nil { ++ common.OptionMap = map[string]string{} ++ } ++ previous, hadValue := common.OptionMap[promptaudit.ConfigOptionKey] ++ common.OptionMap[promptaudit.ConfigOptionKey] = fmt.Sprintf(`{"enabled":true,"blocking_enabled":%t,"store_pass_events":false,"strategy":"priority","worker_count":1,"queue_capacity":100,"scanners":["Jailbreak"],"audit_group_mode":"all","endpoints":[{"id":"guard","name":"guard","base_url":"http://127.0.0.1:18080","timeout_ms":1000,"input_limit":1024,"weight":100,"enabled":true}],"config_version":7}`, blocking) ++ common.OptionMapRWMutex.Unlock() ++ promptaudit.ClearConfigCache() ++ t.Cleanup(func() { ++ common.OptionMapRWMutex.Lock() ++ if hadValue { ++ common.OptionMap[promptaudit.ConfigOptionKey] = previous ++ } else { ++ delete(common.OptionMap, promptaudit.ConfigOptionKey) ++ if !hadMap && len(common.OptionMap) == 0 { ++ common.OptionMap = nil ++ } ++ } ++ common.OptionMapRWMutex.Unlock() ++ promptaudit.ClearConfigCache() ++ }) ++} ++ + func TestTaskLikeTransportsRejectExceededUserConcurrencyBeforeRelayHandler(t *testing.T) { + gin.SetMode(gin.TestMode) + cleanup := setupDirectGatewayAuthDB(t) +@@ -222,7 +434,7 @@ func middlewareCallOrderInFunction(t *testing.T, fileName string, funcName strin + return true + } + switch selector.Sel.Name { +- case "TokenAuth", "UserConcurrencyLimit", "Distribute": ++ case "TokenAuth", "UserConcurrencyLimit", "ModelRequestRateLimit", "HTTPGuardMiddleware", "Distribute", "PriorityAdmission": + order = append(order, selector.Sel.Name) + } + return true +@@ -230,6 +442,23 @@ func middlewareCallOrderInFunction(t *testing.T, fileName string, funcName strin + return order + } + ++func containsOrderedMiddlewareSequence(items []string, sequence []string) bool { ++ if len(sequence) == 0 { ++ return true ++ } ++ matched := 0 ++ for _, item := range items { ++ if item != sequence[matched] { ++ continue ++ } ++ matched++ ++ if matched == len(sequence) { ++ return true ++ } ++ } ++ return false ++} ++ + func firstIndex(items []string, target string) int { + for i, item := range items { + if item == target { +diff --git a/ai-gateway/internal/relay/ws_responses.go b/ai-gateway/internal/relay/ws_responses.go +index 885822ed8deb240a0c4c83c9f36f3a1d60455b65..0cca8b2add6028411eb314c98eb1fc370954dc2f 100644 +--- a/ai-gateway/internal/relay/ws_responses.go ++++ b/ai-gateway/internal/relay/ws_responses.go +@@ -232,6 +232,8 @@ func IsWebSocketUpgradeRequest(r *http.Request) bool { + } + + func WsResponsesHelper(c *gin.Context) *types.AICodexError { ++ // 测试和灰度会替换 hook;请求开始时固定函数快照,避免长连接结束阶段与 hook 回收并发读写。 ++ recordChannelAffinity := responsesWSRecordChannelAffinity + if !IsWebSocketUpgradeRequest(c.Request) { + apiErr := types.NewErrorWithStatusCode( + errors.New("WebSocket upgrade required (Upgrade: websocket)"), +@@ -294,6 +296,19 @@ func WsResponsesHelper(c *gin.Context) *types.AICodexError { + logResponsesWSSetupFailed(c, nil, "first_message", firstErr, coderws.StatusPolicyViolation, "invalid first response.create payload", nil) + return firstErr + } ++ firstGuardCheck := promptaudit.EvaluatePromptGuardBody(c, types.RelayFormatOpenAIResponsesWS, "/v1/responses", firstMessage, "first_turn", false) ++ if !firstGuardCheck.Allowed { ++ closeCode := coderws.StatusTryAgainLater ++ if firstGuardCheck.ErrorCode == "prompt_guard_blocked" { ++ closeCode = coderws.StatusCode(4403) ++ } ++ responsesWSCloseClient(clientConn, closeCode, firstGuardCheck.ErrorCode) ++ apiErr := promptGuardAICodexError(firstGuardCheck) ++ logResponsesWSSetupFailed(c, nil, "prompt_guard_first_turn", apiErr, closeCode, firstGuardCheck.ErrorCode, map[string]any{ ++ "model": requestModel, ++ }) ++ return apiErr ++ } + + prepared, prepErr := dataplaneopenai.PrepareWSForwarding(&dataplaneopenai.WSPrepareRequest{ + RequestModel: requestModel, +@@ -355,7 +370,6 @@ func WsResponsesHelper(c *gin.Context) *types.AICodexError { + return apiErr + } + promptaudit.MaybeEnqueueTurnFromGateway(c, types.RelayFormatOpenAIResponsesWS, firstMessage) +- + dialCtx := c.Request.Context() + cancelDial := func() {} + if runtimeSettings.UpstreamDialTimeout > 0 { +@@ -461,6 +475,10 @@ func WsResponsesHelper(c *gin.Context) *types.AICodexError { + if !isResponsesWSResponseCreatePayload(payload) { + return nil + } ++ guardCheck := promptaudit.EvaluatePromptGuardBody(c, types.RelayFormatOpenAIResponsesWS, "/v1/responses", payload, "subsequent_turn", false) ++ if !guardCheck.Allowed { ++ return promptGuardAICodexError(guardCheck) ++ } + if apiErr := relaycommon.ValidateOpenAIPriorityMode( + gjson.GetBytes(payload, "service_tier").String(), + channelOtherSettings.IsOpenAIPriorityAllowed(), +@@ -524,7 +542,7 @@ func WsResponsesHelper(c *gin.Context) *types.AICodexError { + logger.LogWarnEvent(c.Request.Context(), event, fields) + }, + RecordChannelAffinity: func(channelID int) { +- responsesWSRecordChannelAffinity(c, channelID) ++ recordChannelAffinity(c, channelID) + }, + }) + } +@@ -841,6 +859,14 @@ func mapResponsesWSErrorToCloseCode(err error, stage string, upstreamStatusCode + case errors.Is(err, context.DeadlineExceeded): + return coderws.StatusTryAgainLater, "upstream timeout" + case errors.As(err, &apiErr) && apiErr != nil: ++ switch apiErr.GetErrorCode() { ++ case types.ErrorCodePromptGuardBlocked: ++ return coderws.StatusCode(4403), string(types.ErrorCodePromptGuardBlocked) ++ case types.ErrorCodePromptGuardUnavailable: ++ return coderws.StatusTryAgainLater, string(types.ErrorCodePromptGuardUnavailable) ++ case types.ErrorCodePromptGuardInvalidResponse: ++ return coderws.StatusTryAgainLater, string(types.ErrorCodePromptGuardInvalidResponse) ++ } + if apiErr.GetErrorCode() == types.ErrorCodeSubscriptionPolicyMissing { + return coderws.StatusPolicyViolation, string(types.ErrorCodeSubscriptionPolicyMissing) + } +@@ -874,6 +900,22 @@ func mapResponsesWSErrorToCloseCode(err error, stage string, upstreamStatusCode + return coderws.StatusInternalError, "upstream websocket proxy failed" + } + ++func promptGuardAICodexError(check promptaudit.PromptGuardCheck) *types.AICodexError { ++ code := types.ErrorCodePromptGuardUnavailable ++ message := "提示词安全服务暂时不可用,请稍后重试" ++ if check.ErrorCode == "prompt_guard_blocked" { ++ code = types.ErrorCodePromptGuardBlocked ++ message = "请求因提示词安全策略被阻止" ++ } else if check.ErrorCode == "prompt_guard_invalid_response" { ++ code = types.ErrorCodePromptGuardInvalidResponse ++ } ++ status := check.StatusCode ++ if status == 0 { ++ status = http.StatusServiceUnavailable ++ } ++ return types.NewErrorWithStatusCode(errors.New(message), code, status, types.ErrOptionWithSkipRetry()) ++} ++ + func mapGatewayResponsesWSErrorToCloseCode(apiErr *types.AICodexError) (coderws.StatusCode, string, bool) { + if apiErr == nil { + return 0, "", false +diff --git a/ai-gateway/internal/router/relay-router.go b/ai-gateway/internal/router/relay-router.go +index 6c20b82fdacf19dfa08492fac5ebd2c95cc2722d..6ecf9e3a6b931f8d7c3e97541a83d2b2ef54020e 100644 +--- a/ai-gateway/internal/router/relay-router.go ++++ b/ai-gateway/internal/router/relay-router.go +@@ -8,6 +8,7 @@ import ( + "github.com/mt21625457/aicodex/internal/controller" + "github.com/mt21625457/aicodex/internal/middleware" + "github.com/mt21625457/aicodex/internal/relay" ++ "github.com/mt21625457/aicodex/internal/service/promptaudit" + "github.com/mt21625457/aicodex/internal/types" + + "github.com/gin-gonic/gin" +@@ -103,6 +104,7 @@ func SetRelayRouter(router *gin.Engine) { + httpRouter.Use(middleware.HTTPAuditTrackMultiIP()) + httpRouter.Use(middleware.UserConcurrencyLimit()) + httpRouter.Use(middleware.ModelRequestRateLimit()) ++ httpRouter.Use(promptaudit.HTTPGuardMiddleware()) + httpRouter.Use(middleware.Distribute()) + httpRouter.Use(middleware.PriorityAdmission()) + httpRouter.Use(middleware.HTTPAudit()) +@@ -190,6 +192,7 @@ func SetRelayRouter(router *gin.Engine) { + responsesAliasRouter.Use(middleware.HTTPAuditTrackMultiIP()) + responsesAliasRouter.Use(middleware.UserConcurrencyLimit()) + responsesAliasRouter.Use(middleware.ModelRequestRateLimit()) ++ responsesAliasRouter.Use(promptaudit.HTTPGuardMiddleware(types.RelayFormatOpenAIResponses)) + responsesAliasRouter.Use(middleware.Distribute()) + responsesAliasRouter.Use(middleware.PriorityAdmission()) + responsesAliasRouter.Use(middleware.HTTPAudit()) +@@ -208,7 +211,7 @@ func SetRelayRouter(router *gin.Engine) { + + relaySunoRouter := router.Group("/suno", relayMiddlewares...) + relaySunoRouter.Use(middleware.SystemPerformanceCheck()) +- relaySunoRouter.Use(middleware.TokenAuth(), middleware.UserConcurrencyLimit(), middleware.Distribute(), middleware.PriorityAdmission()) ++ relaySunoRouter.Use(middleware.TokenAuth(), middleware.UserConcurrencyLimit(), promptaudit.HTTPGuardMiddleware(types.RelayFormatTask), middleware.Distribute(), middleware.PriorityAdmission(), promptaudit.HTTPEnqueueMiddleware(types.RelayFormatTask)) + { + relaySunoRouter.POST("/submit/:action", controller.RelayTask) + relaySunoRouter.POST("/fetch", controller.RelayTask) +@@ -221,6 +224,7 @@ func SetRelayRouter(router *gin.Engine) { + relayGeminiRouter.Use(middleware.HTTPAuditTrackMultiIP()) + relayGeminiRouter.Use(middleware.UserConcurrencyLimit()) + relayGeminiRouter.Use(middleware.ModelRequestRateLimit()) ++ relayGeminiRouter.Use(promptaudit.HTTPGuardMiddleware(types.RelayFormatGemini)) + relayGeminiRouter.Use(middleware.Distribute()) + relayGeminiRouter.Use(middleware.PriorityAdmission()) + relayGeminiRouter.Use(middleware.HTTPAudit()) +@@ -237,7 +241,7 @@ func registerMjRouterGroup(relayMjRouter *gin.RouterGroup) { + imageRoute.Use(middleware.UserAuth()) + imageRoute.GET("/image/:id", relay.RelayMidjourneyImage) + +- relayMjRouter.Use(middleware.TokenAuth(), middleware.UserConcurrencyLimit(), middleware.Distribute(), middleware.PriorityAdmission()) ++ relayMjRouter.Use(middleware.TokenAuth(), middleware.UserConcurrencyLimit(), promptaudit.HTTPGuardMiddleware(types.RelayFormatMjProxy), middleware.Distribute(), middleware.PriorityAdmission(), promptaudit.HTTPEnqueueMiddleware(types.RelayFormatMjProxy)) + { + relayMjRouter.POST("/submit/action", controller.RelayMidjourney) + relayMjRouter.POST("/submit/shorten", controller.RelayMidjourney) +diff --git a/ai-gateway/internal/router/video-router.go b/ai-gateway/internal/router/video-router.go +index b6a0168773f73df7995d19f4be63df4c3f9bce4b..29d55934963930d42700eaa8dffb462457b49bee 100644 +--- a/ai-gateway/internal/router/video-router.go ++++ b/ai-gateway/internal/router/video-router.go +@@ -3,13 +3,22 @@ package router + import ( + "github.com/mt21625457/aicodex/internal/controller" + "github.com/mt21625457/aicodex/internal/middleware" ++ "github.com/mt21625457/aicodex/internal/service/promptaudit" ++ "github.com/mt21625457/aicodex/internal/types" + + "github.com/gin-gonic/gin" + ) + + func SetVideoRouter(router *gin.Engine) { + videoV1Router := router.Group("/v1") +- videoV1Router.Use(middleware.TokenAuth(), middleware.Distribute(), middleware.PriorityAdmission()) ++ videoV1Router.Use( ++ middleware.TokenAuth(), ++ middleware.UserConcurrencyLimit(), ++ promptaudit.HTTPGuardMiddleware(types.RelayFormatTask), ++ middleware.Distribute(), ++ middleware.PriorityAdmission(), ++ promptaudit.HTTPEnqueueMiddleware(types.RelayFormatTask), ++ ) + { + videoV1Router.GET("/videos/:task_id/content", controller.VideoProxy) + videoV1Router.POST("/video/generations", controller.RelayTask) +@@ -24,7 +33,15 @@ func SetVideoRouter(router *gin.Engine) { + } + + klingV1Router := router.Group("/kling/v1") +- klingV1Router.Use(middleware.KlingRequestConvert(), middleware.TokenAuth(), middleware.Distribute(), middleware.PriorityAdmission()) ++ klingV1Router.Use( ++ middleware.KlingRequestConvert(), ++ middleware.TokenAuth(), ++ middleware.UserConcurrencyLimit(), ++ promptaudit.HTTPGuardMiddleware(types.RelayFormatTask), ++ middleware.Distribute(), ++ middleware.PriorityAdmission(), ++ promptaudit.HTTPEnqueueMiddleware(types.RelayFormatTask), ++ ) + { + klingV1Router.POST("/videos/text2video", controller.RelayTask) + klingV1Router.POST("/videos/image2video", controller.RelayTask) +@@ -34,7 +51,15 @@ func SetVideoRouter(router *gin.Engine) { + + // Jimeng official API routes - direct mapping to official API format + jimengOfficialGroup := router.Group("jimeng") +- jimengOfficialGroup.Use(middleware.JimengRequestConvert(), middleware.TokenAuth(), middleware.Distribute(), middleware.PriorityAdmission()) ++ jimengOfficialGroup.Use( ++ middleware.JimengRequestConvert(), ++ middleware.TokenAuth(), ++ middleware.UserConcurrencyLimit(), ++ promptaudit.HTTPGuardMiddleware(types.RelayFormatTask), ++ middleware.Distribute(), ++ middleware.PriorityAdmission(), ++ promptaudit.HTTPEnqueueMiddleware(types.RelayFormatTask), ++ ) + { + // Maps to: /?Action=CVSync2AsyncSubmitTask&Version=2022-08-31 and /?Action=CVSync2AsyncGetResult&Version=2022-08-31 + jimengOfficialGroup.POST("/", controller.RelayTask) +diff --git a/ai-gateway/internal/service/promptaudit/client.go b/ai-gateway/internal/service/promptaudit/client.go +index 91e18f5fb4ebb24a8a42c72826506f6bf5152493..941c54c35378d50d71bf49b8c62906bab2663cb3 100644 +--- a/ai-gateway/internal/service/promptaudit/client.go ++++ b/ai-gateway/internal/service/promptaudit/client.go +@@ -47,15 +47,8 @@ func newLLMGuardError(code string, message string, retryable bool, statusCode in + } + + func readSmallResponseBody(body io.Reader) string { +- if body == nil { +- return "" +- } +- data, _ := io.ReadAll(io.LimitReader(body, 4096)) +- message := strings.TrimSpace(string(data)) +- if message == "" { +- return "审计 API 返回非成功状态" +- } +- return message ++ _ = body ++ return "Guard API 返回非成功状态" + } + + func firstNonEmptyString(values ...string) string { +diff --git a/ai-gateway/internal/service/promptaudit/config.go b/ai-gateway/internal/service/promptaudit/config.go +index db835a3704607700a65dad0b0b3cd7cbf4b8e6e0..f15a361f91441ed98f3f6edf4f1cffb7f1530ab1 100644 +--- a/ai-gateway/internal/service/promptaudit/config.go ++++ b/ai-gateway/internal/service/promptaudit/config.go +@@ -5,13 +5,13 @@ import ( + "crypto/sha256" + "encoding/hex" + "encoding/json" +- "errors" + "fmt" +- "net/url" + "os" + "slices" + "strconv" + "strings" ++ "sync" ++ "sync/atomic" + "time" + + "github.com/mt21625457/aicodex/internal/common" +@@ -20,6 +20,12 @@ import ( + + const ConfigOptionKey = "PromptAuditConfigJSON" + ++const ( ++ PromptGuardRequiresAuditEnabled = "prompt_guard_requires_audit_enabled" ++ PromptGuardInvalidStrategy = "prompt_guard_invalid_strategy" ++ PromptAuditConfigInvalid = "prompt_audit_config_invalid" ++) ++ + const ( + AuditGroupModeAll = "all" + AuditGroupModeSelected = "selected" +@@ -32,12 +38,10 @@ const ( + ) + + var ( +- allowedStrategies = map[string]struct{}{ +- "priority": {}, +- "weighted": {}, +- "shadow": {}, +- } +- codeContentScanners = map[string]struct{}{ ++ configSaveMu sync.Mutex ++ strategyMigrationLogged atomic.Bool ++ allowedStrategies = map[string]struct{}{"priority": {}} ++ codeContentScanners = map[string]struct{}{ + "bancode": {}, + "code": {}, + } +@@ -117,6 +121,7 @@ type EndpointConfig struct { + + type Config 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"` +@@ -125,6 +130,9 @@ type Config struct { + AuditGroupMode string `json:"audit_group_mode"` + AuditGroups []string `json:"audit_groups"` + Endpoints []EndpointConfig `json:"endpoints"` ++ ConfigVersion int64 `json:"config_version"` ++ UpdatedBy int `json:"updated_by,omitempty"` ++ ChangeSummary string `json:"change_summary,omitempty"` + UpdatedAt time.Time `json:"updated_at"` + } + +@@ -147,6 +155,7 @@ type EndpointInput struct { + + type SaveConfigRequest 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"` +@@ -155,6 +164,8 @@ type SaveConfigRequest struct { + AuditGroupMode string `json:"audit_group_mode"` + AuditGroups []string `json:"audit_groups"` + Endpoints []EndpointInput `json:"endpoints"` ++ UpdatedBy int `json:"-"` ++ ChangeReason string `json:"change_reason"` + } + + type EndpointPublic struct { +@@ -176,6 +187,7 @@ type EndpointPublic struct { + + type PublicConfig 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"` +@@ -184,11 +196,15 @@ type PublicConfig struct { + AuditGroupMode string `json:"audit_group_mode"` + AuditGroups []string `json:"audit_groups"` + Endpoints []EndpointPublic `json:"endpoints"` ++ ConfigVersion int64 `json:"config_version"` ++ UpdatedBy int `json:"updated_by,omitempty"` ++ ChangeSummary string `json:"change_summary,omitempty"` + UpdatedAt time.Time `json:"updated_at"` + } + + type storedConfig struct { + Enabled bool `json:"enabled"` ++ BlockingEnabled bool `json:"blocking_enabled,omitempty"` + StorePassEvents bool `json:"store_pass_events"` + Strategy string `json:"strategy"` + WorkerCount int `json:"worker_count"` +@@ -197,9 +213,25 @@ type storedConfig struct { + AuditGroupMode string `json:"audit_group_mode,omitempty"` + AuditGroups []string `json:"audit_groups,omitempty"` + Endpoints []storedEndpoint `json:"endpoints"` ++ ConfigVersion int64 `json:"config_version,omitempty"` ++ UpdatedBy int `json:"updated_by,omitempty"` ++ ChangeSummary string `json:"change_summary,omitempty"` + UpdatedAt time.Time `json:"updated_at"` + } + ++// ConfigValidationError 为控制面提供稳定、可测试的配置错误码。 ++type ConfigValidationError struct { ++ Code string ++ Message string ++} ++ ++func (e *ConfigValidationError) Error() string { ++ if e == nil { ++ return "" ++ } ++ return e.Message ++} ++ + type storedEndpoint struct { + ID string `json:"id"` + Name string `json:"name"` +@@ -218,14 +250,16 @@ type storedEndpoint struct { + + func DefaultConfig() Config { + cfg := Config{ +- Enabled: false, +- Strategy: "priority", +- WorkerCount: 4, +- QueueCapacity: 10000, +- Scanners: defaultOpenAICompatibleScanners(), +- AuditGroupMode: AuditGroupModeAll, +- AuditGroups: []string{}, +- Endpoints: []EndpointConfig{}, ++ Enabled: false, ++ BlockingEnabled: false, ++ Strategy: "priority", ++ WorkerCount: 4, ++ QueueCapacity: 10000, ++ Scanners: defaultOpenAICompatibleScanners(), ++ AuditGroupMode: AuditGroupModeAll, ++ AuditGroups: []string{}, ++ Endpoints: []EndpointConfig{}, ++ ConfigVersion: 1, + } + applyEnvDefaults(&cfg) + return cfg +@@ -255,6 +289,8 @@ func (s *ConfigService) Public(ctx context.Context) (PublicConfig, error) { + } + + func (s *ConfigService) Save(ctx context.Context, req SaveConfigRequest) (PublicConfig, error) { ++ configSaveMu.Lock() ++ defer configSaveMu.Unlock() + current, err := s.Load(ctx) + if err != nil { + return PublicConfig{}, err +@@ -274,9 +310,22 @@ func (s *ConfigService) Save(ctx context.Context, req SaveConfigRequest) (Public + if err := s.store.SetOption(ctx, ConfigOptionKey, string(body)); err != nil { + return PublicConfig{}, err + } ++ if configStorePublishesRuntime(s.store) { ++ installConfigSnapshot(cfg) ++ publishConfigInvalidation(ctx, cfg.ConfigVersion) ++ } + return cfg.Public(), nil + } + ++func configStorePublishesRuntime(store OptionStore) bool { ++ switch store.(type) { ++ case ModelOptionStore, *ModelOptionStore: ++ return true ++ default: ++ return false ++ } ++} ++ + func (c Config) Public() PublicConfig { + endpoints := make([]EndpointPublic, 0, len(c.Endpoints)) + for _, endpoint := range c.Endpoints { +@@ -304,6 +353,7 @@ func (c Config) Public() PublicConfig { + } + return PublicConfig{ + Enabled: c.Enabled, ++ BlockingEnabled: c.Enabled && c.BlockingEnabled, + StorePassEvents: c.StorePassEvents, + Strategy: c.Strategy, + WorkerCount: c.WorkerCount, +@@ -312,19 +362,35 @@ func (c Config) Public() PublicConfig { + AuditGroupMode: normalizeAuditGroupMode(c.AuditGroupMode), + AuditGroups: normalizeAuditGroups(c.AuditGroups), + Endpoints: endpoints, ++ ConfigVersion: normalizeConfigVersion(c.ConfigVersion), ++ UpdatedBy: c.UpdatedBy, ++ ChangeSummary: c.ChangeSummary, + UpdatedAt: c.UpdatedAt, + } + } + + func configFromStorage(stored storedConfig) (Config, error) { ++ storedStrategy := strings.ToLower(strings.TrimSpace(stored.Strategy)) ++ if storedStrategy != "" && storedStrategy != "priority" && strategyMigrationLogged.CompareAndSwap(false, true) { ++ LogWarnEvent( ++ "prompt_guard.config_loaded", ++ Field("status", "migrated"), ++ Field("error_code", "historical_strategy_migrated"), ++ Field("strategy", "priority"), ++ ) ++ } + cfg := Config{ + Enabled: stored.Enabled, ++ BlockingEnabled: stored.Enabled && stored.BlockingEnabled, + StorePassEvents: stored.StorePassEvents, + Strategy: normalizeStrategy(stored.Strategy), + WorkerCount: normalizeWorkerCount(stored.WorkerCount), + QueueCapacity: normalizeQueueCapacity(stored.QueueCapacity), + AuditGroupMode: normalizeAuditGroupMode(stored.AuditGroupMode), + AuditGroups: normalizeAuditGroups(stored.AuditGroups), ++ ConfigVersion: normalizeConfigVersion(stored.ConfigVersion), ++ UpdatedBy: stored.UpdatedBy, ++ ChangeSummary: strings.TrimSpace(stored.ChangeSummary), + UpdatedAt: stored.UpdatedAt, + } + for _, endpoint := range stored.Endpoints { +@@ -365,6 +431,7 @@ func configFromStorage(stored storedConfig) (Config, error) { + func configToStorage(cfg Config) (storedConfig, error) { + stored := storedConfig{ + Enabled: cfg.Enabled, ++ BlockingEnabled: cfg.Enabled && cfg.BlockingEnabled, + StorePassEvents: cfg.StorePassEvents, + Strategy: normalizeStrategy(cfg.Strategy), + WorkerCount: normalizeWorkerCount(cfg.WorkerCount), +@@ -373,6 +440,9 @@ func configToStorage(cfg Config) (storedConfig, error) { + AuditGroupMode: normalizeAuditGroupMode(cfg.AuditGroupMode), + AuditGroups: normalizeAuditGroups(cfg.AuditGroups), + Endpoints: make([]storedEndpoint, 0, len(cfg.Endpoints)), ++ ConfigVersion: normalizeConfigVersion(cfg.ConfigVersion), ++ UpdatedBy: cfg.UpdatedBy, ++ ChangeSummary: strings.TrimSpace(cfg.ChangeSummary), + UpdatedAt: cfg.UpdatedAt, + } + for _, endpoint := range cfg.Endpoints { +@@ -400,14 +470,34 @@ func configToStorage(cfg Config) (storedConfig, error) { + } + + func normalizeSaveRequest(req SaveConfigRequest, current Config) (Config, error) { ++ if !req.Enabled && req.BlockingEnabled { ++ return Config{}, &ConfigValidationError{ ++ Code: PromptGuardRequiresAuditEnabled, ++ Message: "启用同步阻止前必须先启用提示词审计", ++ } ++ } ++ strategy := strings.ToLower(strings.TrimSpace(req.Strategy)) ++ if strategy == "" { ++ strategy = "priority" ++ } ++ if _, ok := allowedStrategies[strategy]; !ok { ++ return Config{}, &ConfigValidationError{ ++ Code: PromptGuardInvalidStrategy, ++ Message: "提示词审计调度策略仅支持 priority", ++ } ++ } + cfg := Config{ + Enabled: req.Enabled, ++ BlockingEnabled: req.Enabled && req.BlockingEnabled, + StorePassEvents: req.StorePassEvents, +- Strategy: normalizeStrategy(req.Strategy), ++ Strategy: strategy, + WorkerCount: normalizeWorkerCount(req.WorkerCount), + QueueCapacity: normalizeQueueCapacity(req.QueueCapacity), + AuditGroupMode: normalizeAuditGroupMode(req.AuditGroupMode), + AuditGroups: normalizeAuditGroups(req.AuditGroups), ++ ConfigVersion: normalizeConfigVersion(current.ConfigVersion) + 1, ++ UpdatedBy: req.UpdatedBy, ++ ChangeSummary: buildConfigChangeSummary(req, current), + UpdatedAt: time.Now().UTC(), + } + if strings.TrimSpace(req.AuditGroupMode) == "" { +@@ -426,7 +516,10 @@ func normalizeSaveRequest(req SaveConfigRequest, current Config) (Config, error) + cfg.Endpoints = append(cfg.Endpoints, endpoint) + } + if cfg.Enabled && !hasEnabledEndpointWithScanURL(cfg.Endpoints) { +- return Config{}, errors.New("启用提示词审计前至少需要一个启用且配置了 Base URL 的 OpenAI 兼容审计池") ++ return Config{}, &ConfigValidationError{ ++ Code: PromptAuditConfigInvalid, ++ Message: "启用提示词审计前至少需要一个启用且配置了 Base URL 的 OpenAI 兼容审计池", ++ } + } + cfg.Scanners = normalizeScanners(req.Scanners) + return cfg, nil +@@ -454,11 +547,16 @@ func normalizeEndpointInput(input EndpointInput, index int, currentByID map[stri + } + scanURL = normalizeOpenAIChatCompletionsURL(firstNonEmptyString(baseURL, scanURL)) + if baseURL == "" { +- return EndpointConfig{}, errors.New("OpenAI 兼容审计池 Base URL 不能为空") ++ return EndpointConfig{}, &ConfigValidationError{ ++ Code: PromptAuditConfigInvalid, ++ Message: "OpenAI 兼容审计池 Base URL 不能为空", ++ } + } +- parsed, err := url.Parse(baseURL) +- if err != nil || parsed.Scheme == "" || parsed.Host == "" || (parsed.Scheme != "http" && parsed.Scheme != "https") { +- return EndpointConfig{}, fmt.Errorf("OpenAI 兼容审计 Base URL 无效: %s", baseURL) ++ if err := validateGuardBaseURL(baseURL); err != nil { ++ return EndpointConfig{}, &ConfigValidationError{ ++ Code: PromptAuditConfigInvalid, ++ Message: err.Error(), ++ } + } + if model == "" && hasExisting { + model = existing.Model +@@ -512,6 +610,37 @@ func normalizeStrategy(value string) string { + return "priority" + } + ++func normalizeConfigVersion(value int64) int64 { ++ if value < 1 { ++ return 1 ++ } ++ return value ++} ++ ++func buildConfigChangeSummary(req SaveConfigRequest, current Config) string { ++ reason := strings.TrimSpace(req.ChangeReason) ++ if len([]rune(reason)) > 120 { ++ reason = string([]rune(reason)[:120]) ++ } ++ mode := "异步只审计" ++ if req.Enabled && req.BlockingEnabled { ++ mode = "同步阻止" ++ } else if !req.Enabled { ++ mode = "关闭审计" ++ } ++ previousMode := "异步只审计" ++ if current.Enabled && current.BlockingEnabled { ++ previousMode = "同步阻止" ++ } else if !current.Enabled { ++ previousMode = "关闭审计" ++ } ++ summary := fmt.Sprintf("模式:%s→%s;审计池:%d;类别:%d", previousMode, mode, len(req.Endpoints), len(req.Scanners)) ++ if reason != "" { ++ summary += ";原因:" + reason ++ } ++ return summary ++} ++ + func normalizeWorkerCount(value int) int { + if value < 1 { + return 4 +@@ -709,7 +838,7 @@ func canonicalScannerNameForProtocol(_ string, value string) (string, bool) { + } + + func normalizeScannerKey(value string) string { +- return strings.NewReplacer("_", "", "-", "", " ", "").Replace(strings.ToLower(strings.TrimSpace(value))) ++ return strings.NewReplacer("_", "", "-", "", " ", "", "&", "and").Replace(strings.ToLower(strings.TrimSpace(value))) + } + func hasEnabledEndpointWithScanURL(endpoints []EndpointConfig) bool { + for _, endpoint := range endpoints { +@@ -737,6 +866,9 @@ func applyEnvDefaults(cfg *Config) { + if raw := strings.TrimSpace(os.Getenv("PROMPT_AUDIT_ENABLED")); raw != "" { + cfg.Enabled = parseEnvBool(raw, cfg.Enabled) + } ++ if raw := strings.TrimSpace(os.Getenv("PROMPT_AUDIT_BLOCKING_ENABLED")); raw != "" { ++ cfg.BlockingEnabled = cfg.Enabled && parseEnvBool(raw, cfg.BlockingEnabled) ++ } + if raw := strings.TrimSpace(os.Getenv("PROMPT_AUDIT_STORE_PASS_EVENTS")); raw != "" { + cfg.StorePassEvents = parseEnvBool(raw, cfg.StorePassEvents) + } +diff --git a/ai-gateway/internal/service/promptaudit/config_test.go b/ai-gateway/internal/service/promptaudit/config_test.go +index b9f766d6b289c7f1eb2ab3a4f7dfb2e3fb32e0c7..b20070937a3c6c18a55ea1ea979902b6c17c0a13 100644 +--- a/ai-gateway/internal/service/promptaudit/config_test.go ++++ b/ai-gateway/internal/service/promptaudit/config_test.go +@@ -66,6 +66,7 @@ func clearPromptAuditEnv(t *testing.T) { + t.Helper() + for _, key := range []string{ + "PROMPT_AUDIT_ENABLED", ++ "PROMPT_AUDIT_BLOCKING_ENABLED", + "PROMPT_AUDIT_STORE_PASS_EVENTS", + "PROMPT_AUDIT_STRATEGY", + "PROMPT_AUDIT_WORKER_COUNT", +@@ -153,8 +154,8 @@ func TestConfigServiceLoadsHotOptionMapWithoutDatabaseAndMasksToken(t *testing.T + if err != nil { + t.Fatalf("Load() should use common.OptionMap without touching DB: %v", err) + } +- if !cfg.Enabled || cfg.Strategy != "weighted" || cfg.AuditGroupMode != AuditGroupModeSelected { +- t.Fatalf("hot config not preserved: %+v", cfg) ++ if !cfg.Enabled || cfg.Strategy != "priority" || cfg.AuditGroupMode != AuditGroupModeSelected { ++ t.Fatalf("历史调度策略应迁移为 priority: %+v", cfg) + } + if len(cfg.Endpoints) != 1 { + t.Fatalf("expected one hot endpoint, got %#v", cfg.Endpoints) +@@ -272,6 +273,7 @@ func TestConfigServiceSaveRejectsInvalidEndpointAndCanClearExistingToken(t *test + func TestDefaultConfigAppliesPromptAuditEnvironmentDefaults(t *testing.T) { + clearPromptAuditEnv(t) + t.Setenv("PROMPT_AUDIT_ENABLED", "true") ++ t.Setenv("PROMPT_AUDIT_BLOCKING_ENABLED", "true") + t.Setenv("PROMPT_AUDIT_STORE_PASS_EVENTS", "true") + t.Setenv("PROMPT_AUDIT_STRATEGY", "weighted") + t.Setenv("PROMPT_AUDIT_WORKER_COUNT", "2") +@@ -283,7 +285,7 @@ func TestDefaultConfigAppliesPromptAuditEnvironmentDefaults(t *testing.T) { + + cfg := DefaultConfig() + +- if !cfg.Enabled || !cfg.StorePassEvents || cfg.Strategy != "weighted" { ++ if !cfg.Enabled || !cfg.BlockingEnabled || !cfg.StorePassEvents || cfg.Strategy != "priority" { + t.Fatalf("env flags not applied: %+v", cfg) + } + if cfg.WorkerCount != 2 || cfg.QueueCapacity != 250 { +@@ -508,7 +510,7 @@ func TestConfigServiceSaveEncryptsTokenAndPublicResponseMasksToken(t *testing.T) + + publicCfg, err := svc.Save(context.Background(), SaveConfigRequest{ + Enabled: true, +- Strategy: "weighted", ++ Strategy: "priority", + WorkerCount: 96, + QueueCapacity: 100000, + Scanners: []string{"Jailbreak", "Jailbreak", "PII"}, +@@ -662,7 +664,7 @@ func TestDefaultConfigAppliesPromptAuditEnvOverrides(t *testing.T) { + + cfg := DefaultConfig() + +- if !cfg.Enabled || !cfg.StorePassEvents || cfg.Strategy != "shadow" { ++ if !cfg.Enabled || !cfg.StorePassEvents || cfg.Strategy != "priority" { + t.Fatalf("boolean/strategy env not applied: %#v", cfg) + } + if cfg.WorkerCount != 9 || cfg.QueueCapacity != 12345 { +diff --git a/ai-gateway/internal/service/promptaudit/diagnostics_test.go b/ai-gateway/internal/service/promptaudit/diagnostics_test.go +index bccee339b8f580fa3a15ba59d8dde8062f080d10..078c01f03d2a71e8ebb861a65cab666211eb1d62 100644 +--- a/ai-gateway/internal/service/promptaudit/diagnostics_test.go ++++ b/ai-gateway/internal/service/promptaudit/diagnostics_test.go +@@ -146,7 +146,7 @@ func TestPromptAuditSmallHelpersExposeOperatorSafeDefaults(t *testing.T) { + if promptAuditProbeErrorMessage(nil) != "" { + t.Fatal("nil probe error message should stay empty") + } +- if got := promptAuditProbeErrorMessage(newLLMGuardError("openai_guard_ready_failed", "", true, 503)); !strings.Contains(got, "openai_guard_ready_failed") || !strings.Contains(got, "status=503") { +- t.Fatalf("probe error message should expose stable code and status when message is empty, got %q", got) ++ if got := promptAuditProbeErrorMessage(newLLMGuardError("openai_guard_ready_failed", "", true, 503)); got != "Guard 探测失败" { ++ t.Fatalf("probe error message should remain generic and not expose upstream detail, got %q", got) + } + } +diff --git a/ai-gateway/internal/service/promptaudit/enqueue.go b/ai-gateway/internal/service/promptaudit/enqueue.go +index 6ac5329666fa875ac1984d7402f28afd4a72de3c..df92920a9763c3afff208d6efedc16a872764d0c 100644 +--- a/ai-gateway/internal/service/promptaudit/enqueue.go ++++ b/ai-gateway/internal/service/promptaudit/enqueue.go +@@ -15,7 +15,10 @@ import ( + "github.com/mt21625457/aicodex/internal/types" + ) + +-const promptAuditEnqueueAttemptedKey = "prompt_audit_enqueue_attempted" ++const ( ++ promptAuditEnqueueAttemptedKey = "prompt_audit_enqueue_attempted" ++ promptAuditConfigCacheTTL = 5 * time.Second ++) + + var ( + defaultMemoryPayloadStore = NewMemoryPayloadStore() +@@ -24,8 +27,13 @@ var ( + configCacheMu sync.Mutex + configCacheValue Config + configCacheLoadedAt time.Time ++ configCacheRefreshAfter time.Time ++ configCacheLastError string ++ configCacheLastErrorAt time.Time + ) + ++const promptGuardConfigInvalidationChannel = "aicodex:prompt_guard:config:invalidate" ++ + func SetDefaultsForTesting(repo JobRepository, payloadStore PayloadStore) func() { + previousRepo := defaultRepository + previousPayload := defaultPayloadStore +@@ -48,10 +56,130 @@ func ConfigureDefaultPayloadStore(payloadStore PayloadStore) { + func ClearConfigCache() { + configCacheMu.Lock() + configCacheLoadedAt = time.Time{} ++ configCacheRefreshAfter = time.Time{} + configCacheValue = Config{} ++ configCacheLastError = "" ++ configCacheLastErrorAt = time.Time{} + configCacheMu.Unlock() + } + ++func invalidateConfigCache() { ++ configCacheMu.Lock() ++ configCacheRefreshAfter = time.Time{} ++ configCacheMu.Unlock() ++} ++ ++func installConfigSnapshot(cfg Config) Config { ++ now := time.Now() ++ configCacheMu.Lock() ++ if configCacheValue.ConfigVersion > 0 && normalizeConfigVersion(cfg.ConfigVersion) < normalizeConfigVersion(configCacheValue.ConfigVersion) { ++ current := configCacheValue ++ configCacheMu.Unlock() ++ return current ++ } ++ configCacheValue = cfg ++ configCacheLoadedAt = now ++ configCacheRefreshAfter = now.Add(promptAuditConfigCacheTTL) ++ configCacheLastError = "" ++ configCacheLastErrorAt = time.Time{} ++ configCacheMu.Unlock() ++ LogInfoEvent( ++ "prompt_guard.config_loaded", ++ Field("status", "success"), ++ Field("config_version", normalizeConfigVersion(cfg.ConfigVersion)), ++ Field("blocking_enabled", cfg.Enabled && cfg.BlockingEnabled), ++ ) ++ return cfg ++} ++ ++func publishConfigInvalidation(ctx context.Context, configVersion int64) { ++ if !common.RedisEnabled || common.RDB == nil { ++ LogWarnEvent( ++ "prompt_guard.config_reload_degraded", ++ Field("status", "degraded"), ++ Field("error_code", "redis_unavailable"), ++ Field("config_version", normalizeConfigVersion(configVersion)), ++ ) ++ return ++ } ++ if err := common.RDB.Publish(ctx, promptGuardConfigInvalidationChannel, fmt.Sprintf("%d", normalizeConfigVersion(configVersion))).Err(); err != nil { ++ LogWarnEvent( ++ "prompt_guard.config_reload_degraded", ++ Field("status", "degraded"), ++ Field("error_code", "config_invalidation_publish_failed"), ++ Field("config_version", normalizeConfigVersion(configVersion)), ++ Field("error_kind", "redis_publish_failed"), ++ ) ++ } ++} ++ ++// StartConfigInvalidationSubscriber 监听多实例配置失效通知;Redis 不可用时继续使用 5 秒 TTL。 ++func StartConfigInvalidationSubscriber(ctx context.Context) { ++ if !common.RedisEnabled || common.RDB == nil { ++ return ++ } ++ go func() { ++ pubsub := common.RDB.Subscribe(ctx, promptGuardConfigInvalidationChannel) ++ defer pubsub.Close() ++ if _, err := pubsub.Receive(ctx); err != nil { ++ LogWarnEvent( ++ "prompt_guard.config_reload_degraded", ++ Field("status", "degraded"), ++ Field("error_code", "config_invalidation_subscribe_failed"), ++ Field("error_kind", "redis_subscribe_failed"), ++ ) ++ return ++ } ++ channel := pubsub.Channel() ++ for { ++ select { ++ case <-ctx.Done(): ++ return ++ case _, ok := <-channel: ++ if !ok { ++ LogWarnEvent( ++ "prompt_guard.config_reload_degraded", ++ Field("status", "degraded"), ++ Field("error_code", "config_invalidation_channel_closed"), ++ Field("error_kind", "redis_subscription_closed"), ++ ) ++ return ++ } ++ invalidateConfigCache() ++ if cfg, err := NewConfigService(nil).Load(ctx); err == nil { ++ installConfigSnapshot(cfg) ++ } else { ++ recordConfigLoadError(err) ++ } ++ } ++ } ++ }() ++} ++ ++func recordConfigLoadError(err error) { ++ if err == nil { ++ return ++ } ++ configCacheMu.Lock() ++ configCacheLastError = "提示词审计配置加载失败" ++ configCacheLastErrorAt = time.Now().UTC() ++ lastVersion := configCacheValue.ConfigVersion ++ configCacheMu.Unlock() ++ LogWarnEvent( ++ "prompt_guard.config_reload_degraded", ++ Field("status", "degraded"), ++ Field("error_code", "config_load_failed"), ++ Field("config_version", normalizeConfigVersion(lastVersion)), ++ Field("error_kind", "config_load_failed"), ++ ) ++} ++ ++func configLoadRuntimeState() (Config, time.Time, string, time.Time) { ++ configCacheMu.Lock() ++ defer configCacheMu.Unlock() ++ return configCacheValue, configCacheLoadedAt, configCacheLastError, configCacheLastErrorAt ++} ++ + func MaybeEnqueueFromGateway(c *gin.Context, relayFormat types.RelayFormat, requestBody []byte) { + maybeEnqueueFromGateway(c, relayFormat, requestBody, false) + } +@@ -121,13 +249,16 @@ func EnqueueFromBody(ctx context.Context, relayFormat types.RelayFormat, path st + Field("reason", "config_load_failed"), + Field("request_id", snapshotContext.RequestID), + Field("error_code", "config_load_failed"), +- Field("error_kind", err.Error()), ++ Field("error_kind", "config_load_failed"), + ) + return false, err + } + if !cfg.Enabled { + return false, nil + } ++ if cfg.BlockingEnabled { ++ return false, nil ++ } + if ok, reason := cfg.ShouldAuditGroup(snapshotContext.Group); !ok { + LogWarnEvent( + "prompt_audit.enqueue_dropped", +@@ -380,18 +511,31 @@ func hasTopLevelJSONKey(body []byte, key string) bool { + } + + func loadCachedConfig(ctx context.Context) (Config, error) { ++ now := time.Now() + configCacheMu.Lock() +- defer configCacheMu.Unlock() +- if !configCacheLoadedAt.IsZero() && time.Since(configCacheLoadedAt) < 5*time.Second { +- return configCacheValue, nil ++ if configCacheValue.ConfigVersion > 0 && now.Before(configCacheRefreshAfter) { ++ cfg := configCacheValue ++ configCacheMu.Unlock() ++ return cfg, nil + } ++ lastValid := configCacheValue ++ hasLastValid := !configCacheLoadedAt.IsZero() || lastValid.ConfigVersion > 0 ++ configCacheMu.Unlock() + cfg, err := NewConfigService(nil).Load(ctx) + if err != nil { ++ recordConfigLoadError(err) ++ if hasLastValid { ++ configCacheMu.Lock() ++ if configCacheValue.ConfigVersion == lastValid.ConfigVersion { ++ configCacheRefreshAfter = now.Add(promptAuditConfigCacheTTL) ++ lastValid = configCacheValue ++ } ++ configCacheMu.Unlock() ++ return lastValid, nil ++ } + return Config{}, err + } +- configCacheValue = cfg +- configCacheLoadedAt = time.Now() +- return cfg, nil ++ return installConfigSnapshot(cfg), nil + } + + func enqueueError(code string, format string, args ...any) error { +diff --git a/ai-gateway/internal/service/promptaudit/openai_client.go b/ai-gateway/internal/service/promptaudit/openai_client.go +index c429064815c47c69a88b769828ef1046343bddfc..aa3f94f029dd54b12e33d608c6234136d993dcdf 100644 +--- a/ai-gateway/internal/service/promptaudit/openai_client.go ++++ b/ai-gateway/internal/service/promptaudit/openai_client.go +@@ -19,7 +19,7 @@ type OpenAICompatibleClient struct { + + func NewOpenAICompatibleClient(httpClient *http.Client) *OpenAICompatibleClient { + if httpClient == nil { +- httpClient = &http.Client{} ++ httpClient = newSecureGuardHTTPClient() + } + return &OpenAICompatibleClient{httpClient: httpClient} + } +@@ -90,6 +90,9 @@ func (c *OpenAICompatibleClient) ScanPrompt(ctx context.Context, endpoint Endpoi + Field("latency_ms", result.LatencyMS), + ) + LogInfoEvent("prompt_audit.scan_chunk_completed", fields...) ++ if scanCtx.StopOnBlock && result.Action == "Block" { ++ break ++ } + } + prependAggregatedGuardPolicy(&aggregated, len(chunks), inputChars, inputLimit) + LogInfoEvent( +@@ -118,6 +121,9 @@ func (c *OpenAICompatibleClient) scanPromptChunk(ctx context.Context, endpoint E + if chatURL == "" { + return LLMGuardScanResult{}, newLLMGuardError("openai_guard_not_configured", "OpenAI 兼容审计 Base URL 为空", false, 0) + } ++ if err := validateGuardBaseURL(firstNonEmptyString(endpoint.BaseURL, endpoint.ScanURL)); err != nil { ++ return LLMGuardScanResult{}, newLLMGuardError("openai_guard_endpoint_denied", "Guard Base URL 未通过安全校验", false, 0) ++ } + model := normalizeGuardModel(ProtocolOpenAICompatible, endpoint.Model) + timeout := time.Duration(endpoint.TimeoutMS) * time.Millisecond + if timeout <= 0 { +@@ -141,7 +147,7 @@ func (c *OpenAICompatibleClient) scanPromptChunk(ctx context.Context, endpoint E + } + req, err := http.NewRequestWithContext(callCtx, http.MethodPost, chatURL, bytes.NewReader(body)) + if err != nil { +- return LLMGuardScanResult{}, err ++ return LLMGuardScanResult{}, newLLMGuardError("openai_guard_request_invalid", "创建 Guard 请求失败", false, 0) + } + req.Header.Set("Content-Type", "application/json") + if token := strings.TrimSpace(endpoint.Token); token != "" { +@@ -155,7 +161,7 @@ func (c *OpenAICompatibleClient) scanPromptChunk(ctx context.Context, endpoint E + if errors.Is(callCtx.Err(), context.DeadlineExceeded) || strings.Contains(strings.ToLower(err.Error()), "timeout") { + return LLMGuardScanResult{}, newLLMGuardError("openai_guard_timeout", "OpenAI 兼容审计调用超时", true, 0) + } +- return LLMGuardScanResult{}, newLLMGuardError("openai_guard_request_failed", err.Error(), true, 0) ++ return LLMGuardScanResult{}, newLLMGuardError("openai_guard_request_failed", "OpenAI 兼容审计请求失败", true, 0) + } + defer resp.Body.Close() + if resp.StatusCode < 200 || resp.StatusCode >= 300 { +@@ -169,8 +175,15 @@ func (c *OpenAICompatibleClient) scanPromptChunk(ctx context.Context, endpoint E + return LLMGuardScanResult{}, newLLMGuardError(code, message, retryable, resp.StatusCode) + } + ++ responseBody, err := io.ReadAll(io.LimitReader(resp.Body, maxGuardResponseBytes+1)) ++ if err != nil { ++ return LLMGuardScanResult{}, newLLMGuardError("openai_guard_invalid_response", "读取 Guard 响应失败", false, resp.StatusCode) ++ } ++ if int64(len(responseBody)) > maxGuardResponseBytes { ++ return LLMGuardScanResult{}, newLLMGuardError("openai_guard_invalid_response", "Guard 响应超过大小上限", false, resp.StatusCode) ++ } + var payload map[string]any +- if err := json.NewDecoder(io.LimitReader(resp.Body, 2*1024*1024)).Decode(&payload); err != nil { ++ if err := json.Unmarshal(responseBody, &payload); err != nil { + return LLMGuardScanResult{}, newLLMGuardError("openai_guard_invalid_response", err.Error(), false, resp.StatusCode) + } + content := extractOpenAIChatContent(payload) +@@ -183,6 +196,9 @@ func (c *OpenAICompatibleClient) scanPromptChunk(ctx context.Context, endpoint E + if len(scanners) == 0 { + enabledCategories = parsed.Categories + } ++ if parsed.Safety == SafetyUnsafe && parsed.HasUnknownCategory { ++ enabledCategories = append(enabledCategories, "unknown_unsafe") ++ } + result := buildQwen3GuardScanResult(parsed, enabledCategories, model, latencyMS) + return result, nil + } +@@ -324,6 +340,9 @@ func (c *OpenAICompatibleClient) CheckReady(ctx context.Context, endpoint Endpoi + if base == "" { + return newLLMGuardError("openai_guard_not_configured", "OpenAI 兼容审计 Base URL 为空", false, 0) + } ++ if err := validateGuardBaseURL(base); err != nil { ++ return newLLMGuardError("openai_guard_endpoint_denied", "Guard Base URL 未通过安全校验", false, 0) ++ } + if err := c.checkModels(ctx, endpoint, base); err == nil { + return nil + } else if !shouldFallbackOpenAIReadyCheck(err) { +@@ -342,11 +361,11 @@ func (c *OpenAICompatibleClient) checkModels(ctx context.Context, endpoint Endpo + defer cancel() + modelsURL, err := url.JoinPath(base, "/v1/models") + if err != nil { +- return err ++ return newLLMGuardError("openai_guard_request_invalid", "创建 Guard 探测地址失败", false, 0) + } + req, err := http.NewRequestWithContext(callCtx, http.MethodGet, modelsURL, nil) + if err != nil { +- return err ++ return newLLMGuardError("openai_guard_request_invalid", "创建 Guard 探测请求失败", false, 0) + } + if token := strings.TrimSpace(endpoint.Token); token != "" { + req.Header.Set("Authorization", "Bearer "+token) +@@ -356,7 +375,7 @@ func (c *OpenAICompatibleClient) checkModels(ctx context.Context, endpoint Endpo + if errors.Is(callCtx.Err(), context.DeadlineExceeded) || strings.Contains(strings.ToLower(err.Error()), "timeout") { + return newLLMGuardError("openai_guard_timeout", "OpenAI 兼容 /v1/models 超时", true, 0) + } +- return newLLMGuardError("openai_guard_request_failed", err.Error(), true, 0) ++ return newLLMGuardError("openai_guard_request_failed", "OpenAI 兼容 Guard 探测请求失败", true, 0) + } + defer resp.Body.Close() + if resp.StatusCode < 200 || resp.StatusCode >= 300 { +@@ -384,7 +403,11 @@ func buildQwen3GuardScanResult(parsed qwen3GuardParsed, categories []string, mod + isValid := true + switch safety { + case SafetyUnsafe: +- action = "Block" ++ if len(categories) > 0 || len(parsed.Categories) == 0 || parsed.HasUnknownCategory { ++ action = "Block" ++ } else { ++ action = "Warn" ++ } + isValid = false + case SafetyControversial: + action = "Warn" +@@ -415,14 +438,26 @@ func buildQwen3GuardScanResult(parsed qwen3GuardParsed, categories []string, mod + if len(findings) > 0 { + evidence["_guard_findings"] = findings + } ++ observedFindings := make([]ScannerEvidenceItem, 0, len(parsed.Categories)) ++ for _, category := range parsed.Categories { ++ observedFindings = append(observedFindings, ScannerEvidenceItem{ ++ ScannerID: category, ++ Category: category, ++ Kind: "classification", ++ Severity: safety, ++ }) ++ } ++ if len(observedFindings) > 0 { ++ evidence["_guard_observed_categories"] = observedFindings ++ } + evidence["_guard_policy"] = []ScannerEvidenceItem{{ + Kind: "policy", + Summary: fmt.Sprintf("action=%s safety=%s model=%s", action, safety, model), + Metadata: map[string]any{ +- "safety": safety, +- "categories": categories, +- "model": model, +- "raw": trimForDB(parsed.Raw, 512), ++ "safety": safety, ++ "observed_categories": append([]string(nil), parsed.Categories...), ++ "enforced_categories": append([]string(nil), categories...), ++ "model": model, + }, + }} + +diff --git a/ai-gateway/internal/service/promptaudit/probe.go b/ai-gateway/internal/service/promptaudit/probe.go +index 95d6e8db20a2cb194a0e35f701d1548a1d52709b..d7439473508570da5bdd0c049eb2e8737df5c82e 100644 +--- a/ai-gateway/internal/service/promptaudit/probe.go ++++ b/ai-gateway/internal/service/promptaudit/probe.go +@@ -82,13 +82,24 @@ func promptAuditProbeErrorCode(err error) string { + } + + func promptAuditProbeErrorMessage(err error) string { +- if llmErr := llmGuardErrorFrom(err); llmErr != nil && strings.TrimSpace(llmErr.Message) != "" { +- return llmErr.Message ++ if llmErr := llmGuardErrorFrom(err); llmErr != nil { ++ switch llmErr.Code { ++ case "openai_guard_auth_failed": ++ return "Guard 认证失败" ++ case "openai_guard_timeout": ++ return "Guard 探测超时" ++ case "openai_guard_invalid_response": ++ return "Guard 返回了非法响应" ++ case "openai_guard_endpoint_denied": ++ return "Guard 地址未通过安全校验" ++ default: ++ return "Guard 探测失败" ++ } + } + if err == nil { + return "" + } +- return err.Error() ++ return "Guard 探测失败" + } + + func llmGuardErrorFrom(err error) *LLMGuardError { +diff --git a/ai-gateway/internal/service/promptaudit/probe_test.go b/ai-gateway/internal/service/promptaudit/probe_test.go +index a529981f7f9108b8a86ee4aa1d67e90f8fb54786..4eae84d8c78c7c2a39987e5ecdcb456d0358e23e 100644 +--- a/ai-gateway/internal/service/promptaudit/probe_test.go ++++ b/ai-gateway/internal/service/promptaudit/probe_test.go +@@ -100,7 +100,7 @@ func TestProbeEndpointReturnsStableGuardErrorFields(t *testing.T) { + TimeoutMS: 1000, + }) + +- if result.OK || result.Status != "error" || result.ErrorCode != "openai_guard_auth_failed" || result.Message != "bad api key" { ++ if result.OK || result.Status != "error" || result.ErrorCode != "openai_guard_auth_failed" || result.Message != "Guard 认证失败" { + t.Fatalf("unexpected probe result: %#v", result) + } + if result.HTTPStatus != http.StatusUnauthorized || result.Retryable { +@@ -118,7 +118,7 @@ func TestProbeEndpointReturnsGenericProbeFailureForUnknownError(t *testing.T) { + TimeoutMS: 1000, + }) + +- if result.OK || result.ErrorCode != "prompt_audit_probe_failed" || result.Message != "network unavailable" { ++ if result.OK || result.ErrorCode != "prompt_audit_probe_failed" || result.Message != "Guard 探测失败" { + t.Fatalf("unexpected generic probe failure: %#v", result) + } + } +diff --git a/ai-gateway/internal/service/promptaudit/qwen3guard.go b/ai-gateway/internal/service/promptaudit/qwen3guard.go +index d510b9afd47822f845de7205ed0a1b97fadfed2e..cc8dbc006642e2d95d84273f51f73ab69d885b84 100644 +--- a/ai-gateway/internal/service/promptaudit/qwen3guard.go ++++ b/ai-gateway/internal/service/promptaudit/qwen3guard.go +@@ -5,11 +5,11 @@ import ( + "strings" + ) + +- const ( +- ProtocolOpenAICompatible = "openai_compatible" +- +- DefaultQwen3GuardModel = "sileader/qwen3guard:0.6b" +- ScannerBackendQwen3Guard = "qwen3guard-openai" ++const ( ++ ProtocolOpenAICompatible = "openai_compatible" ++ ++ DefaultQwen3GuardModel = "sileader/qwen3guard:0.6b" ++ ScannerBackendQwen3Guard = "qwen3guard-openai" + + SafetySafe = "Safe" + SafetyControversial = "Controversial" +@@ -42,25 +42,26 @@ func defaultOpenAICompatibleScanners() []string { + return Qwen3GuardCategoryCatalog() + } + +- func normalizeProtocol(value string) string { +- // 提示词审计仅支持 OpenAI 兼容;历史 llm_guard / 空值一律归一。 +- _ = value +- return ProtocolOpenAICompatible +- } +- +- func normalizeGuardModel(_ string, value string) string { +- value = strings.TrimSpace(value) +- if value == "" { +- return DefaultQwen3GuardModel +- } +- return value ++func normalizeProtocol(value string) string { ++ // 提示词审计仅支持 OpenAI 兼容;历史 llm_guard / 空值一律归一。 ++ _ = value ++ return ProtocolOpenAICompatible ++} ++ ++func normalizeGuardModel(_ string, value string) string { ++ value = strings.TrimSpace(value) ++ if value == "" { ++ return DefaultQwen3GuardModel + } ++ return value ++} + + type qwen3GuardParsed struct { +- Safety string +- Categories []string +- Raw string +- Valid bool ++ Safety string ++ Categories []string ++ HasUnknownCategory bool ++ Raw string ++ Valid bool + } + + func parseQwen3GuardOutput(content string) qwen3GuardParsed { +@@ -69,16 +70,46 @@ func parseQwen3GuardOutput(content string) qwen3GuardParsed { + if content == "" { + return parsed + } +- if match := safetyLineRegexp.FindStringSubmatch(content); len(match) == 2 { +- parsed.Safety = canonicalizeSafety(match[1]) ++ nonEmptyLines := make([]string, 0, 2) ++ for _, line := range strings.Split(strings.ReplaceAll(content, "\r\n", "\n"), "\n") { ++ if strings.TrimSpace(line) != "" { ++ nonEmptyLines = append(nonEmptyLines, line) ++ } + } +- if match := categoriesLineRegexp.FindStringSubmatch(content); len(match) == 2 { +- parsed.Categories = splitCategories(match[1]) ++ if len(nonEmptyLines) != 2 { ++ return parsed + } +- parsed.Valid = parsed.Safety != "" ++ safetyMatches := safetyLineRegexp.FindAllStringSubmatch(content, -1) ++ categoryMatches := categoriesLineRegexp.FindAllStringSubmatch(content, -1) ++ if len(safetyMatches) != 1 || len(categoryMatches) != 1 { ++ return parsed ++ } ++ switch strings.TrimSpace(safetyMatches[0][1]) { ++ case SafetySafe, SafetyControversial, SafetyUnsafe: ++ parsed.Safety = strings.TrimSpace(safetyMatches[0][1]) ++ default: ++ return parsed ++ } ++ parsed.Categories = splitCategories(categoryMatches[0][1]) ++ for _, category := range parsed.Categories { ++ if !isKnownQwen3GuardCategory(category) { ++ parsed.HasUnknownCategory = true ++ } ++ } ++ parsed.Valid = true + return parsed + } + ++func isKnownQwen3GuardCategory(category string) bool { ++ key := normalizeScannerKey(category) ++ for _, known := range qwen3GuardCategoryCatalog { ++ if normalizeScannerKey(known) == key { ++ return true ++ } ++ } ++ return false ++} ++ + func canonicalizeSafety(value string) string { + switch strings.ToLower(strings.TrimSpace(value)) { + case "safe": +@@ -128,29 +159,29 @@ func canonicalizeQwen3GuardCategory(value string) string { + return item + } + } +- aliases := map[string]string{ +- "violence": "Violent", +- "violentcontent": "Violent", +- "nonviolentillegalacts": "Non-violent Illegal Acts", +- "illegal": "Non-violent Illegal Acts", +- "sexual": "Sexual Content or Sexual Acts", +- "sexualcontent": "Sexual Content or Sexual Acts", +- "sexualcontentorsexualacts": "Sexual Content or Sexual Acts", +- "selfharm": "Suicide & Self-Harm", +- "suicide": "Suicide & Self-Harm", +- "suicideandselfharm": "Suicide & Self-Harm", +- "unethical": "Unethical Acts", +- "hate": "Unethical Acts", +- "political": "Politically Sensitive Topics", +- "politicallysensitive": "Politically Sensitive Topics", +- "politicallysensitivetopics": "Politically Sensitive Topics", +- "copyright": "Copyright Violation", +- "copyrightviolation": "Copyright Violation", +- "promptinjection": "Jailbreak", +- "llamapromptguard2": "Jailbreak", +- "injection": "Jailbreak", +- "secrets": "PII", +- } ++ aliases := map[string]string{ ++ "violence": "Violent", ++ "violentcontent": "Violent", ++ "nonviolentillegalacts": "Non-violent Illegal Acts", ++ "illegal": "Non-violent Illegal Acts", ++ "sexual": "Sexual Content or Sexual Acts", ++ "sexualcontent": "Sexual Content or Sexual Acts", ++ "sexualcontentorsexualacts": "Sexual Content or Sexual Acts", ++ "selfharm": "Suicide & Self-Harm", ++ "suicide": "Suicide & Self-Harm", ++ "suicideandselfharm": "Suicide & Self-Harm", ++ "unethical": "Unethical Acts", ++ "hate": "Unethical Acts", ++ "political": "Politically Sensitive Topics", ++ "politicallysensitive": "Politically Sensitive Topics", ++ "politicallysensitivetopics": "Politically Sensitive Topics", ++ "copyright": "Copyright Violation", ++ "copyrightviolation": "Copyright Violation", ++ "promptinjection": "Jailbreak", ++ "llamapromptguard2": "Jailbreak", ++ "injection": "Jailbreak", ++ "secrets": "PII", ++ } + if mapped, ok := aliases[key]; ok { + return mapped + } +diff --git a/ai-gateway/internal/service/promptaudit/runtime.go b/ai-gateway/internal/service/promptaudit/runtime.go +index befb3463c092e42230530574ff38f4e5f04b1854..a3f9959d3c2ba17e8a0adb71a83139de6f2b57a8 100644 +--- a/ai-gateway/internal/service/promptaudit/runtime.go ++++ b/ai-gateway/internal/service/promptaudit/runtime.go +@@ -33,21 +33,46 @@ func Runtime(ctx context.Context, repo JobRepository, configSvc *ConfigService) + if configSvc == nil { + configSvc = NewConfigService(nil) + } +- cfg, err := configSvc.Public(ctx) +- if err != nil { +- cfg = DefaultConfig().Public() ++ cfg, configLoadErr := configSvc.Public(ctx) ++ if configLoadErr != nil { ++ recordConfigLoadError(configLoadErr) ++ activeCfg, _, _, _ := configLoadRuntimeState() ++ if activeCfg.ConfigVersion > 0 { ++ cfg = activeCfg.Public() ++ } else { ++ cfg = DefaultConfig().Public() ++ } + } + snapshot := RuntimeSnapshot{ +- Enabled: cfg.Enabled, +- ProcessStatus: "not_started", +- QueueCapacity: cfg.QueueCapacity, +- WorkerTotal: cfg.WorkerCount, +- LLMGuardConnectivity: publicConfigConnectivity(cfg), +- StorageSupported: repo.StorageSupported(), +- Config: cfg, +- QueueBackend: queueBackendName(repo), +- PayloadStore: payloadStoreName(defaultPayloadStore), +- PayloadStoreDegraded: payloadStoreDegraded(defaultPayloadStore), ++ Enabled: cfg.Enabled, ++ BlockingEnabled: cfg.Enabled && cfg.BlockingEnabled, ++ EffectiveMode: effectivePromptAuditMode(cfg.Enabled, cfg.BlockingEnabled), ++ ExpectedConfigVersion: normalizeConfigVersion(cfg.ConfigVersion), ++ ProcessStatus: "not_started", ++ QueueCapacity: cfg.QueueCapacity, ++ WorkerTotal: cfg.WorkerCount, ++ LLMGuardConnectivity: publicConfigConnectivity(cfg), ++ StorageSupported: repo.StorageSupported(), ++ Config: cfg, ++ PromptGuardMetrics: GetPromptGuardMetricsSnapshot(), ++ QueueBackend: queueBackendName(repo), ++ PayloadStore: payloadStoreName(defaultPayloadStore), ++ PayloadStoreDegraded: payloadStoreDegraded(defaultPayloadStore), ++ } ++ activeCfg, loadedAt, loadErr, loadErrAt := configLoadRuntimeState() ++ if activeCfg.ConfigVersion > 0 { ++ snapshot.ActiveConfigVersion = normalizeConfigVersion(activeCfg.ConfigVersion) ++ } ++ if !loadedAt.IsZero() { ++ loadedAt = loadedAt.UTC() ++ snapshot.ConfigLoadedAt = &loadedAt ++ } ++ if loadErr != "" { ++ snapshot.ConfigLoadError = loadErr ++ if !loadErrAt.IsZero() { ++ loadErrAt = loadErrAt.UTC() ++ snapshot.ConfigLoadErrorAt = &loadErrAt ++ } + } + if repo.StorageSupported() { + if stats, err := repo.RuntimeDBStats(ctx); err == nil { +@@ -93,13 +118,21 @@ func Runtime(ctx context.Context, repo JobRepository, configSvc *ConfigService) + } + if cfg.Enabled && !snapshot.StorageSupported { + snapshot.ProcessStatus = "error" +- snapshot.LastErrorCode = "storage_not_supported" + snapshot.LastErrorMessage = "提示词审计日志 Ent 客户端未初始化" ++ if cfg.BlockingEnabled { ++ snapshot.ProcessStatus = "degraded" ++ snapshot.LastErrorMessage += ";同步判定仍可执行,但结果记录降级" ++ } ++ snapshot.LastErrorCode = "storage_not_supported" + } + if cfg.Enabled && (defaultPayloadStore == nil || !defaultPayloadStore.Available()) { + snapshot.ProcessStatus = "error" +- snapshot.LastErrorCode = "payload_store_unavailable" + snapshot.LastErrorMessage = ErrPayloadStoreUnavailable.Error() ++ if cfg.BlockingEnabled { ++ snapshot.ProcessStatus = "degraded" ++ snapshot.LastErrorMessage += ";同步判定不依赖异步载荷存储" ++ } ++ snapshot.LastErrorCode = "payload_store_unavailable" + } + if cfg.Enabled && snapshot.PayloadStoreDegraded && snapshot.ProcessStatus != "error" { + snapshot.ProcessStatus = "degraded" +@@ -108,9 +141,28 @@ func Runtime(ctx context.Context, repo JobRepository, configSvc *ConfigService) + snapshot.LastErrorMessage = "提示词审计正在使用内存 payload store,仅适合单进程开发或测试" + } + } ++ if configLoadErr != nil { ++ if snapshot.ProcessStatus != "error" { ++ snapshot.ProcessStatus = "degraded" ++ } ++ if snapshot.LastErrorCode == "" { ++ snapshot.LastErrorCode = "config_load_failed" ++ snapshot.LastErrorMessage = "提示词审计配置加载失败" ++ } ++ } + return snapshot + } + ++func effectivePromptAuditMode(enabled bool, blockingEnabled bool) string { ++ if !enabled { ++ return "off" ++ } ++ if blockingEnabled { ++ return "blocking" ++ } ++ return "async_audit" ++} ++ + func (r *Runner) heartbeatLoop(ctx context.Context, cfg Config) { + defer r.wg.Done() + ticker := time.NewTicker(10 * time.Second) +diff --git a/ai-gateway/internal/service/promptaudit/runtime_coverage_test.go b/ai-gateway/internal/service/promptaudit/runtime_coverage_test.go +index ebbf045b144ca3c212fbfc19d3606c729264b9e2..6fe7669763d4110c1fa3b1028ea67dffb5d507bb 100644 +--- a/ai-gateway/internal/service/promptaudit/runtime_coverage_test.go ++++ b/ai-gateway/internal/service/promptaudit/runtime_coverage_test.go +@@ -86,6 +86,8 @@ func TestRuntimeMergesDBStatsAndHeartbeatWithObservablePriority(t *testing.T) { + + func TestRuntimeFallsBackToSafeConfigWhenConfigPublicLoadFails(t *testing.T) { + clearRuntimeHeartbeatForTesting() ++ ClearConfigCache() ++ t.Cleanup(ClearConfigCache) + repo := newFakePromptAuditRepo() + svc := NewConfigService(&configurableOptionStore{ + values: map[string]string{}, +@@ -100,8 +102,14 @@ func TestRuntimeFallsBackToSafeConfigWhenConfigPublicLoadFails(t *testing.T) { + if snapshot.WorkerTotal != 4 || snapshot.QueueCapacity != 10000 { + t.Fatalf("runtime safe defaults mismatch: worker=%d capacity=%d", snapshot.WorkerTotal, snapshot.QueueCapacity) + } +- if snapshot.ProcessStatus == "error" || snapshot.LastErrorCode != "" { +- t.Fatalf("config read failure should not fabricate runner error state, got status=%s code=%s", snapshot.ProcessStatus, snapshot.LastErrorCode) ++ if snapshot.ProcessStatus != "degraded" || snapshot.LastErrorCode != "config_load_failed" { ++ t.Fatalf("配置读取失败应明确标记运行态降级,got status=%s code=%s", snapshot.ProcessStatus, snapshot.LastErrorCode) ++ } ++ if snapshot.ConfigLoadError != "提示词审计配置加载失败" || snapshot.ConfigLoadErrorAt == nil { ++ t.Fatalf("配置读取失败应只暴露通用脱敏错误,got error=%q at=%v", snapshot.ConfigLoadError, snapshot.ConfigLoadErrorAt) ++ } ++ if strings.Contains(snapshot.ConfigLoadError, "option storage unavailable") || strings.Contains(snapshot.LastErrorMessage, "option storage unavailable") { ++ t.Fatalf("运行态不得泄露底层配置存储错误: %+v", snapshot) + } + } + +diff --git a/ai-gateway/internal/service/promptaudit/types.go b/ai-gateway/internal/service/promptaudit/types.go +index c6d9f4f0ccd1257ee94614156a597fc7870a5545..ca5254997fafff386e044dfebb1e9e7dea0f1093 100644 +--- a/ai-gateway/internal/service/promptaudit/types.go ++++ b/ai-gateway/internal/service/promptaudit/types.go +@@ -53,16 +53,17 @@ type PayloadStore interface { + } + + type ScanPromptContext struct { +- JobID int64 +- RequestID string +- UserID int +- TokenID int +- ChannelID int +- Endpoint string +- Protocol string +- Model string +- Group string +- Lease *ScanPromptLease ++ JobID int64 ++ RequestID string ++ UserID int ++ TokenID int ++ ChannelID int ++ Endpoint string ++ Protocol string ++ Model string ++ Group string ++ StopOnBlock bool ++ Lease *ScanPromptLease + } + + // ScanPromptLease 允许长文本分片扫描在每片开始前刷新 processing 租约, +@@ -370,28 +371,36 @@ type RuntimeDBStats struct { + } + + type RuntimeSnapshot struct { +- Enabled bool `json:"enabled"` +- ProcessStatus string `json:"process_status"` +- QueueLength int `json:"queue_length"` +- QueueCapacity int `json:"queue_capacity"` +- WorkerTotal int `json:"worker_total"` +- ActiveWorkers int64 `json:"active_workers"` +- Enqueued int64 `json:"enqueued"` +- Dropped int64 `json:"dropped"` +- ProcessedTotal int64 `json:"processed_total"` +- FailedTotal int64 `json:"failed_total"` +- LastErrorCode string `json:"last_error_code"` +- LastErrorMessage string `json:"last_error_message"` +- LLMGuardConnectivity string `json:"llm_guard_connectivity"` +- StorageSupported bool `json:"storage_supported"` +- QueueBackend string `json:"queue_backend"` +- PayloadStore string `json:"payload_store"` +- PayloadStoreDegraded bool `json:"payload_store_degraded"` +- QueuedRows int64 `json:"queued_rows"` +- ProcessingRows int64 `json:"processing_rows"` +- LastEnqueuedAt *time.Time `json:"last_enqueued_at"` +- LastProcessedAt *time.Time `json:"last_processed_at"` +- LastFailedAt *time.Time `json:"last_failed_at"` +- HeartbeatAt *time.Time `json:"heartbeat_at"` +- Config PublicConfig `json:"config"` ++ Enabled bool `json:"enabled"` ++ BlockingEnabled bool `json:"blocking_enabled"` ++ EffectiveMode string `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"` ++ ConfigLoadErrorAt *time.Time `json:"config_load_error_at,omitempty"` ++ ProcessStatus string `json:"process_status"` ++ QueueLength int `json:"queue_length"` ++ QueueCapacity int `json:"queue_capacity"` ++ WorkerTotal int `json:"worker_total"` ++ ActiveWorkers int64 `json:"active_workers"` ++ Enqueued int64 `json:"enqueued"` ++ Dropped int64 `json:"dropped"` ++ ProcessedTotal int64 `json:"processed_total"` ++ FailedTotal int64 `json:"failed_total"` ++ LastErrorCode string `json:"last_error_code"` ++ LastErrorMessage string `json:"last_error_message"` ++ LLMGuardConnectivity string `json:"llm_guard_connectivity"` ++ StorageSupported bool `json:"storage_supported"` ++ QueueBackend string `json:"queue_backend"` ++ PayloadStore string `json:"payload_store"` ++ PayloadStoreDegraded bool `json:"payload_store_degraded"` ++ QueuedRows int64 `json:"queued_rows"` ++ ProcessingRows int64 `json:"processing_rows"` ++ LastEnqueuedAt *time.Time `json:"last_enqueued_at"` ++ LastProcessedAt *time.Time `json:"last_processed_at"` ++ LastFailedAt *time.Time `json:"last_failed_at"` ++ HeartbeatAt *time.Time `json:"heartbeat_at"` ++ Config PublicConfig `json:"config"` ++ PromptGuardMetrics PromptGuardMetricsSnapshot `json:"prompt_guard_metrics"` + } +diff --git a/ai-gateway/internal/service/promptaudit/worker.go b/ai-gateway/internal/service/promptaudit/worker.go +index 0aec513f15a5120ba9ef83f1a5027a6ec87bbd02..f6e059106fe2af03ee44bb86cc891c0766267341 100644 +--- a/ai-gateway/internal/service/promptaudit/worker.go ++++ b/ai-gateway/internal/service/promptaudit/worker.go +@@ -73,10 +73,14 @@ func (r *Runner) Start(ctx context.Context) error { + return ErrPayloadStoreUnavailable + } + workerCount := normalizeWorkerCount(cfg.WorkerCount) ++ installConfigSnapshot(cfg) ++ StartConfigInvalidationSubscriber(ctx) + LogInfoEvent( + "prompt_audit.started", + Field("status", "running"), + Field("enabled", cfg.Enabled), ++ Field("blocking_enabled", cfg.Enabled && cfg.BlockingEnabled), ++ Field("config_version", normalizeConfigVersion(cfg.ConfigVersion)), + Field("queue_capacity", cfg.QueueCapacity), + Field("worker_total", workerCount), + Field("endpoint_count", len(cfg.EnabledEndpoints())), +diff --git a/ai-gateway/internal/types/error.go b/ai-gateway/internal/types/error.go +index 501ee10d7660fb6f6d91baddd44624bbe32120ab..a6a9d55d966b5bb0c50efe8a02ebdef6bfd3252c 100644 +--- a/ai-gateway/internal/types/error.go ++++ b/ai-gateway/internal/types/error.go +@@ -75,14 +75,17 @@ const ( + ErrorCodeBadRequestBody ErrorCode = "bad_request_body" + + // response error +- ErrorCodeReadResponseBodyFailed ErrorCode = "read_response_body_failed" +- ErrorCodeBadResponseStatusCode ErrorCode = "bad_response_status_code" +- ErrorCodeBadResponse ErrorCode = "bad_response" +- ErrorCodeBadResponseBody ErrorCode = "bad_response_body" +- ErrorCodeEmptyResponse ErrorCode = "empty_response" +- ErrorCodeAwsInvokeError ErrorCode = "aws_invoke_error" +- ErrorCodeModelNotFound ErrorCode = "model_not_found" +- ErrorCodePromptBlocked ErrorCode = "prompt_blocked" ++ ErrorCodeReadResponseBodyFailed ErrorCode = "read_response_body_failed" ++ ErrorCodeBadResponseStatusCode ErrorCode = "bad_response_status_code" ++ ErrorCodeBadResponse ErrorCode = "bad_response" ++ ErrorCodeBadResponseBody ErrorCode = "bad_response_body" ++ ErrorCodeEmptyResponse ErrorCode = "empty_response" ++ ErrorCodeAwsInvokeError ErrorCode = "aws_invoke_error" ++ ErrorCodeModelNotFound ErrorCode = "model_not_found" ++ ErrorCodePromptBlocked ErrorCode = "prompt_blocked" ++ ErrorCodePromptGuardBlocked ErrorCode = "prompt_guard_blocked" ++ ErrorCodePromptGuardUnavailable ErrorCode = "prompt_guard_unavailable" ++ ErrorCodePromptGuardInvalidResponse ErrorCode = "prompt_guard_invalid_response" + + // sql error + ErrorCodeQueryDataError ErrorCode = "query_data_error" +diff --git a/deploy/.env.example b/deploy/.env.example +index 01d5fb78a0550abcdcd9444088c38cc1404148bb..19bb22221b9e1520da78e73c94b7ea7e6ea4a705 100644 +--- a/deploy/.env.example ++++ b/deploy/.env.example +@@ -306,22 +306,28 @@ AICODEX_REALTIME_WS_OUTBOUND_QUEUE_SIZE=64 + # AICODEX_REALTIME_WS_ALLOWED_ORIGINS=https://console.example.com + + # --- 用户输入提示词审计(默认关闭)--- +-# 主 aicodex 进程内置审计 worker:模型请求只异步入队,worker 调用外部 LLM Guard HTTP API 扫描。 +-# 开启前必须配置可用的 LLM Guard API,并确保 Redis 可用(完整提示词只以短 TTL 临时载荷写入 Redis,不落库)。 ++# 主 aicodex 进程内置 Qwen3Guard 审计 worker;默认异步只审计,可显式开启同步阻止。 ++# 开启前必须配置可用的 OpenAI 兼容 Guard;异步载荷使用 Redis 短 TTL,同步路径不持久化原文。 + PROMPT_AUDIT_ENABLED=false ++# true=请求在渠道选择、计费和上游调用前同步等待 Guard;Block 或 Guard 不可用时 fail-closed。 ++PROMPT_AUDIT_BLOCKING_ENABLED=false + # 是否持久化 pass 事件;默认 false,仅保存 flag / critical 等风险事件。 + PROMPT_AUDIT_STORE_PASS_EVENTS=false +-# LLM Guard API 选择策略:priority / weighted / shadow ++# Guard 节点调度仅支持 priority(有序故障切换)。 + PROMPT_AUDIT_STRATEGY=priority + # 主进程内审计 worker 数;建议按外部 LLM Guard API 实际吞吐灰度调大。 + PROMPT_AUDIT_WORKER_COUNT=4 + # 队列容量上限;达到上限时主请求继续转发,并输出 prompt_audit.enqueue_dropped。 + PROMPT_AUDIT_QUEUE_CAPACITY=10000 +-# 输入扫描器列表,逗号分隔。 +-PROMPT_AUDIT_SCANNERS=PromptInjection,TokenLimit,Secrets,InvisibleText,Gibberish,Regex +-# Guard Prompt Scan URL,多个地址用逗号分隔;推荐直接填写完整审核路由。 +-# 生产可使用外部审核服务,例如:https://scan.leagsoft.com/v1/scan/prompt +-LLM_GUARD_SCAN_URLS=https://scan.leagsoft.com/v1/scan/prompt ++# Qwen3Guard 输入类别,逗号分隔。 ++PROMPT_AUDIT_SCANNERS=Violent,Non-violent Illegal Acts,Sexual Content or Sexual Acts,PII,Suicide & Self-Harm,Unethical Acts,Politically Sensitive Topics,Copyright Violation,Jailbreak ++# 推荐使用 OpenAI 兼容 Base URL、模型和 API Key;公网地址必须 HTTPS。 ++PROMPT_AUDIT_BASE_URLS= ++PROMPT_AUDIT_MODEL=sileader/qwen3guard:0.6b ++PROMPT_AUDIT_API_KEYS= ++PROMPT_AUDIT_TIMEOUT_MS=30000 ++# 旧变量兼容:按 OpenAI 兼容 Base URL 解释,多个地址用逗号分隔;新部署优先使用 PROMPT_AUDIT_BASE_URLS。 ++LLM_GUARD_SCAN_URLS= + # 旧 LLM Guard API Base URL 兼容变量;为空时优先使用 LLM_GUARD_SCAN_URLS。 + # 如果只填写服务根地址,AICodex 会兼容补齐 /v1/scan/prompt。 + LLM_GUARD_API_BASE_URLS= +diff --git a/docs/constraints/41-ai-readable-logging.md b/docs/constraints/41-ai-readable-logging.md +index 7be9278ffd3ad7468fda6d21118ef387fc2b3075..7f2df768ab3e3d60de6ae7e0f3a0913fb66ab65f 100644 +--- a/docs/constraints/41-ai-readable-logging.md ++++ b/docs/constraints/41-ai-readable-logging.md +@@ -592,3 +592,36 @@ App 下载中心以华为云 OBS/CDN updater metadata 为事实源时,必须 + metadata 缓存事件必须用 `cache_key_hash` 表示缓存键,不得输出完整缓存键、完整 metadata URL、完整 package URL 或 metadata body。`ETag` 与 `Last-Modified` 只允许输出是否存在的布尔值,不得输出 header 原文。 + + 禁止输出 Cookie、Authorization、完整带 query 的 URL、OBS 密钥、CDN 鉴权参数或安装包签名 URL 原文。日志中如需表示 metadata 或 package 来源,只能输出 host、平台键、文件名和允许范围内的稳定状态字段。 ++ ++## 13. Prompt Guard 同步门禁专项约束 ++ ++提示词同步阻止属于请求副作用边界,日志必须能直接回答“使用了哪个配置版本、为什么放行或拒绝、拒绝前是否触发渠道/计费/上游”。 ++ ++稳定事件名: ++ ++- `prompt_guard.config_updated` ++- `prompt_guard.config_loaded` ++- `prompt_guard.config_reload_degraded` ++- `prompt_guard.evaluation_started` ++- `prompt_guard.allowed` ++- `prompt_guard.blocked` ++- `prompt_guard.failed` ++- `prompt_guard.result_record_failed` ++ ++最小字段集合: ++ ++- `request_id`、`user_id`、`token_id`、`group` ++- `protocol`、`endpoint`、`model` ++- `config_version`、`policy_id`、`policy_version`、`guard_endpoint_id` ++- `decision`、`action`、`chunk_total`、`latency_ms` ++- `status`、`error_code`、`stage` ++- `upstream_dispatched`、`billing_preconsumed` ++ ++稳定错误码: ++ ++- `prompt_guard_blocked` ++- `prompt_guard_unavailable` ++- `prompt_guard_invalid_response` ++- `prompt_guard_requires_audit_enabled` ++ ++Block、Unavailable 和非法响应日志必须明确 `upstream_dispatched=false`、`billing_preconsumed=false`。禁止输出完整提示词、原始分片、API Key、Token、Authorization、完整 Guard URL、URL query、Guard 原始响应或内部优先分片边界。分类只允许输出归一化后的类别与稳定 scanner 名称。 +diff --git a/docs/workflows/02-local-dev.md b/docs/workflows/02-local-dev.md +index 2ca1e966f237ae170aebbc15c5e461f2b05f3cde..683b5dcd389508b7e079702c7f5e614fa145db5f 100644 +--- a/docs/workflows/02-local-dev.md ++++ b/docs/workflows/02-local-dev.md +@@ -25,21 +25,24 @@ + + ### 提示词审计本地验证 + +-提示词审计由主 `aicodex` 进程异步投递任务,并在主进程内置 worker 中消费任务、调用外部 Guard Prompt Scan URL。完整原始提示词只会以短 TTL 临时载荷写入 Redis,数据库只保存 hash、脱敏预览、上下文和扫描结果。 ++提示词审计支持两种执行模式:`blocking_enabled=false` 为异步只审计;`blocking_enabled=true` 为同步阻止。同步模式会在渠道选择、计费预扣和上游调用前调用 OpenAI 兼容 Qwen3Guard,命中 Block 返回 403,Guard 不可用或输出非法返回 503。同步结果直接复用到脱敏事件,不重复调用 Guard。 + +-- 推荐在 `deploy/.env` 中设置新的完整审核 URL 和 API Key: ++- 推荐在控制台保存配置;也可在 `deploy/.env` 中设置 OpenAI 兼容 Guard: + - `PROMPT_AUDIT_ENABLED=true` +- - `LLM_GUARD_SCAN_URLS=https://scan.leagsoft.com/v1/scan/prompt` +- - `LLM_GUARD_API_TOKENS=sk-lg_xxx`,需替换为本地私有 API Key;如果外接多个 Guard API,可用逗号分隔,且不得提交真实 key。 ++ - `PROMPT_AUDIT_BLOCKING_ENABLED=false`(先以异步模式建立基线,灰度时再开启) ++ - `PROMPT_AUDIT_BASE_URLS=https://guard.example.com/v1` ++ - `PROMPT_AUDIT_MODEL=sileader/qwen3guard:0.6b` ++ - `PROMPT_AUDIT_API_KEYS=sk_xxx`,只写入本地私有配置,不得提交真实 key。 + - 旧 `laiyer/llm-guard-api` sidecar 仅作为迁移兼容样例保留: + - 随主栈 profile 启动:`cd deploy && docker compose -p aicodex --profile prompt-audit-llm-guard up -d llm-guard-api` + - 使用旧 sidecar 时可设置 `LLM_GUARD_SCAN_URLS=http://127.0.0.1:8000/v1/scan/prompt`;如果仍填写旧 `LLM_GUARD_API_BASE_URLS=http://127.0.0.1:8000`,AICodex 会兼容补齐 `/v1/scan/prompt`。 + - 启动主服务:`cd deploy && docker compose -p aicodex up -d aicodex` +-- 查看主服务审计日志:`cd deploy && docker compose -p aicodex logs -f aicodex | grep prompt_audit` +-- 连通性验证:登录控制台打开“HTTP 审计 → 提示词审计”后点击 endpoint 探测,或调用 `POST /api/prompt-audit/endpoints/probe`;后端会优先检查审核 URL 所属 origin 的 `/health`,必要时用安全探针调用 `scan_url`。 +-- 运行态验证:调用 `GET /api/prompt-audit/runtime`,确认 `enabled=true`、`process_status=running`、`payload_store=redis`、`payload_store_degraded=false`、`llm_guard_connectivity=ok`。常见稳定错误码包括 `llm_guard_auth_failed`、`llm_guard_timeout`、`llm_guard_http_error`、`llm_guard_invalid_response`、`scan_payload_missing` 和 `payload_store_unavailable`。 ++- 查看主服务审计日志:`cd deploy && docker compose -p aicodex logs -f aicodex | grep -E 'prompt_audit|prompt_guard'` ++- 连通性验证:登录控制台打开“提示词审计”后点击审计池探测,或调用 `POST /api/prompt-audit/endpoints/probe`。公网 Guard 只允许 HTTPS;HTTP 仅允许 localhost、单标签内部服务名或显式私网 IP;重定向、link-local 和云元数据地址会被拒绝。 ++- 运行态验证:调用 `GET /api/prompt-audit/runtime`,确认 `effective_mode`、`expected_config_version`、`active_config_version`、`config_loaded_at`、`process_status` 和 `llm_guard_connectivity` 符合预期。版本不一致或 `config_load_error` 非空时不得扩大灰度。 + - 事件落库验证:发送包含 PromptInjection 特征的 `/v1/chat/completions`、`/v1/responses` 或 Claude Messages 请求,再查询 `GET /api/prompt-audit/events?decision=critical`;页面和接口只能展示脱敏预览、hash、scanner 命中和处理元数据。 +-- 回滚方式:将 `PROMPT_AUDIT_ENABLED=false` 后重启主服务;提示词审计是异步旁路,关闭后不会影响主模型请求转发,已写入的审计事件可继续保留用于复核。 ++- 同步验证:先用良性输入确认请求成功,再用 fake Guard 分别返回 `Safety: Unsafe / Categories: Jailbreak`、超时和非法格式,确认 HTTP 分别返回 403/503,且上游调用数、渠道重试数和预扣次数均为 0;Responses WebSocket 首轮和后续 `response.create` 也必须在本轮预扣前检查。 ++- 回滚方式:在控制台关闭“同步阻止”并保存,即刻恢复异步只审计;无需关闭审计或删除历史事件。若需完全停用,再关闭“启用审计”。 + + ## 前端 + +diff --git a/webui/src/api/promptAudit.test.ts b/webui/src/api/promptAudit.test.ts +index 2c1030573522eacc893b4f47c77ee26563ccbeb5..94f2a5f07a1d18bca2de73c7c22e49f49f65394c 100644 +--- a/webui/src/api/promptAudit.test.ts ++++ b/webui/src/api/promptAudit.test.ts +@@ -76,6 +76,7 @@ const eventListParams: PromptAuditEventListParams = { + + const savePayload: PromptAuditConfigSavePayload = { + enabled: true, ++ blocking_enabled: false, + store_pass_events: false, + strategy: 'priority', + worker_count: 16, +diff --git a/webui/src/features/prompt-audit/PromptAuditPage.test.tsx b/webui/src/features/prompt-audit/PromptAuditPage.test.tsx +index ede5204bd9a2d01541b09d9ab2468d3d98f7b4b4..243d722a91d77b2a8ec077825a804341cccd0e9d 100644 +--- a/webui/src/features/prompt-audit/PromptAuditPage.test.tsx ++++ b/webui/src/features/prompt-audit/PromptAuditPage.test.tsx +@@ -80,6 +80,8 @@ const savePromptAuditConfigMock = vi.mocked(savePromptAuditConfig) + + const configResponse: PromptAuditConfigResponse = { + enabled: true, ++ blocking_enabled: false, ++ config_version: 3, + store_pass_events: false, + strategy: 'priority', + worker_count: 16, +@@ -103,6 +105,10 @@ const configResponse: PromptAuditConfigResponse = { + + const runtimeResponse: PromptAuditRuntime = { + enabled: true, ++ blocking_enabled: false, ++ effective_mode: 'async_audit', ++ expected_config_version: 3, ++ active_config_version: 3, + process_status: 'running', + queue_length: 0, + queue_capacity: 10000, +@@ -267,6 +273,7 @@ describe('PromptAuditPage', () => { + expect(savePromptAuditConfigMock).toHaveBeenCalledWith( + expect.objectContaining({ + enabled: true, ++ blocking_enabled: false, + store_pass_events: false, + worker_count: 16, + queue_capacity: 10000, +@@ -290,6 +297,77 @@ describe('PromptAuditPage', () => { + expect(screen.getByLabelText('主审计池 API Key')).toHaveValue('') + }) + ++ it('同步阻止需要确认,保存后关闭审计会自动关闭阻止', async () => { ++ confirmImmediately() ++ const user = userEvent.setup() ++ renderWithRouter() ++ ++ const blockingSwitch = await screen.findByRole('switch', { ++ name: '同步阻止', ++ }) ++ expect(blockingSwitch).not.toBeChecked() ++ await user.click(blockingSwitch) ++ expect(Modal.confirm).toHaveBeenCalledWith( ++ expect.objectContaining({ title: '确认开启同步阻止' }), ++ ) ++ expect(blockingSwitch).toBeChecked() ++ ++ await user.click(screen.getByRole('button', { name: /保存配置/ })) ++ await waitFor(() => { ++ expect(savePromptAuditConfigMock).toHaveBeenCalledWith( ++ expect.objectContaining({ blocking_enabled: true }), ++ ) ++ }) ++ ++ await user.click(screen.getByRole('switch', { name: '启用审计' })) ++ expect(blockingSwitch).not.toBeChecked() ++ expect(blockingSwitch).toBeDisabled() ++ }) ++ ++ it('取消同步阻止风险确认时保持异步只审计草稿', async () => { ++ vi.spyOn(Modal, 'confirm').mockImplementation(() => undefined as never) ++ const user = userEvent.setup() ++ renderWithRouter() ++ ++ const blockingSwitch = await screen.findByRole('switch', { ++ name: '同步阻止', ++ }) ++ await user.click(blockingSwitch) ++ ++ expect(Modal.confirm).toHaveBeenCalledWith( ++ expect.objectContaining({ title: '确认开启同步阻止' }), ++ ) ++ expect(blockingSwitch).not.toBeChecked() ++ expect(screen.getByText('异步只审计')).toBeInTheDocument() ++ expect(savePromptAuditConfigMock).not.toHaveBeenCalled() ++ }) ++ ++ it('未保存草稿不得冒充运行时审计状态', async () => { ++ const user = userEvent.setup() ++ renderWithRouter() ++ ++ expect(await screen.findByText('审计已启用')).toBeInTheDocument() ++ await user.click(screen.getByRole('switch', { name: '启用审计' })) ++ ++ expect(screen.getByText('审计已启用')).toBeInTheDocument() ++ expect(screen.queryByText('审计未启用')).not.toBeInTheDocument() ++ expect(screen.getByText('有未保存更改')).toBeInTheDocument() ++ }) ++ ++ it('运行态配置版本不一致时明确展示降级提示和双版本', async () => { ++ fetchPromptAuditRuntimeMock.mockResolvedValueOnce({ ++ ...runtimeResponse, ++ expected_config_version: 4, ++ active_config_version: 3, ++ config_load_error: 'config_load_failed', ++ }) ++ renderWithRouter() ++ ++ expect(await screen.findByText('配置版本未同步')).toBeInTheDocument() ++ expect(screen.getByText('期望版本: 4')).toBeInTheDocument() ++ expect(screen.getByText('生效版本: 3')).toBeInTheDocument() ++ }) ++ + it('通过参数弹框修改权重、超时和单片输入上限后统一保存', async () => { + const user = userEvent.setup() + renderWithRouter() +diff --git a/webui/src/features/prompt-audit/PromptAuditPage.tsx b/webui/src/features/prompt-audit/PromptAuditPage.tsx +index 2bc1aa5cd959a28dc79675344ba518132a49f104..03228aae0ff26f971e83cba86262d6202ae9259e 100644 +--- a/webui/src/features/prompt-audit/PromptAuditPage.tsx ++++ b/webui/src/features/prompt-audit/PromptAuditPage.tsx +@@ -306,6 +306,31 @@ const PromptAuditPage = () => { + value: PromptAuditConfigState[K], + ) => setConfig((current) => ({ ...current, [key]: value })) + ++ const updateAuditEnabled = (enabled: boolean) => { ++ setConfig((current) => ({ ++ ...current, ++ enabled, ++ blockingEnabled: enabled ? current.blockingEnabled : false, ++ })) ++ } ++ ++ const updateBlockingEnabled = (enabled: boolean) => { ++ if (!enabled) { ++ updateConfig('blockingEnabled', false) ++ return ++ } ++ if (!config.enabled) return ++ Modal.confirm({ ++ title: t('确认开启同步阻止'), ++ content: t( ++ '开启后,请求会在转发前等待 Guard 判定;命中 Block 或 Guard 不可用时不会访问上游,并将分别返回 403 或 503。', ++ ), ++ okText: t('确认开启'), ++ cancelText: t('取消'), ++ onOk: () => updateConfig('blockingEnabled', true), ++ }) ++ } ++ + const updateEndpoint = ( + endpointID: string, + patch: Partial, +@@ -383,7 +408,7 @@ const PromptAuditPage = () => { + const loadRuntime = useCallback(async () => { + const response = await fetchPromptAuditRuntime() + setRuntime(response) +- if (response.config) { ++ if (response.config && savedSnapshotRef.current === '') { + const normalized = normalizePromptAuditConfig(response.config) + setConfig({ + ...normalized, +@@ -512,6 +537,24 @@ const PromptAuditPage = () => { + config.auditGroupMode === 'all' || groupsLoading || groupsLoadFailed + + const queueBacklog = Number(runtime?.queue_length || 0) > 0 ++ const runtimeEffectiveMode = ++ runtime?.effective_mode || ++ (runtime?.enabled ++ ? runtime?.blocking_enabled ++ ? 'blocking' ++ : 'async_audit' ++ : 'off') ++ const runtimeModeLabel = ++ runtimeEffectiveMode === 'blocking' ++ ? t('同步阻止') ++ : runtimeEffectiveMode === 'async_audit' ++ ? t('异步只审计') ++ : t('审计关闭') ++ const runtimeVersionMismatch = ++ Number(runtime?.expected_config_version || 0) > 0 && ++ Number(runtime?.expected_config_version) !== ++ Number(runtime?.active_config_version || 0) ++ const runtimeAuditEnabled = runtime?.enabled ?? config.enabled + + const saveConfig = useCallback(async () => { + if (saving) return +@@ -534,6 +577,7 @@ const PromptAuditPage = () => { + commitSavedSnapshot(buildConfigSnapshot(nextConfig)) + } + Toast.success(t('提示词审计配置已保存')) ++ await loadRuntime() + } catch (error: unknown) { + const message = readPromptAuditErrorMessage( + error, +@@ -544,7 +588,7 @@ const PromptAuditPage = () => { + } finally { + setSaving(false) + } +- }, [commitSavedSnapshot, config, saving, t]) ++ }, [commitSavedSnapshot, config, loadRuntime, saving, t]) + + useEffect(() => { + const onKeyDown = (event: KeyboardEvent) => { +@@ -1216,12 +1260,23 @@ const PromptAuditPage = () => { + icon={} + actions={ +
+- +- {config.enabled ? t('审计已启用') : t('审计未启用')} ++ ++ {runtimeAuditEnabled ? t('审计已启用') : t('审计未启用')} + + + {runtime?.process_status || runtimeStatusTag.label} + ++ ++ {runtimeModeLabel} ++ ++ {runtimeVersionMismatch ? ( ++ ++ {t('配置版本未同步')} ++ ++ ) : null} + {isDirty ? ( + + {t('有未保存更改')} +@@ -1292,7 +1347,9 @@ const PromptAuditPage = () => { +
+ {`${t('Worker 数')}: ${config.workerCount}`} + {`${t('队列容量')}: ${config.queueCapacity}`} +- {`${t('调度策略')}: ${config.strategy === 'round_robin' ? t('轮询') : t('优先级')}`} ++ {`${t('调度策略')}: ${t('优先级故障切换')}`} ++ {`${t('期望版本')}: ${runtime?.expected_config_version || config.configVersion}`} ++ {`${t('生效版本')}: ${runtime?.active_config_version || 0}`} +
+
+ prompt_audit.started +@@ -1348,20 +1405,7 @@ const PromptAuditPage = () => { +
+
+ +- ++ {t('优先级故障切换')} + +@@ -2521,8 +2565,27 @@ const PromptAuditPage = () => { + + updateConfig('enabled', enabled)} ++ onChange={updateAuditEnabled} ++ size='small' ++ aria-label={t('启用审计')} ++ /> ++ ++ +
diff --git a/frontend/src/features/prompt-audit/components/EventWorkspace.vue b/frontend/src/features/prompt-audit/components/EventWorkspace.vue index 1bf166a89..9e6464a5a 100644 --- a/frontend/src/features/prompt-audit/components/EventWorkspace.vue +++ b/frontend/src/features/prompt-audit/components/EventWorkspace.vue @@ -79,7 +79,7 @@ {{ event.snapshot.group_name || '—' }}

{{ event.snapshot.endpoint }}

-

{{ event.snapshot.model }} · {{ event.snapshot.protocol }}

+

{{ event.snapshot.model }} · {{ event.snapshot.protocol }} · {{ event.snapshot.stage || 'http' }}

{{ event.decision }} · {{ event.risk_level }} diff --git a/frontend/src/i18n/locales/en/admin/promptAudit.ts b/frontend/src/i18n/locales/en/admin/promptAudit.ts index cbe8113b3..8e214d116 100644 --- a/frontend/src/i18n/locales/en/admin/promptAudit.ts +++ b/frontend/src/i18n/locales/en/admin/promptAudit.ts @@ -40,7 +40,7 @@ export default { 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.', + detailTitle: 'Prompt audit event details', tabs: { summary: 'Audit summary', risks: 'Specific risks', technical: 'Technical details' }, redactedPreview: 'Irreversible redacted preview', categories: 'Categories', model: 'Model', stage: 'Request stage', 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.' }, diff --git a/frontend/src/i18n/locales/zh/admin/promptAudit.ts b/frontend/src/i18n/locales/zh/admin/promptAudit.ts index c39f33e46..5bc5f6a81 100644 --- a/frontend/src/i18n/locales/zh/admin/promptAudit.ts +++ b/frontend/src/i18n/locales/zh/admin/promptAudit.ts @@ -40,7 +40,7 @@ export default { 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: '本事件没有派生风险摘要。', + detailTitle: '提示词审计事件详情', tabs: { summary: '审计摘要', risks: '具体风险', technical: '技术信息' }, redactedPreview: '不可逆脱敏预览', categories: '分类', model: '模型', stage: '请求阶段', noRisks: '本事件没有派生风险摘要。', deleteConfirmTitle: '删除审计事件?', deleteConfirmMessage: '将永久删除 {count} 条事件及符合条件的孤立任务。', filterDeleteTitle: '确认按筛选删除', filterDeleteCount: '服务端快照匹配 {count} 条事件。', snapshotMax: '快照最大事件 ID', expiresAt: '确认令牌过期时间', filterDeleteWarning: '只删除预览高水位内的事件;预览后产生的新事件会保留。筛选一旦变化,必须重新预览。', confirmFilterDelete: '确认永久删除', }, messages: { saved: '提示词审计配置已保存,明文 API Key 状态已清除。', probeSucceeded: '审计节点连接正常。', deleted: '已删除 {count} 条审计事件。' }, From df9d9e2e4032128c994ac6aa09c28a8cb9eef623 Mon Sep 17 00:00:00 2001 From: mt21625457 Date: Fri, 17 Jul 2026 09:00:14 +0800 Subject: [PATCH 3/7] fix(security-audit): harden role scan, startup, probe, and localhost dial Scan client-injected assistant/tool/model turns, fail closed when config cannot be trusted after startup or stale invalidation, reuse probe tokens only for the same base URL, and restrict localhost dials to loopback addresses. Co-authored-by: Cursor --- backend/cmd/server/main.go | 7 +-- .../securityaudit/prompt_config_store.go | 43 ++++++++++++++++++- .../securityaudit/prompt_config_test.go | 16 +++++++ .../securityaudit/prompt_outbound_security.go | 15 ++++++- .../prompt_outbound_security_test.go | 28 ++++++++++++ .../internal/securityaudit/prompt_service.go | 20 ++++++--- .../internal/securityaudit/prompt_snapshot.go | 20 +++++++-- .../securityaudit/prompt_snapshot_test.go | 28 ++++++------ .../securityaudit/prompt_worker_test.go | 2 +- 9 files changed, 147 insertions(+), 32 deletions(-) diff --git a/backend/cmd/server/main.go b/backend/cmd/server/main.go index c9b2b41a4..02d6f00b6 100644 --- a/backend/cmd/server/main.go +++ b/backend/cmd/server/main.go @@ -155,9 +155,10 @@ func runMainServer() { 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) + // Startup continues so unrelated APIs stay up, but Prompt Audit itself + // fails closed (unavailable) until a later reload installs a trusted + // snapshot—avoiding a silent ModeOff bypass of persisted blocking policy. + log.Printf("Prompt Audit started in degraded fail-closed state: %v", err) } } diff --git a/backend/internal/securityaudit/prompt_config_store.go b/backend/internal/securityaudit/prompt_config_store.go index e67f115fe..795a8476c 100644 --- a/backend/internal/securityaudit/prompt_config_store.go +++ b/backend/internal/securityaudit/prompt_config_store.go @@ -36,6 +36,11 @@ type ConfigManager struct { // independently of whether endpoint credentials or the full config could be // activated. A config version alone cannot distinguish async from blocking. expectedBlocking atomic.Bool + // configUntrusted is set when a load/reload fails before a trustworthy + // snapshot is installed. While set, EffectiveMode fails closed so a + // persisted blocking policy cannot be silently skipped after startup or + // invalidation errors. + configUntrusted atomic.Bool stateMu sync.RWMutex lastLoadError string @@ -63,6 +68,9 @@ func (m *ConfigManager) Start(ctx context.Context) error { m.cancel = cancel m.lifecycleMu.Unlock() loadErr := m.Reload(runCtx) + if loadErr != nil { + m.markConfigUntrusted() + } m.wg.Add(1) go m.refreshLoop(runCtx) if m.redis != nil { @@ -89,17 +97,20 @@ func (m *ConfigManager) Shutdown(_ context.Context) error { func (m *ConfigManager) Reload(ctx context.Context) error { if m == nil || m.settings == nil { + m.markUntrustedIfNoActiveSnapshot() return errors.New("prompt audit setting repository unavailable") } values, err := m.settings.GetMultiple(ctx, []string{SettingKeyPromptAuditConfig, SettingKeyRiskControl}) if err != nil { m.recordLoadError(err) + m.markUntrustedIfNoActiveSnapshot() return err } m.observeExpectedState(values[SettingKeyPromptAuditConfig], values[SettingKeyRiskControl] == "true") storage, err := ParseStorageConfig(values[SettingKeyPromptAuditConfig]) if err != nil { m.recordLoadError(err) + m.markUntrustedIfNoActiveSnapshot() return err } m.expected.Store(storage.ConfigVersion) @@ -107,10 +118,13 @@ func (m *ConfigManager) Reload(ctx context.Context) error { active, err := ActiveFromStorage(storage, values[SettingKeyRiskControl] == "true", m.encryptor) if err != nil { m.recordLoadError(err) + // expectedBlocking may already require fail-closed via BlockingActivationDegraded. + m.markUntrustedIfNoActiveSnapshot() return err } now := m.clock.Now() m.snapshot.Store(&activeConfigSnapshot{storage: cloneStorageConfig(storage), active: cloneActiveConfig(active), loadedAt: now}) + m.configUntrusted.Store(false) m.clearLoadError() LogInfo(EventConfigLoaded, map[string]any{ "config_version": storage.ConfigVersion, "status": "loaded", @@ -130,7 +144,13 @@ func (m *ConfigManager) Active() (ActiveConfig, bool) { } func (m *ConfigManager) BlockingActivationDegraded() bool { - if m == nil || !m.expectedBlocking.Load() { + if m == nil { + return false + } + if m.configUntrusted.Load() { + return true + } + if !m.expectedBlocking.Load() { return false } active, ok := m.Active() @@ -153,6 +173,22 @@ func (m *ConfigManager) EffectiveMode() Mode { return active.EffectiveMode() } +func (m *ConfigManager) markConfigUntrusted() { + if m == nil { + return + } + m.configUntrusted.Store(true) +} + +func (m *ConfigManager) markUntrustedIfNoActiveSnapshot() { + if m == nil { + return + } + if _, ok := m.Active(); !ok { + m.markConfigUntrusted() + } +} + func (m *ConfigManager) Public() PublicConfig { if m == nil { return PublicFromStorage(DefaultStorageConfig(), false) @@ -378,6 +414,11 @@ func (m *ConfigManager) subscribeLoop(ctx context.Context) { } m.expected.Store(version) if err := m.Reload(ctx); err != nil { + // A newer published version failed to activate. Until reload + // succeeds, do not keep serving a potentially stale weaker mode. + if active, ok := m.Active(); !ok || active.ConfigVersion < version { + m.markConfigUntrusted() + } LogWarn(EventConfigReloadDegraded, map[string]any{ "config_version": version, "status": "degraded", "error_code": "config_invalidation_reload_failed", }) diff --git a/backend/internal/securityaudit/prompt_config_test.go b/backend/internal/securityaudit/prompt_config_test.go index 399f5d250..ad25dac9e 100644 --- a/backend/internal/securityaudit/prompt_config_test.go +++ b/backend/internal/securityaudit/prompt_config_test.go @@ -131,6 +131,22 @@ func TestConfigManagerStaleWeakerSnapshotFailsClosedWhenBlockingExpected(t *test require.Equal(t, ErrorCodeUnavailable, guardErr.Code) } +type errorSettingRepository struct{ staticSettingRepository } + +func (errorSettingRepository) GetMultiple(context.Context, []string) (map[string]string, error) { + return nil, errors.New("settings unavailable") +} + +func TestConfigManagerStartupLoadFailureFailsClosedWithoutSnapshot(t *testing.T) { + manager := NewConfigManager(nil, errorSettingRepository{}, nil, prefixEncryptor{}) + err := manager.Start(context.Background()) + require.Error(t, err) + require.True(t, manager.configUntrusted.Load()) + require.True(t, manager.BlockingActivationDegraded()) + require.Equal(t, ModeBlocking, manager.EffectiveMode()) + require.NoError(t, manager.Shutdown(context.Background())) +} + func TestParseLegacyConfigDefaultsMissingFieldsWithoutEnablingBlocking(t *testing.T) { storage, err := ParseStorageConfig(`{"enabled":false,"config_version":9}`) require.NoError(t, err) diff --git a/backend/internal/securityaudit/prompt_outbound_security.go b/backend/internal/securityaudit/prompt_outbound_security.go index 40df98ef0..f1e3ac3fc 100644 --- a/backend/internal/securityaudit/prompt_outbound_security.go +++ b/backend/internal/securityaudit/prompt_outbound_security.go @@ -164,11 +164,22 @@ func secureDialContext(dialer *net.Dialer, resolver DNSResolver, allowPrivate bo } var lastErr error for _, addr := range addresses { - if isBlockedAddress(addr) || (!allowPrivate && (addr.IsPrivate() || addr.IsLoopback())) { + if isBlockedAddress(addr) { lastErr = fmt.Errorf("prompt guard resolved address blocked") continue } - if !addr.IsGlobalUnicast() && !addr.IsPrivate() && !addr.IsLoopback() { + if allowPrivate { + // localhost / *.localhost may only resolve to loopback. A hosts or + // DNS mapping from localhost to RFC1918 must not become an SSRF pivot. + if !addr.IsLoopback() { + lastErr = fmt.Errorf("prompt guard resolved address blocked") + continue + } + } else if addr.IsPrivate() || addr.IsLoopback() { + lastErr = fmt.Errorf("prompt guard resolved address blocked") + continue + } + if !addr.IsGlobalUnicast() && !addr.IsLoopback() { lastErr = fmt.Errorf("prompt guard resolved address blocked") continue } diff --git a/backend/internal/securityaudit/prompt_outbound_security_test.go b/backend/internal/securityaudit/prompt_outbound_security_test.go index 15e56d0bc..76d7f9fdb 100644 --- a/backend/internal/securityaudit/prompt_outbound_security_test.go +++ b/backend/internal/securityaudit/prompt_outbound_security_test.go @@ -49,6 +49,12 @@ func TestSecureDialRejectsDNSRebindingToPrivateAddress(t *testing.T) { require.Error(t, err) } +func TestSecureDialLocalhostAllowlistRejectsRFC1918Resolution(t *testing.T) { + dial := secureDialContext(nil, staticResolver{addresses: []netip.Addr{netip.MustParseAddr("10.0.0.8")}}, true) + _, err := dial(context.Background(), "tcp", "localhost:8080") + require.Error(t, err) +} + func TestSecureHTTPClientDoesNotBypassDestinationValidationThroughEnvironmentProxy(t *testing.T) { client, err := NewSecureHTTPClient(ActiveEndpoint{BaseURL: "https://guard.example.com", TimeoutMS: 1000}) require.NoError(t, err) @@ -207,6 +213,28 @@ func TestPromptAuditProbeModelsFallbackAndResponseSafety(t *testing.T) { }) } +func TestResolveProbeEndpointReusesTokenOnlyForMatchingBaseURL(t *testing.T) { + manager := &ConfigManager{} + manager.snapshot.Store(&activeConfigSnapshot{active: ActiveConfig{Endpoints: []ActiveEndpoint{{ + ID: "guard-1", BaseURL: "https://guard.example.com", Token: "STORED_GUARD_TOKEN", TimeoutMS: 1000, InputLimit: 1024, Enabled: true, + }}}}) + service := &PromptService{config: manager} + + matched, applied, err := service.resolveProbeEndpoint(UpdateEndpoint{ + ID: "guard-1", BaseURL: "https://guard.example.com/v1", TimeoutMS: 1000, InputLimit: 1024, + }) + require.NoError(t, err) + require.True(t, applied) + require.Equal(t, "STORED_GUARD_TOKEN", matched.Token) + + mismatched, applied, err := service.resolveProbeEndpoint(UpdateEndpoint{ + ID: "guard-1", BaseURL: "https://attacker.example.com", TimeoutMS: 1000, InputLimit: 1024, + }) + require.NoError(t, err) + require.False(t, applied) + require.Empty(t, mismatched.Token) +} + func newProbeTestService() *PromptService { return &PromptService{ config: &ConfigManager{}, scanner: NewOpenAICompatibleScanner(), clock: realClock{}, diff --git a/backend/internal/securityaudit/prompt_service.go b/backend/internal/securityaudit/prompt_service.go index 8402e60d9..71d0fadbe 100644 --- a/backend/internal/securityaudit/prompt_service.go +++ b/backend/internal/securityaudit/prompt_service.go @@ -312,21 +312,27 @@ func modelsResponseReady(body []byte, model string) bool { } func (s *PromptService) resolveProbeEndpoint(input UpdateEndpoint) (ActiveEndpoint, bool, error) { + baseURL, err := NormalizeBaseURL(input.BaseURL) + if err != nil { + return ActiveEndpoint{}, false, err + } 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 + if endpoint.ID != strings.TrimSpace(input.ID) { + continue } + // Reuse a stored credential only when the probe targets the same + // normalized base URL. Otherwise an admin probe could exfiltrate + // the Guard token to an attacker-controlled HTTPS host. + if endpoint.BaseURL == baseURL { + 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 diff --git a/backend/internal/securityaudit/prompt_snapshot.go b/backend/internal/securityaudit/prompt_snapshot.go index 13002ded1..6b87ad186 100644 --- a/backend/internal/securityaudit/prompt_snapshot.go +++ b/backend/internal/securityaudit/prompt_snapshot.go @@ -92,7 +92,10 @@ func extractProtocolSegments(protocol string, document any) []string { } } -var clientInstructionRoles = []string{"user", "system", "developer"} +// clientInstructionRoles are roles a client may freely populate. Attackers can +// place jailbreak/PII text in assistant/tool turns, so blocking audit must scan +// them too—not only user/system/developer instructions. +var clientInstructionRoles = []string{"user", "system", "developer", "assistant", "tool"} func extractChatLikeSegments(root map[string]any) []string { if root == nil { @@ -168,7 +171,7 @@ func extractResponses(value any) []string { result = append(result, entry) case map[string]any: role := strings.ToLower(stringValue(entry["role"])) - if role != "" && role != "user" && role != "system" && role != "developer" { + if role != "" && !isClientInstructionRole(role) { continue } if content, exists := entry["content"]; exists { @@ -183,7 +186,7 @@ func extractResponses(value any) []string { return result case map[string]any: role := strings.ToLower(stringValue(typed["role"])) - if role != "" && role != "user" && role != "system" && role != "developer" { + if role != "" && !isClientInstructionRole(role) { return nil } return contentTexts(typed["content"]) @@ -192,6 +195,15 @@ func extractResponses(value any) []string { } } +func isClientInstructionRole(role string) bool { + switch strings.ToLower(strings.TrimSpace(role)) { + case "user", "system", "developer", "assistant", "tool", "model": + return true + default: + return false + } +} + func extractGemini(value any) []string { var contents []any switch typed := value.(type) { @@ -209,7 +221,7 @@ func extractGemini(value any) []string { continue } role := strings.ToLower(stringValue(content["role"])) - if role != "" && role != "user" { + if role != "" && !isClientInstructionRole(role) { continue } parts, _ := content["parts"].([]any) diff --git a/backend/internal/securityaudit/prompt_snapshot_test.go b/backend/internal/securityaudit/prompt_snapshot_test.go index 17ab36371..a1427ba1a 100644 --- a/backend/internal/securityaudit/prompt_snapshot_test.go +++ b/backend/internal/securityaudit/prompt_snapshot_test.go @@ -17,7 +17,7 @@ func TestExtractPromptSnapshotProtocols(t *testing.T) { 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_chat_completions", `{"messages":[{"role":"user","content":"old"},{"role":"assistant","content":"assistant turn"},{"role":"user","content":[{"type":"text","text":"最新😀"}]}]}`, "最新😀", 3}, {"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}, @@ -67,8 +67,8 @@ 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":"assistant","content":"assistant client injection"}, + {"role":"tool","content":"tool client injection"}, {"role":"user","content":[ {"type":"text","text":"最新第一块😀"}, {"type":"image_url","image_url":{"url":"data:image/png;base64,IMAGE_CANARY_BASE64"}}, @@ -78,10 +78,11 @@ func TestPromptSnapshotLatestUserMessageIsOnePrioritizedSegment(t *testing.T) { }`) 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.Equal(t, 4, snapshot.MessageCount) + require.True(t, strings.HasPrefix(snapshot.ScanText, "最新第一块😀\n最新第二块é")) + require.Contains(t, snapshot.ScanText, "历史输入") + require.Contains(t, snapshot.ScanText, "assistant client injection") + require.Contains(t, snapshot.ScanText, "tool client injection") require.NotContains(t, snapshot.ScanText, "IMAGE_CANARY_BASE64") require.Equal(t, utf8.RuneCountInString(snapshot.ScanText), snapshot.PromptLength) } @@ -93,7 +94,7 @@ func TestPromptSnapshotResponsesShapes(t *testing.T) { 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: "message array", body: `{"input":[{"role":"assistant","content":"assistant turn"},{"role":"user","content":[{"type":"input_text","text":"message block"}]}]}`, want: "message block\n\nassistant turn"}, {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"}, } @@ -122,7 +123,7 @@ func TestPromptSnapshotGeminiBatchShapesAndMediaExclusion(t *testing.T) { require.Contains(t, snapshot.ScanText, expected) } require.NotContains(t, snapshot.ScanText, "ROOT_BASE64") - require.NotContains(t, snapshot.ScanText, "ignore model") + require.Contains(t, snapshot.ScanText, "ignore model") } func TestPromptSnapshotMediaOnlyExtractsDeterministicTextPrompts(t *testing.T) { @@ -163,7 +164,7 @@ func TestResponsesWebSocketOnlyAuditsResponseCreateAndPreservesStage(t *testing. } func TestPromptSnapshotEmptyAndLongUnicodeInput(t *testing.T) { - _, err := ExtractPromptSnapshot(Request{Protocol: "openai_chat_completions", Body: []byte(`{"messages":[{"role":"assistant","content":"not user"},{"role":"user","content":" "}]}`)}) + _, err := ExtractPromptSnapshot(Request{Protocol: "openai_chat_completions", Body: []byte(`{"messages":[{"role":"function","content":"not audited role"},{"role":"user","content":" "}]}`)}) require.True(t, errors.Is(err, ErrNoPromptText)) latest := strings.Repeat("最新😀é", 80) @@ -186,10 +187,10 @@ func TestPromptSnapshotIncludesClientControlledInstructions(t *testing.T) { want []string }{ { - name: "openai system and developer", + name: "openai system developer assistant tool", protocol: "openai_chat_completions", - body: `{"messages":[{"role":"system","content":"system jailbreak"},{"role":"developer","content":"developer policy"},{"role":"assistant","content":"ignore"},{"role":"user","content":"hello"}]}`, - want: []string{"system jailbreak", "developer policy", "hello"}, + body: `{"messages":[{"role":"system","content":"system jailbreak"},{"role":"developer","content":"developer policy"},{"role":"assistant","content":"assistant jailbreak"},{"role":"tool","content":"tool payload"},{"role":"user","content":"hello"}]}`, + want: []string{"system jailbreak", "developer policy", "assistant jailbreak", "tool payload", "hello"}, }, { name: "openai system only", @@ -223,7 +224,6 @@ func TestPromptSnapshotIncludesClientControlledInstructions(t *testing.T) { for _, expected := range tt.want { require.Contains(t, snapshot.ScanText, expected) } - require.NotContains(t, snapshot.ScanText, "ignore") }) } } diff --git a/backend/internal/securityaudit/prompt_worker_test.go b/backend/internal/securityaudit/prompt_worker_test.go index 6c7ed7f3b..5c39be8e4 100644 --- a/backend/internal/securityaudit/prompt_worker_test.go +++ b/backend/internal/securityaudit/prompt_worker_test.go @@ -302,7 +302,7 @@ func TestEnqueuerSkipsOffOutOfScopeAndNoText(t *testing.T) { 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"}]}`)}}, + {name: "no user text", cfg: asyncConfig(), req: Request{Protocol: "openai_chat_completions", Body: []byte(`{"messages":[{"role":"function","content":"not audited"}]}`)}}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) { From 18e698bed6bf98e9affda25b3a82314e31e489ed Mon Sep 17 00:00:00 2001 From: mt21625457 Date: Fri, 17 Jul 2026 11:45:14 +0800 Subject: [PATCH 4/7] feat(security-audit): allow admin-managed audit node targets and polish pool UI Let admins configure private/intranet Guard endpoints without destination-class blocking, and fix prompt-audit switch layout so thumbs and labels no longer overlap. Co-authored-by: Cursor --- .../securityaudit/prompt_outbound_security.go | 139 +---------------- .../prompt_outbound_security_test.go | 52 +++---- deploy/build_image.sh | 0 deploy/docker-compose.yml | 2 +- .../features/prompt-audit/PromptAuditView.vue | 142 +++++++++++++----- .../__tests__/PromptAuditView.spec.ts | 30 ++++ .../prompt-audit/components/EndpointPool.vue | 135 ++++++++++------- .../components/EventWorkspace.vue | 18 +-- .../prompt-audit/components/PolicyPanel.vue | 12 +- .../components/RuntimeOverview.vue | 100 ++++++------ .../src/i18n/locales/en/admin/promptAudit.ts | 2 + .../src/i18n/locales/zh/admin/promptAudit.ts | 2 + frontend/src/style.css | 8 + .../design.md | 14 +- .../specs/prompt-input-audit/spec.md | 18 ++- .../tasks.md | 13 ++ 16 files changed, 342 insertions(+), 345 deletions(-) mode change 100644 => 100755 deploy/build_image.sh diff --git a/backend/internal/securityaudit/prompt_outbound_security.go b/backend/internal/securityaudit/prompt_outbound_security.go index f1e3ac3fc..81e987eb4 100644 --- a/backend/internal/securityaudit/prompt_outbound_security.go +++ b/backend/internal/securityaudit/prompt_outbound_security.go @@ -1,13 +1,9 @@ package securityaudit import ( - "context" "crypto/tls" - "errors" - "fmt" "net" "net/http" - "net/netip" "net/url" "strings" "time" @@ -17,40 +13,6 @@ import ( 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) @@ -64,29 +26,10 @@ func NormalizeBaseURL(raw string) (string, error) { if parsed.User != nil || parsed.RawQuery != "" || parsed.Fragment != "" { return "", infraerrors.BadRequest("prompt_audit_unsafe_base_url", "审计节点地址不能包含凭据、查询参数或片段") } - host := strings.ToLower(strings.TrimSuffix(parsed.Hostname(), ".")) + host := strings.TrimSpace(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", "审计节点地址不在允许范围") - } - // Loopback literals remain available for local Guard nodes and tests. - // RFC1918 literals are rejected so an admin session cannot pivot into - // arbitrary private-network services; use a hostname allowlist instead. - if addr.IsPrivate() { - return "", infraerrors.BadRequest("prompt_audit_unsafe_base_url", "审计节点地址不在允许范围") - } - allowPrivate = 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 = "" @@ -113,17 +56,10 @@ func ModelsURL(base string) (string, error) { } func NewSecureHTTPClient(endpoint ActiveEndpoint) (*http.Client, error) { - normalized, err := NormalizeBaseURL(endpoint.BaseURL) + _, 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.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 @@ -138,7 +74,10 @@ func NewSecureHTTPClient(endpoint ActiveEndpoint) (*http.Client, error) { ExpectContinueTimeout: time.Second, TLSClientConfig: &tls.Config{MinVersion: tls.VersionTLS12}, } - transport.DialContext = secureDialContext(dialer, resolver, allowPrivate) + // Endpoint ownership and destination trust are administrator concerns. + // Use the standard dialer so configured private, loopback, reserved, and + // DNS-resolved addresses are all reachable from the service environment. + transport.DialContext = dialer.DialContext timeout := time.Duration(endpoint.TimeoutMS) * time.Millisecond if timeout <= 0 { timeout = DefaultTimeoutMS * time.Millisecond @@ -146,71 +85,5 @@ func NewSecureHTTPClient(endpoint ActiveEndpoint) (*http.Client, error) { 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) { - lastErr = fmt.Errorf("prompt guard resolved address blocked") - continue - } - if allowPrivate { - // localhost / *.localhost may only resolve to loopback. A hosts or - // DNS mapping from localhost to RFC1918 must not become an SSRF pivot. - if !addr.IsLoopback() { - lastErr = fmt.Errorf("prompt guard resolved address blocked") - continue - } - } else if addr.IsPrivate() || addr.IsLoopback() { - lastErr = fmt.Errorf("prompt guard resolved address blocked") - continue - } - if !addr.IsGlobalUnicast() && !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 { - // Only the localhost name family is trusted for private/loopback dials. - // A bare "*.local" suffix is too broad (mDNS/intranet names) and would - // re-open RFC1918 SSRF after literal private IPs were rejected. - return host == "localhost" || strings.HasSuffix(host, ".localhost") -} - -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 -} diff --git a/backend/internal/securityaudit/prompt_outbound_security_test.go b/backend/internal/securityaudit/prompt_outbound_security_test.go index 76d7f9fdb..327504d1b 100644 --- a/backend/internal/securityaudit/prompt_outbound_security_test.go +++ b/backend/internal/securityaudit/prompt_outbound_security_test.go @@ -5,7 +5,6 @@ import ( "encoding/json" "net/http" "net/http/httptest" - "net/netip" "strings" "sync/atomic" "testing" @@ -14,25 +13,20 @@ import ( "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://localhost:8080"} +func TestNormalizeBaseURLAllowsAdministratorConfiguredDestinations(t *testing.T) { + allowed := []string{ + "https://guard.example.com", "https://guard.example.com/v1", "http://guard.example.com", + "http://127.0.0.1:8080", "http://10.0.0.8:8080", "https://172.16.0.5", + "http://169.254.169.254", "https://metadata.google.internal", "https://192.0.2.1", + "http://internal-admin.local", "http://guard.local: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]", - "http://10.0.0.8:8080", "http://192.168.1.10:8080", "https://172.16.0.5", - "http://internal-admin.local", "http://guard.local:8080", + "ftp://guard.example.com", "https://user:pass@guard.example.com", + "https://guard.example.com?q=secret", "https://guard.example.com/#fragment", } for _, raw := range blocked { _, err := NormalizeBaseURL(raw) @@ -43,24 +37,13 @@ func TestNormalizeBaseURLSecurity(t *testing.T) { 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 TestSecureDialLocalhostAllowlistRejectsRFC1918Resolution(t *testing.T) { - dial := secureDialContext(nil, staticResolver{addresses: []netip.Addr{netip.MustParseAddr("10.0.0.8")}}, true) - _, err := dial(context.Background(), "tcp", "localhost:8080") - require.Error(t, err) -} - -func TestSecureHTTPClientDoesNotBypassDestinationValidationThroughEnvironmentProxy(t *testing.T) { +func TestHTTPClientUsesDirectStandardDialer(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) + require.NotNil(t, transport.DialContext) } func TestOpenAICompatibleScannerRequestContract(t *testing.T) { @@ -83,13 +66,16 @@ func TestOpenAICompatibleScannerRequestContract(t *testing.T) { 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) +func TestOpenAICompatibleScannerFollowsRedirectAndRejectsOversize(t *testing.T) { + target := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + _, _ = w.Write([]byte(`{"choices":[{"message":{"content":"Safety: Safe\nCategories: None"}}]}`)) })) + defer target.Close() + redirect := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { http.Redirect(w, r, target.URL, 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) + result, err := NewOpenAICompatibleScanner().Scan(context.Background(), ActiveEndpoint{ID: "redirect", BaseURL: redirect.URL, Model: DefaultGuardModel, TimeoutMS: 1000}, "hello", AllScannerIDs) + require.NoError(t, err) + require.Equal(t, EventPass, result.Decision) oversize := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { _, _ = w.Write([]byte(strings.Repeat("x", int(maxGuardResponseBytes)+1))) })) diff --git a/deploy/build_image.sh b/deploy/build_image.sh old mode 100644 new mode 100755 diff --git a/deploy/docker-compose.yml b/deploy/docker-compose.yml index 6aecdcfa5..637c2762f 100644 --- a/deploy/docker-compose.yml +++ b/deploy/docker-compose.yml @@ -16,7 +16,7 @@ services: # Sub2API Application # =========================================================================== sub2api: - image: weishaw/sub2api:latest + image: sub2api:latest container_name: sub2api restart: unless-stopped ulimits: diff --git a/frontend/src/features/prompt-audit/PromptAuditView.vue b/frontend/src/features/prompt-audit/PromptAuditView.vue index e0491bbdc..939fc52ac 100644 --- a/frontend/src/features/prompt-audit/PromptAuditView.vue +++ b/frontend/src/features/prompt-audit/PromptAuditView.vue @@ -1,6 +1,6 @@