diff --git a/packages/orm/src/client/crud/dialects/postgresql.ts b/packages/orm/src/client/crud/dialects/postgresql.ts index 75ee4f35e..5c6866b23 100644 --- a/packages/orm/src/client/crud/dialects/postgresql.ts +++ b/packages/orm/src/client/crud/dialects/postgresql.ts @@ -2,6 +2,8 @@ import { invariant } from '@zenstackhq/common-helpers'; import type { BuiltinType, FieldDef, SchemaDef } from '@zenstackhq/schema'; import Decimal from 'decimal.js'; import { + ValueNode, + type OperationNode, expressionBuilder, sql, type AliasableExpression, @@ -531,6 +533,17 @@ export class PostgresCrudDialect extends LateralJoinDi ) { const leftResolved = this.resolveFieldSqlType(leftFieldDef); const rightResolved = this.resolveFieldSqlType(rightFieldDef); + + // Fast path for comparing a column with a @db.* native type against a bound value (e.g. + // `userId == auth().id`). Casting the column would defeat index usage, so the value side is + // handled instead: PostgreSQL infers the parameter type from the column, and for uuid the + // value is validated up front so a malformed one yields a constant result rather than an + // "invalid input syntax for type uuid" error. + const valueResult = this.tryBuildNativeTypeValueComparison(left, leftResolved, op, right, rightResolved); + if (valueResult) { + return valueResult; + } + // If the resolved SQL types differ and at least one side carries a @db.* native type override, // cast that side back to its base ZModel SQL type so PostgreSQL doesn't reject the comparison // (e.g. "operator does not exist: uuid = text"). @@ -548,6 +561,65 @@ export class PostgresCrudDialect extends LateralJoinDi return super.buildComparison(left, leftFieldDef, op, right, rightFieldDef); } + // native SQL types that accept any text input, so a bound string value never fails to parse + private static readonly textLikeSqlTypes = new Set(['text', 'varchar', 'bpchar', 'citext']); + + // uuid input formats accepted by PostgreSQL (canonical 8-4-4-4-12 or 32 hex digits). This is a + // pure format check: unlike RFC 4122 validators it doesn't require specific version/variant bits, + // since PostgreSQL stores any 128-bit value (e.g. `00000000-0000-0000-0000-000000000001`). + private static readonly uuidFormatRegex = + /^(?:[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}|[0-9a-f]{32})$/i; + + private tryBuildNativeTypeValueComparison( + left: Expression, + leftResolved: ReturnType, + op: string, + right: Expression, + rightResolved: ReturnType, + ): Expression | undefined { + let valueNode: OperationNode; + let columnResolved: typeof leftResolved; + const leftNode = left.toOperationNode(); + const rightNode = right.toOperationNode(); + if (ValueNode.is(rightNode) && !ValueNode.is(leftNode)) { + valueNode = rightNode; + columnResolved = leftResolved; + } else if (ValueNode.is(leftNode) && !ValueNode.is(rightNode)) { + valueNode = leftNode; + columnResolved = rightResolved; + } else { + return undefined; + } + + if (!columnResolved.hasDbOverride || !columnResolved.sqlType) { + return undefined; + } + + if (PostgresCrudDialect.textLikeSqlTypes.has(columnResolved.sqlType)) { + // text-like column, any string value is valid input + return this.eb(left, op as any, right) as Expression; + } + + if (columnResolved.sqlType === 'uuid' && (op === '=' || op === '!=')) { + const value = (valueNode as ValueNode).value; + if (typeof value === 'string' && PostgresCrudDialect.uuidFormatRegex.test(value)) { + // well-formed uuid, compare natively without casting the column + return this.eb(left, op as any, right) as Expression; + } else if (op === '=') { + // malformed uuid can never equal a uuid column + return this.eb.lit(false) as unknown as Expression; + } else { + // malformed uuid differs from every non-null uuid; a null column must not match, matching + // the SQL semantics of `col != value` (null when col is null) + const column = valueNode === rightNode ? left : right; + return this.eb(column, 'is not', null) as Expression; + } + } + + // other native types fall back to casting the column + return undefined; + } + override getStringCasingBehavior() { // Postgres `LIKE` is case-sensitive, `ILIKE` is case-insensitive return { supportsILike: true, likeCaseSensitive: true }; diff --git a/packages/plugins/policy/src/expression-transformer.ts b/packages/plugins/policy/src/expression-transformer.ts index 9a3761365..777fde082 100644 --- a/packages/plugins/policy/src/expression-transformer.ts +++ b/packages/plugins/policy/src/expression-transformer.ts @@ -41,6 +41,7 @@ import { SelectionNode, SelectQueryNode, TableNode, + UnaryOperationNode, ValueListNode, ValueNode, WhereNode, @@ -91,6 +92,12 @@ export type ExpressionTransformerContext = { */ memberSelect?: SelectionNode; + /** + * In case of transforming a collection predicate's LHS, wraps the innermost relation subquery with + * `exists` or `not exists` so the predicate is evaluated as a semi-join instead of an aggregate + */ + memberExists?: 'exists' | 'not exists'; + /** * In case of transforming a collection predicate's LHS, the table alias to use for the innermost * relation (the one the predicate filter is compiled against) @@ -206,14 +213,28 @@ export class ExpressionTransformer { if (!fieldDef.relation) { return this.createColumnRef(expr.field, context); } else { - const { memberFilter, memberSelect, memberAlias, ...restContext } = context; + const { memberFilter, memberSelect, memberExists, memberAlias, ...restContext } = context; const relation = this.transformRelationAccess(expr.field, fieldDef.type, restContext, memberAlias); - return { - ...relation, - where: this.mergeWhere(relation.where, memberFilter), - selections: memberSelect ? [memberSelect] : relation.selections, - }; + return this.finalizeMemberSubquery( + { + ...relation, + where: this.mergeWhere(relation.where, memberFilter), + selections: memberSelect ? [memberSelect] : relation.selections, + }, + memberExists, + ); + } + } + + // wraps the innermost collection-predicate subquery with `exists`/`not exists` if requested + private finalizeMemberSubquery( + node: SelectQueryNode, + memberExists: 'exists' | 'not exists' | undefined, + ): OperationNode { + if (!memberExists) { + return node; } + return UnaryOperationNode.create(OperatorNode.create(memberExists), node); } private mergeWhere(where: WhereNode | undefined, memberFilter: OperationNode | undefined) { @@ -448,17 +469,23 @@ export class ExpressionTransformer { predicateFilter = logicalNot(this.dialect, predicateFilter); } - const count = FunctionNode.create('count', [ValueNode.createImmediate(1)]); - - const predicateResult = match(expr.op) - .with('?', () => BinaryOperationNode.create(count, OperatorNode.create('>'), ValueNode.createImmediate(0))) - .with('!', () => BinaryOperationNode.create(count, OperatorNode.create('='), ValueNode.createImmediate(0))) - .with('^', () => BinaryOperationNode.create(count, OperatorNode.create('='), ValueNode.createImmediate(0))) + // `?` (some) => exists(select 1 ... where filter) + // `!` (all) => not exists(select 1 ... where not filter) + // `^` (none) => not exists(select 1 ... where filter) + // `exists` lets the database plan the predicate as a semi-join that can use indexes and + // stop at the first match, unlike a correlated `count(1) > 0` aggregate + const memberExists = match(expr.op) + .with('?', () => 'exists' as const) + .with('!', () => 'not exists' as const) + .with('^', () => 'not exists' as const) .exhaustive(); return this.transform(expr.left, { ...context, - memberSelect: SelectionNode.create(AliasNode.create(predicateResult, IdentifierNode.create('_'))), + memberSelect: SelectionNode.create( + AliasNode.create(ValueNode.createImmediate(1), IdentifierNode.create('_')), + ), + memberExists, memberFilter: predicateFilter, memberAlias, }); @@ -743,7 +770,7 @@ export class ExpressionTransformer { let receiver: OperationNode; let receiverAlias: string; let startType: string | undefined; - const { memberFilter, memberSelect, memberAlias, ...restContext } = context; + const { memberFilter, memberSelect, memberExists, memberAlias, ...restContext } = context; if (ExpressionUtils.isThis(expr.receiver)) { if (expr.members.length === 1) { @@ -837,7 +864,7 @@ export class ExpressionTransformer { currAlias = alias; } - let currNode: SelectQueryNode | ColumnNode | ReferenceNode | undefined = undefined; + let currNode: OperationNode | undefined = undefined; for (let i = members.length - 1; i >= 0; i--) { const member = members[i]!; @@ -858,19 +885,23 @@ export class ExpressionTransformer { ); if (currNode) { - currNode = { + const outer: SelectQueryNode = { ...relation, selections: [ SelectionNode.create(AliasNode.create(currNode, IdentifierNode.create(members[i + 1]!))), ], }; + currNode = outer; } else { // inner most member, merge with member filter from the context - currNode = { - ...relation, - where: this.mergeWhere(relation.where, memberFilter), - selections: memberSelect ? [memberSelect] : relation.selections, - }; + currNode = this.finalizeMemberSubquery( + { + ...relation, + where: this.mergeWhere(relation.where, memberFilter), + selections: memberSelect ? [memberSelect] : relation.selections, + }, + memberExists, + ); } } else { invariant(i === members.length - 1, 'plain field access must be the last segment'); @@ -1139,6 +1170,12 @@ export class ExpressionTransformer { // receiver is the first hop, so native-type info (@db.*) on the terminal field is // available for casting in buildComparison. return walkRelationChain(model, [expr.receiver.field, ...expr.members]); + } else if (this.isAuthCall(expr.receiver)) { + // `auth().<...>.field` chain rooted at the auth model. Resolving the terminal + // field lets buildComparison see matching native types on both sides (e.g. + // `userId == auth().id` with both `@db.Uuid`) and skip casting the column, + // which would otherwise defeat index usage. + return walkRelationChain(this.authType, expr.members); } } return undefined; diff --git a/tests/regression/test/issue-2851.test.ts b/tests/regression/test/issue-2851.test.ts new file mode 100644 index 000000000..9b92e5d32 --- /dev/null +++ b/tests/regression/test/issue-2851.test.ts @@ -0,0 +1,332 @@ +import { createPolicyTestClient } from '@zenstackhq/testtools'; +import { randomUUID } from 'node:crypto'; +import { describe, expect, it } from 'vitest'; + +describe('Regression for issue #2851', () => { + async function createClient(schema: string) { + const sqls: string[] = []; + const db = await createPolicyTestClient(schema, { + provider: 'postgresql', + usePrismaPush: true, + log: (event) => { + if (event.level === 'query') { + sqls.push(event.query.sql); + } + }, + }); + return { db, sqls }; + } + + it('does not cast a uuid column when comparing with a uuid auth() field', async () => { + const { db, sqls } = await createClient( + ` +model User { + id String @id @default(uuid()) @db.Uuid + memberships Membership[] + @@allow('all', true) +} + +model Membership { + id String @id @default(uuid()) @db.Uuid + userID String @db.Uuid + user User @relation(fields: [userID], references: [id]) + @@allow('read', userID == auth().id) +} + `, + ); + + const rawDb = db.$unuseAll(); + const user1 = await rawDb.user.create({ data: {} }); + const user2 = await rawDb.user.create({ data: {} }); + await rawDb.membership.create({ data: { userID: user1.id } }); + await rawDb.membership.create({ data: { userID: user2.id } }); + + sqls.length = 0; + const result = await db.$setAuth({ id: user1.id }).membership.findMany(); + expect(result).toHaveLength(1); + expect(result[0]!.userID).toBe(user1.id); + + const query = sqls.find((sql) => sql.includes('from "public"."Membership"')); + expect(query).toBeDefined(); + expect(query).not.toContain('cast('); + expect(query).toContain('"Membership"."userID" = $1'); + }); + + it('does not cast a uuid column when comparing a relation with auth()', async () => { + const { db, sqls } = await createClient( + ` +model User { + id String @id @default(uuid()) @db.Uuid + memberships Membership[] + @@allow('all', true) +} + +model Membership { + id String @id @default(uuid()) @db.Uuid + userID String @db.Uuid + user User @relation(fields: [userID], references: [id]) + @@allow('read', user == auth()) +} + `, + ); + + const rawDb = db.$unuseAll(); + const user1 = await rawDb.user.create({ data: {} }); + const user2 = await rawDb.user.create({ data: {} }); + await rawDb.membership.create({ data: { userID: user1.id } }); + await rawDb.membership.create({ data: { userID: user2.id } }); + + sqls.length = 0; + const result = await db.$setAuth({ id: user1.id }).membership.findMany(); + expect(result).toHaveLength(1); + expect(result[0]!.userID).toBe(user1.id); + + const query = sqls.find((sql) => sql.includes('from "public"."Membership"')); + expect(query).toBeDefined(); + expect(query).not.toContain('as text'); + }); + + it('still works when comparing a plain string column with a uuid auth() field', async () => { + const { db } = await createClient( + ` +model User { + id String @id @default(uuid()) @db.Uuid + @@allow('all', true) +} + +model Item { + id String @id @default(uuid()) @db.Uuid + ownerId String + @@allow('all', ownerId == auth().id) +} + `, + ); + + const rawDb = db.$unuseAll(); + const uid = randomUUID(); + await rawDb.user.create({ data: { id: uid } }); + await rawDb.item.create({ data: { ownerId: uid } }); + await rawDb.item.create({ data: { ownerId: randomUUID() } }); + + const result = await db.$setAuth({ id: uid }).item.findMany(); + expect(result).toHaveLength(1); + expect(result[0]!.ownerId).toBe(uid); + }); + + it('does not cast a uuid column when comparing with a plain string auth() field', async () => { + const { db, sqls } = await createClient( + ` +model User { + id String @id @default(uuid()) + @@allow('all', true) +} + +model Item { + id String @id @default(uuid()) @db.Uuid + ownerId String @db.Uuid + @@allow('all', ownerId == auth().id) +} + `, + ); + + const rawDb = db.$unuseAll(); + const uid = randomUUID(); + await rawDb.user.create({ data: { id: uid } }); + await rawDb.item.create({ data: { ownerId: uid } }); + await rawDb.item.create({ data: { ownerId: randomUUID() } }); + + sqls.length = 0; + const result = await db.$setAuth({ id: uid }).item.findMany(); + expect(result).toHaveLength(1); + expect(result[0]!.ownerId).toBe(uid); + + // the auth value is a bound parameter, PostgreSQL infers its type from the column + const query = sqls.find((sql) => sql.includes('from "public"."Item"')); + expect(query).not.toContain('cast('); + expect(query).toContain('"Item"."ownerId" = $1'); + }); + + it('denies instead of failing when auth() id is not a well-formed uuid', async () => { + const { db } = await createClient( + ` +model User { + id String @id @default(uuid()) + @@allow('all', true) +} + +model Item { + id String @id @default(uuid()) @db.Uuid + ownerId String @db.Uuid + @@allow('read', ownerId == auth().id) +} + +model Note { + id String @id @default(uuid()) @db.Uuid + ownerId String? @db.Uuid + @@allow('read', ownerId != auth().id) +} + `, + ); + + const rawDb = db.$unuseAll(); + const uid = randomUUID(); + await rawDb.item.create({ data: { ownerId: uid } }); + const note = await rawDb.note.create({ data: { ownerId: uid } }); + // a note with null owner: `ownerId != x` is null in SQL, so it must never match + await rawDb.note.create({ data: { ownerId: null } }); + + const badAuthDb = db.$setAuth({ id: 'not-a-uuid' }); + // `==` against a malformed uuid is always false + await expect(badAuthDb.item.findMany()).resolves.toHaveLength(0); + // `!=` against a malformed uuid is true for non-null columns only + const notes = await badAuthDb.note.findMany(); + expect(notes).toHaveLength(1); + expect(notes[0]!.id).toBe(note.id); + + // well-formed uuids are still compared natively regardless of casing or dashes + await expect(db.$setAuth({ id: uid.toUpperCase() }).item.findMany()).resolves.toHaveLength(1); + await expect(db.$setAuth({ id: uid.replace(/-/g, '') }).item.findMany()).resolves.toHaveLength(1); + + // uuids that are not RFC 4122 compliant (no version/variant bits) are still valid for PostgreSQL + const seedId = '00000000-0000-0000-0000-000000000001'; + await rawDb.item.create({ data: { ownerId: seedId } }); + await expect(db.$setAuth({ id: seedId }).item.findMany()).resolves.toHaveLength(1); + }); + + it('does not cast a varchar column when comparing with an auth() field', async () => { + const { db, sqls } = await createClient( + ` +model User { + id String @id @default(uuid()) + @@allow('all', true) +} + +model Item { + id String @id @default(uuid()) + ownerId String @db.VarChar(64) + @@allow('read', ownerId == auth().id) +} + `, + ); + + const rawDb = db.$unuseAll(); + const uid = randomUUID(); + await rawDb.item.create({ data: { ownerId: uid } }); + await rawDb.item.create({ data: { ownerId: randomUUID() } }); + + sqls.length = 0; + const result = await db.$setAuth({ id: uid }).item.findMany(); + expect(result).toHaveLength(1); + + const query = sqls.find((sql) => sql.includes('from "public"."Item"')); + expect(query).not.toContain('cast('); + expect(query).toContain('"Item"."ownerId" = $1'); + }); + + describe('collection predicates compile to EXISTS', () => { + const schema = ` +model User { + id String @id @default(uuid()) @db.Uuid + memberships Membership[] + @@allow('all', true) +} + +model Team { + id String @id @default(uuid()) @db.Uuid + name String + members Membership[] + projects Project[] + @@allow('create', true) + @@allow('read', members?[userID == auth().id]) + @@allow('update', members![userID == auth().id]) + @@allow('delete', members^[userID == auth().id]) +} + +model Membership { + id String @id @default(uuid()) @db.Uuid + teamID String @db.Uuid + team Team @relation(fields: [teamID], references: [id], onDelete: Cascade) + userID String @db.Uuid + user User @relation(fields: [userID], references: [id]) + @@allow('all', true) +} + +model Project { + id String @id @default(uuid()) @db.Uuid + teamID String @db.Uuid + team Team @relation(fields: [teamID], references: [id], onDelete: Cascade) + @@allow('read', team.members?[userID == auth().id]) +} +`; + + async function setup() { + const { db, sqls } = await createClient(schema); + const rawDb = db.$unuseAll(); + const user1 = await rawDb.user.create({ data: {} }); + const user2 = await rawDb.user.create({ data: {} }); + // team1: only user1; team2: user1 and user2; team3: only user2 + const team1 = await rawDb.team.create({ + data: { name: 't1', members: { create: [{ userID: user1.id }] } }, + }); + const team2 = await rawDb.team.create({ + data: { name: 't2', members: { create: [{ userID: user1.id }, { userID: user2.id }] } }, + }); + const team3 = await rawDb.team.create({ + data: { name: 't3', members: { create: [{ userID: user2.id }] } }, + }); + await rawDb.project.create({ data: { teamID: team1.id } }); + await rawDb.project.create({ data: { teamID: team3.id } }); + return { db, rawDb, sqls, user1, user2, team1, team2, team3 }; + } + + it('uses EXISTS for `?` (some)', async () => { + const { db, sqls, user1, team1, team2 } = await setup(); + sqls.length = 0; + const teams = await db.$setAuth({ id: user1.id }).team.findMany(); + expect(teams.map((t: any) => t.id).sort()).toEqual([team1.id, team2.id].sort()); + + const query = sqls.find((sql) => sql.includes('from "public"."Team"')); + expect(query).toContain('exists (select 1'); + expect(query).not.toContain('count('); + expect(query).not.toContain('cast('); + }); + + it('uses EXISTS for nested relation chains', async () => { + const { db, sqls, user1, team1 } = await setup(); + sqls.length = 0; + const projects = await db.$setAuth({ id: user1.id }).project.findMany(); + expect(projects).toHaveLength(1); + expect(projects[0]!.teamID).toBe(team1.id); + + const query = sqls.find((sql) => sql.includes('from "public"."Project"')); + expect(query).toContain('exists (select 1'); + expect(query).not.toContain('count('); + }); + + it('uses NOT EXISTS for `!` (all)', async () => { + const { db, sqls, user1, team1 } = await setup(); + sqls.length = 0; + // only team1 has all members equal to user1 + const result = await db.$setAuth({ id: user1.id }).team.updateMany({ data: { name: 'updated' } }); + expect(result.count).toBe(1); + expect((await db.$unuseAll().team.findUnique({ where: { id: team1.id } }))?.name).toBe('updated'); + + const query = sqls.find((sql) => sql.startsWith('update "public"."Team"')); + expect(query).toContain('not exists (select 1'); + expect(query).not.toContain('count('); + }); + + it('uses NOT EXISTS for `^` (none)', async () => { + const { db, sqls, user1, team3 } = await setup(); + sqls.length = 0; + // only team3 has no member equal to user1 + const result = await db.$setAuth({ id: user1.id }).team.deleteMany(); + expect(result.count).toBe(1); + expect(await db.$unuseAll().team.findUnique({ where: { id: team3.id } })).toBeNull(); + + const query = sqls.find((sql) => sql.startsWith('delete from "public"."Team"')); + expect(query).toContain('not exists (select 1'); + expect(query).not.toContain('count('); + }); + }); +});