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
2 changes: 1 addition & 1 deletion tslang/include/TypeScript/MLIRLogic/MLIRDefines.h
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,6 @@ using SymbolTableScopeT = llvm::ScopedHashTableScope<StringRef, VariablePairT>;
using BoundRefCacheScopeT = llvm::ScopedHashTableScope<mlir::Value, mlir::Value>;

typedef std::pair<mlir::Type, StringRef> SafeTypeKeyType;
using SafeTypesMapScopeT = llvm::ScopedHashTableScope<SafeTypeKeyType, mlir::Value>;
using SafeTypesMapScopeT = llvm::ScopedHashTableScope<SafeTypeKeyType, mlir::Type>;

#endif // MLIR_TYPESCRIPT_MLIRGENLOGIC_MLIRDEFINES_H_
38 changes: 30 additions & 8 deletions tslang/lib/TypeScript/MLIRGenExpressions.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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<mlir_ts::SafeCastOp>(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<Expression>(), genContext);
EXIT_IF_FAILED_OR_NO_VALUE(result)

auto namePtr = MLIRHelper::getName(propertyAccessExpression->name, stringAllocator);
return mlirGenPropertyAccessExpression(location, V(result), namePtr,
!!propertyAccessExpression->questionDotToken, genContext);
}

Expand Down
40 changes: 37 additions & 3 deletions tslang/lib/TypeScript/MLIRGenImpl.h
Original file line number Diff line number Diff line change
Expand Up @@ -3308,7 +3308,7 @@ class MLIRGenImpl
auto propAccess = expr.as<PropertyAccessExpression>();
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());
}
}
}
Expand All @@ -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<mlir_ts::AnyType>(exprType))
{
return builder.create<mlir_ts::UnboxOp>(location, safeType, exprValue);
}

if (auto optType = dyn_cast<mlir_ts::OptionalType>(exprType))
{
if (optType.getElementType() == safeType)
{
return builder.create<mlir_ts::ValueOp>(location, safeType, exprValue);
}
}

if (isa<mlir_ts::UnionType>(exprType))
{
return isa<mlir_ts::UnionType>(safeType)
? builder.create<mlir_ts::CastOp>(location, safeType, exprValue).getResult()
: builder.create<mlir_ts::GetValueFromUnionOp>(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;
Expand Down Expand Up @@ -5459,7 +5487,9 @@ class MLIRGenImpl
genContext);
}

auto result = mlirGen(leftExpression, genContext);
auto result = leftExpression == SyntaxKind::PropertyAccessExpression
? mlirGenAssignedPropertyAccess(leftExpression.as<PropertyAccessExpression>(), genContext)
: mlirGen(leftExpression, genContext);
EXIT_IF_FAILED_OR_NO_VALUE(result)
auto leftExpressionValue = V(result);

Expand Down Expand Up @@ -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);

Expand Down Expand Up @@ -12787,7 +12819,9 @@ class MLIRGenImpl

llvm::ScopedHashTable<StringRef, mlir::LLVM::DIScopeAttr> debugScope;

llvm::ScopedHashTable<SafeTypeKeyType, mlir::Value> 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<SafeTypeKeyType, mlir::Type> safeTypesMap;

// helper to get line number
Parser parser;
Expand Down
45 changes: 44 additions & 1 deletion tslang/lib/TypeScript/MLIRGenStatements.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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;

Expand Down Expand Up @@ -245,6 +247,7 @@ namespace mlirgen
builder.setInsertionPointToStart(&tryOp.getBody().front());

SymbolTableScopeT varScope(symbolTable);
SafeTypesMapScopeT safeTypesMapScope(safeTypesMap);
GenContext tryBodyGenContext(tryGenContext);
tryBodyGenContext.parentBlockContext = &tryGenContext;

Expand Down Expand Up @@ -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<ts::Block>()->statements)
{
if (statementAlwaysExits(blockStatement))
{
return true;
}
}

return false;
case SyntaxKind::IfStatement:
{
auto ifStatement = statement.as<IfStatement>();
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);
Expand Down Expand Up @@ -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);
Expand All @@ -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)
Expand Down Expand Up @@ -556,6 +594,11 @@ namespace mlirgen

builder.setInsertionPointAfter(ifOp);

if (thenExits && elseSafeCase.safeType)
{
addSafeCastStatement(elseSafeCase.expr, elseSafeCase.safeType, false, nullptr, genContext);
}

return mlir::success();
}

Expand Down
3 changes: 3 additions & 0 deletions tslang/test/tester/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down Expand Up @@ -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")
Expand Down Expand Up @@ -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
Expand Down
103 changes: 103 additions & 0 deletions tslang/test/tester/tests/00safe_cast_early_exit.ts
Original file line number Diff line number Diff line change
@@ -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.");
}
Loading