Make DrvFs initialization callback lifetime-safe (#41218)

* Make DrvFs initialization callback lifetime-safe Route DrvFs initialization through the owning user session so VM shutdown and replacement are synchronized with callback execution. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: 7169104c-ee30-445e-b133-49bfa478ef12 * Validate the DrvFs initialization callback in WslCoreVm::Create Reject an empty callback at VM creation so a missing callback fails deterministically instead of throwing std::bad_function_call from a later process creation path. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: 7cf83647-6fbe-47bb-a4df-888a6abed9de --------- Co-authored-by: Ben Hillis <benhill@ntdev.microsoft.com> Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: 7169104c-ee30-445e-b133-49bfa478ef12 Copilot-Session: 7cf83647-6fbe-47bb-a4df-888a6abed9de

Ben Hillis committed Aug 20, 2026 at 10:58 UTC a454dbd799fc25955f029dad1bfcb0af5385ad27
4 files changed +48 -25
src/windows/service/exe/LxssUserSession.cpp
+31 -1
@@ -2978,8 +2978,13 @@ void LxssUserSessionImpl::_CreateVm()
2978
2979 m_vmId.store(vmId);
2980
2981 + const auto weakSession = weak_from_this();
2982 + auto initializeDrvFs = [weakSession, vmId](HANDLE userToken) noexcept {
2983 + return s_InitializeDrvFs(weakSession, vmId, userToken);
2984 + };
2985 +
2986 // Create the utility VM and register for callbacks.
2982 - m_utilityVm = WslCoreVm::Create(m_userToken, std::move(config), vmId);
2987 + m_utilityVm = WslCoreVm::Create(m_userToken, std::move(config), vmId, std::move(initializeDrvFs));
2988
2989 if (m_httpProxyStateTracker)
2990 {
@@ -4235,6 +4240,31 @@ wil::unique_hkey LxssUserSessionImpl::s_OpenLxssUserKey(_In_ HANDLE UserToken)
4240 return wsl::windows::common::registry::OpenLxssUserKey();
4241 }
4242
4243 +LX_INIT_DRVFS_MOUNT LxssUserSessionImpl::s_InitializeDrvFs(_In_ const std::weak_ptr<LxssUserSessionImpl>& Session, _In_ const GUID& VmId, _In_ HANDLE UserToken) noexcept
4244 +{
4245 + try
4246 + {
4247 + const auto session = Session.lock();
4248 + if (!session)
4249 + {
4250 + return LxInitDrvfsMountNone;
4251 + }
4252 +
4253 + std::lock_guard lock(session->m_instanceLock);
4254 + if (!session->m_utilityVm || !IsEqualGUID(session->m_utilityVm->GetRuntimeId(), VmId))
4255 + {
4256 + return LxInitDrvfsMountNone;
4257 + }
4258 +
4259 + return session->m_utilityVm->InitializeDrvFs(UserToken) ? LxInitDrvfsMountElevated : LxInitDrvfsMountNonElevated;
4260 + }
4261 + catch (...)
4262 + {
4263 + LOG_CAUGHT_EXCEPTION();
4264 + return LxInitDrvfsMountNone;
4265 + }
4266 +}
4267 +
4268 bool LxssUserSessionImpl::s_TerminateInstance(_Inout_ LxssUserSessionImpl* UserSession, _In_ GUID DistroGuid, _In_ bool CheckForClients)
4269 {
4270 bool success = true;
src/windows/service/exe/LxssUserSession.h
+3 -1
@@ -313,7 +313,7 @@ private:
313 /// <summary>
314 /// Each user gets its own LxssUserSessionImpl object, This object manages the lifetime of running instances.
315 /// </summary>
316 -class LxssUserSessionImpl
316 +class LxssUserSessionImpl : public std::enable_shared_from_this<LxssUserSessionImpl>
317 {
318 public:
319 LxssUserSessionImpl(_In_ PSID userSid, _In_ DWORD sessionId, _Inout_ wsl::windows::service::PluginManager& pluginManager);
@@ -772,6 +772,8 @@ private:
772 /// </summary>
773 static wil::unique_hkey s_OpenLxssUserKey(_In_ HANDLE UserToken);
774
775 + static LX_INIT_DRVFS_MOUNT s_InitializeDrvFs(_In_ const std::weak_ptr<LxssUserSessionImpl>& Session, _In_ const GUID& VmId, _In_ HANDLE UserToken) noexcept;
776 +
777 /// <summary>
778 /// Ensures the distribution name is valid.
779 /// </summary>
src/windows/service/exe/WslCoreVm.cpp
+8 -19
@@ -75,17 +75,20 @@ RequiredExtraMmioSpaceForPmemFileInMb(_In_ PCWSTR FilePath)
75 }
76 } // namespace
77
78 -WslCoreVm::WslCoreVm(_In_ wsl::core::Config&& VmConfig) :
79 - m_vmConfig(std::move(VmConfig)), m_traceClient(m_vmConfig.EnableTelemetry)
78 +WslCoreVm::WslCoreVm(_In_ wsl::core::Config&& VmConfig, _In_ InitializeDrvFsCallback InitializeDrvFs) :
79 + m_vmConfig(std::move(VmConfig)), m_initializeDrvFs(std::move(InitializeDrvFs)), m_traceClient(m_vmConfig.EnableTelemetry)
80 {
81 // Create a job object that will terminate child processes (wslhost.exe, wslrelay.exe)
82 // when the VM is destroyed.
83 m_processJobObject = wsl::windows::common::helpers::CreateKillOnCloseJob();
84 }
85
86 -std::unique_ptr<WslCoreVm> WslCoreVm::Create(_In_ const wil::shared_handle& UserToken, _In_ wsl::core::Config&& VmConfig, _In_ const GUID& VmId)
86 +std::unique_ptr<WslCoreVm> WslCoreVm::Create(
87 + _In_ const wil::shared_handle& UserToken, _In_ wsl::core::Config&& VmConfig, _In_ const GUID& VmId, _In_ InitializeDrvFsCallback InitializeDrvFs)
88 {
88 - auto newInstance = std::unique_ptr<WslCoreVm>{new WslCoreVm{std::move(VmConfig)}};
89 + THROW_HR_IF(E_INVALIDARG, !InitializeDrvFs);
90 +
91 + auto newInstance = std::unique_ptr<WslCoreVm>{new WslCoreVm{std::move(VmConfig), std::move(InitializeDrvFs)}};
92 try
93 {
94 const auto startTimeMs = GetTickCount64();
@@ -1281,7 +1284,7 @@ std::shared_ptr<LxssRunningInstance> WslCoreVm::CreateInstanceInternal(
1284 localConfig,
1285 DefaultUid,
1286 ClientLifetimeId,
1284 - std::bind(s_InitializeDrvFs, this, std::placeholders::_1),
1287 + m_initializeDrvFs,
1288 featureFlags,
1289 m_vmConfig.DistributionStartTimeout,
1290 m_vmConfig.InstanceIdleTimeout,
@@ -2778,20 +2781,6 @@ std::string WslCoreVm::s_GetMountTargetName(_In_ PCWSTR Disk, _In_opt_ PCWSTR Na
2781 return target;
2782 }
2783
2781 -LX_INIT_DRVFS_MOUNT WslCoreVm::s_InitializeDrvFs(_Inout_ WslCoreVm* VmContext, _In_ HANDLE UserToken)
2782 -{
2783 - try
2784 - {
2785 - return VmContext->InitializeDrvFs(UserToken) ? LxInitDrvfsMountElevated : LxInitDrvfsMountNonElevated;
2786 - }
2787 - catch (...)
2788 - {
2789 - LOG_CAUGHT_EXCEPTION();
2790 -
2791 - return LxInitDrvfsMountNone;
2792 - }
2793 -}
2794 -
2784 void CALLBACK WslCoreVm::s_OnExit(_In_ HCS_EVENT* Event, _In_opt_ void* Context)
2785 try
2786 {
src/windows/service/exe/WslCoreVm.h
+6 -4
@@ -51,7 +51,10 @@ class WslCoreVm
51 void operator=(const WslCoreVm&) = delete;
52
53 public:
54 - static std::unique_ptr<WslCoreVm> Create(_In_ const wil::shared_handle& UserToken, _In_ wsl::core::Config&& VmConfig, _In_ const GUID& VmId);
54 + using InitializeDrvFsCallback = std::function<LX_INIT_DRVFS_MOUNT(HANDLE)>;
55 +
56 + static std::unique_ptr<WslCoreVm> Create(
57 + _In_ const wil::shared_handle& UserToken, _In_ wsl::core::Config&& VmConfig, _In_ const GUID& VmId, _In_ InitializeDrvFsCallback InitializeDrvFs);
58
59 ~WslCoreVm() noexcept;
60
@@ -176,7 +179,7 @@ private:
179 bool operator==(const VirtioFsShare& other) const;
180 };
181
179 - WslCoreVm(_In_ wsl::core::Config&& VmConfig);
182 + WslCoreVm(_In_ wsl::core::Config&& VmConfig, _In_ InitializeDrvFsCallback InitializeDrvFs);
183
184 _Requires_lock_held_(m_guestDeviceLock)
185 void AddDrvFsShare(_In_ bool Admin, _In_ HANDLE UserToken);
@@ -258,8 +261,6 @@ private:
261
262 static std::string s_GetMountTargetName(_In_ PCWSTR Disk, _In_opt_ PCWSTR Name, _In_ int PartitionIndex);
263
261 - static LX_INIT_DRVFS_MOUNT s_InitializeDrvFs(_Inout_ WslCoreVm* VmContext, _In_ HANDLE UserToken);
262 -
264 static void CALLBACK s_OnExit(_In_ HCS_EVENT* Event, _In_opt_ void* Context);
265
266 wil::srwlock m_guestDeviceLock;
@@ -282,6 +283,7 @@ private:
283 std::wstring m_machineId;
284 GUID m_runtimeId;
285 wsl::core::Config m_vmConfig;
286 + InitializeDrvFsCallback m_initializeDrvFs;
287 std::wstring m_comPipe0;
288 std::wstring m_comPipe1;
289 int m_pageReportingOrder;