diff --git a/src/main/java/net/sf/jsqlparser/statement/PrepareStatement.java b/src/main/java/net/sf/jsqlparser/statement/PrepareStatement.java index e4675b340..6d682deb4 100644 --- a/src/main/java/net/sf/jsqlparser/statement/PrepareStatement.java +++ b/src/main/java/net/sf/jsqlparser/statement/PrepareStatement.java @@ -9,8 +9,13 @@ */ package net.sf.jsqlparser.statement; +import java.util.List; +import java.util.function.Consumer; +import net.sf.jsqlparser.statement.create.table.ColDataType; +import net.sf.jsqlparser.statement.select.PlainSelect; + /** - * {@code PREPARE name AS statement}, which stores a parameterised statement for later + * {@code PREPARE name [(types)] AS statement}, which stores a parameterised statement for later * {@code EXECUTE}. * * @see Prepared @@ -19,6 +24,7 @@ public class PrepareStatement implements Statement { private String name; private Statement statement; + private List parameterTypes; public PrepareStatement() {} @@ -53,8 +59,32 @@ public PrepareStatement withStatement(Statement statement) { return this; } + /** Returns declared parameter types, or {@code null} when types are inferred. */ + public List getParameterTypes() { + return parameterTypes; + } + + public void setParameterTypes(List parameterTypes) { + this.parameterTypes = parameterTypes; + } + + public PrepareStatement withParameterTypes(List parameterTypes) { + setParameterTypes(parameterTypes); + return this; + } + public StringBuilder appendTo(StringBuilder builder) { - builder.append("PREPARE ").append(name).append(" AS ").append(statement); + return appendTo(builder, builder::append); + } + + /** Renders the nested statement through the caller's statement writer. */ + public StringBuilder appendTo(StringBuilder builder, Consumer statementPrinter) { + builder.append("PREPARE ").append(name); + if (parameterTypes != null && !parameterTypes.isEmpty()) { + builder.append(PlainSelect.getStringList(parameterTypes, true, true)); + } + builder.append(" AS "); + statementPrinter.accept(statement); return builder; } diff --git a/src/main/java/net/sf/jsqlparser/util/deparser/StatementDeParser.java b/src/main/java/net/sf/jsqlparser/util/deparser/StatementDeParser.java index ffd5fe900..c3f078566 100644 --- a/src/main/java/net/sf/jsqlparser/util/deparser/StatementDeParser.java +++ b/src/main/java/net/sf/jsqlparser/util/deparser/StatementDeParser.java @@ -612,7 +612,7 @@ public StringBuilder visit(DisconnectStatement disconnectStatement, S contex @Override public StringBuilder visit(PrepareStatement prepareStatement, S context) { - prepareStatement.appendTo(builder); + prepareStatement.appendTo(builder, statement -> statement.accept(this, context)); return builder; } diff --git a/src/main/jjtree/net/sf/jsqlparser/parser/JSqlParserCC.jjt b/src/main/jjtree/net/sf/jsqlparser/parser/JSqlParserCC.jjt index 5e6d9c46e..136ade609 100644 --- a/src/main/jjtree/net/sf/jsqlparser/parser/JSqlParserCC.jjt +++ b/src/main/jjtree/net/sf/jsqlparser/parser/JSqlParserCC.jjt @@ -4701,14 +4701,28 @@ PrepareStatement PrepareStatement() #PrepareStatement: PrepareStatement prepareStatement = new PrepareStatement(); ObjectNames name; Statement statement; + List parameterTypes = new ArrayList(); + ColDataType parameterType; } { name=RelObjectNames() { prepareStatement.setName(String.join(".", name.getNames())); } + [ "(" parameterType=PrepareParameterType() { parameterTypes.add(parameterType); } + ( "," parameterType=PrepareParameterType() { parameterTypes.add(parameterType); } )* + ")" { prepareStatement.setParameterTypes(parameterTypes); } ] statement=SingleStatement() { prepareStatement.setStatement(statement); } { return prepareStatement; } } +/** PostgreSQL UNKNOWN leaves a parameter type to be inferred from the prepared statement. */ +ColDataType PrepareParameterType(): +{ ColDataType type; Token token; } +{ + ( token= { type = new ColDataType(token.image); } + | type=ColDataType() ) + { return type; } +} + /** * DEALLOCATE [PREPARE] name. */ diff --git a/src/test/java/net/sf/jsqlparser/statement/PostgreSqlPrepareTest.java b/src/test/java/net/sf/jsqlparser/statement/PostgreSqlPrepareTest.java new file mode 100644 index 000000000..be9642db2 --- /dev/null +++ b/src/test/java/net/sf/jsqlparser/statement/PostgreSqlPrepareTest.java @@ -0,0 +1,95 @@ +/*- + * #%L + * JSQLParser library + * %% + * Copyright (C) 2004 - 2026 JSQLParser + * %% + * Dual licensed under GNU LGPL 2.1 or Apache License 2.0 + * #L% + */ +package net.sf.jsqlparser.statement; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertNull; +import static org.junit.jupiter.api.Assertions.assertThrows; + +import net.sf.jsqlparser.JSQLParserException; +import net.sf.jsqlparser.expression.LongValue; +import net.sf.jsqlparser.parser.AbstractJSqlParser.Dialect; +import net.sf.jsqlparser.parser.CCJSqlParserUtil; +import net.sf.jsqlparser.util.deparser.ExpressionDeParser; +import net.sf.jsqlparser.util.deparser.SelectDeParser; +import net.sf.jsqlparser.util.deparser.StatementDeParser; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.ValueSource; + +class PostgreSqlPrepareTest { + @ParameterizedTest + @ValueSource(strings = { + "PREPARE p(bigint) AS SELECT * FROM users WHERE id = $1", + "PREPARE p(bigint[], text) AS SELECT * FROM users WHERE id = ANY($1) AND name = $2", + "PREPARE p(pg_catalog.int4, numeric(10, 2)) AS SELECT $1 + $2", + "PREPARE \"prepared query\"(timestamp with time zone, double precision) AS SELECT $1, $2", + "PREPARE p(unknown) AS SELECT $1::text", + "PREPARE p(bigint, text) AS INSERT INTO users(id, name) VALUES ($1, $2)", + "PREPARE p(text, bigint) AS UPDATE users SET name = $1 WHERE id = $2", + "PREPARE p(bigint) AS DELETE FROM users WHERE id = $1", + "PREPARE p AS SELECT * FROM users WHERE id = $1" + }) + void preservesDeclaredTypesAndNestedStatements(String sql) throws JSQLParserException { + PrepareStatement prepare = parse(sql); + assertRoundTrip(prepare); + assertEquals(prepare.toString(), CCJSqlParserUtil.parse(prepare.toString()).toString()); + } + + @Test + void exposesMutableTypesAndRetainsInferredForm() throws JSQLParserException { + PrepareStatement prepare = parse("PREPARE p(bigint, text) AS SELECT $1, $2"); + assertEquals(2, prepare.getParameterTypes().size()); + assertEquals("bigint", prepare.getParameterTypes().get(0).getBaseTypeName()); + prepare.getParameterTypes().get(0).setDataType("integer"); + assertEquals("PREPARE p(integer, text) AS SELECT $1, $2", prepare.toString()); + assertRoundTrip(prepare); + prepare.setParameterTypes(null); + assertEquals("PREPARE p AS SELECT $1, $2", prepare.toString()); + assertNull(parse(prepare.toString()).getParameterTypes()); + assertRoundTrip(new PrepareStatement("p", prepare.getStatement())); + } + + @Test + void passesNestedExpressionsToTheConfiguredDeparser() throws JSQLParserException { + PrepareStatement prepare = + parse("PREPARE p(bigint) AS SELECT $1 + 7 FROM users WHERE id > 8"); + StringBuilder buffer = new StringBuilder(); + ExpressionDeParser expressions = new ExpressionDeParser() { + @Override + public StringBuilder visit(LongValue value, S context) { + return getBuilder().append(value.getValue() + 100); + } + }; + prepare.accept(new StatementDeParser(expressions, new SelectDeParser(), buffer), null); + assertEquals("PREPARE p(bigint) AS SELECT $1 + 107 FROM users WHERE id > 108", + buffer.toString()); + assertRoundTrip(parse(buffer.toString())); + } + + @ParameterizedTest + @ValueSource(strings = {"PREPARE p() AS SELECT 1", "PREPARE p(bigint,) AS SELECT $1", + "PREPARE p(bigint text) AS SELECT $1", "PREPARE p(bigint) SELECT $1"}) + void rejectsMalformedTypeLists(String sql) { + assertThrows(JSQLParserException.class, () -> parse(sql)); + } + + private static PrepareStatement parse(String sql) throws JSQLParserException { + return (PrepareStatement) CCJSqlParserUtil.parse(sql, + p -> p.withDialect(Dialect.POSTGRESQL)); + } + + private static void assertRoundTrip(PrepareStatement prepare) throws JSQLParserException { + StringBuilder buffer = new StringBuilder(); + prepare.accept(new StatementDeParser(buffer), null); + assertEquals(prepare.toString(), buffer.toString()); + assertEquals(prepare.toString(), parse(buffer.toString()).toString()); + } +}