Merge pull request #5511 from wucm667/fix/pr-5234-ws-audit-logging

fix(security-audit): restore websocket audit logs
This commit is contained in:
Wesley Liddick
2026-08-12 09:58:24 +08:00
committed by GitHub
2 changed files with 57 additions and 14 deletions
@@ -102,39 +102,52 @@ func runSecurityAudit(c *gin.Context, reqLog *zap.Logger, coordinator *securitya
if entry, ok := cached.(securityAuditWSDedupeEntry); ok &&
entry.stage == request.Stage && entry.turn == turnNo && entry.bodyHash == bodyHash {
decision := entry.decision
logSecurityAuditDone(reqLog, request, decision, true)
return &decision
}
}
logSecurityAuditStart(reqLog, request, len(body), false)
decision := coordinator.Check(c.Request.Context(), request)
if decision.Kind == securityaudit.DecisionAllow {
c.Set(securityAuditWSDedupeContextKey, securityAuditWSDedupeEntry{
stage: request.Stage, turn: turnNo, bodyHash: bodyHash, decision: decision,
})
}
logSecurityAuditDone(reqLog, request, decision, false)
return &decision
}
}
if reqLog != nil {
reqLog.Info("security_audit.gateway_check_start",
zap.String("request_id", request.RequestID), zap.Int64("user_id", request.UserID),
zap.Int64("api_key_id", request.APIKeyID), zap.Int64p("group_id", request.GroupID),
zap.String("endpoint", request.Endpoint), zap.String("provider", request.Provider),
zap.String("protocol", request.Protocol), zap.String("model", request.Model), zap.String("stage", request.Stage),
zap.Int("body_bytes", len(body)))
}
logSecurityAuditStart(reqLog, request, len(body), false)
decision := coordinator.Check(c.Request.Context(), request)
if decision.AllowNextStage && cacheCompletion {
c.Set(securityAuditCompletedContextKey, true)
}
if reqLog != nil {
reqLog.Info("security_audit.gateway_check_done",
zap.String("request_id", request.RequestID), zap.String("decision", string(decision.Kind)),
zap.String("error_code", decision.ErrorCode), zap.Bool("allow_next_stage", decision.AllowNextStage),
zap.String("stage", request.Stage))
}
logSecurityAuditDone(reqLog, request, decision, false)
return &decision
}
func logSecurityAuditStart(reqLog *zap.Logger, request securityaudit.Request, bodyBytes int, cached bool) {
if reqLog == nil {
return
}
reqLog.Info("security_audit.gateway_check_start",
zap.String("request_id", request.RequestID), zap.Int64("user_id", request.UserID),
zap.Int64("api_key_id", request.APIKeyID), zap.Int64p("group_id", request.GroupID),
zap.String("endpoint", request.Endpoint), zap.String("provider", request.Provider),
zap.String("protocol", request.Protocol), zap.String("model", request.Model), zap.String("stage", request.Stage),
zap.Int("body_bytes", bodyBytes), zap.Bool("cached", cached))
}
func logSecurityAuditDone(reqLog *zap.Logger, request securityaudit.Request, decision securityaudit.Decision, cached bool) {
if reqLog == nil {
return
}
reqLog.Info("security_audit.gateway_check_done",
zap.String("request_id", request.RequestID), zap.String("decision", string(decision.Kind)),
zap.String("error_code", decision.ErrorCode), zap.Bool("allow_next_stage", decision.AllowNextStage),
zap.String("stage", request.Stage), zap.Bool("cached", cached))
}
func securityAuditWSTurn(c *gin.Context) (int, bool) {
turn, exists := c.Get(securityAuditWSTurnContextKey)
if !exists {
@@ -11,6 +11,8 @@ import (
middleware2 "github.com/Wei-Shaw/sub2api/internal/server/middleware"
"github.com/gin-gonic/gin"
"github.com/stretchr/testify/require"
"go.uber.org/zap"
"go.uber.org/zap/zaptest/observer"
)
func TestCachesSecurityAuditCompletionSkipsWebSocketStages(t *testing.T) {
@@ -131,6 +133,34 @@ func TestRunSecurityAuditDoesNotCacheFlaggedWebSocketDecision(t *testing.T) {
require.Equal(t, int64(2), engine.evaluates.Load())
}
func TestRunSecurityAuditLogsWebSocketChecksAndCacheHits(t *testing.T) {
gin.SetMode(gin.TestMode)
engine := &turnCountingEngine{mode: securityaudit.ModeBlocking}
coordinator := securityaudit.NewCoordinator(nil, engine)
core, logs := observer.New(zap.InfoLevel)
reqLog := zap.New(core)
recorder := httptest.NewRecorder()
c, _ := gin.CreateTestContext(recorder)
c.Request = httptest.NewRequest(http.MethodPost, "/v1/responses", nil)
c.Set(securityAuditWSTurnContextKey, 2)
payload := []byte(`{"type":"response.create","response":{"input":"same turn"}}`)
runSecurityAudit(c, reqLog, coordinator, nil, nil, middleware2.AuthSubject{UserID: 7}, "openai_responses", "gpt-test", payload, "subsequent_turn")
runSecurityAudit(c, reqLog, coordinator, nil, nil, middleware2.AuthSubject{UserID: 7}, "openai_responses", "gpt-test", payload, "subsequent_turn")
startLogs := logs.FilterMessage("security_audit.gateway_check_start").All()
require.Len(t, startLogs, 1)
require.Equal(t, false, startLogs[0].ContextMap()["cached"])
doneLogs := logs.FilterMessage("security_audit.gateway_check_done").All()
require.Len(t, doneLogs, 2)
require.Equal(t, false, doneLogs[0].ContextMap()["cached"])
require.Equal(t, true, doneLogs[1].ContextMap()["cached"])
require.Equal(t, "allow", doneLogs[1].ContextMap()["decision"])
require.Equal(t, "subsequent_turn", doneLogs[1].ContextMap()["stage"])
require.Equal(t, int64(1), engine.evaluates.Load())
}
type turnCountingEngine struct {
mode securityaudit.Mode
enqueues atomic.Int64