diff --git a/docs/design/datacontracts/RuntimeTypeSystem.md b/docs/design/datacontracts/RuntimeTypeSystem.md index b6bbb7aee425b8..77b047bf63869a 100644 --- a/docs/design/datacontracts/RuntimeTypeSystem.md +++ b/docs/design/datacontracts/RuntimeTypeSystem.md @@ -215,6 +215,7 @@ public enum AsyncMethodFlags : uint IsAsyncVariant = 0x2, Thunk = 0x4, ReturnDroppingThunk = 0x8, + CovariantForwardingThunk = 0x10, } // Identifies one of the runtime's well-known singleton MethodTables, each addressable @@ -1581,6 +1582,7 @@ And the following enumeration definitions IsAsyncVariant = 0x4, Thunk = 0x10, ReturnDroppingThunk = 0x20, + CovariantForwardingThunk = 0x40, } [Flags] @@ -2027,6 +2029,8 @@ Reading a method's Runtime Async flags: result |= AsyncMethodFlags.Thunk; if ((raw & AsyncMethodFlags_1.ReturnDroppingThunk) != 0) result |= AsyncMethodFlags.ReturnDroppingThunk; + if ((raw & AsyncMethodFlags_1.CovariantForwardingThunk) != 0) + result |= AsyncMethodFlags.CovariantForwardingThunk; return result; } ``` diff --git a/src/coreclr/vm/asyncthunks.cpp b/src/coreclr/vm/asyncthunks.cpp index 4f9941f3a5937f..344411d23754bc 100644 --- a/src/coreclr/vm/asyncthunks.cpp +++ b/src/coreclr/vm/asyncthunks.cpp @@ -25,46 +25,56 @@ bool MethodDesc::TryGenerateAsyncThunk(DynamicResolver** resolver, COR_ILMETHOD_ return false; } - MethodDesc* pAsyncOtherVariant = nullptr; + MethodDesc* pThunkTarget = nullptr; if (!IsAsyncMethod()) { // a non-async thunk is implemented in terms of the async variant which has user code - pAsyncOtherVariant = this->GetAsyncVariant(); + pThunkTarget = this->GetAsyncVariant(); + } + else if (IsCovariantForwardingThunk()) + { + // this is an async variant of a method that covariantly returns a type derived from + // Task/Task. It calls the ordinary variant, which has user code, and awaits the result. + pThunkTarget = this->GetOrdinaryVariant(); } else { _ASSERTE(IsReturnDroppingThunk()); // this is a special void-returning async variant that calls // the normal async variant and drops the result - pAsyncOtherVariant = this->GetAsyncVariant(); + pThunkTarget = this->GetAsyncVariant(); } - _ASSERTE(!IsWrapperStub() && !pAsyncOtherVariant->IsWrapperStub()); + _ASSERTE(!IsWrapperStub() && !pThunkTarget->IsWrapperStub()); MetaSig msig(this); - SigTypeContext sigContext(pAsyncOtherVariant); + SigTypeContext sigContext(pThunkTarget); ILStubLinker sl( GetModule(), GetSignature(), &sigContext, - pAsyncOtherVariant, + pThunkTarget, (ILStubLinkerFlags)ILSTUB_LINKER_FLAG_NONE); if (!IsAsyncMethod()) { - EmitTaskReturningThunk(pAsyncOtherVariant, msig, &sl); + EmitTaskReturningThunk(pThunkTarget, msig, &sl); + } + else if (IsCovariantForwardingThunk()) + { + EmitCovariantForwardingThunk(pThunkTarget, msig, &sl); } else { _ASSERTE(IsReturnDroppingThunk()); - EmitReturnDroppingThunk(pAsyncOtherVariant, msig, &sl); + EmitReturnDroppingThunk(pThunkTarget, msig, &sl); } NewHolder ilResolver = new ILStubResolver(); // Initialize the resolver target details. ilResolver->SetStubMethodDesc(this); - ilResolver->SetStubTargetMethodDesc(pAsyncOtherVariant); + ilResolver->SetStubTargetMethodDesc(pThunkTarget); // Generate all IL associated data for JIT *methodILDecoder = ilResolver->FinalizeILStub(&sl); @@ -339,6 +349,13 @@ SigPointer MethodDesc::GetAsyncThunkResultTypeSig() // Task.FromResult, this returns a MethodSpec representing // Task.FromResult>. int MethodDesc::GetTokenForGenericMethodCallWithAsyncReturnType(ILCodeStream* pCode, MethodDesc* md) +{ + return GetTokenForGenericMethodCall(pCode, md, GetAsyncThunkResultTypeSig()); +} + +// Given a method Foo, return a MethodSpec token for Foo instantiated with the type +// described by typeArgSig. +int MethodDesc::GetTokenForGenericMethodCall(ILCodeStream* pCode, MethodDesc* md, SigPointer typeArgSig) { if (!md->HasClassOrMethodInstantiation()) { @@ -351,10 +368,9 @@ int MethodDesc::GetTokenForGenericMethodCallWithAsyncReturnType(ILCodeStream* pC SigBuilder methodSigBuilder; methodSigBuilder.AppendByte(IMAGE_CEE_CS_CALLCONV_GENERICINST); methodSigBuilder.AppendData(1); - SigPointer retTypeSig = GetAsyncThunkResultTypeSig(); PCCOR_SIGNATURE retTypeSigRaw; uint32_t retTypeSigLen; - retTypeSig.GetSignature(&retTypeSigRaw, &retTypeSigLen); + typeArgSig.GetSignature(&retTypeSigRaw, &retTypeSigLen); methodSigBuilder.AppendBlob((const PVOID)retTypeSigRaw, retTypeSigLen); DWORD methodSigLen; @@ -463,3 +479,93 @@ void MethodDesc::EmitReturnDroppingThunk(MethodDesc* pAsyncOtherVariant, MetaSig pCode->EmitPOP(); pCode->EmitRET(); } + +// Returns a SigPointer to the return type in the given signature. +// For example, for "int Foo(string)" this returns the signature representing (int). +static SigPointer GetReturnTypeSig(Signature signature) +{ + SigPointer pSig(signature.GetRawSig(), signature.GetRawSigLen()); + uint32_t callConvInfo; + IfFailThrow(pSig.GetCallingConvInfo(&callConvInfo)); + + if ((callConvInfo & IMAGE_CEE_CS_CALLCONV_GENERIC) != 0) + { + // GenParamCount + IfFailThrow(pSig.GetData(NULL)); + } + + // ParamCount + IfFailThrow(pSig.GetData(NULL)); + + // ReturnType comes now. Skip the modifiers (like modreqs in async signatures). + IfFailThrow(pSig.SkipCustomModifiers()); + + PCCOR_SIGNATURE retTypeSig; + uint32_t tailLength; + pSig.GetSignature(&retTypeSig, &tailLength); + + // Skip to the end of the return type so we can get the length. + IfFailThrow(pSig.SkipExactlyOne()); + + PCCOR_SIGNATURE retTypeSigEnd; + pSig.GetSignature(&retTypeSigEnd, &tailLength); + + return SigPointer(retTypeSig, (DWORD)(retTypeSigEnd - retTypeSig)); +} + +// Provided an ordinary variant that covariantly returns a type derived from Task/Task, +// emits an async variant that calls the ordinary variant and awaits the returned Task. +// A thunk is used (rather than an "async version" of this method's own IL) so that only +// methods that covariantly override a task-returning method need an extra variant; other +// overrides of the same slot keep being treated as ordinary, non-task-returning methods. +void MethodDesc::EmitCovariantForwardingThunk(MethodDesc* pOrdinaryVariant, MetaSig& msig, ILStubLinker* pSL) +{ + _ASSERTE(IsAsyncMethod() && IsAsyncVariantMethod() && IsCovariantForwardingThunk()); + _ASSERTE(!pOrdinaryVariant->IsAsyncVariantMethod()); + _ASSERTE(!IsAsyncVariantForValueTaskReturningMethod()); + + _ASSERTE(this->IsVirtual()); + _ASSERTE(pOrdinaryVariant->IsVirtual()); + _ASSERTE(msig.HasThis()); + + // Implement IL that is effectively the following: + // { + // return await this.ordinary(arg); // CALLVIRT + TransparentAwait + // } + ILCodeStream* pCode = pSL->NewCodeStream(ILStubLinker::kDispatch); + int token = GetTokenForThunkTarget(pCode, pOrdinaryVariant); + + DWORD localArg = 0; + pCode->EmitLDARG(localArg++); + for (UINT iArg = 0; iArg < msig.NumFixedArgs(); iArg++) + { + pCode->EmitLDARG(localArg++); + } + + // ordinary(arg) + // The returned type derives from Task or Task, so it can be passed to the + // matching TransparentAwait overload as-is. + pCode->EmitCALLVIRT(token, localArg, 1); + + // The await below is in tail position ("return await ..."). + pCode->EmitCALL(METHOD__ASYNC_HELPERS__TAIL_AWAIT, 0, 0); + + // await the returned Task + bool returnsVoid = msig.IsReturnTypeVoid(); + int awaitToken; + if (returnsVoid) + { + awaitToken = pCode->GetToken(CoreLibBinder::GetMethod(METHOD__ASYNC_HELPERS__TRANSPARENT_AWAIT_TASK)); + } + else + { + MethodDesc* pAwaitMD = CoreLibBinder::GetMethod(METHOD__ASYNC_HELPERS__TRANSPARENT_AWAIT_TASK_OF_T); + TypeHandle thRetType = msig.GetRetTypeHandleThrowing(); + pAwaitMD = FindOrCreateAssociatedMethodDesc(pAwaitMD, pAwaitMD->GetMethodTable(), FALSE, Instantiation(&thRetType, 1), FALSE); + awaitToken = GetTokenForGenericMethodCall(pCode, pAwaitMD, GetReturnTypeSig(GetSignature())); + } + + pCode->EmitCALL(awaitToken, 1, returnsVoid ? 0 : 1); + // return; + pCode->EmitRET(); +} diff --git a/src/coreclr/vm/method.cpp b/src/coreclr/vm/method.cpp index 845bbb749eeffd..175abdd84112b0 100644 --- a/src/coreclr/vm/method.cpp +++ b/src/coreclr/vm/method.cpp @@ -1537,13 +1537,10 @@ DWORD MethodDesc::GetAttrs() const _ASSERTE(!"If this ever fires, then this method should return HRESULT"); return 0; } - - if (IsReturnDroppingThunk()) + if (IsReturnDroppingThunk() || IsCovariantForwardingThunk()) { - // A return-dropping thunk is synthesized by the runtime and always has an implementation - - // it calls the ordinary async variant virtually and drops the result. - // The metadata method that the thunk is derived from may be abstract (i.e. when the covariant - // override that needs the thunk is abstract), but the thunk itself never is. + // These thunks are synthesized by the runtime and always have an implementation, + // even when the covariant override that needs the thunk is abstract. dwAttributes &= ~mdAbstract; } diff --git a/src/coreclr/vm/method.hpp b/src/coreclr/vm/method.hpp index 1c6e13fd9c1c73..776fa25ecd298d 100644 --- a/src/coreclr/vm/method.hpp +++ b/src/coreclr/vm/method.hpp @@ -75,6 +75,10 @@ enum class AsyncMethodFlags Thunk = 16, // A special thunk to drop return value in covariant return scenario ReturnDroppingThunk = 32, + // A special thunk for an override that covariantly returns a type derived from Task/Task. + // Such a method does not formally return Task/Task, so its IL cannot be compiled as an + // async version. The thunk calls the ordinary variant and awaits the returned Task instead. + CovariantForwardingThunk = 64, // Note: If adding more flags make sure to modify RequiresAsyncContextSaveAndRestore // The rest of the methods that are not in any of the above groups. @@ -2117,10 +2121,20 @@ class MethodDesc return hasAsyncFlags(asyncFlags, AsyncMethodFlags::ReturnDroppingThunk); } + inline bool IsCovariantForwardingThunk() const + { + LIMITED_METHOD_DAC_CONTRACT; + if (!HasAsyncMethodData()) + return false; + + AsyncMethodFlags asyncFlags = GetAddrOfAsyncMethodData()->flags; + return hasAsyncFlags(asyncFlags, AsyncMethodFlags::CovariantForwardingThunk); + } + inline bool SupportsAsyncVersionCodegen() const { LIMITED_METHOD_DAC_CONTRACT; - return IsAsyncThunkMethod() && IsAsyncVariantMethod() && !IsReturnDroppingThunk(); + return IsAsyncThunkMethod() && IsAsyncVariantMethod() && !IsReturnDroppingThunk() && !IsCovariantForwardingThunk(); } inline bool MatchesAsyncVariantLookup(AsyncVariantLookup lookup) const @@ -2376,8 +2390,10 @@ class MethodDesc bool TryGenerateUnsafeAccessor(DynamicResolver** resolver, COR_ILMETHOD_DECODER** methodILDecoder); void EmitTaskReturningThunk(MethodDesc* pAsyncCallVariant, MetaSig& thunkMsig, ILStubLinker* pSL); void EmitReturnDroppingThunk(MethodDesc* pAsyncOtherVariant, MetaSig& msig, ILStubLinker* pSL); + void EmitCovariantForwardingThunk(MethodDesc* pOrdinaryVariant, MetaSig& msig, ILStubLinker* pSL); int GetTokenForThunkTarget(ILCodeStream* pCode, MethodDesc* md); int GetTokenForGenericMethodCallWithAsyncReturnType(ILCodeStream* pCode, MethodDesc* md); + int GetTokenForGenericMethodCall(ILCodeStream* pCode, MethodDesc* md, SigPointer typeArgSig); public: SigPointer GetAsyncThunkResultTypeSig(); static void CreateDerivedTargetSig(MetaSig& msig, SigBuilder* stubSigBuilder); diff --git a/src/coreclr/vm/methodtablebuilder.cpp b/src/coreclr/vm/methodtablebuilder.cpp index 9372036a62f67b..10ee14ffd6f0a6 100644 --- a/src/coreclr/vm/methodtablebuilder.cpp +++ b/src/coreclr/vm/methodtablebuilder.cpp @@ -2655,6 +2655,479 @@ HRESULT MethodTableBuilder::FindMethodDeclarationForMethodImpl( return hr; } +//--------------------------------------------------------------------------------------- +// +// Given the type arguments of a generic instantiation (a sequence of "cInstArgs" types), returns +// the signature of the type argument at the given index. +// +static bool TryGetInstantiationArg( + SigParser instArgs, + DWORD cInstArgs, + DWORD index, + PCCOR_SIGNATURE* ppArg, + DWORD* pcbArg) +{ + STANDARD_VM_CONTRACT; + + if (index >= cInstArgs) + return false; + + for (DWORD i = 0; i < index; i++) + { + if (FAILED(instArgs.SkipExactlyOne())) + return false; + } + + PCCOR_SIGNATURE pArgStart = instArgs.GetPtr(); + if (FAILED(instArgs.SkipExactlyOne())) + return false; + + *ppArg = pArgStart; + *pcbArg = (DWORD)(instArgs.GetPtr() - pArgStart); + return true; +} + +//--------------------------------------------------------------------------------------- +// +// Copies one type signature from pSrc to pDst, replacing references to type variables +// (ELEMENT_TYPE_VAR) with the corresponding type arguments from the given instantiation. +// The instantiation is expected to be in the scope where the resulting signature will be used, +// thus the type arguments are copied verbatim. +// +// Returns false if the signature contains constructs that are not supported here. +// +static bool CopyTypeSigWithSubstitution( + SigParser* pSrc, + SigBuilder* pDst, + SigParser instArgs, + DWORD cInstArgs) +{ + STANDARD_VM_CONTRACT; + + BYTE type; + if (FAILED(pSrc->GetByte(&type))) + return false; + + switch (type) + { + case ELEMENT_TYPE_VOID: + case ELEMENT_TYPE_BOOLEAN: + case ELEMENT_TYPE_CHAR: + case ELEMENT_TYPE_I1: + case ELEMENT_TYPE_U1: + case ELEMENT_TYPE_I2: + case ELEMENT_TYPE_U2: + case ELEMENT_TYPE_I4: + case ELEMENT_TYPE_U4: + case ELEMENT_TYPE_I8: + case ELEMENT_TYPE_U8: + case ELEMENT_TYPE_R4: + case ELEMENT_TYPE_R8: + case ELEMENT_TYPE_STRING: + case ELEMENT_TYPE_OBJECT: + case ELEMENT_TYPE_TYPEDBYREF: + case ELEMENT_TYPE_I: + case ELEMENT_TYPE_U: + pDst->AppendElementType((CorElementType)type); + return true; + + case ELEMENT_TYPE_CLASS: + case ELEMENT_TYPE_VALUETYPE: + { + mdToken token; + if (FAILED(pSrc->GetToken(&token))) + return false; + + pDst->AppendElementType((CorElementType)type); + pDst->AppendToken(token); + return true; + } + + case ELEMENT_TYPE_VAR: + { + // A type variable of the type that declares the overridden method. + // Replace it with the corresponding type argument. + uint32_t index; + if (FAILED(pSrc->GetData(&index))) + return false; + + PCCOR_SIGNATURE pArg; + DWORD cbArg; + if (!TryGetInstantiationArg(instArgs, cInstArgs, index, &pArg, &cbArg)) + return false; + + pDst->AppendBlob((PVOID)pArg, cbArg); + return true; + } + + case ELEMENT_TYPE_MVAR: + { + // A type variable of the method itself. The overriding method has the same + // type parameters, so such references can be copied as-is. + uint32_t index; + if (FAILED(pSrc->GetData(&index))) + return false; + + pDst->AppendElementType((CorElementType)type); + pDst->AppendData(index); + return true; + } + + case ELEMENT_TYPE_SZARRAY: + case ELEMENT_TYPE_PTR: + pDst->AppendElementType((CorElementType)type); + return CopyTypeSigWithSubstitution(pSrc, pDst, instArgs, cInstArgs); + + case ELEMENT_TYPE_CMOD_REQD: + case ELEMENT_TYPE_CMOD_OPT: + { + mdToken token; + if (FAILED(pSrc->GetToken(&token))) + return false; + + pDst->AppendElementType((CorElementType)type); + pDst->AppendToken(token); + return CopyTypeSigWithSubstitution(pSrc, pDst, instArgs, cInstArgs); + } + + case ELEMENT_TYPE_ARRAY: + { + pDst->AppendElementType((CorElementType)type); + if (!CopyTypeSigWithSubstitution(pSrc, pDst, instArgs, cInstArgs)) + return false; + + uint32_t rank; + if (FAILED(pSrc->GetData(&rank))) + return false; + pDst->AppendData(rank); + + if (rank != 0) + { + uint32_t nsizes; + if (FAILED(pSrc->GetData(&nsizes))) + return false; + pDst->AppendData(nsizes); + + while (nsizes--) + { + uint32_t size; + if (FAILED(pSrc->GetData(&size))) + return false; + pDst->AppendData(size); + } + + uint32_t nlbounds; + if (FAILED(pSrc->GetData(&nlbounds))) + return false; + pDst->AppendData(nlbounds); + + while (nlbounds--) + { + PCCOR_SIGNATURE pLowerBound = pSrc->GetPtr(); + if (FAILED(pSrc->GetData(NULL))) + return false; + pDst->AppendBlob((PVOID)pLowerBound, pSrc->GetPtr() - pLowerBound); + } + } + + return true; + } + + case ELEMENT_TYPE_GENERICINST: + { + BYTE classOrValueType; + if (FAILED(pSrc->GetByte(&classOrValueType))) + return false; + + if ((classOrValueType != ELEMENT_TYPE_CLASS) && (classOrValueType != ELEMENT_TYPE_VALUETYPE)) + return false; + + mdToken token; + if (FAILED(pSrc->GetToken(&token))) + return false; + + uint32_t argCnt; + if (FAILED(pSrc->GetData(&argCnt))) + return false; + + pDst->AppendElementType((CorElementType)type); + pDst->AppendElementType((CorElementType)classOrValueType); + pDst->AppendToken(token); + pDst->AppendData(argCnt); + + while (argCnt--) + { + if (!CopyTypeSigWithSubstitution(pSrc, pDst, instArgs, cInstArgs)) + return false; + } + + return true; + } + + default: + // Anything else (function pointers, ...) is not supported here. + return false; + } +} + +//--------------------------------------------------------------------------------------- +// +// A MethodImpl declaration may be a MethodDef token, in which case it refers to the method on the +// generic type definition of the declaring type and carries no instantiation of its own. +// This helper recovers the instantiation of the declaring type by walking the inheritance chain +// from the type that is being built up to the declaring type, composing the instantiations found +// in the "extends" clauses. The resulting type arguments are expressed in the scope of the type +// that is being built and thus can be used in its signatures as-is. +// +// The type arguments are stored in the provided CQuickBytes buffer and are only valid for as long +// as the buffer is alive. +// +// Returns false if the declaring type could not be reached without loading types +// (i.e. it is in another module) or if the case is not supported. +// +static bool TryGetDeclaringTypeInstantiation( + IMDInternalImport* pMDInternalImport, + mdTypeDef tkImplType, + DWORD cImplTypeArgs, + mdTypeDef tkDeclType, + CQuickBytes* pInstBuffer, + SigParser* pInstArgs, + DWORD* pcInstArgs) +{ + STANDARD_VM_CONTRACT; + + // The instantiation of the type that is being built, in its own scope, is the identity - !0, !1, ... + SigBuilder curInstBuilder; + for (DWORD i = 0; i < cImplTypeArgs; i++) + { + curInstBuilder.AppendElementType(ELEMENT_TYPE_VAR); + curInstBuilder.AppendData(i); + } + + DWORD cbCurInst; + PCCOR_SIGNATURE pCurInst = (PCCOR_SIGNATURE)curInstBuilder.GetSignature(&cbCurInst); + DWORD cCurInst = cImplTypeArgs; + + CQuickBytes curInstBuffer; + memcpy(curInstBuffer.AllocThrows(cbCurInst), pCurInst, cbCurInst); + pCurInst = (PCCOR_SIGNATURE)curInstBuffer.Ptr(); + + mdTypeDef tkType = tkImplType; + while (tkType != tkDeclType) + { + mdToken tkExtends; + if (FAILED(pMDInternalImport->GetTypeDefProps(tkType, NULL, &tkExtends))) + return false; + + if (IsNilToken(tkExtends)) + return false; + + if (TypeFromToken(tkExtends) == mdtTypeDef) + { + // A non-generic base type - the instantiation becomes empty. + tkType = tkExtends; + cbCurInst = 0; + cCurInst = 0; + continue; + } + + if (TypeFromToken(tkExtends) != mdtTypeSpec) + { + // Either the base type is in another module (mdtTypeRef), in which case it cannot declare + // a method referred to by a MethodDef token, or the hierarchy ended (mdTypeDefNil). + return false; + } + + PCCOR_SIGNATURE pTypeSpecSig; + ULONG cbTypeSpecSig; + if (FAILED(pMDInternalImport->GetTypeSpecFromToken(tkExtends, &pTypeSpecSig, &cbTypeSpecSig))) + return false; + + // GENERICINST CLASS + SigParser typeSpecSig(pTypeSpecSig, cbTypeSpecSig); + BYTE elemType; + if (FAILED(typeSpecSig.GetByte(&elemType)) || (elemType != ELEMENT_TYPE_GENERICINST)) + return false; + + if (FAILED(typeSpecSig.GetByte(&elemType)) || (elemType != ELEMENT_TYPE_CLASS)) + return false; + + mdToken tkBase; + if (FAILED(typeSpecSig.GetToken(&tkBase)) || (TypeFromToken(tkBase) != mdtTypeDef)) + return false; + + uint32_t argCnt; + if (FAILED(typeSpecSig.GetData(&argCnt))) + return false; + + // The type arguments of the base type may refer to the type variables of the current type, + // which are expressed in terms of the type that is being built by the current instantiation. + SigBuilder nextInstBuilder; + for (uint32_t i = 0; i < argCnt; i++) + { + SigParser curInstArgs(pCurInst, cbCurInst); + if (!CopyTypeSigWithSubstitution(&typeSpecSig, &nextInstBuilder, curInstArgs, cCurInst)) + return false; + } + + DWORD cbNextInst; + PCCOR_SIGNATURE pNextInst = (PCCOR_SIGNATURE)nextInstBuilder.GetSignature(&cbNextInst); + + memcpy(curInstBuffer.AllocThrows(cbNextInst), pNextInst, cbNextInst); + pCurInst = (PCCOR_SIGNATURE)curInstBuffer.Ptr(); + cbCurInst = cbNextInst; + cCurInst = argCnt; + tkType = tkBase; + } + + memcpy(pInstBuffer->AllocThrows(cbCurInst), pCurInst, cbCurInst); + *pInstArgs = SigParser((PCCOR_SIGNATURE)pInstBuffer->Ptr(), cbCurInst); + *pcInstArgs = cCurInst; + return true; +} + +//--------------------------------------------------------------------------------------- +// +// Task and Task are not sealed, thus a covariant override may return a type that derives from +// Task/Task, while not being Task/Task itself. Such a method is not Task-returning on its own, +// but the method that it overrides may well be. Since the overridden method has an Async variant, +// the override must have one as well, or it would not be able to override it. +// +// This helper checks whether the given MethodImpl declaration is a Task-returning method and, +// if so, produces the signature of the "element" type of its return type - the type that the +// Async variant of the overriding method must return. +// A NULL element signature means that the Async variant returns void (the declaration returns Task). +// The element signature, when not NULL, is built in the provided SigBuilder and thus is only valid +// for as long as the SigBuilder is alive. +// +// Returns false if the declaration is not Task-returning or if the case is not supported. +// +static bool TryGetCovariantOverrideAsyncVariantReturnType( + IMDInternalImport* pMDInternalImport, + Module* pModule, + mdTypeDef tkImplType, + DWORD cImplTypeArgs, + mdToken tkDecl, + SigBuilder* pElementSigBuilder, + PCCOR_SIGNATURE* ppElementSig, + ULONG* pcbElementSig) +{ + STANDARD_VM_CONTRACT; + + *ppElementSig = NULL; + *pcbElementSig = 0; + + PCCOR_SIGNATURE pSigDecl = NULL; + ULONG cbSigDecl = 0; + + // Type arguments of the type that declares the overridden method, if that type is generic. + // The type arguments are expressed in the scope of the overriding type and thus can be used + // in its signatures as-is. + SigParser declInstArgs; + DWORD cDeclInstArgs = 0; + CQuickBytes declInstBuffer; + + if (TypeFromToken(tkDecl) == mdtMethodDef) + { + if (FAILED(pMDInternalImport->GetSigOfMethodDef(tkDecl, &cbSigDecl, &pSigDecl))) + return false; + + // A MethodDef declaration refers to the method on the generic type definition of the + // declaring type, so, when that type is generic, the type arguments must be recovered + // from the inheritance chain of the type that is being built. + mdTypeDef tkDeclType; + if (FAILED(pMDInternalImport->GetParentToken(tkDecl, &tkDeclType)) || + (TypeFromToken(tkDeclType) != mdtTypeDef)) + { + return false; + } + + // If the instantiation cannot be recovered, continue with an empty one - it is only + // needed if the return type of the declaration actually refers to the type variables + // of the declaring type. + TryGetDeclaringTypeInstantiation( + pMDInternalImport, tkImplType, cImplTypeArgs, tkDeclType, &declInstBuffer, &declInstArgs, &cDeclInstArgs); + } + else if (TypeFromToken(tkDecl) == mdtMemberRef) + { + mdToken tkParent; + if (FAILED(pMDInternalImport->GetParentToken(tkDecl, &tkParent))) + return false; + + if (TypeFromToken(tkParent) == mdtTypeSpec) + { + // The overridden method is declared by an instantiated generic type. Its signature may + // refer to the type parameters of that type, so we will need to substitute the type + // arguments from the instantiation. + PCCOR_SIGNATURE pTypeSpecSig; + ULONG cbTypeSpecSig; + if (FAILED(pMDInternalImport->GetTypeSpecFromToken(tkParent, &pTypeSpecSig, &cbTypeSpecSig))) + return false; + + // GENERICINST (CLASS | VALUETYPE) + SigParser typeSpecSig(pTypeSpecSig, cbTypeSpecSig); + BYTE elemType; + if (FAILED(typeSpecSig.GetByte(&elemType)) || (elemType != ELEMENT_TYPE_GENERICINST)) + return false; + + if (FAILED(typeSpecSig.SkipExactlyOne())) // the generic type + return false; + + uint32_t argCnt; + if (FAILED(typeSpecSig.GetData(&argCnt))) + return false; + + declInstArgs = typeSpecSig; + cDeclInstArgs = argCnt; + } + + LPCSTR szDeclName; + if (FAILED(pMDInternalImport->GetNameAndSigOfMemberRef(tkDecl, &pSigDecl, &cbSigDecl, &szDeclName))) + return false; + } + else + { + return false; + } + + ULONG declOffsetOfAsyncDetails = 0; + ULONG declElementTypeLength = 0; + bool declReturnsValueTask = false; + MethodReturnKind declReturnKind = ClassifyMethodReturnKind( + SigPointer(pSigDecl, cbSigDecl), pModule, &declOffsetOfAsyncDetails, &declElementTypeLength, &declReturnsValueTask); + + // ValueTask and ValueTask are structs, so they cannot be base types of a covariant return type. + if (declReturnsValueTask) + return false; + + if (declReturnKind == MethodReturnKind::NonGenericTaskReturningMethod) + { + // "Task"-returning declaration. The Async variant returns void. + return true; + } + + if (declReturnKind == MethodReturnKind::GenericTaskReturningMethod) + { + // "Task"-returning declaration. The Async variant returns T. + // E_T_GENERICINST E_T_CLASS 1 + ULONG taskTokenLen = CorSigUncompressedDataSize(&pSigDecl[declOffsetOfAsyncDetails + 2]); + SigParser elementSig(pSigDecl + declOffsetOfAsyncDetails + 2 + taskTokenLen + 1, declElementTypeLength); + + // T may refer to the type parameters of the declaring type, in which case we need to + // substitute the corresponding type arguments to make the signature usable in the scope + // of the overriding method. + if (!CopyTypeSigWithSubstitution(&elementSig, pElementSigBuilder, declInstArgs, cDeclInstArgs)) + return false; + + DWORD cbElementSig; + *ppElementSig = (PCCOR_SIGNATURE)pElementSigBuilder->GetSignature(&cbElementSig); + *pcbElementSig = cbElementSig; + return true; + } + + return false; +} + //--------------------------------------------------------------------------------------- // // Used by BuildMethodTable @@ -3339,6 +3812,42 @@ MethodTableBuilder::EnumerateClassMethods() } } + // A covariant override of a Task-returning method may return a type derived from + // Task/Task and thus not be Task-returning itself. We still need to treat it as + // Task-returning, so that it gets an Async variant that overrides the Async variant + // of the overridden method. The Async variant is always a thunk in such case, since + // the method itself does not formally return a Task and thus cannot be async. + // The return type of the Async variant is the "element" type of the overridden method - + // void when the overridden method returns Task and T when it returns Task. + bool isCovariantTaskOverride = false; + SigBuilder covariantElementSigBuilder; + PCCOR_SIGNATURE pCovariantElementSig = NULL; + ULONG cbCovariantElementSig = 0; + if (bmtMetaData->fHasCovariantOverride && + (implType == METHOD_IMPL) && + !IsTaskReturning(returnKind) && + !IsMiAsync(dwImplFlags) && + IsMdVirtual(dwMemberAttrs)) + { + for (DWORD impls = 0; impls < bmtMethod->dwNumberMethodImpls; impls++) + { + if ((bmtMetaData->rgMethodImplTokens[impls].methodBody == tok) && + bmtMetaData->rgMethodImplTokens[impls].fRequiresCovariantReturnTypeChecking) + { + isCovariantTaskOverride = TryGetCovariantOverrideAsyncVariantReturnType( + pMDInternalImport, + GetModule(), + GetCl(), + bmtGenerics->GetNumGenericArgs(), + bmtMetaData->rgMethodImplTokens[impls].methodDecl, + &covariantElementSigBuilder, + &pCovariantElementSig, + &cbCovariantElementSig); + break; + } + } + } + // For delegates we don't allow any non-runtime implemented bodies // for any of the four special methods if (IsDelegate() && !IsMiRuntime(dwImplFlags)) @@ -3375,7 +3884,7 @@ MethodTableBuilder::EnumerateClassMethods() type, implType); - if (IsTaskReturning(returnKind)) + if (IsTaskReturning(returnKind) || isCovariantTaskOverride) { // Declare a TaskReturning variant method. // In the next pass we will also add an AsyncCall variant that can be called by async @@ -3418,8 +3927,9 @@ MethodTableBuilder::EnumerateClassMethods() ULONG taskTypePrefixSize; ULONG taskTypePrefixReplacementSize; + DWORD asyncMemberAttrs = dwMemberAttrs; AsyncMethodFlags asyncFlags = (AsyncMethodFlags::AsyncCall | AsyncMethodFlags::IsAsyncVariant); - if (returnsValueTask) + if (returnsValueTask && !isCovariantTaskOverride) { asyncFlags |= AsyncMethodFlags::IsAsyncVariantForValueTask; } @@ -3431,13 +3941,45 @@ MethodTableBuilder::EnumerateClassMethods() if (insertCount == 2) asyncFlags |= (AsyncMethodFlags::Thunk | AsyncMethodFlags::ReturnDroppingThunk); + if (isCovariantTaskOverride) + { + // The method itself does not return the well-known Task/Task, so its IL cannot be + // compiled as an async version. The async variant is a thunk that calls the ordinary + // variant and awaits the returned Task. + _ASSERTE(hasAsyncFlags(asyncFlags, AsyncMethodFlags::Thunk)); + asyncFlags |= AsyncMethodFlags::CovariantForwardingThunk; + // The forwarding thunk is concrete even when the ordinary variant is abstract. + asyncMemberAttrs &= ~mdAbstract; + } + // Here we construct the signature of async call variant given its task-returning counterpart. // It is basically just removing the Task/ValueTask part of the return type and keeping // the token for T or inserting void instead. // The rest of the signature stays exactly the same. ULONG taskTokenLen = 0; - if (insertCount == 2) + if (isCovariantTaskOverride) + { + // The method returns a type derived from Task/Task and overrides a + // Task-returning method. The async variant returns the "element" type of + // the overridden method. + + // from ". . . MyTask . . . Method(args);" we construct + // ". . . tk . . . Method(args);" + // (or "void" instead of "tk" when the overridden method returns Task) + + // Compute the size of the return type that we are replacing. + SigParser ownReturnType(pMemberSignature + offsetOfAsyncDetails, cMemberSignature - offsetOfAsyncDetails); + IfFailThrow(ownReturnType.SkipExactlyOne()); + taskTypePrefixSize = (ULONG)(ownReturnType.GetPtr() - (pMemberSignature + offsetOfAsyncDetails)); + + taskTypePrefixReplacementSize = (cbCovariantElementSig == 0) ? + 1 : // ELEMENT_TYPE_VOID + cbCovariantElementSig; + + cAsyncThunkMemberSignature = cMemberSignature - taskTypePrefixSize + taskTypePrefixReplacementSize; + } + else if (insertCount == 2) { // This is a rare case when we need two async variants and this is the second one. // The need arises when a Task-returning method has a Task returning virtual override. @@ -3498,7 +4040,18 @@ MethodTableBuilder::EnumerateClassMethods() _ASSERTE((cMemberSignature - originalRemainingSigOffset) == (cAsyncThunkMemberSignature - newRemainingSigOffset)); memcpy(pNewMemberSignature + newRemainingSigOffset, pMemberSignature + originalRemainingSigOffset, cMemberSignature - originalRemainingSigOffset); - if (returnKind == MethodReturnKind::NonGenericTaskReturningMethod || insertCount == 2) + if (isCovariantTaskOverride) + { + if (cbCovariantElementSig == 0) + { + pNewMemberSignature[newRemainingSigOffset - 1] = ELEMENT_TYPE_VOID; + } + else + { + memcpy(pNewMemberSignature + offsetOfAsyncDetails, pCovariantElementSig, cbCovariantElementSig); + } + } + else if (returnKind == MethodReturnKind::NonGenericTaskReturningMethod || insertCount == 2) { pNewMemberSignature[newRemainingSigOffset - 1] = ELEMENT_TYPE_VOID; } @@ -3516,7 +4069,7 @@ MethodTableBuilder::EnumerateClassMethods() pNewMethod = new (GetStackingAllocator()) bmtMDMethod( bmtInternal->pType, tok, - dwMemberAttrs, + asyncMemberAttrs, dwImplFlags, dwMethodRVA, newMemberSig, @@ -3546,13 +4099,20 @@ MethodTableBuilder::EnumerateClassMethods() bmtVT->dwMaxVtableSize++; // Increment the number of non-abstract declared methods - if (!IsMdAbstract(dwMemberAttrs)) + if (!IsMdAbstract(pNewMethod->GetDeclAttrs())) { bmtMethod->dwNumDeclaredNonAbstractMethods++; } // Normal methods only insert a single method - if (!IsTaskReturning(returnKind)) + if (!IsTaskReturning(returnKind) && !isCovariantTaskOverride) + { + break; + } + + // A covariant override of a Task-returning method needs exactly one async variant - + // the one that matches the async variant of the overridden method. + if (isCovariantTaskOverride && (insertCount == 1)) { break; } @@ -6217,15 +6777,22 @@ MethodTableBuilder::bmtMethodHandle MethodTableBuilder::FindDeclMethodOnClassInH if (variantLookup != AsyncVariantLookup::Ordinary) { - if (pCurMD->ReturnsTaskOrValueTask()) - { - pCurMD = pCurMD->GetAsyncVariant(); - } - else + // NOTE: we cannot use GetAsyncVariant() here. Fetching an associated MethodDesc + // of a generic method may create one and that may load types, which is not + // allowed while building a method table. We only need the slot of the + // variant, so the MethodDesc introduced by the declaring type will do. + MethodDesc* pVariantMD = pCurMD->ReturnsTaskOrValueTask() ? + pCurMD->GetMethodTable()->GetParallelMethodDesc(pCurMD, variantLookup) : + NULL; + + if (pVariantMD == NULL) { + // Other variant may not exist. For example we return Task and the base is generic and returns T. declMethod = {}; break; } + + pCurMD = pVariantMD; } declMethod = (*bmtParent->pSlotTable)[pCurMD->GetSlot()].Decl(); diff --git a/src/coreclr/vm/readytoruninfo.cpp b/src/coreclr/vm/readytoruninfo.cpp index f4646c8b0271c6..37b71015e9a90d 100644 --- a/src/coreclr/vm/readytoruninfo.cpp +++ b/src/coreclr/vm/readytoruninfo.cpp @@ -1324,12 +1324,12 @@ PCODE ReadyToRunInfo::GetEntryPoint(MethodDesc * pMD, PrepareCodeConfig* pConfig if (ReadyToRunCodeDisabled()) goto done; - // Return-dropping async thunks are VM-synthesized. They share the same metadata - // token and signature shape as a regular async variant, so the R2R lookup below + // Return-dropping and covariant forwarding async thunks are VM-synthesized. They share the + // same metadata token and signature shape as a regular async variant, so the R2R lookup below // would incorrectly bind the thunk to the non-thunk's compiled code and bypass // the virtual dispatch the thunk performs. Crossgen2 does not emit R2R code for // these thunks; fall back to transient IL generation in the prestub. - if (pMD->IsReturnDroppingThunk()) + if (pMD->IsReturnDroppingThunk() || pMD->IsCovariantForwardingThunk()) goto done; ETW::MethodLog::GetR2RGetEntryPointStart(pMD); diff --git a/src/native/managed/cdac/Microsoft.Diagnostics.DataContractReader.Abstractions/Contracts/IRuntimeTypeSystem.cs b/src/native/managed/cdac/Microsoft.Diagnostics.DataContractReader.Abstractions/Contracts/IRuntimeTypeSystem.cs index 069f7c266f0aa8..687e311b67bc13 100644 --- a/src/native/managed/cdac/Microsoft.Diagnostics.DataContractReader.Abstractions/Contracts/IRuntimeTypeSystem.cs +++ b/src/native/managed/cdac/Microsoft.Diagnostics.DataContractReader.Abstractions/Contracts/IRuntimeTypeSystem.cs @@ -102,6 +102,7 @@ public enum AsyncMethodFlags : uint IsAsyncVariant = 0x2, Thunk = 0x4, ReturnDroppingThunk = 0x8, + CovariantForwardingThunk = 0x10, } public enum WellKnownMethodTable diff --git a/src/native/managed/cdac/Microsoft.Diagnostics.DataContractReader.Contracts/Contracts/RuntimeTypeSystem_1.cs b/src/native/managed/cdac/Microsoft.Diagnostics.DataContractReader.Contracts/Contracts/RuntimeTypeSystem_1.cs index 200ce0ec9c4899..0f7701c76f99b9 100644 --- a/src/native/managed/cdac/Microsoft.Diagnostics.DataContractReader.Contracts/Contracts/RuntimeTypeSystem_1.cs +++ b/src/native/managed/cdac/Microsoft.Diagnostics.DataContractReader.Contracts/Contracts/RuntimeTypeSystem_1.cs @@ -169,6 +169,7 @@ internal enum AsyncMethodFlags_1 : uint IsAsyncVariant = 0x4, Thunk = 0x10, ReturnDroppingThunk = 0x20, + CovariantForwardingThunk = 0x40, } [Flags] @@ -2306,6 +2307,8 @@ public AsyncMethodFlags GetAsyncMethodFlags(MethodDescHandle methodDescHandle) result |= AsyncMethodFlags.Thunk; if ((raw & AsyncMethodFlags_1.ReturnDroppingThunk) != 0) result |= AsyncMethodFlags.ReturnDroppingThunk; + if ((raw & AsyncMethodFlags_1.CovariantForwardingThunk) != 0) + result |= AsyncMethodFlags.CovariantForwardingThunk; return result; } diff --git a/src/native/managed/cdac/Microsoft.Diagnostics.DataContractReader.Legacy/Dbi/DacDbiImpl.cs b/src/native/managed/cdac/Microsoft.Diagnostics.DataContractReader.Legacy/Dbi/DacDbiImpl.cs index 86b12f4bea666b..e4c04b98aa1b0d 100644 --- a/src/native/managed/cdac/Microsoft.Diagnostics.DataContractReader.Legacy/Dbi/DacDbiImpl.cs +++ b/src/native/managed/cdac/Microsoft.Diagnostics.DataContractReader.Legacy/Dbi/DacDbiImpl.cs @@ -2328,7 +2328,9 @@ private TargetPointer GetAmbientSP(IStackDataFrameHandle handle, CodeBlockHandle private static bool IsDiagnosticsHidden(AsyncMethodFlags flags) => flags.HasFlag(AsyncMethodFlags.Thunk) - && (flags.HasFlag(AsyncMethodFlags.ReturnDroppingThunk) || !flags.HasFlag(AsyncMethodFlags.IsAsyncVariant)); + && (flags.HasFlag(AsyncMethodFlags.ReturnDroppingThunk) + || flags.HasFlag(AsyncMethodFlags.CovariantForwardingThunk) + || !flags.HasFlag(AsyncMethodFlags.IsAsyncVariant)); private static bool HasClassOrMethodInstantiation(IRuntimeTypeSystem rts, MethodDescHandle md) { diff --git a/src/native/managed/cdac/tests/UnitTests/MethodDescTests.cs b/src/native/managed/cdac/tests/UnitTests/MethodDescTests.cs index 43a31e7bfa9e9f..48cd09e1aff5a0 100644 --- a/src/native/managed/cdac/tests/UnitTests/MethodDescTests.cs +++ b/src/native/managed/cdac/tests/UnitTests/MethodDescTests.cs @@ -620,6 +620,7 @@ public void MethodDescClassificationFlags(MockTarget.Architecture arch) TargetPointer asyncThunkWithNativeCodeSlotMethod = TargetPointer.Null; TargetPointer asyncVariantThunkMethod = TargetPointer.Null; TargetPointer asyncReturnDroppingThunkMethod = TargetPointer.Null; + TargetPointer asyncCovariantForwardingThunkMethod = TargetPointer.Null; IRuntimeTypeSystem rts = CreateRuntimeTypeSystemContract(arch, methodDescBuilder => { @@ -773,6 +774,23 @@ public void MethodDescClassificationFlags(MockTarget.Architecture arch) int asyncDataOffset = (int)(md.Address - chunk.Address) + (int)mdBaseSize; helpers.Write(chunk.Memory.Span.Slice(asyncDataOffset, sizeof(uint)), (uint)(RuntimeTypeSystem_1.AsyncMethodFlags_1.Thunk | RuntimeTypeSystem_1.AsyncMethodFlags_1.IsAsyncVariant | RuntimeTypeSystem_1.AsyncMethodFlags_1.ReturnDroppingThunk)); } + + // Async covariant forwarding thunk method + { + uint mdBaseSize = (uint)methodDescBuilder.MethodDescLayout.Size; + uint totalSize = mdBaseSize + methodDescBuilder.AsyncMethodDataSize; + byte chunkSize = (byte)(totalSize / methodDescBuilder.MethodDescAlignment); + MockMethodDescChunk chunk = methodDescBuilder.AddMethodDescChunk("asyncCovariantForwardingThunk", chunkSize); + chunk.MethodTable = methodTable.Value; + chunk.Size = chunkSize; + chunk.Count = 1; + MockMethodDesc md = chunk.GetMethodDescAtChunkIndex(0, methodDescBuilder.MethodDescLayout); + md.Flags = (ushort)((ushort)MethodClassification.IL | (ushort)MethodDescFlags_1.MethodDescFlags.HasAsyncMethodData); + md.Slot = 5; + asyncCovariantForwardingThunkMethod = new TargetPointer(md.Address); + int asyncDataOffset = (int)(md.Address - chunk.Address) + (int)mdBaseSize; + helpers.Write(chunk.Memory.Span.Slice(asyncDataOffset, sizeof(uint)), (uint)(RuntimeTypeSystem_1.AsyncMethodFlags_1.Thunk | RuntimeTypeSystem_1.AsyncMethodFlags_1.IsAsyncVariant | RuntimeTypeSystem_1.AsyncMethodFlags_1.CovariantForwardingThunk)); + } }); // Normal IL method: not hidden by any primitive @@ -843,6 +861,12 @@ public void MethodDescClassificationFlags(MockTarget.Architecture arch) MethodDescHandle handle = rts.GetMethodDescHandle(asyncReturnDroppingThunkMethod); Assert.Equal(AsyncMethodFlags.Thunk | AsyncMethodFlags.IsAsyncVariant | AsyncMethodFlags.ReturnDroppingThunk, rts.GetAsyncMethodFlags(handle)); } + + // Async covariant forwarding thunk: Thunk | IsAsyncVariant | CovariantForwardingThunk + { + MethodDescHandle handle = rts.GetMethodDescHandle(asyncCovariantForwardingThunkMethod); + Assert.Equal(AsyncMethodFlags.Thunk | AsyncMethodFlags.IsAsyncVariant | AsyncMethodFlags.CovariantForwardingThunk, rts.GetAsyncMethodFlags(handle)); + } } [Theory] diff --git a/src/tests/async/async-versions/async-versions.il b/src/tests/async/async-versions/async-versions.il index 74edd05dcd2440..a5e8be30299bbf 100644 --- a/src/tests/async/async-versions/async-versions.il +++ b/src/tests/async/async-versions/async-versions.il @@ -16,6 +16,60 @@ { } +.class public auto ansi beforefieldinit ModifiedTask extends class [System.Runtime]System.Threading.Tasks.Task`1 +{ + .method public hidebysig specialname rtspecialname instance void .ctor() cil managed + { + ldarg.0 + ldnull + ldftn int32 ModifiedTask::GetValue() + newobj instance void class [System.Runtime]System.Func`1::.ctor(object, native int) + call instance void class [System.Runtime]System.Threading.Tasks.Task`1::.ctor(class [System.Runtime]System.Func`1) + ldarg.0 + call instance void [System.Runtime]System.Threading.Tasks.Task::RunSynchronously() + ret + } + + .method private hidebysig static int32 GetValue() cil managed + { + ldc.i4.s 42 + ret + } +} + +.class public auto ansi beforefieldinit ModifiedBase extends [System.Runtime]System.Object +{ + .method public hidebysig specialname rtspecialname instance void .ctor() cil managed + { + ldarg.0 + call instance void [System.Runtime]System.Object::.ctor() + ret + } + + .method public hidebysig newslot virtual instance class [System.Runtime]System.Threading.Tasks.Task`1 Get() cil managed async + { + ldc.i4.1 + ret + } +} + +.class public auto ansi beforefieldinit ModifiedDerived extends ModifiedBase +{ + .method public hidebysig specialname rtspecialname instance void .ctor() cil managed + { + ldarg.0 + call instance void ModifiedBase::.ctor() + ret + } + + .method public hidebysig newslot virtual instance class ModifiedTask Get() cil managed + { + .override method instance class [System.Runtime]System.Threading.Tasks.Task`1 ModifiedBase::Get() + newobj instance void ModifiedTask::.ctor() + ret + } +} + .class public auto ansi abstract sealed beforefieldinit AsyncVersions extends [System.Runtime]System.Object { @@ -157,6 +211,34 @@ // --- Tests ----------------------------------------------------------- + .method private hidebysig static class [System.Runtime]System.Threading.Tasks.Task`1 + AwaitModified() cil managed async + { + .maxstack 1 + newobj instance void ModifiedDerived::.ctor() + callvirt instance class [System.Runtime]System.Threading.Tasks.Task`1 ModifiedBase::Get() + call !!0 [System.Runtime]System.Runtime.CompilerServices.AsyncHelpers::Await(class [System.Runtime]System.Threading.Tasks.Task`1) + ret + } + + .method public hidebysig static void ModifiedGenericArgumentTest() cil managed + { + .custom instance void [xunit.core]Xunit.FactAttribute::.ctor() = ( 01 00 00 00 ) + .maxstack 2 + + call class [System.Runtime]System.Threading.Tasks.Task`1 AsyncVersions::AwaitModified() + callvirt instance !0 class [System.Runtime]System.Threading.Tasks.Task`1::get_Result() + ldc.i4.s 42 + beq.s Ok + + ldstr "Modified generic argument dispatched to the wrong method" + newobj instance void [System.Runtime]System.Exception::.ctor(string) + throw + + Ok: + ret + } + .method public hidebysig static void JmpToSuspendingTest() cil managed { .custom instance void [Microsoft.DotNet.XUnitExtensions]Xunit.ConditionalFactAttribute::.ctor(class [System.Runtime]System.Type, string[]) = { @@ -228,6 +310,7 @@ .try { + call void AsyncVersions::ModifiedGenericArgumentTest() call void AsyncVersions::JmpToSuspendingTest() call void AsyncVersions::JmpToCompletedTest() call void AsyncVersions::TailToSuspendingTest() diff --git a/src/tests/async/custom-task-covariant-return/custom-task-covariant-return.cs b/src/tests/async/custom-task-covariant-return/custom-task-covariant-return.cs new file mode 100644 index 00000000000000..70acd9ba5605f8 --- /dev/null +++ b/src/tests/async/custom-task-covariant-return/custom-task-covariant-return.cs @@ -0,0 +1,641 @@ +// Licensed to the .NET Foundation under one or more agreements. +// The .NET Foundation licenses this file to you under the MIT license. + +using System; +using System.Collections.Generic; +using System.Runtime.CompilerServices; +using System.Threading.Tasks; +using Xunit; + +// Task and Task are not sealed, so a covariant override may return a type that +// derives from Task, but is not itself Task or Task. Such an override is not +// task-returning as far as the runtime is concerned, while the method that it +// overrides may well be a runtime async method. +namespace CustomTaskCovariantReturn +{ + public class Program + { + internal static string Trace; + + public class MyTask : Task + { + public MyTask(Action action) : base(action) => RunSynchronously(); + } + + public class MyTask : Task + { + public MyTask(Func func) : base(func) => RunSynchronously(); + } + + public class Base + { + public virtual async Task M1() + { + Trace += "Base.M1;"; + } + + public virtual async Task M2() + { + Trace += "Base.M2;"; + return 1; + } + } + + public class Derived : Base + { + public override MyTask M1() => new MyTask(() => Trace += "Derived.M1;"); + + public override MyTask M2() => new MyTask(() => + { + Trace += "Derived.M2;"; + return 42; + }); + } + + public class Derived2 : Derived + { + public override MyTask M1() => new MyTask(() => + { + Trace += "Derived2.M1;"; + base.M1().GetAwaiter().GetResult(); + }); + + public override MyTask M2() => new MyTask(() => + { + Trace += "Derived2.M2;"; + return base.M2().GetAwaiter().GetResult() + 1; + }); + } + + public abstract class AbstractDerived : Base + { + public abstract override MyTask M1(); + + public abstract override MyTask M2(); + } + + public class ConcreteDerived : AbstractDerived + { + public override MyTask M1() => new MyTask(() => Trace += "ConcreteDerived.M1;"); + + public override MyTask M2() => new MyTask(() => + { + Trace += "ConcreteDerived.M2;"; + return 42; + }); + } + + // awaiting the result of the call directly + [MethodImpl(MethodImplOptions.NoInlining)] + private static async Task CallM1(Base b) => await b.M1(); + + [MethodImpl(MethodImplOptions.NoInlining)] + private static async Task CallM2(Base b) => await b.M2(); + + // the same, but the returned task is observed as an object as well, + // so the call itself cannot be a runtime async call. + [MethodImpl(MethodImplOptions.NoInlining)] + private static async Task CallM1ViaTask(Base b) + { + Task t = b.M1(); + await t; + return t; + } + + [MethodImpl(MethodImplOptions.NoInlining)] + private static async Task CallM2ViaTask(Base b) + { + Task t = b.M2(); + int result = await t; + Assert.Equal(typeof(MyTask), t.GetType()); + return result; + } + + [Fact] + public static void TestCustomTaskOverrideViaTask() + { + // check year to not be concerned with devirtualization. + Base b = DateTime.Now.Year > 0 ? new Derived() : new Base(); + + Trace = null; + Task t = CallM1ViaTask(b).GetAwaiter().GetResult(); + Assert.Equal("Derived.M1;", Trace); + Assert.IsType(t); + + Trace = null; + Assert.Equal(42, CallM2ViaTask(b).GetAwaiter().GetResult()); + Assert.Equal("Derived.M2;", Trace); + } + + [Fact] + public static void TestCustomTaskOverride() + { + Base b = DateTime.Now.Year > 0 ? new Derived() : new Base(); + + Trace = null; + CallM1(b).GetAwaiter().GetResult(); + Assert.Equal("Derived.M1;", Trace); + + Trace = null; + Assert.Equal(42, CallM2(b).GetAwaiter().GetResult()); + Assert.Equal("Derived.M2;", Trace); + } + + [Fact] + public static void TestCustomTaskOverrideOfCustomTaskOverride() + { + Base b = DateTime.Now.Year > 0 ? new Derived2() : new Base(); + + Trace = null; + CallM1(b).GetAwaiter().GetResult(); + Assert.Equal("Derived2.M1;Derived.M1;", Trace); + + Trace = null; + Assert.Equal(43, CallM2(b).GetAwaiter().GetResult()); + Assert.Equal("Derived2.M2;Derived.M2;", Trace); + } + + [Theory] + [InlineData(false)] + [InlineData(true)] + public static void TestAbstractCustomTaskOverride(bool viaTask) + { + Base b = DateTime.Now.Year > 0 ? new ConcreteDerived() : new Base(); + + Trace = null; + if (viaTask) + { + Assert.IsType(CallM1ViaTask(b).GetAwaiter().GetResult()); + } + else + { + CallM1(b).GetAwaiter().GetResult(); + } + Assert.Equal("ConcreteDerived.M1;", Trace); + + Trace = null; + int result = (viaTask ? CallM2ViaTask(b) : CallM2(b)).GetAwaiter().GetResult(); + Assert.Equal(42, result); + Assert.Equal("ConcreteDerived.M2;", Trace); + } + + [Fact] + public static void TestCustomTaskOverrideCalledDirectly() + { + Derived d = DateTime.Now.Year > 0 ? new Derived2() : new Derived(); + + Trace = null; + d.M1().GetAwaiter().GetResult(); + Assert.Equal("Derived2.M1;Derived.M1;", Trace); + + Trace = null; + Assert.Equal(43, d.M2().GetAwaiter().GetResult()); + Assert.Equal("Derived2.M2;Derived.M2;", Trace); + } + } +} + +// The same as above, but the overridden methods are not runtime async. +namespace CustomTaskCovariantReturnWithoutRuntimeAsync +{ + public class Program + { + internal static string Trace; + + public class MyTask : Task + { + public MyTask(Action action) : base(action) => RunSynchronously(); + } + + public class MyTask : Task + { + public MyTask(Func func) : base(func) => RunSynchronously(); + } + + public class Base + { + [RuntimeAsyncMethodGeneration(false)] + public virtual async Task M1() + { + Trace += "Base.M1;"; + } + + [RuntimeAsyncMethodGeneration(false)] + public virtual async Task M2() + { + Trace += "Base.M2;"; + return 1; + } + } + + public class Derived : Base + { + public override MyTask M1() => new MyTask(() => Trace += "Derived.M1;"); + + public override MyTask M2() => new MyTask(() => + { + Trace += "Derived.M2;"; + return 42; + }); + } + + [MethodImpl(MethodImplOptions.NoInlining)] + private static async Task CallM1(Base b) => await b.M1(); + + [MethodImpl(MethodImplOptions.NoInlining)] + private static async Task CallM2(Base b) => await b.M2(); + + [MethodImpl(MethodImplOptions.NoInlining)] + private static async Task CallM1ViaTask(Base b) + { + Task t = b.M1(); + await t; + Assert.Equal(typeof(MyTask), t.GetType()); + } + + [MethodImpl(MethodImplOptions.NoInlining)] + private static async Task CallM2ViaTask(Base b) + { + Task t = b.M2(); + int result = await t; + Assert.Equal(typeof(MyTask), t.GetType()); + return result; + } + + [Fact] + public static void TestCustomTaskOverrideViaTaskWithoutRuntimeAsync() + { + Base b = DateTime.Now.Year > 0 ? new Derived() : new Base(); + + Trace = null; + CallM1ViaTask(b).GetAwaiter().GetResult(); + Assert.Equal("Derived.M1;", Trace); + + Trace = null; + Assert.Equal(42, CallM2ViaTask(b).GetAwaiter().GetResult()); + Assert.Equal("Derived.M2;", Trace); + } + + [Fact] + public static void TestCustomTaskOverrideWithoutRuntimeAsync() + { + Base b = DateTime.Now.Year > 0 ? new Derived() : new Base(); + + Trace = null; + CallM1(b).GetAwaiter().GetResult(); + Assert.Equal("Derived.M1;", Trace); + + Trace = null; + Assert.Equal(42, CallM2(b).GetAwaiter().GetResult()); + Assert.Equal("Derived.M2;", Trace); + } + } +} + +// The same as above, but the methods and/or their declaring types are generic. +namespace CustomTaskCovariantReturnGenerics +{ + public class Program + { + internal static string Trace; + + public class MyTask : Task + { + public MyTask(Action action) : base(action) => RunSynchronously(); + } + + public class MyTask : Task + { + public MyTask(Func func) : base(func) => RunSynchronously(); + } + + public struct S : IEquatable + { + public S(int num) => this.num = num; + + public int num; + + public bool Equals(S other) => other.num == num; + } + + // a generic method on a non-generic type - the element type is a method type parameter + public class Base + { + public virtual async Task M1(T t) + { + Trace += "Base.M1;"; + return t; + } + + public virtual async Task M2(T t) + { + Trace += "Base.M2;"; + } + } + + public class Derived : Base + { + public override MyTask M1(T t) => new MyTask(() => + { + Trace += "Derived.M1;"; + return t; + }); + + public override MyTask M2(T t) => new MyTask(() => Trace += "Derived.M2;"); + } + + // a generic type - the element type is a type parameter of the declaring type + public class GBase + { + public virtual async Task M1(T t) + { + Trace += "GBase.M1;"; + return t; + } + + public virtual async Task M2(T t) + { + Trace += "GBase.M2;"; + } + + public virtual async Task> M3(T t) + { + Trace += "GBase.M3;"; + return new List { t }; + } + + public virtual async Task M4(U u) + { + Trace += "GBase.M4;"; + return u; + } + + public virtual async Task M5(T t) + { + Trace += "GBase.M5;"; + return new T[] { t }; + } + } + + public class GDerived : GBase + { + public override MyTask M1(T t) => new MyTask(() => + { + Trace += "GDerived.M1;"; + return t; + }); + + public override MyTask M2(T t) => new MyTask(() => Trace += "GDerived.M2;"); + + public override MyTask> M3(T t) => new MyTask>(() => + { + Trace += "GDerived.M3;"; + return new List { t }; + }); + + public override MyTask M4(U u) => new MyTask(() => + { + Trace += "GDerived.M4;"; + return u; + }); + + public override MyTask M5(T t) => new MyTask(() => + { + Trace += "GDerived.M5;"; + return new T[] { t }; + }); + } + + // the derived type closes the instantiation of the base type + public class ClosedDerived : GBase + { + public override MyTask M1(int t) => new MyTask(() => + { + Trace += "ClosedDerived.M1;"; + return t + 1; + }); + } + + // the derived type instantiates the base type with a composed type + public class ListDerived : GBase> + { + public override MyTask> M1(List t) => new MyTask>(() => + { + Trace += "ListDerived.M1;"; + return t; + }); + } + + public class ArrayMid : GBase + { + } + + public class MultiHopDerived : ArrayMid> + { + public override MyTask[]> M1(List[] t) => new MyTask[]>(() => + { + Trace += "MultiHopDerived.M1;"; + return t; + }); + } + + [MethodImpl(MethodImplOptions.NoInlining)] + private static async Task CallM1(Base b, T t) => await b.M1(t); + + [MethodImpl(MethodImplOptions.NoInlining)] + private static async Task CallM2(Base b, T t) => await b.M2(t); + + [MethodImpl(MethodImplOptions.NoInlining)] + private static async Task CallGM1(GBase b, T t) => await b.M1(t); + + [MethodImpl(MethodImplOptions.NoInlining)] + private static async Task CallGM2(GBase b, T t) => await b.M2(t); + + [MethodImpl(MethodImplOptions.NoInlining)] + private static async Task> CallGM3(GBase b, T t) => await b.M3(t); + + [MethodImpl(MethodImplOptions.NoInlining)] + private static async Task CallGM4(GBase b, U u) => await b.M4(u); + + [MethodImpl(MethodImplOptions.NoInlining)] + private static async Task CallGM5(GBase b, T t) => await b.M5(t); + + [Fact] + public static void TestGenericMethodCovariantOverride() + { + // check year to not be concerned with devirtualization. + Base b = DateTime.Now.Year > 0 ? new Derived() : new Base(); + + Trace = null; + Assert.Equal(42, CallM1(b, 42).GetAwaiter().GetResult()); + Assert.Equal("Derived.M1;", Trace); + + Trace = null; + Assert.Equal("hi", CallM1(b, "hi").GetAwaiter().GetResult()); + Assert.Equal("Derived.M1;", Trace); + + Trace = null; + CallM2(b, 42).GetAwaiter().GetResult(); + Assert.Equal("Derived.M2;", Trace); + + Trace = null; + CallM2(b, "hi").GetAwaiter().GetResult(); + Assert.Equal("Derived.M2;", Trace); + } + + // non-generic callers, so that the calls are runtime async calls with + // fully concrete instantiations. + [MethodImpl(MethodImplOptions.NoInlining)] + private static async Task CallGM1Int(GBase b) => await b.M1(42); + + [MethodImpl(MethodImplOptions.NoInlining)] + private static async Task CallGM1String(GBase b) => await b.M1("hi"); + + [MethodImpl(MethodImplOptions.NoInlining)] + private static async Task CallGM2Int(GBase b) => await b.M2(42); + + [MethodImpl(MethodImplOptions.NoInlining)] + private static async Task> CallGM3Int(GBase b) => await b.M3(42); + + [MethodImpl(MethodImplOptions.NoInlining)] + private static async Task CallM1Int(Base b) => await b.M1(42); + + [MethodImpl(MethodImplOptions.NoInlining)] + private static async Task CallM1String(Base b) => await b.M1("hi"); + + [Fact] + public static void TestGenericTypeCovariantOverrideNonGenericCaller() + { + GBase bi = DateTime.Now.Year > 0 ? new GDerived() : new GBase(); + + Trace = null; + Assert.Equal(42, CallGM1Int(bi).GetAwaiter().GetResult()); + Assert.Equal("GDerived.M1;", Trace); + + Trace = null; + CallGM2Int(bi).GetAwaiter().GetResult(); + Assert.Equal("GDerived.M2;", Trace); + + Trace = null; + Assert.Equal(new List { 42 }, CallGM3Int(bi).GetAwaiter().GetResult()); + Assert.Equal("GDerived.M3;", Trace); + + GBase bs = DateTime.Now.Year > 0 ? new GDerived() : new GBase(); + + Trace = null; + Assert.Equal("hi", CallGM1String(bs).GetAwaiter().GetResult()); + Assert.Equal("GDerived.M1;", Trace); + } + + [Fact] + public static void TestGenericMethodCovariantOverrideNonGenericCaller() + { + Base b = DateTime.Now.Year > 0 ? new Derived() : new Base(); + + Trace = null; + Assert.Equal(42, CallM1Int(b).GetAwaiter().GetResult()); + Assert.Equal("Derived.M1;", Trace); + + Trace = null; + Assert.Equal("hi", CallM1String(b).GetAwaiter().GetResult()); + Assert.Equal("Derived.M1;", Trace); + } + + [Fact] + public static void TestGenericTypeCovariantOverride() + { + GBase bi = DateTime.Now.Year > 0 ? new GDerived() : new GBase(); + + Trace = null; + Assert.Equal(42, CallGM1(bi, 42).GetAwaiter().GetResult()); + Assert.Equal("GDerived.M1;", Trace); + + Trace = null; + CallGM2(bi, 42).GetAwaiter().GetResult(); + Assert.Equal("GDerived.M2;", Trace); + + Trace = null; + Assert.Equal(new List { 42 }, CallGM3(bi, 42).GetAwaiter().GetResult()); + Assert.Equal("GDerived.M3;", Trace); + + Trace = null; + Assert.Equal(new int[] { 42 }, CallGM5(bi, 42).GetAwaiter().GetResult()); + Assert.Equal("GDerived.M5;", Trace); + + Trace = null; + Assert.Equal("hi", CallGM4(bi, "hi").GetAwaiter().GetResult()); + Assert.Equal("GDerived.M4;", Trace); + + Trace = null; + Assert.Equal(11, CallGM4(bi, 11).GetAwaiter().GetResult()); + Assert.Equal("GDerived.M4;", Trace); + + GBase bs = DateTime.Now.Year > 0 ? new GDerived() : new GBase(); + + Trace = null; + Assert.Equal("hi", CallGM1(bs, "hi").GetAwaiter().GetResult()); + Assert.Equal("GDerived.M1;", Trace); + + Trace = null; + CallGM2(bs, "hi").GetAwaiter().GetResult(); + Assert.Equal("GDerived.M2;", Trace); + + Trace = null; + Assert.Equal(new List { "hi" }, CallGM3(bs, "hi").GetAwaiter().GetResult()); + Assert.Equal("GDerived.M3;", Trace); + } + + [Fact] + public static void TestGenericTypeCovariantOverrideWithStruct() + { + GBase b = DateTime.Now.Year > 0 ? new GDerived() : new GBase(); + + Trace = null; + Assert.Equal(new S(42), CallGM1(b, new S(42)).GetAwaiter().GetResult()); + Assert.Equal("GDerived.M1;", Trace); + + Trace = null; + Assert.Equal(new S(42), CallGM4(b, new S(42)).GetAwaiter().GetResult()); + Assert.Equal("GDerived.M4;", Trace); + } + + [Fact] + public static void TestClosedGenericBaseCovariantOverride() + { + GBase b = DateTime.Now.Year > 0 ? new ClosedDerived() : new GBase(); + + Trace = null; + Assert.Equal(43, CallGM1(b, 42).GetAwaiter().GetResult()); + Assert.Equal("ClosedDerived.M1;", Trace); + } + + [Fact] + public static void TestComposedGenericBaseCovariantOverride() + { + GBase> b = DateTime.Now.Year > 0 ? new ListDerived() : new GBase>(); + + Trace = null; + Assert.Equal(new List { 42 }, CallGM1(b, new List { 42 }).GetAwaiter().GetResult()); + Assert.Equal("ListDerived.M1;", Trace); + + GBase> bs = DateTime.Now.Year > 0 ? new ListDerived() : new GBase>(); + + Trace = null; + Assert.Equal(new List { "hi" }, CallGM1(bs, new List { "hi" }).GetAwaiter().GetResult()); + Assert.Equal("ListDerived.M1;", Trace); + } + + [Fact] + public static void TestMultiHopComposedGenericBaseCovariantOverride() + { + GBase[]> b = DateTime.Now.Year > 0 ? new MultiHopDerived() : new GBase[]>(); + List[] input = new List[] { new List { 42 } }; + + Trace = null; + List[] result = CallGM1(b, input).GetAwaiter().GetResult(); + Assert.Same(input, result); + Assert.Equal(new List { 42 }, result[0]); + Assert.Equal("MultiHopDerived.M1;", Trace); + } + } +} diff --git a/src/tests/async/custom-task-covariant-return/custom-task-covariant-return.csproj b/src/tests/async/custom-task-covariant-return/custom-task-covariant-return.csproj new file mode 100644 index 00000000000000..5fc7131d728e55 --- /dev/null +++ b/src/tests/async/custom-task-covariant-return/custom-task-covariant-return.csproj @@ -0,0 +1,10 @@ + + + + true + + + + + +