diff --git a/packages/3-extensions/sql-orm-client/src/collection-contract.ts b/packages/3-extensions/sql-orm-client/src/collection-contract.ts index c09d7626999e..2174ee9f0426 100644 --- a/packages/3-extensions/sql-orm-client/src/collection-contract.ts +++ b/packages/3-extensions/sql-orm-client/src/collection-contract.ts @@ -300,8 +300,10 @@ export interface ResolvedIncludeRelation { readonly relatedNamespaceId: string; readonly relatedTableName: string; readonly localTableName: string; - readonly targetColumn: string; - readonly localColumn: string; + /** Target-side join columns, positionally paired with `localColumns`. */ + readonly targetColumns: readonly string[]; + /** Local-side join columns, positionally paired with `targetColumns`. */ + readonly localColumns: readonly string[]; readonly cardinality: RelationCardinalityTag | undefined; readonly through?: IncludeThroughDescriptor; } @@ -336,22 +338,31 @@ export function resolveIncludeRelation( { meta: { model: baseModelName, relation: relationName } }, ); } - const localField = relation.on.localFields[0]; - const targetField = relation.on.targetFields[0]; - if (!localField || !targetField) { + const localFields = relation.on.localFields; + const targetFields = relation.on.targetFields; + const localColumns: string[] = []; + const targetColumns: string[] = []; + const pairCount = Math.min(localFields.length, targetFields.length); + + for (let i = 0; i < pairCount; i++) { + const localField = localFields[i]; + const targetField = targetFields[i]; + if (!localField || !targetField) { + continue; + } + localColumns.push(resolveFieldToColumn(contract, namespaceId, declaringModelName, localField)); + targetColumns.push( + resolveFieldToColumn(contract, relation.toNamespace, relation.to, targetField), + ); + } + + if (localColumns.length === 0) { throw new InternalError( `Relation '${relationName}' on model '${declaringModelName}' has incomplete join metadata (missing localFields or targetFields)`, ); } const relatedTableName = resolveModelTableName(contract, relation.toNamespace, relation.to); - const localColumn = resolveFieldToColumn(contract, namespaceId, declaringModelName, localField); - const targetColumn = resolveFieldToColumn( - contract, - relation.toNamespace, - relation.to, - targetField, - ); let through: IncludeThroughDescriptor | undefined; if (relation.through !== undefined) { @@ -373,8 +384,8 @@ export function resolveIncludeRelation( relatedNamespaceId: relation.toNamespace, relatedTableName, localTableName, - targetColumn, - localColumn, + targetColumns, + localColumns, cardinality: relation.cardinality, ...ifDefined('through', through), }; diff --git a/packages/3-extensions/sql-orm-client/src/collection.ts b/packages/3-extensions/sql-orm-client/src/collection.ts index 922c8ab81dbe..bc486d356a58 100644 --- a/packages/3-extensions/sql-orm-client/src/collection.ts +++ b/packages/3-extensions/sql-orm-client/src/collection.ts @@ -618,8 +618,8 @@ class CollectionImpl< relatedNamespaceId: relation.relatedNamespaceId, relatedTableName: relation.relatedTableName, localTableName: relation.localTableName, - targetColumn: relation.targetColumn, - localColumn: relation.localColumn, + targetColumns: relation.targetColumns, + localColumns: relation.localColumns, cardinality: relation.cardinality, ...ifDefined('through', relation.through), nested: nestedState, diff --git a/packages/3-extensions/sql-orm-client/src/query-plan-select.ts b/packages/3-extensions/sql-orm-client/src/query-plan-select.ts index 0153fcc6b43a..8cf5c32bf5ed 100644 --- a/packages/3-extensions/sql-orm-client/src/query-plan-select.ts +++ b/packages/3-extensions/sql-orm-client/src/query-plan-select.ts @@ -261,7 +261,36 @@ interface IncludeParentSource { } function localColumnsForRowInclude(include: IncludeExpr): readonly string[] { - return include.through?.parentLocalColumns ?? [include.localColumn]; + return include.through?.parentLocalColumns ?? include.localColumns; +} + +/** + * Correlate a child row back to its parent across every column of the + * relation's key. Composite foreign keys contribute one equality per + * column, ANDed together — mirroring the relation-filter join in + * `model-accessor.ts`. Correlating on a prefix of the key would match + * every child sharing that prefix. + */ +function buildIncludeJoinExpr( + include: IncludeExpr, + childTableRef: string, + parentLocalRefs: readonly ColumnRef[], +): AnyExpression { + const joinExprs: AnyExpression[] = []; + const count = Math.min(parentLocalRefs.length, include.targetColumns.length); + + for (let i = 0; i < count; i++) { + const parentLocalRef = parentLocalRefs[i]; + const targetColumn = include.targetColumns[i]; + if (parentLocalRef === undefined || targetColumn === undefined) { + continue; + } + joinExprs.push(BinaryExpr.eq(ColumnRef.of(childTableRef, targetColumn), parentLocalRef)); + } + + const firstExpr = joinExprs[0]; + assertDefined(firstExpr, `Include '${include.relationName}' has no parent-local column ref`); + return joinExprs.length === 1 ? firstExpr : AndExpr.of(joinExprs); } function resolveParentLocalRefs( @@ -578,15 +607,7 @@ function buildIncludeChildRowsSelect( whereExpr = childWhere ? AndExpr.of([artifacts.whereExpr, childWhere]) : artifacts.whereExpr; junctionJoins = [artifacts.junctionJoin]; } else { - const parentLocalRef = parentLocalRefs[0]; - assertDefined( - parentLocalRef, - `Include '${include.relationName}' has no parent-local column ref`, - ); - const joinExpr = BinaryExpr.eq( - ColumnRef.of(childTableRef, include.targetColumn), - parentLocalRef, - ); + const joinExpr = buildIncludeJoinExpr(include, childTableRef, parentLocalRefs); whereExpr = childWhere ? AndExpr.of([joinExpr, childWhere]) : joinExpr; } @@ -1019,15 +1040,7 @@ function buildIncludeChildScalarSelect( whereExpr = childWhere ? AndExpr.of([artifacts.whereExpr, childWhere]) : artifacts.whereExpr; junctionJoins = [artifacts.junctionJoin]; } else { - const parentLocalRef = parentLocalRefs[0]; - assertDefined( - parentLocalRef, - `Include '${include.relationName}' has no parent-local column ref`, - ); - const joinExpr = BinaryExpr.eq( - ColumnRef.of(childTableRef, include.targetColumn), - parentLocalRef, - ); + const joinExpr = buildIncludeJoinExpr(include, childTableRef, parentLocalRefs); whereExpr = childWhere ? AndExpr.of([joinExpr, childWhere]) : joinExpr; } diff --git a/packages/3-extensions/sql-orm-client/src/types.ts b/packages/3-extensions/sql-orm-client/src/types.ts index f8f8efe8f583..238526d1ee0c 100644 --- a/packages/3-extensions/sql-orm-client/src/types.ts +++ b/packages/3-extensions/sql-orm-client/src/types.ts @@ -72,8 +72,10 @@ export interface IncludeExpr { readonly relatedNamespaceId: string; readonly relatedTableName: string; readonly localTableName: string; - readonly targetColumn: string; - readonly localColumn: string; + /** Target-side join columns, positionally paired with `localColumns`. */ + readonly targetColumns: readonly string[]; + /** Local-side join columns, positionally paired with `targetColumns`. */ + readonly localColumns: readonly string[]; readonly cardinality: RelationCardinalityTag | undefined; readonly through?: IncludeThroughDescriptor; readonly nested: CollectionState; diff --git a/packages/3-extensions/sql-orm-client/test/collection-contract.test.ts b/packages/3-extensions/sql-orm-client/test/collection-contract.test.ts index b78e06c5841b..df6b9ab8d51b 100644 --- a/packages/3-extensions/sql-orm-client/test/collection-contract.test.ts +++ b/packages/3-extensions/sql-orm-client/test/collection-contract.test.ts @@ -102,8 +102,38 @@ describe('collection-contract capability detection', () => { relatedNamespaceId: 'public', relatedTableName: 'posts', localTableName: 'users', - targetColumn: 'user_id', - localColumn: 'id', + targetColumns: ['user_id'], + localColumns: ['id'], + cardinality: '1:N', + }); + }); + + it('resolveIncludeRelation() resolves every column of a composite foreign key', () => { + const composite = withPatchedDomainModels(getTestContract(), (models) => { + const user = models['User'] as Record; + return { + ...models, + User: { + ...user, + relations: { + ...(user['relations'] as Record), + posts: { + to: { model: 'Post', namespace: 'public' }, + cardinality: '1:N', + on: { localFields: ['id', 'email'], targetFields: ['userId', 'title'] }, + }, + }, + }, + }; + }); + + expect(resolveIncludeRelation(composite, 'public', 'User', 'posts')).toEqual({ + relatedModelName: 'Post', + relatedNamespaceId: 'public', + relatedTableName: 'posts', + localTableName: 'users', + targetColumns: ['user_id', 'title'], + localColumns: ['id', 'email'], cardinality: '1:N', }); }); diff --git a/packages/3-extensions/sql-orm-client/test/collection-dispatch.test.ts b/packages/3-extensions/sql-orm-client/test/collection-dispatch.test.ts index 551f60101276..0cc5a7dfbd23 100644 --- a/packages/3-extensions/sql-orm-client/test/collection-dispatch.test.ts +++ b/packages/3-extensions/sql-orm-client/test/collection-dispatch.test.ts @@ -31,8 +31,8 @@ function includeFor( relatedTableName: relation.relatedTableName, relatedNamespaceId: relation.relatedNamespaceId, localTableName: relation.localTableName, - targetColumn: relation.targetColumn, - localColumn: relation.localColumn, + targetColumns: relation.targetColumns, + localColumns: relation.localColumns, cardinality: relation.cardinality, nested, scalar: undefined, diff --git a/packages/3-extensions/sql-orm-client/test/collection.state.test.ts b/packages/3-extensions/sql-orm-client/test/collection.state.test.ts index 6c064fb68c86..8f399328493a 100644 --- a/packages/3-extensions/sql-orm-client/test/collection.state.test.ts +++ b/packages/3-extensions/sql-orm-client/test/collection.state.test.ts @@ -188,7 +188,7 @@ describe('Collection', () => { relationName: 'posts', relatedModelName: 'Post', relatedTableName: 'posts', - targetColumn: 'user_id', + targetColumns: ['user_id'], cardinality: '1:N', }); expect(withPosts.state.includes[0]?.nested.filters).toEqual([ @@ -240,8 +240,8 @@ describe('Collection', () => { relationName: 'author', relatedModelName: 'User', relatedTableName: 'users', - targetColumn: 'id', - localColumn: 'user_id', + targetColumns: ['id'], + localColumns: ['user_id'], cardinality: 'N:1', }); diff --git a/packages/3-extensions/sql-orm-client/test/query-plan-select.test.ts b/packages/3-extensions/sql-orm-client/test/query-plan-select.test.ts index a3b96bf86d40..8bde9d2a12be 100644 --- a/packages/3-extensions/sql-orm-client/test/query-plan-select.test.ts +++ b/packages/3-extensions/sql-orm-client/test/query-plan-select.test.ts @@ -144,6 +144,43 @@ describe('compileSelectWithIncludes', () => { ); }); + it('correlates a composite foreign key on every column pair', () => { + const include: IncludeExpr = { + relationName: 'posts', + relatedModelName: 'Post', + relatedNamespaceId: 'public', + relatedTableName: 'posts', + localTableName: 'users', + targetColumns: ['user_id', 'title'], + localColumns: ['id', 'email'], + cardinality: '1:N', + nested: emptyState(), + scalar: undefined, + combine: undefined, + }; + + const plan = compileSelectWithIncludes(baseContract, getTestAggregates(), 'public', 'users', { + ...emptyState(), + includes: [include], + }); + + expectSelectAst(plan.ast); + const postsProjection = plan.ast.projection.find((item) => item.alias === 'posts'); + expectSubqueryExpr(postsProjection?.expr); + + const childRowsSource = postsProjection.expr.query.from; + expectDerivedTableSource(childRowsSource); + + // Correlating on `user_id` alone would match every post sharing it, + // so both pairs of the key have to appear. + expect(childRowsSource.query.where).toEqual( + AndExpr.of([ + BinaryExpr.eq(ColumnRef.of('posts', 'user_id'), ColumnRef.of('users', 'id')), + BinaryExpr.eq(ColumnRef.of('posts', 'title'), ColumnRef.of('users', 'email')), + ]), + ); + }); + it('builds lexicographic cursor filters with distinctOn, limit, and offset', () => { const { collection } = createCollection(); const state = collection @@ -1400,8 +1437,8 @@ describe('compileSelectWithIncludes polymorphic targets', () => { relatedTableName: relation.relatedTableName, relatedNamespaceId: relation.relatedNamespaceId, localTableName: relation.localTableName, - targetColumn: relation.targetColumn, - localColumn: relation.localColumn, + targetColumns: relation.targetColumns, + localColumns: relation.localColumns, cardinality: relation.cardinality, nested, scalar: undefined, diff --git a/packages/3-extensions/sql-orm-client/test/variant-include.collection-contract.test.ts b/packages/3-extensions/sql-orm-client/test/variant-include.collection-contract.test.ts index 6651e02409ac..18b17c165ad4 100644 --- a/packages/3-extensions/sql-orm-client/test/variant-include.collection-contract.test.ts +++ b/packages/3-extensions/sql-orm-client/test/variant-include.collection-contract.test.ts @@ -17,8 +17,8 @@ describe('resolveIncludeRelation() with a selected parent variant', () => { relatedNamespaceId: 'public', relatedTableName: 'assignees', localTableName: 'features', - localColumn: 'assignee_id', - targetColumn: 'id', + localColumns: ['assignee_id'], + targetColumns: ['id'], cardinality: 'N:1', }); }); @@ -37,8 +37,8 @@ describe('resolveIncludeRelation() with a selected parent variant', () => { relatedNamespaceId: 'public', relatedTableName: 'assignees', localTableName: 'tasks', - localColumn: 'assignee_id', - targetColumn: 'id', + localColumns: ['assignee_id'], + targetColumns: ['id'], cardinality: 'N:1', }); }); @@ -57,8 +57,8 @@ describe('resolveIncludeRelation() with a selected parent variant', () => { relatedNamespaceId: 'public', relatedTableName: 'tasks', localTableName: 'tasks', - localColumn: 'id', - targetColumn: 'parent_id', + localColumns: ['id'], + targetColumns: ['parent_id'], cardinality: '1:N', }); }); diff --git a/packages/3-extensions/sql-orm-client/test/variant-include.query-plan-fixtures.ts b/packages/3-extensions/sql-orm-client/test/variant-include.query-plan-fixtures.ts index 44df4657ba4b..82494b992bde 100644 --- a/packages/3-extensions/sql-orm-client/test/variant-include.query-plan-fixtures.ts +++ b/packages/3-extensions/sql-orm-client/test/variant-include.query-plan-fixtures.ts @@ -45,8 +45,8 @@ export function includeExpr(options: { relatedNamespaceId: 'public', relatedTableName: options.relatedTableName, localTableName: options.localTableName, - targetColumn: options.targetColumn, - localColumn: options.localColumn, + targetColumns: [options.targetColumn], + localColumns: [options.localColumn], cardinality: options.cardinality, ...ifDefined('through', options.through), nested: options.nested ?? emptyState(),