diff --git a/core/src/main/kotlin/com/alecstrong/sql/psi/core/psi/mixins/ColumnNameMixin.kt b/core/src/main/kotlin/com/alecstrong/sql/psi/core/psi/mixins/ColumnNameMixin.kt index d771da969..0a34f8f22 100644 --- a/core/src/main/kotlin/com/alecstrong/sql/psi/core/psi/mixins/ColumnNameMixin.kt +++ b/core/src/main/kotlin/com/alecstrong/sql/psi/core/psi/mixins/ColumnNameMixin.kt @@ -3,6 +3,8 @@ package com.alecstrong.sql.psi.core.psi.mixins import com.alecstrong.sql.psi.core.AnnotationException import com.alecstrong.sql.psi.core.SqlAnnotationHolder import com.alecstrong.sql.psi.core.SqlParser +import com.alecstrong.sql.psi.core.psi.SqlColumnDef +import com.alecstrong.sql.psi.core.psi.SqlColumnName import com.alecstrong.sql.psi.core.psi.SqlColumnReference import com.alecstrong.sql.psi.core.psi.SqlNamedElementImpl import com.intellij.icons.AllIcons @@ -32,3 +34,7 @@ internal abstract class ColumnNameMixin( return AllIcons.Nodes.DataColumn } } + +fun SqlColumnName.getColumnDefOrNull(): SqlColumnDef? { + return reference?.resolve()?.parent as? SqlColumnDef +} diff --git a/core/src/main/kotlin/com/alecstrong/sql/psi/core/psi/mixins/CreateTableMixin.kt b/core/src/main/kotlin/com/alecstrong/sql/psi/core/psi/mixins/CreateTableMixin.kt index 4d4c8cc2b..3a6dea5c7 100644 --- a/core/src/main/kotlin/com/alecstrong/sql/psi/core/psi/mixins/CreateTableMixin.kt +++ b/core/src/main/kotlin/com/alecstrong/sql/psi/core/psi/mixins/CreateTableMixin.kt @@ -207,10 +207,11 @@ internal abstract class CreateTableMixin private constructor( } } - tableConstraintList.filter { it.foreignKeyClause != null } + tableConstraintList.filter { it.foreignTableClause != null } .forEach { constraint -> - constraint.foreignKeyClause!!.checkCompositeForeignKey( - constraint.columnNameList, + val foreignTableClause = constraint.foreignTableClause!! + foreignTableClause.foreignKeyClause.checkCompositeForeignKey( + foreignTableClause.columnNameList, ) } } diff --git a/core/src/main/kotlin/com/alecstrong/sql/psi/core/psi/mixins/ForeignKeyClauseMixin.kt b/core/src/main/kotlin/com/alecstrong/sql/psi/core/psi/mixins/ForeignKeyClauseMixin.kt index 351a79986..63894a97e 100644 --- a/core/src/main/kotlin/com/alecstrong/sql/psi/core/psi/mixins/ForeignKeyClauseMixin.kt +++ b/core/src/main/kotlin/com/alecstrong/sql/psi/core/psi/mixins/ForeignKeyClauseMixin.kt @@ -1,10 +1,13 @@ package com.alecstrong.sql.psi.core.psi.mixins import com.alecstrong.sql.psi.core.psi.QueryElement.QueryResult +import com.alecstrong.sql.psi.core.psi.SqlColumnDef import com.alecstrong.sql.psi.core.psi.SqlCompositeElementImpl +import com.alecstrong.sql.psi.core.psi.SqlCreateTableStmt import com.alecstrong.sql.psi.core.psi.SqlForeignKeyClause import com.intellij.lang.ASTNode import com.intellij.psi.PsiElement +import com.intellij.psi.util.parentOfType internal abstract class ForeignKeyClauseMixin( node: ASTNode, @@ -18,3 +21,26 @@ internal abstract class ForeignKeyClauseMixin( return super.queryAvailable(child) } } + +fun SqlColumnDef.isForeignKey(): Boolean { + for (columnConstraint in columnConstraintList) { + if (columnConstraint.foreignKeyClause != null) { + return true + } + } + val createTableStmt: SqlCreateTableStmt? = parentOfType() + if (createTableStmt != null) { + for (tableConstraints in createTableStmt.tableConstraintList) { + val foreignTableClause = tableConstraints.foreignTableClause + if (foreignTableClause != null) { + val columns = foreignTableClause.columnNameList + for (column in columns) { + if (column.reference?.resolve() == columnName) { + return true + } + } + } + } + } + return false +} diff --git a/core/src/main/kotlin/com/alecstrong/sql/psi/core/sql.bnf b/core/src/main/kotlin/com/alecstrong/sql/psi/core/sql.bnf index 10035cc21..039206105 100644 --- a/core/src/main/kotlin/com/alecstrong/sql/psi/core/sql.bnf +++ b/core/src/main/kotlin/com/alecstrong/sql/psi/core/sql.bnf @@ -126,7 +126,8 @@ generated_clause ::= [ GENERATED ] ALWAYS AS LP expr RP check_constraint ::= CHECK LP expr RP default_constraint ::= DEFAULT ( signed_number | literal_value | LP expr RP ) signed_number ::= [ PLUS | MINUS ] numeric_literal -table_constraint ::= [ CONSTRAINT identifier ] ( ( PRIMARY KEY | UNIQUE ) LP indexed_column ( COMMA indexed_column ) * RP conflict_clause | CHECK LP expr RP | FOREIGN KEY LP column_name ( COMMA column_name ) * RP foreign_key_clause ) +table_constraint ::= [ CONSTRAINT identifier ] ( ( PRIMARY KEY | UNIQUE ) LP indexed_column ( COMMA indexed_column ) * RP conflict_clause | CHECK LP expr RP | foreign_table_clause ) +foreign_table_clause ::= FOREIGN KEY LP column_name ( COMMA column_name ) * RP foreign_key_clause foreign_key_clause ::= REFERENCES foreign_table [ LP column_name ( COMMA column_name ) * RP ] [ ( ON ( DELETE | UPDATE ) ( SET NULL | SET DEFAULT | CASCADE | RESTRICT | NO ACTION ) | MATCH identifier ) * ] [ [ NOT ] DEFERRABLE [ INITIALLY DEFERRED | INITIALLY IMMEDIATE ] ] { mixin = "com.alecstrong.sql.psi.core.psi.mixins.ForeignKeyClauseMixin" } diff --git a/core/src/test/kotlin/com/alecstrong/sql/psi/core/GetForeignKeyHelperTest.kt b/core/src/test/kotlin/com/alecstrong/sql/psi/core/GetForeignKeyHelperTest.kt new file mode 100644 index 000000000..49d691331 --- /dev/null +++ b/core/src/test/kotlin/com/alecstrong/sql/psi/core/GetForeignKeyHelperTest.kt @@ -0,0 +1,103 @@ +package com.alecstrong.sql.psi.core + +import com.alecstrong.sql.psi.core.psi.SqlColumnExpr +import com.alecstrong.sql.psi.core.psi.mixins.getColumnDefOrNull +import com.alecstrong.sql.psi.core.psi.mixins.isForeignKey +import com.alecstrong.sql.psi.test.fixtures.compileFile +import com.google.common.truth.Truth.assertThat +import org.junit.After +import org.junit.Before +import org.junit.Test +import java.io.File + +class GetForeignKeyHelperTest { + @Before + fun before() { + File("build/tmp").deleteRecursively() + } + + @After + fun after() { + SqlParserUtil.reset() + File("build/tmp").deleteRecursively() + } + + @Test + fun columnConstraint() { + val sqlFile = compileFile( + """ + |CREATE TABLE foo ( + | id INT PRIMARY KEY + |); + | + |CREATE TABLE bar ( + |a TEXT REFERENCES foo(id) + |); + | + |SELECT a FROM bar; + """.trimMargin(), + ) + val select = sqlFile.sqlStmtList!!.stmtList.last() + val a = (select.compoundSelectStmt!!.selectStmtList.single().resultColumnList.single().expr as SqlColumnExpr).columnName + + val columnDef = a.getColumnDefOrNull() + assertThat(columnDef).isNotNull() + assertThat(columnDef!!.isForeignKey()).isTrue() + } + + @Test + fun tableConstraint() { + val sqlFile = compileFile( + """ + |CREATE TABLE foo ( + | id INT PRIMARY KEY + |); + | + |CREATE TABLE bar ( + |a INT, + |FOREIGN KEY (a) REFERENCES foo(id) + |); + | + |SELECT a FROM bar; + """.trimMargin(), + ) + val select = sqlFile.sqlStmtList!!.stmtList.last() + val a = (select.compoundSelectStmt!!.selectStmtList.single().resultColumnList.single().expr as SqlColumnExpr).columnName + + val columnDef = a.getColumnDefOrNull() + assertThat(columnDef).isNotNull() + assertThat(columnDef!!.isForeignKey()).isTrue() + } + + @Test + fun alterTable() { + compileFile( + """ + |CREATE TABLE foo ( + | id INT PRIMARY KEY + |); + | + |CREATE TABLE bar ( + |b TEXT + |); + """.trimMargin(), + fileName = "1.s", + ) + val alterFile = compileFile( + """ + |ALTER TABLE bar + | ADD COLUMN a INT REFERENCES foo(id) + |; + | + |SELECT a FROM bar; + """.trimMargin(), + fileName = "2.s", + ) + val select = alterFile.sqlStmtList!!.stmtList.last() + val a = (select.compoundSelectStmt!!.selectStmtList.single().resultColumnList.single().expr as SqlColumnExpr).columnName + + val columnDef = a.getColumnDefOrNull() + assertThat(columnDef).isNotNull() + assertThat(columnDef!!.isForeignKey()).isTrue() + } +}