Type Extension

Making it easier to work with shaders


Type Extension

Syntax

Extension declaration:

'extension' type-expr
    [':' bases-clause]
'{' member-list '}'

Generic extension declaration:

'extension' generic-params-decl type-expr
    [':' bases-clause]
    ('where' where-clause)*
'{' member-list '}'

Parameters

Description

An extension declaration extends a struct type, an enum type, or a set of such types. An extension may be used to add static data members, member functions, constructors, properties, subscript operators, function call operators, and conformances to an existing type. An extension may not change the data layout, that is, it cannot be used to append non-static data members.

📝 Remark: An interface type cannot be extended. Doing so would add new requirements to all conforming types, which would invalidate existing conformances.

Struct Extension

A previously defined struct type can be extended using an extension declaration. In the following example, an extension is used to add a new member function.

Example 1:

struct ExampleStruct
{
    uint32_t a;

    uint32_t getASquared()
    {
        return a * a;
    }
}

extension ExampleStruct
{
    // add a member function to ExampleStruct
    [mutating] void addToA(uint32_t x)
    {
        a = a + x;
    }
}

An extension can also be used to provide interface requirements to a struct.

Example 2:

interface IReq
{
    int requiredFunc();
}

struct TestClass : IReq
{
}

extension TestClass
{
    int requiredFunc()
    {
        return 42;
    }
}

[shader("compute")]
void main(uint3 id : SV_DispatchThreadID)
{
    TestClass obj = {  };

    obj.requiredFunc();
}

Finally, an extension can add new interface conformances to a struct.

Example 3:

interface IReq
{
    int requiredFunc();
}

struct TestClass
{
}

extension TestClass : IReq
{
    int requiredFunc()
    {
        return 42;
    }
}

[shader("compute")]
void main(uint3 id : SV_DispatchThreadID)
{
    IReq obj = TestClass();

    obj.requiredFunc();
}

⚠️ Warning: When an extension and the base type contain a member with the same signature, it is currently undefined which member takes effect. (Issue #9660)

Enumeration Extension

Similar to a struct, a previously defined enum type can be extended using an extension declaration. For non-static member functions, this is the value of the enumeration object.

An enumerator is an immutable instance of an enumeration type.

The first example shows basic enumeration extension.

Example 1:

enum TestEnum
{
    VALUE1 = 1,
    VALUE2,
    VALUE3,
    VALUE4,
}

extension TestEnum
{
    static int staticValue = 42;
    static int staticFunction()
    {
        return 123;
    }

    int times2()
    {
        // 'this' is the value of the
        // enumeration object
        return this * 2;
    }
}


RWStructuredBuffer<int> output;

[numthreads(1,1,1)]
void main(uint3 tid : SV_DispatchThreadID)
{
    TestEnum b = TestEnum.VALUE1;

    // static members work the same as in structs
    output[0] = TestEnum.staticValue;      // 42
    output[1] = TestEnum.staticFunction(); // 123

    // member function returning double the current
    // value
    output[2] = b.times2();           // 1 * 2 = 2

    // member functions also work on enumerators
    output[3] = b.VALUE4.times2();    // 4 * 2 = 8
}

The second example shows various ways an enum can be extended.

Example 2:

interface IBase
{
    static const int requiredConstant;
    property int requiredProp { get; set; }
    int requiredFunc();
}

enum TestEnum
{
    Zero = 0,
    One = 1,
    Two = 2,
    Three = 3,
}

enum AnotherEnum
{
    SomeValue = 42,
}

extension TestEnum : IBase
{
    // constant value required by IBase
    static const int requiredConstant = 42;

    // property required by IBase
    property int requiredProp {
        get() { return this; }
        set(int newVal) { this = (TestEnum)newVal; }
    }

    // member function required by IBase
    int requiredFunc()
    {
        return -this;
    }

    // subscript operator
    __subscript(int index) -> int {
        get() { return index * this; }
    }

    // incrementing member function
    [mutating] void increment()
    {
        this = (TestEnum)(this + 1);
    }

    // function call operator
    float operator () (float scale)
    {
        return (float)this * scale;
    }

    __init(AnotherEnum other)
    {
        // first, capture the underlying value to
        // avoid recursion
        int underlying = other;
        this = (TestEnum)(underlying * 2);

        // Note: the following would recursively
        // call this function again
        // this = (TestEnum)other;
    }
}

RWStructuredBuffer<int> output;

[numthreads(1,1,1)]
void main(uint3 threadId : SV_DispatchThreadID)
{
    TestEnum tmp = TestEnum.Two;
    output[0] = tmp.requiredProp; // 2

    tmp.requiredProp = 55;
    output[1] = tmp.requiredProp; // 55

    tmp.increment();
    output[2] = tmp.requiredProp; // 56

    output[3] = TestEnum.Three.requiredFunc(); // -3

    output[4] = tmp[4];           // 56 * 4 = 224
    output[5] = (int)(TestEnum.Two(0.5));  // 1

    tmp = TestEnum(AnotherEnum.SomeValue); // 42 * 2
    output[6] = tmp;                       // 84

    // Illegal, since enumerators are immutable
    // TestEnum.One.increment();

}

Generic Extension

All types conforming to an interface may be extended using a generic extension declaration, which adds new members to all conforming types. If multiple declarations share the same signature, the one in the concrete type takes precedence.

Example:

interface IBase
{
    int getA();
}

struct ConcreteInt16 : IBase
{
    int16_t a;

    int getA()
    {
        return a;
    }
}

struct ConcreteInt32 : IBase
{
    int32_t a;

    int getA()
    {
        return a;
    }
}

extension<T : IBase> T
{
    // added to all types conforming to
    // interface IBase
    int getASquared()
    {
        return getA() * getA();
    }
}

See Generics for further information on generics.