diff --git a/pkg/sql/plan/base_binder.go b/pkg/sql/plan/base_binder.go index 1b6124d900a0f..9a05657f11318 100644 --- a/pkg/sql/plan/base_binder.go +++ b/pkg/sql/plan/base_binder.go @@ -519,10 +519,14 @@ func (b *baseBinder) baseBindColRef(astExpr *tree.UnresolvedName, depth int32, i }) } + preserveSpecialValue := typ != nil && b.mysqlSpecialTargetType != nil && + typ.Enumvalues == b.mysqlSpecialTargetType.Enumvalues && + ((isEnumPlanType(typ) && isEnumPlanType(b.mysqlSpecialTargetType)) || + (isSetPlanType(typ) && isSetPlanType(b.mysqlSpecialTargetType))) // ENUM and SET have distinct storage and display representations. Keep their // display value by default. Numeric and bitwise expression binders explicitly // enable raw storage binding so they follow MySQL's numeric semantics. - if !b.bindRawMySQLSpecialType && isEnumOrSetPlanType(typ) { + if !b.bindRawMySQLSpecialType && !preserveSpecialValue && isEnumOrSetPlanType(typ) { if err != nil { errutil.ReportError(b.GetContext(), err) return diff --git a/pkg/sql/plan/bind_context.go b/pkg/sql/plan/bind_context.go index 91bdd3f996970..cd52850018c9f 100644 --- a/pkg/sql/plan/bind_context.go +++ b/pkg/sql/plan/bind_context.go @@ -26,13 +26,14 @@ import ( func NewBindContext(builder *QueryBuilder, parent *BindContext) *BindContext { bc := &BindContext{ - groupByAst: make(map[string]int32), - groupByParamAst: make(map[string]int32), - aggregateByAst: make(map[string]int32), - sampleByAst: make(map[string]int32), - projectByExpr: make(map[string]int32), - windowByAst: make(map[string]int32), - timeByAst: make(map[string]int32), + outputColumnProvenance: make(map[int32]OutputColumnProvenance), + groupByAst: make(map[string]int32), + groupByParamAst: make(map[string]int32), + aggregateByAst: make(map[string]int32), + sampleByAst: make(map[string]int32), + projectByExpr: make(map[string]int32), + windowByAst: make(map[string]int32), + timeByAst: make(map[string]int32), projectColByAst: make(map[string]int32), @@ -66,6 +67,7 @@ func NewBindContext(builder *QueryBuilder, parent *BindContext) *BindContext { bc.viewChain = append([]string{}, parent.viewChain...) } bc.directView = parent.directView + bc.restoreViewMySQLSpecialTypes = parent.restoreViewMySQLSpecialTypes } return bc diff --git a/pkg/sql/plan/build.go b/pkg/sql/plan/build.go index daab3827030a2..6ff66d0d040b7 100644 --- a/pkg/sql/plan/build.go +++ b/pkg/sql/plan/build.go @@ -39,6 +39,21 @@ func bindAndOptimizeSelectQueryWithValidator( isPrepareStmt bool, skipStats bool, validate func(*Query) error, +) (*Plan, error) { + return bindAndOptimizeSelectQueryWithValidatorAndCapture( + stmtType, ctx, stmt, isPrepareStmt, skipStats, validate, nil, false, + ) +} + +func bindAndOptimizeSelectQueryWithValidatorAndCapture( + stmtType plan.Query_StatementType, + ctx CompilerContext, + stmt *tree.Select, + isPrepareStmt bool, + skipStats bool, + validate func(*Query) error, + capture func(*BindContext), + restoreViewMySQLSpecialTypes bool, ) (*Plan, error) { start := time.Now() defer func() { @@ -47,6 +62,7 @@ func bindAndOptimizeSelectQueryWithValidator( builder := NewQueryBuilder(stmtType, ctx, isPrepareStmt, true) bindCtx := NewBindContext(builder, nil) + bindCtx.restoreViewMySQLSpecialTypes = restoreViewMySQLSpecialTypes if IsSnapshotValid(ctx.GetSnapshot()) { bindCtx.snapshot = ctx.GetSnapshot() } @@ -58,6 +74,9 @@ func bindAndOptimizeSelectQueryWithValidator( builder.skipStats = skipStats rootId = builder.reuseMultiReferenceCTEs(rootId) ctx.SetViews(bindCtx.views) + if capture != nil { + capture(bindCtx) + } builder.qry.Steps = append(builder.qry.Steps, rootId) if validate != nil { diff --git a/pkg/sql/plan/build_ddl.go b/pkg/sql/plan/build_ddl.go index 70b7007b59c99..d78ce66e299a0 100644 --- a/pkg/sql/plan/build_ddl.go +++ b/pkg/sql/plan/build_ddl.go @@ -152,15 +152,24 @@ func genViewTableDef(ctx CompilerContext, stmt *tree.Select, colNames tree.Ident // check view statement var stmtPlan *Plan + var outputColumnProvenance []OutputColumnProvenance + captureColumnTypes := func(bindCtx *BindContext) { + outputColumnProvenance = make([]OutputColumnProvenance, len(bindCtx.headings)) + for i := range outputColumnProvenance { + outputColumnProvenance[i] = bindCtx.outputColumnProvenanceForProject(int32(i)) + } + } var err error switch s := stmt.Select.(type) { case *tree.ParenSelect: - stmtPlan, err = bindAndOptimizeSelectQueryWithValidator(plan.Query_SELECT, ctx, s.Select, false, true, validate) + stmtPlan, err = bindAndOptimizeSelectQueryWithValidatorAndCapture( + plan.Query_SELECT, ctx, s.Select, false, true, validate, captureColumnTypes, true) if err != nil { return nil, err } default: - stmtPlan, err = bindAndOptimizeSelectQueryWithValidator(plan.Query_SELECT, ctx, stmt, false, true, validate) + stmtPlan, err = bindAndOptimizeSelectQueryWithValidatorAndCapture( + plan.Query_SELECT, ctx, stmt, false, true, validate, captureColumnTypes, true) if err != nil { return nil, err } @@ -179,11 +188,17 @@ func genViewTableDef(ctx CompilerContext, stmt *tree.Select, colNames tree.Ident originName = string(colNames[idx]) name = originName } + typ := &expr.Typ + if idx < len(outputColumnProvenance) { + if sourceType := mysqlSpecialTypeFromProvenance(outputColumnProvenance[idx]); sourceType != nil { + typ = sourceType + } + } cols[idx] = &plan.ColDef{ Name: strings.ToLower(name), OriginName: originName, Alg: plan.CompressType_Lz4, - Typ: expr.Typ, + Typ: *typ, Default: &plan.Default{ NullAbility: !expr.Typ.NotNullable, Expr: nil, @@ -246,16 +261,6 @@ func genAsSelectCols(ctx CompilerContext, stmt *tree.Select, isPrepareStmt bool) builder := NewQueryBuilder(plan.Query_SELECT, ctx, isPrepareStmt, false) bindCtx := NewBindContext(builder, nil) - getTblAndColName := func(relPos, colPos int32) (string, string) { - name := builder.nameByColRef[[2]int32{relPos, colPos}] - // name pattern: tableName.colName - splits := strings.Split(name, ".") - if len(splits) < 2 { - return "", "" - } - return splits[0], splits[1] - } - if s, ok := stmt.Select.(*tree.ParenSelect); ok { stmt = s.Select } @@ -269,21 +274,13 @@ func genAsSelectCols(ctx CompilerContext, stmt *tree.Select, isPrepareStmt bool) for i, expr := range rootNode.ProjectList { defaultVal := "" typ := &expr.Typ - switch e := expr.Expr.(type) { - case *plan.Expr_Col: - tblName, colName := getTblAndColName(e.Col.RelPos, e.Col.ColPos) - if binding, ok := bindCtx.bindingByTable[tblName]; ok { - defaultVal = binding.defaults[binding.colIdByName[colName]] - } - case *plan.Expr_F: - // enum - if e.F.Func.ObjName == moEnumCastIndexToValueFun || e.F.Func.ObjName == moSetCastIndexToValueFun { - // cast_index_to_value('apple,banana,orange', cast(col_name as T_uint16)) - colRef := e.F.Args[1].Expr.(*plan.Expr_Col).Col - tblName, colName := getTblAndColName(colRef.RelPos, colRef.ColPos) - if binding, ok := bindCtx.bindingByTable[tblName]; ok { - typ = binding.types[binding.colIdByName[colName]] - } + provenance := bindCtx.outputColumnProvenanceForProject(int32(i)) + if provenance.State == ProvenanceSingleSource && provenance.Source != nil { + if provenance.CanInheritSourceDefault && provenance.Source.Metadata.HasDefault { + defaultVal = provenance.Source.Metadata.DefaultOriginString + } + if isEnumOrSetPlanType(&provenance.Source.Metadata.Typ) { + typ = &provenance.Source.Metadata.Typ } } diff --git a/pkg/sql/plan/build_ddl_test.go b/pkg/sql/plan/build_ddl_test.go index f36730e6b1b2e..636bc6da006f3 100644 --- a/pkg/sql/plan/build_ddl_test.go +++ b/pkg/sql/plan/build_ddl_test.go @@ -331,6 +331,526 @@ func TestBuildCreateViewExplicitColumnList(t *testing.T) { }) } +func addMySQLSpecialTypeColumns(ctx *MockCompilerContext) { + ctx.tables["nation"].Cols = append(ctx.tables["nation"].Cols, + &plan.ColDef{ + Name: "priority", + Typ: plan.Type{ + Id: int32(types.T_enum), + Enumvalues: "low,medium,high", + NotNullable: true, + }, + }, + &plan.ColDef{ + Name: "flags", + Typ: plan.Type{ + Id: int32(types.T_uint64), + Enumvalues: "red,green,blue", + }, + }, + ) +} + +func TestBuildCreateViewPreservesMySQLSpecialColumnTypes(t *testing.T) { + const rootSQL = "create view v (renamed_priority, renamed_flags, renamed_name) as " + + "select priority, flags, n_name from nation" + ctx := &rootSQLCompilerContext{ + MockCompilerContext: NewMockCompilerContext(false), + rootSQL: rootSQL, + } + addMySQLSpecialTypeColumns(ctx.MockCompilerContext) + + stmt, err := parsers.ParseOne(t.Context(), dialect.MYSQL, rootSQL, 1) + require.NoError(t, err) + defer stmt.Free() + + p, err := BuildPlan(ctx, stmt, false) + require.NoError(t, err) + cols := p.GetDdl().GetCreateView().GetTableDef().GetCols() + require.Len(t, cols, 3) + priorityType := cols[0].GetTyp() + flagsType := cols[1].GetTyp() + nameType := cols[2].GetTyp() + require.Equal(t, "renamed_priority", cols[0].GetName()) + require.Equal(t, int32(types.T_enum), priorityType.GetId()) + require.Equal(t, "low,medium,high", priorityType.GetEnumvalues()) + require.True(t, priorityType.GetNotNullable()) + require.Equal(t, "renamed_flags", cols[1].GetName()) + require.Equal(t, int32(types.T_uint64), flagsType.GetId()) + require.Equal(t, "red,green,blue", flagsType.GetEnumvalues()) + require.False(t, flagsType.GetNotNullable()) + require.Equal(t, "renamed_name", cols[2].GetName()) + require.Equal(t, int32(types.T_varchar), nameType.GetId()) +} + +func TestBuildCreateViewTracksMySQLSpecialColumnTypeProvenance(t *testing.T) { + tests := []struct { + name string + selectSQL string + wantSpecialType bool + }{ + {name: "direct", selectSQL: "select priority, flags from nation", wantSpecialType: true}, + {name: "order by", selectSQL: "select priority, flags from nation order by priority, flags", wantSpecialType: true}, + {name: "order by null", selectSQL: "select priority, flags from nation order by null", wantSpecialType: true}, + {name: "group by", selectSQL: "select priority, flags from nation group by priority, flags", wantSpecialType: true}, + {name: "distinct", selectSQL: "select distinct priority, flags from nation", wantSpecialType: true}, + {name: "derived table", selectSQL: "select priority, flags from (select priority, flags from nation) d", wantSpecialType: true}, + {name: "cte", selectSQL: "with d as (select priority, flags from nation) select priority, flags from d", wantSpecialType: true}, + {name: "derived table order by", selectSQL: "select priority, flags from (select priority, flags from nation) d order by flags", wantSpecialType: true}, + {name: "cte order by", selectSQL: "with d as (select priority, flags from nation) select priority, flags from d order by flags", wantSpecialType: true}, + {name: "alias", selectSQL: "select priority as p, flags as f from nation", wantSpecialType: true}, + {name: "same arms union distinct", selectSQL: "select priority, flags from nation union select priority, flags from nation"}, + {name: "union all", selectSQL: "select priority, flags from nation union all select priority, flags from nation"}, + {name: "recursive cte", selectSQL: "with recursive d(priority, flags) as (select priority, flags from nation union all select priority, flags from d where false) select priority, flags from d"}, + {name: "string expressions", selectSQL: "select concat(priority, ''), concat(flags, '') from nation"}, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + rootSQL := "create view v as " + test.selectSQL + ctx := &rootSQLCompilerContext{ + MockCompilerContext: NewMockCompilerContext(false), + rootSQL: rootSQL, + } + addMySQLSpecialTypeColumns(ctx.MockCompilerContext) + stmt, err := parsers.ParseOne(t.Context(), dialect.MYSQL, rootSQL, 1) + require.NoError(t, err) + defer stmt.Free() + + viewPlan, err := BuildPlan(ctx, stmt, false) + require.NoError(t, err) + cols := viewPlan.GetDdl().GetCreateView().GetTableDef().GetCols() + require.Len(t, cols, 2) + if test.wantSpecialType { + require.True(t, isEnumPlanType(&cols[0].Typ)) + require.True(t, isSetPlanType(&cols[1].Typ)) + } else { + require.Equal(t, int32(types.T_varchar), cols[0].Typ.GetId()) + require.Equal(t, int32(types.T_varchar), cols[1].Typ.GetId()) + } + }) + } +} + +func TestBuildCTASPreservesMySQLSpecialColumnTypes(t *testing.T) { + const sql = "create table copied as select priority, flags, n_name from nation" + ctx := NewMockCompilerContext(false) + addMySQLSpecialTypeColumns(ctx) + stmt, err := parsers.ParseOne(t.Context(), dialect.MYSQL, sql, 1) + require.NoError(t, err) + defer stmt.Free() + + p, err := BuildPlan(ctx, stmt, false) + require.NoError(t, err) + cols := p.GetDdl().GetCreateTable().GetTableDef().GetCols() + require.GreaterOrEqual(t, len(cols), 3) + require.True(t, isEnumPlanType(&cols[0].Typ)) + require.Equal(t, "low,medium,high", cols[0].Typ.GetEnumvalues()) + require.True(t, isSetPlanType(&cols[1].Typ)) + require.Equal(t, "red,green,blue", cols[1].Typ.GetEnumvalues()) + require.Equal(t, int32(types.T_varchar), cols[2].Typ.GetId()) +} + +func TestViewRebindPreservesMySQLSpecialColumnSemantics(t *testing.T) { + const createViewSQL = "create view v_enum_set as select priority, flags, n_name from nation" + ctx := NewMockCompilerContext(false) + addMySQLSpecialTypeColumns(ctx) + createCtx := &rootSQLCompilerContext{MockCompilerContext: ctx, rootSQL: createViewSQL} + stmt, err := parsers.ParseOne(t.Context(), dialect.MYSQL, createViewSQL, 1) + require.NoError(t, err) + createPlan, err := BuildPlan(createCtx, stmt, false) + stmt.Free() + require.NoError(t, err) + + viewDef := DeepCopyTableDef(createPlan.GetDdl().GetCreateView().GetTableDef(), true) + viewDef.Name = "v_enum_set" + viewDef.DbName = "tpch" + viewDef.TableType = catalog.SystemViewRel + ctx.tables["v_enum_set"] = viewDef + ctx.objects["v_enum_set"] = &plan.ObjectRef{SchemaName: "tpch", ObjName: "v_enum_set"} + + stmt, err = parsers.ParseOne(t.Context(), dialect.MYSQL, + "select priority from v_enum_set order by priority", 1) + require.NoError(t, err) + selectPlan, err := BuildPlan(ctx, stmt, false) + stmt.Free() + require.NoError(t, err) + + var sortKey *plan.Expr + for _, node := range selectPlan.GetQuery().GetNodes() { + if node.GetNodeType() == plan.Node_SORT { + require.Len(t, node.GetOrderBy(), 1) + sortKey = node.GetOrderBy()[0].GetExpr() + break + } + } + require.NotNil(t, sortKey) + sortType := sortKey.GetTyp() + require.Equal(t, int32(types.T_enum), sortType.GetId()) + require.Equal(t, "low,medium,high", sortType.GetEnumvalues()) + query := selectPlan.GetQuery() + require.Len(t, query.GetSteps(), 1) + resultNode := query.GetNodes()[query.GetSteps()[0]] + require.Len(t, resultNode.GetProjectList(), 1) + resultType := resultNode.GetProjectList()[0].GetTyp() + require.Equal(t, int32(types.T_varchar), resultType.GetId()) + + stmt, err = parsers.ParseOne(t.Context(), dialect.MYSQL, + "select flags from v_enum_set", 1) + require.NoError(t, err) + rawSetPlan, err := BuildPlan(ctx, stmt, false) + stmt.Free() + require.NoError(t, err) + setDisplayFound := false + for _, node := range rawSetPlan.GetQuery().GetNodes() { + for _, project := range node.GetProjectList() { + fn := project.GetF() + if fn == nil { + continue + } + require.NotEqual(t, moSetCastValueToIndexFun, fn.GetFunc().GetObjName(), + "a direct view projection must not round-trip a SET bitmap through its display string") + if fn.GetFunc().GetObjName() == moSetCastIndexToValueFun { + setDisplayFound = true + require.Len(t, fn.GetArgs(), 2) + require.True(t, isSetPlanType(&fn.GetArgs()[1].Typ)) + } + } + } + require.True(t, setDisplayFound) + + stmt, err = parsers.ParseOne(t.Context(), dialect.MYSQL, + "create table copied_from_view as select priority, flags, n_name from v_enum_set", 1) + require.NoError(t, err) + ctasPlan, err := BuildPlan(ctx, stmt, false) + stmt.Free() + require.NoError(t, err) + cols := ctasPlan.GetDdl().GetCreateTable().GetTableDef().GetCols() + require.GreaterOrEqual(t, len(cols), 3) + require.True(t, isEnumPlanType(&cols[0].Typ)) + require.Equal(t, "low,medium,high", cols[0].Typ.GetEnumvalues()) + require.True(t, isSetPlanType(&cols[1].Typ)) + require.Equal(t, "red,green,blue", cols[1].Typ.GetEnumvalues()) + require.Equal(t, int32(types.T_varchar), cols[2].Typ.GetId()) + + ctasDef := DeepCopyTableDef(ctasPlan.GetDdl().GetCreateTable().GetTableDef(), true) + ctasDef.Name = "copied_from_view" + ctasDef.DbName = "tpch" + ctx.tables[ctasDef.Name] = ctasDef + ctx.objects[ctasDef.Name] = &plan.ObjectRef{SchemaName: "tpch", ObjName: ctasDef.Name} + stmt, err = parsers.ParseOne(t.Context(), dialect.MYSQL, + ctasPlan.GetDdl().GetCreateTable().GetCreateAsSelectSql(), 1) + require.NoError(t, err) + insertPlan, err := BuildPlan(ctx, stmt, false) + stmt.Free() + require.NoError(t, err) + for _, node := range insertPlan.GetQuery().GetNodes() { + for _, project := range node.GetProjectList() { + if fn := project.GetF(); fn != nil { + require.NotEqual(t, moSetCastValueToIndexFun, fn.GetFunc().GetObjName(), + "CTAS INSERT must retain the projected SET bitmap: node=%d type=%s expr=%s", + node.GetNodeId(), node.GetNodeType().String(), project.String()) + } + } + } + + stmt, err = parsers.ParseOne(t.Context(), dialect.MYSQL, + "insert into copied_from_view (priority, flags, n_name) "+ + "select priority, concat(flags, ',green'), n_name from v_enum_set", 1) + require.NoError(t, err) + nestedPlan, err := BuildPlan(ctx, stmt, false) + stmt.Free() + require.NoError(t, err) + nestedDisplayFound := false + for _, node := range nestedPlan.GetQuery().GetNodes() { + for _, project := range node.GetProjectList() { + walkPlanExpr(project, func(expr *plan.Expr) { + if fn := expr.GetF(); fn != nil && fn.GetFunc().GetObjName() == moSetCastIndexToValueFun { + nestedDisplayFound = true + } + }) + } + } + require.True(t, nestedDisplayFound, + "a SET column nested in CONCAT must keep its SQL-visible string semantics") +} + +func TestViewRebindPreservesTransparentMySQLSpecialColumnTypes(t *testing.T) { + tests := []struct { + name string + selectSQL string + wantSpecialType bool + }{ + {name: "derived table", selectSQL: "select priority, flags from (select priority, flags from nation) d", wantSpecialType: true}, + {name: "cte", selectSQL: "with d as (select priority, flags from nation) select priority, flags from d", wantSpecialType: true}, + {name: "order by", selectSQL: "select priority, flags from nation order by flags", wantSpecialType: true}, + {name: "derived table order by", selectSQL: "select priority, flags from (select priority, flags from nation) d order by flags", wantSpecialType: true}, + {name: "cte order by", selectSQL: "with d as (select priority, flags from nation) select priority, flags from d order by flags", wantSpecialType: true}, + {name: "union all", selectSQL: "select priority, flags from nation union all select priority, flags from nation"}, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + createViewSQL := "create view v as " + test.selectSQL + ctx := NewMockCompilerContext(false) + addMySQLSpecialTypeColumns(ctx) + createCtx := &rootSQLCompilerContext{MockCompilerContext: ctx, rootSQL: createViewSQL} + stmt, err := parsers.ParseOne(t.Context(), dialect.MYSQL, createViewSQL, 1) + require.NoError(t, err) + createPlan, err := BuildPlan(createCtx, stmt, false) + stmt.Free() + require.NoError(t, err) + + viewDef := DeepCopyTableDef(createPlan.GetDdl().GetCreateView().GetTableDef(), true) + viewDef.Name = "v" + viewDef.DbName = "tpch" + viewDef.TableType = catalog.SystemViewRel + ctx.tables[viewDef.Name] = viewDef + ctx.objects[viewDef.Name] = &plan.ObjectRef{SchemaName: "tpch", ObjName: viewDef.Name} + + stmt, err = parsers.ParseOne(t.Context(), dialect.MYSQL, + "create table copied as select priority, flags from v", 1) + require.NoError(t, err) + ctasPlan, err := BuildPlan(ctx, stmt, false) + stmt.Free() + require.NoError(t, err) + cols := ctasPlan.GetDdl().GetCreateTable().GetTableDef().GetCols() + require.GreaterOrEqual(t, len(cols), 2) + if test.wantSpecialType { + require.True(t, isEnumPlanType(&cols[0].Typ)) + require.True(t, isSetPlanType(&cols[1].Typ)) + for _, node := range ctasPlan.GetQuery().GetNodes() { + for _, project := range node.GetProjectList() { + walkPlanExpr(project, func(expr *plan.Expr) { + if fn := expr.GetF(); fn != nil { + require.NotEqual(t, moSetCastValueToIndexFun, fn.GetFunc().GetObjName(), + "transparent View CTAS must not round-trip a SET bitmap") + } + }) + } + } + + stmt, err = parsers.ParseOne(t.Context(), dialect.MYSQL, + "select cast(flags as unsigned) from v", 1) + require.NoError(t, err) + castPlan, err := BuildPlan(ctx, stmt, false) + stmt.Free() + require.NoError(t, err) + for _, node := range castPlan.GetQuery().GetNodes() { + for _, project := range node.GetProjectList() { + walkPlanExpr(project, func(expr *plan.Expr) { + if fn := expr.GetF(); fn != nil { + require.NotEqual(t, moSetCastIndexToValueFun, fn.GetFunc().GetObjName(), + "numeric View consumer must receive the raw SET bitmap") + } + }) + } + } + } else { + require.Equal(t, int32(types.T_varchar), cols[0].Typ.GetId()) + require.Equal(t, int32(types.T_varchar), cols[1].Typ.GetId()) + } + }) + } +} + +func TestViewSpecialTypeBoundaryCanonicalizesSemanticResults(t *testing.T) { + for _, test := range []struct { + name string + selectSQL string + }{ + {name: "distinct", selectSQL: "select distinct flags from nation"}, + {name: "group by", selectSQL: "select flags from nation group by flags"}, + {name: "group by order", selectSQL: "select flags from nation group by flags order by flags"}, + {name: "derived distinct", selectSQL: "select flags from (select distinct flags from nation) d"}, + } { + t.Run(test.name, func(t *testing.T) { + createViewSQL := "create view v_semantic_set as " + test.selectSQL + ctx := NewMockCompilerContext(false) + addMySQLSpecialTypeColumns(ctx) + createCtx := &rootSQLCompilerContext{MockCompilerContext: ctx, rootSQL: createViewSQL} + stmt, err := parsers.ParseOne(t.Context(), dialect.MYSQL, createViewSQL, 1) + require.NoError(t, err) + createPlan, err := BuildPlan(createCtx, stmt, false) + stmt.Free() + require.NoError(t, err) + + viewDef := DeepCopyTableDef(createPlan.GetDdl().GetCreateView().GetTableDef(), true) + viewDef.Name = "v_semantic_set" + viewDef.DbName = "tpch" + viewDef.TableType = catalog.SystemViewRel + ctx.tables[viewDef.Name] = viewDef + ctx.objects[viewDef.Name] = &plan.ObjectRef{SchemaName: "tpch", ObjName: viewDef.Name} + + stmt, err = parsers.ParseOne(t.Context(), dialect.MYSQL, "select flags from v_semantic_set", 1) + require.NoError(t, err) + queryPlan, err := BuildPlan(ctx, stmt, false) + stmt.Free() + require.NoError(t, err) + + setDisplayProjects := 0 + setCanonicalProjects := 0 + semanticStringInput := false + for _, node := range queryPlan.GetQuery().GetNodes() { + if node.GetNodeType() == plan.Node_AGG { + for _, group := range node.GetGroupBy() { + if types.T(group.Typ.Id).IsMySQLString() { + semanticStringInput = true + } + } + } + for _, project := range node.GetProjectList() { + if fn := project.GetF(); fn != nil { + switch fn.GetFunc().GetObjName() { + case moSetCastIndexToValueFun: + setDisplayProjects++ + case moSetCastValueToIndexFun: + setCanonicalProjects++ + } + } + } + } + require.GreaterOrEqual(t, setDisplayProjects, 1, + "semantic operator must consume the SQL-visible SET value") + require.True(t, semanticStringInput, + "GROUP BY/DISTINCT must operate on the SQL-visible string type") + require.GreaterOrEqual(t, setCanonicalProjects, 1, + "completed semantic View boundary must canonically re-encode SET") + require.True(t, isSetPlanType(&viewDef.Cols[0].Typ)) + }) + } +} + +func TestOutputColumnProvenanceCarriesSourceAndClearsSemanticBoundaries(t *testing.T) { + ctx := NewMockCompilerContext(false) + addMySQLSpecialTypeColumns(ctx) + ctx.tables["nation"].Cols[0].Default = &plan.Default{OriginString: "'ALGERIA'"} + + tests := []struct { + name string + sql string + wantState ProvenanceState + wantDefault string + canInheritDefault bool + }{ + {name: "direct", sql: "select n_nationkey from nation", wantState: ProvenanceSingleSource, wantDefault: "'ALGERIA'", canInheritDefault: true}, + {name: "alias derived", sql: "select k from (select n_nationkey as k from nation) d", wantState: ProvenanceSingleSource, wantDefault: "'ALGERIA'"}, + {name: "non recursive cte", sql: "with d as (select n_nationkey as k from nation) select k from d", wantState: ProvenanceSingleSource, wantDefault: "'ALGERIA'"}, + {name: "expression", sql: "select n_nationkey + 0 from nation", wantState: ProvenanceNone}, + {name: "same arms union distinct", sql: "select n_nationkey from nation union select n_nationkey from nation", wantState: ProvenanceNone}, + {name: "union all", sql: "select n_nationkey from nation union all select n_nationkey from nation", wantState: ProvenanceNone}, + {name: "recursive cte", sql: "with recursive d(k) as (select n_nationkey from nation union all select k from d where false) select k from d", wantState: ProvenanceNone}, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + stmt, err := parsers.ParseOne(t.Context(), dialect.MYSQL, test.sql, 1) + require.NoError(t, err) + defer stmt.Free() + selectStmt := stmt.(*tree.Select) + builder := NewQueryBuilder(plan.Query_SELECT, ctx, false, false) + bindCtx := NewBindContext(builder, nil) + _, err = builder.bindSelect(selectStmt, bindCtx, true) + require.NoError(t, err) + + provenance := bindCtx.outputColumnProvenanceForProject(0) + require.Equal(t, test.wantState, provenance.State) + if test.wantState == ProvenanceSingleSource { + require.NotNil(t, provenance.Source) + require.Equal(t, test.wantDefault, provenance.Source.Metadata.DefaultOriginString) + require.Equal(t, test.canInheritDefault, provenance.CanInheritSourceDefault) + require.NotZero(t, provenance.Source.RelPos) + } else { + require.Nil(t, provenance.Source) + } + }) + } +} + +func TestBuildCTASConsumesOutputColumnProvenance(t *testing.T) { + ctx := NewMockCompilerContext(false) + ctx.tables["nation"].Cols[0].Default = &plan.Default{OriginString: "'ALGERIA'"} + + tests := []struct { + name string + selectSQL string + wantDefault string + }{ + {name: "direct alias", selectSQL: "select n_nationkey as k from nation", wantDefault: "'ALGERIA'"}, + {name: "derived", selectSQL: "select k from (select n_nationkey as k from nation) d"}, + {name: "cte", selectSQL: "with d as (select n_nationkey as k from nation) select k from d"}, + {name: "expression", selectSQL: "select n_nationkey + 0 as k from nation"}, + {name: "union", selectSQL: "select n_nationkey as k from nation union all select n_nationkey from nation"}, + } + + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + sql := "create table copied as " + test.selectSQL + stmt, err := parsers.ParseOne(t.Context(), dialect.MYSQL, sql, 1) + require.NoError(t, err) + defer stmt.Free() + p, err := BuildPlan(ctx, stmt, false) + require.NoError(t, err) + cols := p.GetDdl().GetCreateTable().GetTableDef().GetCols() + require.NotEmpty(t, cols) + require.Equal(t, test.wantDefault, cols[0].GetDefault().GetOriginString()) + }) + } +} + +func TestOutputColumnProvenanceSnapshotsCatalogMetadataOnce(t *testing.T) { + ctx := NewMockCompilerContext(false) + addMySQLSpecialTypeColumns(ctx) + priorityCol := ctx.tables["nation"].Cols[len(ctx.tables["nation"].Cols)-2] + priorityCol.Default = &plan.Default{OriginString: "'low'"} + + stmt, err := parsers.ParseOne(t.Context(), dialect.MYSQL, "select priority from nation", 1) + require.NoError(t, err) + defer stmt.Free() + builder := NewQueryBuilder(plan.Query_SELECT, ctx, false, false) + bindCtx := NewBindContext(builder, nil) + _, err = builder.bindSelect(stmt.(*tree.Select), bindCtx, true) + require.NoError(t, err) + provenance := bindCtx.outputColumnProvenanceForProject(0) + require.Equal(t, ProvenanceSingleSource, provenance.State) + require.NotNil(t, provenance.Source) + + priorityCol.Typ.Enumvalues = "changed" + priorityCol.Default.OriginString = "'changed'" + require.Equal(t, "low,medium,high", provenance.Source.Metadata.Typ.Enumvalues) + require.True(t, provenance.Source.Metadata.HasDefault) + require.Equal(t, "'low'", provenance.Source.Metadata.DefaultOriginString) +} + +func TestTransparentOutputSourceExprRejectsSemanticExpressions(t *testing.T) { + enumType := plan.Type{Id: int32(types.T_enum), Enumvalues: "low,high"} + valid := &plan.Expr{ + Expr: &plan.Expr_F{F: &plan.Function{ + Func: &plan.ObjectRef{ObjName: moEnumCastIndexToValueFun}, + Args: []*plan.Expr{ + {Typ: plan.Type{Id: int32(types.T_varchar)}}, + {Typ: enumType, Expr: &plan.Expr_Col{Col: &plan.ColRef{RelPos: 1, ColPos: 2}}}, + }, + }}, + } + + got, ok := transparentOutputSourceExpr(valid) + require.True(t, ok) + require.Equal(t, enumType, got.Typ) + + for _, mutate := range []func(*plan.Expr){ + func(expr *plan.Expr) { expr.GetF().Args = expr.GetF().Args[:1] }, + func(expr *plan.Expr) { expr.GetF().Args[1].Expr = nil }, + func(expr *plan.Expr) { expr.GetF().Args[1].Typ.Id = int32(types.T_varchar) }, + func(expr *plan.Expr) { expr.GetF().Func.ObjName = "concat" }, + } { + expr := DeepCopyExpr(valid) + mutate(expr) + _, ok = transparentOutputSourceExpr(expr) + require.False(t, ok) + } +} + func TestBuildCreateViewRejectsTemporaryTable(t *testing.T) { tests := []string{ "create view v as select * from nation", diff --git a/pkg/sql/plan/build_test.go b/pkg/sql/plan/build_test.go index 33aa608ff11b9..c9fb0db8771e6 100644 --- a/pkg/sql/plan/build_test.go +++ b/pkg/sql/plan/build_test.go @@ -896,6 +896,19 @@ func TestInsertSelectProjectedSetUsesStoredBitmap(t *testing.T) { require.True(t, planHasPlainUint64ColRef(logicPlan)) } +func TestInsertSelectSetTargetRejectsUnknownSourceColumn(t *testing.T) { + mock := NewMockOptimizer(true) + addSetBitmapDestinationForTest(mock) + mock.ctxt.tables["set_bitmap_destination"].Cols[1].Typ.Enumvalues = "a,b" + + _, err := runOneStmt( + mock, + t, + "insert into set_bitmap_destination(id, bitmap) select n_nationkey, missing from nation", + ) + require.ErrorContains(t, err, "column missing does not exist") +} + func addSetBitmapDestinationForTest(mock *MockOptimizer) { const tableName = "set_bitmap_destination" idType := plan.Type{Id: int32(types.T_int32), NotNullable: true} diff --git a/pkg/sql/plan/mysql_special_types.go b/pkg/sql/plan/mysql_special_types.go index cbe925933a3b6..f7726797545c2 100644 --- a/pkg/sql/plan/mysql_special_types.go +++ b/pkg/sql/plan/mysql_special_types.go @@ -430,6 +430,60 @@ func (bc *BindContext) mysqlSpecialOrderTypeForProject(colPos int32) *plan.Type return bc.mysqlSpecialOrderTypeForExpr(bc.projects[colPos]) } +func (bc *BindContext) setMySQLSpecialCanonicalType(colPos int32, typ *plan.Type) { + if bc.mysqlSpecialCanonicalTypes == nil { + bc.mysqlSpecialCanonicalTypes = make(map[int32]*plan.Type) + } + bc.mysqlSpecialCanonicalTypes[colPos] = DeepCopyType(typ) +} + +func (bc *BindContext) mysqlSpecialCanonicalTypeForExpr(expr *plan.Expr) *plan.Type { + if expr == nil { + return nil + } + col := expr.GetCol() + if col == nil { + return nil + } + if col.RelPos == bc.projectTag && col.ColPos >= 0 && int(col.ColPos) < len(bc.projects) { + if typ, recorded := bc.mysqlSpecialCanonicalTypes[col.ColPos]; recorded { + return DeepCopyType(typ) + } + project := bc.projects[col.ColPos] + if project == nil { + return nil + } + if projectCol := project.GetCol(); projectCol != nil && + projectCol.RelPos == col.RelPos && projectCol.ColPos == col.ColPos { + return nil + } + return bc.mysqlSpecialCanonicalTypeForExpr(project) + } + binding := bc.bindingByTag[col.RelPos] + if binding == nil || col.ColPos < 0 || int(col.ColPos) >= len(binding.mysqlSpecialCanonicalTypes) { + return nil + } + return DeepCopyType(binding.mysqlSpecialCanonicalTypes[col.ColPos]) +} + +func (bc *BindContext) mysqlSpecialCanonicalTypeForProject(colPos int32) *plan.Type { + if typ, recorded := bc.mysqlSpecialCanonicalTypes[colPos]; recorded { + return DeepCopyType(typ) + } + if colPos < 0 || int(colPos) >= len(bc.projects) { + return nil + } + return bc.mysqlSpecialCanonicalTypeForExpr(bc.projects[colPos]) +} + +func mysqlSpecialTypeFromProvenance(provenance OutputColumnProvenance) *plan.Type { + if provenance.State != ProvenanceSingleSource || provenance.Source == nil || + !isEnumOrSetPlanType(&provenance.Source.Metadata.Typ) { + return nil + } + return DeepCopyType(&provenance.Source.Metadata.Typ) +} + func mysqlSpecialOrderTypesCompatible(left, right *plan.Type) bool { return left != nil && right != nil && left.Id == right.Id && left.Enumvalues == right.Enumvalues } diff --git a/pkg/sql/plan/output_column_provenance.go b/pkg/sql/plan/output_column_provenance.go new file mode 100644 index 0000000000000..4e9341f91b324 --- /dev/null +++ b/pkg/sql/plan/output_column_provenance.go @@ -0,0 +1,159 @@ +// Copyright 2026 Matrix Origin +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package plan + +import "github.com/matrixorigin/matrixone/pkg/pb/plan" + +type ProvenanceState uint8 + +const ( + ProvenanceUnknown ProvenanceState = iota + ProvenanceNone + ProvenanceSingleSource +) + +// SourceColumnMetadata is an immutable planner-local snapshot. It deliberately +// contains only metadata consumed by output-schema builders, so transparent +// query boundaries can share it without retaining or repeatedly copying a +// complete catalog ColDef. +type SourceColumnMetadata struct { + Typ plan.Type + HasDefault bool + DefaultOriginString string +} + +type SourceColumn struct { + RelPos int32 + ColPos int32 + TableID uint64 + Metadata SourceColumnMetadata +} + +type OutputColumnProvenance struct { + State ProvenanceState + Source *SourceColumn + CanInheritSourceDefault bool +} + +func snapshotSourceColumnMetadata(col *plan.ColDef) SourceColumnMetadata { + typ := col.Typ + metadata := SourceColumnMetadata{ + Typ: plan.Type{ + Id: typ.Id, + NotNullable: typ.NotNullable, + AutoIncr: typ.AutoIncr, + Width: typ.Width, + Scale: typ.Scale, + Table: typ.Table, + Enumvalues: typ.Enumvalues, + }, + } + if col.Default != nil { + metadata.HasDefault = true + metadata.DefaultOriginString = col.Default.OriginString + } + return metadata +} + +// transparentOutputSourceExpr unwraps planner display adapters that preserve +// the identity of one source column. All other functions are semantic +// expressions and therefore clear lineage. +func transparentOutputSourceExpr(expr *plan.Expr) (*plan.Expr, bool) { + if expr == nil { + return nil, false + } + fn := expr.GetF() + if fn == nil || fn.Func == nil || len(fn.Args) != 2 || fn.Args[1] == nil || fn.Args[1].GetCol() == nil { + return nil, false + } + + sourceExpr := fn.Args[1] + switch fn.Func.ObjName { + case moEnumCastIndexToValueFun: + return sourceExpr, isEnumPlanType(&sourceExpr.Typ) + case moSetCastIndexToValueFun: + return sourceExpr, isSetPlanType(&sourceExpr.Typ) + default: + return nil, false + } +} + +func (bc *BindContext) outputColumnProvenanceForExpr(expr *plan.Expr) OutputColumnProvenance { + if expr == nil { + return OutputColumnProvenance{State: ProvenanceNone} + } + if sourceExpr, ok := transparentOutputSourceExpr(expr); ok { + return bc.outputColumnProvenanceForExpr(sourceExpr) + } + + col := expr.GetCol() + if col == nil { + return OutputColumnProvenance{State: ProvenanceNone} + } + if col.RelPos == bc.projectTag && col.ColPos >= 0 && int(col.ColPos) < len(bc.projects) { + if provenance, recorded := bc.outputColumnProvenance[col.ColPos]; recorded { + return provenance + } + project := bc.projects[col.ColPos] + if project == nil { + return OutputColumnProvenance{State: ProvenanceNone} + } + if projectCol := project.GetCol(); projectCol != nil && + projectCol.RelPos == col.RelPos && projectCol.ColPos == col.ColPos { + return OutputColumnProvenance{State: ProvenanceNone} + } + return bc.outputColumnProvenanceForExpr(project) + } + if bc.groupTag > 0 && col.RelPos == bc.groupTag && col.ColPos >= 0 && int(col.ColPos) < len(bc.groups) { + groupExpr := bc.groups[col.ColPos] + if groupExpr == nil { + return OutputColumnProvenance{State: ProvenanceNone} + } + if groupCol := groupExpr.GetCol(); groupCol != nil && + groupCol.RelPos == col.RelPos && groupCol.ColPos == col.ColPos { + return OutputColumnProvenance{State: ProvenanceNone} + } + return bc.outputColumnProvenanceForExpr(groupExpr) + } + binding := bc.bindingByTag[col.RelPos] + if binding == nil || col.ColPos < 0 || int(col.ColPos) >= len(binding.outputColumnProvenance) { + return OutputColumnProvenance{State: ProvenanceNone} + } + return binding.outputColumnProvenance[col.ColPos] +} + +func (bc *BindContext) outputColumnProvenanceForProject(colPos int32) OutputColumnProvenance { + if provenance, recorded := bc.outputColumnProvenance[colPos]; recorded { + return provenance + } + if colPos < 0 || int(colPos) >= len(bc.projects) { + return OutputColumnProvenance{State: ProvenanceNone} + } + return bc.outputColumnProvenanceForExpr(bc.projects[colPos]) +} + +func (bc *BindContext) outputColumnProvenanceForBoundary() []OutputColumnProvenance { + provenance := make([]OutputColumnProvenance, min(len(bc.headings), len(bc.projects))) + for i := range provenance { + provenance[i] = bc.outputColumnProvenanceForProject(int32(i)) + } + return provenance +} + +func (bc *BindContext) clearOutputColumnProvenance() { + for i := 0; i < min(len(bc.headings), len(bc.projects)); i++ { + bc.outputColumnProvenance[int32(i)] = OutputColumnProvenance{State: ProvenanceNone} + } +} diff --git a/pkg/sql/plan/projection_binder.go b/pkg/sql/plan/projection_binder.go index 56e4ac7f0cb6a..53293f45cc540 100644 --- a/pkg/sql/plan/projection_binder.go +++ b/pkg/sql/plan/projection_binder.go @@ -98,6 +98,12 @@ func (b *ProjectionBinder) BindExpr(astExpr tree.Expr, depth int32, isRoot bool) target := b.numericTargetType b.numericTargetType = nil defer func() { b.numericTargetType = target }() + _, isBareColumn := unwrapParenExpr(astExpr).(*tree.UnresolvedName) + if isBareColumn && isEnumOrSetPlanType(target) { + previousTarget := b.mysqlSpecialTargetType + b.mysqlSpecialTargetType = target + defer func() { b.mysqlSpecialTargetType = previousTarget }() + } if subquery, ok := scalarSubqueryExpr(astExpr); ok && !subquery.Exists { previousSubqueryTarget := b.numericSubqueryTarget b.numericSubqueryTarget = target diff --git a/pkg/sql/plan/query_builder.go b/pkg/sql/plan/query_builder.go index 33dae27a9b349..bd8f658a3ff00 100644 --- a/pkg/sql/plan/query_builder.go +++ b/pkg/sql/plan/query_builder.go @@ -3346,11 +3346,12 @@ func (builder *QueryBuilder) buildUnionWithResultLen( } if len(selectStmts) == 1 { + var nodeID int32 switch sltStmt := selectStmts[0].(type) { case *tree.Select: if sltClause, ok := sltStmt.Select.(*tree.SelectClause); ok { sltClause.Distinct = true - return builder.bindSelect(sltStmt, ctx, isRoot) + nodeID, err = builder.bindSelect(sltStmt, ctx, isRoot) } else { // rewrite sltStmt to select distinct * from (sltStmt) a tmpSltStmt := &tree.Select{ @@ -3376,15 +3377,19 @@ func (builder *QueryBuilder) buildUnionWithResultLen( Limit: astLimit, OrderBy: astOrderBy, } - return builder.bindSelect(tmpSltStmt, ctx, isRoot) + nodeID, err = builder.bindSelect(tmpSltStmt, ctx, isRoot) } case *tree.SelectClause: if !sltStmt.Distinct { sltStmt.Distinct = true } - return builder.bindSelect(&tree.Select{Select: sltStmt, Limit: astLimit, OrderBy: astOrderBy}, ctx, isRoot) + nodeID, err = builder.bindSelect(&tree.Select{Select: sltStmt, Limit: astLimit, OrderBy: astOrderBy}, ctx, isRoot) } + if err == nil { + ctx.clearOutputColumnProvenance() + } + return nodeID, err } // build selects @@ -3624,6 +3629,7 @@ func (builder *QueryBuilder) buildUnionWithResultLen( }, }) } + ctx.clearOutputColumnProvenance() // A set-operation result keeps ENUM/SET definition-order provenance only // when every non-NULL branch is the same pure display contract. A literal // NULL is neutral because it cannot introduce a competing comparison @@ -3981,6 +3987,13 @@ func (builder *QueryBuilder) bindNoRecursiveCte( if err != nil { return } + if subCtx.restoreViewMySQLSpecialTypes { + nodeID, err = builder.appendMySQLSpecialTypeBoundary( + nodeID, subCtx, subCtx.outputColumnProvenanceForBoundary()) + if err != nil { + return + } + } if subCtx.hasSingleRow { ctx.hasSingleRow = true @@ -4319,6 +4332,7 @@ func (builder *QueryBuilder) bindRecursiveCte( //5. bind final statement ctx.sinkTag = initCtx.sinkTag //5.0 add initial statement as table binding into the ctx of main query + initCtx.clearOutputColumnProvenance() err = builder.addBinding(initLastNodeID, *cteRef.ast.Name, ctx) if err != nil { return @@ -4694,6 +4708,20 @@ func (builder *QueryBuilder) bindSelect(stmt *tree.Select, ctx *BindContext, isR if resultLen, notCacheable, err = builder.bindProjection(ctx, projectionBinder, selectList, notCacheable); err != nil { return } + var viewRawProjects []*plan.Expr + if ctx.restoreViewMySQLSpecialTypes && astOrderBy != nil && !ctx.isDistinct && + len(ctx.groups) == 0 && len(ctx.aggregates) == 0 { + for i := 0; i < resultLen; i++ { + rawExpr, ok := transparentOutputSourceExpr(ctx.projects[i]) + if !ok { + continue + } + if viewRawProjects == nil { + viewRawProjects = make([]*plan.Expr, resultLen) + } + viewRawProjects[i] = DeepCopyExpr(rawExpr) + } + } // bind TIME WINDOW var fillType plan.Node_FillType @@ -4719,6 +4747,20 @@ func (builder *QueryBuilder) bindSelect(stmt *tree.Select, ctx *BindContext, isR return } } + if len(boundOrderBys) > 0 && len(viewRawProjects) > 0 { + ctx.mysqlSpecialRawProjectPositions = make(map[int32]int32, len(viewRawProjects)) + for visiblePos, rawExpr := range viewRawProjects { + if rawExpr == nil { + continue + } + rawPos, appendErr := appendOrderByProjectExpr(ctx, rawExpr) + if appendErr != nil { + err = appendErr + return + } + ctx.mysqlSpecialRawProjectPositions[int32(visiblePos)] = rawPos + } + } // bind limit/offset clause var boundOffsetExpr *Expr @@ -4834,6 +4876,13 @@ func (builder *QueryBuilder) bindSelect(stmt *tree.Select, ctx *BindContext, isR } } } + if len(ctx.groups) > 0 || ctx.isDistinct { + for i := 0; i < resultLen; i++ { + if typ := mysqlSpecialTypeFromProvenance(ctx.outputColumnProvenanceForProject(int32(i))); typ != nil { + ctx.setMySQLSpecialCanonicalType(int32(i), typ) + } + } + } // append SORT node (include limit, offset) if len(boundOrderBys) > 0 { @@ -8811,6 +8860,7 @@ func (builder *QueryBuilder) bindView( return 0, nil } viewCtx := NewBindContext(builder, nil) + viewCtx.restoreViewMySQLSpecialTypes = true viewCtx.snapshot = snapshot viewCtx.lower = ctx.lower @@ -8891,6 +8941,11 @@ func (builder *QueryBuilder) bindView( if err != nil { return } + nodeID, err = builder.appendMySQLSpecialTypeBoundary( + nodeID, viewCtx, viewCtx.outputColumnProvenanceForBoundary()) + if err != nil { + return + } if len(viewStmt.ColNames) > 0 { if len(viewStmt.ColNames) != len(viewCtx.headings) { return 0, moerr.NewViewWrongList(builder.GetContext()) @@ -8904,6 +8959,114 @@ func (builder *QueryBuilder) bindView( return } +// appendMySQLSpecialTypeBoundary restores transparent ENUM/SET outputs only +// after a complete query boundary. Semantic operators inside the query must +// continue to consume the SQL-visible value because index-to-value conversion +// is not always injective. +func (builder *QueryBuilder) appendMySQLSpecialTypeBoundary( + nodeID int32, ctx *BindContext, provenance []OutputColumnProvenance, +) (int32, error) { + visibleProjects := min(len(ctx.headings), len(ctx.projects), len(ctx.results)) + rootNode := builder.qry.Nodes[nodeID] + if rootNode.NodeType == plan.Node_PROJECT && len(ctx.mysqlSpecialRawProjectPositions) > 0 { + allSpecialTypesRestored := true + for i := 0; i < visibleProjects; i++ { + if ctx.mysqlSpecialCanonicalTypeForProject(int32(i)) != nil { + allSpecialTypesRestored = false + continue + } + targetType := mysqlSpecialTypeFromProvenance(provenance[i]) + if targetType == nil { + continue + } + rawPos, ok := ctx.mysqlSpecialRawProjectPositions[int32(i)] + if !ok { + allSpecialTypesRestored = false + continue + } + rawExpr := GetColExpr(*targetType, ctx.projectTag, rawPos) + rootNode.ProjectList[i] = rawExpr + ctx.results[i] = rawExpr + } + if allSpecialTypesRestored { + return nodeID, nil + } + } + canExposeRaw := rootNode.NodeType == plan.Node_PROJECT && + len(rootNode.BindingTags) == 1 && rootNode.BindingTags[0] == ctx.projectTag && ctx.resultTag == 0 + if canExposeRaw { + rootVisibleProjects := min(len(ctx.headings), len(rootNode.ProjectList)) + allSpecialTypesRestored := true + for i := 0; i < rootVisibleProjects; i++ { + if ctx.mysqlSpecialCanonicalTypeForProject(int32(i)) != nil { + allSpecialTypesRestored = false + continue + } + sourceExpr, ok := transparentOutputSourceExpr(rootNode.ProjectList[i]) + if !ok { + if i < len(provenance) && mysqlSpecialTypeFromProvenance(provenance[i]) != nil { + allSpecialTypesRestored = false + } + continue + } + restored := DeepCopyExpr(sourceExpr) + rootNode.ProjectList[i] = restored + if i < len(ctx.projects) { + ctx.projects[i] = restored + } + if i < len(ctx.results) { + ctx.results[i] = restored + } + } + if allSpecialTypesRestored { + return nodeID, nil + } + } + + sourceTag := ctx.rootTag() + projects := make([]*plan.Expr, len(ctx.results)) + for i := range ctx.results { + projects[i] = GetColExpr(ctx.results[i].Typ, sourceTag, int32(i)) + } + needsBoundary := false + for i := 0; i < visibleProjects; i++ { + targetType := ctx.mysqlSpecialCanonicalTypeForProject(int32(i)) + if targetType == nil { + continue + } + needsBoundary = true + var castErr error + if isEnumPlanType(targetType) { + projects[i], castErr = funcCastForEnumType(builder.GetContext(), projects[i], *targetType) + } else { + projects[i], castErr = funcCastForSetType(builder.GetContext(), projects[i], *targetType) + } + if castErr != nil { + return 0, castErr + } + } + if needsBoundary { + for i := 0; i < visibleProjects && i < len(provenance); i++ { + ctx.outputColumnProvenance[int32(i)] = provenance[i] + } + ctx.projectTag = builder.genNewBindTag() + ctx.resultTag = 0 + ctx.projects = projects + ctx.results = projects + return builder.appendNode(&plan.Node{ + NodeType: plan.Node_PROJECT, + ProjectList: projects, + Children: []int32{nodeID}, + BindingTags: []int32{ctx.projectTag}, + }, ctx), nil + } + + // Catalog lineage alone cannot prove that a visible string can be reversed + // to its original ENUM ordinal or SET bitmap. Without a raw sidecar or an + // exact display wrapper, preserve the visible value unchanged. + return nodeID, nil +} + // ViewData persisted before parser SQL mode was recorded used MatrixOne's // legacy grammar where || meant concat. const legacyViewParserSQLMode = "PIPES_AS_CONCAT" @@ -9095,6 +9258,13 @@ func (builder *QueryBuilder) buildTable(stmt tree.TableExpr, ctx *BindContext, t if err != nil { return 0, err } + if subCtx.restoreViewMySQLSpecialTypes { + nodeID, err = builder.appendMySQLSpecialTypeBoundary( + nodeID, subCtx, subCtx.outputColumnProvenanceForBoundary()) + if err != nil { + return 0, err + } + } if subCtx.isCorrelated { if err = builder.normalizeTransparentCorrelatedDerivedTable(nodeID, ctx); err != nil { return 0, err @@ -9613,6 +9783,8 @@ func (builder *QueryBuilder) addBinding(nodeID int32, alias tree.AliasClause, ct var colIsHidden []bool var types []*plan.Type var mysqlSpecialOrderTypes []*plan.Type + var mysqlSpecialCanonicalTypes []*plan.Type + var outputColumnProvenance []OutputColumnProvenance var defaultVals []string var binding *Binding var bindingToReplace *Binding @@ -9691,12 +9863,30 @@ func (builder *QueryBuilder) addBinding(nodeID int32, alias tree.AliasClause, ct binding = NewBinding(tag, nodeID, node.TableDef.DbName, table, node.TableDef.TblId, cols, colIsHidden, types, util.TableIsClusterTable(node.TableDef.TableType), defaultVals) binding.originCols = originCols + binding.outputColumnProvenance = make([]OutputColumnProvenance, colLength) + for i, col := range node.TableDef.Cols { + binding.outputColumnProvenance[i] = OutputColumnProvenance{ + State: ProvenanceSingleSource, + CanInheritSourceDefault: true, + Source: &SourceColumn{ + RelPos: tag, + ColPos: int32(i), + TableID: node.TableDef.TblId, + Metadata: snapshotSourceColumnMetadata(col), + }, + } + } } else { // Subquery subCtx := builder.ctxByNode[nodeID] tag := subCtx.rootTag() headings := subCtx.headings projects := subCtx.projects + if subCtx.restoreViewMySQLSpecialTypes && + (len(subCtx.mysqlSpecialRawProjectPositions) > 0 || len(subCtx.mysqlSpecialCanonicalTypes) > 0) && + len(subCtx.results) >= len(headings) { + projects = subCtx.results + } if len(alias.Cols) > len(headings) { return moerr.NewSyntaxErrorf(builder.GetContext(), "table %q has %d columns available but %d columns specified", alias.Alias, len(headings), len(alias.Cols)) @@ -9736,6 +9926,17 @@ func (builder *QueryBuilder) addBinding(nodeID int32, alias tree.AliasClause, ct } mysqlSpecialOrderTypes[i] = orderType } + if canonicalType := subCtx.mysqlSpecialCanonicalTypeForProject(int32(i)); canonicalType != nil { + if mysqlSpecialCanonicalTypes == nil { + mysqlSpecialCanonicalTypes = make([]*plan.Type, colLength) + } + mysqlSpecialCanonicalTypes[i] = canonicalType + } + if outputColumnProvenance == nil { + outputColumnProvenance = make([]OutputColumnProvenance, colLength) + } + outputColumnProvenance[i] = subCtx.outputColumnProvenanceForProject(int32(i)) + outputColumnProvenance[i].CanInheritSourceDefault = false name := table + "." + cols[i] builder.nameByColRef[[2]int32{tag, int32(i)}] = name } @@ -9743,6 +9944,8 @@ func (builder *QueryBuilder) addBinding(nodeID int32, alias tree.AliasClause, ct binding = NewBinding(tag, nodeID, "", table, 0, cols, colIsHidden, types, false, defaultVals) binding.originCols = originCols binding.mysqlSpecialOrderTypes = mysqlSpecialOrderTypes + binding.mysqlSpecialCanonicalTypes = mysqlSpecialCanonicalTypes + binding.outputColumnProvenance = outputColumnProvenance } if bindingToReplace != nil { diff --git a/pkg/sql/plan/types.go b/pkg/sql/plan/types.go index 4c9ab7d605819..59a71afaf40a8 100644 --- a/pkg/sql/plan/types.go +++ b/pkg/sql/plan/types.go @@ -413,6 +413,11 @@ type aliasItem struct { type BindContext struct { binder Binder + // outputColumnProvenance records planner-local lineage overrides by output + // position. An explicit None prevents later transparent-boundary code from + // rediscovering a source after a semantic boundary has cleared it. + outputColumnProvenance map[int32]OutputColumnProvenance + // mysqlSpecialOrderTypes records the storage type behind a visible ENUM/SET // display value. It is planner-local semantic provenance: only a pure // display projection (or a pure column passthrough of one) may populate it. @@ -421,6 +426,18 @@ type BindContext struct { // The generated plan consumes the provenance by materializing an ordinary // numeric sort expression, so this metadata never crosses the plan wire. mysqlSpecialOrderTypes map[int32]*plan.Type + // mysqlSpecialCanonicalTypes records outputs whose SQL-visible value has + // already passed through GROUP BY or DISTINCT and must be canonically + // re-encoded when a persisted View exposes an ENUM/SET catalog type. + mysqlSpecialCanonicalTypes map[int32]*plan.Type + // restoreViewMySQLSpecialTypes is inherited only while rebinding a persisted + // View. It lets transparent derived/CTE query boundaries expose their raw + // ENUM/SET values without changing ordinary query-boundary behavior. + restoreViewMySQLSpecialTypes bool + // mysqlSpecialRawProjectPositions maps a visible output position to a hidden + // raw ENUM/SET sidecar in the query block's PROJECT. It is populated only + // for row-preserving View ORDER BY boundaries. + mysqlSpecialRawProjectPositions map[int32]int32 //cteByName saves all cte definitions in the current stmt cteByName map[string]*CTERef @@ -570,6 +587,7 @@ type baseBinder struct { numericParamType *Type numericSubqueryTarget *Type numericFunctionTarget bool + mysqlSpecialTargetType *Type allowCanonicalNameConstValueCast bool bindRawMySQLSpecialType bool } @@ -693,6 +711,12 @@ type Binding struct { // the string column is a pure display of the recorded ENUM/SET storage // type, and may therefore use definition-order semantics when ordered. mysqlSpecialOrderTypes []*plan.Type + // mysqlSpecialCanonicalTypes is aligned with cols and propagates the + // post-semantic canonical-value contract through transparent bindings. + mysqlSpecialCanonicalTypes []*plan.Type + // outputColumnProvenance is aligned with cols and carries planner-local, + // single-source output lineage. It is never serialized into the plan. + outputColumnProvenance []OutputColumnProvenance refCnts []uint // lower case colIdByName map[string]int32 diff --git a/pkg/tests/issues/issue_26226_test.go b/pkg/tests/issues/issue_26226_test.go new file mode 100644 index 0000000000000..15c8a2bd8d443 --- /dev/null +++ b/pkg/tests/issues/issue_26226_test.go @@ -0,0 +1,216 @@ +// Copyright 2026 Matrix Origin +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package issues + +import ( + "context" + "database/sql" + "fmt" + "strings" + "testing" + "time" + + _ "github.com/go-sql-driver/mysql" + "github.com/matrixorigin/matrixone/pkg/embed" + "github.com/stretchr/testify/require" +) + +func TestIssue26226ViewDistinctUsesVisibleSetValue(t *testing.T) { + embed.RunBaseClusterTests(t, func(c embed.Cluster) { + cn, err := c.GetCNService(0) + require.NoError(t, err) + port := cn.GetServiceConfig().CN.Frontend.Port + dbConn, err := sql.Open("mysql", fmt.Sprintf("dump:111@tcp(127.0.0.1:%d)/", port)) + require.NoError(t, err) + defer dbConn.Close() + ctx, cancel := context.WithTimeout(context.Background(), 3*time.Minute) + defer cancel() + execSQLRequire(t, ctx, dbConn, "set role moadmin") + + const db = "issue_26226" + for _, stmt := range []string{ + "drop database if exists " + db, + "create database " + db, + "create table " + db + ".t (id int primary key, flags set('', 'a'))", + "insert into " + db + ".t values (1, ''), (2, 1)", + "create view " + db + ".v_raw as select flags from " + db + ".t where id = 2", + "create view " + db + ".v as select distinct flags from " + db + ".t", + "create view " + db + ".v_order as select id, flags from " + db + ".t order by flags", + "create view " + db + ".v_order_derived as select id, flags from (select id, flags from " + db + ".t) d order by flags", + "create view " + db + ".v_order_cte as with d as (select id, flags from " + db + ".t) select id, flags from d order by flags", + "create view " + db + ".v_group as select flags from " + db + ".t group by flags", + "create view " + db + ".v_derived as select flags from (select id, flags from " + db + ".t) d where id = 2", + "create view " + db + ".v_cte as with d as (select id, flags from " + db + ".t) select flags from d where id = 2", + "create view " + db + ".v_union as select flags from " + db + ".t union all select flags from " + db + ".t", + "create view " + db + ".v_union_distinct as select flags from " + db + ".t union select flags from " + db + ".t", + "create view " + db + ".v_recursive as with recursive d(flags) as (select flags from " + db + ".t union all select flags from d where false) select flags from d", + "create table " + db + ".copied as select flags from " + db + ".v_raw", + "create table " + db + ".copied_derived as select flags from " + db + ".v_derived", + "create table " + db + ".copied_cte as select flags from " + db + ".v_cte", + "create table " + db + ".copied_union as select flags from " + db + ".v_union", + "create table " + db + ".copied_union_distinct as select flags from " + db + ".v_union_distinct", + "create table " + db + ".copied_recursive as select flags from " + db + ".v_recursive", + "create table " + db + ".inserted (flags set('', 'a'))", + "insert into " + db + ".inserted select flags from " + db + ".v_raw", + "create table " + db + ".expr_src (flags set('a', 'b'))", + "insert into " + db + ".expr_src values ('a')", + "create table " + db + ".expr_dst (flags set('a', 'b'))", + "insert into " + db + ".expr_dst select concat(flags, ',b') from " + db + ".expr_src", + "create table " + db + ".semantic_t (priority enum('low','medium','high'), flags set('', 'a', 'b'))", + "insert into " + db + ".semantic_t values ('low', ''), ('medium', 1), ('high', 'a')", + "create view " + db + ".v_semantic_group as select priority, flags from " + db + ".semantic_t group by priority, flags", + "create view " + db + ".v_semantic_distinct as select distinct priority, flags from " + db + ".semantic_t", + "create view " + db + ".v_semantic_derived as select priority, flags from (select distinct priority, flags from " + db + ".semantic_t) d", + "create table " + db + ".defaults_src (e enum('low','medium','high') default 'medium', s set('', 'a', 'b') default 'a', n int default 7)", + "create table " + db + ".defaults_direct as select e, s, n from " + db + ".defaults_src", + "create table " + db + ".defaults_derived as select e, s, n from (select e, s, n from " + db + ".defaults_src) d", + "create table " + db + ".defaults_cte as with d as (select e, s, n from " + db + ".defaults_src) select e, s, n from d", + "create table " + db + ".defaults_union as select e, s, n from " + db + ".defaults_src union all select e, s, n from " + db + ".defaults_src", + } { + execSQLRequire(t, ctx, dbConn, stmt) + } + defer func() { + cleanupCtx, cleanupCancel := context.WithTimeout(context.Background(), 30*time.Second) + defer cleanupCancel() + execSQLMaybe(t, cleanupCtx, dbConn, "drop database if exists "+db) + }() + + var baseCount, viewCount int + require.NoError(t, dbConn.QueryRowContext(ctx, + "select count(*) from (select distinct flags from "+db+".t) d").Scan(&baseCount)) + require.NoError(t, dbConn.QueryRowContext(ctx, + "select count(*) from "+db+".v").Scan(&viewCount)) + require.Equal(t, 1, baseCount) + require.Equal(t, baseCount, viewCount) + for _, test := range []struct { + tableName string + dataType string + }{ + {tableName: "v_raw", dataType: "set"}, + {tableName: "v", dataType: "set"}, + {tableName: "v_order", dataType: "set"}, + {tableName: "v_order_derived", dataType: "set"}, + {tableName: "v_order_cte", dataType: "set"}, + {tableName: "v_group", dataType: "set"}, + {tableName: "v_derived", dataType: "set"}, + {tableName: "v_cte", dataType: "set"}, + {tableName: "v_union", dataType: "varchar"}, + {tableName: "v_union_distinct", dataType: "varchar"}, + {tableName: "v_recursive", dataType: "varchar"}, + {tableName: "copied_derived", dataType: "set"}, + {tableName: "copied_cte", dataType: "set"}, + {tableName: "copied_union", dataType: "varchar"}, + {tableName: "copied_union_distinct", dataType: "varchar"}, + {tableName: "copied_recursive", dataType: "varchar"}, + } { + var dataType string + require.NoError(t, dbConn.QueryRowContext(ctx, + "select data_type from information_schema.columns "+ + "where table_schema = ? and table_name = ? and column_name = 'flags'", + db, test.tableName).Scan(&dataType)) + require.Equal(t, test.dataType, strings.ToLower(dataType), test.tableName) + } + + for _, query := range []string{ + "select cast(flags as unsigned) from " + db + ".t where id = 2", + "select cast(flags as unsigned) from " + db + ".v_raw", + "select cast(flags as unsigned) from " + db + ".v_order where id = 2", + "select cast(flags as unsigned) from " + db + ".v_order_derived where id = 2", + "select cast(flags as unsigned) from " + db + ".v_order_cte where id = 2", + "select cast(flags as unsigned) from " + db + ".v_derived", + "select cast(flags as unsigned) from " + db + ".v_cte", + "select cast(flags as unsigned) from " + db + ".copied", + "select cast(flags as unsigned) from " + db + ".copied_derived", + "select cast(flags as unsigned) from " + db + ".copied_cte", + "select cast(flags as unsigned) from " + db + ".inserted", + } { + var bitmap uint64 + require.NoError(t, dbConn.QueryRowContext(ctx, query).Scan(&bitmap)) + require.Equal(t, uint64(1), bitmap, query) + } + for _, query := range []string{ + "select concat(flags, 'x') from " + db + ".v_derived", + "select concat(flags, 'x') from " + db + ".v_cte", + } { + var visibleValue string + require.NoError(t, dbConn.QueryRowContext(ctx, query).Scan(&visibleValue)) + require.Equal(t, "x", visibleValue, query) + } + + var nestedBitmap uint64 + require.NoError(t, dbConn.QueryRowContext(ctx, + "select cast(flags as unsigned) from "+db+".expr_dst").Scan(&nestedBitmap)) + require.Equal(t, uint64(3), nestedBitmap) + + for _, view := range []string{ + "v_semantic_group", + "v_semantic_distinct", + "v_semantic_derived", + } { + rows, err := dbConn.QueryContext(ctx, + "select priority, cast(flags as unsigned) from "+db+"."+view+" order by priority") + require.NoError(t, err, view) + defer rows.Close() + var actual [][2]any + for rows.Next() { + var priority string + var flags uint64 + require.NoError(t, rows.Scan(&priority, &flags)) + actual = append(actual, [2]any{priority, flags}) + } + require.NoError(t, rows.Err()) + require.Equal(t, [][2]any{{"low", uint64(0)}, {"medium", uint64(0)}, {"high", uint64(2)}}, actual, view) + } + + for _, test := range []struct { + tableName string + wantDefaults []string + }{ + {tableName: "defaults_direct", wantDefaults: []string{"'medium'", "'a'", "7"}}, + {tableName: "defaults_derived", wantDefaults: []string{"", "", ""}}, + {tableName: "defaults_cte", wantDefaults: []string{"", "", ""}}, + {tableName: "defaults_union", wantDefaults: []string{"", "", ""}}, + } { + rows, err := dbConn.QueryContext(ctx, + "select column_default from information_schema.columns "+ + "where table_schema = ? and table_name = ? and column_name in ('e', 's', 'n') "+ + "order by ordinal_position", + db, test.tableName) + require.NoError(t, err, test.tableName) + defer rows.Close() + var actualDefaults []string + for rows.Next() { + var defaultValue sql.NullString + require.NoError(t, rows.Scan(&defaultValue)) + actualDefaults = append(actualDefaults, defaultValue.String) + } + require.NoError(t, rows.Err()) + require.Equal(t, test.wantDefaults, actualDefaults, test.tableName) + + var name, createSQL string + require.NoError(t, dbConn.QueryRowContext(ctx, + "show create table "+db+"."+test.tableName).Scan(&name, &createSQL)) + if test.tableName == "defaults_direct" { + require.Contains(t, createSQL, "DEFAULT 'medium'") + require.Contains(t, createSQL, "DEFAULT 'a'") + require.Contains(t, createSQL, "DEFAULT 7") + } else { + require.NotContains(t, createSQL, "DEFAULT 'medium'") + require.NotContains(t, createSQL, "DEFAULT 'a'") + require.NotContains(t, createSQL, "DEFAULT 7") + } + } + }) +}