diff --git a/src/ast/__tests__/class-extractor.test.ts b/src/ast/__tests__/class-extractor.test.ts index 16df694..8f930c5 100644 --- a/src/ast/__tests__/class-extractor.test.ts +++ b/src/ast/__tests__/class-extractor.test.ts @@ -124,3 +124,39 @@ describe("extract multiple NormalClassDeclaration correctly", () => { expect(ast).toEqual(expectedAst); }); }); + +describe("extract nested NormalClassDeclaration correctly", () => { + it("extract a static nested class inside a class body correctly", () => { + const programStr = ` + class Outer { + static class Inner {} + } + `; + + const expectedAst: AST = { + kind: "CompilationUnit", + importDeclarations: [], + topLevelClassOrInterfaceDeclarations: [ + { + kind: "NormalClassDeclaration", + classModifier: [], + typeIdentifier: "Outer", + classBody: [ + { + kind: "NormalClassDeclaration", + classModifier: ["static"], + typeIdentifier: "Inner", + classBody: [], + location: expect.anything(), + }, + ], + location: expect.anything(), + }, + ], + location: expect.anything(), + }; + + const ast = parse(programStr); + expect(ast).toEqual(expectedAst); + }); +}); diff --git a/src/ast/astExtractor/class-extractor.ts b/src/ast/astExtractor/class-extractor.ts index cd55cac..3896aa2 100644 --- a/src/ast/astExtractor/class-extractor.ts +++ b/src/ast/astExtractor/class-extractor.ts @@ -91,5 +91,12 @@ export class ClassExtractor extends BaseJavaCstVisitorWithDefaults { this.body.push(methodNode); }) } + if (ctx.classDeclaration) { + ctx.classDeclaration.forEach(x => { + const classExtractor = new ClassExtractor(); + const classNode = classExtractor.extract(x); + this.body.push(classNode); + }) + } } } diff --git a/src/ast/types/classes.ts b/src/ast/types/classes.ts index 576d9aa..38df2ef 100644 --- a/src/ast/types/classes.ts +++ b/src/ast/types/classes.ts @@ -43,7 +43,7 @@ export type ClassModifier = | "strictfp" | "enum"; -export type ClassBodyDeclaration = ClassMemberDeclaration | ConstructorDeclaration | EnumDeclaration; +export type ClassBodyDeclaration = ClassMemberDeclaration | ConstructorDeclaration | EnumDeclaration | NormalClassDeclaration; export type ClassMemberDeclaration = MethodDeclaration | FieldDeclaration; export interface ConstructorDeclaration extends BaseNode { diff --git a/src/compiler/__tests__/classOrdering.test.ts b/src/compiler/__tests__/classOrdering.test.ts index 838e548..22c51a5 100644 --- a/src/compiler/__tests__/classOrdering.test.ts +++ b/src/compiler/__tests__/classOrdering.test.ts @@ -11,7 +11,7 @@ describe('compiled class ordering', () => { } ` const classes = compileFromSource(program) - expect(classes.map(c => c.className)).toEqual(['Main', 'Day']) + expect(classes.map(c => c.className)).toEqual(['Main', 'Main$Day']) }) it('keeps top-level declaration order, with member enums appended', () => { @@ -24,7 +24,9 @@ describe('compiled class ordering', () => { ` const classes = compileFromSource(program) expect(classes[0].className).toBe('Main') - expect(new Set(classes.map(c => c.className))).toEqual(new Set(['Main', 'A', 'B'])) + expect(new Set(classes.map(c => c.className))).toEqual( + new Set(['Main', 'Main$A', 'Main$B']) + ) }) it('is unchanged when there is no enum', () => { diff --git a/src/compiler/__tests__/index.ts b/src/compiler/__tests__/index.ts index bddc2a2..eefc9b3 100644 --- a/src/compiler/__tests__/index.ts +++ b/src/compiler/__tests__/index.ts @@ -10,6 +10,7 @@ import { importTest } from "./tests/import.test"; import { arrayTest } from "./tests/array.test"; import { classTest } from "./tests/class.test"; import { enumTest } from "./tests/enum.test"; +import { nestedClassesTest } from "./tests/nestedClasses.test"; import { typeConversionTest } from "./tests/typeConversion.test"; describe("compiler tests", () => { @@ -25,5 +26,6 @@ describe("compiler tests", () => { arrayTest(); classTest(); enumTest(); + nestedClassesTest(); typeConversionTest(); }) diff --git a/src/compiler/__tests__/tests/class.test.ts b/src/compiler/__tests__/tests/class.test.ts index 4ec2a9b..4a34d8c 100644 --- a/src/compiler/__tests__/tests/class.test.ts +++ b/src/compiler/__tests__/tests/class.test.ts @@ -103,6 +103,19 @@ const testCases: testCase[] = [ `, expectedLines: ["in f"], }, + { + comment: "instance field with an inline initializer", + program: ` + public class Main { + public int one = 1; + public static void main(String[] args) { + Main m = new Main(); + System.out.println(m.one); + } + } + `, + expectedLines: ["1"], + }, { comment: "instance field", program: ` diff --git a/src/compiler/__tests__/tests/nestedClasses.test.ts b/src/compiler/__tests__/tests/nestedClasses.test.ts new file mode 100644 index 0000000..7b6714b --- /dev/null +++ b/src/compiler/__tests__/tests/nestedClasses.test.ts @@ -0,0 +1,99 @@ +import { + runTest, + testCase, +} from "../__utils__/test-utils"; + +const testCases: testCase[] = [ + { + comment: "basic static nested class instantiation and usage", + program: ` + public class Main { + public static class Inner { + public int value; + public Inner(int v) { + this.value = v; + } + public int getValue() { + return this.value; + } + } + + public static void main(String[] args) { + Inner inner = new Inner(42); + System.out.println(inner.getValue()); + } + } + `, + expectedLines: ["42"], + }, + { + comment: "one nested class field typed as a sibling nested class", + program: ` + public class Main { + public static class Box { + public int value; + public Box(int v) { + this.value = v; + } + } + + public static class Holder { + public Box box; + public Holder(Box b) { + this.box = b; + } + } + + public static void main(String[] args) { + Box box = new Box(9); + Holder holder = new Holder(box); + System.out.println(holder.box.value); + } + } + `, + expectedLines: ["9"], + }, + { + comment: "nested class instance field with an inline initializer", + program: ` + public class Main { + static class Inner1 { + public int one = 1; + } + public static void main(String[] args) { + Inner1 inner1 = new Inner1(); + int n = inner1.one; + System.out.println(n); + } + } + `, + expectedLines: ["1"], + }, + { + comment: "two levels of static nested class", + program: ` + public class Main { + public static class Middle { + public static class Inner { + public String greet() { + return "hello"; + } + } + } + + public static void main(String[] args) { + Inner inner = new Inner(); + System.out.println(inner.greet()); + } + } + `, + expectedLines: ["hello"], + }, +]; + +export const nestedClassesTest = () => describe("static nested classes", () => { + for (let testCase of testCases) { + const { comment: comment, program: program, expectedLines: expectedLines } = testCase; + it(comment, () => runTest(program, expectedLines)); + } +}); diff --git a/src/compiler/code-generator.ts b/src/compiler/code-generator.ts index 2a2225e..5cc812d 100644 --- a/src/compiler/code-generator.ts +++ b/src/compiler/code-generator.ts @@ -1098,7 +1098,13 @@ const codeGenerators: { [type: string]: (node: Node, cg: CodeGenerator) => Compi } } const res = compile(expr, cg) - const classInfoIndex = cg.constantPoolManager.indexClassInfo(ct) + let castClassName = ct + try { + castClassName = cg.symbolTable.queryClass(ct).name + } catch (e) { + castClassName = ct.includes('/') ? ct : ct.replace(/\./g, '/') + } + const classInfoIndex = cg.constantPoolManager.indexClassInfo(castClassName) cg.code.push(OPCODE.CHECKCAST, 0, classInfoIndex) return res }, @@ -1107,7 +1113,18 @@ const codeGenerators: { [type: string]: (node: Node, cg: CodeGenerator) => Compi const { identifier: id, argumentList: argLst } = node as ClassInstanceCreationExpression let maxStack = 2 - cg.code.push(OPCODE.NEW, 0, cg.constantPoolManager.indexClassInfo(id), OPCODE.DUP) + let instantiatedClassName = id + try { + instantiatedClassName = cg.symbolTable.queryClass(id).name + } catch (e) { + instantiatedClassName = id.includes('/') ? id : id.replace(/\./g, '/') + } + cg.code.push( + OPCODE.NEW, + 0, + cg.constantPoolManager.indexClassInfo(instantiatedClassName), + OPCODE.DUP + ) const argTypes: Array = [] argLst.forEach((x, i) => { @@ -1120,7 +1137,10 @@ const codeGenerators: { [type: string]: (node: Node, cg: CodeGenerator) => Compi const methodInfos = cg.symbolTable.queryMethod('') as MethodInfos for (let i = 0; i < methodInfos.length; i++) { const methodInfo = methodInfos[i] - if (methodInfo.typeDescriptor.includes(argDescriptor) && methodInfo.className == id) { + if ( + methodInfo.typeDescriptor.includes(argDescriptor) && + methodInfo.className == instantiatedClassName + ) { const method = cg.constantPoolManager.indexMethodrefInfo( methodInfo.className, methodInfo.name, @@ -1131,7 +1151,7 @@ const codeGenerators: { [type: string]: (node: Node, cg: CodeGenerator) => Compi } } - return { stackSize: maxStack, resultType: id } + return { stackSize: maxStack, resultType: instantiatedClassName } }, ArrayAccess: (node: Node, cg: CodeGenerator) => { @@ -1768,7 +1788,11 @@ const codeGenerators: { [type: string]: (node: Node, cg: CodeGenerator) => Compi return field.ordinal })() : label.expression.kind === 'ExpressionName' - ? (() => { throw new Error(`Identifier case labels are only supported for enum switch selectors: ${label.expression.name}`) })() + ? (() => { + throw new Error( + `Identifier case labels are only supported for enum switch selectors: ${label.expression.name}` + ) + })() : parseInt((label.expression as Literal).literalType.value) caseValues.push(value) caseLabelMap.set(value, caseLabels[index]) diff --git a/src/compiler/compiler.ts b/src/compiler/compiler.ts index ee65704..2892487 100644 --- a/src/compiler/compiler.ts +++ b/src/compiler/compiler.ts @@ -6,11 +6,13 @@ import { ConstructorDeclaration, EnumDeclaration, FieldDeclaration, - MethodDeclaration + MethodDeclaration, + NormalClassDeclaration } from '../ast/types/classes' import { AttributeInfo } from '../ClassFile/types/attributes' import { FieldInfo } from '../ClassFile/types/fields' import { MethodInfo } from '../ClassFile/types/methods' +import { ConstructNotSupportedError } from './error' import { ConstantPoolManager } from './constant-pool-manager' import { generateClassAccessFlags, @@ -55,32 +57,37 @@ export class Compiler { compile(ast: AST) { this.setup() this.symbolTable.handleImports(ast.importDeclarations) - const declarations = [ - ...ast.topLevelClassOrInterfaceDeclarations, - ...ast.topLevelClassOrInterfaceDeclarations.flatMap(declaration => - declaration.kind === 'NormalClassDeclaration' - ? this.getMemberEnums(declaration.classBody) - : [] - ) - ] + const topLevelDeclarations = ast.topLevelClassOrInterfaceDeclarations + const memberDeclarations = topLevelDeclarations.flatMap(declaration => + declaration.kind === 'NormalClassDeclaration' + ? this.getMemberTypes(declaration.classBody, declaration.typeIdentifier) + : [] + ) + const declarations = [...topLevelDeclarations, ...memberDeclarations] - // Enums are compiled first so their synthetic members are in the symbol - // table before the classes that reference them are compiled. - const compilationOrder = [ - ...declarations.filter(declaration => declaration.kind === 'EnumDeclaration'), - ...declarations.filter(declaration => declaration.kind !== 'EnumDeclaration') - ] + // Member (nested) declarations are compiled before the top-level + // declarations that may reference them, so their fields/methods are + // already in the symbol table by the time the enclosing class compiles. + const compilationOrder = [...memberDeclarations, ...topLevelDeclarations] declarations.forEach(decl => { const className = decl.typeIdentifier const parentClassName = 'sclass' in decl && decl.sclass ? decl.sclass : 'java/lang/Object' const accessFlags = generateClassAccessFlags(decl.classModifier) - this.symbolTable.insertClassInfo({ - name: className, - accessFlags: accessFlags, - parentClassName: parentClassName, - isEnum: decl.kind === 'EnumDeclaration' - }) + // Nested types are registered under both their simple name (so + // unqualified references from within the compilation unit resolve) and + // their qualified binary name (so a descriptor like `LOuter$Inner;` + // can be resolved back to this class's info via queryClass(...)). + const simpleName = className.split('$').pop() as string + this.symbolTable.insertClassInfo( + { + name: className, + accessFlags: accessFlags, + parentClassName: parentClassName, + isEnum: decl.kind === 'EnumDeclaration' + }, + [simpleName, className] + ) this.symbolTable.returnToRoot() }) @@ -97,10 +104,38 @@ export class Compiler { return declarations.map(decl => compiled.get(decl) as Class) } - private getMemberEnums(classBody: Array): Array { + /** + * Flattens nested enum and static nested class declarations out of a class + * body, rewriting each `typeIdentifier` to its JVM binary name + * (`Outer$Inner`, `Outer$Middle$Inner`, ...) so the rest of the compiler can + * treat them exactly like top-level declarations. + */ + private getMemberTypes( + classBody: Array, + enclosingBinaryName: string + ): Array { return classBody.flatMap(declaration => { - if (declaration.kind !== 'EnumDeclaration') return [] - return [declaration, ...this.getMemberEnums(declaration.enumBody.bodyMembers || [])] + if (declaration.kind === 'EnumDeclaration') { + const qualified = { + ...declaration, + typeIdentifier: enclosingBinaryName + '$' + declaration.typeIdentifier + } + return [ + qualified, + ...this.getMemberTypes(qualified.enumBody.bodyMembers || [], qualified.typeIdentifier) + ] + } + if (declaration.kind === 'NormalClassDeclaration') { + if (!declaration.classModifier.includes('static')) { + throw new ConstructNotSupportedError('non-static nested class') + } + const qualified = { + ...declaration, + typeIdentifier: enclosingBinaryName + '$' + declaration.typeIdentifier + } + return [qualified, ...this.getMemberTypes(qualified.classBody, qualified.typeIdentifier)] + } + return [] }) } @@ -110,7 +145,10 @@ export class Compiler { this.parentClassName = sclass ? sclass : 'java/lang/Object' const accessFlags = generateClassAccessFlags(classNode.classModifier) this.symbolTable.extend() - this.symbolTable.insertClassInfo({ name: this.className, accessFlags: accessFlags }) + this.symbolTable.insertClassInfo({ name: this.className, accessFlags: accessFlags }, [ + this.className.split('$').pop() as string, + this.className + ]) const superClassIndex = this.constantPoolManager.indexClassInfo(this.parentClassName) const thisClassIndex = this.constantPoolManager.indexClassInfo(this.className) @@ -146,7 +184,10 @@ export class Compiler { this.parentClassName = 'java/lang/Enum' const accessFlags = generateClassAccessFlags(enumNode.classModifier) | 0x4000 // ACC_ENUM this.symbolTable.extend() - this.symbolTable.insertClassInfo({ name: this.className, accessFlags: accessFlags }) + this.symbolTable.insertClassInfo({ name: this.className, accessFlags: accessFlags }, [ + this.className.split('$').pop() as string, + this.className + ]) const superClassIndex = this.constantPoolManager.indexClassInfo(this.parentClassName) const thisClassIndex = this.constantPoolManager.indexClassInfo(this.className) @@ -536,29 +577,47 @@ export class Compiler { nonStaticMethods.forEach(m => this.compileMethod(m)) staticMethods.forEach(m => this.compileMethod(m)) this.compileStaticFieldInitializers(staticFields) - constructors.forEach(c => this.compileConstructor(c)) + constructors.forEach(c => this.compileConstructor(c, nonStaticFields)) } /** - * Emits a `` that runs the initialiser expression of each static - * field, in declaration order (`static T f = expr;` -> `f = expr;`). + * Builds synthetic `f = expr;` assignment statements from each field's + * initialiser, in declaration order. Instance-field targets are qualified + * as `this.f` - the Assignment code generator only emits the ALOAD_0 + * needed before PUTFIELD when it sees that prefix; a bare name silently + * skips loading the receiver. */ - private compileStaticFieldInitializers(staticFields: Array) { + private buildFieldInitializerStatements( + fields: Array, + qualifyWithThis: boolean = false + ): any[] { const blockStatements: any[] = [] - for (const field of staticFields) { + for (const field of fields) { for (const declarator of field.variableDeclaratorList) { if (declarator.variableInitializer === undefined) continue + const name = qualifyWithThis + ? 'this.' + declarator.variableDeclaratorId + : declarator.variableDeclaratorId blockStatements.push({ kind: 'ExpressionStatement', stmtExp: { kind: 'Assignment', - left: { kind: 'ExpressionName', name: declarator.variableDeclaratorId }, + left: { kind: 'ExpressionName', name }, operator: '=', right: declarator.variableInitializer } }) } } + return blockStatements + } + + /** + * Emits a `` that runs the initialiser expression of each static + * field, in declaration order (`static T f = expr;` -> `f = expr;`). + */ + private compileStaticFieldInitializers(staticFields: Array) { + const blockStatements = this.buildFieldInitializerStatements(staticFields) if (blockStatements.length === 0) return this.compileMethod({ @@ -646,7 +705,14 @@ export class Compiler { }) } - private compileConstructor(constructor: ConstructorDeclaration) { + private compileConstructor( + constructor: ConstructorDeclaration, + instanceFields: Array = [] + ) { + // Instance field initialisers run at the start of every constructor body + // (right after the implicit super() call that generateCode() emits), + // mirroring how compileStaticFieldInitializers seeds . + const fieldInitializers = this.buildFieldInitializerStatements(instanceFields, true) const methodNode: MethodDeclaration = { kind: 'MethodDeclaration', methodModifier: constructor.constructorModifier, @@ -655,7 +721,10 @@ export class Compiler { formalParameterList: constructor.constructorDeclarator.formalParameterList, result: 'void' }, - methodBody: constructor.constructorBody + methodBody: { + kind: 'Block', + blockStatements: [...fieldInitializers, ...constructor.constructorBody.blockStatements] + } } this.compileMethod(methodNode) diff --git a/src/compiler/symbol-table.ts b/src/compiler/symbol-table.ts index d1442e9..e121e28 100644 --- a/src/compiler/symbol-table.ts +++ b/src/compiler/symbol-table.ts @@ -245,18 +245,24 @@ export class SymbolTable { this.curTable = this.tables[this.curIdx] } - insertClassInfo(info: ClassInfo) { - const key = generateSymbol(info.name, SymbolType.CLASS) - - if (this.curTable.has(key)) { - throw new SymbolRedeclarationError(info.name) - } - + // A nested class is registered under both its simple name (so unqualified + // source references like `Inner` resolve) and its qualified binary name + // (so a `Ljava/lang/Outer$Inner;`-style descriptor can be resolved back to + // its member table) - both keys share the same SymbolNode/children table. + insertClassInfo(info: ClassInfo, lookupNames: Array = [info.name]) { + const names = [...new Set(lookupNames)] const symbolNode: SymbolNode = { info: info, children: this.getNewTable() } - this.curTable.set(key, symbolNode) + + for (const name of names) { + const key = generateSymbol(name, SymbolType.CLASS) + if (this.curTable.has(key)) { + throw new SymbolRedeclarationError(info.name) + } + this.curTable.set(key, symbolNode) + } this.tables[++this.curIdx] = symbolNode.children this.curTable = this.tables[this.curIdx] diff --git a/src/types/checker/__tests__/nestedClasses.test.ts b/src/types/checker/__tests__/nestedClasses.test.ts new file mode 100644 index 0000000..e294c38 --- /dev/null +++ b/src/types/checker/__tests__/nestedClasses.test.ts @@ -0,0 +1,99 @@ +import { check } from '..' +import { parse } from '../../ast' +import { TypeCheckerError, UnsupportedNestedClassError } from '../../errors' +import { Type } from '../../types/type' + +const testcases: { + input: string + result: { type: Type | null; errors: Error[] } + only?: boolean +}[] = [ + { + input: ` + class Outer { + static class Inner { + int value; + Inner(int v) { value = v; } + int getValue() { return value; } + } + + public static void main(String[] args) { + Inner inner = new Inner(5); + inner.getValue(); + } + } + `, + result: { type: null, errors: [] } + }, + { + input: ` + class Outer { + static int counter = 1; + + static class Inner { + int read() { return counter; } + } + + public static void main(String[] args) { + Inner inner = new Inner(); + inner.read(); + } + } + `, + result: { type: null, errors: [] } + }, + { + input: ` + class Outer { + static class Middle { + static class Inner { + void hello() {} + } + } + + public static void main(String[] args) { + Inner inner = new Inner(); + inner.hello(); + } + } + `, + result: { type: null, errors: [] } + }, + { + input: ` + class Outer { + class Inner {} + + public static void main(String[] args) {} + } + `, + result: { type: null, errors: [new UnsupportedNestedClassError()] } + } +] + +describe('Type Checker', () => { + testcases.map(testcase => { + let it = test + if (testcase.only) it = test.only + it(`Checking nested classes for '${testcase.input}'`, () => { + const program = testcase.input + const ast = parse(program) + if (!ast) throw new Error('Program parsing returns null.') + if (ast instanceof TypeCheckerError) throw new Error('Test case is invalid.') + const result = check(ast) + if (result.currentType === null) expect(result.currentType).toBe(testcase.result.type) + else expect(result.currentType).toBeInstanceOf(testcase.result.type) + if (testcase.result.errors.length > result.errors.length) { + testcase.result.errors.forEach((error, index) => { + if (!result.errors[index]) expect('').toBe(error.message) + expect(result.errors[index].message).toBe(error.message) + }) + } else { + result.errors.forEach((error, index) => { + if (!testcase.result.errors[index]) expect(error.message).toBe('') + expect(error.message).toBe(testcase.result.errors[index].message) + }) + } + }) + }) +}) diff --git a/src/types/checker/index.ts b/src/types/checker/index.ts index e9147bc..da07a0a 100644 --- a/src/types/checker/index.ts +++ b/src/types/checker/index.ts @@ -552,6 +552,7 @@ export const typeCheckBody = (node: Node, frame: Frame = Frame.globalFrame()): R let numFieldDeclarations = 0 let numMethodDeclarations = 0 + let numNestedClassDeclarations = 0 for (let i = 0; i < node.classBody.classBodyDeclarations.length; i++) { const bodyDeclaration = node.classBody.classBodyDeclarations[i] @@ -559,7 +560,7 @@ export const typeCheckBody = (node: Node, frame: Frame = Frame.globalFrame()): R case 'ConstructorDeclaration': { const methodFrame = classFrame.newChildFrame() const constructor = classType.getConstructor( - i - numFieldDeclarations - numMethodDeclarations + i - numFieldDeclarations - numMethodDeclarations - numNestedClassDeclarations ) const constructorMethodErrors: TypeCheckerError[] = [] constructor.mapParameters((name, type, isVarargs) => { @@ -634,6 +635,113 @@ export const typeCheckBody = (node: Node, frame: Frame = Frame.globalFrame()): R if (checkErrors.length > 0) errors.push(...checkErrors) break } + case 'NormalClassDeclaration': { + const { errors: checkErrors } = typeCheckBody(bodyDeclaration, classFrame) + if (checkErrors.length > 0) errors.push(...checkErrors) + break + } + } + + if (bodyDeclaration.kind === 'FieldDeclaration') numFieldDeclarations += 1 + if (bodyDeclaration.kind === 'MethodDeclaration') numMethodDeclarations += 1 + if (bodyDeclaration.kind === 'NormalClassDeclaration') numNestedClassDeclarations += 1 + } + 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.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 diff --git a/src/types/checker/prechecks.ts b/src/types/checker/prechecks.ts index c0a08ac..2a6304d 100644 --- a/src/types/checker/prechecks.ts +++ b/src/types/checker/prechecks.ts @@ -2,7 +2,12 @@ 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' -import { CyclicInheritanceError, DuplicateClassError, TypeCheckerError } from '../errors' +import { + CyclicInheritanceError, + DuplicateClassError, + TypeCheckerError, + UnsupportedNestedClassError +} from '../errors' import { Method } from '../types/methods' import { Frame } from './environment' import { newResult, OK_RESULT, Result } from '.' @@ -57,6 +62,24 @@ export const addClasses = (node: Node, frame: Frame): Result => { node.typeIdentifier.location ) if (error instanceof Error) return newResult(null, [new DuplicateClassError(node.location)]) + + // Register static nested class declarations found directly in this class's body. + for (const bodyDeclaration of node.classBody.classBodyDeclarations) { + if (bodyDeclaration.kind !== 'NormalClassDeclaration') continue + const isStatic = bodyDeclaration.classModifiers.some( + modifier => modifier.identifier === 'static' + ) + if (!isStatic) { + errors.push(new UnsupportedNestedClassError(bodyDeclaration.location)) + continue + } + const nestedResult = addClasses(bodyDeclaration, frame) + if (nestedResult.hasErrors) errors.push(...nestedResult.errors) + else if (nestedResult.currentType instanceof ClassType) + nestedResult.currentType.setEnclosingClass(classType) + } + if (errors.length > 0) return newResult(null, errors) + return newResult(classType) } case 'EnumDeclaration': { @@ -123,6 +146,73 @@ export const addClassMethods = (node: Node, frame: Frame): Result => { } const classType = createClassFieldsAndMethods(node, frame, createMethod, createMethod) if (classType instanceof TypeCheckerError) return newResult(null, [classType]) + + const errors: TypeCheckerError[] = [] + for (const bodyDeclaration of node.classBody.classBodyDeclarations) { + if (bodyDeclaration.kind !== 'NormalClassDeclaration') continue + const nestedResult = addClassMethods(bodyDeclaration, frame) + if (nestedResult.hasErrors) errors.push(...nestedResult.errors) + } + if (errors.length > 0) return newResult(null, errors) + + 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) + 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) + if (methodSignature instanceof TypeCheckerError) return newResult(null, [methodSignature]) + const methodName = (bodyNode).methodHeader.methodDeclarator.identifier + const error = classType.addMethod(methodName.identifier, methodSignature, methodName.location) + if (error instanceof TypeCheckerError) return newResult(null, [error]) + break + } + } + } + return newResult(classType) } case 'EnumDeclaration': { @@ -218,6 +308,27 @@ export const addClassParents = (node: Node, frame: Frame): Result => { } classType.setParentClass(extendsType) } + + const errors: TypeCheckerError[] = [] + for (const bodyDeclaration of node.classBody.classBodyDeclarations) { + if (bodyDeclaration.kind !== 'NormalClassDeclaration') continue + const nestedResult = addClassParents(bodyDeclaration, frame) + if (nestedResult.hasErrors) errors.push(...nestedResult.errors) + } + if (errors.length > 0) return newResult(null, errors) + + 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) } case 'EnumDeclaration': { diff --git a/src/types/errors.ts b/src/types/errors.ts index b1dc889..6eb17ed 100644 --- a/src/types/errors.ts +++ b/src/types/errors.ts @@ -171,3 +171,9 @@ export class UnhandledExceptionError extends TypeCheckerError { super('unhandled exception', location) } } + +export class UnsupportedNestedClassError extends TypeCheckerError { + constructor(location?: Location) { + super('only static nested classes are supported', location) + } +} diff --git a/src/types/types/classes.ts b/src/types/types/classes.ts index 8334f2e..7bc237f 100644 --- a/src/types/types/classes.ts +++ b/src/types/types/classes.ts @@ -26,6 +26,7 @@ export class ClassType extends ClassOrInterfaceType implements Class { public readonly name: string private _modifiers = new Modifiers() private _parent: Class = new ObjectClass() + private _enclosingClass: Class | null = null private _constructors: Method[] = [] private _fields = new Map() @@ -39,13 +40,27 @@ export class ClassType extends ClassOrInterfaceType implements Class { public accessField(_name: string, location: Location): Type | TypeCheckerError { const field = this._fields.get(_name) if (field) return field - return this._parent.accessField(_name, location) + const parentResult = this._parent.accessField(_name, location) + if (!(parentResult instanceof TypeCheckerError)) return parentResult + if (this._enclosingClass) return this._enclosingClass.accessField(_name, location) + return parentResult } public accessMethod(name: string, location: Location): Method[] | TypeCheckerError { const method = this._methods.get(name) if (method) return method - return this._parent.accessMethod(name, location) + const parentResult = this._parent.accessMethod(name, location) + if (!(parentResult instanceof TypeCheckerError)) return parentResult + if (this._enclosingClass) return this._enclosingClass.accessMethod(name, location) + return parentResult + } + + public getEnclosingClass(): Class | null { + return this._enclosingClass + } + + public setEnclosingClass(enclosingClass: Class): void { + this._enclosingClass = enclosingClass } public addConstructor(method: Method, location: Location): void | TypeCheckerError {