diff --git a/agents/drivers/oracle-go/main.go b/agents/drivers/oracle-go/main.go index 68e21a03a..0c1d24d81 100644 --- a/agents/drivers/oracle-go/main.go +++ b/agents/drivers/oracle-go/main.go @@ -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 { diff --git a/agents/drivers/oracle-go/main_test.go b/agents/drivers/oracle-go/main_test.go index 4e689e3f8..b4ceb9a2a 100644 --- a/agents/drivers/oracle-go/main_test.go +++ b/agents/drivers/oracle-go/main_test.go @@ -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"` diff --git a/apps/desktop/src/lib/__tests__/sql/sqlStatementRanges.spec.ts b/apps/desktop/src/lib/__tests__/sql/sqlStatementRanges.spec.ts index d6e97ce9c..b3804dad2 100644 --- a/apps/desktop/src/lib/__tests__/sql/sqlStatementRanges.spec.ts +++ b/apps/desktop/src/lib/__tests__/sql/sqlStatementRanges.spec.ts @@ -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"]); }); diff --git a/crates/dbx-core/src/sql.rs b/crates/dbx-core/src/sql.rs index d731aae80..59aec8e24 100644 --- a/crates/dbx-core/src/sql.rs +++ b/crates/dbx-core/src/sql.rs @@ -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 = "\