From df9d9e2e4032128c994ac6aa09c28a8cb9eef623 Mon Sep 17 00:00:00 2001 From: mt21625457 Date: Fri, 17 Jul 2026 09:00:14 +0800 Subject: [PATCH] fix(security-audit): harden role scan, startup, probe, and localhost dial Scan client-injected assistant/tool/model turns, fail closed when config cannot be trusted after startup or stale invalidation, reuse probe tokens only for the same base URL, and restrict localhost dials to loopback addresses. Co-authored-by: Cursor --- backend/cmd/server/main.go | 7 +-- .../securityaudit/prompt_config_store.go | 43 ++++++++++++++++++- .../securityaudit/prompt_config_test.go | 16 +++++++ .../securityaudit/prompt_outbound_security.go | 15 ++++++- .../prompt_outbound_security_test.go | 28 ++++++++++++ .../internal/securityaudit/prompt_service.go | 20 ++++++--- .../internal/securityaudit/prompt_snapshot.go | 20 +++++++-- .../securityaudit/prompt_snapshot_test.go | 28 ++++++------ .../securityaudit/prompt_worker_test.go | 2 +- 9 files changed, 147 insertions(+), 32 deletions(-) diff --git a/backend/cmd/server/main.go b/backend/cmd/server/main.go index c9b2b41a46..02d6f00b6c 100644 --- a/backend/cmd/server/main.go +++ b/backend/cmd/server/main.go @@ -155,9 +155,10 @@ func runMainServer() { defer app.Cleanup() if app.PromptAudit != nil { if err := app.PromptAudit.Start(context.Background()); err != nil { - // Prompt Audit is default-off and isolated. Startup degradation must be - // observable but must not take unrelated APIs down. - log.Printf("Prompt Audit started in degraded state: %v", err) + // Startup continues so unrelated APIs stay up, but Prompt Audit itself + // fails closed (unavailable) until a later reload installs a trusted + // snapshot—avoiding a silent ModeOff bypass of persisted blocking policy. + log.Printf("Prompt Audit started in degraded fail-closed state: %v", err) } } diff --git a/backend/internal/securityaudit/prompt_config_store.go b/backend/internal/securityaudit/prompt_config_store.go index e67f115fe1..795a8476c4 100644 --- a/backend/internal/securityaudit/prompt_config_store.go +++ b/backend/internal/securityaudit/prompt_config_store.go @@ -36,6 +36,11 @@ type ConfigManager struct { // independently of whether endpoint credentials or the full config could be // activated. A config version alone cannot distinguish async from blocking. expectedBlocking atomic.Bool + // configUntrusted is set when a load/reload fails before a trustworthy + // snapshot is installed. While set, EffectiveMode fails closed so a + // persisted blocking policy cannot be silently skipped after startup or + // invalidation errors. + configUntrusted atomic.Bool stateMu sync.RWMutex lastLoadError string @@ -63,6 +68,9 @@ func (m *ConfigManager) Start(ctx context.Context) error { m.cancel = cancel m.lifecycleMu.Unlock() loadErr := m.Reload(runCtx) + if loadErr != nil { + m.markConfigUntrusted() + } m.wg.Add(1) go m.refreshLoop(runCtx) if m.redis != nil { @@ -89,17 +97,20 @@ func (m *ConfigManager) Shutdown(_ context.Context) error { func (m *ConfigManager) Reload(ctx context.Context) error { if m == nil || m.settings == nil { + m.markUntrustedIfNoActiveSnapshot() return errors.New("prompt audit setting repository unavailable") } values, err := m.settings.GetMultiple(ctx, []string{SettingKeyPromptAuditConfig, SettingKeyRiskControl}) if err != nil { m.recordLoadError(err) + m.markUntrustedIfNoActiveSnapshot() return err } m.observeExpectedState(values[SettingKeyPromptAuditConfig], values[SettingKeyRiskControl] == "true") storage, err := ParseStorageConfig(values[SettingKeyPromptAuditConfig]) if err != nil { m.recordLoadError(err) + m.markUntrustedIfNoActiveSnapshot() return err } m.expected.Store(storage.ConfigVersion) @@ -107,10 +118,13 @@ func (m *ConfigManager) Reload(ctx context.Context) error { active, err := ActiveFromStorage(storage, values[SettingKeyRiskControl] == "true", m.encryptor) if err != nil { m.recordLoadError(err) + // expectedBlocking may already require fail-closed via BlockingActivationDegraded. + m.markUntrustedIfNoActiveSnapshot() return err } now := m.clock.Now() m.snapshot.Store(&activeConfigSnapshot{storage: cloneStorageConfig(storage), active: cloneActiveConfig(active), loadedAt: now}) + m.configUntrusted.Store(false) m.clearLoadError() LogInfo(EventConfigLoaded, map[string]any{ "config_version": storage.ConfigVersion, "status": "loaded", @@ -130,7 +144,13 @@ func (m *ConfigManager) Active() (ActiveConfig, bool) { } func (m *ConfigManager) BlockingActivationDegraded() bool { - if m == nil || !m.expectedBlocking.Load() { + if m == nil { + return false + } + if m.configUntrusted.Load() { + return true + } + if !m.expectedBlocking.Load() { return false } active, ok := m.Active() @@ -153,6 +173,22 @@ func (m *ConfigManager) EffectiveMode() Mode { return active.EffectiveMode() } +func (m *ConfigManager) markConfigUntrusted() { + if m == nil { + return + } + m.configUntrusted.Store(true) +} + +func (m *ConfigManager) markUntrustedIfNoActiveSnapshot() { + if m == nil { + return + } + if _, ok := m.Active(); !ok { + m.markConfigUntrusted() + } +} + func (m *ConfigManager) Public() PublicConfig { if m == nil { return PublicFromStorage(DefaultStorageConfig(), false) @@ -378,6 +414,11 @@ func (m *ConfigManager) subscribeLoop(ctx context.Context) { } m.expected.Store(version) if err := m.Reload(ctx); err != nil { + // A newer published version failed to activate. Until reload + // succeeds, do not keep serving a potentially stale weaker mode. + if active, ok := m.Active(); !ok || active.ConfigVersion < version { + m.markConfigUntrusted() + } LogWarn(EventConfigReloadDegraded, map[string]any{ "config_version": version, "status": "degraded", "error_code": "config_invalidation_reload_failed", }) diff --git a/backend/internal/securityaudit/prompt_config_test.go b/backend/internal/securityaudit/prompt_config_test.go index 399f5d2500..ad25dac9e2 100644 --- a/backend/internal/securityaudit/prompt_config_test.go +++ b/backend/internal/securityaudit/prompt_config_test.go @@ -131,6 +131,22 @@ func TestConfigManagerStaleWeakerSnapshotFailsClosedWhenBlockingExpected(t *test require.Equal(t, ErrorCodeUnavailable, guardErr.Code) } +type errorSettingRepository struct{ staticSettingRepository } + +func (errorSettingRepository) GetMultiple(context.Context, []string) (map[string]string, error) { + return nil, errors.New("settings unavailable") +} + +func TestConfigManagerStartupLoadFailureFailsClosedWithoutSnapshot(t *testing.T) { + manager := NewConfigManager(nil, errorSettingRepository{}, nil, prefixEncryptor{}) + err := manager.Start(context.Background()) + require.Error(t, err) + require.True(t, manager.configUntrusted.Load()) + require.True(t, manager.BlockingActivationDegraded()) + require.Equal(t, ModeBlocking, manager.EffectiveMode()) + require.NoError(t, manager.Shutdown(context.Background())) +} + func TestParseLegacyConfigDefaultsMissingFieldsWithoutEnablingBlocking(t *testing.T) { storage, err := ParseStorageConfig(`{"enabled":false,"config_version":9}`) require.NoError(t, err) diff --git a/backend/internal/securityaudit/prompt_outbound_security.go b/backend/internal/securityaudit/prompt_outbound_security.go index 40df98ef0a..f1e3ac3fc5 100644 --- a/backend/internal/securityaudit/prompt_outbound_security.go +++ b/backend/internal/securityaudit/prompt_outbound_security.go @@ -164,11 +164,22 @@ func secureDialContext(dialer *net.Dialer, resolver DNSResolver, allowPrivate bo } var lastErr error for _, addr := range addresses { - if isBlockedAddress(addr) || (!allowPrivate && (addr.IsPrivate() || addr.IsLoopback())) { + if isBlockedAddress(addr) { lastErr = fmt.Errorf("prompt guard resolved address blocked") continue } - if !addr.IsGlobalUnicast() && !addr.IsPrivate() && !addr.IsLoopback() { + if allowPrivate { + // localhost / *.localhost may only resolve to loopback. A hosts or + // DNS mapping from localhost to RFC1918 must not become an SSRF pivot. + if !addr.IsLoopback() { + lastErr = fmt.Errorf("prompt guard resolved address blocked") + continue + } + } else if addr.IsPrivate() || addr.IsLoopback() { + lastErr = fmt.Errorf("prompt guard resolved address blocked") + continue + } + if !addr.IsGlobalUnicast() && !addr.IsLoopback() { lastErr = fmt.Errorf("prompt guard resolved address blocked") continue } diff --git a/backend/internal/securityaudit/prompt_outbound_security_test.go b/backend/internal/securityaudit/prompt_outbound_security_test.go index 15e56d0bc5..76d7f9fdb1 100644 --- a/backend/internal/securityaudit/prompt_outbound_security_test.go +++ b/backend/internal/securityaudit/prompt_outbound_security_test.go @@ -49,6 +49,12 @@ func TestSecureDialRejectsDNSRebindingToPrivateAddress(t *testing.T) { require.Error(t, err) } +func TestSecureDialLocalhostAllowlistRejectsRFC1918Resolution(t *testing.T) { + dial := secureDialContext(nil, staticResolver{addresses: []netip.Addr{netip.MustParseAddr("10.0.0.8")}}, true) + _, err := dial(context.Background(), "tcp", "localhost:8080") + require.Error(t, err) +} + func TestSecureHTTPClientDoesNotBypassDestinationValidationThroughEnvironmentProxy(t *testing.T) { client, err := NewSecureHTTPClient(ActiveEndpoint{BaseURL: "https://guard.example.com", TimeoutMS: 1000}) require.NoError(t, err) @@ -207,6 +213,28 @@ func TestPromptAuditProbeModelsFallbackAndResponseSafety(t *testing.T) { }) } +func TestResolveProbeEndpointReusesTokenOnlyForMatchingBaseURL(t *testing.T) { + manager := &ConfigManager{} + manager.snapshot.Store(&activeConfigSnapshot{active: ActiveConfig{Endpoints: []ActiveEndpoint{{ + ID: "guard-1", BaseURL: "https://guard.example.com", Token: "STORED_GUARD_TOKEN", TimeoutMS: 1000, InputLimit: 1024, Enabled: true, + }}}}) + service := &PromptService{config: manager} + + matched, applied, err := service.resolveProbeEndpoint(UpdateEndpoint{ + ID: "guard-1", BaseURL: "https://guard.example.com/v1", TimeoutMS: 1000, InputLimit: 1024, + }) + require.NoError(t, err) + require.True(t, applied) + require.Equal(t, "STORED_GUARD_TOKEN", matched.Token) + + mismatched, applied, err := service.resolveProbeEndpoint(UpdateEndpoint{ + ID: "guard-1", BaseURL: "https://attacker.example.com", TimeoutMS: 1000, InputLimit: 1024, + }) + require.NoError(t, err) + require.False(t, applied) + require.Empty(t, mismatched.Token) +} + func newProbeTestService() *PromptService { return &PromptService{ config: &ConfigManager{}, scanner: NewOpenAICompatibleScanner(), clock: realClock{}, diff --git a/backend/internal/securityaudit/prompt_service.go b/backend/internal/securityaudit/prompt_service.go index 8402e60d99..71d0fadbe5 100644 --- a/backend/internal/securityaudit/prompt_service.go +++ b/backend/internal/securityaudit/prompt_service.go @@ -312,21 +312,27 @@ func modelsResponseReady(body []byte, model string) bool { } func (s *PromptService) resolveProbeEndpoint(input UpdateEndpoint) (ActiveEndpoint, bool, error) { + baseURL, err := NormalizeBaseURL(input.BaseURL) + if err != nil { + return ActiveEndpoint{}, false, err + } token := strings.TrimSpace(input.Token) if token == "" { if cfg, ok := s.config.Active(); ok { for _, endpoint := range cfg.Endpoints { - if endpoint.ID == strings.TrimSpace(input.ID) { - token = endpoint.Token - break + if endpoint.ID != strings.TrimSpace(input.ID) { + continue } + // Reuse a stored credential only when the probe targets the same + // normalized base URL. Otherwise an admin probe could exfiltrate + // the Guard token to an attacker-controlled HTTPS host. + if endpoint.BaseURL == baseURL { + token = endpoint.Token + } + break } } } - baseURL, err := NormalizeBaseURL(input.BaseURL) - if err != nil { - return ActiveEndpoint{}, false, err - } model := strings.TrimSpace(input.Model) if model == "" { model = DefaultGuardModel diff --git a/backend/internal/securityaudit/prompt_snapshot.go b/backend/internal/securityaudit/prompt_snapshot.go index 13002ded14..6b87ad1865 100644 --- a/backend/internal/securityaudit/prompt_snapshot.go +++ b/backend/internal/securityaudit/prompt_snapshot.go @@ -92,7 +92,10 @@ func extractProtocolSegments(protocol string, document any) []string { } } -var clientInstructionRoles = []string{"user", "system", "developer"} +// clientInstructionRoles are roles a client may freely populate. Attackers can +// place jailbreak/PII text in assistant/tool turns, so blocking audit must scan +// them too—not only user/system/developer instructions. +var clientInstructionRoles = []string{"user", "system", "developer", "assistant", "tool"} func extractChatLikeSegments(root map[string]any) []string { if root == nil { @@ -168,7 +171,7 @@ func extractResponses(value any) []string { result = append(result, entry) case map[string]any: role := strings.ToLower(stringValue(entry["role"])) - if role != "" && role != "user" && role != "system" && role != "developer" { + if role != "" && !isClientInstructionRole(role) { continue } if content, exists := entry["content"]; exists { @@ -183,7 +186,7 @@ func extractResponses(value any) []string { return result case map[string]any: role := strings.ToLower(stringValue(typed["role"])) - if role != "" && role != "user" && role != "system" && role != "developer" { + if role != "" && !isClientInstructionRole(role) { return nil } return contentTexts(typed["content"]) @@ -192,6 +195,15 @@ func extractResponses(value any) []string { } } +func isClientInstructionRole(role string) bool { + switch strings.ToLower(strings.TrimSpace(role)) { + case "user", "system", "developer", "assistant", "tool", "model": + return true + default: + return false + } +} + func extractGemini(value any) []string { var contents []any switch typed := value.(type) { @@ -209,7 +221,7 @@ func extractGemini(value any) []string { continue } role := strings.ToLower(stringValue(content["role"])) - if role != "" && role != "user" { + if role != "" && !isClientInstructionRole(role) { continue } parts, _ := content["parts"].([]any) diff --git a/backend/internal/securityaudit/prompt_snapshot_test.go b/backend/internal/securityaudit/prompt_snapshot_test.go index 17ab363711..a1427ba1aa 100644 --- a/backend/internal/securityaudit/prompt_snapshot_test.go +++ b/backend/internal/securityaudit/prompt_snapshot_test.go @@ -17,7 +17,7 @@ func TestExtractPromptSnapshotProtocols(t *testing.T) { protocol, body, first string count int }{ - {"openai_chat_completions", `{"messages":[{"role":"user","content":"old"},{"role":"assistant","content":"ignore"},{"role":"user","content":[{"type":"text","text":"最新😀"}]}]}`, "最新😀", 2}, + {"openai_chat_completions", `{"messages":[{"role":"user","content":"old"},{"role":"assistant","content":"assistant turn"},{"role":"user","content":[{"type":"text","text":"最新😀"}]}]}`, "最新😀", 3}, {"openai_responses", `{"input":[{"role":"user","content":[{"type":"input_text","text":"response text"}]}]}`, "response text", 1}, {"anthropic_messages", `{"messages":[{"role":"user","content":[{"type":"text","text":"claude"}]}]}`, "claude", 1}, {"gemini", `{"contents":[{"role":"user","parts":[{"text":"gemini"},{"inline_data":{"data":"BASE64"}}]}]}`, "gemini", 1}, @@ -67,8 +67,8 @@ func TestPromptSnapshotLatestUserMessageIsOnePrioritizedSegment(t *testing.T) { body := []byte(`{ "messages":[ {"role":"user","content":"历史输入"}, - {"role":"assistant","content":"assistant output must be ignored"}, - {"role":"tool","content":"tool output must be ignored"}, + {"role":"assistant","content":"assistant client injection"}, + {"role":"tool","content":"tool client injection"}, {"role":"user","content":[ {"type":"text","text":"最新第一块😀"}, {"type":"image_url","image_url":{"url":"data:image/png;base64,IMAGE_CANARY_BASE64"}}, @@ -78,10 +78,11 @@ func TestPromptSnapshotLatestUserMessageIsOnePrioritizedSegment(t *testing.T) { }`) snapshot, err := ExtractPromptSnapshot(Request{Protocol: "openai_chat_completions", Body: body}) require.NoError(t, err) - require.Equal(t, 2, snapshot.MessageCount) - require.Equal(t, "最新第一块😀\n最新第二块é\n\n历史输入", snapshot.ScanText) - require.NotContains(t, snapshot.ScanText, "assistant output") - require.NotContains(t, snapshot.ScanText, "tool output") + require.Equal(t, 4, snapshot.MessageCount) + require.True(t, strings.HasPrefix(snapshot.ScanText, "最新第一块😀\n最新第二块é")) + require.Contains(t, snapshot.ScanText, "历史输入") + require.Contains(t, snapshot.ScanText, "assistant client injection") + require.Contains(t, snapshot.ScanText, "tool client injection") require.NotContains(t, snapshot.ScanText, "IMAGE_CANARY_BASE64") require.Equal(t, utf8.RuneCountInString(snapshot.ScanText), snapshot.PromptLength) } @@ -93,7 +94,7 @@ func TestPromptSnapshotResponsesShapes(t *testing.T) { want string }{ {name: "string", body: `{"input":"plain response input"}`, want: "plain response input"}, - {name: "message array", body: `{"input":[{"role":"assistant","content":"ignore"},{"role":"user","content":[{"type":"input_text","text":"message block"}]}]}`, want: "message block"}, + {name: "message array", body: `{"input":[{"role":"assistant","content":"assistant turn"},{"role":"user","content":[{"type":"input_text","text":"message block"}]}]}`, want: "message block\n\nassistant turn"}, {name: "direct input text", body: `{"input":[{"type":"input_text","text":"direct block"}]}`, want: "direct block"}, {name: "single object", body: `{"input":{"role":"user","content":[{"type":"input_text","text":"single object"}]}}`, want: "single object"}, } @@ -122,7 +123,7 @@ func TestPromptSnapshotGeminiBatchShapesAndMediaExclusion(t *testing.T) { require.Contains(t, snapshot.ScanText, expected) } require.NotContains(t, snapshot.ScanText, "ROOT_BASE64") - require.NotContains(t, snapshot.ScanText, "ignore model") + require.Contains(t, snapshot.ScanText, "ignore model") } func TestPromptSnapshotMediaOnlyExtractsDeterministicTextPrompts(t *testing.T) { @@ -163,7 +164,7 @@ func TestResponsesWebSocketOnlyAuditsResponseCreateAndPreservesStage(t *testing. } func TestPromptSnapshotEmptyAndLongUnicodeInput(t *testing.T) { - _, err := ExtractPromptSnapshot(Request{Protocol: "openai_chat_completions", Body: []byte(`{"messages":[{"role":"assistant","content":"not user"},{"role":"user","content":" "}]}`)}) + _, err := ExtractPromptSnapshot(Request{Protocol: "openai_chat_completions", Body: []byte(`{"messages":[{"role":"function","content":"not audited role"},{"role":"user","content":" "}]}`)}) require.True(t, errors.Is(err, ErrNoPromptText)) latest := strings.Repeat("最新😀é", 80) @@ -186,10 +187,10 @@ func TestPromptSnapshotIncludesClientControlledInstructions(t *testing.T) { want []string }{ { - name: "openai system and developer", + name: "openai system developer assistant tool", protocol: "openai_chat_completions", - body: `{"messages":[{"role":"system","content":"system jailbreak"},{"role":"developer","content":"developer policy"},{"role":"assistant","content":"ignore"},{"role":"user","content":"hello"}]}`, - want: []string{"system jailbreak", "developer policy", "hello"}, + body: `{"messages":[{"role":"system","content":"system jailbreak"},{"role":"developer","content":"developer policy"},{"role":"assistant","content":"assistant jailbreak"},{"role":"tool","content":"tool payload"},{"role":"user","content":"hello"}]}`, + want: []string{"system jailbreak", "developer policy", "assistant jailbreak", "tool payload", "hello"}, }, { name: "openai system only", @@ -223,7 +224,6 @@ func TestPromptSnapshotIncludesClientControlledInstructions(t *testing.T) { for _, expected := range tt.want { require.Contains(t, snapshot.ScanText, expected) } - require.NotContains(t, snapshot.ScanText, "ignore") }) } } diff --git a/backend/internal/securityaudit/prompt_worker_test.go b/backend/internal/securityaudit/prompt_worker_test.go index 6c7ed7f3b9..5c39be8e44 100644 --- a/backend/internal/securityaudit/prompt_worker_test.go +++ b/backend/internal/securityaudit/prompt_worker_test.go @@ -302,7 +302,7 @@ func TestEnqueuerSkipsOffOutOfScopeAndNoText(t *testing.T) { cfg.GroupIDs = []int64{9} return cfg }(), req: asyncRequest()}, - {name: "no user text", cfg: asyncConfig(), req: Request{Protocol: "openai_chat_completions", Body: []byte(`{"messages":[{"role":"assistant","content":"ignore"}]}`)}}, + {name: "no user text", cfg: asyncConfig(), req: Request{Protocol: "openai_chat_completions", Body: []byte(`{"messages":[{"role":"function","content":"not audited"}]}`)}}, } for _, tt := range tests { t.Run(tt.name, func(t *testing.T) {