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) {