Skip to content

Commit ceb3fb5

Browse files
committed
Fix mutable callback state propagation
1 parent e61f853 commit ceb3fb5

3 files changed

Lines changed: 110 additions & 3 deletions

File tree

‎PdVm.Runtime/PdVmProgramBase.cs‎

Lines changed: 74 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,8 @@ public abstract class PdVmProgramBase : IPdVmCallableProgram
77
private readonly List<PdVmValue> _stack = new();
88
private readonly List<PdVmValue> _locals = new();
99
private readonly Dictionary<int, PdVmCaptureCell> _captureCells = new();
10+
private readonly Dictionary<int, PdVmCaptureCell> _mutableBorrowAliases = new();
11+
private readonly HashSet<PdVmCaptureCell> _mutableBorrowCells = new(ReferenceEqualityComparer.Instance);
1012
private readonly List<PdVmExecutionFrame> _executionFrames = new();
1113
private readonly SemaphoreSlim _managedExecutionGate = new(1, 1);
1214
private readonly object _callbackQueueLock = new();
@@ -24,6 +26,7 @@ public abstract class PdVmProgramBase : IPdVmCallableProgram
2426
private bool _shutdown;
2527
private bool _callbackRunnerActive;
2628
private CancellationTokenSource _callbackRunnerCancellation = new();
29+
private (PdVmValue Value, PdVmCaptureCell Cell)? _lastBorrowedCapture;
2730

2831
private sealed record QueuedCallbackWork(
2932
object Generation,
@@ -363,13 +366,49 @@ protected PdVmValue PopValue()
363366
protected PdVmValue LoadLocalValue(byte index)
364367
{
365368
var absolute = ResolveLocalIndex(index);
366-
return _captureCells.TryGetValue(absolute, out var cell) ? cell.Value : _locals[absolute];
369+
if (_captureCells.TryGetValue(absolute, out var cell))
370+
{
371+
_lastBorrowedCapture = IsInsideScriptCallable() && _mutableBorrowCells.Contains(cell)
372+
? (cell.Value, cell)
373+
: null;
374+
return cell.Value;
375+
}
376+
377+
_lastBorrowedCapture = null;
378+
return _locals[absolute];
367379
}
368380

381+
private bool IsInsideScriptCallable() =>
382+
GetActiveFrameOrDefault() is { PrototypeId: not null };
383+
369384
protected void StoreLocalValue(byte index, PdVmValue value)
370385
{
371386
ArgumentNullException.ThrowIfNull(value);
372387
var absolute = ResolveLocalIndex(index);
388+
if (_mutableBorrowAliases.TryGetValue(absolute, out var aliasCell))
389+
{
390+
// The CLR compiler reuses the alias local as a temporary while
391+
// lowering Set. Only the container returned by Set can publish a
392+
// borrowed map/array update; staging scalars (usually null) must
393+
// remain local temporaries.
394+
if (value.Kind is not (PdVmValueKind.Map or PdVmValueKind.Array))
395+
{
396+
_locals[absolute] = value;
397+
_lastBorrowedCapture = null;
398+
return;
399+
}
400+
401+
if (ReferencesCaptureCell(value, aliasCell, new HashSet<PdVmCaptureCell>(ReferenceEqualityComparer.Instance)))
402+
{
403+
throw new InvalidOperationException("callable capture ownership cycle is unsupported");
404+
}
405+
406+
aliasCell.Value = value;
407+
_locals[absolute] = value;
408+
_lastBorrowedCapture = null;
409+
return;
410+
}
411+
373412
if (_captureCells.TryGetValue(absolute, out var cell))
374413
{
375414
if (ReferencesCaptureCell(value, cell, new HashSet<PdVmCaptureCell>(ReferenceEqualityComparer.Instance)))
@@ -379,8 +418,15 @@ protected void StoreLocalValue(byte index, PdVmValue value)
379418

380419
cell.Value = value;
381420
}
421+
else if (_lastBorrowedCapture is { } borrowed &&
422+
ReferenceEquals(value, borrowed.Value))
423+
{
424+
_mutableBorrowAliases[absolute] = borrowed.Cell;
425+
borrowed.Cell.Value = value;
426+
}
382427

383428
_locals[absolute] = value;
429+
_lastBorrowedCapture = null;
384430
}
385431

386432
protected PdVmValue[] GetLocalValues() => _locals.ToArray();
@@ -534,6 +580,12 @@ protected bool CompleteActiveFrame()
534580
{
535581
_captureCells.Remove(absolute);
536582
}
583+
foreach (var absolute in _mutableBorrowAliases.Keys
584+
.Where(index => index >= frame.LocalBase && index < frameEnd)
585+
.ToArray())
586+
{
587+
_mutableBorrowAliases.Remove(absolute);
588+
}
537589

538590
if (frameEnd != _locals.Count)
539591
{
@@ -570,6 +622,9 @@ protected void ResetRuntimeForReuse()
570622
_stack.Clear();
571623
_locals.Clear();
572624
_captureCells.Clear();
625+
_mutableBorrowAliases.Clear();
626+
_mutableBorrowCells.Clear();
627+
_lastBorrowedCapture = null;
573628
_executionFrames.Clear();
574629
_managedCallableResult = null;
575630
_mapIterators.Clear();
@@ -584,6 +639,9 @@ protected void ShutdownRuntime()
584639
_stack.Clear();
585640
_locals.Clear();
586641
_captureCells.Clear();
642+
_mutableBorrowAliases.Clear();
643+
_mutableBorrowCells.Clear();
644+
_lastBorrowedCapture = null;
587645
_executionFrames.Clear();
588646
_managedCallableResult = null;
589647
_mapIterators.Clear();
@@ -760,6 +818,10 @@ private void EnterScriptCallable(
760818
if (prototype.SelfSlot != slot)
761819
{
762820
_captureCells[absolute] = cell;
821+
if (prototype.CaptureModes[index] == PdVmRuntimeCaptureBindingMode.BorrowMut)
822+
{
823+
_mutableBorrowCells.Add(cell);
824+
}
763825
}
764826
}
765827
}
@@ -841,6 +903,11 @@ private PdVmValue BindCallable(IReadOnlyList<PdVmValue> args)
841903
_captureCells[absolute] = shared;
842904
}
843905

906+
if (mode == PdVmRuntimeCaptureBindingMode.BorrowMut)
907+
{
908+
_mutableBorrowCells.Add(shared);
909+
}
910+
844911
_locals[absolute] = shared.Value;
845912
cells[index] = shared;
846913
}
@@ -866,6 +933,7 @@ private PdVmCallOutcome DetachLocal(IReadOnlyList<PdVmValue> args)
866933

867934
var absolute = ResolveLocalIndex((byte)slot);
868935
_captureCells.Remove(absolute);
936+
_mutableBorrowAliases.Remove(absolute);
869937
_locals[absolute] = PdVmValue.Null();
870938
return PdVmCallOutcome.Returned(PdVmCallReturn.None);
871939
}
@@ -1494,6 +1562,11 @@ private void AbortManagedInvocation(int stackBase, int localBase)
14941562
{
14951563
_captureCells.Remove(absolute);
14961564
}
1565+
foreach (var absolute in _mutableBorrowAliases.Keys.Where(index => index >= localBase).ToArray())
1566+
{
1567+
_mutableBorrowAliases.Remove(absolute);
1568+
}
1569+
_lastBorrowedCapture = null;
14971570

14981571
if (_locals.Count > localBase)
14991572
{

‎PdVm.Tests/PdVmTypedDotNetInteropTests.cs‎

Lines changed: 29 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -175,6 +175,35 @@ public async Task ManagedCallbackCanCallAClosureKeptInRootLocals()
175175
Assert.Equal(1, result.AsInt());
176176
}
177177

178+
[Fact]
179+
public async Task MutableBorrowAliasPublishesContainerUpdates()
180+
{
181+
using var fixture = new SourceFixture(
182+
"let mut state: map<int> = { value: 0 };\n" +
183+
"pub fn increment() -> int {\n" +
184+
" let mut current = &mut state;\n" +
185+
" current.value = current.value + 1;\n" +
186+
" state.value\n" +
187+
"}\n");
188+
var output = PdVmDotNetSourceCompiler.CompileFile(fixture.SourcePath, fixture.OutputPath);
189+
var program = Assert.IsAssignableFrom<IPdVmCallableProgram>(
190+
PdVmAssemblyLoader.CreateProgram(Assembly.Load(File.ReadAllBytes(output))));
191+
var host = PdVmDefaultHost.CreateConsoleHost();
192+
_ = PdVmExecution.Run(program, host);
193+
194+
var first = await program.InvokeCallableAsync(
195+
program.ResolveCallable("increment"),
196+
Array.Empty<PdVmValue>(),
197+
host);
198+
var second = await program.InvokeCallableAsync(
199+
program.ResolveCallable("increment"),
200+
Array.Empty<PdVmValue>(),
201+
host);
202+
203+
Assert.Equal(1, first.AsInt());
204+
Assert.Equal(2, second.AsInt());
205+
}
206+
178207
[Fact]
179208
public async Task CallbackAdapterSchemaIsValidatedBeforeScriptExecution()
180209
{

‎examples/dotnet-minesweeper.rss‎

Lines changed: 7 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -423,14 +423,19 @@ let mut display = new_board_array(64);
423423
let mut flash_cells = new_board_array(64);
424424

425425
fn render_board(render_rows: int, render_columns: int, render_paused: bool, render_flash_until: int, render_pressed_index: int) -> null {
426+
// These arrays are resized when the difficulty changes. Borrow the root
427+
// slots so this renderer observes the replacement arrays instead of the
428+
// 8x8 values captured during startup.
429+
let current_display = &display;
430+
let current_flash_cells = &flash_cells;
426431
let mut render_styles = new_board_array(render_rows * render_columns);
427432
let mut render_index = 0;
428433
while render_index < render_rows * render_columns {
429-
let mut render_style = display[render_index].copy();
434+
let mut render_style = current_display[render_index].copy();
430435
if render_paused {
431436
render_style = 0;
432437
}
433-
if render_flash_until > 0 && flash_cells[render_index].copy() == 1 {
438+
if render_flash_until > 0 && current_flash_cells[render_index].copy() == 1 {
434439
render_style = 2;
435440
}
436441
if render_index == render_pressed_index && render_style == 0 {

0 commit comments

Comments
 (0)