From fad12f7271763974bc44525ace6e78c68b211a6b Mon Sep 17 00:00:00 2001 From: Stone_Red <56473591+Stone-Red-Code@users.noreply.github.com> Date: Fri, 19 Dec 2025 21:48:37 +0100 Subject: [PATCH] Add dependency resolution --- RemoteExec.Client/RemoteExecutionException.cs | 6 +- RemoteExec.Client/RemoteExecutor.cs | 40 ++-- RemoteExec.Server/GlobalSuppressions.cs | 8 + RemoteExec.Server/Hubs/RemoteExecutionHub.cs | 171 ++++++++++++++++-- RemoteExec.Server/Program.cs | 4 +- .../RemoteJobAssemblyLoadContext.cs | 13 +- RemoteExec.Shared/RemoteExecutionRequest.cs | 2 +- RemoteExec/Program.cs | 6 +- 8 files changed, 196 insertions(+), 54 deletions(-) create mode 100644 RemoteExec.Server/GlobalSuppressions.cs diff --git a/RemoteExec.Client/RemoteExecutionException.cs b/RemoteExec.Client/RemoteExecutionException.cs index 53453a8..edf792b 100644 --- a/RemoteExec.Client/RemoteExecutionException.cs +++ b/RemoteExec.Client/RemoteExecutionException.cs @@ -1,5 +1,5 @@ - -[Serializable] -internal class RemoteExecutionException(string message) : Exception(message) +namespace RemoteExec.Client; + +public class RemoteExecutionException(string message) : Exception(message) { } \ No newline at end of file diff --git a/RemoteExec.Client/RemoteExecutor.cs b/RemoteExec.Client/RemoteExecutor.cs index 31c3096..0157eed 100644 --- a/RemoteExec.Client/RemoteExecutor.cs +++ b/RemoteExec.Client/RemoteExecutor.cs @@ -5,6 +5,7 @@ using RemoteExec.Shared; using System.Diagnostics.CodeAnalysis; using System.Reflection; using System.Text.Json; +using System.Threading.Channels; namespace RemoteExec.Client; @@ -18,6 +19,29 @@ public class RemoteExecutor(string url) public async Task StartAsync(CancellationToken cancellationToken = default) { + _ = _connection.On($"RequestAssembly", async (string assemblyName, Guid requestId) => + { + Assembly? assembly = AppDomain.CurrentDomain.GetAssemblies().FirstOrDefault(a => a.GetName().FullName == assemblyName); + + if (assembly == null) + { + assembly = Assembly.Load(new AssemblyName(assemblyName)); + } + + byte[] dllBytes = await File.ReadAllBytesAsync(assembly.Location!); + + Channel channel = Channel.CreateUnbounded(); + + foreach (byte b in dllBytes) + { + await channel.Writer.WriteAsync(b); + } + + channel.Writer.Complete(); + + await _connection.InvokeAsync("ProvideAssembly", requestId, channel.Reader); + }); + await _connection.StartAsync(cancellationToken); } @@ -73,28 +97,16 @@ public class RemoteExecutor(string url) { MethodInfo method = del.Method; Type declaringType = method.DeclaringType!; - Assembly asm = declaringType.Assembly; + Assembly assembly = declaringType.Assembly; if (!method.IsStatic) { throw new InvalidOperationException("Only static methods supported"); } - if (asm.IsDynamic) - { - throw new InvalidOperationException("Dynamic assemblies are not supported"); - } - - if (string.IsNullOrEmpty(asm.Location)) - { - throw new InvalidOperationException("Assembly location is not available"); - } - - byte[] dllBytes = File.ReadAllBytes(asm.Location); - RemoteExecutionRequest request = new RemoteExecutionRequest { - AssemblyBytes = dllBytes, + AssemblyName = assembly.GetName().FullName, TypeName = declaringType.FullName!, MethodName = method.Name, ArgumentTypes = [.. method.GetParameters().Select(p => p.ParameterType.AssemblyQualifiedName!)], diff --git a/RemoteExec.Server/GlobalSuppressions.cs b/RemoteExec.Server/GlobalSuppressions.cs new file mode 100644 index 0000000..cbef956 --- /dev/null +++ b/RemoteExec.Server/GlobalSuppressions.cs @@ -0,0 +1,8 @@ +// This file is used by Code Analysis to maintain SuppressMessage +// attributes that are applied to this project. +// Project-level suppressions either have no target or are given +// a specific target and scoped to a namespace, type, member, etc. + +using System.Diagnostics.CodeAnalysis; + +[assembly: SuppressMessage("Minor Code Smell", "S2325:Methods and properties that don't access instance data should be static", Justification = "", Scope = "type", Target = "~T:RemoteExec.Server.Hubs.RemoteExecutionHub")] diff --git a/RemoteExec.Server/Hubs/RemoteExecutionHub.cs b/RemoteExec.Server/Hubs/RemoteExecutionHub.cs index aae93eb..199ef7d 100644 --- a/RemoteExec.Server/Hubs/RemoteExecutionHub.cs +++ b/RemoteExec.Server/Hubs/RemoteExecutionHub.cs @@ -2,30 +2,64 @@ using RemoteExec.Shared; +using System.Collections.Concurrent; using System.Reflection; using System.Runtime.Loader; using System.Text.Json; +using System.Threading.Channels; namespace RemoteExec.Server.Hubs; -public class RemoteExecutionHub : Hub +public class RemoteExecutionHub(ILogger logger) : Hub { - private readonly Dictionary connections = []; + private static readonly ConcurrentDictionary connections = new(); + + private static readonly ConcurrentDictionary> pendingAssemblyRequests = new(); public override Task OnConnectedAsync() { - connections.Add(Context.ConnectionId, new RemoteJobAssemblyLoadContext($"RemoteJob_{Guid.NewGuid()}", true)); + RemoteJobAssemblyLoadContext assemblyLoadContext = new RemoteJobAssemblyLoadContext($"RemoteJob_{Guid.NewGuid()}"); + + assemblyLoadContext.Resolving += AssemblyLoadContext_Resolving; + + _ = connections.TryAdd(Context.ConnectionId, assemblyLoadContext); return base.OnConnectedAsync(); } + private Assembly? AssemblyLoadContext_Resolving(AssemblyLoadContext assemblyLoadContext, AssemblyName assemblyName) + { + logger.LogWarning("Assembly resolution triggered synchronously for {AssemblyName}. This should have been pre-loaded.", assemblyName.FullName); + + // Return null to let other resolution mechanisms try + return null; + } + + public async Task ProvideAssembly(Guid requestId, ChannelReader stream) + { + using MemoryStream ms = new MemoryStream(); + + while (await stream.WaitToReadAsync()) + { + while (stream.TryRead(out byte item)) + { + ms.WriteByte(item); + } + } + + byte[] assemblyBytes = ms.ToArray(); + + if (pendingAssemblyRequests.TryRemove(requestId, out TaskCompletionSource? tcs)) + { + tcs.SetResult(assemblyBytes); + } + } + public override Task OnDisconnectedAsync(Exception? exception) { - if (connections.TryGetValue(Context.ConnectionId, out RemoteJobAssemblyLoadContext assemblyLoadContext)) + if (connections.TryRemove(Context.ConnectionId, out RemoteJobAssemblyLoadContext? assemblyLoadContext)) { assemblyLoadContext.Unload(); - - _ = connections.Remove(Context.ConnectionId); } return base.OnDisconnectedAsync(exception); @@ -35,16 +69,24 @@ public class RemoteExecutionHub : Hub { try { - AssemblyLoadContext alc = new AssemblyLoadContext( - name: $"RemoteJob_{Guid.NewGuid()}", - isCollectible: true); + if (!connections.TryGetValue(Context.ConnectionId, out RemoteJobAssemblyLoadContext? assemblyLoadContext)) + { + throw new InvalidOperationException("Connection not found"); + } - using MemoryStream ms = new MemoryStream(req.AssemblyBytes); - Assembly asm = alc.LoadFromStream(ms); + Assembly? assembly = assemblyLoadContext.Assemblies.FirstOrDefault(a => a.GetName().FullName == req.AssemblyName); - Type type = asm.GetType(req.TypeName, throwOnError: true)!; + assembly ??= await RequestAssemblyAsync(req.AssemblyName); - Type?[] argTypes = req.ArgumentTypes + if (!assemblyLoadContext.Assemblies.Contains(assembly)) + { + using MemoryStream ms = new MemoryStream(await GetAssemblyBytesAsync(assembly)); + assembly = assemblyLoadContext.LoadFromStream(ms); + } + + Type type = assembly.GetType(req.TypeName, throwOnError: true)!; + + Type[] argTypes = req.ArgumentTypes .Select(Type.GetType) .ToArray()!; @@ -53,14 +95,13 @@ public class RemoteExecutionHub : Hub BindingFlags.Static | BindingFlags.Public | BindingFlags.NonPublic, binder: null, argTypes, - modifiers: null); + modifiers: null) ?? throw new MissingMethodException(req.TypeName, req.MethodName); - if (method == null) - { - throw new MissingMethodException(req.TypeName, req.MethodName); - } + // Pre-load all referenced assemblies to avoid triggering Resolving event during Invoke + await PreLoadReferencedAssembliesAsync(assemblyLoadContext, assembly); ParameterInfo[] parameters = method.GetParameters(); + if (parameters.Length != req.Arguments.Length) { throw new ArgumentException("Argument count mismatch"); @@ -95,8 +136,6 @@ public class RemoteExecutionHub : Hub object? result = method.Invoke(null, invokeArgs); - alc.Unload(); - return new RemoteExecutionResult { Result = result @@ -110,4 +149,96 @@ public class RemoteExecutionHub : Hub }; } } + + private async Task RequestAssemblyAsync(string assemblyName) + { + try + { + Guid guid = Guid.NewGuid(); + TaskCompletionSource tcs = new TaskCompletionSource(); + + _ = pendingAssemblyRequests.TryAdd(guid, tcs); + + logger.LogInformation("Requesting assembly {Assembly} with request ID {RequestId}", assemblyName, guid); + + await Clients.Caller.SendAsync("RequestAssembly", assemblyName, guid); + + logger.LogInformation("Waiting for assembly {Assembly} with request ID {RequestId}", assemblyName, guid); + + // Wait for the assembly with a timeout + byte[] assemblyBytes = await tcs.Task.WaitAsync(TimeSpan.FromSeconds(30)); + + logger.LogInformation("Received assembly {Assembly} with request ID {RequestId}", assemblyName, guid); + + // Return a temporary assembly just for metadata inspection + using MemoryStream ms = new MemoryStream(assemblyBytes); + return Assembly.Load(assemblyBytes); + } + catch (Exception ex) + { + logger.LogError(ex, "Error requesting assembly {Assembly}", assemblyName); + throw; + } + } + + private async Task GetAssemblyBytesAsync(Assembly assembly) + { + string assemblyName = assembly.GetName().FullName!; + + Guid guid = Guid.NewGuid(); + TaskCompletionSource tcs = new TaskCompletionSource(); + + _ = pendingAssemblyRequests.TryAdd(guid, tcs); + + logger.LogInformation("Requesting assembly bytes for {Assembly} with request ID {RequestId}", assemblyName, guid); + + await Clients.Caller.SendAsync("RequestAssembly", assemblyName, guid); + + byte[] assemblyBytes = await tcs.Task.WaitAsync(TimeSpan.FromSeconds(30)); + + logger.LogInformation("Received assembly bytes for {Assembly} with request ID {RequestId}", assemblyName, guid); + + return assemblyBytes; + } + + private async Task PreLoadReferencedAssembliesAsync(RemoteJobAssemblyLoadContext assemblyLoadContext, Assembly assembly) + { + AssemblyName[] referencedAssemblies = assembly.GetReferencedAssemblies(); + + foreach (AssemblyName referencedAssembly in referencedAssemblies) + { + try + { + // Try to load from the assembly load context first + Assembly? loadedAssembly = assemblyLoadContext.Assemblies.FirstOrDefault(a => a.GetName().FullName == referencedAssembly.FullName); + + if (loadedAssembly != null) + { + continue; // Already loaded in the context + } + + // Try to load from default context (BCL assemblies) + try + { + _ = assemblyLoadContext.LoadFromAssemblyName(referencedAssembly); + continue; // Successfully loaded from default context + } + catch + { + // If not in default context, request from client + logger.LogInformation("Pre-loading referenced assembly {Assembly}", referencedAssembly.FullName); + + Assembly tempAssembly = await RequestAssemblyAsync(referencedAssembly.FullName!); + byte[] assemblyBytes = await GetAssemblyBytesAsync(tempAssembly); + + using MemoryStream ms = new MemoryStream(assemblyBytes); + _ = assemblyLoadContext.LoadFromStream(ms); + } + } + catch (Exception ex) + { + logger.LogWarning(ex, "Could not pre-load referenced assembly {Assembly}", referencedAssembly.FullName); + } + } + } } \ No newline at end of file diff --git a/RemoteExec.Server/Program.cs b/RemoteExec.Server/Program.cs index ea540e2..8348be2 100644 --- a/RemoteExec.Server/Program.cs +++ b/RemoteExec.Server/Program.cs @@ -5,7 +5,7 @@ WebApplicationBuilder builder = WebApplication.CreateBuilder(args); // Add services to the container. builder.Services.AddControllers(); -builder.Services.AddSignalR(options => options.MaximumReceiveMessageSize = null); +builder.Services.AddSignalR(); builder.Services.AddOpenApi(); WebApplication app = builder.Build(); @@ -24,4 +24,4 @@ app.UseAuthorization(); app.MapControllers(); -app.Run(); +await app.RunAsync(); diff --git a/RemoteExec.Server/RemoteJobAssemblyLoadContext.cs b/RemoteExec.Server/RemoteJobAssemblyLoadContext.cs index 0febaba..92399c6 100644 --- a/RemoteExec.Server/RemoteJobAssemblyLoadContext.cs +++ b/RemoteExec.Server/RemoteJobAssemblyLoadContext.cs @@ -1,18 +1,7 @@ -using System.Reflection; -using System.Runtime.Loader; +using System.Runtime.Loader; namespace RemoteExec.Server; public class RemoteJobAssemblyLoadContext(string name) : AssemblyLoadContext(name, true) { - public event EventHandler? RequestAssembly; - - protected override Assembly? Load(AssemblyName assemblyName) - { - RequestAssemblyEventArgs requestAssemblyEventArgs = new RequestAssemblyEventArgs(assemblyName); - - RequestAssembly?.Invoke(this, requestAssemblyEventArgs); - - return requestAssemblyEventArgs.GetAssemblyAsync().ConfigureAwait(false).GetAwaiter().GetResult(); - } } diff --git a/RemoteExec.Shared/RemoteExecutionRequest.cs b/RemoteExec.Shared/RemoteExecutionRequest.cs index 2c47d11..dca8ec7 100644 --- a/RemoteExec.Shared/RemoteExecutionRequest.cs +++ b/RemoteExec.Shared/RemoteExecutionRequest.cs @@ -2,7 +2,7 @@ public sealed class RemoteExecutionRequest { - public required byte[] AssemblyBytes { get; set; } + public required string AssemblyName { get; set; } public required string TypeName { get; set; } public required string MethodName { get; set; } public required string[] ArgumentTypes { get; set; } diff --git a/RemoteExec/Program.cs b/RemoteExec/Program.cs index 558dfbc..00a075b 100644 --- a/RemoteExec/Program.cs +++ b/RemoteExec/Program.cs @@ -1,4 +1,6 @@ -using RemoteExec.Client; +using CuteUtils.FluentMath.TypeExtensions; + +using RemoteExec.Client; RemoteExecutor remoteExecutor = new RemoteExecutor("https://localhost:7109/remote"); await remoteExecutor.StartAsync(); @@ -13,5 +15,5 @@ Console.WriteLine($"Result: {result}"); static int Multiply(int x, int y) { - return x * y; + return x.Multiply(y); } \ No newline at end of file