Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
160 changes: 122 additions & 38 deletions internal/sql_workbench/service/sql_workbench_service.go
Original file line number Diff line number Diff line change
Expand Up @@ -104,6 +104,7 @@ type SqlWorkbenchService struct {
sqlWorkbenchDatasourceRepo biz.SqlWorkbenchDatasourceRepo
proxyTargetRepo biz.ProxyTargetRepo
cbOperationLogUsecase *biz.CbOperationLogUsecase
maintenanceTimeUsecase *biz.MaintenanceTimeUsecase
sqlResultMasker sqlresultmasker.SQLResultMasker
}

Expand Down Expand Up @@ -184,6 +185,7 @@ func NewAndInitSqlWorkbenchService(logger utilLog.Logger, opts *conf.DMSOptions)
// 初始化操作日志相关
cbOperationLogRepo := storage.NewCbOperationLogRepo(logger, st)
cbOperationLogUsecase := biz.NewCbOperationLogUsecase(logger, cbOperationLogRepo, opPermissionVerifyUsecase, proxyTargetRepo, biz.NewSystemVariableUsecase(logger, storage.NewSystemVariableRepo(logger, st)))
maintenanceTimeUsecase := biz.NewMaintenanceTimeUsecase(logger, opPermissionVerifyUsecase)

return &SqlWorkbenchService{
cfg: opts.SqlWorkBenchOpts,
Expand All @@ -197,6 +199,7 @@ func NewAndInitSqlWorkbenchService(logger utilLog.Logger, opts *conf.DMSOptions)
sqlWorkbenchDatasourceRepo: sqlWorkbenchDatasourceRepo,
proxyTargetRepo: proxyTargetRepo,
cbOperationLogUsecase: cbOperationLogUsecase,
maintenanceTimeUsecase: maintenanceTimeUsecase,
}, nil
}

Expand Down Expand Up @@ -1183,7 +1186,7 @@ func (sqlWorkbenchService *SqlWorkbenchService) AuditMiddleware() echo.Middlewar
// 注意:解析仅服务于审核辅助路径,解析失败不应直接阻塞用户的 SQL 执行;
// 否则一旦中间件辅助能力出错(如 sid 解码失败),用户连查询都跑不了。
// 真正的「未启用审核 / 审核失败」等强策略仍由后续分支按既有 fail-closed 处理。
sql, sidInfo, err := sqlWorkbenchService.parseStreamExecuteRequest(bodyBytes)
sql, sidInfo, isExecuteAnyway, err := sqlWorkbenchService.parseStreamExecuteRequest(bodyBytes)
if err != nil {
sqlWorkbenchService.log.Warnf("failed to parse streamExecute request, skipping audit: %v", err)
return next(c)
Expand Down Expand Up @@ -1245,7 +1248,13 @@ func (sqlWorkbenchService *SqlWorkbenchService) AuditMiddleware() echo.Middlewar
}

// 拦截响应并添加审核结果
return sqlWorkbenchService.interceptAndAddAuditResult(c, next, dmsUserId, auditResult, dbService)
execCtx := &streamExecuteRequestContext{
sql: sql,
datasourceID: datasourceID,
schemaName: schemaName,
isExecuteAnyway: isExecuteAnyway,
}
return sqlWorkbenchService.interceptAndAddAuditResult(c, next, dmsUserId, auditResult, dbService, execCtx)
}
}
}
Expand All @@ -1257,11 +1266,11 @@ type streamExecuteSidInfo struct {
dbID int
}

// parseStreamExecuteRequest 解析 streamExecute 请求体,提取 SQL 与会话 sid 信息
func (sqlWorkbenchService *SqlWorkbenchService) parseStreamExecuteRequest(bodyBytes []byte) (sql string, sidInfo *streamExecuteSidInfo, err error) {
// parseStreamExecuteRequest 解析 streamExecute 请求体,提取 SQL、会话 sid 与 isExecuteAnyway
func (sqlWorkbenchService *SqlWorkbenchService) parseStreamExecuteRequest(bodyBytes []byte) (sql string, sidInfo *streamExecuteSidInfo, isExecuteAnyway bool, err error) {
var requestBody map[string]interface{}
if err := json.Unmarshal(bodyBytes, &requestBody); err != nil {
return "", nil, fmt.Errorf("failed to unmarshal request body: %v", err)
return "", nil, false, fmt.Errorf("failed to unmarshal request body: %v", err)
}

if sqlVal, ok := requestBody["sql"]; ok {
Expand All @@ -1270,6 +1279,8 @@ func (sqlWorkbenchService *SqlWorkbenchService) parseStreamExecuteRequest(bodyBy
}
}

isExecuteAnyway = parseIsExecuteAnywayFromRequest(requestBody)

if sidVal, ok := requestBody["sid"]; ok {
if sidStr, ok := sidVal.(string); ok {
sidInfo, err = sqlWorkbenchService.parseStreamExecuteSid(sidStr)
Expand All @@ -1279,7 +1290,26 @@ func (sqlWorkbenchService *SqlWorkbenchService) parseStreamExecuteRequest(bodyBy
}
}

return sql, sidInfo, nil
return sql, sidInfo, isExecuteAnyway, nil
}

func parseIsExecuteAnywayFromRequest(requestBody map[string]interface{}) bool {
isExecuteAnywayVal, ok := requestBody["isExecuteAnyway"]
if !ok || isExecuteAnywayVal == nil {
return false
}

switch val := isExecuteAnywayVal.(type) {
case bool:
return val
case string:
parsed, err := strconv.ParseBool(val)
if err == nil {
return parsed
}
}

return false
}

// parseStreamExecuteSid 解析 ODC sid(与 pathUtil.generateDatabaseSid 对齐):
Expand Down Expand Up @@ -1441,6 +1471,19 @@ func (sqlWorkbenchService *SqlWorkbenchService) isEnableSQLAudit(dbService *biz.
return dbService.SQLEConfig.AuditEnabled && dbService.SQLEConfig.SQLQueryConfig.AuditEnabled
}

// normalizeSQLEAuditSchemaName 将 ODC SQL Server 的 database.schema 组合归一化为 SQLE 连接库名。
// ODC 将 SQL Server 的 database.schema 平铺为 schema 名(如 TestDB.dbo),
// 而 SQLE 直连审核会把 schema_name 当作连接的数据库名。
func normalizeSQLEAuditSchemaName(dbType, schemaName string) string {
if dbType != string(pkgConst.DBTypeSQLServer) {
return schemaName
}
if idx := strings.Index(schemaName, "."); idx > 0 {
return schemaName[:idx]
}
return schemaName
}

// callSQLEAudit 调用 SQLE 直接审核接口
func (sqlWorkbenchService *SqlWorkbenchService) callSQLEAudit(ctx context.Context, sql string, dbService *biz.DBService, schemaName string) (*cloudbeaver.AuditSQLReply, error) {
// 获取 SQLE 服务地址
Expand All @@ -1450,6 +1493,7 @@ func (sqlWorkbenchService *SqlWorkbenchService) callSQLEAudit(ctx context.Contex
}

sqleAddr := fmt.Sprintf("%s/v2/sql_audit", target.URL.String())
schemaName = normalizeSQLEAuditSchemaName(dbService.DBType, schemaName)

auditReq := cloudbeaver.AuditSQLReq{
InstanceType: dbService.DBType,
Expand Down Expand Up @@ -1480,17 +1524,22 @@ func (sqlWorkbenchService *SqlWorkbenchService) callSQLEAudit(ctx context.Contex
}

// interceptAndAddAuditResult 拦截响应并添加审核结果
func (sqlWorkbenchService *SqlWorkbenchService) interceptAndAddAuditResult(c echo.Context, next echo.HandlerFunc, userId string, auditResult *cloudbeaver.AuditSQLReply, dbService *biz.DBService) error {
// 判断是否需要审批
func (sqlWorkbenchService *SqlWorkbenchService) interceptAndAddAuditResult(c echo.Context, next echo.HandlerFunc, userId string, auditResult *cloudbeaver.AuditSQLReply, dbService *biz.DBService, execCtx *streamExecuteRequestContext) error {
allowQueryWhenLessThanAuditLevel := dbService.GetAllowQueryWhenLessThanAuditLevel()
needApproval := sqlWorkbenchService.shouldRequireApproval(auditResult.Data.SQLResults, allowQueryWhenLessThanAuditLevel)
auditPassed := isAuditSuccessful(auditResult, allowQueryWhenLessThanAuditLevel)

// 如果需要审批,直接返回审核结果,不请求真实的 streamExecute 接口
if needApproval {
if !auditPassed && !execCtx.isExecuteAnyway {
return sqlWorkbenchService.buildAuditResponseWithoutExecution(c, userId, auditResult, dbService)
}

// 不需要审批,执行真实请求并添加审核结果
if blocked, err := sqlWorkbenchService.checkMaintenanceTime(c, userId, auditResult.Data.SQLResults, dbService); blocked || err != nil {
return err
}

if sqlWorkbenchService.shouldExecuteByWorkflow(dbService, auditResult.Data.SQLResults) {
return sqlWorkbenchService.executeNonDQLByWorkflow(c, userId, execCtx, auditResult, dbService)
}

return sqlWorkbenchService.executeAndAddAuditResult(c, next, auditResult, dbService)
}

Expand Down Expand Up @@ -1799,35 +1848,68 @@ func (sqlWorkbenchService *SqlWorkbenchService) mergeSQLEAuditResults(data *Stre
}
}

// shouldRequireApproval 根据审核放行等级判断是否需要审批
func (sqlWorkbenchService *SqlWorkbenchService) shouldRequireApproval(sqlResults []cloudbeaver.AuditSQLResV2, allowQueryWhenLessThanAuditLevel string) bool {
// 如果没有设置审核放行等级,那么直接放行
if allowQueryWhenLessThanAuditLevel == "" {
// isAuditSuccessful 判断审核是否放行,语义对齐 CloudBeaver AuditSQL / IsSuccess
func isAuditSuccessful(auditResult *cloudbeaver.AuditSQLReply, allowQueryWhenLessThanAuditLevel string) bool {
if auditResult == nil || auditResult.Data == nil {
return false
}

// 遍历所有 SQL 审核结果
for _, sqlResult := range sqlResults {
// 检查是否有执行失败的审核项
for _, auditItem := range sqlResult.AuditResult {
if auditItem.ExecutionFailed {
return true
for _, sqlResult := range auditResult.Data.SQLResults {
for _, res := range sqlResult.AuditResult {
if res.ExecutionFailed {
return false
}
}
}

// 比较审核等级:如果 SQL 的审核等级大于允许的等级,则需要审批
// 使用 RuleLevel 的 LessOrEqual 方法进行比较
sqlAuditLevel := dbmodel.RuleLevel(sqlResult.AuditLevel)
allowedLevel := dbmodel.RuleLevel(allowQueryWhenLessThanAuditLevel)
if auditResult.Data.PassRate == 0 {
return cloudbeaver.AllowQuery(allowQueryWhenLessThanAuditLevel, auditResult.Data.SQLResults)
}

// 如果 SQL 的审核等级大于允许的等级(即 !LessOrEqual),则需要审批
if !sqlAuditLevel.LessOrEqual(allowedLevel) {
return true
}
return true
}

// shouldRequireApproval 根据审核放行语义判断是否需要审批
func (sqlWorkbenchService *SqlWorkbenchService) shouldRequireApproval(auditResult *cloudbeaver.AuditSQLReply, allowQueryWhenLessThanAuditLevel string) bool {
return !isAuditSuccessful(auditResult, allowQueryWhenLessThanAuditLevel)
}

// checkMaintenanceTime 检查运维时间管控(ODC 工作台)
// 返回 blocked=true 表示已构造拦截响应,调用方应立即返回
func (sqlWorkbenchService *SqlWorkbenchService) checkMaintenanceTime(c echo.Context, userID string, auditResults []cloudbeaver.AuditSQLResV2, dbService *biz.DBService) (blocked bool, err error) {
if sqlWorkbenchService.maintenanceTimeUsecase == nil {
return false, fmt.Errorf("maintenance time usecase is nil")
}

// 所有 SQL 的审核等级都小于等于允许的等级,不需要审批
return false
if userID == "" {
return false, fmt.Errorf("current user uid is empty")
}

sqlTypes := make([]string, 0, len(auditResults))
for _, r := range auditResults {
sqlTypes = append(sqlTypes, r.SQLType)
}

var sqlQueryConfig *biz.SQLQueryConfig
if dbService != nil && dbService.SQLEConfig != nil {
sqlQueryConfig = dbService.SQLEConfig.SQLQueryConfig
}

allowed, message, checkErr := sqlWorkbenchService.maintenanceTimeUsecase.CheckSQLExecutionAllowed(
c.Request().Context(),
userID,
sqlTypes,
time.Now(),
sqlQueryConfig,
)
if checkErr != nil {
sqlWorkbenchService.log.Errorf("check maintenance time failed: %v", checkErr)
return false, checkErr
}
if !allowed {
return true, sqlWorkbenchService.buildWorkflowErrorResponse(c, message)
}
return false, nil
}

// convertSQLEAuditToViolatedRules 将 SQLE 审核结果转换为 violatedRules 格式
Expand Down Expand Up @@ -1896,12 +1978,14 @@ type StreamExecuteResponse struct {

// StreamExecuteData streamExecute 响应中的 data 字段
type StreamExecuteData struct {
ApprovalRequired bool `json:"approvalRequired"`
LogicalSQL bool `json:"logicalSql"`
RequestID *string `json:"requestId"`
SQLs []StreamExecuteSQLItem `json:"sqls"`
UnauthorizedDBResources interface{} `json:"unauthorizedDBResources"`
ViolatedRules []interface{} `json:"violatedRules"`
ApprovalRequired bool `json:"approvalRequired"`
LogicalSQL bool `json:"logicalSql"`
RequestID *string `json:"requestId"`
SQLs []StreamExecuteSQLItem `json:"sqls"`
UnauthorizedDBResources interface{} `json:"unauthorizedDBResources"`
ViolatedRules []interface{} `json:"violatedRules"`
WorkflowInfo *StreamExecuteWorkflowInfo `json:"workflowInfo,omitempty"`
ErrorMessage string `json:"errorMessage,omitempty"`
}

// StreamExecuteSQLItem SQL 条目
Expand Down
Loading