diff --git a/src/sqlancer/postgres/PostgresSchema.java b/src/sqlancer/postgres/PostgresSchema.java index c99c8648e..ac908a007 100644 --- a/src/sqlancer/postgres/PostgresSchema.java +++ b/src/sqlancer/postgres/PostgresSchema.java @@ -31,7 +31,7 @@ public class PostgresSchema extends AbstractSchema dataTypes = new ArrayList<>(Arrays.asList(values())); @@ -43,9 +43,46 @@ public static PostgresDataType getRandomType() { dataTypes.remove(PostgresDataType.RANGE); dataTypes.remove(PostgresDataType.MONEY); dataTypes.remove(PostgresDataType.BIT); + dataTypes.remove(PostgresDataType.TIMESTAMP); + dataTypes.remove(PostgresDataType.DATE); + dataTypes.remove(PostgresDataType.TIME); } return Randomly.fromList(dataTypes); } + + @Override + public String toString() { + switch (this) { + case INT: + return "INTEGER"; + case BOOLEAN: + return "BOOLEAN"; + case TEXT: + return "TEXT"; + case DECIMAL: + return "DECIMAL"; + case FLOAT: + return "FLOAT"; + case REAL: + return "REAL"; + case RANGE: + return "INT4RANGE"; + case MONEY: + return "MONEY"; + case BIT: + return "BIT"; + case INET: + return "INET"; + case TIMESTAMP: + return "TIMESTAMP"; + case DATE: + return "DATE"; + case TIME: + return "TIME WITH TIME ZONE"; + default: + throw new AssertionError(this); + } + } } public static class PostgresColumn extends AbstractTableColumn { @@ -141,6 +178,17 @@ public static PostgresDataType getColumnType(String typeString) { return PostgresDataType.BIT; case "inet": return PostgresDataType.INET; + + case "timestamp": + case "timestamp with time zone": + case "timestamp without time zone": + return PostgresDataType.TIMESTAMP; + case "date": + return PostgresDataType.DATE; + case "time": + case "time with time zone": + case "time without time zone": + return PostgresDataType.TIME; default: throw new AssertionError(typeString); } diff --git a/src/sqlancer/postgres/PostgresToStringVisitor.java b/src/sqlancer/postgres/PostgresToStringVisitor.java index 87bd3c429..1b4c499bd 100644 --- a/src/sqlancer/postgres/PostgresToStringVisitor.java +++ b/src/sqlancer/postgres/PostgresToStringVisitor.java @@ -197,15 +197,28 @@ public void visit(PostgresOrderByTerm term) { @Override public void visit(PostgresFunction f) { - sb.append(f.getFunctionName()); - sb.append("("); - int i = 0; - for (PostgresExpression arg : f.getArguments()) { - if (i++ != 0) { - sb.append(", "); + if (f.isExtractFunction()) { + visitExtractFunction(f); + } else { + sb.append(f.getFunctionName()); + sb.append("("); + int i = 0; + for (PostgresExpression arg : f.getArguments()) { + if (i++ != 0) { + sb.append(", "); + } + visit(arg); } - visit(arg); + sb.append(")"); } + } + + private void visitExtractFunction(PostgresFunction f) { + sb.append(f.getFunctionName()); + sb.append("("); + visit(f.getArguments()[0]); + sb.append(" FROM "); + visit(f.getArguments()[1]); sb.append(")"); } @@ -258,11 +271,15 @@ private void appendType(PostgresCastOperation cast) { break; case BIT: sb.append("BIT"); - // if (Randomly.getBoolean()) { - // sb.append("("); - // sb.append(Randomly.getNotCachedInteger(1, 100)); - // sb.append(")"); - // } + break; + case TIMESTAMP: + sb.append("TIMESTAMP"); + break; + case DATE: + sb.append("DATE"); + break; + case TIME: + sb.append("TIME"); break; default: throw new AssertionError(cast.getType()); diff --git a/src/sqlancer/postgres/ast/PostgresFunction.java b/src/sqlancer/postgres/ast/PostgresFunction.java index 5fe2968ab..9d638d699 100644 --- a/src/sqlancer/postgres/ast/PostgresFunction.java +++ b/src/sqlancer/postgres/ast/PostgresFunction.java @@ -31,6 +31,82 @@ public PostgresExpression[] getArguments() { return args.clone(); } + public boolean isExtractFunction() { + return false; + } + + public String getArgString() { + StringBuilder sb = new StringBuilder(); + for (int i = 0; i < args.length; i++) { + if (i != 0) { + sb.append(", "); + } + sb.append(formatArgumentForPostgresFunction(args[i])); + } + return sb.toString(); + } + + /** + * Formats an argument for use in PostgreSQL function calls. Date or time values need explicit CAST statements + * because they are stored as text but PostgreSQL requires proper type annotations for function parameters. + * + * @param arg + * the expression to format + * + * @return the formatted argument string + */ + private String formatArgumentForPostgresFunction(PostgresExpression arg) { + if (arg.getExpressionType() == PostgresDataType.TIME || arg.getExpressionType() == PostgresDataType.TIMESTAMP + || arg.getExpressionType() == PostgresDataType.DATE) { + return String.format("CAST(%s AS %s)", arg, arg.getExpressionType().toString()); + } else { + return arg.toString(); + } + } + + public static class PostgresExtractFunction extends PostgresFunction { + + public PostgresExtractFunction(PostgresFunctionWithUnknownResult f, PostgresDataType returnType, + PostgresExpression... args) { + super(f, returnType, args); + } + + @Override + public String getFunctionName() { + return "EXTRACT"; + } + + @Override + public boolean isExtractFunction() { + return true; + } + + @Override + public String getArgString() { + return String.format("%s FROM %s", getArguments()[0], getArguments()[1]); + } + } + + @Override + public PostgresConstant getExpectedValue() { + if (functionWithKnownResult == null) { + return null; + } + PostgresConstant[] constants = new PostgresConstant[args.length]; + for (int i = 0; i < constants.length; i++) { + constants[i] = args[i].getExpectedValue(); + if (constants[i] == null) { + return null; + } + } + return functionWithKnownResult.apply(constants, args); + } + + @Override + public PostgresDataType getExpressionType() { + return returnType; + } + public enum PostgresFunctionWithResult { ABS(1, "abs") { @@ -76,7 +152,6 @@ public boolean supportsReturnType(PostgresDataType type) { public PostgresDataType[] getInputTypesForReturnType(PostgresDataType returnType, int nrArguments) { return new PostgresDataType[] { PostgresDataType.TEXT }; } - }, LENGTH(1, "length") { @Override @@ -265,24 +340,4 @@ public boolean checkArguments(PostgresExpression... constants) { } - @Override - public PostgresConstant getExpectedValue() { - if (functionWithKnownResult == null) { - return null; - } - PostgresConstant[] constants = new PostgresConstant[args.length]; - for (int i = 0; i < constants.length; i++) { - constants[i] = args[i].getExpectedValue(); - if (constants[i] == null) { - return null; - } - } - return functionWithKnownResult.apply(constants, args); - } - - @Override - public PostgresDataType getExpressionType() { - return returnType; - } - } diff --git a/src/sqlancer/postgres/ast/PostgresFunctionWithUnknownResult.java b/src/sqlancer/postgres/ast/PostgresFunctionWithUnknownResult.java index 287f46784..79465cbc5 100644 --- a/src/sqlancer/postgres/ast/PostgresFunctionWithUnknownResult.java +++ b/src/sqlancer/postgres/ast/PostgresFunctionWithUnknownResult.java @@ -141,7 +141,33 @@ public PostgresExpression[] getArguments(PostgresDataType returnType, PostgresEx RANGE_MERGE("range_merge", PostgresDataType.RANGE, PostgresDataType.RANGE, PostgresDataType.RANGE), // // https://www.postgresql.org/docs/13/functions-admin.html#FUNCTIONS-ADMIN-DBSIZE - GET_COLUMN_SIZE("get_column_size", PostgresDataType.INT, PostgresDataType.TEXT); + GET_COLUMN_SIZE("get_column_size", PostgresDataType.INT, PostgresDataType.TEXT), + + // Extract function implementation + EXTRACT("extract", PostgresDataType.FLOAT, PostgresDataType.TEXT, PostgresDataType.TIMESTAMP, PostgresDataType.DATE, + PostgresDataType.TIME) { + @Override + public PostgresExpression[] getArguments(PostgresDataType returnType, PostgresExpressionGenerator gen, + int depth) { + PostgresExpression[] args = new PostgresExpression[2]; + PostgresDataType sourceType = Randomly.fromOptions(PostgresDataType.TIMESTAMP, PostgresDataType.DATE, + PostgresDataType.TIME); + args[1] = gen.generateExpression(depth + 1, sourceType); + + String[] validFields; + if (sourceType == PostgresDataType.DATE) { + validFields = new String[] { "YEAR", "MONTH", "DAY", "DECADE", "CENTURY", "MILLENNIUM", "QUARTER", + "WEEK", "DOY", "DOW", "ISODOW", "ISOYEAR" }; + } else if (sourceType == PostgresDataType.TIME) { + validFields = new String[] { "HOUR", "MINUTE", "SECOND" }; + } else { + validFields = new String[] { "YEAR", "MONTH", "DAY", "HOUR", "MINUTE", "SECOND", "DECADE", "CENTURY", + "MILLENNIUM", "QUARTER", "WEEK", "DOY", "DOW", "ISODOW", "ISOYEAR", "EPOCH" }; + } + args[0] = PostgresConstant.createTextConstant(Randomly.fromOptions(validFields)); + return args; + } + }; // PG_DATABASE_SIZE("pg_database_size", PostgresDataType.INT, PostgresDataType.INT); // PG_SIZE_BYTES("pg_size_bytes", PostgresDataType.INT, PostgresDataType.TEXT); diff --git a/src/sqlancer/postgres/gen/PostgresAlterTableGenerator.java b/src/sqlancer/postgres/gen/PostgresAlterTableGenerator.java index c8842f0a4..6e0c436ee 100644 --- a/src/sqlancer/postgres/gen/PostgresAlterTableGenerator.java +++ b/src/sqlancer/postgres/gen/PostgresAlterTableGenerator.java @@ -53,9 +53,7 @@ protected enum Action { ALTER_VIEW_RENAME_COLUMN // RENAME COLUMN old_name TO new_name (for views) } - private static final List VIEW_ACTIONS = List.of( - Action.ALTER_VIEW_RENAME_COLUMN - ); + private static final List VIEW_ACTIONS = List.of(Action.ALTER_VIEW_RENAME_COLUMN); public PostgresAlterTableGenerator(PostgresTable randomTable, PostgresGlobalState globalState, boolean generateOnlyKnown) { diff --git a/src/sqlancer/postgres/gen/PostgresCommon.java b/src/sqlancer/postgres/gen/PostgresCommon.java index 63bc885cb..81f03e661 100644 --- a/src/sqlancer/postgres/gen/PostgresCommon.java +++ b/src/sqlancer/postgres/gen/PostgresCommon.java @@ -287,6 +287,21 @@ public static boolean appendDataType(PostgresDataType type, StringBuilder sb, bo case INET: sb.append("inet"); break; + case DATE: + sb.append("DATE"); + break; + case TIME: + sb.append("TIME"); + if (!generateOnlyKnown && Randomly.getBoolean()) { + sb.append(" WITH TIME ZONE"); + } + break; + case TIMESTAMP: + sb.append("TIMESTAMP"); + if (!generateOnlyKnown && Randomly.getBoolean()) { + sb.append(" WITH TIME ZONE"); + } + break; default: throw new AssertionError(type); } diff --git a/src/sqlancer/postgres/gen/PostgresExpressionGenerator.java b/src/sqlancer/postgres/gen/PostgresExpressionGenerator.java index bad87affc..a24d73940 100644 --- a/src/sqlancer/postgres/gen/PostgresExpressionGenerator.java +++ b/src/sqlancer/postgres/gen/PostgresExpressionGenerator.java @@ -142,7 +142,14 @@ private PostgresExpression generateFunctionWithUnknownResult(int depth, Postgres throw new IgnoreMeException(); } PostgresFunctionWithUnknownResult randomFunction = Randomly.fromList(supportedFunctions); - return new PostgresFunction(randomFunction, type, randomFunction.getArguments(type, this, depth + 1)); + PostgresExpression[] args = randomFunction.getArguments(type, this, depth + 1); + + // Create the appropriate function type based on the function name + if (randomFunction == PostgresFunctionWithUnknownResult.EXTRACT) { + return new PostgresFunction.PostgresExtractFunction(randomFunction, type, args); + } else { + return new PostgresFunction(randomFunction, type, args); + } } private PostgresExpression generateFunctionWithKnownResult(int depth, PostgresDataType type) { @@ -336,6 +343,9 @@ private PostgresExpression generateExpressionInternal(int depth, PostgresDataTyp case FLOAT: case MONEY: case INET: + case DATE: + case TIMESTAMP: + case TIME: return generateConstant(r, dataType); case BIT: return generateBitExpression(depth); @@ -357,6 +367,9 @@ private static PostgresCompoundDataType getCompoundDataType(PostgresDataType typ case RANGE: case REAL: case INET: + case TIME: + case DATE: + case TIMESTAMP: return PostgresCompoundDataType.create(type); case TEXT: // TODO case BIT: @@ -371,7 +384,6 @@ private static PostgresCompoundDataType getCompoundDataType(PostgresDataType typ default: throw new AssertionError(type); } - } private enum RangeExpression { @@ -484,8 +496,17 @@ private String selectWindowFunctionName() { } private PostgresExpression generateConcat(int depth) { - PostgresExpression left = generateExpression(depth + 1, PostgresDataType.TEXT); - PostgresExpression right = generateExpression(depth + 1); + // Don't allow temporal types in concatenation to avoid invalid time zone strings + List allowedTypes = Arrays.asList(PostgresDataType.TEXT, PostgresDataType.INT, + PostgresDataType.BOOLEAN); + PostgresDataType leftType = PostgresDataType.TEXT; + PostgresDataType rightType = PostgresDataType.TEXT; + if (!PostgresProvider.generateOnlyKnown) { + leftType = Randomly.fromList(allowedTypes); + rightType = Randomly.fromList(allowedTypes); + } + PostgresExpression left = generateExpression(depth + 1, leftType); + PostgresExpression right = generateExpression(depth + 1, rightType); return new PostgresConcatOperation(left, right); } @@ -596,6 +617,31 @@ public static PostgresExpression generateConstant(Randomly r, PostgresDataType t return PostgresConstant.createInetConstant(getRandomInet(r)); case BIT: return PostgresConstant.createBitConstant(r.getInteger()); + case DATE: + return PostgresConstant.createTextConstant(String.format("%d-%02d-%02d", r.getInteger(1000, 2100), // year + r.getInteger(1, 12), // month + r.getInteger(1, 28) // day - using 28 to avoid invalid dates + )); + case TIMESTAMP: + return PostgresConstant + .createTextConstant(String.format("%d-%02d-%02d %02d:%02d:%02d", r.getInteger(1000, 2100), // year + r.getInteger(1, 12), // month + r.getInteger(1, 28), // day + r.getInteger(0, 23), // hour + r.getInteger(0, 59), // minute + r.getInteger(0, 59) // second + )); + case TIME: + // Generate valid timezone offset in format +/-HH:MM + String sign = Randomly.fromOptions("+", "-"); + int hours = r.getInteger(0, 12); // Valid timezone hours: 0-12 + int minutes = r.getInteger(0, 59); + String tz = String.format("%s%02d:%02d", sign, hours, minutes); + return PostgresConstant.createTextConstant(String.format("%02d:%02d:%02d%s", r.getInteger(0, 23), // hour + r.getInteger(0, 59), // minute + r.getInteger(0, 59), // second + tz // time zone + )); default: throw new AssertionError(type); }