fix(security-audit): restore websocket audit logs
This commit is contained in:
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user