From e941932d5e75810cc4bdcc13943ba9a2be692001 Mon Sep 17 00:00:00 2001 From: kjw142857 Date: Tue, 11 Aug 2026 06:09:53 +0800 Subject: [PATCH 1/2] tighten selector criteria --- .../__tests__/switchStatements.test.ts | 23 ++++++++++++++++++- src/types/checker/statements.ts | 4 ++-- 2 files changed, 24 insertions(+), 3 deletions(-) diff --git a/src/types/checker/__tests__/switchStatements.test.ts b/src/types/checker/__tests__/switchStatements.test.ts index 8889e5bf..6344ebce 100644 --- a/src/types/checker/__tests__/switchStatements.test.ts +++ b/src/types/checker/__tests__/switchStatements.test.ts @@ -1,6 +1,6 @@ import { check } from '..' import { parse } from '../../ast' -import { IncompatibleTypesError, TypeCheckerError } from '../../errors' +import { IncompatibleTypesError, SelectorTypeNotAllowedError, TypeCheckerError } from '../../errors' import { Type } from '../../types/type' const createProgram = (statement: string) => ` @@ -27,6 +27,27 @@ const testcases: { `, result: { type: null, errors: [] } }, + { + input: ` + String selector = "Tuesday"; + switch(selector) { + case "Tuesday": { + selector = "Wednesday"; + } + default: + } + `, + result: { type: null, errors: [] } + }, + { + input: ` + Boolean selector = true; + switch(selector) { + default: {} + } + `, + result: { type: null, errors: [new SelectorTypeNotAllowedError()] } + }, { input: ` int selector = 1; diff --git a/src/types/checker/statements.ts b/src/types/checker/statements.ts index ef812dad..fe88fcef 100644 --- a/src/types/checker/statements.ts +++ b/src/types/checker/statements.ts @@ -12,7 +12,7 @@ import { isPrimitiveIntegralType, isPrimitiveLongType, isReferenceBooleanType, - isReferenceType + isStringType } from '../types/utils' export const checkDoExpression = ( @@ -28,7 +28,7 @@ export const checkSwitchExpression = ( location: Location ): null | TypeCheckerError => { if (isPrimitiveIntegralType(expressionType) && !isPrimitiveLongType(expressionType)) return null - if (isReferenceType(expressionType)) return null + if (isStringType(expressionType)) return null return new SelectorTypeNotAllowedError(location) } From 89ccf3067db10b398fdd1215b3c15d841f2ff6fd Mon Sep 17 00:00:00 2001 From: kjw142857 Date: Tue, 11 Aug 2026 07:49:09 +0800 Subject: [PATCH 2/2] add enum support --- src/compiler/code-generator.ts | 28 +++- src/compiler/compiler-utils.ts | 3 +- .../__tests__/switchStatements.test.ts | 25 ++++ src/types/checker/environment.ts | 4 +- src/types/checker/index.ts | 108 +++++++++++++-- src/types/checker/prechecks.ts | 125 +++++++++++++++++- src/types/checker/statements.ts | 2 + src/types/types/classes.ts | 2 + 8 files changed, 278 insertions(+), 19 deletions(-) diff --git a/src/compiler/code-generator.ts b/src/compiler/code-generator.ts index 773bde90..d4d5d101 100644 --- a/src/compiler/code-generator.ts +++ b/src/compiler/code-generator.ts @@ -1,4 +1,5 @@ import { OPCODE } from '../ClassFile/constants/instructions' +import { ACCESS_FLAGS } from '../ClassFile/types' import { ExceptionHandler, AttributeInfo } from '../ClassFile/types/attributes' import { FIELD_FLAGS } from '../ClassFile/types/fields' import { METHOD_FLAGS } from '../ClassFile/types/methods' @@ -1375,6 +1376,27 @@ const codeGenerators: { [type: string]: (node: Node, cg: CodeGenerator) => Compi const { stackSize: exprStackSize, resultType } = compile(expression, cg) let maxStack = exprStackSize + // If the expression is an enum type, invoke ordinal() to convert to int and then continue + let _resultType = resultType + if (_resultType && _resultType.startsWith('L') && _resultType !== 'Ljava/lang/String;') { + const clean = _resultType.replace(/^L|;$/g, '') + try { + const classInfo = cg.symbolTable.queryClass(clean) + if (classInfo.accessFlags & ACCESS_FLAGS.ACC_ENUM) { + // call java.lang.Enum.ordinal() (returns int) + cg.code.push( + OPCODE.INVOKEVIRTUAL, + 0, + cg.constantPoolManager.indexMethodrefInfo('java/lang/Enum', 'ordinal', '()I') + ) + _resultType = 'I' + maxStack = Math.max(maxStack, exprStackSize + 1) + } + } catch (e) { + // ignore: not a known class + } + } + const caseLabels: Label[] = cases.map(() => cg.generateNewLabel()) const defaultLabel = cg.generateNewLabel() const endLabel = cg.generateNewLabel() @@ -1382,7 +1404,7 @@ const codeGenerators: { [type: string]: (node: Node, cg: CodeGenerator) => Compi // Track the switch statement's end label cg.switchLabels.push(endLabel) - if (['I', 'B', 'S', 'C'].includes(resultType)) { + if (['I', 'B', 'S', 'C'].includes(_resultType)) { const caseValues: number[] = [] const caseLabelMap: Map = new Map() let hasDefault = false @@ -1556,7 +1578,7 @@ const codeGenerators: { [type: string]: (node: Node, cg: CodeGenerator) => Compi } endLabel.offset = cg.code.length - } else if (resultType === 'Ljava/lang/String;') { + } else if (_resultType === 'Ljava/lang/String;') { // **String Switch Handling** const hashCaseMap: Map = new Map() @@ -1708,7 +1730,7 @@ const codeGenerators: { [type: string]: (node: Node, cg: CodeGenerator) => Compi endLabel.offset = cg.code.length } else { throw new Error( - `Switch statements only support byte, short, int, char, or String types. Found: ${resultType}` + `Switch statements only support byte, short, int, char, String, or enum types. Found: ${_resultType}` ) } diff --git a/src/compiler/compiler-utils.ts b/src/compiler/compiler-utils.ts index 9adcae6d..4bff9609 100644 --- a/src/compiler/compiler-utils.ts +++ b/src/compiler/compiler-utils.ts @@ -6,7 +6,8 @@ import { ClassModifier, FieldModifier, MethodModifier } from '../ast/types/class const classAccessFlagMap = new Map([ ['public', ACCESS_FLAGS.ACC_PUBLIC], ['final', ACCESS_FLAGS.ACC_FINAL], - ['abstract', ACCESS_FLAGS.ACC_ABSTRACT] + ['abstract', ACCESS_FLAGS.ACC_ABSTRACT], + ['enum', ACCESS_FLAGS.ACC_ENUM] ]) export function generateClassAccessFlags(modifiers: Array) { diff --git a/src/types/checker/__tests__/switchStatements.test.ts b/src/types/checker/__tests__/switchStatements.test.ts index 6344ebce..8fb29ca2 100644 --- a/src/types/checker/__tests__/switchStatements.test.ts +++ b/src/types/checker/__tests__/switchStatements.test.ts @@ -73,6 +73,31 @@ const testcases: { } `, result: { type: null, errors: [new IncompatibleTypesError()] } + }, + { + input: ` + enum Color { RED, BLUE } + Color selector = Color.RED; + switch(selector) { + case Color.RED: { + selector = Color.BLUE; + } + default: {} + } + `, + result: { type: null, errors: [] } + }, + { + input: ` + enum Color { RED, BLUE } + enum Other { X } + Color selector = Color.RED; + switch(selector) { + case Other.X: {} + default: {} + } + `, + result: { type: null, errors: [new IncompatibleTypesError()] } } ] diff --git a/src/types/checker/environment.ts b/src/types/checker/environment.ts index 6ef2ad2b..92e12851 100644 --- a/src/types/checker/environment.ts +++ b/src/types/checker/environment.ts @@ -41,7 +41,9 @@ const GLOBAL_TYPE_ENVIRONMENT: { [key: string]: Type } = { // Hard coded variables System: SYSTEM_CLASS, Throwable: new NonPrimitives.Throwable(), - Exception: new NonPrimitives.Exception() + Exception: new NonPrimitives.Exception(), + // enum base type + Enum: new ClassType('Enum') } export class Frame { diff --git a/src/types/checker/index.ts b/src/types/checker/index.ts index 77491719..0f3c39fe 100644 --- a/src/types/checker/index.ts +++ b/src/types/checker/index.ts @@ -69,7 +69,6 @@ const isCastCompatible = (fromType: Type, toType: Type): boolean => { const fromName = fromType.constructor.name; const toName = toType.constructor.name; - console.log(fromName, toName); return !(fromName === 'char' && toName !== 'int'); } @@ -384,7 +383,6 @@ export const typeCheckBody = (node: Node, frame: Frame = Frame.globalFrame()): R return newResult(null, errors) } case 'InstanceofExpression': { - console.log(node) return OK_RESULT } case 'BinaryLiteral': @@ -584,6 +582,88 @@ export const typeCheckBody = (node: Node, frame: Frame = Frame.globalFrame()): R } return newResult(null, errors) } + case 'EnumDeclaration': { + const errors: TypeCheckerError[] = [] + const classType = frame.getType(node.typeIdentifier.identifier, node.typeIdentifier.location) + if (classType instanceof TypeCheckerError) return newResult(null, [classType]) + if (!(classType instanceof ClassType)) throw new Error('enum type retrieved should be ClassImpl') + + const classFrame = frame.newChildFrame() + classFrame.setClass(classType) + classType.mapFields((name, type) => { + const error = classFrame.setVariable(name, type, { startLine: -1, startOffset: -1 }) + if (error) errors.push(error) + }) + if (errors.length > 0) return newResult(null, errors) + + const bodyDecls = node.enumBody.enumBodyDeclarations?.classBodyDeclaration || [] + let numFieldDeclarations = 0 + let numMethodDeclarations = 0 + for (let i = 0; i < bodyDecls.length; i++) { + const bodyDeclaration = bodyDecls[i] + switch (bodyDeclaration.kind) { + case 'ConstructorDeclaration': { + const methodFrame = classFrame.newChildFrame() + const constructor = classType.getConstructor(i - numFieldDeclarations - numMethodDeclarations) + const constructorMethodErrors: TypeCheckerError[] = [] + constructor.mapParameters((name, type, isVarargs) => { + const error = methodFrame.setVariable(name, type, { startLine: -1, startOffset: -1 }) + if (error) constructorMethodErrors.push(error) + }) + if (constructorMethodErrors.length > 0) { + errors.push(...constructorMethodErrors) + break + } + const { errors: checkErrors } = typeCheckBody(bodyDeclaration.constructorBody, methodFrame) + if (checkErrors.length > 0) errors.push(...checkErrors) + break + } + case 'FieldDeclaration': { + for (const variableDeclarator of (bodyDeclaration as any).variableDeclaratorList.variableDeclarators) { + const field = classType.accessField(variableDeclarator.variableDeclaratorId.identifier.identifier, variableDeclarator.variableDeclaratorId.identifier.location) + if (field instanceof TypeCheckerError) throw new Error('field should exist in enum') + const initializer = variableDeclarator.variableInitializer + if (initializer) { + const type = createArrayType(field, initializer, expression => { + const result = typeCheckBody(expression, frame) + if (result.errors.length > 0) return result.errors[0] + if (!result.currentType) throw new Error('array initializer expression should have a type') + return result.currentType + }) + if (type instanceof TypeCheckerError) errors.push(type) + } + } + break + } + case 'MethodDeclaration': { + const methodIdentifier = (bodyDeclaration as any).methodHeader.methodDeclarator.identifier + const methodName = methodIdentifier.identifier + const overloadIndex = bodyDecls + .filter((n: any) => n.kind === 'MethodDeclaration' && (n as any).methodHeader.methodDeclarator.identifier.identifier === methodName) + .findIndex(n => n === bodyDeclaration) + const method = classType.getMethod(methodName)[overloadIndex] + const methodFrame = classFrame.newChildFrame() + const methodErrors: TypeCheckerError[] = [] + methodFrame.setReturnType(method.getReturnType()) + method.mapParameters((name, type, isVarargs) => { + const error = methodFrame.setVariable(name, type, { startLine: -1, startOffset: -1 }) + if (error) methodErrors.push(error) + }) + if (methodErrors.length > 0) { + errors.push(...methodErrors) + break + } + const { errors: checkErrors } = typeCheckBody((bodyDeclaration as any).methodBody, methodFrame) + if (checkErrors.length > 0) errors.push(...checkErrors) + break + } + } + + if (bodyDeclaration.kind === 'FieldDeclaration') numFieldDeclarations += 1 + if (bodyDeclaration.kind === 'MethodDeclaration') numMethodDeclarations += 1 + } + return newResult(null, errors) + } case 'OrdinaryCompilationUnit': { const typeCheckErrors = node.topLevelClassOrInterfaceDeclarations .map(declaration => typeCheckBody(declaration, frame)) @@ -644,16 +724,20 @@ export const typeCheckBody = (node: Node, frame: Frame = Frame.globalFrame()): R const switchBlockFrame = frame.newChildFrame() for (const group of node.switchBlock.switchBlockStatementGroups) { for (const switchLabel of group.switchLabels) { - if ('caseConstant' in switchLabel) { - const checkResult = typeCheckBody( - switchLabel.caseConstant as CaseConstant, - switchBlockFrame - ) - if (checkResult.hasErrors) return checkResult - if (!checkResult.currentType) - throw new TypeCheckerInternalError('Switch case constant should have a type.') - if (expressionCheck.currentType.canBeAssigned(checkResult.currentType)) continue - return newResult(null, [new IncompatibleTypesError(switchLabel.location)]) + // Support both singular 'caseConstant' and plural 'caseConstants' AST shapes + const caseConstants: CaseConstant[] = [] + if ('caseConstant' in switchLabel && (switchLabel as any).caseConstant) caseConstants.push((switchLabel as any).caseConstant as CaseConstant) + if ('caseConstants' in switchLabel && (switchLabel as any).caseConstants) caseConstants.push(...((switchLabel as any).caseConstants as CaseConstant[])) + if (caseConstants.length > 0) { + for (const caseConst of caseConstants) { + const checkResult = typeCheckBody(caseConst, switchBlockFrame) + if (checkResult.hasErrors) return checkResult + if (!checkResult.currentType) + throw new TypeCheckerInternalError('Switch case constant should have a type.') + const assignable = expressionCheck.currentType.canBeAssigned(checkResult.currentType) + if (assignable) continue + return newResult(null, [new IncompatibleTypesError(switchLabel.location)]) + } } } if (group.blockStatements) { diff --git a/src/types/checker/prechecks.ts b/src/types/checker/prechecks.ts index fe7df448..9ec0a9c6 100644 --- a/src/types/checker/prechecks.ts +++ b/src/types/checker/prechecks.ts @@ -1,4 +1,4 @@ -import { Class, ClassType, ObjectClass } from '../types/classes' +import { Class, ClassType, EnumClass, ObjectClass } from '../types/classes' import { ConstructorDeclaration, MethodDeclaration, Node } from '../ast/specificationTypes' import { createClassFieldsAndMethods } from '../typeFactories/classFactory' import { createMethod } from '../typeFactories/methodFactory' @@ -15,6 +15,31 @@ export const addClasses = (node: Node, frame: Frame): Result => { const typeCheckErrors = node.topLevelClassOrInterfaceDeclarations .map(declaration => addClasses(declaration, frame)) .reduce((errors, result) => (result.hasErrors ? [...errors, ...result.errors] : errors), []) + + // Register any nested enum declarations found anywhere in the compilation unit + const registerNestedEnums = (obj: any) => { + if (!obj || typeof obj !== 'object') return + if (Array.isArray(obj)) { + obj.forEach(registerNestedEnums) + return + } + if (obj.kind === 'EnumDeclaration') { + try { + const enumType = new EnumClass(obj.typeIdentifier.identifier) + const err = frame.setType(obj.typeIdentifier.identifier, enumType, obj.typeIdentifier.location) + if (err instanceof Error) { + // duplicate class — add as error + typeCheckErrors.push(new DuplicateClassError(obj.location)) + } + } catch (e) { + // ignore + } + return + } + Object.keys(obj).forEach(k => registerNestedEnums(obj[k])) + } + node.topLevelClassOrInterfaceDeclarations.forEach(registerNestedEnums) + return newResult(null, typeCheckErrors) } case 'NormalClassDeclaration': { @@ -35,7 +60,16 @@ export const addClasses = (node: Node, frame: Frame): Result => { return newResult(classType) } case 'EnumDeclaration': { - throw new Error('Not implemented') + const enumType = new EnumClass(node.typeIdentifier.identifier) + const errors: TypeCheckerError[] = [] + if (errors.length > 0) return newResult(null, errors) + const error = frame.setType( + node.typeIdentifier.identifier, + enumType, + node.typeIdentifier.location + ) + if (error instanceof Error) return newResult(null, [new DuplicateClassError(node.location)]) + return newResult(enumType) } case 'RecordDeclaration': { throw new Error('Not implemented') @@ -54,6 +88,23 @@ export const addClassMethods = (node: Node, frame: Frame): Result => { const typeCheckErrors = node.topLevelClassOrInterfaceDeclarations .map(declaration => addClassMethods(declaration, frame)) .reduce((errors, result) => (result.hasErrors ? [...errors, ...result.errors] : errors), []) + + // Also process any nested enum declarations (e.g., enums declared inside methods) + const processNestedEnums = (obj: any) => { + if (!obj || typeof obj !== 'object') return + if (Array.isArray(obj)) { + obj.forEach(processNestedEnums) + return + } + if (obj.kind === 'EnumDeclaration') { + const res = addClassMethods(obj, frame) + if (res.hasErrors) typeCheckErrors.push(...res.errors) + return + } + Object.keys(obj).forEach(k => processNestedEnums(obj[k])) + } + node.topLevelClassOrInterfaceDeclarations.forEach(processNestedEnums) + return newResult(null, typeCheckErrors) } case 'ConstructorDeclaration': @@ -74,6 +125,64 @@ export const addClassMethods = (node: Node, frame: Frame): Result => { if (classType instanceof TypeCheckerError) return newResult(null, [classType]) return newResult(classType) } + case 'EnumDeclaration': { + const createMethodLocal = ( + node: ConstructorDeclaration | MethodDeclaration + ): Method | TypeCheckerError => { + const result = addClassMethods(node, frame) + if (result.errors.length > 0) return result.errors[0] + return result.currentType as Method + } + + // Populate enum constants and any class-body declarations (fields/methods/constructors) + const classType = frame.getType(node.typeIdentifier.identifier, node.typeIdentifier.location) + if (classType instanceof TypeCheckerError) return newResult(null, [classType]) + if (!(classType instanceof ClassType)) throw new Error('enum type should be a ClassImpl') + + // Add enum constants as fields of the enum type + const enumConstants = node.enumBody.enumConstantList?.enumConstants || [] + for (const constant of enumConstants) { + const fieldError = classType.addField(constant.identifier.identifier, classType, constant.location) + if (fieldError instanceof TypeCheckerError) return newResult(null, [fieldError]) + } + + // Process body declarations similar to class body + const bodyDecls = node.enumBody.enumBodyDeclarations?.classBodyDeclaration || [] + for (const bodyNode of bodyDecls) { + switch (bodyNode.kind) { + case 'ConstructorDeclaration': { + const constructorMethod = createMethodLocal(bodyNode as ConstructorDeclaration) + if (constructorMethod instanceof TypeCheckerError) return newResult(null, [constructorMethod]) + const error = classType.addConstructor(constructorMethod, bodyNode.location) + if (error instanceof TypeCheckerError) return newResult(null, [error]) + break + } + case 'FieldDeclaration': { + const fieldType = frame.getType( + (bodyNode as any).unannType ? (bodyNode as any).unannType : (bodyNode as any).fieldType, + bodyNode.location + ) + if (fieldType instanceof TypeCheckerError) return newResult(null, [fieldType]) + for (const declarator of (bodyNode as any).variableDeclaratorList.variableDeclarators) { + const fieldIdentifier = declarator.variableDeclaratorId.identifier + const error = classType.addField(fieldIdentifier.identifier, fieldType, fieldIdentifier.location) + if (error instanceof TypeCheckerError) return newResult(null, [error]) + } + break + } + case 'MethodDeclaration': { + const methodSignature = createMethodLocal(bodyNode as MethodDeclaration) + if (methodSignature instanceof TypeCheckerError) return newResult(null, [methodSignature]) + const methodName = (bodyNode as MethodDeclaration).methodHeader.methodDeclarator.identifier + const error = classType.addMethod(methodName.identifier, methodSignature, methodName.location) + if (error instanceof TypeCheckerError) return newResult(null, [error]) + break + } + } + } + + return newResult(classType) + } default: return OK_RESULT } @@ -111,6 +220,18 @@ export const addClassParents = (node: Node, frame: Frame): Result => { } return newResult(classType) } + case 'EnumDeclaration': { + const classType = frame.getType(node.typeIdentifier.identifier, node.typeIdentifier.location) + if (classType instanceof Error) return newResult(null, [classType]) + if (!(classType instanceof ClassType)) throw new Error('enum type should be a ClassImpl') + + // Enums implicitly extend java.lang.Enum (represented here as 'Enum' in the type environment) + const enumBase = frame.getType('Enum', node.typeIdentifier.location) + if (enumBase instanceof Error) return newResult(null, [enumBase]) + if (!(enumBase instanceof ClassType)) throw new Error('Enum base should be a ClassImpl') + classType.setParentClass(enumBase) + return newResult(classType) + } default: return OK_RESULT } diff --git a/src/types/checker/statements.ts b/src/types/checker/statements.ts index fe88fcef..3ce1e192 100644 --- a/src/types/checker/statements.ts +++ b/src/types/checker/statements.ts @@ -6,6 +6,7 @@ import { TypeCheckerError } from '../errors' import { Throwable } from '../types/references' +import { EnumClass } from '../types/classes' import { Type } from '../types/type' import { isPrimitiveBooleanType, @@ -29,6 +30,7 @@ export const checkSwitchExpression = ( ): null | TypeCheckerError => { if (isPrimitiveIntegralType(expressionType) && !isPrimitiveLongType(expressionType)) return null if (isStringType(expressionType)) return null + if (expressionType instanceof EnumClass) return null return new SelectorTypeNotAllowedError(location) } diff --git a/src/types/types/classes.ts b/src/types/types/classes.ts index 00b4ec0d..8334f2e5 100644 --- a/src/types/types/classes.ts +++ b/src/types/types/classes.ts @@ -144,6 +144,8 @@ export class ClassType extends ClassOrInterfaceType implements Class { } } +export class EnumClass extends ClassType {} + export class ObjectClass extends ClassOrInterfaceType implements Class { public readonly name: string = 'Object' public constructor() {