fix(oracle): preserve plsql block terminators

This commit is contained in:
LRcoding 2026-07-07 15:29:43 +08:00 committed by GitHub
parent ef39c2bfe8
commit fc01cea80f
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
4 changed files with 153 additions and 1 deletions

View File

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

View File

@ -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"`

View File

@ -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"]);
});

View File

@ -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 = "\