fix(oracle): preserve quoted view names

This commit is contained in:
t8y2 2026-07-27 22:57:40 +08:00
parent ea26a64b25
commit 789f0118f5
No known key found for this signature in database
2 changed files with 222 additions and 17 deletions

View File

@ -2112,16 +2112,7 @@ func (s *server) getObjectSource(schema, name, objectType string) (map[string]an
}
upperType := strings.ToUpper(objectType)
if upperType == "VIEW" {
// ALL_VIEWS.TEXT for views — ALL_SOURCE doesn't contain views, and
// DBMS_METADATA.GET_DDL fails on XE editions.
var source string
err = s.db.QueryRow(
"SELECT TEXT FROM ALL_VIEWS WHERE OWNER = :1 AND VIEW_NAME = :2",
schema, strings.ToUpper(name),
).Scan(&source)
if errors.Is(err, sql.ErrNoRows) {
return map[string]any{"name": name, "object_type": objectType, "schema": schema, "source": ""}, nil
}
source, err := s.getViewSource(schema, name)
if err != nil {
return nil, err
}
@ -2219,22 +2210,51 @@ func normalizeDDLObjectType(value string) string {
}
func (s *server) buildViewDDL(schema, name string) (string, error) {
source, err := s.getViewSource(schema, name)
if err != nil {
return "", err
}
trimmed := strings.TrimSpace(source)
upperSource := strings.ToUpper(trimmed)
if strings.HasPrefix(upperSource, "CREATE ") || strings.HasPrefix(upperSource, "ALTER ") {
return trimmed, nil
}
return fmt.Sprintf("CREATE OR REPLACE VIEW %s.%s AS\n%s", quoteIdentifier(schema), quoteIdentifier(name), trimmed), nil
}
func (s *server) getViewSource(schema, name string) (string, error) {
db, err := s.requireDB()
if err != nil {
return "", err
}
viewName := strings.TrimSpace(name)
var ddl string
metadataErr := db.QueryRow(
"SELECT DBMS_METADATA.GET_DDL('VIEW', :1, :2) FROM DUAL",
viewName, schema,
).Scan(&ddl)
if metadataErr == nil && strings.TrimSpace(ddl) != "" {
return strings.TrimSpace(ddl), nil
}
var source string
err = db.QueryRow(
fallbackErr := db.QueryRow(
"SELECT TEXT FROM ALL_VIEWS WHERE OWNER = :1 AND VIEW_NAME = :2",
schema, strings.ToUpper(name),
schema, viewName,
).Scan(&source)
if errors.Is(err, sql.ErrNoRows) {
return "", fmt.Errorf("view not found: %s.%s", schema, name)
if fallbackErr == nil && strings.TrimSpace(source) != "" {
return strings.TrimSpace(source), nil
}
if err != nil {
return "", err
if fallbackErr != nil && !errors.Is(fallbackErr, sql.ErrNoRows) {
if metadataErr != nil {
return "", fmt.Errorf(
"failed to load view source for %s.%s: DBMS_METADATA: %v; ALL_VIEWS: %w",
schema, viewName, metadataErr, fallbackErr,
)
}
return "", fmt.Errorf("failed to load view source for %s.%s from ALL_VIEWS: %w", schema, viewName, fallbackErr)
}
return fmt.Sprintf("CREATE OR REPLACE VIEW %s.%s AS\n%s", quoteIdentifier(schema), quoteIdentifier(name), strings.TrimSpace(source)), nil
return "", fmt.Errorf("view source not found: %s.%s", schema, viewName)
}
func (s *server) buildTableDDL(schema, table string) (string, error) {

View File

@ -1210,6 +1210,92 @@ func fakeOracleColumnLoader(columns []oracleColumnMeta) oracleColumnMetaLoader {
}
}
func TestGetObjectSourceUsesOriginalViewNameWithDBMSMetadata(t *testing.T) {
for _, viewName := range []string{"vEnginWJZ", "V_ENGINE_WJZ"} {
t.Run(viewName, func(t *testing.T) {
ddl := `CREATE OR REPLACE FORCE VIEW "ZTZS_ERP2"."` + viewName + `" AS SELECT source_id FROM "ZTZS_ERP2"."SOURCE_TABLE"`
db, scripted := openOracleViewSourceTestDB(t, []oracleViewSourceQueryStep{
{
queryContains: "DBMS_METADATA.GET_DDL('VIEW'",
args: []driver.Value{viewName, "ZTZS_ERP2"},
rows: [][]driver.Value{{ddl}},
},
})
s := newServer()
s.db = db
result, err := s.getObjectSource("ZTZS_ERP2", viewName, "VIEW")
if err != nil {
t.Fatal(err)
}
if result["source"] != ddl {
t.Fatalf("unexpected view source: %#v", result["source"])
}
if scripted.next != len(scripted.steps) {
t.Fatalf("expected %d queries, got %d", len(scripted.steps), scripted.next)
}
})
}
}
func TestGetObjectSourceFallsBackToAllViewsWithOriginalName(t *testing.T) {
const viewName = "vEnginWJZ"
const source = `SELECT source_id FROM "ZTZS_ERP2"."SOURCE_TABLE"`
db, scripted := openOracleViewSourceTestDB(t, []oracleViewSourceQueryStep{
{
queryContains: "DBMS_METADATA.GET_DDL('VIEW'",
args: []driver.Value{viewName, "ZTZS_ERP2"},
err: errors.New("ORA-31603: object not found"),
},
{
queryContains: "FROM ALL_VIEWS",
args: []driver.Value{"ZTZS_ERP2", viewName},
rows: [][]driver.Value{{source}},
},
})
s := newServer()
s.db = db
result, err := s.getObjectSource("ZTZS_ERP2", viewName, "VIEW")
if err != nil {
t.Fatal(err)
}
if result["source"] != source {
t.Fatalf("unexpected fallback source: %#v", result["source"])
}
if scripted.next != len(scripted.steps) {
t.Fatalf("expected %d queries, got %d", len(scripted.steps), scripted.next)
}
}
func TestGetObjectSourceRejectsMissingViewSource(t *testing.T) {
db, scripted := openOracleViewSourceTestDB(t, []oracleViewSourceQueryStep{
{
queryContains: "DBMS_METADATA.GET_DDL('VIEW'",
args: []driver.Value{"vEnginWJZ", "ZTZS_ERP2"},
err: errors.New("ORA-31603: object not found"),
},
{
queryContains: "FROM ALL_VIEWS",
args: []driver.Value{"ZTZS_ERP2", "vEnginWJZ"},
rows: nil,
},
})
s := newServer()
s.db = db
result, err := s.getObjectSource("ZTZS_ERP2", "vEnginWJZ", "VIEW")
if err == nil || !strings.Contains(err.Error(), "view source not found") {
t.Fatalf("expected missing view source error, got result=%#v error=%v", result, err)
}
if result != nil {
t.Fatalf("missing view source must not return a successful empty result: %#v", result)
}
if scripted.next != len(scripted.steps) {
t.Fatalf("expected %d queries, got %d", len(scripted.steps), scripted.next)
}
}
func contains(values []string, target string) bool {
for _, value := range values {
if value == target {
@ -1219,6 +1305,105 @@ func contains(values []string, target string) bool {
return false
}
type oracleViewSourceQueryStep struct {
queryContains string
args []driver.Value
rows [][]driver.Value
err error
}
type oracleViewSourceDriver struct {
steps []oracleViewSourceQueryStep
next int
}
func (d *oracleViewSourceDriver) Open(string) (driver.Conn, error) {
return &oracleViewSourceConn{driver: d}, nil
}
type oracleViewSourceConn struct {
driver *oracleViewSourceDriver
}
func (c *oracleViewSourceConn) Prepare(string) (driver.Stmt, error) {
return nil, errors.New("use QueryContext directly")
}
func (c *oracleViewSourceConn) Close() error {
return nil
}
func (c *oracleViewSourceConn) Begin() (driver.Tx, error) {
return nil, errors.New("not supported")
}
func (c *oracleViewSourceConn) QueryContext(
_ context.Context,
query string,
args []driver.NamedValue,
) (driver.Rows, error) {
if c.driver.next >= len(c.driver.steps) {
return nil, errors.New("unexpected extra query: " + query)
}
step := c.driver.steps[c.driver.next]
c.driver.next++
if !strings.Contains(query, step.queryContains) {
return nil, errors.New("unexpected query: " + query)
}
values := make([]driver.Value, len(args))
for index, arg := range args {
values[index] = arg.Value
}
if !reflect.DeepEqual(values, step.args) {
return nil, errors.New("unexpected query arguments")
}
if step.err != nil {
return nil, step.err
}
return &oracleViewSourceRows{columns: []string{"SOURCE"}, values: step.rows}, nil
}
type oracleViewSourceRows struct {
columns []string
values [][]driver.Value
next int
}
func (r *oracleViewSourceRows) Columns() []string {
return r.columns
}
func (r *oracleViewSourceRows) Close() error {
return nil
}
func (r *oracleViewSourceRows) Next(dest []driver.Value) error {
if r.next >= len(r.values) {
return io.EOF
}
copy(dest, r.values[r.next])
r.next++
return nil
}
func openOracleViewSourceTestDB(
t *testing.T,
steps []oracleViewSourceQueryStep,
) (*sql.DB, *oracleViewSourceDriver) {
t.Helper()
driverName := "oracle-test-view-source-" + strings.ReplaceAll(t.Name(), "/", "-") + "-" + time.Now().Format("150405.000000000")
scripted := &oracleViewSourceDriver{steps: steps}
sql.Register(driverName, scripted)
db, err := sql.Open(driverName, "")
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() {
_ = db.Close()
})
return db, scripted
}
// -- fake drivers for timeout tests --
func init() {