struct ControlPoint
{
    float4 objNormal : TEXCOORD10; // 对象空间法线
    float4 objPos : SV_Position0; // 对象空间位置
    float4 color1 : COLOR1;
    float3 color0 : COLOR0;
    float4 tex0 : TEXCOORD0;
    float4 tex1 : TEXCOORD1;
    float4 tex2 : TEXCOORD2;
    float4 objTangent : TEXCOORD3; // xyz=对象空间切线 w=tangent.w handedness
    float4 objBinormal : TEXCOORD4; // xyz=对象空间副切线（VS直传）
    float4 viewProjRow0 : TEXCOORD6; // gViewProjection第0行
    float4 viewProjRow1 : TEXCOORD7; // gViewProjection第1行
    float4 viewProjRow2 : TEXCOORD8; // gViewProjection第2行
    float4 viewProjRow3 : TEXCOORD9; // gViewProjection第3行
    float4 viewRow0 : TEXCOORD14; // gViewMatrix第0行
    float4 viewRow1 : TEXCOORD15; // gViewMatrix第1行
    float4 viewRow2 : TEXCOORD16; // gViewMatrix第2行
    float4 viewRow3 : TEXCOORD17; // gViewMatrix第3行
    float4 skinRow0 : TEXCOORD18; // skin变换第0行
    float4 skinRow1 : TEXCOORD19; // skin变换第1行
    float4 skinRow2 : TEXCOORD20; // skin变换第2行
    float4 skinTranslate : TEXCOORD21; // skin变换平移
    float4 posEdgeCP1 : TEXCOORD11; // PN三角形：位置边贝塞尔控制点1（对象空间）
    float4 posEdgeCP2 : TEXCOORD12; // PN三角形：位置边贝塞尔控制点2（对象空间）
};

struct TessFactors
{
    float edge[3] : SV_TessFactor;
    float inside  : SV_InsideTessFactor;
};

struct DSOutput
{
    float4 clipPos : SV_Position0;
    float4 color1 : COLOR1;
    float3 color0 : COLOR0;
    float4 tex0 : TEXCOORD0;
    float4 tex1 : TEXCOORD1;
    float4 tex2 : TEXCOORD2;
    float4 viewTangent : TEXCOORD3;
    float4 viewBinormal : TEXCOORD4;
    float4 viewNormal : TEXCOORD5;
};

Texture1D<float4> IniParams : register(t120);

float3 CalculateCubicBezierPosition(float3 controlPoints[10], float3 bary, float smoothing = 0.75)
{
    float3 flatSurfacePosition = bary.x * controlPoints[0] + bary.y * controlPoints[3] + bary.z * controlPoints[6];
    float3 bezierSurfacePosition =
        bary.x * bary.x * bary.x * controlPoints[0] + 3.0 * bary.x * bary.x * bary.y * controlPoints[1] + 3.0 * bary.x * bary.y * bary.y * controlPoints[2] +
        bary.y * bary.y * bary.y * controlPoints[3] + 3.0 * bary.y * bary.y * bary.z * controlPoints[4] + 3.0 * bary.y * bary.z * bary.z * controlPoints[5] +
        bary.z * bary.z * bary.z * controlPoints[6] + 3.0 * bary.x * bary.z * bary.z * controlPoints[7] + 3.0 * bary.x * bary.x * bary.z * controlPoints[8] +
        6.0 * bary.x * bary.y * bary.z * controlPoints[9];

    return lerp(flatSurfacePosition, bezierSurfacePosition, smoothing);
}

[domain("tri")]
DSOutput main(TessFactors factors, float3 bary : SV_DomainLocation, OutputPatch<ControlPoint, 3> patch)
{
    DSOutput output;

    output.color1 = patch[0].color1*bary.x + patch[1].color1*bary.y + patch[2].color1*bary.z;
    output.color0 = patch[0].color0*bary.x + patch[1].color0*bary.y + patch[2].color0*bary.z;
    output.tex0 = patch[0].tex0*bary.x + patch[1].tex0*bary.y + patch[2].tex0*bary.z;
    output.tex1 = patch[0].tex1*bary.x + patch[1].tex1*bary.y + patch[2].tex1*bary.z;
    output.tex2 = patch[0].tex2*bary.x + patch[1].tex2*bary.y + patch[2].tex2*bary.z;

    // === PN三角形曲面求值：在对象空间计算 ===
    float3 posControlPoints[10];
    float3 avgBezier = 0.0, avgVertices = 0.0;
    [unroll]
    for (int i = 0; i < 3; i++) {
        posControlPoints[i * 3] = patch[i].objPos.xyz;
        posControlPoints[i * 3 + 1] = patch[i].posEdgeCP1.xyz;
        posControlPoints[i * 3 + 2] = patch[i].posEdgeCP2.xyz;
        avgBezier += patch[i].posEdgeCP1.xyz + patch[i].posEdgeCP2.xyz;
        avgVertices += patch[i].objPos.xyz;
    }
    avgBezier /= 6.0;
    avgVertices /= 3.0;
    posControlPoints[9] = avgBezier + (avgBezier - avgVertices) * 0.5;

    float _Smoothing = IniParams[0].x;
    float3 finalPosOS = CalculateCubicBezierPosition(posControlPoints, bary, _Smoothing);

    // === 插值 skin 变换矩阵，对象空间→世界空间 ===
    float3 skinRow0 = patch[0].skinRow0.xyz * bary.x + patch[1].skinRow0.xyz * bary.y + patch[2].skinRow0.xyz * bary.z;
    float3 skinRow1 = patch[0].skinRow1.xyz * bary.x + patch[1].skinRow1.xyz * bary.y + patch[2].skinRow1.xyz * bary.z;
    float3 skinRow2 = patch[0].skinRow2.xyz * bary.x + patch[1].skinRow2.xyz * bary.y + patch[2].skinRow2.xyz * bary.z;
    float3 skinT = patch[0].skinTranslate.xyz * bary.x + patch[1].skinTranslate.xyz * bary.y + patch[2].skinTranslate.xyz * bary.z;

    float3 finalPosWS;
    finalPosWS.x = dot(skinRow0, finalPosOS) + skinT.x;
    finalPosWS.y = dot(skinRow1, finalPosOS) + skinT.y;
    finalPosWS.z = dot(skinRow2, finalPosOS) + skinT.z;

    // 裁剪空间变换（行向量 × row_major 矩阵）
    float4 m0 = patch[0].viewProjRow0;
    float4 m1 = patch[0].viewProjRow1;
    float4 m2 = patch[0].viewProjRow2;
    float4 m3 = patch[0].viewProjRow3;
    float4 worldPosH = float4(finalPosWS, 1.0);
    float4 clipPos;
    clipPos = m1 * worldPosH.yyyy;
    clipPos = m0 * worldPosH.xxxx + clipPos;
    clipPos = m2 * worldPosH.zzzz + clipPos;
    clipPos = m3 + clipPos;
    output.clipPos = clipPos;

    // === TBN：对象空间重心插值 → skin 旋转(3x3)→ 世界空间 → 视图空间 ===
    float3 objN = patch[0].objNormal.xyz   * bary.x + patch[1].objNormal.xyz   * bary.y + patch[2].objNormal.xyz   * bary.z;
    float3 objT = patch[0].objTangent.xyz  * bary.x + patch[1].objTangent.xyz  * bary.y + patch[2].objTangent.xyz  * bary.z;
    float3 objB = patch[0].objBinormal.xyz * bary.x + patch[1].objBinormal.xyz * bary.y + patch[2].objBinormal.xyz * bary.z;

    // skin旋转（仅3x3，法线/切线不需要平移）
    float3 worldNormal, worldTangent, worldBinormal;
    worldNormal.x   = dot(skinRow0, objN); worldNormal.y   = dot(skinRow1, objN); worldNormal.z   = dot(skinRow2, objN);
    worldTangent.x  = dot(skinRow0, objT); worldTangent.y  = dot(skinRow1, objT); worldTangent.z  = dot(skinRow2, objT);
    worldBinormal.x = dot(skinRow0, objB); worldBinormal.y = dot(skinRow1, objB); worldBinormal.z = dot(skinRow2, objB);

    // 视图空间变换（行向量 × row_major矩阵，与原始VS一致）
    float3 viewRow0 = patch[0].viewRow0.xyz;
    float3 viewRow1 = patch[0].viewRow1.xyz;
    float3 viewRow2 = patch[0].viewRow2.xyz;
    float3 viewSpaceNormal   = worldNormal.x   * viewRow0 + worldNormal.y   * viewRow1 + worldNormal.z   * viewRow2;
    float3 viewSpaceTangent  = worldTangent.x  * viewRow0 + worldTangent.y  * viewRow1 + worldTangent.z  * viewRow2;
    float3 viewSpaceBinormal = worldBinormal.x * viewRow0 + worldBinormal.y * viewRow1 + worldBinormal.z * viewRow2;

    output.viewNormal   = float4(viewSpaceNormal,   0);
    output.viewTangent  = float4(viewSpaceTangent,  0);
    output.viewBinormal = float4(viewSpaceBinormal, 0);

    return output;
}
