Files
sub2api/backend/internal/securityaudit/prompt_guard_test.go
T
Nick 291a737422 test: stop four concurrency tests from failing on a busy machine
All four pass on a quiet box and fail on a loaded one, each for its own
reason. None of them is testing the clock, so none of them should be
failing on it.

- ollama_cloud_usage_test.go: the second caller was released with
  `close(release)` right after its goroutine was started, not after it
  had reached the singleflight group. When the first refresh won that
  race, the second became a new singleflight execution, re-read the
  account, saw the LastAttemptAt the first one had just written, and
  came back with the 30-second manual-refresh 429 at line 685. It now
  counts account loads and waits for the second caller's own load,
  which happens right before it joins the group. Adds a counting
  GetByID to the test repo.

- gateway_hotpath_optimization_test.go: a 20ms sleep was meant to let
  all 12 callers reach the cache before the loader was released. A
  caller that arrived after the load had finished got a hit, not a
  miss, so the miss count came out 11 of 12. It now waits on the miss
  counter itself, which is the value the test asserts on.

- token_refresh_pool_health_test.go: the floor was `configuredSpacing`
  minus 10ms, i.e. 40ms out of 50ms. Each start timestamp is taken
  after the rate gate releases the goroutine, so scheduler delay can
  compress one observed gap with the gate behaving correctly — seen at
  37ms and again at 13ms. The floor is now a tenth of the configured
  spacing. Measured with providerQPS=20, 8 attempts, concurrency 2:
  gate at 50ms gives a minimum gap of 49.97ms, gate at 0 gives 22µs.
  So an unpaced gate sits three orders of magnitude under the 5ms floor
  and is still caught, while jitter has room to move. A comment warns
  against replacing this with an assertion on the total span of the
  starts: the span is set by how long each attempt takes under the
  concurrency limit, not by the gate — 471ms paced against 241ms
  unpaced — so a span check passes with the gate disabled.

- prompt_guard_test.go: the bound only has to show the failover shared
  the first endpoint's 70ms deadline and did not take the second
  endpoint's own 500ms one. An unshared deadline lands near 535ms, so
  350ms still fails loudly (seen: 224ms against a 180ms bound).

Tests only; no product code is touched. Each fix was checked in both
directions: it passes with the behaviour intact, and it still fails when
the behaviour is broken on purpose (for the QPS one, by swapping the
shared rate gate for a zero-interval one).
2026-07-25 21:59:20 +07:00

301 lines
14 KiB
Go

package securityaudit
import (
"context"
"errors"
"strings"
"sync"
"testing"
"time"
"github.com/stretchr/testify/require"
)
type scriptedScanner struct {
mu sync.Mutex
calls []string
block <-chan struct{}
entered chan<- struct{}
}
func (s *scriptedScanner) Scan(ctx context.Context, endpoint ActiveEndpoint, _ string, _ []string) (*NormalizedResult, error) {
s.mu.Lock()
s.calls = append(s.calls, endpoint.ID)
s.mu.Unlock()
if s.entered != nil {
select {
case s.entered <- struct{}{}:
default:
}
}
if s.block != nil {
select {
case <-s.block:
case <-ctx.Done():
return nil, &GuardError{Code: ErrorCodeUnavailable, Retryable: true, Timeout: true, Cause: ctx.Err()}
}
}
if endpoint.ID == "bad" {
return nil, &GuardError{Code: ErrorCodeUnavailable, Retryable: true}
}
if endpoint.ID == "invalid" {
return nil, &GuardError{Code: ErrorCodeInvalidResponse}
}
return &NormalizedResult{Decision: EventPass, RiskLevel: RiskLow, Action: ActionAllow, Safety: "Safe", ScannerScores: map[string]float64{}, ScannerEvidence: map[string]string{}, GuardEndpointID: endpoint.ID}, nil
}
func guardConfig(endpoints ...ActiveEndpoint) ActiveConfig {
return ActiveConfig{RiskControlEnabled: true, Enabled: true, BlockingEnabled: true, ConfigVersion: 2, Scanners: AllScannerIDs, Endpoints: endpoints}
}
func TestGuardEvaluatorOrderedFailoverAndInvalidTerminal(t *testing.T) {
scanner := &scriptedScanner{}
metrics := NewAtomicMetrics()
evaluator := newGuardEvaluator(scanner, nil, metrics, 4, 2)
snapshot := PromptSnapshot{RequestID: "r", ScanText: "hello", PromptLength: 5}
decision, err := evaluator.Evaluate(context.Background(), guardConfig(
ActiveEndpoint{ID: "bad", Enabled: true, TimeoutMS: 1000, InputLimit: 100},
ActiveEndpoint{ID: "good", Enabled: true, TimeoutMS: 1000, InputLimit: 100},
), snapshot)
require.NoError(t, err)
require.Equal(t, DecisionAllow, decision.Kind)
require.Equal(t, int64(1), metrics.Snapshot().Failovers)
_, err = evaluator.Evaluate(context.Background(), guardConfig(
ActiveEndpoint{ID: "invalid", Enabled: true, TimeoutMS: 1000, InputLimit: 100},
ActiveEndpoint{ID: "good", Enabled: true, TimeoutMS: 1000, InputLimit: 100},
), snapshot)
var guardErr *GuardError
require.ErrorAs(t, err, &guardErr)
require.Equal(t, ErrorCodeInvalidResponse, guardErr.Code)
snapshotMetrics := metrics.Snapshot()
require.Equal(t, int64(2), snapshotMetrics.Total)
require.Equal(t, int64(1), snapshotMetrics.Allowed)
require.Equal(t, int64(1), snapshotMetrics.Invalid)
}
func TestGuardEvaluatorGlobalBulkheadIsNonBlocking(t *testing.T) {
release := make(chan struct{})
entered := make(chan struct{}, 1)
scanner := &scriptedScanner{block: release, entered: entered}
metrics := NewAtomicMetrics()
evaluator := newGuardEvaluator(scanner, nil, metrics, 1, 1)
cfg := guardConfig(ActiveEndpoint{ID: "good", Enabled: true, TimeoutMS: 2000, InputLimit: 100})
done := make(chan error, 1)
go func() {
_, err := evaluator.Evaluate(context.Background(), cfg, PromptSnapshot{ScanText: "one", PromptLength: 3})
done <- err
}()
select {
case <-entered:
case <-time.After(time.Second):
t.Fatal("first evaluation did not enter scanner")
}
start := time.Now()
_, err := evaluator.Evaluate(context.Background(), cfg, PromptSnapshot{ScanText: "two", PromptLength: 3})
require.Error(t, err)
require.Less(t, time.Since(start), 200*time.Millisecond)
require.Equal(t, int64(1), metrics.Snapshot().BulkheadFull)
close(release)
require.NoError(t, <-done)
snapshotMetrics := metrics.Snapshot()
require.Equal(t, int64(2), snapshotMetrics.Total)
require.Equal(t, int64(1), snapshotMetrics.Allowed)
require.Equal(t, int64(1), snapshotMetrics.Unavailable)
}
func TestGuardEvaluatorPerNodeBulkheadIsNonBlocking(t *testing.T) {
release := make(chan struct{})
entered := make(chan struct{}, 1)
scanner := &scriptedScanner{block: release, entered: entered}
metrics := NewAtomicMetrics()
evaluator := newGuardEvaluator(scanner, nil, metrics, 2, 1)
cfg := guardConfig(ActiveEndpoint{ID: "same-node", Enabled: true, TimeoutMS: 2000, InputLimit: 100})
done := make(chan error, 1)
go func() {
_, err := evaluator.Evaluate(context.Background(), cfg, PromptSnapshot{ScanText: "one", PromptLength: 3})
done <- err
}()
select {
case <-entered:
case <-time.After(time.Second):
t.Fatal("first evaluation did not enter scanner")
}
started := time.Now()
_, err := evaluator.Evaluate(context.Background(), cfg, PromptSnapshot{ScanText: "two", PromptLength: 3})
require.Error(t, err)
require.Less(t, time.Since(started), 200*time.Millisecond)
require.GreaterOrEqual(t, metrics.Snapshot().BulkheadFull, int64(1))
close(release)
require.NoError(t, <-done)
}
func TestGuardEvaluatorLastChunkFailureNeverAllows(t *testing.T) {
call := 0
scanner := PromptScannerFunc(func(context.Context, ActiveEndpoint, string, []string) (*NormalizedResult, error) {
call++
if call == 2 {
return nil, &GuardError{Code: ErrorCodeUnavailable, Retryable: true, Cause: errors.New("down")}
}
return &NormalizedResult{Decision: EventPass, RiskLevel: RiskLow, Action: ActionAllow, ScannerScores: map[string]float64{}, ScannerEvidence: map[string]string{}}, nil
})
metrics := NewAtomicMetrics()
evaluator := newGuardEvaluator(scanner, nil, metrics, 2, 2)
_, err := evaluator.Evaluate(context.Background(), guardConfig(ActiveEndpoint{ID: "one", Enabled: true, TimeoutMS: 1000, InputLimit: 3}), PromptSnapshot{ScanText: "abcdef", PromptLength: 6})
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) {
calls++
return &NormalizedResult{
Decision: EventCritical, RiskLevel: RiskCritical, Action: ActionBlock, Safety: "Unsafe",
Categories: []string{"jailbreak"}, MatchedScanners: []string{"jailbreak"},
ScannerScores: map[string]float64{"jailbreak": 1}, ScannerEvidence: map[string]string{"jailbreak": "Jailbreak"},
}, nil
})
metrics := NewAtomicMetrics()
evaluator := newGuardEvaluator(scanner, nil, metrics, 2, 2)
decision, err := evaluator.Evaluate(context.Background(), guardConfig(
ActiveEndpoint{ID: "one", Enabled: true, TimeoutMS: 1000, InputLimit: 3},
), PromptSnapshot{ScanText: "abcdefghi", PromptLength: 9})
require.NoError(t, err)
require.Equal(t, DecisionBlock, decision.Kind)
require.Equal(t, 1, calls)
require.Equal(t, 3, decision.Result.ChunkTotal)
require.Equal(t, int64(1), metrics.Snapshot().Blocked)
}
func TestGuardEvaluatorFlagSharedDeadlineFailClosedAndContextCancel(t *testing.T) {
t.Run("flag allows next stage", func(t *testing.T) {
metrics := NewAtomicMetrics()
evaluator := newGuardEvaluator(PromptScannerFunc(func(context.Context, ActiveEndpoint, string, []string) (*NormalizedResult, error) {
return &NormalizedResult{Decision: EventFlag, RiskLevel: RiskMedium, Action: ActionWarn, Safety: "Controversial", Categories: []string{"violent"}, MatchedScanners: []string{"violent"}, ScannerScores: map[string]float64{"violent": .5}, ScannerEvidence: map[string]string{"violent": "Violent"}}, nil
}), nil, metrics, 2, 2)
decision, err := evaluator.Evaluate(context.Background(), guardConfig(ActiveEndpoint{ID: "one", Enabled: true, TimeoutMS: 1000, InputLimit: 100}), PromptSnapshot{ScanText: "review", PromptLength: 6})
require.NoError(t, err)
require.Equal(t, DecisionFlag, decision.Kind)
require.True(t, decision.AllowNextStage)
require.Equal(t, int64(1), metrics.Snapshot().Flagged)
})
t.Run("all failovers share first endpoint deadline", func(t *testing.T) {
calls := 0
scanner := PromptScannerFunc(func(ctx context.Context, endpoint ActiveEndpoint, _ string, _ []string) (*NormalizedResult, error) {
calls++
if endpoint.ID == "first" {
select {
case <-time.After(35 * time.Millisecond):
return nil, &GuardError{Code: ErrorCodeUnavailable, Retryable: true}
case <-ctx.Done():
return nil, &GuardError{Code: ErrorCodeUnavailable, Retryable: true, Timeout: true, Cause: ctx.Err()}
}
}
<-ctx.Done()
return nil, &GuardError{Code: ErrorCodeUnavailable, Retryable: true, Timeout: true, Cause: ctx.Err()}
})
metrics := NewAtomicMetrics()
evaluator := newGuardEvaluator(scanner, nil, metrics, 2, 2)
started := time.Now()
_, err := evaluator.Evaluate(context.Background(), guardConfig(
ActiveEndpoint{ID: "first", Enabled: true, TimeoutMS: 70, InputLimit: 100},
ActiveEndpoint{ID: "second", Enabled: true, TimeoutMS: 500, InputLimit: 100},
), PromptSnapshot{ScanText: "deadline", PromptLength: 8})
elapsed := time.Since(started)
require.Error(t, err)
require.Equal(t, 2, calls)
// The bound only has to prove the failover shared the first endpoint's
// 70ms deadline instead of taking the second endpoint's own 500ms one.
// An unshared deadline lands at ~535ms, so 350ms still fails loudly
// while leaving room for scheduler delay on a busy CI machine. A
// tighter bound made this test flaky, not stricter.
require.Less(t, elapsed, 350*time.Millisecond)
require.GreaterOrEqual(t, elapsed, 50*time.Millisecond)
require.Equal(t, int64(1), metrics.Snapshot().Failovers)
require.Equal(t, int64(1), metrics.Snapshot().Timeouts)
})
t.Run("canceled parent never allows", func(t *testing.T) {
ctx, cancel := context.WithCancel(context.Background())
cancel()
evaluator := newGuardEvaluator(PromptScannerFunc(func(ctx context.Context, _ ActiveEndpoint, _ string, _ []string) (*NormalizedResult, error) {
<-ctx.Done()
return nil, &GuardError{Code: ErrorCodeUnavailable, Retryable: true, Cause: ctx.Err()}
}), nil, NewAtomicMetrics(), 2, 2)
decision, err := evaluator.Evaluate(ctx, guardConfig(ActiveEndpoint{ID: "one", Enabled: true, TimeoutMS: 1000, InputLimit: 100}), PromptSnapshot{ScanText: "cancel", PromptLength: 6})
require.Error(t, err)
require.Nil(t, decision)
})
}
func TestGuardEvaluatorRecordsExistingResultOnceAndRecordFailureDoesNotChangeDecision(t *testing.T) {
for _, recordErr := range []error{nil, errors.New("database unavailable")} {
repo := &fakeJobRepository{recordBlockingErr: recordErr}
metrics := NewAtomicMetrics()
scannerCalls := 0
evaluator := newGuardEvaluator(PromptScannerFunc(func(context.Context, ActiveEndpoint, string, []string) (*NormalizedResult, error) {
scannerCalls++
return &NormalizedResult{Decision: EventCritical, RiskLevel: RiskCritical, Action: ActionBlock, Safety: "Unsafe", Categories: []string{"pii"}, MatchedScanners: []string{"pii"}, ScannerScores: map[string]float64{"pii": 1}, ScannerEvidence: map[string]string{"pii": "PII"}}, nil
}), repo, metrics, 2, 2)
decision, err := evaluator.Evaluate(context.Background(), guardConfig(ActiveEndpoint{ID: "one", Enabled: true, TimeoutMS: 1000, InputLimit: 100}), PromptSnapshot{ScanText: "raw prompt", RedactedPreview: "raw***", PromptLength: 10})
require.NoError(t, err)
require.Equal(t, DecisionBlock, decision.Kind)
require.Equal(t, 1, scannerCalls)
require.Equal(t, 1, repo.recordBlockingCalls)
require.Empty(t, repo.recordBlockingSnapshot.ScanText)
require.Same(t, decision.Result, repo.recordBlockingResult)
if recordErr != nil {
require.Equal(t, int64(1), metrics.Snapshot().RecordFailed)
} else {
require.Zero(t, metrics.Snapshot().RecordFailed)
}
}
}
func TestGuardEvaluatorNilResultAndScannerPanicBecomeStableFailures(t *testing.T) {
tests := []struct {
name string
scan PromptScannerFunc
code string
}{
{name: "nil result", scan: func(context.Context, ActiveEndpoint, string, []string) (*NormalizedResult, error) { return nil, nil }, code: ErrorCodeInvalidResponse},
{name: "panic", scan: func(context.Context, ActiveEndpoint, string, []string) (*NormalizedResult, error) {
panic("raw prompt canary")
}, code: ErrorCodeUnavailable},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
evaluator := newGuardEvaluator(tt.scan, nil, NewAtomicMetrics(), 2, 2)
_, err := evaluator.Evaluate(context.Background(), guardConfig(ActiveEndpoint{ID: "one", Enabled: true, TimeoutMS: 1000, InputLimit: 100}), PromptSnapshot{ScanText: "input", PromptLength: 5})
var guardErr *GuardError
require.ErrorAs(t, err, &guardErr)
require.Equal(t, tt.code, guardErr.Code)
require.NotContains(t, err.Error(), "canary")
})
}
}
type PromptScannerFunc func(context.Context, ActiveEndpoint, string, []string) (*NormalizedResult, error)
func (f PromptScannerFunc) Scan(ctx context.Context, endpoint ActiveEndpoint, chunk string, scanners []string) (*NormalizedResult, error) {
return f(ctx, endpoint, chunk, scanners)
}