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
100 changes: 100 additions & 0 deletions tslang/lib/TypeScript/MLIRGenCast.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -1565,9 +1565,109 @@ namespace mlirgen
}
}

if (mlir::failed(verifyUnrelatedClassCast(location, type, valueType)))
{
return mlir::failure();
}

return mlir::success();
}

// Two classes where neither extends the other are compatible by their members, as in
// TypeScript: `const b: B = new A()` with `A { v: number }` and `B { v: string }` compiled to a
// plain cast, and reading `b.v` read a number as a string and crashed. Like an assertion, the
// cast needs one side's fields to be found in the other with types that extend them.
mlir::LogicalResult MLIRGenImpl::verifyUnrelatedClassCast(mlir::Location location, mlir::Type type, mlir::Type valueType)
{
std::string mismatch;
if (areIncompatibleUnrelatedClasses(location, type, valueType, mismatch))
{
emitError(location, "type ") << to_print(valueType) << " is not assignable to type " << to_print(type) << ": " << mismatch;
return mlir::failure();
}

return mlir::success();
}

bool MLIRGenImpl::areIncompatibleUnrelatedClasses(mlir::Location location, mlir::Type type, mlir::Type valueType, std::string &mismatch)
{
auto classType = dyn_cast<mlir_ts::ClassType>(type);
auto valueClassType = dyn_cast<mlir_ts::ClassType>(valueType);
if (!classType || !valueClassType || classType == valueClassType
|| mth.isGenericType(classType) || mth.isGenericType(valueClassType))
{
return false;
}

auto classInfo = getClassInfoByFullName(classType.getName().getValue());
auto valueClassInfo = getClassInfoByFullName(valueClassType.getName().getValue());
if (!classInfo || !valueClassInfo
|| classInfo->hasBase(valueClassType) || valueClassInfo->hasBase(classType)
// two specializations of one generic compare by their type arguments
|| (classInfo->originClassType && classInfo->originClassType == valueClassInfo->originClassType))
{
return false;
}

// the data fields, inherited ones included; not the internal ones (.vtbl) nor the storage
// a derived class embeds for its base
std::function<void(ClassInfo::TypePtr, llvm::StringMap<mlir::Type> &)> collectFields =
[&](ClassInfo::TypePtr info, llvm::StringMap<mlir::Type> &fields) {
for (auto &base : info->baseClasses)
{
collectFields(base, fields);
}

auto storageType = dyn_cast<mlir_ts::ClassStorageType>(info->classType.getStorageType());
if (!storageType)
{
return;
}

for (auto &field : storageType.getFields())
{
auto strId = dyn_cast_or_null<mlir::StringAttr>(field.id);
if (!strId || strId.getValue().starts_with(".")
|| llvm::any_of(info->baseClasses, [&](auto &base) { return strId.getValue() == base->fullName; }))
{
continue;
}

fields[strId.getValue()] = field.type;
}
};

llvm::StringMap<mlir::Type> fields;
llvm::StringMap<mlir::Type> valueFields;
collectFields(classInfo, fields);
collectFields(valueClassInfo, valueFields);

// every field of `to` is in `from`, of a type that extends it
auto fits = [&](llvm::StringMap<mlir::Type> &from, llvm::StringMap<mlir::Type> &to, std::string &mismatch) {
for (auto &field : to)
{
auto found = from.find(field.getKey());
if (found == from.end())
{
mismatch = "'" + field.getKey().str() + "' is missing";
return false;
}

llvm::StringMap<std::pair<ts::TypeParameterDOM::TypePtr,mlir::Type>> typeParamsWithArgs;
if (!isTrue(mth.extendsType(location, found->getValue(), field.getValue(), typeParamsWithArgs)))
{
mismatch = "'" + field.getKey().str() + "' is " + to_print(found->getValue()) + ", not " + to_print(field.getValue());
return false;
}
}

return true;
};

std::string reverseMismatch;
return !fits(valueFields, fields, mismatch) && !fits(fields, valueFields, reverseMismatch);
}

ValueOrLogicalResult MLIRGenImpl::castPrimitiveTypeFromAny(mlir::Location location, mlir::Type type, mlir::Value value, const GenContext &genContext)
{
// info, we add "_" extra as scanner append "_" in front of "__";
Expand Down
17 changes: 15 additions & 2 deletions tslang/lib/TypeScript/MLIRGenImpl.h
Original file line number Diff line number Diff line change
Expand Up @@ -3441,8 +3441,19 @@ class MLIRGenImpl
return mlir::success();
}

CAST_A(result, location, safeType, exprValue, genContext);
castedValue = V(result);
// A type guard narrows an A to an unrelated class B (`isB(a)`) where the value is
// checked at run time to be one - TypeScript narrows it to A & B. The cast that
// a plain assignment would be refused is right here.
std::string mismatch;
if (areIncompatibleUnrelatedClasses(location, safeType, exprValue.getType(), mismatch))
{
castedValue = builder.create<mlir_ts::CastOp>(location, safeType, exprValue);
}
else
{
CAST_A(result, location, safeType, exprValue, genContext);
castedValue = V(result);
}
}

LLVM_DEBUG(llvm::dbgs() << "\n!! Safe Type: [" << parameterName << "] is [" << safeType << "]\n");
Expand Down Expand Up @@ -11268,6 +11279,8 @@ class MLIRGenImpl
// wrong casts
// TODO: put it into Cast::Verify
mlir::LogicalResult verifyCastCompatibility(mlir::Location location, mlir::Type type, mlir::Type valueType);
mlir::LogicalResult verifyUnrelatedClassCast(mlir::Location location, mlir::Type type, mlir::Type valueType);
bool areIncompatibleUnrelatedClasses(mlir::Location location, mlir::Type type, mlir::Type valueType, std::string &mismatch);

ValueOrLogicalResult castPrimitiveTypeFromAny(mlir::Location location, mlir::Type type, mlir::Value value, const GenContext &genContext);

Expand Down
12 changes: 12 additions & 0 deletions tslang/test/tester/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -605,6 +605,7 @@ tslang_add_test(NAME test-compile-01-symbol COMMAND test-runner "${PROJECT_SOURC
tslang_add_test(NAME test-compile-00-as COMMAND test-runner "${PROJECT_SOURCE_DIR}/test/tester/tests/00as.ts")
tslang_add_test(NAME test-compile-00-as-const COMMAND test-runner "${PROJECT_SOURCE_DIR}/test/tester/tests/00as_const.ts")
tslang_add_test(NAME test-compile-00-type-guard-function COMMAND test-runner "${PROJECT_SOURCE_DIR}/test/tester/tests/00type_guard_function.ts")
tslang_add_test(NAME test-compile-00-class-cast-structural COMMAND test-runner "${PROJECT_SOURCE_DIR}/test/tester/tests/00class_cast_structural.ts")
tslang_add_test(NAME test-compile-00-names-conflict COMMAND test-runner "${PROJECT_SOURCE_DIR}/test/tester/tests/00names_conflict.ts")
tslang_add_test(NAME test-compile-00-generic-arguments-name-conflict COMMAND test-runner "${PROJECT_SOURCE_DIR}/test/tester/tests/00generic_arguments_name_conflict.ts")
tslang_add_test(NAME test-compile-00-decorators COMMAND test-runner "${PROJECT_SOURCE_DIR}/test/tester/tests/00decorators.ts")
Expand Down Expand Up @@ -1063,6 +1064,7 @@ tslang_add_test(NAME test-jit-01-symbol COMMAND test-runner -jit "${PROJECT_SOUR
tslang_add_test(NAME test-jit-00-as COMMAND test-runner -jit "${PROJECT_SOURCE_DIR}/test/tester/tests/00as.ts")
tslang_add_test(NAME test-jit-00-as-const COMMAND test-runner -jit "${PROJECT_SOURCE_DIR}/test/tester/tests/00as_const.ts")
tslang_add_test(NAME test-jit-00-type-guard-function COMMAND test-runner -jit "${PROJECT_SOURCE_DIR}/test/tester/tests/00type_guard_function.ts")
tslang_add_test(NAME test-jit-00-class-cast-structural COMMAND test-runner -jit "${PROJECT_SOURCE_DIR}/test/tester/tests/00class_cast_structural.ts")
tslang_add_test(NAME test-jit-00-names-conflict COMMAND test-runner -jit "${PROJECT_SOURCE_DIR}/test/tester/tests/00names_conflict.ts")
tslang_add_test(NAME test-jit-00-generic-arguments-name-conflict COMMAND test-runner -jit "${PROJECT_SOURCE_DIR}/test/tester/tests/00generic_arguments_name_conflict.ts")
tslang_add_test(NAME test-jit-00-decorators COMMAND test-runner -jit "${PROJECT_SOURCE_DIR}/test/tester/tests/00decorators.ts")
Expand Down Expand Up @@ -1825,6 +1827,7 @@ set(TSLANG_CORPUS
00tuple.ts
00type_aliases_in_generics.ts
00type_guard_function.ts
00class_cast_structural.ts
00typed_array.ts
00typeof_function_narrowing.ts
00typeof_static_fold.ts
Expand Down Expand Up @@ -2563,6 +2566,15 @@ set_tests_properties(test-compile-export-all-no-locals
PROPERTIES PASS_REGULAR_EXPRESSION "let moduleLevel"
FAIL_REGULAR_EXPRESSION "namespace [.]f_|Stack dump|error:")

# A cast between unrelated classes whose fields do not fit is an error; it compiled and crashed.
add_test(NAME test-compile-class-cast-unrelated-error
COMMAND $<TARGET_FILE:tslang> --emit=obj --no-default-lib -mm=none
"${PROJECT_SOURCE_DIR}/test/tester/class-cast/unrelated.ts"
-o "${CMAKE_CURRENT_BINARY_DIR}/class-cast-unrelated.obj")
set_tests_properties(test-compile-class-cast-unrelated-error
PROPERTIES PASS_REGULAR_EXPRESSION "type A is not assignable to type B: 'v' is number, not string"
FAIL_REGULAR_EXPRESSION "Stack dump|Assertion failed")

# The shared-component tier under the other two models. A shared library records the model
# it was built under, so both halves of a pair are built with the same flag - which is what
# these run. The file pairs are the default model's, verbatim.
Expand Down
15 changes: 15 additions & 0 deletions tslang/test/tester/class-cast/unrelated.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,15 @@
// A and B are unrelated and their `v` fields have different types; the cast compiled, and reading
// b.v read a number as a string.
class A {
constructor(public v: number) {}
}

class B {
constructor(public v: string) {}
}

function main() {
const a = new A(1);
const b: B = a;
print(b.v);
}
46 changes: 46 additions & 0 deletions tslang/test/tester/tests/00class_cast_structural.ts
Original file line number Diff line number Diff line change
@@ -0,0 +1,46 @@
// Classes where neither extends the other are compatible by their fields, as in TypeScript.
class Point {
constructor(public x: number, public y: number) {}
}

class Vec {
constructor(public x: number, public y: number) {}
}

class Named {
name = "n";
}

class Dog {
bark() {
return 1;
}
}

class Cat {
meow() {
return 2;
}
}

function isCat(p: Dog | Cat): p is Cat {
return p instanceof Cat;
}

function main() {
// the same fields
const p = new Point(1, 2);
const v: Vec = p;
assert(v.x == 1 && v.y == 2, "same shape");

// no fields on either side
const d = new Dog();
if (isCat(d)) {
assert(false, "a dog is not a cat");
}

const n = new Named();
assert(n.name == "n", "named");

print("done.");
}
Loading