fix(oracle): preserve plsql block terminators
This commit is contained in:
parent
ef39c2bfe8
commit
fc01cea80f
|
|
@ -22,6 +22,13 @@ import (
|
|||
|
||||
const protocolVersion = 1
|
||||
const defaultMaxRows = 1000
|
||||
|
||||
var (
|
||||
oraclePlSQLBlockStartRegexp = regexp.MustCompile(`(?is)^\s*(?:DECLARE|BEGIN|CREATE\s+(?:OR\s+REPLACE\s+)?(?:(?:EDITIONABLE|NONEDITIONABLE)\s+)?(?:FUNCTION|PROCEDURE|TRIGGER|PACKAGE(?:\s+BODY)?|TYPE(?:\s+BODY)?))\b`)
|
||||
oraclePlSQLBlockEndRegexp = regexp.MustCompile(`(?is)\bEND\s*;\s*$`)
|
||||
oracleNamedPlSQLBlockEndRegexp = regexp.MustCompile(`(?is)\bEND\s+([A-Z0-9_$#]+)\s*;\s*$`)
|
||||
)
|
||||
|
||||
const oracleListDatabasesSQL = `
|
||||
SELECT username AS owner
|
||||
FROM all_users
|
||||
|
|
@ -2327,7 +2334,52 @@ func parseURLParams(raw string) map[string]string {
|
|||
}
|
||||
|
||||
func trimStatementSQL(sqlText string) string {
|
||||
return strings.TrimRight(strings.TrimSpace(sqlText), "; \t\r\n")
|
||||
trimmed := stripTrailingSlashDelimiter(strings.TrimSpace(sqlText))
|
||||
if isOraclePlSQLBlock(trimmed) {
|
||||
return trimmed
|
||||
}
|
||||
return strings.TrimRight(trimmed, "; \t\r\n")
|
||||
}
|
||||
|
||||
func stripTrailingSlashDelimiter(sqlText string) string {
|
||||
trimmed := strings.TrimSpace(sqlText)
|
||||
if !strings.HasSuffix(trimmed, "/") {
|
||||
return trimmed
|
||||
}
|
||||
slashStart := len(trimmed) - 1
|
||||
lineStart := strings.LastIndex(trimmed[:slashStart], "\n") + 1
|
||||
if strings.TrimSpace(trimmed[lineStart:slashStart]) != "" {
|
||||
return trimmed
|
||||
}
|
||||
beforeSlash := strings.TrimSpace(trimmed[:lineStart])
|
||||
// SQL*Plus uses a standalone slash to execute PL/SQL blocks; go-ora needs
|
||||
// only the block text and not that client-side delimiter.
|
||||
if isOraclePlSQLBlock(beforeSlash) {
|
||||
return beforeSlash
|
||||
}
|
||||
return trimmed
|
||||
}
|
||||
|
||||
func isOraclePlSQLBlock(sqlText string) bool {
|
||||
trimmed := strings.TrimSpace(sqlText)
|
||||
start := trimLeadingSQLComments(trimmed)
|
||||
if !oraclePlSQLBlockStartRegexp.MatchString(start) {
|
||||
return false
|
||||
}
|
||||
upper := strings.ToUpper(trimmed)
|
||||
if oraclePlSQLBlockEndRegexp.MatchString(upper) {
|
||||
return true
|
||||
}
|
||||
matches := oracleNamedPlSQLBlockEndRegexp.FindStringSubmatch(upper)
|
||||
if len(matches) != 2 {
|
||||
return false
|
||||
}
|
||||
switch matches[1] {
|
||||
case "IF", "LOOP", "CASE":
|
||||
return false
|
||||
default:
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
func isQuerySQL(sqlText string) bool {
|
||||
|
|
|
|||
|
|
@ -168,6 +168,55 @@ func TestIsQuerySQLRequiresKeywordBoundary(t *testing.T) {
|
|||
}
|
||||
}
|
||||
|
||||
func TestTrimStatementSQLPreservesAnonymousPLSQLBlockTerminator(t *testing.T) {
|
||||
sqlText := `DECLARE
|
||||
PRE_TRD_DATE INTEGER ;
|
||||
BEGIN
|
||||
SELECT 1 + 2 INTO PRE_TRD_DATE FROM DUAL;
|
||||
END;`
|
||||
|
||||
if got := trimStatementSQL(sqlText); got != sqlText {
|
||||
t.Fatalf("trimStatementSQL() = %q, want full PL/SQL block %q", got, sqlText)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTrimStatementSQLStripsSlashDelimiterAfterPLSQLBlock(t *testing.T) {
|
||||
sqlText := "BEGIN\n NULL;\nEND;\n/"
|
||||
want := "BEGIN\n NULL;\nEND;"
|
||||
|
||||
if got := trimStatementSQL(sqlText); got != want {
|
||||
t.Fatalf("trimStatementSQL() = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTrimStatementSQLPreservesCreatePLSQLObjectTerminator(t *testing.T) {
|
||||
tests := []string{
|
||||
"CREATE OR REPLACE PROCEDURE p AS\nBEGIN\n NULL;\nEND;",
|
||||
"CREATE OR REPLACE FUNCTION f RETURN NUMBER AS\nBEGIN\n RETURN 1;\nEND;",
|
||||
"CREATE OR REPLACE PACKAGE pkg_utils AS\n FUNCTION get_version RETURN VARCHAR2;\nEND pkg_utils;",
|
||||
}
|
||||
for _, sqlText := range tests {
|
||||
if got := trimStatementSQL(sqlText); got != sqlText {
|
||||
t.Fatalf("trimStatementSQL() = %q, want full PL/SQL object %q", got, sqlText)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestTrimStatementSQLStripsSlashDelimiterAfterCreatePLSQLObject(t *testing.T) {
|
||||
sqlText := "CREATE OR REPLACE PROCEDURE p AS\nBEGIN\n NULL;\nEND;\n/"
|
||||
want := "CREATE OR REPLACE PROCEDURE p AS\nBEGIN\n NULL;\nEND;"
|
||||
|
||||
if got := trimStatementSQL(sqlText); got != want {
|
||||
t.Fatalf("trimStatementSQL() = %q, want %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTrimStatementSQLRemovesRegularStatementSemicolon(t *testing.T) {
|
||||
if got := trimStatementSQL("SELECT 1 FROM DUAL;"); got != "SELECT 1 FROM DUAL" {
|
||||
t.Fatalf("trimStatementSQL() = %q, want regular statement without semicolon", got)
|
||||
}
|
||||
}
|
||||
|
||||
func protocolContract(t *testing.T) struct {
|
||||
ProtocolVersion int `json:"protocolVersion"`
|
||||
AllCapabilities []string `json:"allCapabilities"`
|
||||
|
|
|
|||
|
|
@ -65,6 +65,12 @@ END;
|
|||
/
|
||||
SELECT 1;`;
|
||||
|
||||
const oracleIssue2405PlSql = `DECLARE
|
||||
PRE_TRD_DATE INTEGER ;
|
||||
BEGIN
|
||||
SELECT 1 + 2 INTO PRE_TRD_DATE FROM DUAL;
|
||||
END;`;
|
||||
|
||||
const mysqlRoutineFixture = `CREATE PROCEDURE p()
|
||||
BEGIN
|
||||
SELECT 1;
|
||||
|
|
@ -184,6 +190,10 @@ describe("splitSqlStatementRanges", () => {
|
|||
expect(ranges[0].sql).not.toContain("\n/");
|
||||
});
|
||||
|
||||
it("keeps issue #2405 Oracle PL/SQL block together without a slash delimiter", () => {
|
||||
expect(rangeSqlTexts(splitSqlStatementRanges(oracleIssue2405PlSql, "oracle"))).toEqual([oracleIssue2405PlSql]);
|
||||
});
|
||||
|
||||
it("keeps SAP HANA DO blocks together", () => {
|
||||
const ranges = splitSqlStatementRanges(sapHanaDoBlockFixture, "saphana");
|
||||
|
||||
|
|
@ -438,6 +448,12 @@ WHERE request_json LIKE '%"paperFlag":null%';`;
|
|||
expect(range?.sql.trim()).toBe(oraclePlSqlFixture.slice(0, oraclePlSqlFixture.indexOf("\n/")));
|
||||
});
|
||||
|
||||
it("returns the full issue #2405 Oracle PL/SQL block for cursors inside the block", () => {
|
||||
for (const cursor of [indexOf(oracleIssue2405PlSql, "PRE_TRD_DATE"), indexOf(oracleIssue2405PlSql, "SELECT 1 + 2"), indexOf(oracleIssue2405PlSql, "END;")]) {
|
||||
expect(statementRangeAtCursor(oracleIssue2405PlSql, cursor, "oracle")?.sql.trim()).toBe(oracleIssue2405PlSql);
|
||||
}
|
||||
});
|
||||
|
||||
it("returns the full SAP HANA DO block for cursors inside nested statements", () => {
|
||||
const range = statementRangeAtCursor(sapHanaDoBlockFixture, indexOf(sapHanaDoBlockFixture, "Result"), "saphana");
|
||||
|
||||
|
|
@ -477,6 +493,10 @@ describe("executableStatementRanges", () => {
|
|||
expect(rangeSqlTexts(executableStatementRanges(oraclePlSqlFixture, "oracle"))).toEqual([oraclePlSqlFixture.slice(0, oraclePlSqlFixture.indexOf("\n/")), "SELECT 1"]);
|
||||
});
|
||||
|
||||
it("returns the issue #2405 Oracle PL/SQL block as one executable range", () => {
|
||||
expect(rangeSqlTexts(executableStatementRanges(oracleIssue2405PlSql, "oracle"))).toEqual([oracleIssue2405PlSql]);
|
||||
});
|
||||
|
||||
it("does not split executable MySQL routine ranges at inner statements", () => {
|
||||
expect(rangeSqlTexts(executableStatementRanges(mysqlRoutineFixture, "mysql"))).toEqual([mysqlRoutineFixture.slice(0, mysqlRoutineFixture.indexOf("\nSELECT 2;")).replace(/;$/, "").trim(), "SELECT 2"]);
|
||||
});
|
||||
|
|
|
|||
|
|
@ -3014,6 +3014,37 @@ END;";
|
|||
assert_eq!(split_sql_statements_for_database(sql, DatabaseType::Gaussdb), vec![sql.to_string()]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn oracle_like_split_keeps_issue_2405_anonymous_plsql_block_together() {
|
||||
let sql = "\
|
||||
DECLARE
|
||||
PRE_TRD_DATE INTEGER ;
|
||||
BEGIN
|
||||
SELECT 1 + 2 INTO PRE_TRD_DATE FROM DUAL;
|
||||
END;";
|
||||
|
||||
assert_eq!(split_sql_statements_for_database(sql, DatabaseType::Oracle), vec![sql.to_string()]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn oracle_like_current_statement_keeps_issue_2405_anonymous_plsql_block_together() {
|
||||
let sql = "\
|
||||
DECLARE
|
||||
PRE_TRD_DATE INTEGER ;
|
||||
BEGIN
|
||||
SELECT 1 + 2 INTO PRE_TRD_DATE FROM DUAL;
|
||||
END;";
|
||||
let cursors = [
|
||||
sql[..sql.find("PRE_TRD_DATE").unwrap()].encode_utf16().count(),
|
||||
sql[..sql.find("SELECT 1 + 2").unwrap()].encode_utf16().count(),
|
||||
sql[..sql.find("END;").unwrap()].encode_utf16().count(),
|
||||
];
|
||||
|
||||
for cursor in cursors {
|
||||
assert_eq!(find_statement_at_cursor_for_database(sql, cursor, DatabaseType::Oracle), sql);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn oracle_like_split_treats_slash_line_as_plsql_delimiter() {
|
||||
let sql = "\
|
||||
|
|
|
|||
Loading…
Reference in New Issue