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 = () => { - - updateConfig('strategy', String(value)) - } - style={{ width: 180 }} - > - - {t('优先级')} - - - {t('轮询')} - - + {t('优先级故障切换')} } onClick={addEndpoint}> {t('添加审计池')} @@ -2521,8 +2565,27 @@ const PromptAuditPage = () => { updateConfig('enabled', enabled)} + onChange={updateAuditEnabled} + size='small' + aria-label={t('启用审计')} + /> + + + + + {t('同步阻止')} + + diff --git a/webui/src/features/prompt-audit/promptAuditViewModel.test.ts b/webui/src/features/prompt-audit/promptAuditViewModel.test.ts index b90cc77727e54a1c22552a1d9f054ef700732ed4..0163323fa9603bbe7854ebaf58cdf84cfd88d126 100644 --- a/webui/src/features/prompt-audit/promptAuditViewModel.test.ts +++ b/webui/src/features/prompt-audit/promptAuditViewModel.test.ts @@ -68,6 +68,8 @@ import type { const configResponse = { enabled: true, + blocking_enabled: true, + config_version: 7, store_pass_events: false, strategy: 'priority', worker_count: 16, @@ -333,6 +335,8 @@ describe('promptAuditViewModel', () => { const config = normalizePromptAuditConfig(configResponse) expect(config).toMatchObject({ enabled: true, + blockingEnabled: true, + configVersion: 7, workerCount: '16', queueCapacity: '10000', scanners: ['Jailbreak', 'PII'], @@ -349,11 +353,13 @@ describe('promptAuditViewModel', () => { const snapshot = buildConfigSnapshot(config) expect(parsePromptAuditSnapshot(snapshot)).toMatchObject({ enabled: true, + blockingEnabled: true, auditGroupMode: 'selected', }) expect(parsePromptAuditSnapshot('bad json')).toBeNull() expect(buildPromptAuditConfigPayload(config)).toMatchObject({ enabled: true, + blocking_enabled: true, worker_count: 16, queue_capacity: 10000, audit_group_mode: 'selected', @@ -371,6 +377,26 @@ describe('promptAuditViewModel', () => { }) }) + it('旧配置响应缺少同步字段时安全回退到异步只审计', () => { + const legacy = normalizePromptAuditConfig({ + ...configResponse, + blocking_enabled: undefined, + config_version: undefined, + strategy: 'round_robin', + }) + expect(legacy).toMatchObject({ + enabled: true, + blockingEnabled: false, + configVersion: 1, + strategy: 'priority', + }) + expect(buildPromptAuditConfigPayload(legacy)).toMatchObject({ + enabled: true, + blocking_enabled: false, + strategy: 'priority', + }) + }) + it('运行态、探测结果和端点选项保持旧语义', () => { const config = normalizePromptAuditConfig(configResponse) expect(endpointStatusFromConnectivity('ok', true)).toBe('healthy') diff --git a/webui/src/features/prompt-audit/promptAuditViewModel.ts b/webui/src/features/prompt-audit/promptAuditViewModel.ts index 9150a61a005585a8a98e5d2aef7c79bd466fa375..fd52612c7cc0761f9f00d09727c4f7f1abd5363c 100644 --- a/webui/src/features/prompt-audit/promptAuditViewModel.ts +++ b/webui/src/features/prompt-audit/promptAuditViewModel.ts @@ -151,6 +151,7 @@ export const DEFAULT_PROMPT_AUDIT_ENDPOINTS: PromptAuditEndpointState[] = [ export const DEFAULT_PROMPT_AUDIT_CONFIG: PromptAuditConfigState = { enabled: false, + blockingEnabled: false, storePassEvents: false, strategy: 'priority', workerCount: '4', @@ -159,6 +160,7 @@ export const DEFAULT_PROMPT_AUDIT_CONFIG: PromptAuditConfigState = { auditGroupMode: 'all', auditGroups: [], endpoints: DEFAULT_PROMPT_AUDIT_ENDPOINTS, + configVersion: 1, } export const DEFAULT_PROMPT_AUDIT_FILTERS: PromptAuditFilters = { @@ -793,8 +795,9 @@ export const normalizePromptAuditConfig = ( } return { enabled: Boolean(data.enabled), + blockingEnabled: Boolean(data.enabled && data.blocking_enabled), storePassEvents: Boolean(data.store_pass_events), - strategy: data.strategy || 'priority', + strategy: 'priority', workerCount: String(data.worker_count || 4), queueCapacity: String(data.queue_capacity || 10000), scanners: @@ -804,6 +807,7 @@ export const normalizePromptAuditConfig = ( auditGroupMode: normalizeAuditGroupMode(data.audit_group_mode), auditGroups: normalizeStringList(data.audit_groups), endpoints: data.endpoints.map(normalizeEndpointFromAPI), + configVersion: Math.max(1, Number(data.config_version || 1)), } } @@ -834,14 +838,16 @@ const endpointEditableSnapshot = (endpoint: PromptAuditEndpointState) => ({ export const buildConfigSnapshot = (config: PromptAuditConfigState): string => JSON.stringify({ enabled: config.enabled, + blockingEnabled: config.enabled && config.blockingEnabled, storePassEvents: config.storePassEvents, - strategy: config.strategy || 'priority', + strategy: 'priority', workerCount: String(config.workerCount ?? ''), queueCapacity: String(config.queueCapacity ?? ''), scanners: normalizeStringList(config.scanners).sort(), auditGroupMode: normalizeAuditGroupMode(config.auditGroupMode), auditGroups: normalizeStringList(config.auditGroups).sort(), endpoints: config.endpoints.map(endpointEditableSnapshot), + configVersion: Math.max(1, Number(config.configVersion || 1)), }) export const parsePromptAuditSnapshot = ( @@ -876,8 +882,9 @@ export const buildPromptAuditConfigPayload = ( config: PromptAuditConfigState, ): PromptAuditConfigSavePayload => ({ enabled: config.enabled, + blocking_enabled: config.enabled && config.blockingEnabled, store_pass_events: config.storePassEvents, - strategy: config.strategy, + strategy: 'priority', worker_count: toPositiveInt(config.workerCount, 4), queue_capacity: toPositiveInt(config.queueCapacity, 10000), scanners: normalizeStringList(config.scanners), diff --git a/webui/src/types/promptAudit.ts b/webui/src/types/promptAudit.ts index c99e817e57766960b3f0467572554ddf9bd48771..de68778d5b4af479952f455dd3c495fb49bb5be6 100644 --- a/webui/src/types/promptAudit.ts +++ b/webui/src/types/promptAudit.ts @@ -81,6 +81,7 @@ export interface PromptAuditEndpointState { export interface PromptAuditConfigResponse { enabled?: boolean + blocking_enabled?: boolean store_pass_events?: boolean strategy?: string worker_count?: number @@ -89,10 +90,14 @@ export interface PromptAuditConfigResponse { audit_group_mode?: string audit_groups?: string[] endpoints?: PromptAuditEndpointAPI[] + config_version?: number + updated_at?: string + change_summary?: string } export interface PromptAuditConfigState { enabled: boolean + blockingEnabled: boolean storePassEvents: boolean strategy: string workerCount: string @@ -101,6 +106,7 @@ export interface PromptAuditConfigState { auditGroupMode: PromptAuditGroupMode auditGroups: string[] endpoints: PromptAuditEndpointState[] + configVersion: number } export interface PromptAuditEndpointSavePayload { @@ -119,6 +125,7 @@ export interface PromptAuditEndpointSavePayload { export interface PromptAuditConfigSavePayload { enabled: boolean + blocking_enabled: boolean store_pass_events: boolean strategy: string worker_count: number @@ -131,6 +138,13 @@ export interface PromptAuditConfigSavePayload { export interface PromptAuditRuntime { enabled?: boolean + blocking_enabled?: boolean + effective_mode?: 'off' | 'async_audit' | 'blocking' | string + expected_config_version?: number + active_config_version?: number + config_loaded_at?: string + config_load_error?: string + config_load_error_at?: string process_status?: string queue_length?: number queue_capacity?: number
prompt_audit.started