fix(security-audit): isolate latest user prompt chunk

This commit is contained in:
mt21625457
2026-07-17 12:12:41 +08:00
parent 18e698bed6
commit 7ed4e7e5e3
4 changed files with 176 additions and 70 deletions
@@ -3,6 +3,7 @@ package securityaudit
import (
"context"
"errors"
"strings"
"sync"
"testing"
"time"
@@ -143,6 +144,24 @@ func TestGuardEvaluatorLastChunkFailureNeverAllows(t *testing.T) {
require.Error(t, err)
}
func TestGuardEvaluatorScansLatestUserPromptAsIndependentFirstChunk(t *testing.T) {
latest := "请帮我编写一篇黄色小说 名字你来取"
history := strings.Repeat("# AGENTS.md instructions 项目安全规则。", 30)
seen := make([]string, 0, 4)
scanner := PromptScannerFunc(func(_ context.Context, _ ActiveEndpoint, prompt string, _ []string) (*NormalizedResult, error) {
seen = append(seen, prompt)
return &NormalizedResult{Decision: EventPass, RiskLevel: RiskLow, Action: ActionAllow, ScannerScores: map[string]float64{}, ScannerEvidence: map[string]string{}}, nil
})
evaluator := newGuardEvaluator(scanner, nil, NewAtomicMetrics(), 2, 2)
_, err := evaluator.Evaluate(context.Background(), guardConfig(
ActiveEndpoint{ID: "one", Enabled: true, TimeoutMS: 1000, InputLimit: 128},
), PromptSnapshot{ScanText: latest + promptAuditPrioritySeparator + history, PromptLength: len([]rune(latest + history))})
require.NoError(t, err)
require.Greater(t, len(seen), 1)
require.Equal(t, latest, seen[0])
require.Equal(t, history, strings.Join(seen[1:], ""))
}
func TestGuardEvaluatorBlockStopsRemainingChunksButReportsPlannedTotal(t *testing.T) {
calls := 0
scanner := PromptScannerFunc(func(context.Context, ActiveEndpoint, string, []string) (*NormalizedResult, error) {
@@ -3,6 +3,7 @@ package securityaudit
import (
"errors"
"sort"
"strings"
"time"
)
@@ -10,17 +11,17 @@ func SplitRunes(value string, limit int) []string {
if limit <= 0 {
return nil
}
runes := []rune(value)
if len(runes) == 0 {
return nil
}
chunks := make([]string, 0, (len(runes)+limit-1)/limit)
for start := 0; start < len(runes); start += limit {
end := start + limit
if end > len(runes) {
end = len(runes)
segments := strings.Split(value, promptAuditPrioritySeparator)
chunks := make([]string, 0, len(segments))
for _, segment := range segments {
runes := []rune(segment)
for start := 0; start < len(runes); start += limit {
end := start + limit
if end > len(runes) {
end = len(runes)
}
chunks = append(chunks, string(runes[start:end]))
}
chunks = append(chunks, string(runes[start:end]))
}
return chunks
}
@@ -21,18 +21,25 @@ var (
phonePattern = regexp.MustCompile(`(?:\+?\d[\d\s().-]{8,}\d)`)
)
const promptAuditPrioritySeparator = "\x00SUB2API_PROMPT_AUDIT_PRIORITY_END\x00"
type promptSegment struct {
text string
user bool
}
func ExtractPromptSnapshot(req Request) (PromptSnapshot, error) {
var document any
if err := json.Unmarshal(req.Body, &document); err != nil {
return PromptSnapshot{}, errors.New("prompt audit request JSON is invalid")
}
segments := extractProtocolSegments(req.Protocol, document)
segments = normalizeSegmentsLatestFirst(segments)
extracted := extractProtocolSegments(req.Protocol, document)
segments := normalizeSegmentsLatestUserFirst(extracted)
if len(segments) == 0 {
return PromptSnapshot{}, ErrNoPromptText
}
scanText := strings.Join(segments, "\n\n")
digest := sha256.Sum256([]byte(scanText))
scanText, metadataText := buildPrioritizedScanText(segments)
digest := sha256.Sum256([]byte(metadataText))
stage := strings.TrimSpace(req.Stage)
if stage == "" {
stage = "http"
@@ -42,8 +49,8 @@ func ExtractPromptSnapshot(req Request) (PromptSnapshot, error) {
UserEmailSnapshot: req.UserEmail, APIKeyID: req.APIKeyID, APIKeyNameSnapshot: req.APIKeyName,
GroupID: cloneInt64Ptr(req.GroupID), GroupName: req.GroupName, Provider: req.Provider,
Endpoint: req.Endpoint, Protocol: req.Protocol, Model: req.Model,
PromptHash: hex.EncodeToString(digest[:]), RedactedPreview: BuildPromptPreview(scanText, DefaultPromptPreviewMaxRunes),
PromptLength: utf8.RuneCountInString(scanText), MessageCount: len(segments), Stage: stage,
PromptHash: hex.EncodeToString(digest[:]), RedactedPreview: BuildPromptPreview(metadataText, DefaultPromptPreviewMaxRunes),
PromptLength: utf8.RuneCountInString(metadataText), MessageCount: len(segments), Stage: stage,
ScanText: scanText,
}, nil
}
@@ -52,7 +59,7 @@ func ExtractPromptSnapshot(req Request) (PromptSnapshot, error) {
// considered before BuildPromptPreview withholds the majority for storage/UI.
const DefaultPromptPreviewMaxRunes = 96
func extractProtocolSegments(protocol string, document any) []string {
func extractProtocolSegments(protocol string, document any) []promptSegment {
root, _ := document.(map[string]any)
protocol = strings.ToLower(strings.TrimSpace(protocol))
switch protocol {
@@ -77,7 +84,7 @@ func extractProtocolSegments(protocol string, document any) []string {
}
return append(extractInstructions(root["instructions"]), extractResponses(root["input"])...)
case "openai_images", "grok_media", "media", "images":
return extractMediaPrompts(root)
return userPromptSegments(extractMediaPrompts(root))
default:
if segments := extractChatLikeSegments(root); len(segments) > 0 {
return segments
@@ -88,7 +95,7 @@ func extractProtocolSegments(protocol string, document any) []string {
if gemini := extractGeminiRoot(root); len(gemini) > 0 {
return gemini
}
return extractMediaPrompts(root)
return userPromptSegments(extractMediaPrompts(root))
}
}
@@ -97,14 +104,14 @@ func extractProtocolSegments(protocol string, document any) []string {
// them too—not only user/system/developer instructions.
var clientInstructionRoles = []string{"user", "system", "developer", "assistant", "tool"}
func extractChatLikeSegments(root map[string]any) []string {
func extractChatLikeSegments(root map[string]any) []promptSegment {
if root == nil {
return nil
}
return extractMessages(root["messages"], clientInstructionRoles...)
}
func extractMessages(value any, wantedRoles ...string) []string {
func extractMessages(value any, wantedRoles ...string) []promptSegment {
items, ok := value.([]any)
if !ok {
return nil
@@ -113,7 +120,7 @@ func extractMessages(value any, wantedRoles ...string) []string {
for _, role := range wantedRoles {
wanted[strings.ToLower(strings.TrimSpace(role))] = struct{}{}
}
result := make([]string, 0, len(items))
result := make([]promptSegment, 0, len(items))
for _, item := range items {
message, ok := item.(map[string]any)
if !ok {
@@ -124,62 +131,62 @@ func extractMessages(value any, wantedRoles ...string) []string {
continue
}
texts := contentTexts(message["content"])
if len(texts) > 0 {
result = append(result, strings.Join(texts, "\n"))
for _, text := range texts {
result = append(result, promptSegment{text: text, user: role == "user"})
}
}
return result
}
func extractInstructions(value any) []string {
func extractInstructions(value any) []promptSegment {
switch typed := value.(type) {
case string:
if text := strings.TrimSpace(typed); text != "" {
return []string{text}
return []promptSegment{{text: text}}
}
case []any:
return contentTexts(typed)
return systemPromptSegments(contentTexts(typed))
case map[string]any:
return contentTexts(typed)
return systemPromptSegments(contentTexts(typed))
}
return nil
}
func extractAnthropicSystem(value any) []string {
func extractAnthropicSystem(value any) []promptSegment {
switch typed := value.(type) {
case string:
if text := strings.TrimSpace(typed); text != "" {
return []string{text}
return []promptSegment{{text: text}}
}
case []any:
return contentTexts(typed)
return systemPromptSegments(contentTexts(typed))
case map[string]any:
return contentTexts(typed)
return systemPromptSegments(contentTexts(typed))
}
return nil
}
func extractResponses(value any) []string {
func extractResponses(value any) []promptSegment {
switch typed := value.(type) {
case string:
return []string{typed}
return []promptSegment{{text: typed, user: true}}
case []any:
result := make([]string, 0, len(typed))
result := make([]promptSegment, 0, len(typed))
for _, item := range typed {
switch entry := item.(type) {
case string:
result = append(result, entry)
result = append(result, promptSegment{text: entry, user: true})
case map[string]any:
role := strings.ToLower(stringValue(entry["role"]))
if role != "" && !isClientInstructionRole(role) {
continue
}
if content, exists := entry["content"]; exists {
if texts := contentTexts(content); len(texts) > 0 {
result = append(result, strings.Join(texts, "\n"))
for _, text := range contentTexts(content) {
result = append(result, promptSegment{text: text, user: role == "" || role == "user"})
}
} else if text := stringValue(entry["text"]); text != "" {
result = append(result, text)
result = append(result, promptSegment{text: text, user: role == "" || role == "user"})
}
}
}
@@ -189,7 +196,7 @@ func extractResponses(value any) []string {
if role != "" && !isClientInstructionRole(role) {
return nil
}
return contentTexts(typed["content"])
return promptSegmentsForRole(contentTexts(typed["content"]), role)
default:
return nil
}
@@ -204,7 +211,7 @@ func isClientInstructionRole(role string) bool {
}
}
func extractGemini(value any) []string {
func extractGemini(value any) []promptSegment {
var contents []any
switch typed := value.(type) {
case []any:
@@ -214,7 +221,7 @@ func extractGemini(value any) []string {
default:
return nil
}
result := make([]string, 0, len(contents))
result := make([]promptSegment, 0, len(contents))
for _, item := range contents {
content, ok := item.(map[string]any)
if !ok {
@@ -228,7 +235,7 @@ func extractGemini(value any) []string {
for _, part := range parts {
if object, ok := part.(map[string]any); ok {
if text := stringValue(object["text"]); text != "" {
result = append(result, text)
result = append(result, promptSegment{text: text, user: role == "" || role == "user"})
}
}
}
@@ -236,7 +243,7 @@ func extractGemini(value any) []string {
return result
}
func extractGeminiRoot(root map[string]any) []string {
func extractGeminiRoot(root map[string]any) []promptSegment {
if root == nil {
return nil
}
@@ -261,41 +268,45 @@ func extractGeminiRoot(root map[string]any) []string {
return result
}
func extractGeminiSystemInstruction(value any) []string {
func extractGeminiSystemInstruction(value any) []promptSegment {
switch typed := value.(type) {
case string:
if text := strings.TrimSpace(typed); text != "" {
return []string{text}
return []promptSegment{{text: text}}
}
case map[string]any:
if parts, ok := typed["parts"].([]any); ok {
result := make([]string, 0, len(parts))
result := make([]promptSegment, 0, len(parts))
for _, part := range parts {
if object, ok := part.(map[string]any); ok {
if text := stringValue(object["text"]); text != "" {
result = append(result, text)
result = append(result, promptSegment{text: text})
}
}
}
return result
}
return contentTexts(typed)
return systemPromptSegments(contentTexts(typed))
case []any:
return extractGemini(typed)
segments := extractGemini(typed)
for index := range segments {
segments[index].user = false
}
return segments
}
return nil
}
func extractGeminiInstances(value any) []string {
func extractGeminiInstances(value any) []promptSegment {
instances, ok := value.([]any)
if !ok {
return nil
}
result := make([]string, 0, len(instances))
result := make([]promptSegment, 0, len(instances))
for _, item := range instances {
if instance, ok := item.(map[string]any); ok {
if prompt := stringValue(instance["prompt"]); prompt != "" {
result = append(result, prompt)
result = append(result, promptSegment{text: prompt, user: true})
}
}
}
@@ -402,24 +413,58 @@ func contentTexts(value any) []string {
return nil
}
func normalizeSegmentsLatestFirst(values []string) []string {
normalized := make([]string, 0, len(values))
func normalizeSegmentsLatestUserFirst(values []promptSegment) []string {
normalized := make([]promptSegment, 0, len(values))
for _, value := range values {
value = strings.TrimSpace(value)
if value != "" {
value.text = strings.TrimSpace(value.text)
if value.text != "" {
normalized = append(normalized, value)
}
}
if len(normalized) <= 1 {
return normalized
if len(normalized) == 0 {
return nil
}
priorityIndex := len(normalized) - 1
for index := len(normalized) - 1; index >= 0; index-- {
if normalized[index].user {
priorityIndex = index
break
}
}
latest := normalized[len(normalized)-1]
result := make([]string, 0, len(normalized))
result = append(result, latest)
result = append(result, normalized[:len(normalized)-1]...)
result = append(result, normalized[priorityIndex].text)
for index, segment := range normalized {
if index != priorityIndex {
result = append(result, segment.text)
}
}
return result
}
func buildPrioritizedScanText(segments []string) (scanText string, metadataText string) {
metadataText = strings.Join(segments, "\n\n")
if len(segments) <= 1 {
return metadataText, metadataText
}
return segments[0] + promptAuditPrioritySeparator + strings.Join(segments[1:], "\n\n"), metadataText
}
func promptSegmentsForRole(texts []string, role string) []promptSegment {
result := make([]promptSegment, 0, len(texts))
for _, text := range texts {
result = append(result, promptSegment{text: text, user: role == "" || role == "user"})
}
return result
}
func userPromptSegments(texts []string) []promptSegment {
return promptSegmentsForRole(texts, "user")
}
func systemPromptSegments(texts []string) []promptSegment {
return promptSegmentsForRole(texts, "system")
}
func RedactPreview(value string, maxRunes int) string {
value = bearerPattern.ReplaceAllString(value, "Bearer ***")
value = apiKeyPattern.ReplaceAllStringFunc(value, func(match string) string {
@@ -30,7 +30,7 @@ func TestExtractPromptSnapshotProtocols(t *testing.T) {
require.NoError(t, err)
require.True(t, strings.HasPrefix(snapshot.ScanText, tt.first))
require.Equal(t, tt.count, snapshot.MessageCount)
require.Equal(t, utf8.RuneCountInString(snapshot.ScanText), snapshot.PromptLength)
require.Equal(t, utf8.RuneCountInString(metadataTextForTest(snapshot.ScanText)), snapshot.PromptLength)
require.NotEmpty(t, snapshot.PromptHash)
require.NotContains(t, snapshot.ScanText, "BASE64SECRET")
})
@@ -49,7 +49,7 @@ func TestSnapshotRedactsCanariesAndPreservesHashOfScanText(t *testing.T) {
require.NotContains(t, snapshot.RedactedPreview, "138 0013 8000")
require.Contains(t, snapshot.ScanText, "PROMPT_CANARY_ABC123")
require.NotEqual(t, snapshot.ScanText, snapshot.RedactedPreview)
digest := sha256.Sum256([]byte(snapshot.ScanText))
digest := sha256.Sum256([]byte(metadataTextForTest(snapshot.ScanText)))
require.Equal(t, hex.EncodeToString(digest[:]), snapshot.PromptHash)
require.Empty(t, snapshot.Redacted().ScanText)
}
@@ -63,7 +63,19 @@ func TestSplitRunesDoesNotSplitUTF8(t *testing.T) {
require.Equal(t, "中文😀éabc", strings.Join(chunks, ""))
}
func TestPromptSnapshotLatestUserMessageIsOnePrioritizedSegment(t *testing.T) {
func TestSplitRunesKeepsPrioritySegmentIndependent(t *testing.T) {
latest := "请帮我编写一篇黄色小说 名字你来取"
history := strings.Repeat("AGENTS.md 项目约束。", 40)
chunks := SplitRunes(latest+promptAuditPrioritySeparator+history, 128)
require.Greater(t, len(chunks), 2)
require.Equal(t, latest, chunks[0])
require.Equal(t, history, strings.Join(chunks[1:], ""))
for _, chunk := range chunks {
require.NotContains(t, chunk, promptAuditPrioritySeparator)
}
}
func TestPromptSnapshotLatestUserTextBlockIsOnePrioritizedSegment(t *testing.T) {
body := []byte(`{
"messages":[
{"role":"user","content":"历史输入"},
@@ -78,13 +90,37 @@ func TestPromptSnapshotLatestUserMessageIsOnePrioritizedSegment(t *testing.T) {
}`)
snapshot, err := ExtractPromptSnapshot(Request{Protocol: "openai_chat_completions", Body: body})
require.NoError(t, err)
require.Equal(t, 4, snapshot.MessageCount)
require.True(t, strings.HasPrefix(snapshot.ScanText, "最新第一块😀\n最新第二块é"))
require.Equal(t, 5, snapshot.MessageCount)
require.True(t, strings.HasPrefix(snapshot.ScanText, "最新第二块é"+promptAuditPrioritySeparator))
require.Contains(t, snapshot.ScanText, "最新第一块😀")
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)
require.Equal(t, utf8.RuneCountInString(metadataTextForTest(snapshot.ScanText)), snapshot.PromptLength)
}
func TestPromptSnapshotSeparatesAnthropicUserPromptFromHarnessBlocks(t *testing.T) {
latest := "请帮我编写一篇黄色小说 名字你来取"
agents := "# AGENTS.md instructions\n<INSTRUCTIONS>" + strings.Repeat("安全约束。", 80) + "</INSTRUCTIONS>"
environment := "<environment_context><cwd>/workspace</cwd></environment_context>"
body := []byte(`{"system":"system policy","messages":[{"role":"user","content":[` +
`{"type":"text","text":` + string(mustJSON(t, agents)) + `},` +
`{"type":"text","text":` + string(mustJSON(t, environment)) + `},` +
`{"type":"text","text":` + string(mustJSON(t, latest)) + `}` +
`]}]}`)
snapshot, err := ExtractPromptSnapshot(Request{Protocol: "anthropic_messages", Body: body})
require.NoError(t, err)
require.Equal(t, 4, snapshot.MessageCount)
require.True(t, strings.HasPrefix(snapshot.ScanText, latest+promptAuditPrioritySeparator))
require.True(t, strings.HasPrefix(snapshot.RedactedPreview, "请帮我编写一篇黄色小说"))
chunks := SplitRunes(snapshot.ScanText, 128)
require.Equal(t, latest, chunks[0])
require.Contains(t, strings.Join(chunks[1:], ""), "# AGENTS.md instructions")
require.Contains(t, strings.Join(chunks[1:], ""), "<environment_context>")
require.NotContains(t, strings.Join(chunks, ""), promptAuditPrioritySeparator)
}
func TestPromptSnapshotResponsesShapes(t *testing.T) {
@@ -102,7 +138,7 @@ func TestPromptSnapshotResponsesShapes(t *testing.T) {
t.Run(tt.name, func(t *testing.T) {
snapshot, err := ExtractPromptSnapshot(Request{Protocol: "openai_responses", Body: []byte(tt.body)})
require.NoError(t, err)
require.Equal(t, tt.want, snapshot.ScanText)
require.Equal(t, tt.want, metadataTextForTest(snapshot.ScanText))
})
}
}
@@ -174,7 +210,8 @@ func TestPromptSnapshotEmptyAndLongUnicodeInput(t *testing.T) {
require.NoError(t, err)
require.True(t, strings.HasPrefix(snapshot.ScanText, latest))
chunks := SplitRunes(snapshot.ScanText, 127)
require.Equal(t, snapshot.ScanText, strings.Join(chunks, ""))
require.Equal(t, strings.Replace(snapshot.ScanText, promptAuditPrioritySeparator, "", 1), strings.Join(chunks, ""))
require.Equal(t, latest, chunks[0]+strings.Join(chunks[1:len(SplitRunes(latest, 127))], ""))
for _, chunk := range chunks {
require.LessOrEqual(t, len([]rune(chunk)), 127)
require.True(t, utf8.ValidString(chunk))
@@ -252,3 +289,7 @@ func mustJSON(t *testing.T, value string) []byte {
require.NoError(t, err)
return raw
}
func metadataTextForTest(scanText string) string {
return strings.Replace(scanText, promptAuditPrioritySeparator, "\n\n", 1)
}