fix(oracle): preserve quoted view names
This commit is contained in:
parent
ea26a64b25
commit
789f0118f5
|
|
@ -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) {
|
||||
|
|
|
|||
|
|
@ -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() {
|
||||
|
|
|
|||
Loading…
Reference in New Issue