From 627775d5d385e7ec633b605cf10fcb4e8f2a8943 Mon Sep 17 00:00:00 2001 From: hu <187184415@qq.com> Date: Sun, 9 Aug 2026 00:35:51 +0800 Subject: [PATCH 1/3] fix(prompt-filter): sync and remove failed review keys --- frontend/src/locales/en.json | 4 +- frontend/src/locales/zh.json | 4 +- frontend/src/pages/PromptFilter.tsx | 86 ++++++++++++++++++++++++++--- 3 files changed, 85 insertions(+), 9 deletions(-) diff --git a/frontend/src/locales/en.json b/frontend/src/locales/en.json index 2b2a96df..ca7007e5 100644 --- a/frontend/src/locales/en.json +++ b/frontend/src/locales/en.json @@ -2402,11 +2402,13 @@ "reviewApiKeyConfigured": "{{n}} key(s) configured; leave blank to keep", "reviewApiKeyHint": "One key per line. Multiple keys are round-robin balanced and failed over on rate limits, invalid keys, or server errors.", "reviewKeyListTitle": "Saved review keys", - "reviewKeyListHint": "Only masked suffixes are shown. Connection test results map to each key so a failing key can be removed individually.", + "reviewKeyListHint": "Only masked saved keys are shown. The list refreshes after saving, and connection test status maps back to each saved key.", "reviewKeyDeleteTitle": "Delete this review key?", "reviewKeyDeleteConfirm": "This removes {{key}} without changing the other review keys. The last key cannot be deleted while model review is enabled.", "reviewKeyDeleteAria": "Delete review key {{key}}", "reviewKeyDeleted": "Deleted review key {{key}}", + "reviewKeyRemoveDraftAria": "Remove review key {{key}} from the pending list", + "reviewKeyRemovedFromDraft": "Removed {{key}} from the pending list; save settings to apply the change", "reviewRequestMode": "Request Adapter", "reviewModeChat": "Chat Completions compatible", "reviewModeModerations": "Moderations compatible", diff --git a/frontend/src/locales/zh.json b/frontend/src/locales/zh.json index 61695df6..8f306f69 100644 --- a/frontend/src/locales/zh.json +++ b/frontend/src/locales/zh.json @@ -2402,11 +2402,13 @@ "reviewApiKeyConfigured": "已配置 {{n}} 个 key,留空保持不变", "reviewApiKeyHint": "每行一个 key。配置多个可轮询分摊额度,遇到限流、无效 key 或服务端错误时自动切换。", "reviewKeyListTitle": "已保存的审查 Key", - "reviewKeyListHint": "仅显示脱敏尾号;连接测试结果会对应到具体 Key,可单独删除异常项。", + "reviewKeyListHint": "仅显示已保存 Key 的脱敏尾号;保存设置后列表会自动刷新,连接测试状态会对应到具体 Key。", "reviewKeyDeleteTitle": "删除这个审查 Key?", "reviewKeyDeleteConfirm": "将删除 {{key}}。其他审查 Key 不受影响;启用模型复核时不能删除最后一个 Key。", "reviewKeyDeleteAria": "删除审查 Key {{key}}", "reviewKeyDeleted": "已删除审查 Key {{key}}", + "reviewKeyRemoveDraftAria": "从待保存列表移除审查 Key {{key}}", + "reviewKeyRemovedFromDraft": "已从待保存列表移除 {{key}},请保存设置使修改生效", "reviewRequestMode": "请求适配模式", "reviewModeChat": "Chat Completions 兼容接口", "reviewModeModerations": "Moderations 兼容接口", diff --git a/frontend/src/pages/PromptFilter.tsx b/frontend/src/pages/PromptFilter.tsx index af77d652..d4df0c1d 100644 --- a/frontend/src/pages/PromptFilter.tsx +++ b/frontend/src/pages/PromptFilter.tsx @@ -16,7 +16,7 @@ import { formatBeijingTime, formatRelativeTime } from '../utils/time' import { getErrorMessage } from '../utils/error' import { getPromptFilterScoreBand, normalizePromptFilterScore } from '../lib/promptFilterScore' import { parseAdvancedConfigDocument, patchAdvancedConfigDocument, readAdvancedConfigPath } from '../types' -import type { AdvancedConfigObject, AdvancedConfigPatch, PromptFilterLog, PromptFilterMatch, PromptFilterRule, PromptFilterRulesResponse, PromptFilterTestResponse, PromptGuardConfig, PromptGuardLayer, PromptGuardMode, PromptGuardProfile, PromptGuardProvider, PromptIdentityUpdateMode, PromptIntelligenceAIAnalysisResponse, PromptIntelligenceAIProvider, PromptIntelligenceCandidate, PromptIntelligenceEvidenceResponse, PromptIntelligenceGatewayKey, PromptIntelligenceRun, PromptPolicyIncident, PromptPolicyIncidentDetailResponse, PromptReviewAPIKeyDescriptor, PromptReviewTestResponse, PromptRiskProfile, PromptRiskProfileDetailResponse, SystemSettings } from '../types' +import type { AdvancedConfigObject, AdvancedConfigPatch, PromptFilterLog, PromptFilterMatch, PromptFilterRule, PromptFilterRulesResponse, PromptFilterTestResponse, PromptGuardConfig, PromptGuardLayer, PromptGuardMode, PromptGuardProfile, PromptGuardProvider, PromptIdentityUpdateMode, PromptIntelligenceAIAnalysisResponse, PromptIntelligenceAIProvider, PromptIntelligenceCandidate, PromptIntelligenceEvidenceResponse, PromptIntelligenceGatewayKey, PromptIntelligenceRun, PromptPolicyIncident, PromptPolicyIncidentDetailResponse, PromptReviewAPIKeyDescriptor, PromptReviewKeyTestResult, PromptReviewTestResponse, PromptRiskProfile, PromptRiskProfileDetailResponse, SystemSettings } from '../types' import { Badge } from '@/components/ui/badge' import { Button } from '@/components/ui/button' import { Card, CardContent } from '@/components/ui/card' @@ -441,6 +441,18 @@ const defaultForm: PromptFilterForm = { prompt_filter_review_fail_closed: true, } +function parsePromptReviewAPIKeyInput(raw: string): string[] { + const seen = new Set() + return raw + .split(/[\s,;]+/) + .map((key) => key.trim()) + .filter((key) => { + if (!key || seen.has(key)) return false + seen.add(key) + return true + }) +} + const emptyFilters: LogFilters = { action: '', source: '', @@ -543,6 +555,7 @@ export default function PromptFilter() { const { toast, showToast } = useToast() const [form, setForm] = useState(defaultForm) const [saving, setSaving] = useState(false) + const [settingsSaveRevision, setSettingsSaveRevision] = useState(0) const advancedConfigError = useMemo( () => parseAdvancedConfigDocument(form.prompt_filter_advanced_config).error, [form.prompt_filter_advanced_config], @@ -622,6 +635,7 @@ export default function PromptFilter() { setForm(normalizePromptFilterForm(updated)) setData((current) => ({ ...current, settings: updated })) + setSettingsSaveRevision((revision) => revision + 1) setSaving(false) const quarantines = updated.prompt_filter_pattern_quarantines ?? [] if (quarantines.length > 0) { @@ -745,6 +759,7 @@ export default function PromptFilter() { testResult={testResult} runTest={runTest} advancedConfigError={advancedConfigError} + settingsSaveRevision={settingsSaveRevision} onSave={() => void saveSettings()} /> ) : null} @@ -2707,6 +2722,7 @@ function OverviewView({ testResult, runTest, advancedConfigError, + settingsSaveRevision, onSave, }: { form: PromptFilterForm @@ -2727,6 +2743,7 @@ function OverviewView({ testResult: PromptFilterTestResponse | null runTest: () => void advancedConfigError: string | null + settingsSaveRevision: number onSave: () => void }) { const { t } = useTranslation() @@ -2767,10 +2784,11 @@ function OverviewView({ const reviewStrategy = !form.prompt_filter_review_enabled ? 'off' : form.prompt_filter_review_fail_closed ? 'fail_closed' : 'fail_open' - const draftReviewKeyCount = (form.prompt_filter_review_api_key ?? '') - .split(/[\s,]+/) - .filter(Boolean).length - const reviewKeyCount = draftReviewKeyCount || form.prompt_filter_review_api_key_count + const draftReviewKeys = useMemo( + () => parsePromptReviewAPIKeyInput(form.prompt_filter_review_api_key ?? ''), + [form.prompt_filter_review_api_key], + ) + const reviewKeyCount = draftReviewKeys.length || form.prompt_filter_review_api_key_count const enabledAdvancedFeatures = [ advancedProtection.normalization.enabled ? t('promptFilter.enabledFeatures.normalization') : null, advancedProtection.context_discount.enabled ? t('promptFilter.enabledFeatures.contextDiscount') : null, @@ -2883,7 +2901,7 @@ function OverviewView({ if (!cancelled) setReviewKeysLoading(false) }) return () => { cancelled = true } - }, [reviewSettingsOpen, showToast]) + }, [reviewSettingsOpen, settingsSaveRevision, showToast]) const deleteReviewKey = async (keyID: string, masked: string) => { const approved = await confirm({ title: t('promptFilter.reviewKeyDeleteTitle'), @@ -2914,6 +2932,26 @@ function OverviewView({ setDeletingReviewKeyID(null) } } + const removeFailedReviewTestKey = async (item: PromptReviewKeyTestResult) => { + if (!item.key_id) return + const masked = item.key_masked || `Key #${item.key_index}` + if (draftReviewKeys.length > 0) { + const draftIndex = item.key_index - 1 + if (draftIndex < 0 || draftIndex >= draftReviewKeys.length) return + const remaining = draftReviewKeys.filter((_, index) => index !== draftIndex) + setForm((current) => ({ ...current, prompt_filter_review_api_key: remaining.join('\n') })) + setReviewTestResult((current) => current ? { + ...current, + key_count: remaining.length, + results: current.results?.filter((result) => result.key_id !== item.key_id), + } : null) + showToast(t('promptFilter.reviewKeyRemovedFromDraft', { key: masked })) + return + } + if (configuredReviewKeys.some((key) => key.id === item.key_id)) { + await deleteReviewKey(item.key_id, masked) + } + } const applyRecommendedProtection = () => { setForm((current) => { const patched = patchAdvancedConfigDocument(current.prompt_filter_advanced_config, [ @@ -3244,7 +3282,41 @@ function OverviewView({
{t('promptFilter.reviewModel')}: {reviewTestResult.model}
{reviewTestResult.highest_category ?
{t('promptFilter.reviewTestHighestCategory')}: {reviewTestResult.highest_category}
: null} {reviewTestResult.reason ?
{t('promptFilter.reviewTestReason')}: {reviewTestResult.reason}
: null} - {reviewTestResult.results?.length ?
{reviewTestResult.results.map((item) =>
{item.key_masked || `Key #${item.key_index}`}{item.ok ? t('common.success') : t('common.failed')}
{item.latency_ms} ms · {item.ok ? item.confidence.toFixed(2) : item.error || '-'}
)}
: null} + {reviewTestResult.results?.length ? ( +
+ {reviewTestResult.results.map((item) => { + const canRemoveFailedKey = !item.ok && Boolean(item.key_id) && ( + draftReviewKeys.length > 0 || configuredReviewKeys.some((key) => key.id === item.key_id) + ) + const masked = item.key_masked || `Key #${item.key_index}` + return ( +
+
+ {masked} +
+ {item.ok ? t('common.success') : t('common.failed')} + {canRemoveFailedKey ? ( + + ) : null} +
+
+
{item.latency_ms} ms · {item.ok ? item.confidence.toFixed(2) : item.error || '-'}
+
+ ) + })} +
+ ) : null} ) : null} From 4554c0ac81c96cff16b919753954f086b0f96c34 Mon Sep 17 00:00:00 2001 From: hu <187184415@qq.com> Date: Sun, 9 Aug 2026 21:22:13 +0800 Subject: [PATCH 2/3] fix(prompt-filter): preserve CY evidence and reduce review latency --- admin/prompt_intelligence_ai.go | 195 ++++++++++++++++++++++-- admin/prompt_intelligence_ai_test.go | 48 ++++++ database/prompt_policy_incident.go | 118 +++++++++++++- database/prompt_policy_incident_test.go | 51 +++++++ proxy/cyber_policy_test.go | 91 ++++++++++- proxy/prompt_risk_trust.go | 62 +++++++- proxy/prompt_risk_trust_test.go | 64 +++++++- proxy/prompt_rule_evidence.go | 146 +++++++++++++++++- 8 files changed, 735 insertions(+), 40 deletions(-) diff --git a/admin/prompt_intelligence_ai.go b/admin/prompt_intelligence_ai.go index d913511d..38b7bc5e 100644 --- a/admin/prompt_intelligence_ai.go +++ b/admin/prompt_intelligence_ai.go @@ -72,12 +72,36 @@ type promptIntelligenceAIAnalysisMetadata struct { APIKeyName string `json:"api_key_name,omitempty"` ReviewSystemPromptHash string `json:"review_system_prompt_hash"` UpstreamEvidenceCount int `json:"upstream_evidence_count"` + LearnableEvidenceCount int `json:"learnable_evidence_count"` Result promptIntelligenceAIDecision `json:"result"` RuleValidationError string `json:"rule_validation_error,omitempty"` IdentityValidationError string `json:"identity_validation_error,omitempty"` RawOutputPreview string `json:"raw_output_preview,omitempty"` } +type promptIntelligenceLearningContext struct { + Origin string `json:"origin"` + Role string `json:"role,omitempty"` + Text string `json:"text"` + Linked bool `json:"linked,omitempty"` + Truncated bool `json:"truncated,omitempty"` + Trust string `json:"trust,omitempty"` +} + +type promptIntelligenceLearningEvidence struct { + Version int `json:"version"` + Quality string `json:"quality"` + PromptText string `json:"prompt_text,omitempty"` + Context []promptIntelligenceLearningContext `json:"context,omitempty"` + UpstreamError string `json:"upstream_error,omitempty"` + Transport string `json:"transport,omitempty"` + StatusCode int `json:"status_code,omitempty"` + AttemptIndex int `json:"attempt_index,omitempty"` + ReviewModel string `json:"review_model,omitempty"` + ReviewFlagged bool `json:"review_flagged,omitempty"` + ReviewError string `json:"review_error,omitempty"` +} + type promptIdentityUpdateResult struct { Mode string `json:"mode"` Suggested bool `json:"suggested"` @@ -271,12 +295,18 @@ func (h *Handler) AnalyzePromptIntelligenceCandidate(c *gin.Context) { writeError(c, http.StatusConflict, "该候选没有可供分析的上游 CY 证据") return } + learnableEvidence := selectPromptIntelligenceLearnableEvidence(upstreamEvidence, 20) + if len(learnableEvidence) == 0 { + writeError(c, http.StatusConflict, "该候选只有证据不足的 CY 记录,尚未提取到可学习的 Prompt 或关联上下文;已停止调用外部模型") + return + } + learnableEvidenceCount := countPromptIntelligenceLearnableEvidence(upstreamEvidence) cfg := h.store.GetPromptFilterConfig() reviewCfg := promptfilter.NormalizeReviewConfig(cfg.Review) reviewSystemPrompt := promptfilter.NormalizeReviewAdapterConfig(reviewCfg.Adapter).SystemPrompt analysisSystemPrompt := buildPromptIntelligenceAIIdentity(reviewSystemPrompt) - analysisInput := buildPromptIntelligenceAIEvidenceInput(candidate, upstreamEvidence) + analysisInput := buildPromptIntelligenceAIEvidenceInput(candidate, learnableEvidence) rawOutput, attribution, err := h.callPromptIntelligenceAI(c.Request.Context(), request, reviewCfg, analysisSystemPrompt, analysisInput) if err != nil { writeError(c, http.StatusBadGateway, err.Error()) @@ -287,7 +317,7 @@ func (h *Handler) AnalyzePromptIntelligenceCandidate(c *gin.Context) { writeError(c, http.StatusBadGateway, err.Error()) return } - coverage := summarizePromptIntelligenceCoverage(upstreamEvidence) + coverage := summarizePromptIntelligenceCoverage(learnableEvidence) if err := validatePromptIntelligenceAICoverageDecision(decision, coverage); err != nil { writeError(c, http.StatusBadGateway, err.Error()) return @@ -297,7 +327,7 @@ func (h *Handler) AnalyzePromptIntelligenceCandidate(c *gin.Context) { Version: 1, Provider: attribution.Provider, Model: attribution.Model, APIKeyID: attribution.APIKeyID, APIKeyName: attribution.APIKeyName, ReviewSystemPromptHash: promptfilter.StableEvidenceFingerprint("review-system-prompt", reviewSystemPrompt), - UpstreamEvidenceCount: len(upstreamEvidence), Result: decision, + UpstreamEvidenceCount: len(upstreamEvidence), LearnableEvidenceCount: learnableEvidenceCount, Result: decision, RawOutputPreview: promptfilter.RedactedPreview(promptfilter.RedactSensitive(rawOutput), 4000), } if decision.Rule != nil { @@ -341,9 +371,10 @@ func (h *Handler) AnalyzePromptIntelligenceCandidate(c *gin.Context) { if metadata.IdentityValidationError != "" { response.IdentityUpdate.BlockReason = metadata.IdentityValidationError } else if request.IdentityUpdateMode == promptIdentityUpdateModeGuardedAuto { - response.IdentityUpdate.Eligible = decision.Confidence >= promptIdentityAutoMinConfidence && len(upstreamEvidence) >= promptIdentityAutoMinUpstreamEvidence + directEvidenceCount := countPromptIntelligenceDirectEvidence(upstreamEvidence) + response.IdentityUpdate.Eligible = decision.Confidence >= promptIdentityAutoMinConfidence && directEvidenceCount >= promptIdentityAutoMinUpstreamEvidence if !response.IdentityUpdate.Eligible { - response.IdentityUpdate.BlockReason = fmt.Sprintf("受控自动应用要求置信度至少 %.2f 且同类上游证据至少 %d 条", promptIdentityAutoMinConfidence, promptIdentityAutoMinUpstreamEvidence) + response.IdentityUpdate.BlockReason = fmt.Sprintf("受控自动应用要求置信度至少 %.2f 且同类完整 Prompt 证据至少 %d 条", promptIdentityAutoMinConfidence, promptIdentityAutoMinUpstreamEvidence) } else { applied, applyErr := h.applyPromptIntelligenceIdentityPatch(c.Request.Context(), candidateID, analysisEvidence.ID, "guarded_auto") if applyErr != nil { @@ -446,13 +477,23 @@ unsafe to generalize.` func buildPromptIntelligenceAIEvidenceInput(candidate *database.PromptRuleCandidate, evidence []*database.PromptRuleCandidateEvidence) string { type safeEvidence struct { - SourceKind string `json:"source_kind"` - SamplePreview string `json:"sample_preview"` - Protocol string `json:"protocol,omitempty"` - Provider string `json:"provider,omitempty"` - Model string `json:"model,omitempty"` - ObservedAt time.Time `json:"observed_at"` - Context map[string]any `json:"context,omitempty"` + SourceKind string `json:"source_kind"` + EvidenceQuality string `json:"evidence_quality"` + SamplePreview string `json:"sample_preview,omitempty"` + PromptText string `json:"prompt_text,omitempty"` + RelatedContext []promptIntelligenceLearningContext `json:"related_context,omitempty"` + UpstreamError string `json:"upstream_error,omitempty"` + Transport string `json:"transport,omitempty"` + StatusCode int `json:"status_code,omitempty"` + AttemptIndex int `json:"attempt_index,omitempty"` + ReviewModel string `json:"review_model,omitempty"` + ReviewFlagged bool `json:"review_flagged,omitempty"` + ReviewError string `json:"review_error,omitempty"` + Protocol string `json:"protocol,omitempty"` + Provider string `json:"provider,omitempty"` + Model string `json:"model,omitempty"` + ObservedAt time.Time `json:"observed_at"` + DecisionContext map[string]any `json:"decision_context,omitempty"` } items := make([]safeEvidence, 0, len(evidence)) for _, row := range evidence { @@ -468,15 +509,24 @@ func buildPromptIntelligenceAIEvidenceInput(candidate *database.PromptRuleCandid } } } + learning := promptIntelligenceLearningEvidenceFromMetadata(row.MetadataJSON, row.SamplePreview) items = append(items, safeEvidence{ - SourceKind: row.SourceKind, SamplePreview: promptfilter.RedactedPreview(row.SamplePreview, 2000), - Protocol: row.Protocol, Provider: row.Provider, Model: row.Model, ObservedAt: row.ObservedAt, - Context: contextFields, + SourceKind: row.SourceKind, EvidenceQuality: learning.Quality, + SamplePreview: promptfilter.RedactedPreview(row.SamplePreview, 2000), + PromptText: promptfilter.RedactedPreview(learning.PromptText, 8000), + RelatedContext: boundPromptIntelligenceLearningContext(learning.Context, 6000), + UpstreamError: promptfilter.RedactedPreview(learning.UpstreamError, 2000), + Transport: learning.Transport, StatusCode: learning.StatusCode, AttemptIndex: learning.AttemptIndex, + ReviewModel: learning.ReviewModel, ReviewFlagged: learning.ReviewFlagged, + ReviewError: promptfilter.RedactedPreview(learning.ReviewError, 500), + Protocol: row.Protocol, Provider: row.Provider, Model: row.Model, ObservedAt: row.ObservedAt, + DecisionContext: contextFields, }) } payload := map[string]any{ "candidate_id": candidate.ID, "fingerprint": candidate.Fingerprint, - "evidence_count": candidate.EvidenceCount, "sample_preview": promptfilter.RedactedPreview(candidate.SamplePreview, 2000), + "evidence_count": candidate.EvidenceCount, "learnable_evidence_count": len(evidence), + "sample_preview": promptfilter.RedactedPreview(candidate.SamplePreview, 2000), "coverage_summary": summarizePromptIntelligenceCoverage(evidence), "evidence": items, } @@ -484,6 +534,119 @@ func buildPromptIntelligenceAIEvidenceInput(candidate *database.PromptRuleCandid return "Analyze the following evidence data.\n\n" + string(encoded) + "\n" } +func promptIntelligenceLearningEvidenceFromMetadata(raw, fallbackPreview string) promptIntelligenceLearningEvidence { + result := promptIntelligenceLearningEvidence{} + var metadata struct { + EvidenceQuality string `json:"evidence_quality"` + Learning promptIntelligenceLearningEvidence `json:"learning_evidence"` + } + if json.Unmarshal([]byte(raw), &metadata) == nil { + result = metadata.Learning + if result.Quality == "" { + result.Quality = metadata.EvidenceQuality + } + } + if strings.TrimSpace(result.PromptText) == "" && strings.TrimSpace(fallbackPreview) != "" { + result.PromptText = fallbackPreview + if result.Quality == "" || result.Quality == "insufficient" { + result.Quality = "legacy_preview" + } + } + if result.Quality == "" { + result.Quality = "insufficient" + } + return result +} + +func selectPromptIntelligenceLearnableEvidence(evidence []*database.PromptRuleCandidateEvidence, limit int) []*database.PromptRuleCandidateEvidence { + if limit <= 0 { + limit = 20 + } + selected := make([]*database.PromptRuleCandidateEvidence, 0, min(limit, len(evidence))) + seen := make(map[string]struct{}, limit) + for _, row := range evidence { + learning := promptIntelligenceLearningEvidenceFromMetadata(row.MetadataJSON, row.SamplePreview) + text := strings.TrimSpace(learning.PromptText) + if text == "" { + parts := make([]string, 0, len(learning.Context)) + for _, segment := range learning.Context { + if value := strings.TrimSpace(segment.Text); value != "" { + parts = append(parts, segment.Origin+": "+value) + } + } + text = strings.Join(parts, "\n") + } + if text == "" || learning.Quality == "insufficient" { + continue + } + fingerprint := promptfilter.PromptEvidenceFingerprint(text) + if fingerprint == "" { + fingerprint = row.SourceRefHash + } + if _, exists := seen[fingerprint]; exists { + continue + } + seen[fingerprint] = struct{}{} + selected = append(selected, row) + if len(selected) >= limit { + break + } + } + return selected +} + +func countPromptIntelligenceDirectEvidence(evidence []*database.PromptRuleCandidateEvidence) int { + count := 0 + for _, row := range evidence { + learning := promptIntelligenceLearningEvidenceFromMetadata(row.MetadataJSON, row.SamplePreview) + if strings.TrimSpace(learning.PromptText) != "" && learning.Quality != "context_only" && learning.Quality != "insufficient" { + count++ + } + } + return count +} + +func countPromptIntelligenceLearnableEvidence(evidence []*database.PromptRuleCandidateEvidence) int { + count := 0 + for _, row := range evidence { + learning := promptIntelligenceLearningEvidenceFromMetadata(row.MetadataJSON, row.SamplePreview) + if learning.Quality == "insufficient" { + continue + } + if strings.TrimSpace(learning.PromptText) != "" { + count++ + continue + } + for _, context := range learning.Context { + if strings.TrimSpace(context.Text) != "" { + count++ + break + } + } + } + return count +} + +func boundPromptIntelligenceLearningContext(contexts []promptIntelligenceLearningContext, maxRunes int) []promptIntelligenceLearningContext { + if maxRunes <= 0 { + return nil + } + result := make([]promptIntelligenceLearningContext, 0, min(len(contexts), 8)) + remaining := maxRunes + for _, context := range contexts { + if remaining <= 0 || len(result) >= 8 { + break + } + context.Text = promptfilter.RedactedPreview(context.Text, remaining) + if strings.TrimSpace(context.Text) == "" { + continue + } + remaining -= len([]rune(context.Text)) + result = append(result, context) + } + return result +} + func summarizePromptIntelligenceCoverage(evidence []*database.PromptRuleCandidateEvidence) promptIntelligenceCoverageSummary { summary := promptIntelligenceCoverageSummary{EffectiveCoverage: "unknown", UpstreamEvidence: len(evidence)} for _, row := range evidence { diff --git a/admin/prompt_intelligence_ai_test.go b/admin/prompt_intelligence_ai_test.go index 36d01d63..1a8c2e3f 100644 --- a/admin/prompt_intelligence_ai_test.go +++ b/admin/prompt_intelligence_ai_test.go @@ -86,6 +86,54 @@ func TestPromptIntelligenceCoverageAllowsNoChangeWhenEveryCYWasBlocked(t *testin } } +func TestPromptIntelligenceLearnableEvidenceSelectionRejectsInsufficientAndDeduplicates(t *testing.T) { + insufficient := []*database.PromptRuleCandidateEvidence{ + {SourceRefHash: "one", MetadataJSON: `{"evidence_quality":"insufficient","learning_evidence":{"version":1,"quality":"insufficient"}}`}, + {SourceRefHash: "two", MetadataJSON: `{"evidence_quality":"insufficient","learning_evidence":{"version":1,"quality":"insufficient"}}`}, + } + if selected := selectPromptIntelligenceLearnableEvidence(insufficient, 20); len(selected) != 0 { + t.Fatalf("insufficient evidence selected for AI: %#v", selected) + } + evidence := []*database.PromptRuleCandidateEvidence{ + {SourceRefHash: "one", MetadataJSON: `{"evidence_quality":"complete","learning_evidence":{"version":1,"quality":"complete","prompt_text":"same request"}}`}, + {SourceRefHash: "two", MetadataJSON: `{"evidence_quality":"complete","learning_evidence":{"version":1,"quality":"complete","prompt_text":"same request"}}`}, + {SourceRefHash: "three", MetadataJSON: `{"evidence_quality":"context_only","learning_evidence":{"version":1,"quality":"context_only","context":[{"origin":"history","text":"linked context"}]}}`}, + } + selected := selectPromptIntelligenceLearnableEvidence(evidence, 20) + if len(selected) != 2 { + t.Fatalf("representative evidence len=%d want=2", len(selected)) + } + if direct := countPromptIntelligenceDirectEvidence(selected); direct != 1 { + t.Fatalf("direct evidence count=%d want=1", direct) + } +} + +func TestPromptIntelligenceEvidenceInputIncludesDurableLearningBundle(t *testing.T) { + evidence := []*database.PromptRuleCandidateEvidence{{ + SourceKind: database.PromptRuleCandidateSourceUpstreamCyberPolicy, + SamplePreview: "preview", + MetadataJSON: `{ + "local_action":"allow","local_outcome":"no_hit","local_comparison":"confirmed_miss", + "evidence_quality":"complete","learning_evidence":{ + "version":1,"quality":"complete","prompt_text":"full request Authorization: Bearer secret-token", + "context":[{"origin":"history","text":"linked context"}], + "upstream_error":"cyber_policy details","transport":"sse","status_code":400,"attempt_index":2, + "review_model":"deepseek-test","review_flagged":false,"review_error":"timeout" + } + }`, + Protocol: "responses", Provider: "openai", Model: "gpt-5.6-sol", ObservedAt: time.Now(), + }} + input := buildPromptIntelligenceAIEvidenceInput(&database.PromptRuleCandidate{ID: 7, EvidenceCount: 1}, evidence) + for _, expected := range []string{"full request", "linked context", "cyber_policy details", "deepseek-test", `"status_code":400`, `"learnable_evidence_count":1`} { + if !strings.Contains(input, expected) { + t.Fatalf("AI evidence input missing %q: %s", expected, input) + } + } + if strings.Contains(input, "secret-token") || !strings.Contains(input, "[REDACTED]") { + t.Fatalf("AI evidence input was not redacted: %s", input) + } +} + func TestPromptIntelligenceReviewProviderUsesBoundedParallelKeys(t *testing.T) { var active atomic.Int32 var maximum atomic.Int32 diff --git a/database/prompt_policy_incident.go b/database/prompt_policy_incident.go index 2ecad47b..e5713a5e 100644 --- a/database/prompt_policy_incident.go +++ b/database/prompt_policy_incident.go @@ -484,7 +484,45 @@ func reconcileStoredPromptPolicyIncidentFromShadowTx(ctx context.Context, tx *sq input.AuditScore, input.ReasonCode, input.PrimaryOrigin, input.MatchedPatterns, strings.TrimSpace(input.RequestCorrelationID)); err != nil { return err } - _, err := tx.ExecContext(ctx, `UPDATE prompt_risk_events SET + rows, err := tx.QueryContext(ctx, `SELECT evidence.id, evidence.metadata_json + FROM prompt_rule_candidate_evidence evidence + JOIN prompt_policy_incidents incident ON incident.incident_id=evidence.prompt_policy_incident_id + WHERE incident.request_correlation_id=$1 AND incident.upstream_error_code='cyber_policy' + AND evidence.source_kind=$2`, strings.TrimSpace(input.RequestCorrelationID), PromptRuleCandidateSourceUpstreamCyberPolicy) + if err != nil { + return err + } + type evidenceMetadataUpdate struct { + id int64 + metadata string + } + updates := make([]evidenceMetadataUpdate, 0, 2) + for rows.Next() { + var id int64 + var raw string + if scanErr := rows.Scan(&id, &raw); scanErr != nil { + rows.Close() + return scanErr + } + updated, updateErr := mergePromptPolicyCandidateEvidenceMetadata(raw, PromptPolicyOutcomeAuditHit, PromptPolicyComparisonLocalDetected, + input.AuditScore, input.ReasonCode, input.PrimaryOrigin, input.MatchedPatterns) + if updateErr != nil { + rows.Close() + return updateErr + } + updates = append(updates, evidenceMetadataUpdate{id: id, metadata: updated}) + } + if err = rows.Err(); err != nil { + rows.Close() + return err + } + rows.Close() + for _, update := range updates { + if _, err = tx.ExecContext(ctx, `UPDATE prompt_rule_candidate_evidence SET metadata_json=$1 WHERE id=$2`, update.metadata, update.id); err != nil { + return err + } + } + _, err = tx.ExecContext(ctx, `UPDATE prompt_risk_events SET event_kind='upstream_cy_local_detected', request_risk_score=28, evidence_confidence=85, local_outcome=$1, local_comparison=$2, reason_code=$3 WHERE source_type=$4 AND source_id IN ( @@ -493,6 +531,79 @@ func reconcileStoredPromptPolicyIncidentFromShadowTx(ctx context.Context, tx *sq return err } +func mergePromptPolicyCandidateEvidenceMetadata(raw, outcome, comparison string, auditScore int, reasonCode, primaryOrigin, matchedPatterns string) (string, error) { + metadata := map[string]any{} + if strings.TrimSpace(raw) != "" { + if err := json.Unmarshal([]byte(raw), &metadata); err != nil { + return "", err + } + } + var matches any = []any{} + if strings.TrimSpace(matchedPatterns) != "" { + if err := json.Unmarshal([]byte(matchedPatterns), &matches); err != nil { + return "", err + } + } + metadata["local_outcome"] = outcome + metadata["local_comparison"] = comparison + metadata["local_audit_score"] = auditScore + metadata["local_reason_code"] = reasonCode + metadata["local_primary_origin"] = primaryOrigin + metadata["local_matches"] = matches + learning, _ := metadata["learning_evidence"].(map[string]any) + if learning == nil { + learning = map[string]any{"version": 1} + } + learning["shadow_audit"] = map[string]any{ + "audit_score": auditScore, "reason_code": reasonCode, + "primary_origin": primaryOrigin, "matches": matches, + } + metadata["learning_evidence"] = learning + matchCount := 0 + if values, ok := matches.([]any); ok { + matchCount = len(values) + } + encoded, err := json.Marshal(metadata) + if err != nil { + return "", err + } + if len(encoded) > 64*1024 { + delete(metadata, "local_matches") + metadata["local_matches_count"] = matchCount + learning["shadow_audit"] = map[string]any{ + "audit_score": auditScore, "reason_code": reasonCode, + "primary_origin": primaryOrigin, "match_count": matchCount, + } + encoded, err = json.Marshal(metadata) + if err != nil { + return "", err + } + } + if len(encoded) > 64*1024 { + // Old evidence may already sit at the portable metadata ceiling. Keep the + // durable learning bundle and the reconciled decision, while dropping + // redundant incident fields that remain available from the incident row. + metadata = map[string]any{ + "evidence_quality": metadata["evidence_quality"], + "learning_evidence": learning, + "local_outcome": outcome, + "local_comparison": comparison, + "local_audit_score": auditScore, + "local_reason_code": reasonCode, + "local_primary_origin": primaryOrigin, + "local_matches_count": matchCount, + } + encoded, err = json.Marshal(metadata) + if err != nil { + return "", err + } + } + if len(encoded) > 64*1024 { + return "", errors.New("reconciled candidate learning evidence exceeds 64 KiB") + } + return string(encoded), nil +} + func (db *DB) PersistPromptPolicyIncident(ctx context.Context, rawIncident PromptPolicyIncidentInput, rawCandidate PromptRuleCandidateInput, rawEvidence PromptRuleCandidateEvidenceInput) error { if db == nil { return errors.New("database is nil") @@ -543,6 +654,11 @@ func (db *DB) PersistPromptPolicyIncident(ctx context.Context, rawIncident Promp } if shadowEvidence != nil { mergePromptPolicyShadowEvidence(&incident, shadowEvidence) + evidence.MetadataJSON, err = mergePromptPolicyCandidateEvidenceMetadata(evidence.MetadataJSON, incident.LocalOutcome, + incident.LocalComparison, *incident.LocalAuditScore, incident.LocalReasonCode, incident.LocalPrimaryOrigin, incident.LocalMatchedPatterns) + if err != nil { + return err + } if _, execErr := tx.ExecContext(ctx, `UPDATE prompt_policy_incidents SET local_comparison=$1, local_outcome=$2, local_audit_score=$3, local_reason_code=$4, local_primary_origin=$5, local_matched_patterns=$6 diff --git a/database/prompt_policy_incident_test.go b/database/prompt_policy_incident_test.go index 18c539e3..d443abc7 100644 --- a/database/prompt_policy_incident_test.go +++ b/database/prompt_policy_incident_test.go @@ -6,6 +6,7 @@ import ( "database/sql" "database/sql/driver" "encoding/hex" + "encoding/json" "fmt" "path/filepath" "strings" @@ -171,6 +172,29 @@ func TestPromptPolicyIncidentReconcilesAsyncShadowEvidenceInEitherWriteOrder(t * if eventKind != "upstream_cy_local_detected" || comparison != PromptPolicyComparisonLocalDetected { t.Fatalf("risk event was not reconciled: kind=%q comparison=%q", eventKind, comparison) } + evidenceRows, err := db.ListPromptRuleCandidateEvidence(context.Background(), got.CandidateID, 100) + if err != nil { + t.Fatalf("query reconciled candidate evidence len=%d err=%v", len(evidenceRows), err) + } + var evidenceMetadata string + for _, row := range evidenceRows { + if row.PromptPolicyIncidentID == incidentID { + evidenceMetadata = row.MetadataJSON + break + } + } + if evidenceMetadata == "" { + t.Fatalf("candidate evidence link for incident %q not found", incidentID) + } + var metadata map[string]any + if err := json.Unmarshal([]byte(evidenceMetadata), &metadata); err != nil { + t.Fatalf("decode reconciled evidence metadata: %v", err) + } + learning, _ := metadata["learning_evidence"].(map[string]any) + shadow, _ := learning["shadow_audit"].(map[string]any) + if metadata["local_comparison"] != PromptPolicyComparisonLocalDetected || metadata["local_outcome"] != PromptPolicyOutcomeAuditHit || metadata["local_reason_code"] != "prompt_policy_shadow_async" || shadow["audit_score"] != float64(20) { + t.Fatalf("candidate learning evidence was not reconciled: %#v", metadata) + } } t.Run("shadow_before_incident", func(t *testing.T) { @@ -232,6 +256,33 @@ func TestPromptPolicyIncidentReconcilesAsyncShadowEvidenceInEitherWriteOrder(t * }) } +func TestMergePromptPolicyCandidateEvidenceMetadataCompactsNearLimit(t *testing.T) { + rawBytes, err := json.Marshal(map[string]any{ + "padding": strings.Repeat("x", 65200), + "evidence_quality": "complete", + "learning_evidence": map[string]any{ + "version": 1, "quality": "complete", "prompt_text": "durable prompt", + }, + }) + if err != nil { + t.Fatal(err) + } + merged, err := mergePromptPolicyCandidateEvidenceMetadata(string(rawBytes), PromptPolicyOutcomeAuditHit, + PromptPolicyComparisonLocalDetected, 80, "prompt_policy_shadow_async", "current_user", + `[{"name":"rule","weight":80}]`) + if err != nil { + t.Fatal(err) + } + if len(merged) > 64*1024 || !json.Valid([]byte(merged)) { + t.Fatalf("compacted metadata bytes=%d valid=%t", len(merged), json.Valid([]byte(merged))) + } + if strings.Contains(merged, `"padding"`) || !strings.Contains(merged, `"prompt_text":"durable prompt"`) || !strings.Contains(merged, `"local_comparison":"local_detected"`) { + t.Fatalf("unexpected compacted metadata bytes=%d padding=%t prompt=%t comparison=%t", len(merged), + strings.Contains(merged, `"padding"`), strings.Contains(merged, `"prompt_text":"durable prompt"`), + strings.Contains(merged, `"local_comparison":"local_detected"`)) + } +} + func TestClearPromptFilterLogsKeepsIncidentsAndCandidateEvidence(t *testing.T) { db := newPromptPolicySQLiteTestDB(t) ctx := context.Background() diff --git a/proxy/cyber_policy_test.go b/proxy/cyber_policy_test.go index c7acc303..b5150d72 100644 --- a/proxy/cyber_policy_test.go +++ b/proxy/cyber_policy_test.go @@ -317,21 +317,96 @@ func TestPromptPolicyIncidentUsesStableFingerprintWhenPromptUnavailable(t *testi } defer db.Close() handler := NewHandler(auth.NewStore(nil, nil, &database.SystemSettings{PromptFilterEnabled: true}), db, nil, nil) + incidentIDs := make([]string, 0, 2) + for _, requestID := range []string{"unavailable-one", "unavailable-two"} { + ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) + ctx.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil) + ctx.Request.Header.Set("X-Request-ID", requestID) + incidentID, accepted := handler.logUpstreamCyberPolicy(ctx, "/v1/responses", "gpt-5.6-sol", []byte(`{"error":{"code":"cyber_policy"}}`)) + if !accepted || incidentID == "" { + t.Fatalf("incident enqueue accepted=%t id=%q", accepted, incidentID) + } + incidentIDs = append(incidentIDs, incidentID) + } + waitPromptFilterAuditIdle(t, db) + want := promptfilter.StableEvidenceFingerprint("cyber-insufficient", "/v1/responses\x00\x00\x00cyber_policy") + for _, incidentID := range incidentIDs { + incident, err := db.GetPromptPolicyIncident(context.Background(), incidentID) + if err != nil { + t.Fatalf("GetPromptPolicyIncident: %v", err) + } + if incident.PromptFingerprint != want { + t.Fatalf("unavailable prompt fingerprint = %q, want %q", incident.PromptFingerprint, want) + } + } + candidate, err := db.GetPromptRuleCandidateByFingerprint(context.Background(), want) + if err != nil { + t.Fatalf("GetPromptRuleCandidateByFingerprint: %v", err) + } + if candidate.EvidenceCount != 2 || candidate.SamplePreview != "" { + t.Fatalf("insufficient evidence quarantine candidate=%#v", candidate) + } + evidence, err := db.ListPromptRuleCandidateEvidence(context.Background(), candidate.ID, 10) + if err != nil || len(evidence) != 2 { + t.Fatalf("ListPromptRuleCandidateEvidence len=%d err=%v", len(evidence), err) + } + for _, row := range evidence { + if !strings.Contains(row.MetadataJSON, `"evidence_quality":"insufficient"`) || !strings.Contains(row.MetadataJSON, `"quality":"insufficient"`) { + t.Fatalf("insufficient evidence quality metadata missing: %s", row.MetadataJSON) + } + } +} + +func TestPromptPolicyLearningEvidenceIncludesBoundedContextAndReview(t *testing.T) { + gin.SetMode(gin.TestMode) + db, err := database.New("sqlite", filepath.Join(t.TempDir(), "cyber-learning-bundle.db")) + if err != nil { + t.Fatal(err) + } + defer db.Close() + handler := NewHandler(auth.NewStore(nil, nil, &database.SystemSettings{PromptFilterEnabled: true}), db, nil, nil) ctx, _ := gin.CreateTestContext(httptest.NewRecorder()) ctx.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil) - - incidentID, accepted := handler.logUpstreamCyberPolicy(ctx, "/v1/responses", "gpt-5.6-sol", []byte(`{"error":{"code":"cyber_policy"}}`)) - if !accepted || incidentID == "" { - t.Fatalf("incident enqueue accepted=%t id=%q", accepted, incidentID) + handler.capturePromptRuleLearningEvidence(ctx, "/v1/responses", "gpt-5.6-sol", promptGuardEvaluation{ + Envelope: promptfilter.RequestEnvelope{Protocol: promptfilter.ProtocolResponses, ModelFamily: promptfilter.ModelFamilyOpenAI, Segments: []promptfilter.Segment{ + {Origin: promptfilter.OriginHistory, Role: "user", Text: "linked defensive context", Linked: true, Trust: promptfilter.SegmentTrustClientSupplied}, + {Origin: promptfilter.OriginCurrentUser, Role: "user", Text: "current security request", Trust: promptfilter.SegmentTrustClientSupplied}, + {Origin: promptfilter.OriginSystem, Role: "system", Text: "fixed application boilerplate", Trust: promptfilter.SegmentTrustServerInjected}, + }}, + Decision: promptfilter.Decision{Action: promptfilter.ActionAllow, PrimaryOrigin: promptfilter.OriginCurrentUser}, + Verdict: promptfilter.Verdict{Enabled: true, Action: promptfilter.ActionAllow, ReviewModel: "review-model", ReviewError: "review timeout"}, + }) + incidentID, accepted := handler.logUpstreamCyberPolicy(ctx, "/v1/responses", "gpt-5.6-sol", []byte(`{"error":{"code":"cyber_policy","message":"blocked evidence"}}`)) + if !accepted { + t.Fatal("learning evidence was not enqueued") } waitPromptFilterAuditIdle(t, db) incident, err := db.GetPromptPolicyIncident(context.Background(), incidentID) if err != nil { - t.Fatalf("GetPromptPolicyIncident: %v", err) + t.Fatal(err) + } + evidence, err := db.ListPromptRuleCandidateEvidence(context.Background(), incident.CandidateID, 10) + if err != nil || len(evidence) != 1 { + t.Fatalf("evidence len=%d err=%v", len(evidence), err) + } + metadata := evidence[0].MetadataJSON + for _, expected := range []string{`"evidence_quality":"complete"`, `"prompt_text":"current security request"`, `"linked defensive context"`, `"review_model":"review-model"`, `"review_error":"review timeout"`, `"upstream_error"`} { + if !strings.Contains(metadata, expected) { + t.Fatalf("learning evidence metadata missing %s: %s", expected, metadata) + } + } + if strings.Contains(metadata, "fixed application boilerplate") { + t.Fatalf("server-injected boilerplate leaked into learning evidence: %s", metadata) + } +} + +func TestPromptPolicyRedactedLearningTextHonorsByteBudget(t *testing.T) { + got := promptPolicyRedactedLearningText(strings.Repeat("测试🙂", 10000), 20000, 20000) + if len(got) > 20000 { + t.Fatalf("learning text bytes = %d, want <= 20000", len(got)) } - want := promptfilter.StableEvidenceFingerprint("cyber-unavailable", incident.RequestCorrelationID+"\x00/v1/responses\x00gpt-5.6-sol") - if incident.PromptFingerprint != want { - t.Fatalf("unavailable prompt fingerprint = %q, want %q", incident.PromptFingerprint, want) + if !utf8.ValidString(got) { + t.Fatal("learning text was truncated inside a UTF-8 rune") } } diff --git a/proxy/prompt_risk_trust.go b/proxy/prompt_risk_trust.go index 0596db8d..5820269f 100644 --- a/proxy/prompt_risk_trust.go +++ b/proxy/prompt_risk_trust.go @@ -12,9 +12,13 @@ import ( "github.com/gin-gonic/gin" ) -const promptRiskTrustBypassAuditInterval = 10 * time.Minute +const ( + promptRiskTrustBypassAuditInterval = 10 * time.Minute + promptRiskTrustReviewLeaseDuration = 30 * time.Second +) -var promptRiskTrustBypassAudit sync.Map // subject key -> time.Time +var promptRiskTrustBypassAudit sync.Map // subject key -> time.Time +var promptRiskTrustReviewLeases sync.Map // subject key -> time.Time lease expiry func (h *Handler) promptRiskTrustPolicyForRequest(c *gin.Context) (database.PromptRiskTrustPolicy, string, bool) { if h == nil || h.store == nil || c == nil { @@ -37,10 +41,36 @@ func promptRiskTrustCanBypassReview(decision promptfilter.Decision, verdict prom if reviewText == "" || decision.Action != promptfilter.ActionAllow || verdict.Action != promptfilter.ActionAllow { return false } - if decision.AuditScore > 0 || decision.AuditRawScore > 0 || len(decision.Signals) > 0 || len(verdict.Matched) > 0 { + if len(decision.Errors) > 0 || verdict.ReviewError != "" || decision.Terminal || verdict.StrictHit || verdict.TerminalStrictHit || verdict.TerminalCategoryHit || verdict.SensitiveIntent { return false } - return len(decision.Errors) == 0 && verdict.ReviewError == "" + threshold := verdict.Threshold + if threshold <= 0 { + threshold = promptfilter.DefaultThreshold + } + if decision.AuditScore >= threshold || decision.AuditRawScore >= threshold*2 { + return false + } + return promptRiskTrustHasOnlyLowImpactAuditSignals(decision.Signals, verdict.Matched) +} + +func promptRiskTrustHasOnlyLowImpactAuditSignals(signals []promptfilter.Signal, matches []promptfilter.Match) bool { + for _, match := range matches { + if !match.SignalOnly || match.Strict { + return false + } + } + for _, signal := range signals { + if signal.TerminalCandidate || signal.StrikeEligible || signal.SuggestedAction == promptfilter.ActionWarn || signal.SuggestedAction == promptfilter.ActionBlock || len(signal.Matches) == 0 { + return false + } + for _, match := range signal.Matches { + if !match.SignalOnly || match.Strict { + return false + } + } + } + return true } func promptRiskTrustShouldSuspend(decision promptfilter.Decision, verdict promptfilter.Verdict) bool { @@ -65,14 +95,34 @@ func promptRiskTrustReviewRequired(c *gin.Context, cfg promptfilter.Config, poli forceInterval := time.Duration(adaptive.ForceReviewIntervalMinutes) * time.Minute forceDue := policy.LastModelReviewAt == nil || forceInterval <= 0 || now.Sub(policy.LastModelReviewAt.UTC()) >= forceInterval if forceDue { - return true + return promptRiskTrustAcquireReviewLease(subjectKey, now) } if adaptive.SamplePercent <= 0 { return false } correlationID := ensurePromptPolicyRequestCorrelationID(c) bucket := crc32.ChecksumIEEE([]byte(subjectKey+"\x00"+correlationID)) % 100 - return int(bucket) < adaptive.SamplePercent + return int(bucket) < adaptive.SamplePercent && promptRiskTrustAcquireReviewLease(subjectKey, now) +} + +func promptRiskTrustAcquireReviewLease(subjectKey string, now time.Time) bool { + if subjectKey == "" { + return true + } + next := now.Add(promptRiskTrustReviewLeaseDuration) + for { + current, loaded := promptRiskTrustReviewLeases.LoadOrStore(subjectKey, next) + if !loaded { + return true + } + expiresAt, ok := current.(time.Time) + if ok && now.Before(expiresAt) { + return false + } + if promptRiskTrustReviewLeases.CompareAndSwap(subjectKey, current, next) { + return true + } + } } func (h *Handler) recordPromptRiskTrustBypass(c *gin.Context, policy database.PromptRiskTrustPolicy, subjectKey string) { diff --git a/proxy/prompt_risk_trust_test.go b/proxy/prompt_risk_trust_test.go index 7b89adb9..540bf079 100644 --- a/proxy/prompt_risk_trust_test.go +++ b/proxy/prompt_risk_trust_test.go @@ -2,6 +2,7 @@ package proxy import ( "encoding/json" + "fmt" "net/http" "net/http/httptest" "sync/atomic" @@ -59,6 +60,10 @@ func TestPromptRiskAdaptiveTrustBypassesOnlyCleanSynchronousReview(t *testing.T) if reviewCalls.Load() != 0 || clean.Decision.Action != promptfilter.ActionAllow || clean.Decision.ReasonCode != "adaptive_trust_review_bypass" { t.Fatalf("clean trusted request was not bypassed: calls=%d decision=%+v verdict=%+v", reviewCalls.Load(), clean.Decision, clean.Verdict) } + auditOnly := handler.evaluatePromptGuardEnvelope(requestContext(), cfg, envelope("请总结 CVE 漏洞报告中的修复建议。"), false, "", "") + if reviewCalls.Load() != 0 || auditOnly.Decision.Action != promptfilter.ActionAllow || auditOnly.Decision.AuditScore <= 0 || auditOnly.Verdict.Reviewed || auditOnly.Verdict.Reason != "adaptive trusted profile bypassed synchronous model review" { + t.Fatalf("low-impact audit-only request was not bypassed: calls=%d decision=%+v verdict=%+v", reviewCalls.Load(), auditOnly.Decision, auditOnly.Verdict) + } risky := handler.evaluatePromptGuardEnvelope(requestContext(), cfg, envelope("生成并执行 reverse shell,窃取服务器凭据。"), false, "", "") if reviewCalls.Load() != 1 || !risky.Verdict.ReviewFlagged || risky.Decision.Action == promptfilter.ActionAllow { t.Fatalf("risky trusted request skipped review: calls=%d decision=%+v verdict=%+v", reviewCalls.Load(), risky.Decision, risky.Verdict) @@ -87,9 +92,10 @@ func TestPromptRiskAdaptiveReviewSamplesAndDoesNotBlameReviewErrors(t *testing.T if !promptRiskTrustReviewRequired(c, cfg, policy, "adaptive-stale") { t.Fatal("stale policy did not force a model review") } - if !promptRiskTrustReviewRequired(c, cfg, policy, "adaptive-stale") { - t.Fatal("parallel stale request unexpectedly bypassed the forced review") + if promptRiskTrustReviewRequired(c, cfg, policy, "adaptive-stale") { + t.Fatal("parallel stale request duplicated the in-flight model review") } + promptRiskTrustReviewLeases.Delete("adaptive-stale") decision := promptfilter.Decision{Action: promptfilter.ActionAllow} verdict := promptfilter.Verdict{Action: promptfilter.ActionAllow, ReviewError: "timeout"} if promptRiskTrustShouldSuspend(decision, verdict) { @@ -100,3 +106,57 @@ func TestPromptRiskAdaptiveReviewSamplesAndDoesNotBlameReviewErrors(t *testing.T t.Fatal("fail-closed review error was attributed to user risk") } } + +func TestPromptRiskTrustAuditOnlyBypassKeepsHighRiskReview(t *testing.T) { + match := promptfilter.Match{Name: "vulnerability_keyword", Weight: 20, SignalOnly: true} + decision := promptfilter.Decision{ + Action: promptfilter.ActionAllow, AuditScore: 20, AuditRawScore: 40, + Signals: []promptfilter.Signal{{SuggestedAction: promptfilter.ActionAllow, Matches: []promptfilter.Match{match}}}, + } + verdict := promptfilter.Verdict{Action: promptfilter.ActionAllow, Threshold: 50, Matched: []promptfilter.Match{match}} + if !promptRiskTrustCanBypassReview(decision, verdict, "review text") { + t.Fatal("low-impact signal-only evidence should retain adaptive review bypass") + } + decision.AuditScore = 50 + if promptRiskTrustCanBypassReview(decision, verdict, "review text") { + t.Fatal("threshold-level audit evidence must still receive model review") + } + decision.AuditScore = 20 + decision.Signals[0].Matches[0].Strict = true + if promptRiskTrustCanBypassReview(decision, verdict, "review text") { + t.Fatal("strict audit evidence must still receive model review") + } +} + +func TestPromptRiskAdaptiveReviewCoalescesConcurrentForcedReview(t *testing.T) { + const subjectKey = "adaptive-concurrent-review" + promptRiskTrustReviewLeases.Delete(subjectKey) + t.Cleanup(func() { promptRiskTrustReviewLeases.Delete(subjectKey) }) + stale := time.Now().UTC().Add(-24 * time.Hour) + policy := database.PromptRiskTrustPolicy{ID: 11, Source: database.PromptRiskTrustSourceAutomatic, LastModelReviewAt: &stale} + cfg := promptfilter.DefaultConfig() + cfg.Advanced.AdaptiveReview.Enabled = true + cfg.Advanced.AdaptiveReview.ForceReviewIntervalMinutes = 360 + + start := make(chan struct{}) + results := make(chan bool, 32) + for i := 0; i < cap(results); i++ { + go func(index int) { + c, _ := gin.CreateTestContext(httptest.NewRecorder()) + c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil) + c.Set(promptPolicyRequestCorrelationContextKey, fmt.Sprintf("concurrent-review-%d", index)) + <-start + results <- promptRiskTrustReviewRequired(c, cfg, policy, subjectKey) + }(i) + } + close(start) + required := 0 + for i := 0; i < cap(results); i++ { + if <-results { + required++ + } + } + if required != 1 { + t.Fatalf("concurrent forced reviews = %d, want exactly 1", required) + } +} diff --git a/proxy/prompt_rule_evidence.go b/proxy/prompt_rule_evidence.go index 1a1c3fac..13520717 100644 --- a/proxy/prompt_rule_evidence.go +++ b/proxy/prompt_rule_evidence.go @@ -4,6 +4,7 @@ import ( "encoding/json" "strings" "time" + "unicode/utf8" "github.com/codex2api/database" "github.com/codex2api/security/promptfilter" @@ -36,6 +37,43 @@ type promptRuleLearningEvidence struct { ReviewFlagged bool ReviewError string Matches []promptfilter.Match + Envelope promptfilter.RequestEnvelope +} + +const ( + promptPolicyEvidenceQualityComplete = "complete" + promptPolicyEvidenceQualityContextOnly = "context_only" + promptPolicyEvidenceQualityInsufficient = "insufficient" + promptPolicyLearningPromptRunes = 20000 + promptPolicyLearningPromptBytes = 20000 + promptPolicyLearningContextRunes = 12000 + promptPolicyLearningContextBytes = 12000 + promptPolicyLearningUpstreamErrorRunes = 4000 + promptPolicyLearningUpstreamErrorBytes = 4000 + promptPolicyLearningMetadataBytes = 60 * 1024 +) + +type promptPolicyLearningContextSegment struct { + Origin string `json:"origin"` + Role string `json:"role,omitempty"` + Text string `json:"text"` + Linked bool `json:"linked,omitempty"` + Truncated bool `json:"truncated,omitempty"` + Trust string `json:"trust,omitempty"` +} + +type promptPolicyLearningBundle struct { + Version int `json:"version"` + Quality string `json:"quality"` + PromptText string `json:"prompt_text,omitempty"` + Context []promptPolicyLearningContextSegment `json:"context,omitempty"` + UpstreamError string `json:"upstream_error,omitempty"` + Transport string `json:"transport,omitempty"` + StatusCode int `json:"status_code,omitempty"` + AttemptIndex int `json:"attempt_index,omitempty"` + ReviewModel string `json:"review_model,omitempty"` + ReviewFlagged bool `json:"review_flagged,omitempty"` + ReviewError string `json:"review_error,omitempty"` } type upstreamCyberPolicyAttempt struct { @@ -114,6 +152,7 @@ func (h *Handler) capturePromptRuleLearningEvidence(c *gin.Context, endpoint, mo ReviewFlagged: evaluation.Verdict.ReviewFlagged, ReviewError: evaluation.Verdict.ReviewError, Matches: append([]promptfilter.Match(nil), evaluation.Verdict.Matched...), + Envelope: evaluation.Envelope, }) } @@ -149,12 +188,32 @@ func (h *Handler) enqueueUpstreamCyberPolicyEvidence(c *gin.Context, endpoint, m if len(matchesJSON) == 0 || string(matchesJSON) == "null" { matchesJSON = []byte("[]") } - fingerprint := promptfilter.PromptEvidenceFingerprint(captured.Text) - if fingerprint == "" { - fingerprint = promptfilter.StableEvidenceFingerprint("cyber-unavailable", requestCorrelationID+"\x00"+endpoint+"\x00"+model) - } preview := promptfilter.RedactedPreview(captured.Text, 2000) checkText := promptfilter.RedactedPreview(captured.Text, promptFilterFullTextMaxRunes) + learningSourceText := strings.TrimSpace(envelopeDirectCurrentUserText(captured.Envelope)) + if learningSourceText == "" && captured.PrimaryOrigin == string(promptfilter.OriginApplicationCandidate) { + learningSourceText = strings.TrimSpace(captured.Text) + } + if learningSourceText == "" && len(captured.Envelope.Segments) == 0 { + learningSourceText = strings.TrimSpace(captured.Text) + } + learningPrompt := promptPolicyRedactedLearningText(learningSourceText, promptPolicyLearningPromptRunes, promptPolicyLearningPromptBytes) + learningContext, learningContextText := promptPolicyLearningContext(captured.Envelope) + evidenceQuality := promptPolicyEvidenceQualityInsufficient + fingerprintText := learningSourceText + if fingerprintText != "" { + evidenceQuality = promptPolicyEvidenceQualityComplete + } else if strings.TrimSpace(learningContextText) != "" { + evidenceQuality = promptPolicyEvidenceQualityContextOnly + fingerprintText = learningContextText + } + fingerprint := promptfilter.PromptEvidenceFingerprint(fingerprintText) + if fingerprint == "" { + // Evidence without any extractable request text belongs to one operational + // quarantine bucket. Using the request correlation ID here created one + // permanently unlearnable candidate per CY incident. + fingerprint = promptfilter.StableEvidenceFingerprint("cyber-insufficient", endpoint+"\x00"+protocol+"\x00"+provider+"\x00"+errorCode) + } state := captured.EvaluationState if state == "" { state = database.PromptPolicyEvaluationUnavailable @@ -221,7 +280,14 @@ func (h *Handler) enqueueUpstreamCyberPolicyEvidence(c *gin.Context, endpoint, m incident.LocalAuditScore = promptPolicyInt(captured.AuditScore) incident.LocalAuditRawScore = promptPolicyInt(captured.AuditRawScore) } - metadata, _ := json.Marshal(map[string]any{ + learningBundle := promptPolicyLearningBundle{ + Version: 1, Quality: evidenceQuality, PromptText: learningPrompt, Context: learningContext, + UpstreamError: promptPolicyRedactedLearningText(promptfilter.RedactSensitive(string(body)), promptPolicyLearningUpstreamErrorRunes, promptPolicyLearningUpstreamErrorBytes), + Transport: transport, StatusCode: attempt.StatusCode, AttemptIndex: attempt.AttemptIndex, + ReviewModel: captured.ReviewModel, ReviewFlagged: captured.ReviewFlagged, + ReviewError: promptfilter.RedactedPreview(captured.ReviewError, 1000), + } + metadataFields := map[string]any{ "incident_id": incidentID, "request_correlation_id": requestCorrelationID, "error_code": errorCode, "endpoint": endpoint, "local_evaluation_state": state, "local_outcome": outcome, "local_action": captured.Action, "local_score": incident.LocalScore, @@ -229,13 +295,27 @@ func (h *Handler) enqueueUpstreamCyberPolicyEvidence(c *gin.Context, endpoint, m "local_matches": captured.Matches, "platform": platform, "prompt_available": available, "local_comparison": localComparison, "account_id": attempt.AccountID, "account_groups": routing.AccountGroupNames, "newapi_policy_status": audit.NewAPIPolicyStatus, "newapi_platform": audit.NewAPIPlatform, - }) + "evidence_quality": evidenceQuality, "learning_evidence": learningBundle, + } + metadata, _ := json.Marshal(metadataFields) + if len(metadata) > promptPolicyLearningMetadataBytes { + // The incident retains the complete match JSON. Candidate evidence favors + // the durable learning bundle when the portable 64 KiB metadata limit is + // approached, especially for multibyte Prompt text. + delete(metadataFields, "local_matches") + metadataFields["local_matches_count"] = len(captured.Matches) + metadata, _ = json.Marshal(metadataFields) + } + rationale := "上游返回 cyber_policy,等待归因和候选规则审核" + if evidenceQuality == promptPolicyEvidenceQualityInsufficient { + rationale = "上游返回 cyber_policy,但请求文本证据不足;仅归档并等待补证,不得用于自动学习" + } candidate := database.PromptRuleCandidateInput{ Fingerprint: fingerprint, Kind: database.PromptRuleCandidateKindEvidence, Source: database.PromptRuleCandidateSourceUpstreamCyberPolicy, SamplePreview: preview, - Rationale: "上游返回 cyber_policy,等待归因和候选规则审核", + Rationale: rationale, } evidence := database.PromptRuleCandidateEvidenceInput{ SourceKind: database.PromptRuleCandidateSourceUpstreamCyberPolicy, @@ -314,6 +394,7 @@ func (h *Handler) capturePromptRuleLearningEvidenceOnUpstreamFailure(c *gin.Cont } fallback.EvaluationState = database.PromptPolicyEvaluationNotRun fallback.Text = text + fallback.Envelope = envelope fallback.Mode = cfg.Mode fallback.Threshold = cfg.Threshold if envelope.Protocol != promptfilter.ProtocolUnknown { @@ -325,6 +406,57 @@ func (h *Handler) capturePromptRuleLearningEvidenceOnUpstreamFailure(c *gin.Cont return fallback, true } +func promptPolicyLearningContext(envelope promptfilter.RequestEnvelope) ([]promptPolicyLearningContextSegment, string) { + segments := make([]promptPolicyLearningContextSegment, 0, 8) + parts := make([]string, 0, 8) + remainingBytes := promptPolicyLearningContextBytes + for _, segment := range envelope.Segments { + if remainingBytes <= 0 || len(segments) >= 8 || !promptPolicyLearningContextSegmentEligible(segment) { + continue + } + text := promptPolicyRedactedLearningText(segment.Text, promptPolicyLearningContextRunes, remainingBytes) + if text == "" { + continue + } + remainingBytes -= len(text) + segments = append(segments, promptPolicyLearningContextSegment{ + Origin: string(segment.Origin), Role: segment.Role, Text: text, + Linked: segment.Linked, Truncated: segment.Truncated, Trust: string(segment.Trust), + }) + parts = append(parts, string(segment.Origin)+": "+text) + } + return segments, strings.Join(parts, "\n") +} + +func promptPolicyRedactedLearningText(text string, maxRunes, maxBytes int) string { + value := strings.TrimSpace(promptfilter.RedactedPreview(text, maxRunes)) + if value == "" || maxBytes <= 0 { + return "" + } + if len(value) <= maxBytes { + return value + } + cut := maxBytes + for cut > 0 && !utf8.RuneStart(value[cut]) { + cut-- + } + return strings.TrimSpace(value[:cut]) +} + +func promptPolicyLearningContextSegmentEligible(segment promptfilter.Segment) bool { + if strings.TrimSpace(segment.Text) == "" || segment.Origin == promptfilter.OriginCurrentUser || segment.Origin == promptfilter.OriginApplicationCandidate { + return false + } + switch segment.Origin { + case promptfilter.OriginHistory: + return segment.Linked + case promptfilter.OriginToolOutput, promptfilter.OriginToolArguments, promptfilter.OriginSessionContext, promptfilter.OriginAttachmentContent: + return segment.Trust != promptfilter.SegmentTrustServerInjected + default: + return false + } +} + func promptPolicyLocalOutcome(captured promptRuleLearningEvidence) string { switch captured.Action { case promptfilter.ActionBlock: From 474d4b28ec3b706adad7b919cf3cd5a0bf76233f Mon Sep 17 00:00:00 2001 From: hu <187184415@qq.com> Date: Sun, 9 Aug 2026 21:37:39 +0800 Subject: [PATCH 3/3] perf(prompt-filter): avoid materializing risk history --- database/prompt_risk_profile.go | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/database/prompt_risk_profile.go b/database/prompt_risk_profile.go index b2602cee..e4f5f055 100644 --- a/database/prompt_risk_profile.go +++ b/database/prompt_risk_profile.go @@ -948,7 +948,11 @@ func (db *DB) ListPromptRiskProfiles(ctx context.Context, query PromptRiskProfil AND (LOWER(pri.external_user_id) LIKE $%d OR LOWER(pri.user_name) LIKE $%d OR LOWER(pri.user_email) LIKE $%d OR LOWER(pri.user_group) LIKE $%d) ))`, i, i, i, i, i, i, i, i, i, i)) } - rows, err := db.conn.QueryContext(ctx, `WITH filtered_events AS MATERIALIZED ( + // Keep the shared filter inline. Materializing the full 30-day event rows + // duplicates a large TEXT-heavy working set before both aggregate passes; + // SQLite production databases with dense clean-review history can exhaust + // the admin request budget even though the final profile set is small. + rows, err := db.conn.QueryContext(ctx, `WITH filtered_events AS NOT MATERIALIZED ( SELECT * FROM prompt_risk_events WHERE `+strings.Join(clauses, " AND ")+` ), profile_aggregates AS ( SELECT subject_type, subject_key,