From 7ed4e7e5e37c579f11d1b714d2945809dea1ea5b Mon Sep 17 00:00:00 2001 From: mt21625457 Date: Fri, 17 Jul 2026 12:12:41 +0800 Subject: [PATCH] fix(security-audit): isolate latest user prompt chunk --- .../securityaudit/prompt_guard_test.go | 19 +++ .../internal/securityaudit/prompt_scanner.go | 21 +-- .../internal/securityaudit/prompt_snapshot.go | 149 ++++++++++++------ .../securityaudit/prompt_snapshot_test.go | 57 ++++++- 4 files changed, 176 insertions(+), 70 deletions(-) diff --git a/backend/internal/securityaudit/prompt_guard_test.go b/backend/internal/securityaudit/prompt_guard_test.go index c48993681..76e9dea05 100644 --- a/backend/internal/securityaudit/prompt_guard_test.go +++ b/backend/internal/securityaudit/prompt_guard_test.go @@ -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) { diff --git a/backend/internal/securityaudit/prompt_scanner.go b/backend/internal/securityaudit/prompt_scanner.go index 26fca2f94..103d10045 100644 --- a/backend/internal/securityaudit/prompt_scanner.go +++ b/backend/internal/securityaudit/prompt_scanner.go @@ -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 } diff --git a/backend/internal/securityaudit/prompt_snapshot.go b/backend/internal/securityaudit/prompt_snapshot.go index 6b87ad186..237eedfdd 100644 --- a/backend/internal/securityaudit/prompt_snapshot.go +++ b/backend/internal/securityaudit/prompt_snapshot.go @@ -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 { diff --git a/backend/internal/securityaudit/prompt_snapshot_test.go b/backend/internal/securityaudit/prompt_snapshot_test.go index a1427ba1a..eb432f256 100644 --- a/backend/internal/securityaudit/prompt_snapshot_test.go +++ b/backend/internal/securityaudit/prompt_snapshot_test.go @@ -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" + strings.Repeat("安全约束。", 80) + "" + environment := "/workspace" + 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:], ""), "") + 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) +}