From 644209120c268597ab9a70857d69c9f317ce5378 Mon Sep 17 00:00:00 2001 From: Mecharlance Wei Date: Mon, 30 Mar 2026 23:36:08 +0800 Subject: [PATCH] feat: support scoped DSL conditions in join on --- .../plans/2026-03-30-join-on-scoped-dsl.md | 68 ++++++++++ .../2026-03-30-join-on-scoped-dsl-design.md | 97 +++++++++++++++ readme.md | 70 +++++++++++ .../cn/beagile/dslquery/ColumnFields.java | 25 +++- .../java/cn/beagile/dslquery/DSLQuery.java | 15 +++ .../cn/beagile/dslquery/DSLSQLBuilder.java | 32 +++-- .../java/cn/beagile/dslquery/JoinField.java | 44 +++++-- src/main/java/cn/beagile/dslquery/JoinOn.java | 12 ++ .../cn/beagile/dslquery/JoinOnSQLBuilder.java | 96 +++++++++++++++ .../java/cn/beagile/dslquery/SQLBuilder.java | 8 ++ .../cn/beagile/dslquery/SingleExpression.java | 19 ++- .../cn/beagile/dslquery/DSLQueryTest.java | 10 ++ .../cn/beagile/dslquery/JoinOnApiTest.java | 21 ++++ .../cn/beagile/dslquery/JoinOnQueryTest.java | 116 ++++++++++++++++++ 14 files changed, 609 insertions(+), 24 deletions(-) create mode 100644 docs/superpowers/plans/2026-03-30-join-on-scoped-dsl.md create mode 100644 docs/superpowers/specs/2026-03-30-join-on-scoped-dsl-design.md create mode 100644 src/main/java/cn/beagile/dslquery/JoinOn.java create mode 100644 src/main/java/cn/beagile/dslquery/JoinOnSQLBuilder.java create mode 100644 src/test/java/cn/beagile/dslquery/JoinOnApiTest.java create mode 100644 src/test/java/cn/beagile/dslquery/JoinOnQueryTest.java diff --git a/docs/superpowers/plans/2026-03-30-join-on-scoped-dsl.md b/docs/superpowers/plans/2026-03-30-join-on-scoped-dsl.md new file mode 100644 index 0000000..3a0737c --- /dev/null +++ b/docs/superpowers/plans/2026-03-30-join-on-scoped-dsl.md @@ -0,0 +1,68 @@ +# Join On Scoped DSL Implementation Plan + +> **For agentic workers:** REQUIRED: Use superpowers:subagent-driven-development (if subagents available) or superpowers:executing-plans to implement this plan. Steps use checkbox (`- [ ]`) syntax for tracking. + +**Goal:** Add optional scoped DSL conditions for `join on` while preserving existing annotation-driven join behavior. + +**Architecture:** Introduce a small join-on annotation and a runtime query API that both flow into `JoinField`. Render extra `on` predicates with a dedicated scoped SQL builder so existing `where` parsing and SQL generation remain stable. + +**Tech Stack:** Java, JUnit 5, Gradle + +--- + +## Chunk 1: API and failing coverage + +### Task 1: Add design-facing tests for annotation and runtime join-on + +**Files:** +- Modify: `src/test/java/cn/beagile/dslquery/DeepJoinTest.java` +- Modify: `src/test/java/cn/beagile/dslquery/DSLQueryTest.java` + +- [ ] **Step 1: Write the failing tests** +- [ ] **Step 2: Run the focused test command and confirm the new tests fail for the expected reason** +- [ ] **Step 3: Keep existing snapshots unchanged** + +### Task 2: Add scoped resolution coverage + +**Files:** +- Create: `src/test/java/cn/beagile/dslquery/JoinOnSQLBuilderTest.java` + +- [ ] **Step 1: Write failing tests for `self`, `parent`, `root`, and unknown-field handling** +- [ ] **Step 2: Run the focused test command and confirm the failures are correct** + +## Chunk 2: Minimal implementation + +### Task 3: Add the new annotation and runtime API + +**Files:** +- Create: `src/main/java/cn/beagile/dslquery/JoinOn.java` +- Modify: `src/main/java/cn/beagile/dslquery/DSLQuery.java` + +- [ ] **Step 1: Add `@JoinOn` with runtime retention on fields** +- [ ] **Step 2: Add `DSLQuery.joinOn(path, dsl)` and storage for runtime join-on conditions** +- [ ] **Step 3: Run focused tests** + +### Task 4: Render join-on DSL in `JoinField` + +**Files:** +- Modify: `src/main/java/cn/beagile/dslquery/ColumnFields.java` +- Modify: `src/main/java/cn/beagile/dslquery/JoinField.java` +- Create: `src/main/java/cn/beagile/dslquery/JoinOnSQLBuilder.java` +- Modify: `src/main/java/cn/beagile/dslquery/SQLBuilder.java` +- Modify: `src/main/java/cn/beagile/dslquery/SingleExpression.java` + +- [ ] **Step 1: Thread merged join-on strings into each `JoinField`** +- [ ] **Step 2: Build scoped field lookup for `self`, `parent`, and `root`** +- [ ] **Step 3: Support `@fieldRef` values only when the builder opts in** +- [ ] **Step 4: Run focused tests and make them pass** + +## Chunk 3: Verification + +### Task 5: Regression verification + +**Files:** +- Modify: `src/test/java/cn/beagile/dslquery/ColumnFieldsTest.java` + +- [ ] **Step 1: Add a regression test proving joins are unchanged without join-on** +- [ ] **Step 2: Run focused join-related test suites** +- [ ] **Step 3: Run a broader repository test command if practical** diff --git a/docs/superpowers/specs/2026-03-30-join-on-scoped-dsl-design.md b/docs/superpowers/specs/2026-03-30-join-on-scoped-dsl-design.md new file mode 100644 index 0000000..7d673b4 --- /dev/null +++ b/docs/superpowers/specs/2026-03-30-join-on-scoped-dsl-design.md @@ -0,0 +1,97 @@ +# Join On Scoped DSL Design + +**Goal:** Add extra `join on` conditions without breaking the existing annotation model or string DSL style. + +## Summary + +The current join path always renders a fixed equality predicate from `@JoinColumn` or `@JoinColumns`. This design adds optional scoped DSL conditions that are appended to `on`, while keeping existing queries unchanged. + +The new capability has two entry points: + +- Field annotation: `@JoinOn("(and(enabled eq true))")` +- Runtime API: `DSLQuery.joinOn("org", "(and(type eq SALES))")` + +Both feed the same internal model and are merged with `and`. + +## DSL Semantics + +The feature reuses the existing DSL syntax: + +```java +(and(enabled eq true)(tenantId eq @parent.tenantId)) +``` + +Field resolution is scoped to the current join: + +- Bare field name: current join target (`self`) +- `self.xxx`: current join target +- `parent.xxx`: immediate owner of the join +- `root.xxx`: root query object + +Values that start with `@` are treated as field references instead of bound parameters: + +- `@parent.tenantId` +- `@root.companyId` +- `@self.code` + +This is only enabled for join-on rendering. Normal `where` behavior stays unchanged. + +## Compatibility Rules + +- No `@JoinOn` and no `DSLQuery.joinOn(...)` means generated SQL stays byte-for-byte compatible with the current implementation. +- Existing `where`, `sort`, `deepJoinIncludes`, and `selectIgnores` behavior remains unchanged. +- The new syntax does not modify `@JoinColumn` or `@JoinColumns`. + +## SQL Rendering + +Base join equality remains first, then extra predicates are appended with `and`. + +```sql +left join t_org org_ + on org_.id = t_user.org_id + and org_.enabled = :j0_p0 + and org_.tenant_id = t_user.tenant_id +``` + +## Boundaries + +- Extra `on` predicates can reference mapped `@Column` fields from `self`, `parent`, or `root`. +- They do not reference raw join-key column names unless those columns are also modeled as `@Column` fields. +- For `@JoinColumns`, v1 appends scoped predicates to the final target-table join, not to intermediate bridge joins. +- Unknown scoped fields fail fast instead of silently rendering `true`. + +## Internal Design + +### New API surface + +- Add `@JoinOn` +- Add `DSLQuery.joinOn(String path, String dsl)` + +### Builder flow + +```mermaid +flowchart LR + A["@JoinOn"] --> D["ColumnFields"] + B["DSLQuery.joinOn(path, dsl)"] --> D + D --> E["JoinField"] + E --> F["JoinOnSQLBuilder"] + F --> G["join SQL fragment"] + F --> H["shared params map"] +``` + +### Main implementation pieces + +- `DSLQuery`: store runtime join-on DSL strings by join path +- `ColumnFields`: pass merged join-on definitions into each `JoinField` +- `JoinField`: render base equality plus extra scoped predicates +- `JoinOnSQLBuilder`: resolve scoped field names and bind params into the main query param map +- `SingleExpression`: support builder-opted field-reference values for join-on rendering only + +## Test Strategy + +- Annotation-driven join-on adds constant predicate to `on` +- Runtime `joinOn(path, dsl)` adds predicate to `on` +- Annotation and runtime conditions merge with `and` +- `@parent.xxx` renders a column-to-column comparison +- Unknown scoped field throws +- Existing join SQL snapshots remain unchanged when no join-on is configured diff --git a/readme.md b/readme.md index 63d2d1c..116e2a1 100644 --- a/readme.md +++ b/readme.md @@ -11,6 +11,7 @@ - **类型安全**:基于JPA注解的强类型映射 - **自动SQL生成**:自动将DSL转换为优化的SQL查询 - **深度关联查询**:支持多级JOIN和OneToMany关系 +- **Join On扩展**:支持在关联ON子句中追加DSL条件 - **灵活配置**:支持字段忽略、深度关联控制、时区转换 - **分页支持**:内置分页查询功能 - **数据库无关**:通过QueryExecutor接口适配不同数据库 @@ -250,6 +251,22 @@ private Contact contact; private Org org; ``` +#### @JoinOn +为关联的 `join on` 子句追加DSL条件 + +```java +@JoinColumn(name = "org_id", referencedColumnName = "id") +@JoinOn("(and(enabled eq true)(tenantId eq @parent.tenantId))") +private Org org; +``` + +支持以下作用域: + +- 裸字段名或 `self.xxx`:当前join目标对象 +- `parent.xxx`:当前join的上一级对象 +- `root.xxx`:根查询对象 +- `@fieldPath`:把值解释为字段引用,而不是绑定参数 + #### @OneToMany 一对多关系(JPA标准注解) @@ -366,6 +383,59 @@ List result = new DSLQuery<>(executor, Person.class) // where area.name = 'Beijing' ``` +### 关联Join On扩展条件 + +```java +@View("person") +public class Person { + @Column(name = "tenant_id") + private Long tenantId; + + @JoinColumn(name = "org_id", referencedColumnName = "id") + @JoinOn("(and(enabled eq true)(tenantId eq @parent.tenantId))") + private Org org; +} + +@View("org") +public class Org { + @Column(name = "tenant_id") + private Long tenantId; + + @Column(name = "enabled") + private Boolean enabled; + + @Column(name = "type") + private String type; +} + +// 注解条件 + 运行时条件一起追加到ON +List result = new DSLQuery<>(executor, Person.class) + .joinOn("org", "(and(type eq SALES))") + .query(); + +// 生成的SQL类似: +// select person.tenant_id, org.tenant_id, org.enabled, org.type +// from person +// left join org on org.id = person.org_id +// and org.enabled = true +// and org.tenant_id = person.tenant_id +// and org.type = 'SALES' +``` + +深层关联同样支持作用域字段引用: + +```java +@View("org") +public class Org { + @Column(name = "tenant_id") + private Long tenantId; + + @JoinColumn(name = "area_id", referencedColumnName = "id") + @JoinOn("(and(code eq @root.areaCode)(tenantId eq @parent.tenantId))") + private Area area; +} +``` + ### OneToMany关系 ```java diff --git a/src/main/java/cn/beagile/dslquery/ColumnFields.java b/src/main/java/cn/beagile/dslquery/ColumnFields.java index a3ee4cc..2fc15b8 100644 --- a/src/main/java/cn/beagile/dslquery/ColumnFields.java +++ b/src/main/java/cn/beagile/dslquery/ColumnFields.java @@ -119,12 +119,29 @@ private void readJoinFields(Class clz, List parents, Class parents, Field field) { List newParents = newParents(parents, field); - joinFields.add(new JoinField(field, newParents)); + joinFields.add(new JoinField(field, newParents, joinOnConditions(field, newParents), joinFields.size())); readJoinColumnFields(field, newParents); readEmbeddedFields(field.getType(), newParents); readJoins(field.getType(), newParents); } + private List joinOnConditions(Field field, List parents) { + List result = new ArrayList<>(); + if (field.isAnnotationPresent(JoinOn.class)) { + result.addAll(Arrays.asList(field.getAnnotation(JoinOn.class).value())); + } + @SuppressWarnings("unchecked") + List outerJoinOns = (List) dslQuery.getJoinOns().get(pathOf(parents)); + if (outerJoinOns != null) { + result.addAll(outerJoinOns); + } + return result; + } + + private String pathOf(List parents) { + return parents.stream().map(Field::getName).collect(Collectors.joining(".")); + } + private void readJoinColumnFields(Field field, List newParents) { Arrays.stream(field.getType().getDeclaredFields()) .filter(f -> f.isAnnotationPresent(Column.class)) @@ -183,7 +200,11 @@ public List joined() { } public String joins() { - return joinFields.stream().map(JoinField::joinStatement).collect(Collectors.joining("\n")); + return joins(new HashMap<>(), 0); + } + + public String joins(Map params, int timezoneOffset) { + return joinFields.stream().map(joinField -> joinField.joinStatement(params, timezoneOffset)).collect(Collectors.joining("\n")); } public boolean hasField(Field field, List parents) { diff --git a/src/main/java/cn/beagile/dslquery/DSLQuery.java b/src/main/java/cn/beagile/dslquery/DSLQuery.java index 8c984d9..b289d38 100644 --- a/src/main/java/cn/beagile/dslquery/DSLQuery.java +++ b/src/main/java/cn/beagile/dslquery/DSLQuery.java @@ -2,7 +2,9 @@ import java.util.ArrayList; import java.util.Arrays; +import java.util.LinkedHashMap; import java.util.List; +import java.util.Map; public class DSLQuery { private final QueryExecutor queryExecutor; @@ -14,6 +16,7 @@ public class DSLQuery { private int timezoneOffset; private List deepJoins = new ArrayList<>(); private List selectIgnores = new ArrayList<>(); + private Map> joinOns = new LinkedHashMap<>(); private NullsOrder nullsOrder; private WhereParser whereParser = new WhereParser(); ; @@ -124,4 +127,16 @@ public DSLQuery selectIgnores(String... selectIgnores) { public List getSelectIgnores() { return selectIgnores; } + + public DSLQuery joinOn(String path, String dsl) { + if (path == null || path.isEmpty() || dsl == null || dsl.isEmpty()) { + return this; + } + joinOns.computeIfAbsent(path, key -> new ArrayList<>()).add(dsl); + return this; + } + + public Map> getJoinOns() { + return joinOns; + } } diff --git a/src/main/java/cn/beagile/dslquery/DSLSQLBuilder.java b/src/main/java/cn/beagile/dslquery/DSLSQLBuilder.java index 6039325..745e981 100644 --- a/src/main/java/cn/beagile/dslquery/DSLSQLBuilder.java +++ b/src/main/java/cn/beagile/dslquery/DSLSQLBuilder.java @@ -60,9 +60,7 @@ public void addParamArray(String paramName, String fieldName, String value) { } private List castValueToList(String value, Field field) { - return Stream.of(new Gson().fromJson(value, String[].class)) - .map(v -> castValueByField(v, field)) - .collect(Collectors.toList()); + return castValueToList(value, field, this.dslQuery.getTimezoneOffset()); } private void setParam(String paramName, String fieldName, String value, BiFunction valueConverter) { @@ -72,23 +70,33 @@ private void setParam(String paramName, String fieldName, String value, BiFuncti } private Object castValueByField(String value, Field field) { + return castValueByField(value, field, this.dslQuery.getTimezoneOffset()); + } + + static List castValueToList(String value, Field field, int timezoneOffset) { + return Stream.of(new Gson().fromJson(value, String[].class)) + .map(v -> castValueByField(v, field, timezoneOffset)) + .collect(Collectors.toList()); + } + + static Object castValueByField(String value, Field field, int timezoneOffset) { if (isInstant(field.getType())) { - return getInstantValue(value, field); + return getInstantValue(value, field, timezoneOffset); } if (isTimestampAsDate(field)) { if (field.getType().equals(Timestamp.class)) { - return Timestamp.from(getInstantValue(value, field)); + return Timestamp.from(getInstantValue(value, field, timezoneOffset)); } - return getInstantValue(value, field).toEpochMilli(); + return getInstantValue(value, field, timezoneOffset).toEpochMilli(); } return FIELD_CAST_MAP.get(field.getType()).apply(value); } - private boolean isInstant(Class type) { + private static boolean isInstant(Class type) { return type.equals(Instant.class); } - private boolean isTimestampAsDate(Field field) { + private static boolean isTimestampAsDate(Field field) { if (field.getType().equals(Long.class)) { return field.isAnnotationPresent(DateFormat.class); } @@ -101,10 +109,10 @@ private boolean isTimestampAsDate(Field field) { return false; } - private Instant getInstantValue(String value, Field field) { + private static Instant getInstantValue(String value, Field field, int timezoneOffset) { String dateFormat = field.getAnnotation(DateFormat.class).value(); DateTimeFormatter formatter = DateTimeFormatter.ofPattern(dateFormat); - ZoneId zoneId = ZoneOffset.ofHours(this.dslQuery.getTimezoneOffset()).normalized(); + ZoneId zoneId = ZoneOffset.ofHours(timezoneOffset).normalized(); return LocalDateTime.parse(value, formatter).atZone(zoneId).toInstant(); } @@ -160,7 +168,7 @@ private List getWhereList() { private String getSelectSQL() { String select = "select" + columnFields.distinct() + columnFields.selectFields().stream().map(ColumnField::expression).collect(Collectors.joining(",")) + " from " + columnFields.from(); - String join = columnFields.joins(); + String join = columnFields.joins(this.params, this.dslQuery.getTimezoneOffset()); return String.join("\n", select, join); } @@ -171,7 +179,7 @@ private String getSortSQL() { private String getCountSQL() { List lines = new ArrayList<>(); String countField = getCountField(); - lines.add(String.format("select count(" + countField + ") from %s\n%s", columnFields.from(), columnFields.joins())); + lines.add(String.format("select count(" + countField + ") from %s\n%s", columnFields.from(), columnFields.joins(this.params, this.dslQuery.getTimezoneOffset()))); if (!getWhereList().isEmpty()) { lines.add(getWhereSQL()); } diff --git a/src/main/java/cn/beagile/dslquery/JoinField.java b/src/main/java/cn/beagile/dslquery/JoinField.java index ca8efd6..d4b20e1 100644 --- a/src/main/java/cn/beagile/dslquery/JoinField.java +++ b/src/main/java/cn/beagile/dslquery/JoinField.java @@ -6,26 +6,35 @@ import java.lang.reflect.Field; import java.util.ArrayList; import java.util.List; +import java.util.Map; import java.util.stream.Collectors; import java.util.stream.Stream; public class JoinField { private Field field; private List parents; + private List joinOns; + private int joinIndex; - public JoinField(Field field, List parents) { + public JoinField(Field field, List parents, List joinOns, int joinIndex) { this.field = field; this.parents = parents; + this.joinOns = joinOns; + this.joinIndex = joinIndex; } public String joinStatement() { + return joinStatement(new java.util.HashMap<>(), 0); + } + + public String joinStatement(Map params, int timezoneOffset) { if (field.isAnnotationPresent(JoinColumn.class)) { - return singleJoinStatement(); + return singleJoinStatement(params, timezoneOffset); } - return multiJoinStatement(); + return multiJoinStatement(params, timezoneOffset); } - private String multiJoinStatement() { + private String multiJoinStatement(Map params, int timezoneOffset) { List result = new ArrayList<>(); JoinColumn[] columns = field.getAnnotation(JoinColumns.class).value(); for (int i = 0; i < columns.length; i++) { @@ -42,7 +51,9 @@ private String multiJoinStatement() { .joinTableAlias(getTableAlias()) .joinField(joinColumn.referencedColumnName()) .onTable(preColumn.table()) - .onField(joinColumn.name()).build()); + .onField(joinColumn.name()) + .joinOnClause(joinOnClause(params, timezoneOffset)) + .build()); } } return result.stream().collect(Collectors.joining("\n")); @@ -52,14 +63,23 @@ private String getTableAlias() { return parents.stream().map(Field::getName).collect(Collectors.joining("_", "", "_")); } - private String singleJoinStatement() { + private String singleJoinStatement(Map params, int timezoneOffset) { JoinColumn joinColumn = field.getAnnotation(JoinColumn.class); return JoinBuilder.joinBuilder() .joinTable(field.getType().getAnnotation(View.class).value()) .joinTableAlias(getTableAlias()) .joinField(joinColumn.referencedColumnName()) .onTable(getJoinTable()) - .onField(joinColumn.name()).build(); + .onField(joinColumn.name()) + .joinOnClause(joinOnClause(params, timezoneOffset)) + .build(); + } + + private String joinOnClause(Map params, int timezoneOffset) { + if (joinOns.isEmpty()) { + return ""; + } + return " and " + new JoinOnSQLBuilder(field, parents, "j" + joinIndex + "_", params, timezoneOffset).build(joinOns); } private String getJoinTable() { @@ -83,6 +103,7 @@ public static class JoinBuilder { private String onTable; private String onField; private String joinField; + private String joinOnClause = ""; public static JoinBuilder joinBuilder() { return new JoinBuilder(); @@ -113,13 +134,18 @@ public JoinBuilder joinField(String joinField) { return this; } + public JoinBuilder joinOnClause(String joinOnClause) { + this.joinOnClause = joinOnClause; + return this; + } + public String build() { if (joinTableAlias == null) { return "left join " + joinTable + " on " + joinTable + "." + joinField - + " = " + onTable + "." + onField; + + " = " + onTable + "." + onField + joinOnClause; } return "left join " + joinTable + " " + joinTableAlias + " on " + joinTableAlias + "." + joinField - + " = " + onTable + "." + onField; + + " = " + onTable + "." + onField + joinOnClause; } } } diff --git a/src/main/java/cn/beagile/dslquery/JoinOn.java b/src/main/java/cn/beagile/dslquery/JoinOn.java new file mode 100644 index 0000000..a17c9de --- /dev/null +++ b/src/main/java/cn/beagile/dslquery/JoinOn.java @@ -0,0 +1,12 @@ +package cn.beagile.dslquery; + +import java.lang.annotation.ElementType; +import java.lang.annotation.Retention; +import java.lang.annotation.RetentionPolicy; +import java.lang.annotation.Target; + +@Retention(RetentionPolicy.RUNTIME) +@Target(ElementType.FIELD) +public @interface JoinOn { + String[] value() default {}; +} diff --git a/src/main/java/cn/beagile/dslquery/JoinOnSQLBuilder.java b/src/main/java/cn/beagile/dslquery/JoinOnSQLBuilder.java new file mode 100644 index 0000000..653f988 --- /dev/null +++ b/src/main/java/cn/beagile/dslquery/JoinOnSQLBuilder.java @@ -0,0 +1,96 @@ +package cn.beagile.dslquery; + +import com.google.gson.Gson; + +import javax.persistence.Column; +import java.lang.reflect.Field; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; +import java.util.stream.Collectors; +import java.util.stream.Stream; + +class JoinOnSQLBuilder implements SQLBuilder { + private final Map fields = new LinkedHashMap<>(); + private final Map params; + private final String paramPrefix; + private final int timezoneOffset; + + JoinOnSQLBuilder(Field joinField, List parents, String paramPrefix, Map params, int timezoneOffset) { + this.params = params; + this.paramPrefix = paramPrefix; + this.timezoneOffset = timezoneOffset; + initFields(joinField, parents); + } + + String build(List joinOns) { + if (joinOns.isEmpty()) { + return ""; + } + WhereParser parser = new WhereParser(); + return joinOns.stream() + .map(parser::parse) + .map(expression -> expression.toSQL(this)) + .collect(Collectors.joining(" and ")); + } + + private void initFields(Field joinField, List parents) { + Class rootClass = parents.get(0).getDeclaringClass(); + addScope("", readColumns(joinField.getType(), joinField.getType(), parents, true)); + addScope("self.", readColumns(joinField.getType(), joinField.getType(), parents, true)); + if (parents.size() == 1) { + addScope("parent.", readColumns(rootClass, rootClass, new ArrayList<>(), false)); + } else { + addScope("parent.", readColumns(joinField.getDeclaringClass(), joinField.getDeclaringClass(), parents.subList(0, parents.size() - 1), true)); + } + addScope("root.", readColumns(rootClass, rootClass, new ArrayList<>(), false)); + } + + private void addScope(String prefix, List scopedFields) { + scopedFields.forEach(field -> fields.put(prefix + field.getField().getName(), field)); + } + + private List readColumns(Class ownerClass, Class rootClass, List parents, boolean joined) { + return Arrays.stream(ownerClass.getDeclaredFields()) + .filter(field -> field.isAnnotationPresent(Column.class)) + .map(field -> new ColumnField(field, rootClass, new ArrayList<>(parents), field.getAnnotation(Column.class), joined)) + .collect(Collectors.toList()); + } + + @Override + public void addParamArray(String paramName, String fieldName, String value) { + params.put(paramName, Stream.of(new Gson().fromJson(value, String[].class)) + .map(it -> DSLSQLBuilder.castValueByField(it, findField(fieldName).getField(), timezoneOffset)) + .collect(Collectors.toList())); + } + + @Override + public void addParam(String paramName, String fieldName, String value) { + params.put(paramName, DSLSQLBuilder.castValueByField(value, findField(fieldName).getField(), timezoneOffset)); + } + + @Override + public String aliasOf(String field) { + return findField(field).selectName(); + } + + @Override + public String paramName(String rawParamName) { + return paramPrefix + rawParamName; + } + + @Override + public boolean supportsFieldReferenceValues() { + return true; + } + + private ColumnField findField(String field) { + ColumnField columnField = fields.get(field); + if (columnField == null) { + throw new RuntimeException("field not found: " + field); + } + return columnField; + } +} diff --git a/src/main/java/cn/beagile/dslquery/SQLBuilder.java b/src/main/java/cn/beagile/dslquery/SQLBuilder.java index 0248650..68a5680 100644 --- a/src/main/java/cn/beagile/dslquery/SQLBuilder.java +++ b/src/main/java/cn/beagile/dslquery/SQLBuilder.java @@ -7,4 +7,12 @@ public interface SQLBuilder { void addParam(String paramName, String fieldName, String value); String aliasOf(String field); + + default String paramName(String rawParamName) { + return rawParamName; + } + + default boolean supportsFieldReferenceValues() { + return false; + } } diff --git a/src/main/java/cn/beagile/dslquery/SingleExpression.java b/src/main/java/cn/beagile/dslquery/SingleExpression.java index a4d4255..20ba114 100644 --- a/src/main/java/cn/beagile/dslquery/SingleExpression.java +++ b/src/main/java/cn/beagile/dslquery/SingleExpression.java @@ -4,6 +4,7 @@ import java.io.UnsupportedEncodingException; import java.net.URLEncoder; import java.nio.charset.StandardCharsets; +import java.util.Arrays; import java.util.List; import java.util.Objects; @@ -63,13 +64,29 @@ public String toString() { public String toSQL(SQLBuilder sqlBuilder) { Operator operatorEnum = Operators.byName(this.operator); - String[] paramNames = operatorEnum.params(paramName); + String[] paramNames = Arrays.stream(operatorEnum.params(paramName)) + .map(sqlBuilder::paramName) + .toArray(String[]::new); + if (operatorEnum.requireValue && sqlBuilder.supportsFieldReferenceValues() && isFieldReferenceValue()) { + return toFieldReferenceSQL(sqlBuilder, operatorEnum); + } if (operatorEnum.requireValue) { addParams(sqlBuilder, operatorEnum, paramNames); } return String.format(operatorEnum.whereFormat(), sqlBuilder.aliasOf(field), operatorEnum.operator, paramNames[0], paramNames[1]); } + private boolean isFieldReferenceValue() { + return value != null && value.startsWith("@"); + } + + private String toFieldReferenceSQL(SQLBuilder sqlBuilder, Operator operatorEnum) { + if (operatorEnum.isArray() || operatorEnum == Operator.Between) { + throw new RuntimeException("field reference value not supported for operator:" + operatorEnum.keyword); + } + return String.format("(%s %s %s)", sqlBuilder.aliasOf(field), operatorEnum.operator, sqlBuilder.aliasOf(value.substring(1))); + } + private void addParams(SQLBuilder sqlBuilder, Operator operatorEnum, String[] paramNames) { if (operatorEnum.isArray()) { sqlBuilder.addParamArray(paramNames[0], field, operatorEnum.transferValue(this.value)); diff --git a/src/test/java/cn/beagile/dslquery/DSLQueryTest.java b/src/test/java/cn/beagile/dslquery/DSLQueryTest.java index c2a3476..4a9a1e4 100644 --- a/src/test/java/cn/beagile/dslquery/DSLQueryTest.java +++ b/src/test/java/cn/beagile/dslquery/DSLQueryTest.java @@ -55,6 +55,16 @@ public void should_execute_query_with_where() { assertEquals(1, params.size()); } + @Test + public void should_treat_at_prefixed_where_value_as_literal() { + DSLQuery dslQuery = new DSLQuery(queryExecutor, QueryResultBean.class); + dslQuery.where("(and(name equals @bob))").query(); + verify(queryExecutor).list(any(), sqlQueryArgumentCaptor.capture()); + SQLQuery sqlQuery = sqlQueryArgumentCaptor.getValue(); + expectSqlWithoudEnter("select " + fields + " from view_query where ((view_query.name = :p0))", sqlQuery.getSql()); + assertEquals("@bob", sqlQuery.getParams().get("p0")); + } + @Test public void should_execute_query_with_where_between() { DSLQuery dslQuery = new DSLQuery(queryExecutor, QueryResultBean.class); diff --git a/src/test/java/cn/beagile/dslquery/JoinOnApiTest.java b/src/test/java/cn/beagile/dslquery/JoinOnApiTest.java new file mode 100644 index 0000000..816522f --- /dev/null +++ b/src/test/java/cn/beagile/dslquery/JoinOnApiTest.java @@ -0,0 +1,21 @@ +package cn.beagile.dslquery; + +import org.junit.jupiter.api.Test; + +import java.lang.reflect.Method; + +import static org.junit.jupiter.api.Assertions.assertDoesNotThrow; +import static org.junit.jupiter.api.Assertions.assertNotNull; + +class JoinOnApiTest { + @Test + void should_expose_join_on_annotation() { + assertDoesNotThrow(() -> Class.forName("cn.beagile.dslquery.JoinOn")); + } + + @Test + void should_expose_runtime_join_on_api() { + Method method = assertDoesNotThrow(() -> DSLQuery.class.getMethod("joinOn", String.class, String.class)); + assertNotNull(method); + } +} diff --git a/src/test/java/cn/beagile/dslquery/JoinOnQueryTest.java b/src/test/java/cn/beagile/dslquery/JoinOnQueryTest.java new file mode 100644 index 0000000..82d277a --- /dev/null +++ b/src/test/java/cn/beagile/dslquery/JoinOnQueryTest.java @@ -0,0 +1,116 @@ +package cn.beagile.dslquery; + +import org.junit.jupiter.api.Test; + +import javax.persistence.Column; +import javax.persistence.JoinColumn; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertThrows; + +class JoinOnQueryTest { + private final String nullsOrder = ""; + + @View("t_user") + public static class User { + @Column(name = "tenant_id") + private Long tenantId; + + @JoinColumn(name = "org_id", referencedColumnName = "id") + @JoinOn("(and(enabled eq true)(tenantId eq @parent.tenantId))") + private Org org; + } + + @View("t_org") + public static class Org { + @Column(name = "id") + private Long id; + + @Column(name = "enabled") + private Boolean enabled; + + @Column(name = "tenant_id") + private Long tenantId; + + @Column(name = "type") + private String type; + } + + @View("t_user") + @DeepJoinIncludes("org.area") + public static class DeepUser { + @Column(name = "area_code") + private String areaCode; + + @JoinColumn(name = "org_id", referencedColumnName = "id") + private DeepOrg org; + } + + @View("t_org") + public static class DeepOrg { + @Column(name = "tenant_id") + private Long tenantId; + + @JoinColumn(name = "area_id", referencedColumnName = "id") + @JoinOn("(and(code eq @root.areaCode)(tenantId eq @parent.tenantId))") + private Area area; + } + + @View("t_area") + public static class Area { + @Column(name = "id") + private Long id; + + @Column(name = "code") + private String code; + + @Column(name = "tenant_id") + private Long tenantId; + + @Column(name = "active") + private Boolean active; + } + + @View("t_user") + public static class InvalidUser { + @JoinColumn(name = "org_id", referencedColumnName = "id") + @JoinOn("(and(missing eq true))") + private Org org; + } + + @Test + void should_append_annotation_and_runtime_conditions_to_join_on() { + DSLQuery dslQuery = new DSLQuery<>(null, User.class) + .joinOn("org", "(and(type eq SALES))"); + + DSLSQLBuilder sqlBuilder = new DSLSQLBuilder<>(dslQuery); + SQLQuery sqlQuery = sqlBuilder.build(nullsOrder); + + assertEquals("select t_user.tenant_id tenantId_,org_.id org_id_,org_.enabled org_enabled_,org_.tenant_id org_tenantId_,org_.type org_type_ from t_user\n" + + "left join t_org org_ on org_.id = t_user.org_id and ((org_.enabled = :j0_p0) and (org_.tenant_id = t_user.tenant_id)) and ((org_.type = :j0_p2))", sqlQuery.getSql()); + assertEquals(true, sqlQuery.getParams().get("j0_p0")); + assertEquals("SALES", sqlQuery.getParams().get("j0_p2")); + } + + @Test + void should_resolve_parent_and_root_fields_for_deep_join_conditions() { + DSLQuery dslQuery = new DSLQuery<>(null, DeepUser.class) + .joinOn("org.area", "(and(active eq true))"); + + DSLSQLBuilder sqlBuilder = new DSLSQLBuilder<>(dslQuery); + SQLQuery sqlQuery = sqlBuilder.build(nullsOrder); + + assertEquals("select t_user.area_code areaCode_,org_.tenant_id org_tenantId_,org_area_.id org_area_id_,org_area_.code org_area_code_,org_area_.tenant_id org_area_tenantId_,org_area_.active org_area_active_ from t_user\n" + + "left join t_org org_ on org_.id = t_user.org_id\n" + + "left join t_area org_area_ on org_area_.id = org_.area_id and ((org_area_.code = t_user.area_code) and (org_area_.tenant_id = org_.tenant_id)) and ((org_area_.active = :j1_p2))", sqlQuery.getSql()); + assertEquals(true, sqlQuery.getParams().get("j1_p2")); + } + + @Test + void should_fail_fast_when_join_on_field_is_unknown() { + DSLSQLBuilder sqlBuilder = new DSLSQLBuilder<>(new DSLQuery<>(null, InvalidUser.class)); + + RuntimeException exception = assertThrows(RuntimeException.class, () -> sqlBuilder.build(nullsOrder)); + assertEquals("field not found: missing", exception.getMessage()); + } +}