kra-new/internal/data/data_scope.go

266 lines
9.4 KiB
Go

package data
import (
"database/sql"
"errors"
"kra/internal/biz/system"
"log/slog"
"reflect"
"strings"
"gorm.io/gorm"
"gorm.io/gorm/clause"
"gorm.io/gorm/schema"
)
// registerDataScopeCallbacks installs the global GORM data-scope engine.
// System tables are deliberately excluded: their access is controlled by
// Casbin, while ownership columns on business tables are row-level scope.
type dataScopeAuditEnqueue func(dataAccessLogPO)
var (
errDataScopeRequired = errors.New("受控表访问缺少数据权限上下文")
errInvalidDataScope = errors.New("数据权限范围不合法")
)
func registerDataScopeCallbacks(db *gorm.DB, enqueue dataScopeAuditEnqueue) {
if db == nil {
return
}
q := db.Callback().Query()
if q.Get("data_scope:query") == nil {
_ = q.Before("gorm:query").Register("data_scope:query", applyDataScope("query", enqueue))
}
u := db.Callback().Update()
if u.Get("data_scope:update") == nil {
_ = u.Before("gorm:update").Register("data_scope:update", applyDataScope("update", enqueue))
_ = u.Before("gorm:update").Register("data_scope:stamp_update", stampUpdatedBy)
_ = u.After("gorm:update").Register("data_scope:audit_update", auditBlockedWrite("update", enqueue))
}
d := db.Callback().Delete()
if d.Get("data_scope:delete") == nil {
_ = d.Before("gorm:delete").Register("data_scope:delete", applyDataScope("delete", enqueue))
_ = d.Before("gorm:delete").After("data_scope:delete").Register("data_scope:stamp_delete", stampDeletedBy)
_ = d.After("gorm:delete").Register("data_scope:audit_delete", auditBlockedWrite("delete", enqueue))
}
c := db.Callback().Create()
if c.Get("data_scope:stamp") == nil {
_ = c.Before("gorm:create").Register("data_scope:stamp", stampOwnership(enqueue))
}
}
func hasScopeField(db *gorm.DB, name string) bool {
return db.Statement.Schema != nil && db.Statement.Schema.LookUpField(name) != nil
}
func isControlledTable(db *gorm.DB) bool {
return db.Statement.Schema != nil && !strings.HasPrefix(db.Statement.Table, "sys_") && (hasScopeField(db, "dept_id") || hasScopeField(db, "created_by"))
}
func skipDataScope(db *gorm.DB) bool {
skip, ok := db.Get("data_scope:skip")
value, _ := skip.(bool)
return ok && value
}
func validDataScope(scope system.DataScope) bool {
return scope.UserID != 0 && scope.AuthorityID != 0 && scope.Scope >= 1 && scope.Scope <= 5 && scope.All == (scope.Scope == 1)
}
func applyDataScope(operation string, enqueue dataScopeAuditEnqueue) func(*gorm.DB) {
return func(db *gorm.DB) {
if !isControlledTable(db) {
return
}
if _, done := db.Statement.Clauses["data_scope:applied"]; done {
return
}
if skipDataScope(db) {
return
}
scope, ok := system.DataScopeFromContext(db.Statement.Context)
if !ok {
slog.WarnContext(db.Statement.Context, "数据权限: 业务表访问无身份上下文, 已拒绝", "mod", "data-scope", "table", db.Statement.Table)
recordDataScopeEvent(db, enqueue, "no_identity", operation, "无身份上下文访问受控表, 已拒绝", system.DataScope{})
_ = db.AddError(errDataScopeRequired)
return
}
if !validDataScope(scope) {
recordDataScopeEvent(db, enqueue, "invalid_scope", operation, "数据权限上下文不合法, 已拒绝", scope)
_ = db.AddError(errInvalidDataScope)
return
}
if (operation == "update" || operation == "delete") && !db.AllowGlobalUpdate && !hasWriteConditions(db) {
return
}
db.Statement.Clauses["data_scope:applied"] = clause.Clause{}
if scope.All {
return
}
table := db.Statement.Table
if scope.OwnerUserID != 0 && hasScopeField(db, "created_by") {
db.Where(table+".created_by = ?", scope.OwnerUserID)
return
}
if hasScopeField(db, "dept_id") {
ids := scope.DepartmentIDs
if len(ids) == 0 {
db.Where("1 = 0")
return
}
db.Where(table+".dept_id IN ?", ids)
return
}
db.Where("1 = 0")
}
}
func recordDataScopeEvent(db *gorm.DB, enqueue dataScopeAuditEnqueue, eventType, operation, detail string, scope system.DataScope) {
if enqueue == nil {
return
}
enqueue(dataAccessLogPO{EventType: eventType, TargetTable: db.Statement.Table, Operation: operation, UserID: scope.UserID, AuthorityID: scope.AuthorityID, Scope: scope.Scope, RequestID: scope.RequestID, Method: scope.Method, Path: scope.Path, Detail: detail})
}
func auditBlockedWrite(operation string, enqueue dataScopeAuditEnqueue) func(*gorm.DB) {
return func(db *gorm.DB) {
if _, applied := db.Statement.Clauses["data_scope:applied"]; !applied || db.Error != nil || db.RowsAffected != 0 {
return
}
if scope, ok := system.DataScopeFromContext(db.Statement.Context); ok && !scope.All {
recordDataScopeEvent(db, enqueue, "blocked_write", operation, "数据范围过滤后写操作影响 0 行(疑似越权尝试)", scope)
}
}
}
func stampOwnership(enqueue dataScopeAuditEnqueue) func(*gorm.DB) {
return func(db *gorm.DB) {
if !isControlledTable(db) || skipDataScope(db) {
return
}
scope, ok := system.DataScopeFromContext(db.Statement.Context)
if !ok {
recordDataScopeEvent(db, enqueue, "no_identity", "create", "无身份上下文访问受控表, 已拒绝", system.DataScope{})
_ = db.AddError(errDataScopeRequired)
return
}
if !validDataScope(scope) {
recordDataScopeEvent(db, enqueue, "invalid_scope", "create", "数据权限上下文不合法, 已拒绝", scope)
_ = db.AddError(errInvalidDataScope)
return
}
if hasScopeField(db, "created_by") {
db.Statement.SetColumn("created_by", scope.UserID, true)
}
if hasScopeField(db, "dept_id") {
db.Statement.SetColumn("dept_id", scope.PrimaryDeptID, true)
}
}
}
func stampUpdatedBy(db *gorm.DB) {
stmt := db.Statement
if !isControlledTable(db) || stmt.SkipHooks || !hasScopeField(db, "updated_by") {
return
}
scope, ok := system.DataScopeFromContext(db.Statement.Context)
if !ok || scope.UserID == 0 {
return
}
if _, isMap := stmt.Dest.(map[string]interface{}); !isMap {
if _, isMaps := stmt.Dest.([]map[string]interface{}); !isMaps && stmt.ReflectValue.Kind() == reflect.Struct && !stmt.ReflectValue.CanAddr() {
return
}
}
for _, omitted := range stmt.Omits {
if omitted == "updated_by" {
return
}
}
if len(stmt.Selects) > 0 {
found := false
for _, selected := range stmt.Selects {
if selected == "updated_by" || selected == "*" {
found = true
break
}
}
if !found {
stmt.Selects = append(stmt.Selects, "updated_by")
}
}
stmt.SetColumn("updated_by", scope.UserID, true)
}
func stampDeletedBy(db *gorm.DB) {
stmt := db.Statement
if db.Error != nil || !isControlledTable(db) || stmt.SQL.Len() != 0 || stmt.Unscoped || !hasScopeField(db, "deleted_by") {
return
}
deletedAtType := reflect.TypeOf(gorm.DeletedAt{})
var deletedAt *schema.Field
for _, field := range stmt.Schema.Fields {
if field.FieldType == deletedAtType || field.IndirectFieldType == deletedAtType {
deletedAt = field
break
}
}
if deletedAt == nil {
return
}
if _, customZero := deletedAt.TagSettings["ZEROVALUE"]; customZero {
return
}
scope, ok := system.DataScopeFromContext(stmt.Context)
if !ok || scope.UserID == 0 {
return
}
var primaryExpressions []clause.Expression
_, queryValues := schema.GetIdentityFieldValuesMap(stmt.Context, stmt.ReflectValue, stmt.Schema.PrimaryFields)
column, values := schema.ToQueryValues(stmt.Table, stmt.Schema.PrimaryFieldDBNames, queryValues)
if len(values) > 0 {
primaryExpressions = append(primaryExpressions, clause.IN{Column: column, Values: values})
}
if stmt.ReflectValue.CanAddr() && stmt.Dest != stmt.Model && stmt.Model != nil {
_, queryValues = schema.GetIdentityFieldValuesMap(stmt.Context, reflect.ValueOf(stmt.Model), stmt.Schema.PrimaryFields)
column, values = schema.ToQueryValues(stmt.Table, stmt.Schema.PrimaryFieldDBNames, queryValues)
if len(values) > 0 {
primaryExpressions = append(primaryExpressions, clause.IN{Column: column, Values: values})
}
}
if _, hasWhere := stmt.Clauses["WHERE"]; !hasWhere && len(primaryExpressions) == 0 && !db.AllowGlobalUpdate {
return
}
now := db.NowFunc()
stmt.AddClause(clause.Set{{Column: clause.Column{Name: deletedAt.DBName}, Value: now}, {Column: clause.Column{Name: "deleted_by"}, Value: scope.UserID}})
stmt.SetColumn(deletedAt.DBName, now, true)
stmt.SetColumn("deleted_by", scope.UserID, true)
if len(primaryExpressions) > 0 {
stmt.AddClause(clause.Where{Exprs: primaryExpressions})
}
gorm.SoftDeleteQueryClause{ZeroValue: sql.NullString{Valid: false}, Field: deletedAt}.ModifyStatement(stmt)
stmt.AddClauseIfNotExists(clause.Update{})
stmt.Build(stmt.DB.Callback().Update().Clauses...)
}
func hasWriteConditions(db *gorm.DB) bool {
if value, ok := db.Statement.Clauses["WHERE"]; ok {
if where, ok := value.Expression.(clause.Where); ok && len(where.Exprs) > 0 {
return true
}
}
if db.Statement.Schema == nil || !db.Statement.ReflectValue.IsValid() {
return false
}
_, values := schema.GetIdentityFieldValuesMap(db.Statement.Context, db.Statement.ReflectValue, db.Statement.Schema.PrimaryFields)
if _, query := schema.ToQueryValues(db.Statement.Table, db.Statement.Schema.PrimaryFieldDBNames, values); len(query) > 0 {
return true
}
if db.Statement.ReflectValue.CanAddr() && db.Statement.Dest != db.Statement.Model && db.Statement.Model != nil {
_, values = schema.GetIdentityFieldValuesMap(db.Statement.Context, reflect.ValueOf(db.Statement.Model), db.Statement.Schema.PrimaryFields)
_, query := schema.ToQueryValues(db.Statement.Table, db.Statement.Schema.PrimaryFieldDBNames, values)
return len(query) > 0
}
return false
}