diff --git a/agents/drivers/oracle-go/main.go b/agents/drivers/oracle-go/main.go index 5453ce938..440ecac43 100644 --- a/agents/drivers/oracle-go/main.go +++ b/agents/drivers/oracle-go/main.go @@ -171,6 +171,13 @@ type querySession struct { remaining int } +type oracleColumnMeta struct { + Name string + DataType string +} + +type oracleColumnMetaLoader func(schema, table string) ([]oracleColumnMeta, error) + type databaseInfo struct { Name string `json:"name"` } @@ -773,6 +780,32 @@ ORDER BY c.COLUMN_ID`, []any{schema, table}) return emptyIfNil(result), rows.Err() } +func (s *server) loadOracleColumnMeta(schema, table string) ([]oracleColumnMeta, error) { + schema, err := s.normalizeSchema(schema) + if err != nil { + return nil, err + } + table = strings.ToUpper(strings.TrimSpace(table)) + rows, err := s.queryRows(` +SELECT COLUMN_NAME, DATA_TYPE +FROM ALL_TAB_COLUMNS +WHERE OWNER = :1 AND TABLE_NAME = :2 +ORDER BY COLUMN_ID`, []any{schema, table}) + if err != nil { + return nil, err + } + defer rows.Close() + var result []oracleColumnMeta + for rows.Next() { + var item oracleColumnMeta + if err := rows.Scan(&item.Name, &item.DataType); err != nil { + return nil, err + } + result = append(result, item) + } + return result, rows.Err() +} + func (s *server) listIndexes(schema, table string) ([]indexInfo, error) { schema, err := s.normalizeSchema(schema) if err != nil { @@ -1201,6 +1234,10 @@ 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) if err != nil { return queryPageResult{}, err @@ -1266,6 +1303,10 @@ 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) if err != nil { return queryPageResult{}, err @@ -1417,6 +1458,11 @@ 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) if err != nil { return queryResult{}, err @@ -1456,6 +1502,577 @@ func scanRow(rows *sql.Rows, columnCount int) ([]any, error) { return values, nil } +func (s *server) rewriteXMLTypeSelectSQL(sqlText string) (string, error) { + return rewriteOracleXMLTypeSelectSQL(sqlText, s.loadOracleColumnMeta) +} + +func rewriteOracleXMLTypeSelectSQL(sqlText string, loadColumns oracleColumnMetaLoader) (string, error) { + rewritten, _, err := rewriteOracleXMLTypeSelectSQLDepth(sqlText, loadColumns, 0) + return rewritten, err +} + +func rewriteOracleXMLTypeSelectSQLDepth(sqlText string, loadColumns oracleColumnMetaLoader, depth int) (string, bool, error) { + if depth > 8 { + return sqlText, false, nil + } + if rewritten, changed, handled, err := rewriteDirectOracleXMLTypeSelectSQL(sqlText, loadColumns); handled || err != nil { + return rewritten, changed, err + } + rewritten, changed, err := rewriteNestedOracleSelects(sqlText, loadColumns, depth) + return rewritten, changed, err +} + +func rewriteNestedOracleSelects(sqlText string, loadColumns oracleColumnMetaLoader, depth int) (string, bool, error) { + var builder strings.Builder + changed := false + last := 0 + for pos := 0; pos < len(sqlText); pos++ { + switch sqlText[pos] { + case '\'': + pos = skipSingleQuotedSQL(sqlText, pos) + case '"': + pos = skipDoubleQuotedSQL(sqlText, pos) + case '-': + if pos+1 < len(sqlText) && sqlText[pos+1] == '-' { + pos = skipLineCommentSQL(sqlText, pos) + } + case '/': + if pos+1 < len(sqlText) && sqlText[pos+1] == '*' { + pos = skipBlockCommentSQL(sqlText, pos) + } + case '(': + end := findMatchingSQLParen(sqlText, pos) + if end < 0 { + return sqlText, false, nil + } + inner := sqlText[pos+1 : end] + if startsWithSQLKeyword(trimLeadingSQLComments(inner), "select") { + rewrittenInner, innerChanged, err := rewriteOracleXMLTypeSelectSQLDepth(inner, loadColumns, depth+1) + if err != nil { + return "", false, err + } + if innerChanged { + builder.WriteString(sqlText[last : pos+1]) + builder.WriteString(rewrittenInner) + last = end + changed = true + } + } + pos = end + } + } + if !changed { + return sqlText, false, nil + } + builder.WriteString(sqlText[last:]) + return builder.String(), true, nil +} + +func rewriteDirectOracleXMLTypeSelectSQL(sqlText string, loadColumns oracleColumnMetaLoader) (string, bool, bool, error) { + selectStart := leadingSQLSelectListStart(sqlText) + if selectStart < 0 { + return sqlText, false, false, nil + } + fromIdx := findTopLevelSQLKeyword(sqlText, selectStart, "from") + if fromIdx < 0 { + return sqlText, false, false, nil + } + selectListPrefix, selectList := splitOracleSelectListModifier(sqlText[selectStart:fromIdx]) + tableRef, ok := parseSingleOracleTableRef(sqlText[fromIdx+len("from"):]) + if !ok { + return sqlText, false, false, nil + } + items := splitTopLevelSQLList(selectList) + if len(items) == 0 || !oracleSelectListMayReferenceXMLType(items) { + return sqlText, false, true, nil + } + columns, err := loadColumns(tableRef.Schema, tableRef.Table) + if err != nil { + return "", false, true, err + } + if !oracleColumnsHaveXMLType(columns) { + return sqlText, false, true, nil + } + rewrittenItems, changed := rewriteOracleSelectItemsForXMLType(items, columns, tableRef) + if !changed { + return sqlText, false, true, nil + } + var builder strings.Builder + builder.WriteString(sqlText[:selectStart]) + builder.WriteString(selectListPrefix) + builder.WriteString(strings.Join(rewrittenItems, ", ")) + builder.WriteByte(' ') + builder.WriteString(sqlText[fromIdx:]) + return builder.String(), true, true, nil +} + +type oracleTableRef struct { + Schema string + Table string + Alias string + AliasText string +} + +type oracleIdentifierToken struct { + Name string + Text string + Quoted bool +} + +func parseSingleOracleTableRef(fromSQL string) (oracleTableRef, bool) { + pos := skipSQLWhitespace(fromSQL, 0) + if pos >= len(fromSQL) || fromSQL[pos] == '(' { + return oracleTableRef{}, false + } + first, next, ok := readOracleIdentifierToken(fromSQL, pos) + if !ok { + return oracleTableRef{}, false + } + ref := oracleTableRef{Table: first.Name} + pos = skipSQLWhitespace(fromSQL, next) + if pos < len(fromSQL) && fromSQL[pos] == '.' { + second, afterSecond, ok := readOracleIdentifierToken(fromSQL, skipSQLWhitespace(fromSQL, pos+1)) + if !ok { + return oracleTableRef{}, false + } + ref.Schema = first.Name + ref.Table = second.Name + pos = skipSQLWhitespace(fromSQL, afterSecond) + } + if pos < len(fromSQL) { + if strings.HasPrefix(strings.TrimLeft(fromSQL[pos:], " \t\r\n"), ",") { + return oracleTableRef{}, false + } + if nextKeywordIsOracleClause(fromSQL[pos:]) { + return ref, true + } + if startsWithSQLKeyword(fromSQL[pos:], "join") || + startsWithSQLKeyword(fromSQL[pos:], "inner") || + startsWithSQLKeyword(fromSQL[pos:], "left") || + startsWithSQLKeyword(fromSQL[pos:], "right") || + startsWithSQLKeyword(fromSQL[pos:], "full") || + startsWithSQLKeyword(fromSQL[pos:], "cross") { + return oracleTableRef{}, false + } + alias, afterAlias, ok := readOracleIdentifierToken(fromSQL, pos) + if ok && !oracleIdentifierIsClause(alias.Name) { + ref.Alias = alias.Name + ref.AliasText = alias.Text + pos = skipSQLWhitespace(fromSQL, afterAlias) + } + if strings.HasPrefix(strings.TrimLeft(fromSQL[pos:], " \t\r\n"), ",") || + startsWithSQLKeyword(fromSQL[pos:], "join") || + startsWithSQLKeyword(fromSQL[pos:], "inner") || + startsWithSQLKeyword(fromSQL[pos:], "left") || + startsWithSQLKeyword(fromSQL[pos:], "right") || + startsWithSQLKeyword(fromSQL[pos:], "full") || + startsWithSQLKeyword(fromSQL[pos:], "cross") { + return oracleTableRef{}, false + } + } + return ref, true +} + +func splitOracleSelectListModifier(selectList string) (string, string) { + trimmedLeft := strings.TrimLeft(selectList, " \t\r\n") + prefixLen := len(selectList) - len(trimmedLeft) + for _, keyword := range []string{"distinct", "all"} { + if startsWithSQLKeyword(trimmedLeft, keyword) { + modifierEnd := prefixLen + len(keyword) + for modifierEnd < len(selectList) && isSQLWhitespace(selectList[modifierEnd]) { + modifierEnd++ + } + return selectList[:modifierEnd], selectList[modifierEnd:] + } + } + return selectList[:prefixLen], selectList[prefixLen:] +} + +func oracleSelectListMayReferenceXMLType(items []string) bool { + for _, item := range items { + if _, ok := parseOracleStarSelectItem(item); ok { + return true + } + if _, _, _, ok := parseOracleColumnSelectItem(item); ok { + return true + } + } + return false +} + +func rewriteOracleSelectItemsForXMLType(items []string, columns []oracleColumnMeta, tableRef oracleTableRef) ([]string, bool) { + xmlColumns := map[string]oracleColumnMeta{} + for _, column := range columns { + if isOracleXMLType(column.DataType) { + xmlColumns[oracleIdentifierKey(column.Name)] = column + } + } + rewritten := make([]string, 0, len(items)) + changed := false + for _, item := range items { + if qualifier, ok := parseOracleStarSelectItem(item); ok && oracleQualifierMatchesTable(qualifier, tableRef) { + for _, column := range columns { + rewritten = append(rewritten, oracleSelectExpressionForColumn(column, tableRef, xmlColumns)) + } + changed = true + continue + } + qualifier, column, alias, ok := parseOracleColumnSelectItem(item) + if ok && oracleQualifierMatchesTable(qualifier, tableRef) { + if meta, isXML := xmlColumns[oracleIdentifierKey(column.Name)]; isXML { + outputAlias := alias + if outputAlias == "" { + outputAlias = quoteIdentifier(meta.Name) + } + rewritten = append(rewritten, oracleXMLSerializeExpression(oracleColumnRef(qualifier, meta.Name), outputAlias)) + changed = true + continue + } + } + rewritten = append(rewritten, item) + } + return rewritten, changed +} + +func oracleSelectExpressionForColumn(column oracleColumnMeta, tableRef oracleTableRef, xmlColumns map[string]oracleColumnMeta) string { + qualifier := "" + if tableRef.AliasText != "" { + qualifier = tableRef.AliasText + } + if _, isXML := xmlColumns[oracleIdentifierKey(column.Name)]; isXML { + return oracleXMLSerializeExpression(oracleColumnRef(qualifier, column.Name), quoteIdentifier(column.Name)) + } + return oracleColumnRef(qualifier, column.Name) +} + +func oracleXMLSerializeExpression(columnRef, alias string) string { + // go-ora v2.9.0 does not fully decode Oracle XMLTYPE result payloads, + // especially when 11g switches larger values to locator-based transfer. + return fmt.Sprintf("XMLSERIALIZE(CONTENT %s AS CLOB) AS %s", columnRef, alias) +} + +func oracleColumnRef(qualifier, column string) string { + if strings.TrimSpace(qualifier) == "" { + return quoteIdentifier(column) + } + return qualifier + "." + quoteIdentifier(column) +} + +func parseOracleStarSelectItem(item string) (string, bool) { + trimmed := strings.TrimSpace(item) + if trimmed == "*" { + return "", true + } + qualifier, pos, ok := readOracleIdentifierToken(trimmed, 0) + if !ok { + return "", false + } + pos = skipSQLWhitespace(trimmed, pos) + if pos >= len(trimmed) || trimmed[pos] != '.' { + return "", false + } + pos = skipSQLWhitespace(trimmed, pos+1) + if pos < len(trimmed) && trimmed[pos] == '*' && strings.TrimSpace(trimmed[pos+1:]) == "" { + return qualifier.Text, true + } + return "", false +} + +func parseOracleColumnSelectItem(item string) (qualifier string, column oracleIdentifierToken, alias string, ok bool) { + trimmed := strings.TrimSpace(item) + first, pos, ok := readOracleIdentifierToken(trimmed, 0) + if !ok { + return "", oracleIdentifierToken{}, "", false + } + column = first + pos = skipSQLWhitespace(trimmed, pos) + if pos < len(trimmed) && trimmed[pos] == '.' { + second, afterSecond, ok := readOracleIdentifierToken(trimmed, skipSQLWhitespace(trimmed, pos+1)) + if !ok { + return "", oracleIdentifierToken{}, "", false + } + qualifier = first.Text + column = second + pos = skipSQLWhitespace(trimmed, afterSecond) + } + if pos >= len(trimmed) { + return qualifier, column, "", true + } + if startsWithSQLKeyword(trimmed[pos:], "as") { + aliasToken, afterAlias, ok := readOracleIdentifierToken(trimmed, skipSQLWhitespace(trimmed, pos+len("as"))) + if !ok || strings.TrimSpace(trimmed[afterAlias:]) != "" { + return "", oracleIdentifierToken{}, "", false + } + return qualifier, column, aliasToken.Text, true + } + aliasToken, afterAlias, ok := readOracleIdentifierToken(trimmed, pos) + if !ok || strings.TrimSpace(trimmed[afterAlias:]) != "" { + return "", oracleIdentifierToken{}, "", false + } + return qualifier, column, aliasToken.Text, true +} + +func oracleQualifierMatchesTable(qualifier string, tableRef oracleTableRef) bool { + if strings.TrimSpace(qualifier) == "" { + return true + } + key := oracleIdentifierKey(unquoteOracleIdentifierText(qualifier)) + if tableRef.Alias != "" && key == oracleIdentifierKey(tableRef.Alias) { + return true + } + return key == oracleIdentifierKey(tableRef.Table) +} + +func oracleColumnsHaveXMLType(columns []oracleColumnMeta) bool { + for _, column := range columns { + if isOracleXMLType(column.DataType) { + return true + } + } + return false +} + +func isOracleXMLType(dataType string) bool { + normalized := strings.ToUpper(strings.TrimSpace(dataType)) + return normalized == "XMLTYPE" || normalized == "SYS.XMLTYPE" +} + +func leadingSQLSelectListStart(sqlText string) int { + trimmed := trimLeadingSQLComments(sqlText) + prefixLen := len(sqlText) - len(trimmed) + if !startsWithSQLKeyword(trimmed, "select") { + return -1 + } + return prefixLen + len("select") +} + +func splitTopLevelSQLList(value string) []string { + var result []string + start := 0 + depth := 0 + for pos := 0; pos < len(value); pos++ { + switch value[pos] { + case '\'': + pos = skipSingleQuotedSQL(value, pos) + case '"': + pos = skipDoubleQuotedSQL(value, pos) + case '-': + if pos+1 < len(value) && value[pos+1] == '-' { + pos = skipLineCommentSQL(value, pos) + } + case '/': + if pos+1 < len(value) && value[pos+1] == '*' { + pos = skipBlockCommentSQL(value, pos) + } + case '(': + depth++ + case ')': + if depth > 0 { + depth-- + } + case ',': + if depth == 0 { + result = append(result, strings.TrimSpace(value[start:pos])) + start = pos + 1 + } + } + } + tail := strings.TrimSpace(value[start:]) + if tail != "" { + result = append(result, tail) + } + return result +} + +func findTopLevelSQLKeyword(sqlText string, start int, keyword string) int { + depth := 0 + for pos := start; pos < len(sqlText); pos++ { + switch sqlText[pos] { + case '\'': + pos = skipSingleQuotedSQL(sqlText, pos) + case '"': + pos = skipDoubleQuotedSQL(sqlText, pos) + case '-': + if pos+1 < len(sqlText) && sqlText[pos+1] == '-' { + pos = skipLineCommentSQL(sqlText, pos) + } + case '/': + if pos+1 < len(sqlText) && sqlText[pos+1] == '*' { + pos = skipBlockCommentSQL(sqlText, pos) + } + case '(': + depth++ + case ')': + if depth > 0 { + depth-- + } + default: + if depth == 0 && sqlKeywordAt(sqlText, pos, keyword) { + return pos + } + } + } + return -1 +} + +func findMatchingSQLParen(sqlText string, open int) int { + depth := 0 + for pos := open; pos < len(sqlText); pos++ { + switch sqlText[pos] { + case '\'': + pos = skipSingleQuotedSQL(sqlText, pos) + case '"': + pos = skipDoubleQuotedSQL(sqlText, pos) + case '-': + if pos+1 < len(sqlText) && sqlText[pos+1] == '-' { + pos = skipLineCommentSQL(sqlText, pos) + } + case '/': + if pos+1 < len(sqlText) && sqlText[pos+1] == '*' { + pos = skipBlockCommentSQL(sqlText, pos) + } + case '(': + depth++ + case ')': + depth-- + if depth == 0 { + return pos + } + } + } + return -1 +} + +func readOracleIdentifierToken(value string, pos int) (oracleIdentifierToken, int, bool) { + pos = skipSQLWhitespace(value, pos) + if pos >= len(value) { + return oracleIdentifierToken{}, pos, false + } + if value[pos] == '"' { + end := pos + 1 + var builder strings.Builder + for end < len(value) { + if value[end] == '"' { + if end+1 < len(value) && value[end+1] == '"' { + builder.WriteByte('"') + end += 2 + continue + } + return oracleIdentifierToken{Name: builder.String(), Text: value[pos : end+1], Quoted: true}, end + 1, true + } + builder.WriteByte(value[end]) + end++ + } + return oracleIdentifierToken{}, pos, false + } + if !isOracleIdentifierStart(value[pos]) { + return oracleIdentifierToken{}, pos, false + } + end := pos + 1 + for end < len(value) && isOracleIdentifierPart(value[end]) { + end++ + } + text := value[pos:end] + return oracleIdentifierToken{Name: strings.ToUpper(text), Text: text}, end, true +} + +func unquoteOracleIdentifierText(value string) string { + value = strings.TrimSpace(value) + if len(value) >= 2 && value[0] == '"' && value[len(value)-1] == '"' { + return strings.ReplaceAll(value[1:len(value)-1], `""`, `"`) + } + return strings.ToUpper(value) +} + +func oracleIdentifierKey(value string) string { + return strings.ToUpper(strings.TrimSpace(value)) +} + +func oracleIdentifierIsClause(value string) bool { + switch oracleIdentifierKey(value) { + case "WHERE", "GROUP", "ORDER", "HAVING", "CONNECT", "START", "MODEL", "FETCH", "OFFSET", "UNION", "MINUS", "INTERSECT": + return true + default: + return false + } +} + +func nextKeywordIsOracleClause(value string) bool { + trimmed := strings.TrimSpace(value) + if trimmed == "" { + return true + } + token, _, ok := readOracleIdentifierToken(trimmed, 0) + return ok && oracleIdentifierIsClause(token.Name) +} + +func isOracleIdentifierStart(ch byte) bool { + return (ch >= 'a' && ch <= 'z') || (ch >= 'A' && ch <= 'Z') || ch == '_' || ch == '$' || ch == '#' +} + +func isOracleIdentifierPart(ch byte) bool { + return isOracleIdentifierStart(ch) || (ch >= '0' && ch <= '9') +} + +func skipSQLWhitespace(value string, pos int) int { + for pos < len(value) && isSQLWhitespace(value[pos]) { + pos++ + } + return pos +} + +func isSQLWhitespace(ch byte) bool { + return ch == ' ' || ch == '\t' || ch == '\r' || ch == '\n' +} + +func skipSingleQuotedSQL(value string, pos int) int { + pos++ + for pos < len(value) { + if value[pos] == '\'' { + if pos+1 < len(value) && value[pos+1] == '\'' { + pos += 2 + continue + } + return pos + } + pos++ + } + return len(value) - 1 +} + +func skipDoubleQuotedSQL(value string, pos int) int { + pos++ + for pos < len(value) { + if value[pos] == '"' { + if pos+1 < len(value) && value[pos+1] == '"' { + pos += 2 + continue + } + return pos + } + pos++ + } + return len(value) - 1 +} + +func skipLineCommentSQL(value string, pos int) int { + for pos < len(value) { + if value[pos] == '\n' || value[pos] == '\r' { + return pos + } + pos++ + } + return len(value) - 1 +} + +func skipBlockCommentSQL(value string, pos int) int { + end := strings.Index(value[pos+2:], "*/") + if end < 0 { + return len(value) - 1 + } + return pos + end + 3 +} + func (s *server) setSchema(schema string) error { db, err := s.requireDB() if err != nil { @@ -1577,6 +2194,19 @@ func startsWithSQLKeyword(sqlText, keyword string) bool { return !((next >= 'a' && next <= 'z') || (next >= 'A' && next <= 'Z') || (next >= '0' && next <= '9') || next == '_' || next == '$') } +func sqlKeywordAt(sqlText string, pos int, keyword string) bool { + if pos < 0 || pos+len(keyword) > len(sqlText) || !strings.EqualFold(sqlText[pos:pos+len(keyword)], keyword) { + return false + } + if pos > 0 && isOracleIdentifierPart(sqlText[pos-1]) { + return false + } + if pos+len(keyword) >= len(sqlText) { + return true + } + return !isOracleIdentifierPart(sqlText[pos+len(keyword)]) +} + func quoteIdentifier(value string) string { return `"` + strings.ReplaceAll(value, `"`, `""`) + `"` } diff --git a/agents/drivers/oracle-go/main_test.go b/agents/drivers/oracle-go/main_test.go index d7b552651..7c5d9e876 100644 --- a/agents/drivers/oracle-go/main_test.go +++ b/agents/drivers/oracle-go/main_test.go @@ -381,6 +381,86 @@ func TestIsOraclePGALimitError(t *testing.T) { } } +func TestRewriteOracleXMLTypeSelectStar(t *testing.T) { + sqlText, err := rewriteOracleXMLTypeSelectSQL( + `SELECT * FROM TEST_LOBS`, + fakeOracleColumnLoader([]oracleColumnMeta{ + {Name: "ID", DataType: "NUMBER"}, + {Name: "XML_CONTENT", DataType: "XMLTYPE"}, + {Name: "TEST_NAME", DataType: "VARCHAR2"}, + }), + ) + if err != nil { + t.Fatal(err) + } + want := `SELECT "ID", XMLSERIALIZE(CONTENT "XML_CONTENT" AS CLOB) AS "XML_CONTENT", "TEST_NAME" FROM TEST_LOBS` + if sqlText != want { + t.Fatalf("rewriteOracleXMLTypeSelectSQL() = %s, want %s", sqlText, want) + } +} + +func TestRewriteOracleXMLTypeExplicitColumn(t *testing.T) { + sqlText, err := rewriteOracleXMLTypeSelectSQL( + `SELECT t.ID, t.XML_CONTENT AS xml_doc FROM TEST_LOBS t WHERE t.ID = 1`, + fakeOracleColumnLoader([]oracleColumnMeta{ + {Name: "ID", DataType: "NUMBER"}, + {Name: "XML_CONTENT", DataType: "SYS.XMLTYPE"}, + }), + ) + if err != nil { + t.Fatal(err) + } + want := `SELECT t.ID, XMLSERIALIZE(CONTENT t."XML_CONTENT" AS CLOB) AS xml_doc FROM TEST_LOBS t WHERE t.ID = 1` + if sqlText != want { + t.Fatalf("rewriteOracleXMLTypeSelectSQL() = %s, want %s", sqlText, want) + } +} + +func TestRewriteOracleXMLTypeNestedRownumQuery(t *testing.T) { + sqlText, err := rewriteOracleXMLTypeSelectSQL( + `SELECT * FROM (SELECT "ID", "XML_CONTENT" FROM "DBX"."TEST_LOBS") WHERE ROWNUM <= 100`, + fakeOracleColumnLoader([]oracleColumnMeta{ + {Name: "ID", DataType: "NUMBER"}, + {Name: "XML_CONTENT", DataType: "XMLTYPE"}, + }), + ) + if err != nil { + t.Fatal(err) + } + if !strings.Contains(sqlText, `XMLSERIALIZE(CONTENT "XML_CONTENT" AS CLOB) AS "XML_CONTENT"`) { + t.Fatalf("expected nested XMLTYPE column to be serialized, got: %s", sqlText) + } +} + +func TestRewriteOracleXMLTypeSkipsJoins(t *testing.T) { + called := false + sqlText, err := rewriteOracleXMLTypeSelectSQL( + `SELECT * FROM TEST_LOBS l JOIN OTHER_TABLE o ON o.ID = l.ID`, + func(schema, table string) ([]oracleColumnMeta, error) { + called = true + return nil, nil + }, + ) + if err != nil { + t.Fatal(err) + } + if called { + t.Fatal("join query should not load table metadata") + } + if sqlText != `SELECT * FROM TEST_LOBS l JOIN OTHER_TABLE o ON o.ID = l.ID` { + t.Fatalf("join query should not be rewritten, got: %s", sqlText) + } +} + +func fakeOracleColumnLoader(columns []oracleColumnMeta) oracleColumnMetaLoader { + return func(schema, table string) ([]oracleColumnMeta, error) { + if strings.ToUpper(table) != "TEST_LOBS" { + return nil, nil + } + return columns, nil + } +} + func contains(values []string, target string) bool { for _, value := range values { if value == target {