Skip to content
Merged
120 changes: 120 additions & 0 deletions packages/cli/test/ts-schema-gen.test.ts
Original file line number Diff line number Diff line change
Expand Up @@ -852,6 +852,126 @@ model Post {
});
});

it('supports implicit conversions from enums to arrays', async () => {
const { schema } = await generateTsSchema(`
enum PostStatus {
DRAFT
ACTIVE
CANCELLED
}

model User {
id Int @id @default(autoincrement())
}

model Post {
id String @id
status String

@@validate(status in PostStatus)
}
`);

expect(schema.models['Post']?.attributes).toMatchObject([
{
name: '@@validate',
args: [
{
name: 'value',
value: {
kind: 'binary',
op: 'in',
left: {
kind: 'field',
field: 'status',
},
right: {
kind: 'array',
type: 'PostStatus',
items: [
{
kind: 'literal',
value: 'DRAFT',
},
{
kind: 'literal',
value: 'ACTIVE',
},
{
kind: 'literal',
value: 'CANCELLED',
},
],
},
binding: undefined,
},
},
],
},
]);
});

it('emits enum names (not @map values) when converting enums to arrays', async () => {
// enum values are represented by their names at runtime (TS values, auth(), Zod);
// the ORM's name mapper translates them to `@map`-ed values at SQL execution time
const { schema } = await generateTsSchema(`
enum PostStatus {
DRAFT @map('draft')
ACTIVE @map('active')
CANCELLED @map('cancelled')
}

model User {
id Int @id @default(autoincrement())
}

model Post {
id String @id
status String

@@validate(status in PostStatus)
}
`);

expect(schema.models['Post']?.attributes).toMatchObject([
{
name: '@@validate',
args: [
{
name: 'value',
value: {
kind: 'binary',
op: 'in',
left: {
kind: 'field',
field: 'status',
},
right: {
kind: 'array',
type: 'PostStatus',
items: [
{
kind: 'literal',
value: 'DRAFT',
},
{
kind: 'literal',
value: 'ACTIVE',
},
{
kind: 'literal',
value: 'CANCELLED',
},
],
},
binding: undefined,
},
},
],
},
]);
});

it('supports @@strict for type defs', async () => {
const { schema } = await generateTsSchema(`
model User {
Expand Down
4 changes: 2 additions & 2 deletions packages/language/src/generated/ast.ts
Original file line number Diff line number Diff line change
Expand Up @@ -902,7 +902,7 @@ export function isReferenceExpr(item: unknown): item is ReferenceExpr {
return reflection.isInstance(item, ReferenceExpr.$type);
}

export type ReferenceTarget = CollectionPredicateBinding | DataField | EnumField | FunctionParam;
export type ReferenceTarget = CollectionPredicateBinding | DataField | Enum | EnumField | FunctionParam;

export const ReferenceTarget = {
$type: 'ReferenceTarget'
Expand Down Expand Up @@ -1427,7 +1427,7 @@ export class ZModelAstReflection extends langium.AbstractAstReflection {
name: Enum.name
}
},
superTypes: [AbstractDeclaration.$type, TypeDeclaration.$type]
superTypes: [AbstractDeclaration.$type, ReferenceTarget.$type, TypeDeclaration.$type]
},
EnumField: {
name: EnumField.$type,
Expand Down
6 changes: 6 additions & 0 deletions packages/language/src/generated/grammar.ts
Original file line number Diff line number Diff line change
Expand Up @@ -4008,6 +4008,12 @@ export const ZModelGrammar = (): Grammar => loadedZModelGrammar ?? (loadedZModel
"typeRef": {
"$ref": "#/rules@30"
}
},
{
"$type": "SimpleType",
"typeRef": {
"$ref": "#/rules@46"
}
}
]
}
Expand Down
14 changes: 14 additions & 0 deletions packages/language/src/validators/expression-validator.ts
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ import {
Expression,
isArrayExpr,
isCollectionPredicateBinding,
isDataFieldAttribute,
isDataModel,
isDataModelAttribute,
isEnum,
Expand Down Expand Up @@ -87,6 +88,11 @@ export default class ExpressionValidator implements AstValidator<Expression> {
node: expr,
});
}
if (isEnum(expr.target.ref) && !this.isInPolicyOrValidationAttribute(expr)) {
accept('error', 'Enum reference can only be used with policy and validation attributes', {
node: expr,
});
}
}

private validateMemberAccessExpr(expr: MemberAccessExpr, accept: ValidationAcceptor) {
Expand Down Expand Up @@ -281,6 +287,14 @@ export default class ExpressionValidator implements AstValidator<Expression> {
return findUpAst(node, (n) => isDataModelAttribute(n) && n.decl.$refText === '@@validate');
}

private isInPolicyOrValidationAttribute(node: AstNode) {
const attrs = ['@allow', '@@allow', '@deny', '@@deny', '@@validate'];
return findUpAst(
node,
(n) => (isDataModelAttribute(n) || isDataFieldAttribute(n)) && attrs.includes(n.decl.$refText),
);
}

private isNotModelFieldExpr(expr: Expression): boolean {
return (
// literal
Expand Down
6 changes: 6 additions & 0 deletions packages/language/src/zmodel-linker.ts
Original file line number Diff line number Diff line change
Expand Up @@ -265,6 +265,12 @@ export class ZModelLinker extends DefaultLinker {
} else if (isDataField(target) || isFunctionParam(target)) {
// other references are resolved to their declared type
this.resolveToDeclaredType(node, target.type);
} else if (isEnum(target)) {
node.$resolvedType = {
decl: target,
array: true,
nullable: false,
};
}
}
}
Expand Down
2 changes: 1 addition & 1 deletion packages/language/src/zmodel.langium
Original file line number Diff line number Diff line change
Expand Up @@ -66,7 +66,7 @@ ConfigArrayExpr:
ConfigExpr:
LiteralExpr | InvocationExpr | ConfigArrayExpr;

type ReferenceTarget = FunctionParam | DataField | EnumField | CollectionPredicateBinding;
type ReferenceTarget = FunctionParam | DataField | EnumField | CollectionPredicateBinding | Enum;

ThisExpr:
value='this';
Expand Down
92 changes: 92 additions & 0 deletions packages/language/test/enum.test.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,92 @@
import { describe, it } from 'vitest';
import { loadSchema, loadSchemaWithError } from './utils';

describe('Enum tests', () => {
describe('implicit array conversions', () => {
it('supports usage with policy and validation attributes', async () => {
await loadSchema(`
datasource db {
provider = 'sqlite'
url = 'file:./dev.db'
}

enum PostStatus {
DRAFT
ACTIVE
CANCELLED
}

model User {
id Int @id @default(autoincrement())
}

model Post {
id String @id
status String @allow('update', status in PostStatus) @deny('read', !(status in PostStatus))

@@allow('read', status in PostStatus)
@@deny('read', status in PostStatus)
@@validate(status in PostStatus)
}
`);
});

it('rejects usage with non-policy and non-validation attributes', async () => {
await loadSchemaWithError(
`
datasource db {
provider = 'sqlite'
url = 'file:./dev.db'
}

enum PostStatus {
DRAFT
ACTIVE
CANCELLED
}

model User {
id Int @id @default(autoincrement())
}

model Post {
id String @id
status PostStatus[] @default(PostStatus)
}
`,
/Enum reference can only be used with policy and validation attributes/,
);
});

it('rejects usage in other contexts', async () => {
await loadSchemaWithError(
`
datasource db {
provider = 'sqlite'
url = 'file:./dev.db'
}

enum PostStatus {
DRAFT
ACTIVE
CANCELLED
}

model User {
id Int @id @default(autoincrement())
}

model Post {
id String @id
status String
}

function Test(status: String): Void {
status in PostStatus
}
`,
/Enum reference can only be used with policy and validation attributes/,
);
});
});
});
17 changes: 16 additions & 1 deletion packages/orm/src/client/executor/name-mapper.ts
Original file line number Diff line number Diff line change
Expand Up @@ -283,7 +283,9 @@ export class QueryNameMapper extends OperationNodeTransformer {
if (
ReferenceNode.is(node.leftOperand) &&
ColumnNode.is(node.leftOperand.column) &&
(ValueNode.is(node.rightOperand) || PrimitiveValueListNode.is(node.rightOperand))
(ValueNode.is(node.rightOperand) ||
PrimitiveValueListNode.is(node.rightOperand) ||
ValueListNode.is(node.rightOperand))
) {
const columnNode = node.leftOperand.column;

Expand Down Expand Up @@ -311,6 +313,19 @@ export class QueryNameMapper extends OperationNodeTransformer {
valueNode.values,
),
);
} else if (ValueListNode.is(valueNode)) {
// list value: column IN (EnumValue, EnumValue2)
resultValue = ValueListNode.create(
valueNode.values.map((v) =>
ValueNode.is(v)
? (this.processEnumMappingForValue(
resolvedScope.model!,
columnNode,
v,
) as OperationNode)
: v,
),
);
}

return super.transformBinaryOperation(
Expand Down
24 changes: 22 additions & 2 deletions packages/plugins/policy/src/expression-transformer.ts
Original file line number Diff line number Diff line change
Expand Up @@ -274,14 +274,29 @@ export class ExpressionTransformer<Schema extends SchemaDef> {

const { normalizedLeft, normalizedRight } = this.normalizeBinaryOperationOperands(expr, context);
const left = this.transform(normalizedLeft, context);
const right = this.transform(normalizedRight, context);

if (op === 'in') {
if (this.isNullNode(left)) {
return this.transformValue(false, 'Boolean');
} else {
if (this.isLiteralArray(normalizedRight)) {
// `in` with a list of literal values, e.g. `field in [1, 2, 3]` or
// `field in [ENUM_A, ENUM_B]`: emit a plain SQL `IN (...)` list on every
// dialect (instead of a dialect-specific array value) so that it stays
// index-friendly and the name mapper can translate `@map`-ed enum values
if (normalizedRight.items.length === 0) {
return this.transformValue(false, 'Boolean');
}
return BinaryOperationNode.create(
left,
OperatorNode.create('in'),
ValueListNode.create(normalizedRight.items.map((item) => this.transform(item, context))),
);
}

const right = this.transform(normalizedRight, context);
if (ValueListNode.is(right)) {
// simple `in` operator with a list of values, e.g. `field in [1, 2, 3]`
// simple `in` operator with a list of values, e.g. `field in [auth().x, 2]`
return BinaryOperationNode.create(left, OperatorNode.create('in'), right);
} else {
// array contains
Expand Down Expand Up @@ -317,6 +332,7 @@ export class ExpressionTransformer<Schema extends SchemaDef> {
}
}

const right = this.transform(normalizedRight, context);
if (this.isNullNode(right)) {
return this.transformNullCheck(left, expr.op);
} else if (this.isNullNode(left)) {
Expand Down Expand Up @@ -358,6 +374,10 @@ export class ExpressionTransformer<Schema extends SchemaDef> {
}
}

private isLiteralArray(expr: Expression): expr is ArrayExpression {
return expr.kind === 'array' && expr.items.every((item) => item.kind === 'literal');
}

private normalizeBinaryOperationOperands(expr: BinaryExpression, context: ExpressionTransformerContext) {
// If relation fields are used directly in a comparison, normalize both sides to the
// first id field (used for multiple). This applies whether the relation is SQL-backed
Expand Down
Loading
Loading