diff --git a/pkg/agent/guardrail/prompt_injection.go b/pkg/agent/guardrail/prompt_injection.go index 1ed0a0dc9..787f1c6d4 100644 --- a/pkg/agent/guardrail/prompt_injection.go +++ b/pkg/agent/guardrail/prompt_injection.go @@ -59,20 +59,20 @@ func (g *PromptInjectionGuardrail) Check(ctx context.Context, messages []llm.Mes resp, err := g.client.ChatCompletion( ctx, &llm.ChatCompletionRequest{ - Model: "gpt-4o-mini", - Messages: []llm.Message{ - { - Role: llm.RoleSystem, - Parts: []llm.Part{llm.TextPart{Text: promptInjectionClassifierPrompt}}, - }, - { - Role: llm.RoleUser, - Parts: []llm.Part{llm.TextPart{Text: userText}}, + Model: "gpt-4o-mini", + Messages: []llm.Message{ + { + Role: llm.RoleSystem, + Parts: []llm.Part{llm.TextPart{Text: promptInjectionClassifierPrompt}}, + }, + { + Role: llm.RoleUser, + Parts: []llm.Part{llm.TextPart{Text: userText}}, + }, }, + MaxTokens: new(10), + Temperature: new(0.0), }, - MaxTokens: new(10), - Temperature: new(0.0), - }, ) if err != nil { // If the classifier fails, allow the message through rather than diff --git a/pkg/agent/guardrail/sensitive_data.go b/pkg/agent/guardrail/sensitive_data.go index 96ecd778f..b3f4cbaaa 100644 --- a/pkg/agent/guardrail/sensitive_data.go +++ b/pkg/agent/guardrail/sensitive_data.go @@ -58,7 +58,8 @@ func (g *SensitiveDataGuardrail) Check(_ context.Context, message llm.Message) ( "sk_test_", // LLM provider keys - "sk-", + "sk-proj-", // OpenAI + "sk-ant-", // Anthropic // JWT tokens "eyj", // base64-encoded JSON header diff --git a/pkg/agent/guardrail/sensitive_data_test.go b/pkg/agent/guardrail/sensitive_data_test.go index 080714c35..fdcd4a31a 100644 --- a/pkg/agent/guardrail/sensitive_data_test.go +++ b/pkg/agent/guardrail/sensitive_data_test.go @@ -66,6 +66,8 @@ func TestSensitiveDataGuardrail_Check(t *testing.T) { // LLM provider keys {"openai key", "The API key is sk-proj-abc123", true}, + {"anthropic key", "Key: sk-ant-api03-abc123", true}, + {"sk prefix not a false positive", "This is a risk-based approach to task-management.", false}, // JWT tokens {"jwt token", "Token: eyJhbGciOiJIUzI1NiJ9.payload.sig", true}, diff --git a/pkg/agent/guardrail/system_prompt_leak.go b/pkg/agent/guardrail/system_prompt_leak.go index 61e69295c..271234895 100644 --- a/pkg/agent/guardrail/system_prompt_leak.go +++ b/pkg/agent/guardrail/system_prompt_leak.go @@ -27,9 +27,12 @@ type SystemPromptLeakGuardrail struct { } func NewSystemPromptLeakGuardrail(fingerprints []string) *SystemPromptLeakGuardrail { - lowered := make([]string, len(fingerprints)) - for i, f := range fingerprints { - lowered[i] = strings.ToLower(f) + lowered := make([]string, 0, len(fingerprints)) + for _, f := range fingerprints { + if f == "" { + continue + } + lowered = append(lowered, strings.ToLower(f)) } return &SystemPromptLeakGuardrail{fingerprints: lowered} diff --git a/pkg/agent/guardrail/system_prompt_leak_test.go b/pkg/agent/guardrail/system_prompt_leak_test.go index 9967783be..ad2e762ae 100644 --- a/pkg/agent/guardrail/system_prompt_leak_test.go +++ b/pkg/agent/guardrail/system_prompt_leak_test.go @@ -66,4 +66,14 @@ func TestSystemPromptLeakGuardrail_Check(t *testing.T) { require.NoError(t, err) assert.False(t, result.Tripwire) }) + + t.Run("empty fingerprints are ignored", func(t *testing.T) { + t.Parallel() + + g := guardrail.NewSystemPromptLeakGuardrail([]string{"", "secret phrase", ""}) + result, err := g.Check(context.Background(), assistantMessage("hello world")) + + require.NoError(t, err) + assert.False(t, result.Tripwire) + }) }