diff --git a/tslang/include/TypeScript/MLIRLogic/MLIRDefines.h b/tslang/include/TypeScript/MLIRLogic/MLIRDefines.h index adab23935..24656f028 100644 --- a/tslang/include/TypeScript/MLIRLogic/MLIRDefines.h +++ b/tslang/include/TypeScript/MLIRLogic/MLIRDefines.h @@ -50,6 +50,6 @@ using SymbolTableScopeT = llvm::ScopedHashTableScope; using BoundRefCacheScopeT = llvm::ScopedHashTableScope; typedef std::pair SafeTypeKeyType; -using SafeTypesMapScopeT = llvm::ScopedHashTableScope; +using SafeTypesMapScopeT = llvm::ScopedHashTableScope; #endif // MLIR_TYPESCRIPT_MLIRGENLOGIC_MLIRDEFINES_H_ \ No newline at end of file diff --git a/tslang/lib/TypeScript/MLIRGenExpressions.cpp b/tslang/lib/TypeScript/MLIRGenExpressions.cpp index bdd9cd7c1..b680590a3 100644 --- a/tslang/lib/TypeScript/MLIRGenExpressions.cpp +++ b/tslang/lib/TypeScript/MLIRGenExpressions.cpp @@ -895,16 +895,38 @@ namespace mlirgen auto namePtr = MLIRHelper::getName(propertyAccessExpression->name, stringAllocator); auto propAccessStrRef = mlir::StringRef(print(propertyAccessExpression)).copy(stringAllocator); - // check if we have safe type mapped value - auto safeTypedValue = safeTypesMap.lookup({ expressionValue.getType(), propAccessStrRef }); - if (safeTypedValue) - { - LLVM_DEBUG(llvm::dbgs() << "\n\t...safe type fieldname: \t " - << propAccessStrRef << "." << namePtr << "type: " << expressionValue.getType() << " = " << safeTypedValue;); - return safeTypedValue; + auto fieldResult = mlirGenPropertyAccessExpression(location, expressionValue, namePtr, + !!propertyAccessExpression->questionDotToken, genContext); + + // a narrowed field: read as it is now, cast to the type the test narrowed it to + auto safeType = safeTypesMap.lookup({ expressionValue.getType(), propAccessStrRef }); + if (safeType && !fieldResult.failed_or_no_value() && V(fieldResult).getType() != safeType) + { + LLVM_DEBUG(llvm::dbgs() << "\n\t...safe type fieldname: \t " + << propAccessStrRef << "." << namePtr << "type: " << expressionValue.getType() << " = " << safeType;); + auto fieldValue = V(fieldResult); + auto castValue = castToNarrowedType(location, fieldValue, safeType, genContext); + if (!castValue) + { + return mlir::failure(); + } + + return V(builder.create(location, safeType, castValue, fieldValue)); } - return mlirGenPropertyAccessExpression(location, expressionValue, namePtr, + return fieldResult; + } + + ValueOrLogicalResult MLIRGenImpl::mlirGenAssignedPropertyAccess(PropertyAccessExpression propertyAccessExpression, const GenContext &genContext) + { + // what an assignment writes is the field itself, never the value a narrowing cast from it + auto location = loc(propertyAccessExpression); + + auto result = mlirGen(propertyAccessExpression->expression.as(), genContext); + EXIT_IF_FAILED_OR_NO_VALUE(result) + + auto namePtr = MLIRHelper::getName(propertyAccessExpression->name, stringAllocator); + return mlirGenPropertyAccessExpression(location, V(result), namePtr, !!propertyAccessExpression->questionDotToken, genContext); } diff --git a/tslang/lib/TypeScript/MLIRGenImpl.h b/tslang/lib/TypeScript/MLIRGenImpl.h index 98fb27683..95a52c174 100644 --- a/tslang/lib/TypeScript/MLIRGenImpl.h +++ b/tslang/lib/TypeScript/MLIRGenImpl.h @@ -3308,7 +3308,7 @@ class MLIRGenImpl auto propAccess = expr.as(); auto objType = evaluate(propAccess->expression, genContext); LLVM_DEBUG(llvm::dbgs() << "\n!! Safe Type map for: " << nameStr << " of " << objType << " is [" << safeValue.getType() << "]\n"); - safeTypesMap.insert({ objType, nameStr }, safeValue); + safeTypesMap.insert({ objType, nameStr }, safeValue.getType()); } } } @@ -3317,6 +3317,34 @@ class MLIRGenImpl return result2; } + // the value a narrowing to `safeType` gives `exprValue`, cast the way addSafeCastStatement casts it + mlir::Value castToNarrowedType(mlir::Location location, mlir::Value exprValue, mlir::Type safeType, const GenContext &genContext) + { + auto exprType = exprValue.getType(); + if (isa(exprType)) + { + return builder.create(location, safeType, exprValue); + } + + if (auto optType = dyn_cast(exprType)) + { + if (optType.getElementType() == safeType) + { + return builder.create(location, safeType, exprValue); + } + } + + if (isa(exprType)) + { + return isa(safeType) + ? builder.create(location, safeType, exprValue).getResult() + : builder.create(location, safeType, exprValue).getResult(); + } + + auto result = cast(location, safeType, exprValue, genContext); + return result.failed_or_no_value() ? mlir::Value() : V(result); + } + mlir::LogicalResult addSafeCastStatement(mlir::Location location, StringRef parameterName, mlir::Value exprValue, mlir::Type safeType, bool inverse, ElseSafeCase* elseSafeCase, const GenContext &genContext) { mlir::Value castedValue; @@ -5459,7 +5487,9 @@ class MLIRGenImpl genContext); } - auto result = mlirGen(leftExpression, genContext); + auto result = leftExpression == SyntaxKind::PropertyAccessExpression + ? mlirGenAssignedPropertyAccess(leftExpression.as(), genContext) + : mlirGen(leftExpression, genContext); EXIT_IF_FAILED_OR_NO_VALUE(result) auto leftExpressionValue = V(result); @@ -6378,6 +6408,8 @@ class MLIRGenImpl ValueOrLogicalResult mlirGen(PropertyAccessExpression propertyAccessExpression, const GenContext &genContext); + ValueOrLogicalResult mlirGenAssignedPropertyAccess(PropertyAccessExpression propertyAccessExpression, const GenContext &genContext); + ValueOrLogicalResult mlirGenPropertyAccessExpression(mlir::Location location, mlir::Value objectValue, mlir::StringRef name, const GenContext &genContext); @@ -12787,7 +12819,9 @@ class MLIRGenImpl llvm::ScopedHashTable debugScope; - llvm::ScopedHashTable safeTypesMap; + // a narrowed field (`this.head` after `this.head !== null`) and the type it is narrowed to; a + // read of it loads the field again and casts it, so no value outlives the region it is in + llvm::ScopedHashTable safeTypesMap; // helper to get line number Parser parser; diff --git a/tslang/lib/TypeScript/MLIRGenStatements.cpp b/tslang/lib/TypeScript/MLIRGenStatements.cpp index d28fba91a..cb181ef26 100644 --- a/tslang/lib/TypeScript/MLIRGenStatements.cpp +++ b/tslang/lib/TypeScript/MLIRGenStatements.cpp @@ -173,6 +173,8 @@ namespace mlirgen auto location = loc(blockAST); SymbolTableScopeT varScope(symbolTable); + // a narrowing made in the block (after an early exit) ends with it + SafeTypesMapScopeT safeTypesMapScope(safeTypesMap); GenContext genContextUsing(genContext); genContextUsing.parentBlockContext = &genContext; @@ -245,6 +247,7 @@ namespace mlirgen builder.setInsertionPointToStart(&tryOp.getBody().front()); SymbolTableScopeT varScope(symbolTable); + SafeTypesMapScopeT safeTypesMapScope(safeTypesMap); GenContext tryBodyGenContext(tryGenContext); tryBodyGenContext.parentBlockContext = &tryGenContext; @@ -489,6 +492,38 @@ namespace mlirgen return mlir::success(); } + // whether control never runs past the statement: it returns, throws, breaks or continues on + // every path. Conservative - a statement it cannot tell about does not count. + static bool statementAlwaysExits(Statement statement) + { + switch ((SyntaxKind)statement) + { + case SyntaxKind::ReturnStatement: + case SyntaxKind::ThrowStatement: + case SyntaxKind::BreakStatement: + case SyntaxKind::ContinueStatement: + return true; + case SyntaxKind::Block: + for (auto blockStatement : statement.as()->statements) + { + if (statementAlwaysExits(blockStatement)) + { + return true; + } + } + + return false; + case SyntaxKind::IfStatement: + { + auto ifStatement = statement.as(); + return ifStatement->elseStatement && statementAlwaysExits(ifStatement->thenStatement) + && statementAlwaysExits(ifStatement->elseStatement); + } + default: + return false; + } + } + mlir::LogicalResult MLIRGenImpl::mlirGen(IfStatement ifStatementAST, const GenContext &genContext) { auto location = loc(ifStatementAST); @@ -518,6 +553,9 @@ namespace mlirgen // narrowed to `string` by `typeof x === "string"`. Under --di the narrowed variable's debug record // keeps such a cast alive into LLVM lowering even though the branch body was skipped. ElseSafeCase elseSafeCase{}; + // `if (x === null) return;` narrows x for the rest of the block, as an else branch would + auto thenExits = !hasElse && !literalValue.has_value() && genContext.funcOp + && statementAlwaysExits(ifStatementAST->thenStatement); { SymbolTableScopeT varScope(symbolTable); SafeTypesMapScopeT safeTypesMapScope(safeTypesMap); @@ -526,7 +564,7 @@ namespace mlirgen if (processIf) { // check if we do safe-cast here - checkSafeCast(ifStatementAST->expression, V(result), hasElse ? &elseSafeCase : nullptr, genContext); + checkSafeCast(ifStatementAST->expression, V(result), hasElse || thenExits ? &elseSafeCase : nullptr, genContext); auto result = mlirGen(ifStatementAST->thenStatement, genContext); EXIT_IF_FAILED(result) @@ -556,6 +594,11 @@ namespace mlirgen builder.setInsertionPointAfter(ifOp); + if (thenExits && elseSafeCase.safeType) + { + addSafeCastStatement(elseSafeCase.expr, elseSafeCase.safeType, false, nullptr, genContext); + } + return mlir::success(); } diff --git a/tslang/test/tester/CMakeLists.txt b/tslang/test/tester/CMakeLists.txt index 630d58d79..11817fefc 100644 --- a/tslang/test/tester/CMakeLists.txt +++ b/tslang/test/tester/CMakeLists.txt @@ -529,6 +529,7 @@ tslang_add_test(NAME test-compile-00-safe-cast-typeof COMMAND test-runner "${PRO tslang_add_test(NAME test-compile-00-safe-cast-while COMMAND test-runner "${PROJECT_SOURCE_DIR}/test/tester/tests/00safe_cast_while.ts") tslang_add_test(NAME test-compile-01-safe-cast-while COMMAND test-runner "${PROJECT_SOURCE_DIR}/test/tester/tests/01safe_cast_while.ts") tslang_add_test(NAME test-compile-00-safe-cast-field-access COMMAND test-runner "${PROJECT_SOURCE_DIR}/test/tester/tests/00safe_cast_field_access.ts") +tslang_add_test(NAME test-compile-00-safe-cast-early-exit COMMAND test-runner "${PROJECT_SOURCE_DIR}/test/tester/tests/00safe_cast_early_exit.ts") tslang_add_test(NAME test-compile-00-safe-cast-null-field COMMAND test-runner "${PROJECT_SOURCE_DIR}/test/tester/tests/00safe_cast_null_field.ts") tslang_add_test(NAME test-compile-00-safe-cast-else-scope COMMAND test-runner "${PROJECT_SOURCE_DIR}/test/tester/tests/00safe_cast_else_scope.ts") tslang_add_test(NAME test-compile-00-safe-cast-bug COMMAND test-runner "${PROJECT_SOURCE_DIR}/test/tester/tests/00safe_cast_bug.ts") @@ -973,6 +974,7 @@ tslang_add_test(NAME test-jit-00-safe-cast-typeof COMMAND test-runner -jit "${PR tslang_add_test(NAME test-jit-00-safe-cast-while COMMAND test-runner -jit "${PROJECT_SOURCE_DIR}/test/tester/tests/00safe_cast_while.ts") tslang_add_test(NAME test-jit-01-safe-cast-while COMMAND test-runner -jit "${PROJECT_SOURCE_DIR}/test/tester/tests/01safe_cast_while.ts") tslang_add_test(NAME test-jit-00-safe-cast-field-access COMMAND test-runner -jit "${PROJECT_SOURCE_DIR}/test/tester/tests/00safe_cast_field_access.ts") +tslang_add_test(NAME test-jit-00-safe-cast-early-exit COMMAND test-runner -jit "${PROJECT_SOURCE_DIR}/test/tester/tests/00safe_cast_early_exit.ts") tslang_add_test(NAME test-jit-00-safe-cast-null-field COMMAND test-runner -jit "${PROJECT_SOURCE_DIR}/test/tester/tests/00safe_cast_null_field.ts") tslang_add_test(NAME test-jit-00-safe-cast-else-scope COMMAND test-runner -jit "${PROJECT_SOURCE_DIR}/test/tester/tests/00safe_cast_else_scope.ts") tslang_add_test(NAME test-jit-00-safe-cast-bug COMMAND test-runner -jit "${PROJECT_SOURCE_DIR}/test/tester/tests/00safe_cast_bug.ts") @@ -1759,6 +1761,7 @@ set(TSLANG_CORPUS 00reference_ref_deref.ts 00safe_cast_bug.ts 00safe_cast_field_access.ts + 00safe_cast_early_exit.ts 00safe_cast_null_field.ts 00safe_cast_else_scope.ts 00safe_cast_typeof.ts diff --git a/tslang/test/tester/tests/00safe_cast_early_exit.ts b/tslang/test/tester/tests/00safe_cast_early_exit.ts new file mode 100644 index 000000000..9a48b2937 --- /dev/null +++ b/tslang/test/tester/tests/00safe_cast_early_exit.ts @@ -0,0 +1,103 @@ +// `if (x === null) return;` narrows x for the rest of the block, as an else branch would. A narrowed +// field is read from the field every time, so an assignment to it writes the field (it was "saving +// to constant object") and a read after the assignment sees the new value (#231). + +class N { + v: number; + next: N | null = null; + + constructor(v: number) { + this.v = v; + } +} + +class List { + head: N | null = null; + count = 0; + + first(): number { + if (this.head === null) return -1; + return this.head.v; + } + + add(v: number) { + const n = new N(v); + n.next = this.head; + this.head = n; + this.count++; + } + + remove(v: number): void { + if (this.head === null) { + return; + } + + if (this.head.v === v) { + this.head = this.head.next; + this.count--; + return; + } + + let cur = this.head; + while (cur.next !== null) { + if (cur.next.v === v) { + cur.next = cur.next.next; + this.count--; + return; + } + + cur = cur.next; + } + } + + replaceHead(v: number): number { + if (this.head !== null) { + this.head = new N(v); + // the field as it is now, not the value the test narrowed + return this.head.v; + } + + return -1; + } +} + +function valueOf(h: N | null): number { + if (h === null) { + throw "no value"; + } + + return h.v; +} + +function sum(items: (N | null)[]): number { + let s = 0; + for (const i of items) { + if (i === null) continue; + s += i.v; + } + + return s; +} + +function main() { + const l = new List(); + assert(l.first() == -1); + l.remove(3); + + l.add(1); + l.add(2); + l.add(3); + assert(l.first() == 3); + + l.remove(3); + assert(l.first() == 2 && l.count == 2); + l.remove(1); + assert(l.first() == 2 && l.count == 1); + + assert(l.replaceHead(9) == 9 && l.first() == 9); + + assert(valueOf(new N(4)) == 4); + assert(sum([new N(1), null, new N(5)]) == 6); + + print("done."); +}