diff --git a/agents/drivers/oracle-go/main.go b/agents/drivers/oracle-go/main.go index 3668aa931..6d973fab9 100644 --- a/agents/drivers/oracle-go/main.go +++ b/agents/drivers/oracle-go/main.go @@ -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) { diff --git a/agents/drivers/oracle-go/main_test.go b/agents/drivers/oracle-go/main_test.go index b749ac402..b42770df8 100644 --- a/agents/drivers/oracle-go/main_test.go +++ b/agents/drivers/oracle-go/main_test.go @@ -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() {