Skip to content

Commit d4ed8cc

Browse files
committed
Support multiple inheritance with C3 method resolution order
1 parent 5ca6a0b commit d4ed8cc

8 files changed

Lines changed: 210 additions & 33 deletions

File tree

‎src/DotPython.Compiler/DotPythonBytecodeFormat.cs‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -4,5 +4,5 @@ namespace DotPython.Compiler;
44
public static class DotPythonBytecodeFormat
55
{
66
/// <summary>Gets the current bytecode format version.</summary>
7-
public const int CurrentVersion = 26;
7+
public const int CurrentVersion = 27;
88
}

‎src/DotPython.Compiler/PythonCompiler.cs‎

Lines changed: 6 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -1806,22 +1806,18 @@ private void CompileClassDefinition(PythonClassDefinitionStatement @class)
18061806
var constantIndex = AddConstant(
18071807
new PythonConstant(PythonConstantType.CodeObject, childCode)
18081808
);
1809-
if (@class.Bases.Count > 1)
1810-
{
1811-
Report(
1812-
"DPY3114",
1813-
"Multiple inheritance is not supported in this runtime slice.",
1814-
@class.Span
1815-
);
1816-
}
1817-
18181809
if (@class.Bases.Count == 0)
18191810
{
18201811
Emit(PythonOpCode.MakeClass, constantIndex, @class.Span);
18211812
}
18221813
else
18231814
{
1824-
CompileExpression(@class.Bases[0]);
1815+
foreach (var baseExpression in @class.Bases)
1816+
{
1817+
CompileExpression(baseExpression);
1818+
}
1819+
1820+
Emit(PythonOpCode.BuildTuple, @class.Bases.Count, @class.Span);
18251821
Emit(PythonOpCode.MakeClassWithBases, constantIndex, @class.Span);
18261822
}
18271823

‎src/DotPython.Runtime.Managed/Execution/ManagedObjectProtocols.cs‎

Lines changed: 41 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -105,10 +105,7 @@ when module.Globals.TryGetValue(name, out var moduleValue):
105105

106106
throw MissingAttribute(instance.Type.Name, name, span);
107107
case PythonSuperProxyValue proxy:
108-
if (
109-
proxy.DefiningType.BaseType is { } baseType
110-
&& TryGetTypeAttribute(baseType, name, out var inherited)
111-
)
108+
if (TryResolveSuperAttribute(proxy, name, out var inherited))
112109
{
113110
return inherited is PythonFunctionValue inheritedMethod
114111
? new PythonBoundUserMethodValue(name, proxy.Instance, inheritedMethod)
@@ -143,6 +140,10 @@ proxy.DefiningType.BaseType is { } baseType
143140
return value;
144141
case PythonManagedTypeValue type when name == "__name__":
145142
return new PythonTextValue(type.Name);
143+
case PythonManagedTypeValue type when name == "__mro__":
144+
return new PythonTupleValue([.. type.Mro.Cast<PythonValue>()]);
145+
case PythonManagedTypeValue type when name == "__bases__":
146+
return new PythonTupleValue([.. type.Bases.Cast<PythonValue>()]);
146147
case PythonManagedTypeValue type:
147148
throw MissingAttribute(type.Name, name, span);
148149
case PythonExternalObjectValue external:
@@ -1628,13 +1629,48 @@ PythonManagedObjectValue instance
16281629
_ => value,
16291630
};
16301631

1632+
/// <summary>
1633+
/// `super()` resolution: search the attribute in the classes after the defining
1634+
/// class in the instance's dynamic-type MRO (falling back to the defining class's
1635+
/// own MRO when the instance is not a managed object of a related type).
1636+
/// </summary>
1637+
private static bool TryResolveSuperAttribute(
1638+
PythonSuperProxyValue proxy,
1639+
string name,
1640+
out PythonValue value
1641+
)
1642+
{
1643+
var mro =
1644+
proxy.Instance is PythonManagedObjectValue managed
1645+
&& managed.Type.Mro.Contains(proxy.DefiningType)
1646+
? managed.Type.Mro
1647+
: proxy.DefiningType.Mro;
1648+
var searching = false;
1649+
foreach (var current in mro)
1650+
{
1651+
if (!searching)
1652+
{
1653+
searching = ReferenceEquals(current, proxy.DefiningType);
1654+
continue;
1655+
}
1656+
1657+
if (current.Attributes.TryGetValue(name, out value!))
1658+
{
1659+
return true;
1660+
}
1661+
}
1662+
1663+
value = null!;
1664+
return false;
1665+
}
1666+
16311667
internal static bool TryGetTypeAttribute(
16321668
PythonManagedTypeValue type,
16331669
string name,
16341670
out PythonValue value
16351671
)
16361672
{
1637-
for (var current = type; current is not null; current = current.BaseType)
1673+
foreach (var current in type.Mro)
16381674
{
16391675
if (current.Attributes.TryGetValue(name, out value!))
16401676
{

‎src/DotPython.Runtime.Managed/Execution/PythonValue.cs‎

Lines changed: 24 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -327,18 +327,40 @@ internal PythonManagedTypeValue(
327327
PythonManagedTypeValue? baseType = null,
328328
Func<IReadOnlyList<PythonValue>, PythonValue>? construct = null,
329329
string? exceptionBaseName = null
330+
)
331+
: this(
332+
name,
333+
baseType is null ? [] : [baseType],
334+
baseType?.Mro,
335+
construct,
336+
exceptionBaseName
337+
) { }
338+
339+
internal PythonManagedTypeValue(
340+
string name,
341+
IReadOnlyList<PythonManagedTypeValue> bases,
342+
IReadOnlyList<PythonManagedTypeValue>? linearizedBases,
343+
Func<IReadOnlyList<PythonValue>, PythonValue>? construct = null,
344+
string? exceptionBaseName = null
330345
)
331346
{
332347
ArgumentException.ThrowIfNullOrWhiteSpace(name);
333348
Name = name;
334-
BaseType = baseType;
349+
Bases = bases;
350+
Mro = [this, .. linearizedBases ?? []];
335351
Construct = construct;
336352
ExceptionBaseName = exceptionBaseName;
337353
}
338354

339355
internal Dictionary<string, PythonValue> Attributes { get; } = new(StringComparer.Ordinal);
340356

341-
internal PythonManagedTypeValue? BaseType { get; }
357+
/// <summary>The declared base classes, in source order.</summary>
358+
internal IReadOnlyList<PythonManagedTypeValue> Bases { get; }
359+
360+
/// <summary>The C3 method resolution order, starting with this type (no `object` terminus).</summary>
361+
internal IReadOnlyList<PythonManagedTypeValue> Mro { get; }
362+
363+
internal PythonManagedTypeValue? BaseType => Bases.Count == 0 ? null : Bases[0];
342364

343365
/// <summary>The builtin exception type this class derives from, when it is an exception class.</summary>
344366
internal string? ExceptionBaseName { get; }

‎src/DotPython.Runtime.Managed/Execution/PythonVirtualMachine.cs‎

Lines changed: 127 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -4893,6 +4893,18 @@ private void MakeClass(PythonInstruction instruction, PythonValue? baseValue = n
48934893
closure[index] = CurrentFrame.Cells[cellIndex];
48944894
}
48954895

4896+
if (baseValue is PythonTupleValue { Elements.Length: > 1 } baseTuple)
4897+
{
4898+
var multiType = CreateMultiBaseClass(code.Definition.Name, baseTuple, instruction.Span);
4899+
PushClassBodyFrame(multiType, code, closure);
4900+
return;
4901+
}
4902+
4903+
if (baseValue is PythonTupleValue singleTuple)
4904+
{
4905+
baseValue = singleTuple.Elements.Length == 1 ? singleTuple.Elements[0] : null;
4906+
}
4907+
48964908
PythonManagedTypeValue type;
48974909
switch (baseValue)
48984910
{
@@ -4937,8 +4949,17 @@ private void MakeClass(PythonInstruction instruction, PythonValue? baseValue = n
49374949
_exceptionBaseOverlay[type.Name] = type.ExceptionBaseName;
49384950
}
49394951

4952+
PushClassBodyFrame(type, code, closure);
4953+
}
4954+
4955+
private void PushClassBodyFrame(
4956+
PythonManagedTypeValue type,
4957+
PreparedPythonCode code,
4958+
PythonCell[] closure
4959+
)
4960+
{
49404961
var hasReturnLocalContinuation = CaptureReturnLocalContinuation();
4941-
var cells = CreateCells(code, closure, instruction.Span);
4962+
var cells = CreateCells(code, closure, default);
49424963
PushFrame(
49434964
code,
49444965
CurrentFrame.Globals,
@@ -4951,6 +4972,110 @@ private void MakeClass(PythonInstruction instruction, PythonValue? baseValue = n
49514972
);
49524973
}
49534974

4975+
private static PythonManagedTypeValue CreateMultiBaseClass(
4976+
string name,
4977+
PythonTupleValue baseTuple,
4978+
TextSpan span
4979+
)
4980+
{
4981+
var bases = new List<PythonManagedTypeValue>(baseTuple.Elements.Length);
4982+
foreach (var element in baseTuple.Elements)
4983+
{
4984+
switch (element)
4985+
{
4986+
case PythonManagedTypeValue { ExceptionBaseName: null } managedBase:
4987+
if (bases.Contains(managedBase))
4988+
{
4989+
throw Fault(
4990+
"DPY4034",
4991+
$"duplicate base class {managedBase.Name}",
4992+
span,
4993+
"TypeError"
4994+
);
4995+
}
4996+
4997+
bases.Add(managedBase);
4998+
break;
4999+
case PythonManagedTypeValue or PythonExceptionTypeValue:
5000+
throw Fault(
5001+
"DPY3114",
5002+
"Multiple inheritance with exception bases is not supported in this runtime slice.",
5003+
span,
5004+
"TypeError"
5005+
);
5006+
default:
5007+
throw Fault(
5008+
"DPY4034",
5009+
$"'{ManagedObjectProtocols.GetTypeName(element)}' is not an "
5010+
+ "acceptable base type.",
5011+
span,
5012+
"TypeError"
5013+
);
5014+
}
5015+
}
5016+
5017+
return new PythonManagedTypeValue(name, bases, LinearizeBases(bases, span));
5018+
}
5019+
5020+
/// <summary>
5021+
/// C3 merge over the bases' linearizations plus the base list itself: repeatedly
5022+
/// take the first head that appears in no other sequence's tail.
5023+
/// </summary>
5024+
private static List<PythonManagedTypeValue> LinearizeBases(
5025+
List<PythonManagedTypeValue> bases,
5026+
TextSpan span
5027+
)
5028+
{
5029+
var sequences = bases
5030+
.Select(baseType => new List<PythonManagedTypeValue>(baseType.Mro))
5031+
.ToList();
5032+
sequences.Add([.. bases]);
5033+
var result = new List<PythonManagedTypeValue>();
5034+
while (sequences.Any(sequence => sequence.Count != 0))
5035+
{
5036+
PythonManagedTypeValue? selected = null;
5037+
foreach (var sequence in sequences)
5038+
{
5039+
if (sequence.Count == 0)
5040+
{
5041+
continue;
5042+
}
5043+
5044+
var head = sequence[0];
5045+
var appearsInTail = sequences.Any(other =>
5046+
other.Count != 0 && other.IndexOf(head) > 0
5047+
);
5048+
if (!appearsInTail)
5049+
{
5050+
selected = head;
5051+
break;
5052+
}
5053+
}
5054+
5055+
if (selected is null)
5056+
{
5057+
throw Fault(
5058+
"DPY4034",
5059+
"Cannot create a consistent method resolution order (MRO) for bases "
5060+
+ string.Join(", ", bases.Select(baseType => baseType.Name)),
5061+
span,
5062+
"TypeError"
5063+
);
5064+
}
5065+
5066+
result.Add(selected);
5067+
foreach (var sequence in sequences)
5068+
{
5069+
if (sequence.Count != 0 && ReferenceEquals(sequence[0], selected))
5070+
{
5071+
sequence.RemoveAt(0);
5072+
}
5073+
}
5074+
}
5075+
5076+
return result;
5077+
}
5078+
49545079
private static PythonCell[] CreateCells(
49555080
PreparedPythonCode code,
49565081
PythonCell[] closure,
@@ -5242,15 +5367,7 @@ private bool MatchesClassInfo(PythonValue value, PythonValue classInfo, TextSpan
52425367
return false;
52435368
}
52445369

5245-
for (var current = instance.Type; current is not null; current = current.BaseType)
5246-
{
5247-
if (ReferenceEquals(current, managedType))
5248-
{
5249-
return true;
5250-
}
5251-
}
5252-
5253-
return false;
5370+
return instance.Type.Mro.Any(current => ReferenceEquals(current, managedType));
52545371
}
52555372
case PythonExternalObjectValue externalType:
52565373
return externalType.Protocol.IsInstanceOf(value, span);

‎tests/DotPython.CompilerTests/BytecodeFormatTests.cs‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -6,8 +6,8 @@ namespace DotPython.CompilerTests;
66
public sealed class BytecodeFormatTests
77
{
88
[Fact]
9-
public void CurrentVersion_TracksExceptionGroupInstructions()
9+
public void CurrentVersion_TracksMultipleInheritanceClassBases()
1010
{
11-
Assert.Equal(26, DotPythonBytecodeFormat.CurrentVersion);
11+
Assert.Equal(27, DotPythonBytecodeFormat.CurrentVersion);
1212
}
1313
}

‎tests/DotPython.CompilerTests/DotPythonModuleArtifactTests.cs‎

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -60,7 +60,7 @@ [new DotPythonModuleExport("calculate", "Calculate", DotPythonExportKind.Functio
6060

6161
Assert.Equal(
6262
"{\"formatVersion\":4,\"moduleName\":\"pricing\",\"languageVersion\":\"3.14\","
63-
+ "\"bytecodeFormatVersion\":26,\"exports\":[{\"pythonName\":\"calculate\","
63+
+ "\"bytecodeFormatVersion\":27,\"exports\":[{\"pythonName\":\"calculate\","
6464
+ "\"contractName\":\"Calculate\",\"kind\":\"function\"}]}",
6565
json
6666
);
@@ -503,10 +503,10 @@ public void ManifestJson_RejectsLanguageVersionsOutsideTheSupportedArtifactSet()
503503
{
504504
const string unsupportedLanguage =
505505
"{\"formatVersion\":4,\"moduleName\":\"sample\",\"languageVersion\":\"3.13\","
506-
+ "\"bytecodeFormatVersion\":26,\"exports\":[]}";
506+
+ "\"bytecodeFormatVersion\":27,\"exports\":[]}";
507507
const string nonCanonicalLanguage =
508508
"{\"formatVersion\":4,\"moduleName\":\"sample\",\"languageVersion\":\"3.14.0\","
509-
+ "\"bytecodeFormatVersion\":26,\"exports\":[]}";
509+
+ "\"bytecodeFormatVersion\":27,\"exports\":[]}";
510510

511511
var unsupportedFailure = Assert.Throws<InvalidDataException>(() =>
512512
DotPythonModuleManifestJson.Deserialize(unsupportedLanguage)

‎tests/DotPython.DifferentialTests/ManagedCliDifferentialTests.cs‎

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -477,6 +477,12 @@ public sealed class ManagedCliDifferentialTests
477477
[InlineData(
478478
"class Fut:\n def __await__(self):\n got = yield 'tick'\n return got\nasync def paced():\n for i in range(2):\n v = await Fut()\n yield i * 100 + len(v)\nasync def suspended_comp():\n values = [x async for x in paced()]\n print('collected', values)\ns = suspended_comp()\nprint(s.send(None))\nprint(s.send('ab'))\ntry:\n s.send('xyz')\nexcept StopIteration:\n pass\nasync def nums():\n for i in range(4):\n yield i\nasync def nested_comp():\n grid = [[y async for y in nums()] for _ in range(2)]\n print(grid)\nn = nested_comp()\ntry:\n n.send(None)\nexcept StopIteration:\n pass"
479479
)]
480+
[InlineData(
481+
"class A:\n def who(self): return 'A'\n tag = 'a'\nclass B(A):\n def who(self): return 'B'\nclass C(A):\n def who(self): return 'C'\n tag = 'c'\nclass D(B, C):\n pass\nprint([t.__name__ for t in D.__mro__ if t.__name__ != 'object'])\nprint([t.__name__ for t in D.__bases__])\nd = D()\nprint(d.who(), d.tag, D().who())\nprint(isinstance(d, A), isinstance(d, B), isinstance(d, C))\ntry:\n class Bad(A, A):\n pass\nexcept TypeError as e:\n print('e1:', e)\nclass X: pass\nclass Y(X): pass\ntry:\n class Z(X, Y):\n pass\nexcept TypeError as e:\n print('e2:', e)\nmatch d:\n case C():\n print('pattern-matches-C')"
482+
)]
483+
[InlineData(
484+
"class WA:\n def __init__(self):\n self.trail = ['A']\nclass WB(WA):\n def __init__(self):\n super().__init__()\n self.trail.append('B')\nclass WC(WA):\n def __init__(self):\n super().__init__()\n self.trail.append('C')\nclass WD(WB, WC):\n def __init__(self):\n super().__init__()\n self.trail.append('D')\nprint(WD().trail)\nclass M1:\n def name(self): return 'M1'\nclass M2(M1):\n def name(self): return 'M2+' + super().name()\nclass M3(M1):\n def name(self): return 'M3+' + super().name()\nclass M4(M2, M3):\n def name(self): return 'M4+' + super().name()\nprint(M4().name())\nclass Mix:\n def helper(self): return 'mix:' + self.base()\nclass Base:\n def base(self): return 'base'\nclass App(Mix, Base):\n pass\nprint(App().helper())\nclass Wide(M2, M3, Mix, Base):\n pass\nprint(Wide().name(), Wide().helper())\nprint([t.__name__ for t in Wide.__mro__ if t.__name__ != 'object'])"
485+
)]
480486
public void CommandExecution_MatchesReferencePythonForSupportedSubset(string code)
481487
{
482488
var python = FindReferencePython();

0 commit comments

Comments
 (0)