Skip to content
Merged
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
72 changes: 72 additions & 0 deletions packages/orm/src/client/crud/dialects/postgresql.ts
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -531,6 +533,17 @@ export class PostgresCrudDialect<Schema extends SchemaDef> 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").
Expand All @@ -548,6 +561,65 @@ export class PostgresCrudDialect<Schema extends SchemaDef> 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<unknown>,
leftResolved: ReturnType<typeof this.resolveFieldSqlType>,
op: string,
right: Expression<unknown>,
rightResolved: ReturnType<typeof this.resolveFieldSqlType>,
): Expression<SqlBool> | 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<SqlBool>;
}

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<SqlBool>;
} else if (op === '=') {
// malformed uuid can never equal a uuid column
return this.eb.lit(false) as unknown as Expression<SqlBool>;
} 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<SqlBool>;
}
}

// 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 };
Expand Down
79 changes: 58 additions & 21 deletions packages/plugins/policy/src/expression-transformer.ts
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,7 @@ import {
SelectionNode,
SelectQueryNode,
TableNode,
UnaryOperationNode,
ValueListNode,
ValueNode,
WhereNode,
Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -206,14 +213,28 @@ export class ExpressionTransformer<Schema extends SchemaDef> {
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) {
Expand Down Expand Up @@ -448,17 +469,23 @@ export class ExpressionTransformer<Schema extends SchemaDef> {
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,
});
Expand Down Expand Up @@ -743,7 +770,7 @@ export class ExpressionTransformer<Schema extends SchemaDef> {
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) {
Expand Down Expand Up @@ -837,7 +864,7 @@ export class ExpressionTransformer<Schema extends SchemaDef> {
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]!;
Expand All @@ -858,19 +885,23 @@ export class ExpressionTransformer<Schema extends SchemaDef> {
);

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');
Expand Down Expand Up @@ -1139,6 +1170,12 @@ export class ExpressionTransformer<Schema extends SchemaDef> {
// 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;
Expand Down
Loading
Loading