#pragma once
// Standalone reference math extension; not used by the two minimal executables.
// Requires C++20 and the Windows SDK's DirectXMath. Native compile is unverified.
#include <DirectXMath.h>
#include <algorithm>
#include <cstddef>
#include <cstdint>
#include <span>
#include <stdexcept>
#include <vector>
namespace qubic_reference {
using namespace DirectX;
using std::size_t;
template<class T> struct Key { float time; T value; };
struct JointPose {
    XMFLOAT3 translation{0,0,0};
    XMFLOAT4 rotation{0,0,0,1};
    XMFLOAT3 scale{1,1,1};
};
struct KeyInterval { size_t a, b; float alpha; };
template<class T> KeyInterval FindInterval(std::span<const Key<T>> keys, float time) {
    // Import contract: finite values, strictly increasing key times.
    if (keys.empty()) throw std::invalid_argument("Empty track has no interval");
    if (time <= keys.front().time) return {0,0,0};
    if (time >= keys.back().time) return {keys.size()-1,keys.size()-1,0};
    auto upper = std::upper_bound(keys.begin(),keys.end(),time,
        [](float t,const Key<T>& key){return t < key.time;});
    size_t b = static_cast<size_t>(upper-keys.begin()), a = b-1;
    const float duration = keys[b].time-keys[a].time;
    if (duration <= 0) throw std::invalid_argument("Track times must increase");
    return {a,b,std::clamp((time-keys[a].time)/duration,0.0f,1.0f)};
}
inline XMFLOAT3 SampleVector(std::span<const Key<XMFLOAT3>> keys, float time,
                            const XMFLOAT3& fallback) {
    if (keys.empty()) return fallback;
    auto interval = FindInterval(keys,time);
    XMFLOAT3 out;
    XMStoreFloat3(&out,XMVectorLerp(XMLoadFloat3(&keys[interval.a].value),
        XMLoadFloat3(&keys[interval.b].value),interval.alpha));
    return out;
}
inline XMFLOAT4 BlendRotation(const XMFLOAT4& from,const XMFLOAT4& to,float weight) {
    XMVECTOR a = XMQuaternionNormalize(XMLoadFloat4(&from));
    XMVECTOR b = XMQuaternionNormalize(XMLoadFloat4(&to));
    // q and -q represent one orientation; choose the same hemisphere.
    if (XMVectorGetX(XMQuaternionDot(a,b)) < 0) b = XMVectorNegate(b);
    XMFLOAT4 out;
    XMStoreFloat4(&out,XMQuaternionNormalize(XMQuaternionSlerp(a,b,weight)));
    return out;
}
inline XMFLOAT4 SampleRotation(std::span<const Key<XMFLOAT4>> keys,float time,
                              const XMFLOAT4& fallback) {
    if (keys.empty()) return fallback;
    auto interval = FindInterval(keys,time);
    return BlendRotation(keys[interval.a].value,keys[interval.b].value,interval.alpha);
}
inline JointPose BlendPose(const JointPose& a,const JointPose& b,float weight) {
    weight = std::clamp(weight,0.0f,1.0f);
    JointPose out;
    XMStoreFloat3(&out.translation,XMVectorLerp(XMLoadFloat3(&a.translation),XMLoadFloat3(&b.translation),weight));
    XMStoreFloat3(&out.scale,XMVectorLerp(XMLoadFloat3(&a.scale),XMLoadFloat3(&b.scale),weight));
    out.rotation = BlendRotation(a.rotation,b.rotation,weight);
    return out;
}
inline std::vector<XMFLOAT4X4> EvaluatePalette(std::span<const JointPose> pose,
        std::span<const int32_t> parents,std::span<const XMFLOAT4X4> inverseBind) {
    if (pose.size()!=parents.size() || pose.size()!=inverseBind.size())
        throw std::invalid_argument("Skeleton/palette size mismatch");
    std::vector<XMFLOAT4X4> globals(pose.size()),palette(pose.size());
    for (size_t i=0;i<pose.size();++i) {
        if (parents[i] < -1 || parents[i] >= static_cast<int64_t>(i))
            throw std::invalid_argument("Skeleton must be parent-before-child");
        const auto& p=pose[i];
        XMMATRIX local=XMMatrixScaling(p.scale.x,p.scale.y,p.scale.z)
            * XMMatrixRotationQuaternion(XMQuaternionNormalize(XMLoadFloat4(&p.rotation)))
            * XMMatrixTranslation(p.translation.x,p.translation.y,p.translation.z);
        XMMATRIX global=parents[i]>=0?local*XMLoadFloat4x4(&globals[parents[i]]):local;
        XMStoreFloat4x4(&globals[i],global);
        XMStoreFloat4x4(&palette[i],XMLoadFloat4x4(&inverseBind[i])*global);
    }
    return palette;
}
} // namespace qubic_reference
