diff --git a/nifi-extension-bundles/nifi-standard-bundle/nifi-standard-processors/src/main/java/org/apache/nifi/processors/standard/db/impl/MSSQL2008DatabaseDialectService.java b/nifi-extension-bundles/nifi-standard-bundle/nifi-standard-processors/src/main/java/org/apache/nifi/processors/standard/db/impl/MSSQL2008DatabaseDialectService.java new file mode 100644 index 000000000000..7985dbfec85e --- /dev/null +++ b/nifi-extension-bundles/nifi-standard-bundle/nifi-standard-processors/src/main/java/org/apache/nifi/processors/standard/db/impl/MSSQL2008DatabaseDialectService.java @@ -0,0 +1,143 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.apache.nifi.processors.standard.db.impl; + +import org.apache.nifi.annotation.documentation.CapabilityDescription; +import org.apache.nifi.annotation.documentation.Tags; +import org.apache.nifi.database.dialect.service.api.PageRequest; +import org.apache.nifi.database.dialect.service.api.QueryStatementRequest; +import org.apache.nifi.database.dialect.service.api.StandardStatementResponse; +import org.apache.nifi.database.dialect.service.api.StatementRequest; +import org.apache.nifi.database.dialect.service.api.StatementResponse; +import org.apache.nifi.database.dialect.service.api.StatementType; +import org.apache.nifi.database.dialect.service.api.TableDefinition; + +import java.util.Optional; + +@Tags({"mssql", "sqlserver", "database", "dialect"}) +@CapabilityDescription("Microsoft SQL Server 2008 Database Dialect Service providing SELECT with ROW_NUMBER() paging, UPSERT using MERGE, and basic ALTER/CREATE DDL generation.") +public class MSSQL2008DatabaseDialectService extends MSSQLDatabaseDialectService { + + @Override + public StatementResponse getStatement(final StatementRequest statementRequest) { + if (statementRequest.statementType() == StatementType.SELECT) { + return new StandardStatementResponse(buildSelect2008(statementRequest)); + } + return super.getStatement(statementRequest); + } + + private String buildSelect2008(final StatementRequest statementRequest) { + if (!(statementRequest instanceof QueryStatementRequest query)) { + throw new IllegalArgumentException("Query Statement Request not found [" + statementRequest.getClass() + "]"); + } + + final TableDefinition table = statementRequest.tableDefinition(); + final String qualifiedTableName = qualifyTableName(table); + + final Optional derivedTable = query.derivedTable(); + if (derivedTable.isPresent()) { + final String tableAlias = "AS " + table.tableName(); + return "SELECT * FROM (" + derivedTable.get() + ") " + tableAlias; + } + + final String selectColumns = buildSelectColumns(table.columns()); + + final Optional page = query.pageRequest(); + final Long limit; + final Long offset; + final String indexColumnName; + if (page.isPresent()) { + final PageRequest p = page.get(); + limit = p.limit().isPresent() ? p.limit().getAsLong() : null; + offset = p.offset(); + indexColumnName = p.indexColumnName().orElse(null); + } else { + limit = null; + offset = null; + indexColumnName = null; + } + + final String whereClause = query.whereClause().orElse(null); + final String orderByClause = query.orderByClause().orElse(null); + + final boolean partitioned = indexColumnName != null && !indexColumnName.isBlank(); + final boolean hasOrder = orderByClause != null && !orderByClause.isBlank(); + // Use window paging only when an offset > 0 is requested; when offset == 0, prefer TOP with ORDER BY for efficiency + final boolean useWindowPaging = limit != null && !partitioned && offset != null && offset > 0; + + final StringBuilder sql = new StringBuilder("SELECT "); + + if (limit != null && !partitioned) { + if (useWindowPaging) { + sql.append("* FROM (SELECT "); + } + final long effectiveOffset = (offset == null) ? 0 : offset; + if (effectiveOffset + limit >= 0) { + sql.append("TOP ").append(effectiveOffset + limit).append(' '); + } + } + + sql.append(selectColumns); + + if (useWindowPaging && hasOrder) { + sql.append(", ROW_NUMBER() OVER(ORDER BY ") + .append(orderByClause) + .append(" asc) rnum"); + } + + sql.append(" FROM ").append(qualifiedTableName); + + boolean whereAdded = false; + if (whereClause != null && !whereClause.isBlank()) { + sql.append(" WHERE ").append(whereClause); + whereAdded = true; + } + + if (partitioned) { + sql.append(whereAdded ? " AND " : " WHERE "); + sql.append(indexColumnName) + .append(" >= ") + .append(offset != null ? offset : 0); + if (limit != null) { + sql.append(" AND ") + .append(indexColumnName) + .append(" < ") + .append((offset == null ? 0 : offset) + limit); + } + } + + if (!partitioned && orderByClause != null && !orderByClause.isBlank()) { + if (!useWindowPaging) { + sql.append(" ORDER BY ").append(orderByClause); + } + } + + if (useWindowPaging) { + if (offset != null && offset > 0 && !hasOrder) { + throw new IllegalArgumentException("Order by clause required for pagination when offset > 0"); + } + sql.append(") A WHERE rnum > ") + .append(offset) + .append(" AND rnum <= ") + .append(offset + limit); + } + + return sql.toString(); + } + +} + diff --git a/nifi-extension-bundles/nifi-standard-bundle/nifi-standard-processors/src/main/java/org/apache/nifi/processors/standard/db/impl/MSSQLDatabaseDialectService.java b/nifi-extension-bundles/nifi-standard-bundle/nifi-standard-processors/src/main/java/org/apache/nifi/processors/standard/db/impl/MSSQLDatabaseDialectService.java new file mode 100644 index 000000000000..f50908404da6 --- /dev/null +++ b/nifi-extension-bundles/nifi-standard-bundle/nifi-standard-processors/src/main/java/org/apache/nifi/processors/standard/db/impl/MSSQLDatabaseDialectService.java @@ -0,0 +1,259 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.apache.nifi.processors.standard.db.impl; + +import org.apache.nifi.annotation.documentation.CapabilityDescription; +import org.apache.nifi.annotation.documentation.Tags; +import org.apache.nifi.controller.AbstractControllerService; +import org.apache.nifi.database.dialect.service.api.ColumnDefinition; +import org.apache.nifi.database.dialect.service.api.DatabaseDialectService; +import org.apache.nifi.database.dialect.service.api.PageRequest; +import org.apache.nifi.database.dialect.service.api.QueryStatementRequest; +import org.apache.nifi.database.dialect.service.api.StandardStatementResponse; +import org.apache.nifi.database.dialect.service.api.StatementRequest; +import org.apache.nifi.database.dialect.service.api.StatementResponse; +import org.apache.nifi.database.dialect.service.api.StatementType; +import org.apache.nifi.database.dialect.service.api.TableDefinition; + +import java.sql.JDBCType; +import java.util.ArrayList; +import java.util.EnumSet; +import java.util.List; +import java.util.Objects; +import java.util.Optional; +import java.util.Set; +import java.util.StringJoiner; + +@Tags({"mssql", "sqlserver", "database", "dialect"}) +@CapabilityDescription("Microsoft SQL Server 2012+ Database Dialect Service providing SELECT with paging, UPSERT using MERGE, and basic ALTER/CREATE DDL generation.") +public class MSSQLDatabaseDialectService extends AbstractControllerService implements DatabaseDialectService { + @Override + public StatementResponse getStatement(final StatementRequest statementRequest) { + Objects.requireNonNull(statementRequest, "Statement Request required"); + + final StatementType statementType = statementRequest.statementType(); + final TableDefinition tableDefinition = statementRequest.tableDefinition(); + + final String sql; + switch (statementType) { // NOPMD: ExhaustiveSwitchHasDefault - All cases are handled, default is for unsupported types + case SELECT: + sql = buildSelect(statementRequest); + break; + case UPSERT: + sql = buildMerge(tableDefinition); + break; + case ALTER: + sql = buildAlter(tableDefinition); + break; + case CREATE: + sql = buildCreate(tableDefinition); + break; + case INSERT_IGNORE: + default: + throw new UnsupportedOperationException("Statement Type [" + statementType + "] not supported"); + } + + return new StandardStatementResponse(sql); + } + + @Override + public Set getSupportedStatementTypes() { + return EnumSet.of(StatementType.ALTER, StatementType.CREATE, StatementType.SELECT, StatementType.UPSERT); + } + + private String buildSelect(final StatementRequest statementRequest) { + if (!(statementRequest instanceof QueryStatementRequest query)) { + throw new IllegalArgumentException("Query Statement Request not found [" + statementRequest.getClass() + "]"); + } + + final TableDefinition table = statementRequest.tableDefinition(); + final String qualifiedTableName = qualifyTableName(table); + + final Optional derivedTable = query.derivedTable(); + if (derivedTable.isPresent()) { + final String tableAlias = "AS " + table.tableName(); + return "SELECT * FROM (" + derivedTable.get() + ") " + tableAlias; + } + + final String selectColumns = buildSelectColumns(table.columns()); + + final Optional page = query.pageRequest(); + final Long limit; + final Long offset; + final String indexColumnName; + if (page.isPresent()) { + final PageRequest p = page.get(); + limit = p.limit().isPresent() ? p.limit().getAsLong() : null; + offset = p.offset(); + indexColumnName = p.indexColumnName().orElse(null); + } else { + limit = null; + offset = null; + indexColumnName = null; + } + + final String whereClause = query.whereClause().orElse(null); + final String orderByClause = query.orderByClause().orElse(null); + + final StringBuilder sql = new StringBuilder("SELECT "); + + final boolean partitioned = indexColumnName != null && !indexColumnName.isBlank(); + final boolean orderByBlank = (orderByClause == null || orderByClause.isBlank()); + if (limit != null && !partitioned && (offset == null || (offset == 0 && orderByBlank))) { + sql.append("TOP ").append(limit).append(' '); + } + + sql.append(selectColumns) + .append(" FROM ") + .append(qualifiedTableName); + + boolean whereAdded = false; + if (whereClause != null && !whereClause.isBlank()) { + sql.append(" WHERE ").append(whereClause); + whereAdded = true; + } + + if (partitioned) { + sql.append(whereAdded ? " AND " : " WHERE "); + sql.append(indexColumnName) + .append(" >= ") + .append(offset != null ? offset : 0); + if (limit != null) { + sql.append(" AND ") + .append(indexColumnName) + .append(" < ") + .append((offset == null ? 0 : offset) + limit); + } + } + + if (!partitioned && orderByClause != null && !orderByClause.isBlank()) { + sql.append(" ORDER BY ").append(orderByClause); + } + + if (!partitioned && limit != null && offset != null) { + if (orderByBlank) { + if (offset > 0) { + throw new IllegalArgumentException("Order by clause cannot be null or empty when using row paging"); + } + } else { + sql.append(" OFFSET ").append(offset).append(" ROWS FETCH NEXT ").append(limit).append(" ROWS ONLY"); + } + } + + return sql.toString(); + } + + private String buildMerge(final TableDefinition table) { + final String tableName = qualifyTableName(table); + + final List columnNames = table.columns().stream().map(ColumnDefinition::columnName).toList(); + final List keyColumnNames = table.columns().stream().filter(ColumnDefinition::primaryKey).map(ColumnDefinition::columnName).toList(); + + if (tableName == null || tableName.isBlank()) { + throw new IllegalArgumentException("Table name cannot be null or blank"); + } + if (columnNames == null || columnNames.isEmpty()) { + throw new IllegalArgumentException("Column names cannot be null or empty"); + } + if (keyColumnNames == null || keyColumnNames.isEmpty()) { + throw new IllegalArgumentException("Key column names cannot be null or empty"); + } + + final String sourceColumns = String.join(", ", columnNames); + final StringJoiner valuesJoiner = new StringJoiner(", "); + columnNames.forEach(col -> valuesJoiner.add("?")); + final String sourceValues = valuesJoiner.toString(); + + final String onClause = String.join(" AND ", keyColumnNames.stream() + .map(k -> "target." + k + " = source." + k) + .toList()); + + final List nonKeyColumns = new ArrayList<>(columnNames); + nonKeyColumns.removeAll(keyColumnNames); + final String updateSetClause = String.join(", ", nonKeyColumns.stream() + .map(c -> c + " = source." + c) + .toList()); + + final String insertValues = String.join(", ", columnNames.stream().map(c -> "source." + c).toList()); + + final StringBuilder sql = new StringBuilder(); + sql.append("MERGE INTO ").append(tableName).append(" AS target ") + .append("USING (VALUES (").append(sourceValues).append(")) AS source (").append(sourceColumns).append(") ") + .append("ON ").append(onClause).append(' '); + + if (!nonKeyColumns.isEmpty()) { + sql.append("WHEN MATCHED THEN UPDATE SET ").append(updateSetClause).append(' '); + } + sql.append("WHEN NOT MATCHED THEN INSERT (").append(sourceColumns).append(") VALUES (").append(insertValues).append(");"); + + return sql.toString(); + } + + private String buildAlter(final TableDefinition table) { + final String tableName = qualifyTableName(table); + final List columnAdds = new ArrayList<>(); + for (ColumnDefinition c : table.columns()) { + final String dataType = JDBCType.valueOf(c.dataType()).getName(); + columnAdds.add(c.columnName() + ' ' + dataType); + } + return "ALTER TABLE " + tableName + " ADD " + String.join(", ", columnAdds); + } + + private String buildCreate(final TableDefinition table) { + final String tableName = qualifyTableName(table); + final List defs = new ArrayList<>(); + for (ColumnDefinition c : table.columns()) { + final String dataType = JDBCType.valueOf(c.dataType()).getName(); + final StringBuilder d = new StringBuilder() + .append(c.columnName()) + .append(' ') + .append(dataType); + if (c.nullable() == ColumnDefinition.Nullable.NO) { + d.append(" NOT NULL"); + } + if (c.primaryKey()) { + d.append(" PRIMARY KEY"); + } + defs.add(d.toString()); + } + return "IF OBJECT_ID('" + tableName + "', 'U') IS NULL CREATE TABLE " + tableName + " (" + String.join(", ", defs) + ")"; + } + + protected String buildSelectColumns(final List columns) { + if (columns == null || columns.isEmpty()) { + return "*"; + } + final StringBuilder sb = new StringBuilder(); + for (int i = 0; i < columns.size(); i++) { + if (i > 0) { + sb.append(", "); + } + sb.append(columns.get(i).columnName()); + } + return sb.toString(); + } + + protected String qualifyTableName(final TableDefinition table) { + final StringBuilder name = new StringBuilder(); + table.catalog().ifPresent(c -> name.append(c).append('.')); + table.schemaName().ifPresent(s -> name.append(s).append('.')); + name.append(table.tableName()); + return name.toString(); + } +} + + diff --git a/nifi-extension-bundles/nifi-standard-bundle/nifi-standard-processors/src/main/resources/META-INF/services/org.apache.nifi.controller.ControllerService b/nifi-extension-bundles/nifi-standard-bundle/nifi-standard-processors/src/main/resources/META-INF/services/org.apache.nifi.controller.ControllerService new file mode 100644 index 000000000000..0b3ff5e72f1f --- /dev/null +++ b/nifi-extension-bundles/nifi-standard-bundle/nifi-standard-processors/src/main/resources/META-INF/services/org.apache.nifi.controller.ControllerService @@ -0,0 +1,16 @@ +# Licensed to the Apache Software Foundation (ASF) under one or more +# contributor license agreements. See the NOTICE file distributed with +# this work for additional information regarding copyright ownership. +# The ASF licenses this file to You under the Apache License, Version 2.0 +# (the "License"); you may not use this file except in compliance with +# the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +org.apache.nifi.processors.standard.db.impl.MSSQLDatabaseDialectService +org.apache.nifi.processors.standard.db.impl.MSSQL2008DatabaseDialectService \ No newline at end of file diff --git a/nifi-extension-bundles/nifi-standard-bundle/nifi-standard-processors/src/test/java/org/apache/nifi/processors/standard/db/impl/TestMSSQL2008DatabaseDialectService.java b/nifi-extension-bundles/nifi-standard-bundle/nifi-standard-processors/src/test/java/org/apache/nifi/processors/standard/db/impl/TestMSSQL2008DatabaseDialectService.java new file mode 100644 index 000000000000..477095feb279 --- /dev/null +++ b/nifi-extension-bundles/nifi-standard-bundle/nifi-standard-processors/src/test/java/org/apache/nifi/processors/standard/db/impl/TestMSSQL2008DatabaseDialectService.java @@ -0,0 +1,281 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.apache.nifi.processors.standard.db.impl; + +import org.apache.nifi.database.dialect.service.api.ColumnDefinition; +import org.apache.nifi.database.dialect.service.api.PageRequest; +import org.apache.nifi.database.dialect.service.api.StandardColumnDefinition; +import org.apache.nifi.database.dialect.service.api.StandardPageRequest; +import org.apache.nifi.database.dialect.service.api.StandardQueryStatementRequest; +import org.apache.nifi.database.dialect.service.api.StandardStatementRequest; +import org.apache.nifi.database.dialect.service.api.StatementRequest; +import org.apache.nifi.database.dialect.service.api.StatementType; +import org.apache.nifi.database.dialect.service.api.TableDefinition; +import org.junit.jupiter.api.Assertions; +import org.junit.jupiter.api.Test; + +import java.util.ArrayList; +import java.util.List; +import java.util.Optional; +import java.util.OptionalLong; + +public class TestMSSQL2008DatabaseDialectService { + private final MSSQL2008DatabaseDialectService service = new MSSQL2008DatabaseDialectService(); + + @Test + public void testPagingQuery2008() { + final TableDefinition table = table("database.tablename", columns(sampleColumns(), new boolean[sampleColumns().size()])); + final StatementRequest req = selectReq(table, sampleColumns(), "", "contain", 100L, 0L, null); + final String sql = service.getStatement(req).sql(); + Assertions.assertEquals("SELECT TOP 100 some(set), of(columns), that, might, contain, methods, a.* FROM database.tablename ORDER BY contain", sql); + } + + @Test + public void testTopOnly2008() { + final TableDefinition table = table("database.tablename", columns(sampleColumns(), new boolean[sampleColumns().size()])); + final StatementRequest req = selectReq(table, sampleColumns(), "", "", 50L, null, null); + final String sql = service.getStatement(req).sql(); + Assertions.assertEquals("SELECT TOP 50 some(set), of(columns), that, might, contain, methods, a.* FROM database.tablename", sql); + } + + private static List sampleColumns() { + return List.of("some(set)", "of(columns)", "that", "might", "contain", "methods", "a.*"); + } + + private static TableDefinition table(final String tableName, final List cols) { + return new TableDefinition(Optional.empty(), Optional.empty(), tableName, cols); + } + + private static List columns(final List names, final boolean[] primaryKeys) { + final List defs = new ArrayList<>(); + for (int i = 0; i < names.size(); i++) { + defs.add(new StandardColumnDefinition(names.get(i), java.sql.Types.VARCHAR, ColumnDefinition.Nullable.YES, primaryKeys.length > i && primaryKeys[i])); + } + return defs; + } + @Test + public void testSelectGeneration2008() { + final TableDefinition table = table("database.tablename", columns(sampleColumns(), new boolean[sampleColumns().size()])); + final StatementRequest req = selectReq(table, sampleColumns(), null, null, null, null, null); + final String sql = service.getStatement(req).sql(); + Assertions.assertEquals("SELECT some(set), of(columns), that, might, contain, methods, a.* FROM database.tablename", sql); + } + + @Test + public void testSelectWhereAndOrder2008() { + final TableDefinition table = table("database.tablename", columns(sampleColumns(), new boolean[sampleColumns().size()])); + final StatementRequest req = selectReq(table, sampleColumns(), "that='some\"' value'", "might DESC", null, null, null); + final String sql = service.getStatement(req).sql(); + Assertions.assertEquals("SELECT some(set), of(columns), that, might, contain, methods, a.* FROM database.tablename WHERE that='some\"' value' ORDER BY might DESC", sql); + } + + @Test + public void testSelectWithDerivedTable2008() { + final TableDefinition table = table("derived_alias", columns(sampleColumns(), new boolean[sampleColumns().size()])); + final StatementRequest req = selectReqWithDerivedTable(table, "SELECT * FROM base_table WHERE condition = 'value'"); + final String sql = service.getStatement(req).sql(); + Assertions.assertEquals("SELECT * FROM (SELECT * FROM base_table WHERE condition = 'value') AS derived_alias", sql); + } + + @Test + public void testSelectWithEmptyColumns2008() { + final TableDefinition table = table("database.tablename", List.of()); + final StatementRequest req = selectReq(table, List.of(), null, null, null, null, null); + final String sql = service.getStatement(req).sql(); + Assertions.assertEquals("SELECT * FROM database.tablename", sql); + } + + @Test + public void testPagingWithOffsetAndOrder2008() { + final TableDefinition table = table("database.tablename", columns(sampleColumns(), new boolean[sampleColumns().size()])); + final StatementRequest req = selectReq(table, sampleColumns(), "status='active'", "id", 10L, 20L, null); + final String sql = service.getStatement(req).sql(); + Assertions.assertEquals( + "SELECT * FROM (SELECT TOP 30 some(set), of(columns), that, might, contain, methods, a.*, ROW_NUMBER() OVER(ORDER BY id asc) rnum " + + "FROM database.tablename WHERE status='active') A WHERE rnum > 20 AND rnum <= 30", + sql + ); + } + + @Test + public void testPagingNoOrderByWithOffset2008() { + final TableDefinition table = table("database.tablename", columns(sampleColumns(), new boolean[sampleColumns().size()])); + final StatementRequest req = selectReq(table, sampleColumns(), "", "", 10L, 5L, null); + final IllegalArgumentException e = Assertions.assertThrows(IllegalArgumentException.class, () -> service.getStatement(req)); + Assertions.assertTrue(e.getMessage().contains("Order by clause required for pagination when offset > 0")); + } + + @Test + public void testPartitionedPaging2008() { + final TableDefinition table = table("database.tablename", columns(sampleColumns(), new boolean[sampleColumns().size()])); + final StatementRequest req1 = selectReq(table, sampleColumns(), "1=1", "contain", 100L, 0L, "contain"); + final String sql1 = service.getStatement(req1).sql(); + Assertions.assertEquals("SELECT some(set), of(columns), that, might, contain, methods, a.* FROM database.tablename WHERE 1=1 AND contain >= 0 AND contain < 100", sql1); + + final StatementRequest req2 = selectReq(table, sampleColumns(), "1=1", "contain", 10000L, 123456L, "contain"); + final String sql2 = service.getStatement(req2).sql(); + Assertions.assertEquals("SELECT some(set), of(columns), that, might, contain, methods, a.* FROM database.tablename WHERE 1=1 AND contain >= 123456 AND contain < 133456", sql2); + } + + @Test + public void testPartitionedPagingWithoutLimit2008() { + final TableDefinition table = table("database.tablename", columns(sampleColumns(), new boolean[sampleColumns().size()])); + final StatementRequest req = selectReq(table, sampleColumns(), "status='active'", "id", null, 1000L, "id"); + final String sql = service.getStatement(req).sql(); + Assertions.assertEquals("SELECT some(set), of(columns), that, might, contain, methods, a.* FROM database.tablename WHERE status='active' AND id >= 1000", sql); + } + + @Test + public void testPartitionedPagingWithBlankIndexColumn2008() { + final TableDefinition table = table("database.tablename", columns(sampleColumns(), new boolean[sampleColumns().size()])); + final StatementRequest req = selectReq(table, sampleColumns(), null, "id", 100L, 50L, " "); + final String sql = service.getStatement(req).sql(); + Assertions.assertEquals( + "SELECT * FROM (SELECT TOP 150 some(set), of(columns), that, might, contain, methods, a.*, ROW_NUMBER() OVER(ORDER BY id asc) rnum " + + "FROM database.tablename) A WHERE rnum > 50 AND rnum <= 150", + sql + ); + } + + @Test + public void testLimitZeroWithOffset2008() { + final TableDefinition table = table("database.tablename", columns(sampleColumns(), new boolean[sampleColumns().size()])); + final StatementRequest req = selectReq(table, sampleColumns(), "", "id", 0L, 10L, null); + final String sql = service.getStatement(req).sql(); + Assertions.assertEquals( + "SELECT * FROM (SELECT TOP 10 some(set), of(columns), that, might, contain, methods, a.*, ROW_NUMBER() OVER(ORDER BY id asc) rnum " + + "FROM database.tablename) A WHERE rnum > 10 AND rnum <= 10", + sql + ); + } + + @Test + public void testLimitZeroWithoutOffset2008() { + final TableDefinition table = table("database.tablename", columns(sampleColumns(), new boolean[sampleColumns().size()])); + final StatementRequest req = selectReq(table, sampleColumns(), "", "", 0L, null, null); + final String sql = service.getStatement(req).sql(); + Assertions.assertEquals("SELECT TOP 0 some(set), of(columns), that, might, contain, methods, a.* FROM database.tablename", sql); + } + + @Test + public void testInheritedUpsertFunctionality2008() { + final List cols = List.of("column1", "column2", "column3", "column4"); + final boolean[] pk = new boolean[]{false, true, false, true}; + final TableDefinition table = table("table", columns(cols, pk)); + final StatementRequest req = upsertReq(table); + final String expected = "MERGE INTO table AS target USING (VALUES (?, ?, ?, ?)) AS source (column1, column2, column3, column4) " + + "ON target.column2 = source.column2 AND target.column4 = source.column4 " + + "WHEN MATCHED THEN UPDATE SET column1 = source.column1, column3 = source.column3 " + + "WHEN NOT MATCHED THEN INSERT (column1, column2, column3, column4) VALUES (source.column1, source.column2, source.column3, source.column4);"; + Assertions.assertEquals(expected, service.getStatement(req).sql()); + } + + @Test + public void testInheritedCreateFunctionality2008() { + final List cols = List.of("id", "name", "active"); + final boolean[] pk = new boolean[]{true, false, false}; + final TableDefinition table = table("users", columns(cols, pk)); + final StatementRequest req = createReq(table); + final String expected = "IF OBJECT_ID('users', 'U') IS NULL CREATE TABLE users (id VARCHAR PRIMARY KEY, name VARCHAR, active VARCHAR)"; + Assertions.assertEquals(expected, service.getStatement(req).sql()); + } + + @Test + public void testInheritedAlterFunctionality2008() { + final List cols = List.of("new_column1", "new_column2"); + final boolean[] pk = new boolean[]{false, false}; + final TableDefinition table = table("existing_table", columns(cols, pk)); + final StatementRequest req = alterReq(table); + final String expected = "ALTER TABLE existing_table ADD new_column1 VARCHAR, new_column2 VARCHAR"; + Assertions.assertEquals(expected, service.getStatement(req).sql()); + } + + @Test + public void testInheritedUnsupportedStatementType2008() { + final List cols = List.of("col1"); + final boolean[] pk = new boolean[]{false}; + final TableDefinition table = table("table", columns(cols, pk)); + final StatementRequest req = insertIgnoreReq(table); + final UnsupportedOperationException e = Assertions.assertThrows(UnsupportedOperationException.class, () -> service.getStatement(req)); + Assertions.assertTrue(e.getMessage().contains("Statement Type [INSERT_IGNORE] not supported")); + } + + @Test + public void testSelectWithNonQueryStatementRequest2008() { + final List cols = List.of("col1"); + final boolean[] pk = new boolean[]{false}; + final TableDefinition table = table("table", columns(cols, pk)); + final StatementRequest req = new StandardStatementRequest(StatementType.SELECT, table); + final IllegalArgumentException e = Assertions.assertThrows(IllegalArgumentException.class, () -> service.getStatement(req)); + Assertions.assertTrue(e.getMessage().contains("Query Statement Request not found")); + } + + private static StatementRequest selectReq(final TableDefinition table, + final List columnNames, + final String where, + final String orderBy, + final Long limit, + final Long offset, + final String indexColumn) { + return new StandardQueryStatementRequest( + StatementType.SELECT, + table, + Optional.empty(), + Optional.ofNullable(where).filter(s -> !s.isEmpty()), + Optional.ofNullable(orderBy).filter(s -> !s.isEmpty()), + pageRequest(limit, offset, indexColumn) + ); + } + + private static Optional pageRequest(final Long limit, final Long offset, final String indexColumn) { + if (limit == null && offset == null && indexColumn == null) { + return Optional.empty(); + } + return Optional.of(new StandardPageRequest( + offset == null ? 0L : offset, + limit == null ? OptionalLong.empty() : OptionalLong.of(limit), + Optional.ofNullable(indexColumn) + )); + } + + private static StatementRequest selectReqWithDerivedTable(final TableDefinition table, final String derivedTableSql) { + return new StandardQueryStatementRequest( + StatementType.SELECT, + table, + Optional.of(derivedTableSql), + Optional.empty(), + Optional.empty(), + Optional.empty() + ); + } + + private static StatementRequest upsertReq(final TableDefinition table) { + return new StandardStatementRequest(StatementType.UPSERT, table); + } + + private static StatementRequest createReq(final TableDefinition table) { + return new StandardStatementRequest(StatementType.CREATE, table); + } + + private static StatementRequest alterReq(final TableDefinition table) { + return new StandardStatementRequest(StatementType.ALTER, table); + } + + private static StatementRequest insertIgnoreReq(final TableDefinition table) { + return new StandardStatementRequest(StatementType.INSERT_IGNORE, table); + } +} diff --git a/nifi-extension-bundles/nifi-standard-bundle/nifi-standard-processors/src/test/java/org/apache/nifi/processors/standard/db/impl/TestMSSQLDatabaseDialectService.java b/nifi-extension-bundles/nifi-standard-bundle/nifi-standard-processors/src/test/java/org/apache/nifi/processors/standard/db/impl/TestMSSQLDatabaseDialectService.java new file mode 100644 index 000000000000..9d8f8b9168b3 --- /dev/null +++ b/nifi-extension-bundles/nifi-standard-bundle/nifi-standard-processors/src/test/java/org/apache/nifi/processors/standard/db/impl/TestMSSQLDatabaseDialectService.java @@ -0,0 +1,338 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +package org.apache.nifi.processors.standard.db.impl; + +import org.apache.nifi.database.dialect.service.api.ColumnDefinition; +import org.apache.nifi.database.dialect.service.api.PageRequest; +import org.apache.nifi.database.dialect.service.api.StandardColumnDefinition; +import org.apache.nifi.database.dialect.service.api.StandardPageRequest; +import org.apache.nifi.database.dialect.service.api.StandardQueryStatementRequest; +import org.apache.nifi.database.dialect.service.api.StandardStatementRequest; +import org.apache.nifi.database.dialect.service.api.StatementRequest; +import org.apache.nifi.database.dialect.service.api.StatementType; +import org.apache.nifi.database.dialect.service.api.TableDefinition; +import org.junit.jupiter.api.Test; + +import java.util.ArrayList; +import java.util.EnumSet; +import java.util.List; +import java.util.Optional; +import java.util.OptionalLong; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; + +public class TestMSSQLDatabaseDialectService { + private final MSSQLDatabaseDialectService service = new MSSQLDatabaseDialectService(); + + @Test + public void testSelectGeneration() { + final TableDefinition table = table("database.tablename", columns(sampleColumns(), new boolean[sampleColumns().size()])); + final StatementRequest req = selectReq(table, sampleColumns(), null, null, null, null, null); + final String sql = service.getStatement(req).sql(); + assertEquals("SELECT some(set), of(columns), that, might, contain, methods, a.* FROM database.tablename", sql); + } + + @Test + public void testSelectWhereAndOrder() { + final TableDefinition table = table("database.tablename", columns(sampleColumns(), new boolean[sampleColumns().size()])); + final StatementRequest req = selectReq(table, sampleColumns(), "that='some\"' value'", "might DESC", null, null, null); + final String sql = service.getStatement(req).sql(); + assertEquals("SELECT some(set), of(columns), that, might, contain, methods, a.* FROM database.tablename WHERE that='some\"' value' ORDER BY might DESC", sql); + } + + @Test + public void testTopQuery() { + final TableDefinition table = table("database.tablename", columns(sampleColumns(), new boolean[sampleColumns().size()])); + final StatementRequest req = selectReq(table, sampleColumns(), "", "", 100L, null, null); + final String sql = service.getStatement(req).sql(); + assertEquals("SELECT TOP 100 some(set), of(columns), that, might, contain, methods, a.* FROM database.tablename", sql); + } + + @Test + public void testPagingNoOrderBy() { + final TableDefinition table = table("database.tablename", columns(sampleColumns(), new boolean[sampleColumns().size()])); + final StatementRequest req = selectReq(table, sampleColumns(), "", "", 10L, 1L, null); + final IllegalArgumentException e = assertThrows(IllegalArgumentException.class, () -> service.getStatement(req)); + assertTrue(e.getMessage().contains("Order by clause cannot be null or empty when using row paging")); + } + + @Test + public void testPagingQuery() { + final TableDefinition table = table("database.tablename", columns(sampleColumns(), new boolean[sampleColumns().size()])); + final StatementRequest req = selectReq(table, sampleColumns(), "", "contain", 100L, 0L, null); + final String sql = service.getStatement(req).sql(); + assertEquals("SELECT some(set), of(columns), that, might, contain, methods, a.* FROM database.tablename ORDER BY contain OFFSET 0 ROWS FETCH NEXT 100 ROWS ONLY", sql); + } + + @Test + public void testPartitionedPaging() { + final TableDefinition table = table("database.tablename", columns(sampleColumns(), new boolean[sampleColumns().size()])); + final StatementRequest req1 = selectReq(table, sampleColumns(), "1=1", "contain", 100L, 0L, "contain"); + final String sql1 = service.getStatement(req1).sql(); + assertEquals("SELECT some(set), of(columns), that, might, contain, methods, a.* FROM database.tablename WHERE 1=1 AND contain >= 0 AND contain < 100", sql1); + + final StatementRequest req2 = selectReq(table, sampleColumns(), "1=1", "contain", 10000L, 123456L, "contain"); + final String sql2 = service.getStatement(req2).sql(); + assertEquals("SELECT some(set), of(columns), that, might, contain, methods, a.* FROM database.tablename WHERE 1=1 AND contain >= 123456 AND contain < 133456", sql2); + } + + @Test + public void testSupportsUpsert() { + assertTrue(service.getSupportedStatementTypes().containsAll(EnumSet.of(StatementType.UPSERT, StatementType.SELECT))); + } + + @Test + public void testUpsertMerge() { + final List cols = List.of("column1", "column2", "column3", "column4"); + final boolean[] pk = new boolean[]{ + false, true, false, true}; + final TableDefinition table = table("table", columns(cols, pk)); + final StatementRequest req = upsertReq(table); + final String expected = "MERGE INTO table AS target USING (VALUES (?, ?, ?, ?)) AS source (column1, column2, column3, column4) " + + "ON target.column2 = source.column2 AND target.column4 = source.column4 " + + "WHEN MATCHED THEN UPDATE SET column1 = source.column1, column3 = source.column3 " + + "WHEN NOT MATCHED THEN INSERT (column1, column2, column3, column4) VALUES (source.column1, source.column2, source.column3, source.column4);"; + assertEquals(expected, service.getStatement(req).sql()); + } + + @Test + public void testUpsertValidation() { + final List cols = List.of("column1"); + final boolean[] pk = new boolean[]{ + false}; + final TableDefinition t1 = table(" ", columns(cols, pk)); + final IllegalArgumentException e1 = assertThrows(IllegalArgumentException.class, () -> service.getStatement(upsertReq(t1))); + assertTrue(e1.getMessage().contains("Table name cannot be null or blank")); + } + @Test + public void testCreateStatement() { + final List cols = List.of("id", "name", "email", "active"); + final boolean[] pk = new boolean[]{true, false, false, false}; + final TableDefinition table = table("users", columns(cols, pk)); + final StatementRequest req = createReq(table); + final String expected = "IF OBJECT_ID('users', 'U') IS NULL CREATE TABLE users (id VARCHAR PRIMARY KEY, name VARCHAR, email VARCHAR, active VARCHAR)"; + assertEquals(expected, service.getStatement(req).sql()); + } + + @Test + public void testCreateStatementWithNullableColumns() { + final List cols = List.of("id", "name"); + final boolean[] pk = new boolean[]{true, false}; + final boolean[] nullable = new boolean[]{false, true}; + final TableDefinition table = table("test_table", columns(cols, pk, nullable)); + final StatementRequest req = createReq(table); + final String expected = "IF OBJECT_ID('test_table', 'U') IS NULL CREATE TABLE test_table (id VARCHAR NOT NULL PRIMARY KEY, name VARCHAR)"; + assertEquals(expected, service.getStatement(req).sql()); + } + + @Test + public void testCreateStatementWithQualifiedTableName() { + final List cols = List.of("id", "data"); + final boolean[] pk = new boolean[]{true, false}; + final TableDefinition table = qualifiedTable("catalog", "schema", "table", columns(cols, pk)); + final StatementRequest req = createReq(table); + final String expected = "IF OBJECT_ID('catalog.schema.table', 'U') IS NULL CREATE TABLE catalog.schema.table (id VARCHAR PRIMARY KEY, data VARCHAR)"; + assertEquals(expected, service.getStatement(req).sql()); + } + + @Test + public void testAlterStatement() { + final List cols = List.of("new_column1", "new_column2"); + final boolean[] pk = new boolean[]{false, false}; + final TableDefinition table = table("existing_table", columns(cols, pk)); + final StatementRequest req = alterReq(table); + final String expected = "ALTER TABLE existing_table ADD new_column1 VARCHAR, new_column2 VARCHAR"; + assertEquals(expected, service.getStatement(req).sql()); + } + + @Test + public void testAlterStatementWithQualifiedTableName() { + final List cols = List.of("new_col"); + final boolean[] pk = new boolean[]{false}; + final TableDefinition table = qualifiedTable("db", "dbo", "table1", columns(cols, pk)); + final StatementRequest req = alterReq(table); + final String expected = "ALTER TABLE db.dbo.table1 ADD new_col VARCHAR"; + assertEquals(expected, service.getStatement(req).sql()); + } + + @Test + public void testUnsupportedStatementType() { + final List cols = List.of("col1"); + final boolean[] pk = new boolean[]{false}; + final TableDefinition table = table("table", columns(cols, pk)); + final StatementRequest req = insertIgnoreReq(table); + final UnsupportedOperationException e = assertThrows(UnsupportedOperationException.class, () -> service.getStatement(req)); + assertTrue(e.getMessage().contains("Statement Type [INSERT_IGNORE] not supported")); + } + + @Test + public void testSelectWithDerivedTable() { + final TableDefinition table = table("derived_alias", columns(sampleColumns(), new boolean[sampleColumns().size()])); + final StatementRequest req = selectReqWithDerivedTable(table, "SELECT * FROM base_table WHERE condition = 'value'"); + final String sql = service.getStatement(req).sql(); + assertEquals("SELECT * FROM (SELECT * FROM base_table WHERE condition = 'value') AS derived_alias", sql); + } + + @Test + public void testSelectWithEmptyColumns() { + final TableDefinition table = table("database.tablename", List.of()); + final StatementRequest req = selectReq(table, List.of(), null, null, null, null, null); + final String sql = service.getStatement(req).sql(); + assertEquals("SELECT * FROM database.tablename", sql); + } + + @Test + public void testPartitionedPagingWithoutLimit() { + final TableDefinition table = table("database.tablename", columns(sampleColumns(), new boolean[sampleColumns().size()])); + final StatementRequest req = selectReq(table, sampleColumns(), "status='active'", "id", null, 1000L, "id"); + final String sql = service.getStatement(req).sql(); + assertEquals("SELECT some(set), of(columns), that, might, contain, methods, a.* FROM database.tablename WHERE status='active' AND id >= 1000", sql); + } + + @Test + public void testPartitionedPagingWithBlankIndexColumn() { + final TableDefinition table = table("database.tablename", columns(sampleColumns(), new boolean[sampleColumns().size()])); + final StatementRequest req = selectReq(table, sampleColumns(), null, "id", 100L, 0L, " "); + final String sql = service.getStatement(req).sql(); + assertEquals("SELECT some(set), of(columns), that, might, contain, methods, a.* FROM database.tablename ORDER BY id OFFSET 0 ROWS FETCH NEXT 100 ROWS ONLY", sql); + } + + @Test + public void testUpsertWithAllPrimaryKeyColumns() { + final List cols = List.of("key1", "key2"); + final boolean[] pk = new boolean[]{true, true}; + final TableDefinition table = table("lookup_table", columns(cols, pk)); + final StatementRequest req = upsertReq(table); + final String expected = "MERGE INTO lookup_table AS target USING (VALUES (?, ?)) AS source (key1, key2) " + + "ON target.key1 = source.key1 AND target.key2 = source.key2 " + + "WHEN NOT MATCHED THEN INSERT (key1, key2) VALUES (source.key1, source.key2);"; + assertEquals(expected, service.getStatement(req).sql()); + } + + @Test + public void testUpsertValidationEmptyColumns() { + final TableDefinition table = table("table", List.of()); + final IllegalArgumentException e = assertThrows(IllegalArgumentException.class, () -> service.getStatement(upsertReq(table))); + assertTrue(e.getMessage().contains("Column names cannot be null or empty")); + } + + @Test + public void testUpsertValidationNoPrimaryKeys() { + final List cols = List.of("col1", "col2"); + final boolean[] pk = new boolean[]{false, false}; + final TableDefinition table = table("table", columns(cols, pk)); + final IllegalArgumentException e = assertThrows(IllegalArgumentException.class, () -> service.getStatement(upsertReq(table))); + assertTrue(e.getMessage().contains("Key column names cannot be null or empty")); + } + + @Test + public void testSelectWithNonQueryStatementRequest() { + final List cols = List.of("col1"); + final boolean[] pk = new boolean[]{false}; + final TableDefinition table = table("table", columns(cols, pk)); + final StatementRequest req = new StandardStatementRequest(StatementType.SELECT, table); + final IllegalArgumentException e = assertThrows(IllegalArgumentException.class, () -> service.getStatement(req)); + assertTrue(e.getMessage().contains("Query Statement Request not found")); + } + + private static List sampleColumns() { + return List.of("some(set)", "of(columns)", "that", "might", "contain", "methods", "a.*"); + } + + private static TableDefinition table(final String tableName, final List cols) { + return new TableDefinition(Optional.empty(), Optional.empty(), tableName, cols); + } + + private static List columns(final List names, final boolean[] primaryKeys) { + final List defs = new ArrayList<>(); + for (int i = 0; i < names.size(); i++) { + defs.add(new StandardColumnDefinition(names.get(i), java.sql.Types.VARCHAR, ColumnDefinition.Nullable.YES, primaryKeys.length > i && primaryKeys[i])); + } + return defs; + } + + private static StatementRequest selectReq(final TableDefinition table, + final List columnNames, + final String where, + final String orderBy, + final Long limit, + final Long offset, + final String indexColumn) { + return new StandardQueryStatementRequest( + StatementType.SELECT, + table, + Optional.empty(), + Optional.ofNullable(where).filter(s -> !s.isEmpty()), + Optional.ofNullable(orderBy).filter(s -> !s.isEmpty()), + pageRequest(limit, offset, indexColumn) + ); + } + + private static Optional pageRequest(final Long limit, final Long offset, final String indexColumn) { + if (limit == null && offset == null && indexColumn == null) { + return Optional.empty(); + } + return Optional.of(new StandardPageRequest( + offset == null ? 0L : offset, + limit == null ? OptionalLong.empty() : OptionalLong.of(limit), + Optional.ofNullable(indexColumn) + )); + } + + private static StatementRequest upsertReq(final TableDefinition table) { + return new StandardStatementRequest(StatementType.UPSERT, table); + } + private static TableDefinition qualifiedTable(final String catalog, final String schema, final String tableName, final List cols) { + return new TableDefinition(Optional.ofNullable(catalog), Optional.ofNullable(schema), tableName, cols); + } + + private static List columns(final List names, final boolean[] primaryKeys, final boolean[] nullable) { + final List defs = new ArrayList<>(); + for (int i = 0; i < names.size(); i++) { + final ColumnDefinition.Nullable nullableValue = (nullable.length > i && !nullable[i]) + ? ColumnDefinition.Nullable.NO + : ColumnDefinition.Nullable.YES; + defs.add(new StandardColumnDefinition(names.get(i), java.sql.Types.VARCHAR, nullableValue, primaryKeys.length > i && primaryKeys[i])); + } + return defs; + } + + private static StatementRequest createReq(final TableDefinition table) { + return new StandardStatementRequest(StatementType.CREATE, table); + } + + private static StatementRequest alterReq(final TableDefinition table) { + return new StandardStatementRequest(StatementType.ALTER, table); + } + + private static StatementRequest insertIgnoreReq(final TableDefinition table) { + return new StandardStatementRequest(StatementType.INSERT_IGNORE, table); + } + + private static StatementRequest selectReqWithDerivedTable(final TableDefinition table, final String derivedTableSql) { + return new StandardQueryStatementRequest( + StatementType.SELECT, + table, + Optional.of(derivedTableSql), + Optional.empty(), + Optional.empty(), + Optional.empty() + ); + } +}