package sqlancer.postgres; import java.io.IOException; import java.net.URI; import java.net.URISyntaxException; import java.sql.Connection; import java.sql.DriverManager; import java.sql.SQLException; import java.sql.Statement; import java.util.ArrayList; import java.util.Arrays; import java.util.LinkedList; import java.util.List; import java.util.Queue; import java.util.stream.Collectors; import com.fasterxml.jackson.databind.JsonNode; import com.fasterxml.jackson.databind.ObjectMapper; import com.google.auto.service.AutoService; import sqlancer.AbstractAction; import sqlancer.DatabaseProvider; import sqlancer.IgnoreMeException; import sqlancer.MainOptions; import sqlancer.Randomly; import sqlancer.SQLConnection; import sqlancer.SQLProviderAdapter; import sqlancer.StatementExecutor; import sqlancer.common.DBMSCommon; import sqlancer.common.query.SQLQueryAdapter; import sqlancer.common.query.SQLQueryProvider; import sqlancer.common.query.SQLancerResultSet; import sqlancer.postgres.gen.PostgresAlterTableGenerator; import sqlancer.postgres.gen.PostgresAnalyzeGenerator; import sqlancer.postgres.gen.PostgresClusterGenerator; import sqlancer.postgres.gen.PostgresCommentGenerator; import sqlancer.postgres.gen.PostgresDeleteGenerator; import sqlancer.postgres.gen.PostgresDiscardGenerator; import sqlancer.postgres.gen.PostgresDropIndexGenerator; import sqlancer.postgres.gen.PostgresExplainGenerator; import sqlancer.postgres.gen.PostgresIndexGenerator; import sqlancer.postgres.gen.PostgresInsertGenerator; import sqlancer.postgres.gen.PostgresNotifyGenerator; import sqlancer.postgres.gen.PostgresReindexGenerator; import sqlancer.postgres.gen.PostgresSequenceGenerator; import sqlancer.postgres.gen.PostgresSetGenerator; import sqlancer.postgres.gen.PostgresStatisticsGenerator; import sqlancer.postgres.gen.PostgresTableGenerator; import sqlancer.postgres.gen.PostgresTableSpaceGenerator; import sqlancer.postgres.gen.PostgresTransactionGenerator; import sqlancer.postgres.gen.PostgresTruncateGenerator; import sqlancer.postgres.gen.PostgresUpdateGenerator; import sqlancer.postgres.gen.PostgresVacuumGenerator; import sqlancer.postgres.gen.PostgresViewGenerator; // EXISTS // IN @AutoService(DatabaseProvider.class) public class PostgresProvider extends SQLProviderAdapter { /** * Generate only data types and expressions that are understood by PQS. */ public static boolean generateOnlyKnown; protected String entryURL; protected String username; protected String password; protected String entryPath; protected String host; protected int port; protected String testURL; protected String databaseName; protected String createDatabaseCommand; protected String extensionsList; public PostgresProvider() { super(PostgresGlobalState.class, PostgresOptions.class); } protected PostgresProvider(Class globalClass, Class optionClass) { super(globalClass, optionClass); } public enum Action implements AbstractAction { ANALYZE(PostgresAnalyzeGenerator::create), // ALTER_TABLE(g -> PostgresAlterTableGenerator.create(g.getSchema().getRandomTable(t -> !t.isView()), g, generateOnlyKnown)), // CLUSTER(PostgresClusterGenerator::create), // COMMIT(g -> { SQLQueryAdapter query; if (Randomly.getBoolean()) { query = new SQLQueryAdapter("COMMIT", true); } else if (Randomly.getBoolean()) { query = PostgresTransactionGenerator.executeBegin(); } else { query = new SQLQueryAdapter("ROLLBACK", true); } return query; }), // CREATE_STATISTICS(PostgresStatisticsGenerator::insert), // DROP_STATISTICS(PostgresStatisticsGenerator::remove), // ALTER_STATISTICS(PostgresStatisticsGenerator::alter), // DELETE(PostgresDeleteGenerator::create), // DISCARD(PostgresDiscardGenerator::create), // DROP_INDEX(PostgresDropIndexGenerator::create), // INSERT(PostgresInsertGenerator::insert), // UPDATE(PostgresUpdateGenerator::create), // TRUNCATE(PostgresTruncateGenerator::create), // VACUUM(PostgresVacuumGenerator::create), // REINDEX(PostgresReindexGenerator::create), // SET(PostgresSetGenerator::create), // CREATE_INDEX(PostgresIndexGenerator::generate), // SET_CONSTRAINTS((g) -> { StringBuilder sb = new StringBuilder(); sb.append("SET CONSTRAINTS ALL "); sb.append(Randomly.fromOptions("DEFERRED", "IMMEDIATE")); return new SQLQueryAdapter(sb.toString()); }), // RESET_ROLE((g) -> new SQLQueryAdapter("RESET ROLE")), // COMMENT_ON(PostgresCommentGenerator::generate), // RESET((g) -> new SQLQueryAdapter("RESET ALL") /* * https://www.postgresql.org/docs/13/sql-reset.html TODO: also * configuration parameter */), // NOTIFY(PostgresNotifyGenerator::createNotify), // LISTEN((g) -> PostgresNotifyGenerator.createListen()), // UNLISTEN((g) -> PostgresNotifyGenerator.createUnlisten()), // CREATE_SEQUENCE(PostgresSequenceGenerator::createSequence), // CREATE_VIEW(PostgresViewGenerator::create), // CREATE_TABLESPACE(PostgresTableSpaceGenerator::generate); private final SQLQueryProvider sqlQueryProvider; Action(SQLQueryProvider sqlQueryProvider) { this.sqlQueryProvider = sqlQueryProvider; } @Override public SQLQueryAdapter getQuery(PostgresGlobalState state) throws Exception { return sqlQueryProvider.getQuery(state); } } protected static int mapActions(PostgresGlobalState globalState, Action a) { Randomly r = globalState.getRandomly(); int nrPerformed; switch (a) { case CREATE_INDEX: case CLUSTER: nrPerformed = r.getInteger(0, 3); break; case CREATE_STATISTICS: nrPerformed = r.getInteger(0, 5); break; case ALTER_STATISTICS: nrPerformed = r.getInteger(0, 2); break; case DISCARD: case DROP_INDEX: nrPerformed = r.getInteger(0, 5); break; case COMMIT: nrPerformed = r.getInteger(0, 0); break; case ALTER_TABLE: nrPerformed = r.getInteger(0, 5); break; case REINDEX: case RESET: nrPerformed = r.getInteger(0, 3); break; case DELETE: case RESET_ROLE: case SET: nrPerformed = r.getInteger(0, 5); break; case ANALYZE: nrPerformed = r.getInteger(0, 3); break; case VACUUM: case SET_CONSTRAINTS: case COMMENT_ON: case NOTIFY: case LISTEN: case UNLISTEN: case CREATE_SEQUENCE: case DROP_STATISTICS: case TRUNCATE: nrPerformed = r.getInteger(0, 2); break; case CREATE_VIEW: nrPerformed = r.getInteger(0, 2); break; case CREATE_TABLESPACE: nrPerformed = r.getInteger(0, 2); break; case UPDATE: nrPerformed = r.getInteger(0, 10); break; case INSERT: nrPerformed = r.getInteger(0, globalState.getOptions().getMaxNumberInserts()); break; default: throw new AssertionError(a); } return nrPerformed; } @Override public void generateDatabase(PostgresGlobalState globalState) throws Exception { readFunctions(globalState); createTables(globalState, Randomly.fromOptions(4, 5, 6)); prepareTables(globalState); extensionsList = globalState.getDbmsSpecificOptions().extensions; if (!extensionsList.isEmpty()) { String[] extensionNames = extensionsList.split(","); /* * To avoid of a test interference with an extension objects, create them in a separate schema. Of course, * they must be truly relocatable. */ globalState.executeStatement(new SQLQueryAdapter("CREATE SCHEMA extensions;", true)); for (int i = 0; i < extensionNames.length; i++) { globalState.executeStatement(new SQLQueryAdapter( "CREATE EXTENSION " + extensionNames[i] + " WITH SCHEMA extensions;", true)); } } } @Override public SQLConnection createDatabase(PostgresGlobalState globalState) throws SQLException { if (globalState.getDbmsSpecificOptions().getTestOracleFactory().stream() .anyMatch((o) -> o == PostgresOracleFactory.PQS)) { generateOnlyKnown = true; } username = globalState.getOptions().getUserName(); password = globalState.getOptions().getPassword(); host = globalState.getOptions().getHost(); port = globalState.getOptions().getPort(); entryPath = "/test"; entryURL = globalState.getDbmsSpecificOptions().connectionURL; // trim URL to exclude "jdbc:" if (entryURL.startsWith("jdbc:")) { entryURL = entryURL.substring(5); } String entryDatabaseName = entryPath.substring(1); databaseName = globalState.getDatabaseName(); try { URI uri = new URI(entryURL); String userInfoURI = uri.getUserInfo(); String pathURI = uri.getPath(); if (userInfoURI != null) { // username and password specified in URL take precedence if (userInfoURI.contains(":")) { String[] userInfo = userInfoURI.split(":", 2); username = userInfo[0]; password = userInfo[1]; } else { username = userInfoURI; password = null; } int userInfoIndex = entryURL.indexOf(userInfoURI); String preUserInfo = entryURL.substring(0, userInfoIndex); String postUserInfo = entryURL.substring(userInfoIndex + userInfoURI.length() + 1); entryURL = preUserInfo + postUserInfo; } if (pathURI != null) { entryPath = pathURI; } if (host == null) { host = uri.getHost(); } if (port == MainOptions.NO_SET_PORT) { port = uri.getPort(); } entryURL = String.format("%s://%s:%d/%s", uri.getScheme(), host, port, entryDatabaseName); } catch (URISyntaxException e) { throw new AssertionError(e); } Connection con = DriverManager.getConnection("jdbc:" + entryURL, username, password); globalState.getState().logStatement(String.format("\\c %s;", entryDatabaseName)); globalState.getState().logStatement("DROP DATABASE IF EXISTS " + databaseName); createDatabaseCommand = getCreateDatabaseCommand(globalState); globalState.getState().logStatement(createDatabaseCommand); try (Statement s = con.createStatement()) { s.execute("DROP DATABASE IF EXISTS " + databaseName); } try (Statement s = con.createStatement()) { s.execute(createDatabaseCommand); } con.close(); int databaseIndex = entryURL.indexOf(entryDatabaseName); String preDatabaseName = entryURL.substring(0, databaseIndex); String postDatabaseName = entryURL.substring(databaseIndex + entryDatabaseName.length()); testURL = preDatabaseName + databaseName + postDatabaseName; globalState.getState().logStatement(String.format("\\c %s;", databaseName)); con = DriverManager.getConnection("jdbc:" + testURL, username, password); return new SQLConnection(con); } protected void readFunctions(PostgresGlobalState globalState) throws SQLException { SQLQueryAdapter query = new SQLQueryAdapter("SELECT proname, provolatile FROM pg_proc;"); SQLancerResultSet rs = query.executeAndGet(globalState); while (rs.next()) { String functionName = rs.getString(1); Character functionType = rs.getString(2).charAt(0); globalState.addFunctionAndType(functionName, functionType); } } protected void createTables(PostgresGlobalState globalState, int numTables) throws Exception { while (globalState.getSchema().getDatabaseTables().size() < numTables) { try { String tableName = DBMSCommon.createTableName(globalState.getSchema().getDatabaseTables().size()); SQLQueryAdapter createTable = PostgresTableGenerator.generate(tableName, globalState.getSchema(), generateOnlyKnown, globalState); globalState.executeStatement(createTable); } catch (IgnoreMeException e) { } } } protected void prepareTables(PostgresGlobalState globalState) throws Exception { StatementExecutor se = new StatementExecutor<>(globalState, Action.values(), PostgresProvider::mapActions, (q) -> { if (globalState.getSchema().getDatabaseTables().isEmpty()) { throw new IgnoreMeException(); } }); se.executeStatements(); globalState.executeStatement(new SQLQueryAdapter("COMMIT", true)); globalState.executeStatement(new SQLQueryAdapter("SET SESSION statement_timeout = 5000;\n")); } private String getCreateDatabaseCommand(PostgresGlobalState state) { StringBuilder sb = new StringBuilder(); sb.append("CREATE DATABASE " + databaseName + " "); if (((PostgresOptions) state.getDbmsSpecificOptions()).testCollations) { if (Randomly.getBoolean()) { if (Randomly.getBoolean()) { sb.append("WITH ENCODING '"); sb.append(Randomly.fromOptions("utf8")); sb.append("' "); } if (Randomly.getBoolean() && !state.getCollates().isEmpty()) { sb.append(String.format(" LOCALE = '%s' ", Randomly.fromList(state.getCollates()))); } else { for (String lc : Arrays.asList("LC_COLLATE", "LC_CTYPE")) { if (!state.getCollates().isEmpty() && Randomly.getBoolean()) { sb.append(String.format(" %s = '%s'", lc, Randomly.fromList(state.getCollates()))); } } } sb.append(" TEMPLATE template0"); } } else { sb.append("WITH ENCODING 'UTF8' TEMPLATE template0"); } return sb.toString(); } @Override public String getDBMSName() { return "postgres"; } @Override public String getQueryPlan(String selectStr, PostgresGlobalState globalState) throws Exception { String queryPlan = ""; if (globalState.getOptions().logEachSelect()) { globalState.getLogger().writeCurrent(selectStr); try { globalState.getLogger().getCurrentFileWriter().flush(); } catch (IOException e) { e.printStackTrace(); } } SQLQueryAdapter q = new SQLQueryAdapter(PostgresExplainGenerator.explain(selectStr), null); try (SQLancerResultSet rs = q.executeAndGet(globalState)) { while (rs.next()) { queryPlan += rs.getString(1); } } catch (SQLException | AssertionError e) { queryPlan = ""; } return formatQueryPlan(queryPlan); } @Override protected double[] initializeWeightedAverageReward() { return new double[PostgresProvider.Action.values().length]; } @Override protected void executeMutator(int index, PostgresGlobalState globalState) throws Exception { SQLQueryAdapter queryMutateTable = PostgresProvider.Action.values()[index].getQuery(globalState); globalState.executeStatement(queryMutateTable); } @Override protected boolean addRowsToAllTables(PostgresGlobalState globalState) throws Exception { List tablesNoRow = globalState.getSchema().getDatabaseTables().stream() .filter(t -> t.getNrRows(globalState) == 0).collect(Collectors.toList()); for (PostgresSchema.PostgresTable table : tablesNoRow) { SQLQueryAdapter queryAddRows = PostgresInsertGenerator.insertRows(globalState, table); globalState.executeStatement(queryAddRows); } return true; } public String formatQueryPlan(String queryPlan) throws IOException { ObjectMapper mapper = new ObjectMapper(); JsonNode root = mapper.readTree(queryPlan).get(0).get("Plan"); // Extract nodes using BFS algorithm List nodeTypes = extractNodeTypesIterative(root); return String.join(" ", nodeTypes); } // BFS algorithm for traversing the Json Query Plan private static List extractNodeTypesIterative(JsonNode root) { List result = new ArrayList<>(); Queue queue = new LinkedList<>(); queue.add(root); while (!queue.isEmpty()) { JsonNode node = queue.poll(); if (node.has("Node Type")) { result.add(node.get("Node Type").asText()); } if (node.has("Plans") && node.get("Plans").isArray()) { for (JsonNode plan : node.get("Plans")) { queue.add(plan); } } } return result; } }