diff --git a/agents/drivers/oracle-go/main.go b/agents/drivers/oracle-go/main.go index 1f5b2020b..68e21a03a 100644 --- a/agents/drivers/oracle-go/main.go +++ b/agents/drivers/oracle-go/main.go @@ -1379,11 +1379,7 @@ func (s *server) executeQueryPage(opts queryOptions, pageSize int) (queryPageRes HasMore: false, }, err } - sqlText, err := s.rewriteXMLTypeSelectSQL(sqlText) - if err != nil { - return queryPageResult{}, err - } - rows, err := s.queryRows(sqlText, nil) + rows, err := s.queryRowsWithXMLTypeRewriteIfNeeded(sqlText) if err != nil { return queryPageResult{}, err } @@ -1448,11 +1444,7 @@ func (s *server) startTableRead(opts queryOptions, pageSize int) (queryPageResul if !isQuerySQL(sqlText) { return queryPageResult{}, errors.New("table read requires a SELECT query") } - sqlText, err := s.rewriteXMLTypeSelectSQL(sqlText) - if err != nil { - return queryPageResult{}, err - } - rows, err := s.queryRows(sqlText, nil) + rows, err := s.queryRowsWithXMLTypeRewriteIfNeeded(sqlText) if err != nil { return queryPageResult{}, err } @@ -1603,12 +1595,7 @@ func (s *server) executeQuery(opts queryOptions) (queryResult, error) { } func (s *server) executeSelect(sqlText string, maxRows int) (queryResult, error) { - var err error - sqlText, err = s.rewriteXMLTypeSelectSQL(sqlText) - if err != nil { - return queryResult{}, err - } - rows, err := s.queryRows(sqlText, nil) + rows, err := s.queryRowsWithXMLTypeRewriteIfNeeded(sqlText) if err != nil { return queryResult{}, err } @@ -1647,6 +1634,49 @@ func scanRow(rows *sql.Rows, columnCount int) ([]any, error) { return values, nil } +func (s *server) queryRowsWithXMLTypeRewriteIfNeeded(sqlText string) (*sql.Rows, error) { + rows, err := s.queryRows(sqlText, nil) + if err != nil { + return nil, err + } + if !rowsContainOracleXMLType(rows) { + return rows, nil + } + rewritten, err := s.rewriteXMLTypeSelectSQL(sqlText) + if err != nil { + rows.Close() + return nil, err + } + if rewritten == sqlText { + return rows, nil + } + // Only pay the ALL_TAB_COLUMNS rewrite cost when the result metadata shows + // XMLTYPE. Ordinary Oracle queries should not run dictionary probes first. + rows.Close() + return s.queryRows(rewritten, nil) +} + +func rowsContainOracleXMLType(rows *sql.Rows) bool { + types, err := rows.ColumnTypes() + if err != nil { + return false + } + typeNames := make([]string, 0, len(types)) + for _, columnType := range types { + typeNames = append(typeNames, columnType.DatabaseTypeName()) + } + return oracleColumnTypeNamesContainXMLType(typeNames) +} + +func oracleColumnTypeNamesContainXMLType(typeNames []string) bool { + for _, typeName := range typeNames { + if isOracleXMLType(typeName) { + return true + } + } + return false +} + func (s *server) rewriteXMLTypeSelectSQL(sqlText string) (string, error) { return rewriteOracleXMLTypeSelectSQL(sqlText, s.loadOracleColumnMeta) } diff --git a/agents/drivers/oracle-go/main_test.go b/agents/drivers/oracle-go/main_test.go index 018759345..4e689e3f8 100644 --- a/agents/drivers/oracle-go/main_test.go +++ b/agents/drivers/oracle-go/main_test.go @@ -514,6 +514,26 @@ func TestRewriteOracleXMLTypeSkipsJoins(t *testing.T) { } } +func TestOracleColumnTypeNamesContainXMLType(t *testing.T) { + tests := []struct { + name string + typeNames []string + want bool + }{ + {name: "plain xmltype", typeNames: []string{"NUMBER", "XMLTYPE"}, want: true}, + {name: "qualified xmltype", typeNames: []string{"SYS.XMLTYPE"}, want: true}, + {name: "case and spaces", typeNames: []string{" varchar2 ", "sys.xmltype"}, want: true}, + {name: "ordinary columns", typeNames: []string{"NUMBER", "VARCHAR2", "DATE"}, want: false}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := oracleColumnTypeNamesContainXMLType(tt.typeNames); got != tt.want { + t.Fatalf("oracleColumnTypeNamesContainXMLType(%v) = %v, want %v", tt.typeNames, got, tt.want) + } + }) + } +} + func fakeOracleColumnLoader(columns []oracleColumnMeta) oracleColumnMetaLoader { return func(schema, table string) ([]oracleColumnMeta, error) { if strings.ToUpper(table) != "TEST_LOBS" {