Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
19 changes: 18 additions & 1 deletion packages/orm/src/client/crud-types.ts
Original file line number Diff line number Diff line change
Expand Up @@ -353,6 +353,20 @@ export type BatchResult = { count: number };

//#region Common structures

/**
* Context object passed to `$expr` filters.
*/
export type ExprFilterContext<Schema extends SchemaDef, Model extends GetModels<Schema>> = {
/**
* The alias that can be used to refer to the filtered model. It's the model name for top-level
* filters, but relation filters select the model under a generated alias.
*
* Typed as the model name to match the expression builder's scope, so qualified references
* like `` eb.ref(`${modelAlias}.field`) `` type-check.
*/
modelAlias: Model;
};

export type WhereInput<
Schema extends SchemaDef,
Model extends GetModels<Schema>,
Expand All @@ -370,7 +384,10 @@ export type WhereInput<
{ args: ComputedFieldArgs<Schema, Model, Key> } & FieldFilter<Schema, Model, Key, Options, WithAggregations>
: FieldFilter<Schema, Model, Key, Options, WithAggregations>;
} & {
$expr?: (eb: ExpressionBuilder<ToKyselySchema<Schema>, Model>) => OperandExpression<SqlBool>;
$expr?: (
eb: ExpressionBuilder<ToKyselySchema<Schema>, Model>,
context: ExprFilterContext<Schema, Model>,
) => OperandExpression<SqlBool>;
} & {
AND?: OrArray<WhereInput<Schema, Model, Options, ScalarOnly>>;
OR?: WhereInput<Schema, Model, Options, ScalarOnly>[];
Expand Down
2 changes: 1 addition & 1 deletion packages/orm/src/client/crud/dialects/base-dialect.ts
Original file line number Diff line number Diff line change
Expand Up @@ -295,7 +295,7 @@ export abstract class BaseCrudDialect<Schema extends SchemaDef> {

// call expression builder and combine the results
if ('$expr' in _where && typeof _where['$expr'] === 'function') {
result = this.and(result, _where['$expr'](this.eb));
result = this.and(result, _where['$expr'](this.eb, { modelAlias }));
}

return result;
Expand Down
26 changes: 26 additions & 0 deletions tests/e2e/orm/client-api/find.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -1277,4 +1277,30 @@ describe('Client find tests ', () => {
}),
).resolves.toHaveLength(0);
});

it('supports qualified references with the $expr model alias', async () => {
const user1 = await createUser(client, 'yiming@zenstack.dev');
const user2 = await createUser(client, 'yiming@gmail.com');
await createPosts(client, user1.id);
await createPosts(client, user2.id);

await expect(
client.user.findMany({
where: {
$expr: (eb, { modelAlias }) => eb(eb.ref(`${modelAlias}.email`), 'like', '%@zenstack.dev'),
},
}),
).resolves.toHaveLength(1);

// relation filters select the related model under a generated alias
const posts = await client.post.findMany({
where: {
author: {
$expr: (eb, { modelAlias }) => eb(eb.ref(`${modelAlias}.email`), 'like', '%@zenstack.dev'),
},
},
});
expect(posts.length).toBeGreaterThan(0);
expect(posts.every((p) => p.authorId === user1.id)).toBe(true);
});
});
97 changes: 97 additions & 0 deletions tests/regression/test/issue-2863.test.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,97 @@
import { createTestClient } from '@zenstackhq/testtools';
import { describe, expect, it } from 'vitest';

describe('Regression for issue #2863', () => {
const schema = `
model User {
id Int @id @default(autoincrement())
email String @unique
posts Post[]
}

model Post {
id Int @id @default(autoincrement())
title String
authorId Int
author User @relation(fields: [authorId], references: [id])
}
`;

async function setup() {
const db = await createTestClient(schema);
await db.user.create({
data: { email: 'u1@zenstack.dev', posts: { create: [{ title: 'p1' }] } },
});
await db.user.create({
data: { email: 'u2@example.com', posts: { create: [{ title: 'p2' }, { title: 'p3' }] } },
});
return db;
}

it('passes the model alias to top-level $expr', async () => {
const db = await setup();
let alias: string | undefined;

const users = await db.user.findMany({
where: {
$expr: (eb: any, { modelAlias }: { modelAlias: string }) => {
alias = modelAlias;
return eb(eb.ref(`${modelAlias}.email`), 'like', '%@zenstack.dev');
},
},
});

expect(alias).toBe('User');
expect(users.map((u: any) => u.email)).toEqual(['u1@zenstack.dev']);
});

it('passes the generated alias to $expr inside a to-one relation filter', async () => {
const db = await setup();
let alias: string | undefined;

const posts = await db.post.findMany({
where: {
author: {
$expr: (eb: any, { modelAlias }: { modelAlias: string }) => {
alias = modelAlias;
return eb(eb.ref(`${modelAlias}.email`), 'like', '%@zenstack.dev');
},
},
},
});

expect(alias).not.toBe('User');
expect(posts.map((p: any) => p.title)).toEqual(['p1']);
});

it('passes the generated alias to $expr inside a to-many relation filter', async () => {
const db = await setup();

const users = await db.user.findMany({
where: {
posts: {
some: {
$expr: (eb: any, { modelAlias }: { modelAlias: string }) =>
eb(eb.ref(`${modelAlias}.title`), '=', 'p3'),
},
},
},
});

expect(users.map((u: any) => u.email)).toEqual(['u2@example.com']);
});

it('keeps supporting unqualified references inside relation filters', async () => {
const db = await setup();

const posts = await db.post.findMany({
where: {
author: {
$expr: (eb: any) => eb('email', 'like', '%@zenstack.dev'),
},
},
});

expect(posts.map((p: any) => p.title)).toEqual(['p1']);
});
});