package securityaudit import ( "context" "errors" "strconv" "strings" infraerrors "github.com/Wei-Shaw/sub2api/internal/pkg/errors" "github.com/Wei-Shaw/sub2api/internal/pkg/response" "github.com/Wei-Shaw/sub2api/internal/server/middleware" "github.com/gin-gonic/gin" ) type PromptAdminService interface { GetConfig() (PublicConfig, error) SaveConfig(context.Context, UpdateConfigRequest, int64) (PublicConfig, error) Probe(context.Context, ProbeRequest) ProbeResult Runtime(context.Context) RuntimeSnapshot ListEvents(context.Context, EventFilter, int, int) (*EventPage, error) GetEvent(context.Context, int64) (*Event, error) DeleteEvent(context.Context, int64) (*DeleteResult, error) DeleteEventsByIDs(context.Context, []int64) (*DeleteResult, error) PreviewDelete(context.Context, EventFilter, int64) (*DeletePreview, error) DeleteByFilter(context.Context, DeleteByFilterRequest, int64) (*DeleteResult, error) } type PromptAdminHandler struct{ service PromptAdminService } func NewPromptAdminHandler(service PromptAdminService) *PromptAdminHandler { return &PromptAdminHandler{service: service} } func (h *PromptAdminHandler) GetConfig(c *gin.Context) { config, err := h.service.GetConfig() if err != nil { response.ErrorFrom(c, err) return } response.Success(c, config) } func (h *PromptAdminHandler) UpdateConfig(c *gin.Context) { var request UpdateConfigRequest if err := c.ShouldBindJSON(&request); err != nil { setPromptAdminAudit(c, "failed", "prompt_audit_invalid_config_request", nil) response.ErrorFrom(c, infraerrors.BadRequest("prompt_audit_invalid_config_request", "提示词审计配置请求无效")) return } config, err := h.service.SaveConfig(c.Request.Context(), request, adminID(c)) if err != nil { setPromptAdminAudit(c, "failed", infraerrors.Reason(err), configAuditFields(request, nil)) response.ErrorFrom(c, err) return } setPromptAdminAudit(c, "success", "", configAuditFields(request, &config)) response.Success(c, config) } func (h *PromptAdminHandler) ProbeEndpoint(c *gin.Context) { var request ProbeRequest if err := c.ShouldBindJSON(&request); err != nil { setPromptAdminAudit(c, "failed", "prompt_audit_invalid_probe_request", nil) response.ErrorFrom(c, infraerrors.BadRequest("prompt_audit_invalid_probe_request", "审计节点探测请求无效")) return } result := h.service.Probe(c.Request.Context(), request) status := "failed" if result.OK { status = "success" } setPromptAdminAudit(c, status, result.ErrorCode, map[string]any{ "guard_endpoint_id": request.Endpoint.ID, "http_status": result.HTTPStatus, "latency_ms": result.LatencyMS, "token_applied": result.TokenApplied, "retryable": result.Retryable, }) response.Success(c, result) } func (h *PromptAdminHandler) GetRuntime(c *gin.Context) { response.Success(c, h.service.Runtime(c.Request.Context())) } func (h *PromptAdminHandler) ListEvents(c *gin.Context) { page, err := positiveIntQuery(c, "page", 1, 0) if err != nil { response.ErrorFrom(c, err) return } pageSize, err := positiveIntQuery(c, "page_size", 20, 100) if err != nil { response.ErrorFrom(c, err) return } filter, err := eventFilterFromQuery(c) if err != nil { response.ErrorFrom(c, err) return } result, err := h.service.ListEvents(c.Request.Context(), filter, page, pageSize) if err != nil { response.ErrorFrom(c, err) return } response.Success(c, result) } func (h *PromptAdminHandler) GetEvent(c *gin.Context) { id, err := strconv.ParseInt(c.Param("id"), 10, 64) if err != nil || id <= 0 { response.ErrorFrom(c, infraerrors.BadRequest("prompt_audit_invalid_event_id", "事件 ID 无效")) return } event, err := h.service.GetEvent(c.Request.Context(), id) if errors.Is(err, ErrEventNotFound) { response.ErrorFrom(c, infraerrors.NotFound("prompt_audit_event_not_found", "提示词审计事件不存在")) return } if err != nil { response.ErrorFrom(c, err) return } response.Success(c, event) } func (h *PromptAdminHandler) DeleteEvent(c *gin.Context) { id, err := strconv.ParseInt(c.Param("id"), 10, 64) if err != nil || id <= 0 { setPromptAdminAudit(c, "failed", "prompt_audit_invalid_event_id", nil) response.ErrorFrom(c, infraerrors.BadRequest("prompt_audit_invalid_event_id", "事件 ID 无效")) return } result, err := h.service.DeleteEvent(c.Request.Context(), id) if err != nil { setPromptAdminAudit(c, "failed", infraerrors.Reason(err), map[string]any{"event_id": id}) response.ErrorFrom(c, err) return } setPromptAdminAudit(c, "success", "", deleteAuditFields(result, map[string]any{"event_id": id})) LogWarn(EventEventDeleted, map[string]any{"user_id": adminID(c), "event_id": id, "status": "deleted"}) response.Success(c, result) } type batchDeleteRequest struct { IDs []int64 `json:"ids" binding:"required"` } func (h *PromptAdminHandler) BatchDelete(c *gin.Context) { var request batchDeleteRequest if err := c.ShouldBindJSON(&request); err != nil || len(request.IDs) == 0 || len(request.IDs) > 500 { setPromptAdminAudit(c, "failed", "prompt_audit_invalid_delete_batch", nil) response.ErrorFrom(c, infraerrors.BadRequest("prompt_audit_invalid_delete_batch", "批量删除必须包含 1-500 个事件 ID")) return } for _, id := range request.IDs { if id <= 0 { setPromptAdminAudit(c, "failed", "prompt_audit_invalid_event_id", map[string]any{"requested_count": len(request.IDs)}) response.ErrorFrom(c, infraerrors.BadRequest("prompt_audit_invalid_event_id", "事件 ID 无效")) return } } result, err := h.service.DeleteEventsByIDs(c.Request.Context(), request.IDs) if err != nil { setPromptAdminAudit(c, "failed", infraerrors.Reason(err), map[string]any{"requested_count": len(request.IDs)}) response.ErrorFrom(c, err) return } setPromptAdminAudit(c, "success", "", deleteAuditFields(result, map[string]any{"requested_count": len(request.IDs)})) LogWarn(EventEventsDeleted, map[string]any{"user_id": adminID(c), "status": "deleted"}) response.Success(c, result) } func (h *PromptAdminHandler) DeletePreview(c *gin.Context) { var filter EventFilter if err := c.ShouldBindJSON(&filter); err != nil { setPromptAdminAudit(c, "failed", "prompt_audit_delete_preview_invalid", nil) response.ErrorFrom(c, infraerrors.BadRequest("prompt_audit_delete_preview_invalid", "删除预览筛选无效")) return } preview, err := h.service.PreviewDelete(c.Request.Context(), filter, adminID(c)) if err != nil { setPromptAdminAudit(c, "failed", "prompt_audit_delete_preview_invalid", nil) response.ErrorFrom(c, infraerrors.BadRequest("prompt_audit_delete_preview_invalid", "删除预览筛选无效")) return } setPromptAdminAudit(c, "success", "", map[string]any{ "matched_count": preview.MatchedCount, "snapshot_max_id": preview.SnapshotMaxID, "filter_hash": preview.FilterHash, }) response.Success(c, preview) } func (h *PromptAdminHandler) DeleteByFilter(c *gin.Context) { var request DeleteByFilterRequest if err := c.ShouldBindJSON(&request); err != nil { setPromptAdminAudit(c, "failed", "prompt_audit_delete_confirmation_invalid", nil) response.ErrorFrom(c, infraerrors.BadRequest("prompt_audit_delete_confirmation_invalid", "删除确认无效或已过期")) return } result, err := h.service.DeleteByFilter(c.Request.Context(), request, adminID(c)) if err != nil { setPromptAdminAudit(c, "failed", "prompt_audit_delete_confirmation_invalid", map[string]any{ "snapshot_max_id": request.SnapshotMaxID, "filter_hash": request.FilterHash, "confirm": request.Confirm, }) response.ErrorFrom(c, infraerrors.BadRequest("prompt_audit_delete_confirmation_invalid", "删除确认无效或已过期")) return } setPromptAdminAudit(c, "success", "", deleteAuditFields(result, map[string]any{ "snapshot_max_id": request.SnapshotMaxID, "filter_hash": request.FilterHash, "confirm": request.Confirm, })) response.Success(c, result) } func setPromptAdminAudit(c *gin.Context, result, errorCode string, fields map[string]any) { details := make(map[string]any, len(fields)+2) details["result"] = result if strings.TrimSpace(errorCode) != "" { details["error_code"] = errorCode } for key, value := range fields { details[key] = value } middleware.SetAuditExtra(c, details) } func configAuditFields(request UpdateConfigRequest, saved *PublicConfig) map[string]any { version := request.ExpectedConfigVersion if saved != nil { version = saved.ConfigVersion } return map[string]any{ "enabled": request.Enabled, "blocking_enabled": request.BlockingEnabled, "blocking_latest_turn_only": request.BlockingLatestTurnOnly, "config_version": version, "endpoint_count": len(request.Endpoints), "scanner_count": len(request.Scanners), "all_groups": request.AllGroups, "group_count": len(request.GroupIDs), } } func deleteAuditFields(result *DeleteResult, base map[string]any) map[string]any { fields := make(map[string]any, len(base)+2) for key, value := range base { fields[key] = value } if result != nil { fields["deleted_events"] = result.DeletedEvents fields["deleted_jobs"] = result.DeletedJobs } return fields } func adminID(c *gin.Context) int64 { subject, ok := middleware.GetAuthSubjectFromContext(c) if !ok { return 0 } return subject.UserID } func eventFilterFromQuery(c *gin.Context) (EventFilter, error) { groupID, err := optionalPositiveInt64Query(c, "group_id") if err != nil { return EventFilter{}, err } userID, err := optionalPositiveInt64Query(c, "user_id") if err != nil { return EventFilter{}, err } apiKeyID, err := optionalPositiveInt64Query(c, "api_key_id") if err != nil { return EventFilter{}, err } filter := EventFilter{ Decision: c.Query("decision"), RiskLevel: c.Query("risk_level"), Endpoint: c.Query("endpoint"), GroupID: groupID, UserID: userID, APIKeyID: apiKeyID, RequestID: c.Query("request_id"), PromptHash: c.Query("prompt_hash"), Keyword: c.Query("keyword"), } if value := strings.TrimSpace(c.Query("start_at")); value != "" { filter.StartAt = parseTimeQuery(value) if filter.StartAt == nil { return EventFilter{}, infraerrors.BadRequest("prompt_audit_invalid_time", "开始时间无效") } } if value := strings.TrimSpace(c.Query("end_at")); value != "" { filter.EndAt = parseTimeQuery(value) if filter.EndAt == nil { return EventFilter{}, infraerrors.BadRequest("prompt_audit_invalid_time", "结束时间无效") } } return filter, nil } func optionalPositiveInt64Query(c *gin.Context, key string) (*int64, error) { value := strings.TrimSpace(c.Query(key)) if value == "" { return nil, nil } parsed, err := strconv.ParseInt(value, 10, 64) if err != nil || parsed <= 0 { return nil, infraerrors.BadRequest("prompt_audit_invalid_filter_id", "事件筛选 ID 无效") } return &parsed, nil } func positiveIntQuery(c *gin.Context, key string, defaultValue, maxValue int) (int, error) { value := strings.TrimSpace(c.Query(key)) if value == "" { return defaultValue, nil } parsed, err := strconv.Atoi(value) if err != nil || parsed <= 0 || (maxValue > 0 && parsed > maxValue) { return 0, infraerrors.BadRequest("prompt_audit_invalid_pagination", "分页参数无效") } return parsed, nil }