fix(oracle): serialize xmltype query columns

This commit is contained in:
t8y2 2026-07-01 01:27:22 +08:00
parent 1a8b52410f
commit f5d8ae5894
2 changed files with 710 additions and 0 deletions

View File

@ -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, `"`, `""`) + `"`
}

View File

@ -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 {