diff --git a/Tools/VulkanWrapperGenerator/Generator.cs b/Tools/VulkanWrapperGenerator/Generator.cs new file mode 100644 index 0000000..74bfb02 --- /dev/null +++ b/Tools/VulkanWrapperGenerator/Generator.cs @@ -0,0 +1,735 @@ +using System; +using System.Collections.Generic; +using System.Collections.ObjectModel; +using System.IO; +using System.Linq; +using System.Text; +using System.Text.RegularExpressions; + +public sealed class Generator { + private readonly VkSpec _spec; + private readonly string _outDir; + private readonly ReadOnlyDictionary _sugar; + + public Generator(VkSpec spec, string outDir, IEnumerable? sugarSpecs = null) { + _spec = spec; + _outDir = outDir; + _sugar = new ReadOnlyDictionary( + (sugarSpecs ?? new List()) + .Where(s => !string.IsNullOrWhiteSpace(s.TargetClass)) + .GroupBy(s => s.TargetClass) + .ToDictionary(g => g.Key, g => { + var merged = new SugarClassSpec(g.Key) { + Includes = g.SelectMany(x => x.Includes).Distinct().ToList(), + Methods = g.SelectMany(x => x.Methods).ToList() + }; + return merged; + }) + ); + } + + public void GenerateAll() { + GenerateVulkanFunctions(); + GenerateHandleWrappers(); + GenerateProducerMemberMethods(); + GenerateMemberMethods(); + } + + // ---------- helpers ---------- + static string ToClass(string vkType) => vkType.StartsWith("Vk") ? vkType.Substring(2) : vkType; + + /// + /// Build a pretty method name from a vk* command for a given class. + /// Rules: + /// - Strip leading 'vk'. + /// - Strip the class token at BEGIN or END (Queue, Device, CommandBuffer...). + /// - For CommandBuffer also strip "Cmd". + /// - Remove first occurrence of class token inside the name if still present (e.g., PhysicalDevice). + /// - LowerCamelCase. + /// + static string MakeMethodName(string className, string vkFunc) { + string s = vkFunc.StartsWith("vk") ? vkFunc.Substring(2) : vkFunc; + + var tokens = new List { className }; + if (className == "CommandBuffer") tokens.Add("Cmd"); + + foreach (var t in tokens) { + if (s.StartsWith(t, StringComparison.Ordinal)) { s = s.Substring(t.Length); break; } + } + foreach (var t in tokens) { + if (s.EndsWith(t, StringComparison.Ordinal)) { s = s.Substring(0, s.Length - t.Length); break; } + } + int idx = s.IndexOf(className, StringComparison.Ordinal); + if (idx >= 0) s = s.Remove(idx, className.Length); + + if (s.Length == 0) return ""; + return char.ToLowerInvariant(s[0]) + s.Substring(1); + } + + static string Render(string tpl, Dictionary vars) { + foreach (var kv in vars) tpl = tpl.Replace("${" + kv.Key + "}", kv.Value); + return tpl; + } + + // ========================================= + // VulkanFunctions — function pointers loader + // ========================================= + void GenerateVulkanFunctions() { + var sbH = new StringBuilder(); + sbH.AppendLine("#pragma once"); + sbH.AppendLine("#include "); + sbH.AppendLine("#include "); + sbH.AppendLine("#include "); + sbH.AppendLine(); + sbH.AppendLine("struct VulkanFunctions {"); + + foreach (var cmd in _spec.Commands.Values.OrderBy(c => c.Name)) + sbH.AppendLine(" PFN_" + cmd.Name + " " + cmd.Name + " = nullptr;"); + + sbH.AppendLine(); + sbH.AppendLine(" void loadGlobal();"); + sbH.AppendLine(" void loadInstance(VkInstance instance);"); + sbH.AppendLine(" void loadDevice(VkDevice device);"); + sbH.AppendLine("};"); + sbH.AppendLine("extern VulkanFunctions vk;"); + File.WriteAllText(Path.Combine(_outDir, "VulkanFunctions.hpp"), sbH.ToString()); + + var sbCpp = new StringBuilder(); + sbCpp.AppendLine("#include \"VulkanFunctions.hpp\""); + sbCpp.AppendLine("#if defined(_WIN32)"); + sbCpp.AppendLine("#include "); + sbCpp.AppendLine("static HMODULE g_vkLib = LoadLibraryA(\"vulkan-1.dll\");"); + sbCpp.AppendLine("static void* LoadSymbol(const char* n){ return (void*)GetProcAddress(g_vkLib, n); }"); + sbCpp.AppendLine("#else"); + sbCpp.AppendLine("#include "); + sbCpp.AppendLine("static void* g_vkLib = dlopen(\"libvulkan.so.1\", RTLD_NOW | RTLD_LOCAL);"); + sbCpp.AppendLine("static void* LoadSymbol(const char* n){ return dlsym(g_vkLib, n); }"); + sbCpp.AppendLine("#endif"); + sbCpp.AppendLine("VulkanFunctions vk;"); + sbCpp.AppendLine(); + sbCpp.AppendLine("void VulkanFunctions::loadGlobal(){"); + sbCpp.AppendLine(" vkGetInstanceProcAddr = (PFN_vkGetInstanceProcAddr)LoadSymbol(\"vkGetInstanceProcAddr\");"); + foreach (var cmd in _spec.Commands.Values.Where(c => _spec.ClassifyLoadLevel(c) == VkSpec.LoadLevel.Global).OrderBy(c => c.Name)) + sbCpp.AppendLine(" " + cmd.Name + " = (PFN_" + cmd.Name + ")vkGetInstanceProcAddr(nullptr, \"" + cmd.Name + "\");"); + sbCpp.AppendLine("}"); + sbCpp.AppendLine(); + sbCpp.AppendLine("void VulkanFunctions::loadInstance(VkInstance instance){"); + foreach (var cmd in _spec.Commands.Values.Where(c => _spec.ClassifyLoadLevel(c) == VkSpec.LoadLevel.Instance).OrderBy(c => c.Name)) + sbCpp.AppendLine(" " + cmd.Name + " = (PFN_" + cmd.Name + ")vkGetInstanceProcAddr(instance, \"" + cmd.Name + "\");"); + sbCpp.AppendLine("}"); + sbCpp.AppendLine(); + sbCpp.AppendLine("void VulkanFunctions::loadDevice(VkDevice device){"); + sbCpp.AppendLine(" vkGetDeviceProcAddr = (PFN_vkGetDeviceProcAddr)vkGetInstanceProcAddr(nullptr, \"vkGetDeviceProcAddr\");"); + foreach (var cmd in _spec.Commands.Values.Where(c => _spec.ClassifyLoadLevel(c) == VkSpec.LoadLevel.Device).OrderBy(c => c.Name)) + sbCpp.AppendLine(" " + cmd.Name + " = (PFN_" + cmd.Name + ")vkGetDeviceProcAddr(device, \"" + cmd.Name + "\");"); + sbCpp.AppendLine("}"); + File.WriteAllText(Path.Combine(_outDir, "VulkanFunctions.cpp"), sbCpp.ToString()); + } + + // ========================================= + // RAII wrappers (owned + borrowed) + // ========================================= + void GenerateHandleWrappers() { + var handlesToGenerate = new HashSet(_spec.Pairs.Select(p => p.HandleType)); + foreach (var pr in _spec.Producers) handlesToGenerate.Add(pr.HandleType); + + foreach (var handleType in handlesToGenerate) { + var className = ToClass(handleType); + + VkSpec.Pair? pair = _spec.Pairs.FirstOrDefault(p => p.HandleType == handleType); + var hasOwned = pair != null; + var primaryOwners = new List(); + + if (hasOwned) { + primaryOwners.Add(pair!.ParentType); + primaryOwners.AddRange(pair!.ExtraOwnersForDestroy.Where(o => o != pair!.ParentType)); + } else { + var parents = _spec.Producers.Where(p => p.HandleType == handleType) + .Select(p => p.ParentType) + .Distinct().ToList(); + if (parents.Count > 0) primaryOwners.AddRange(parents); + } + primaryOwners = primaryOwners.Distinct().ToList(); + + var lifetimeUses = new List(); + if (hasOwned) { + var createCmd = _spec.Commands[pair!.CreateOrAllocate]; + lifetimeUses = _spec.AnalyzeCreateInfoOwners(createCmd, handleType) + .Where(u => !primaryOwners.Contains(u.Type)) + .ToList(); + } + + var header = new StringBuilder(); + header.AppendLine("#pragma once"); + header.AppendLine("#include "); + header.AppendLine("#include "); + header.AppendLine("#include "); + header.AppendLine("#include \"VulkanFunctions.hpp\""); + header.AppendLine(); + + header.AppendLine("// === Lifetime summary (deterministic) ==="); + header.AppendLine("// Handle: " + handleType); + var holdsStr = primaryOwners.Count > 0 ? string.Join(", ", primaryOwners.Select(ToClass)) : "—"; + header.AppendLine("// Holds (destroy/free owners + creation parent, or producer parent): " + holdsStr); + var keepList = lifetimeUses.Select(u => ToClass(u.Type) + (u.IsMany ? "[]" : "") + " via " + u.SourceStruct + "::" + u.SourceMember); + var keepsStr = keepList.Any() ? string.Join("; ", keepList) : "—"; + header.AppendLine("// Keeps (from *CreateInfo/*AllocateInfo): " + keepsStr); + header.AppendLine("// Destructor calls vkDestroy*/vkFree* only if _owned == true."); + header.AppendLine(); + + header.AppendLine("class " + className + " : public RefCounted {"); + header.AppendLine("public:"); + header.AppendLine(" " + handleType + " getHandle() const noexcept { return _handle; }"); + + var usedGetterNames = new HashSet(StringComparer.Ordinal); + for (int i = 0; i < primaryOwners.Count; i++) { + var oc = ToClass(primaryOwners[i]); + var getter = "get" + oc; var unique = getter; int k = 2; + while (!usedGetterNames.Add(unique)) unique = getter + k++; + header.AppendLine(" IntrusivePtr<" + oc + "> " + unique + "() const noexcept { return _owner" + i + "; }"); + } + foreach (var u in lifetimeUses) { + var oc = ToClass(u.Type); + var baseName = "get" + oc + (u.IsMany ? "s" : ""); + var unique = baseName; int k = 2; + while (!usedGetterNames.Add(unique)) unique = baseName + k++; + if (u.IsMany) + header.AppendLine(" const std::vector>& " + unique + "() const noexcept { return _lifetime_owners_" + oc + "; }"); + else + header.AppendLine(" IntrusivePtr<" + oc + "> " + unique + "() const noexcept { return _lifetime_owner_" + oc + "; }"); + } + + // adopt raw handle (borrowed) + if (primaryOwners.Count > 0) { + header.Append("public:\n static IntrusivePtr<" + className + "> fromHandle("); + for (int i = 0; i < primaryOwners.Count; i++) { + if (i > 0) header.Append(", "); + header.Append("IntrusivePtr<" + ToClass(primaryOwners[i]) + "> owner" + i); + } + if (primaryOwners.Count > 0) header.Append(", "); + header.Append(handleType + " handle"); + header.AppendLine(") {"); + header.Append(" return IntrusivePtr<" + className + ">(new " + className + "("); + for (int i = 0; i < primaryOwners.Count; i++) { + if (i > 0) header.Append(", "); + header.Append("owner" + i); + } + if (primaryOwners.Count > 0) header.Append(", "); + header.Append("handle, false"); + foreach (var u in lifetimeUses) { + var cls = ToClass(u.Type); + header.Append(", "); + header.Append(u.IsMany ? "std::vector>()" : "IntrusivePtr<" + cls + ">()"); + } + header.AppendLine("));"); + header.AppendLine(" }"); + } + + // sugar methods + if (_sugar.TryGetValue(className, out var sugarSpec)) { + foreach (var m in sugarSpec.Methods) { + if (!m.RequiresFunctions.All(fn => _spec.Commands.ContainsKey(fn))) continue; + if (!m.RequiresExtensions.All(ext => _spec.Extensions.Contains(ext))) continue; + + var (sig, prelude, vars, guard) = BuildSugarSignatureAndContext(m, primaryOwners, className); + + if (!string.IsNullOrEmpty(guard)) header.AppendLine("#if " + guard); + header.AppendLine(" " + sig + " {"); + if (!string.IsNullOrEmpty(prelude)) + foreach (var line in prelude.Split('\n')) + if (line.Length > 0) header.AppendLine(" " + line); + foreach (var ln in m.BodyLines) // FIX: only the current sugar method body + header.AppendLine(" " + Render(ln, vars)); + header.AppendLine(" }"); + if (!string.IsNullOrEmpty(guard)) header.AppendLine("#endif"); + } + } + + // owned factories + if (hasOwned) { + var multiOut = pair!.OutHandleParam.LenAttr != null; + var fac = BuildFactory(_spec.Commands[pair!.CreateOrAllocate], pair!, className, ToClass(pair!.ParentType), primaryOwners, lifetimeUses, multiOut); + header.Append(fac.decl); + header.Append(fac.body); + header.AppendLine(); + } + + // private ctor/dtor/fields + header.AppendLine("private:"); + header.Append(" explicit " + className + "("); + for (int i = 0; i < primaryOwners.Count; i++) { + if (i > 0) header.Append(", "); + header.Append("IntrusivePtr<" + ToClass(primaryOwners[i]) + "> owner" + i); + } + if (primaryOwners.Count > 0) header.Append(", "); + header.Append(handleType + " handle, bool owned"); + foreach (var u in lifetimeUses) { + header.Append(", "); + var cls = ToClass(u.Type); + header.Append(u.IsMany + ? "const std::vector>& owners_" + cls + : "IntrusivePtr<" + cls + "> owner_" + cls); + } + header.AppendLine(") : _handle(handle), _owned(owned)"); + for (int i = 0; i < primaryOwners.Count; i++) + header.AppendLine(" , _owner" + i + "(std::move(owner" + i + "))"); + foreach (var u in lifetimeUses) { + var cls = ToClass(u.Type); + header.AppendLine(u.IsMany + ? " , _lifetime_owners_" + cls + "(owners_" + cls + ")" + : " , _lifetime_owner_" + cls + "(std::move(owner_" + cls + "))"); + } + header.AppendLine(" {}"); + + header.AppendLine(" ~" + className + "() override {"); + header.AppendLine(" if (_handle && _owned) {"); + if (hasOwned) { + var destroyCmd = _spec.Commands[pair!.DestroyOrFree]; + var args = new List(); + int ownerIdx = 0; + foreach (var p in destroyCmd.Params) { + if (p.Type == handleType && p.IsPointer) { args.Add("1"); args.Add("&_handle"); continue; } + if (p.Type == handleType && !p.IsPointer) { args.Add("_handle"); continue; } + if (p.Name == "pAllocator") { args.Add("nullptr"); continue; } + if (p.IsHandle) { args.Add("_owner" + (ownerIdx++) + "->getHandle()"); continue; } + if (p.Name.EndsWith("Count")) args.Add("0"); else args.Add("0"); + } + header.AppendLine(" (void)vk." + destroyCmd.Name + "(" + string.Join(", ", args) + ");"); + header.AppendLine(" _handle = VK_NULL_HANDLE;"); + } + header.AppendLine(" }"); + header.AppendLine(" }"); + + header.AppendLine(" " + handleType + " _handle{VK_NULL_HANDLE};"); + header.AppendLine(" bool _owned{false};"); + for (int i = 0; i < primaryOwners.Count; i++) + header.AppendLine(" IntrusivePtr<" + ToClass(primaryOwners[i]) + "> _owner" + i + ";"); + foreach (var u in lifetimeUses) { + var cls = ToClass(u.Type); + header.AppendLine(u.IsMany + ? " std::vector> _lifetime_owners_" + cls + ";" + : " IntrusivePtr<" + cls + "> _lifetime_owner_" + cls + ";"); + } + + header.AppendLine("};"); + + File.WriteAllText(Path.Combine(_outDir, className + ".hpp"), header.ToString()); + } + } + + // Build sugar signature + prelude + placeholders + guard + (string sig, string prelude, Dictionary vars, string guard) + BuildSugarSignatureAndContext(SugarMethodSpec m, List primaryOwners, string className) { + var vars = new Dictionary { + ["class"] = className, + ["handle"] = "_handle", + ["vk"] = "vk", + ["selfOwners"] = string.Join(", ", Enumerable.Range(0, primaryOwners.Count).Select(i => "_owner" + i)) + }; + for (int i = 0; i < primaryOwners.Count; i++) { + vars["owner" + i] = "_owner" + i; + vars["owner" + i + ".handle"] = "_owner" + i + "->getHandle()"; + } + + var prelude = new StringBuilder(); + var paramList = string.Join(", ", m.Parameters); + + var rxScalar = new Regex(@"\bIntrusivePtr<\s*(\w+)\s*>\s+(\w+)\b"); + foreach (Match mt in rxScalar.Matches(paramList)) { + var varName = mt.Groups[2].Value; + vars[varName + ".handle"] = varName + "->getHandle()"; + } + + var rxVector = new Regex(@"\bstd::vector\s*<\s*IntrusivePtr<\s*(\w+)\s*>\s*>\s*&?\s*(\w+)\b"); + foreach (Match mt in rxVector.Matches(paramList)) { + var elemClass = mt.Groups[1].Value; + var varName = mt.Groups[2].Value; + var vkType = "Vk" + elemClass; + + var rawName = "_raw_" + varName; + prelude.AppendLine("std::vector<" + vkType + "> " + rawName + "; " + rawName + ".reserve(" + varName + ".size());"); + prelude.AppendLine("for (auto const& it : " + varName + ") " + rawName + ".push_back(it->getHandle());"); + + vars[varName + ".handles"] = rawName + ".data()"; + vars[varName + ".size"] = varName + ".size()"; + vars[varName + ".count"] = varName + ".size()"; + } + + var quals = new StringBuilder(); + if (m.IsConst) quals.Append(" const"); + if (m.IsNoexcept) quals.Append(" noexcept"); + + var sig = (m.IsStatic ? "static " : "") + m.ReturnType + " " + m.Name + "(" + string.Join(", ", m.Parameters) + ")" + quals.ToString(); + + string guard = ""; + var parts = new List(); + if (!string.IsNullOrWhiteSpace(m.Guard)) parts.Add(m.Guard!); + if (m.RequiresExtensions.Count > 0) + parts.Add(string.Join(" && ", m.RequiresExtensions.Select(e => "defined(" + e + ")"))); + if (parts.Count > 0) guard = string.Join(" && ", parts); + + return (sig, prelude.ToString(), vars, guard); + } + + // Owned factory: declaration + body (+ comment with vkCreate*) + (string decl, string body) BuildFactory( + VkSpec.Command create, + VkSpec.Pair pair, + string className, + string parentClass, + List primaryOwners, + List lifetimeUses, + bool multiOut) { + var sbDecl = new StringBuilder(); + var sbBody = new StringBuilder(); + + var ps = create.Params.Skip(1).ToList(); + var outLenName = pair.OutHandleParam.LenAttr; + var lenNamesToSkip = new HashSet(); + + var sigParams = new List(); + var prep = new StringBuilder(); + var callArgs = new List(); + + sbDecl.Append("public:\n "); + // <-- Comment with original vk command + sbDecl.Append("// Original: " + create.Name + "\n "); + if (!multiOut) + sbDecl.Append("static IntrusivePtr<" + className + "> create(IntrusivePtr<" + parentClass + "> parent"); + else + sbDecl.Append("static std::vector> createMany(IntrusivePtr<" + parentClass + "> parent"); + + foreach (var p in ps) { + if (p.Name == "pAllocator") continue; + if (p.IsPointer && p.Type == pair.HandleType) continue; + if (outLenName != null && p.Name == outLenName) continue; + + if (p.IsHandle) { + var cls = ToClass(p.Type); + if (p.IsPointer && p.LenAttr != null) { + sigParams.Add("const std::vector>& " + p.Name); + var rawVec = "_raw_" + p.Name; + prep.AppendLine(" std::vector<" + p.Type + "> " + rawVec + "; " + rawVec + ".reserve(" + p.Name + ".size());"); + prep.AppendLine(" for (auto const& it : " + p.Name + ") " + rawVec + ".push_back(it->getHandle());"); + lenNamesToSkip.Add(p.LenAttr); + callArgs.Add(p.Name + ".size()"); + callArgs.Add(rawVec + ".data()"); + } else { + sigParams.Add("IntrusivePtr<" + cls + "> " + p.Name); + callArgs.Add(p.Name + "->getHandle()"); + } + } else { + if (lenNamesToSkip.Contains(p.Name)) continue; + sigParams.Add((p.IsConst ? "const " : "") + p.Type + (p.IsPointer ? "*" : "") + " " + p.Name); + callArgs.Add(p.Name); + } + } + + var lifeSig = new List(); + foreach (var u in lifetimeUses) { + var cls = ToClass(u.Type); + lifeSig.Add(u.IsMany + ? "const std::vector>& owners_" + cls + : "IntrusivePtr<" + cls + "> owner_" + cls); + } + + if (sigParams.Count > 0) sbDecl.Append(", "); + sbDecl.Append(string.Join(", ", sigParams)); + if (lifeSig.Count > 0) { + if (sigParams.Count > 0) sbDecl.Append(", "); + sbDecl.Append(string.Join(", ", lifeSig)); + } + sbDecl.AppendLine(") {"); + + sbBody.Append(prep.ToString()); + + callArgs.Add("nullptr"); + + var outName = "_out"; + if (!multiOut) { + sbBody.AppendLine(" " + pair.HandleType + " " + outName + " = VK_NULL_HANDLE;"); + callArgs.Add("&" + outName); + } else { + var lenName = outLenName ?? "count"; + sbBody.AppendLine(" std::vector<" + pair.HandleType + "> " + outName + ";"); + sbBody.AppendLine(" " + outName + ".resize(" + lenName + ");"); + callArgs.Add(outName + ".data()"); + } + + var argsFinal = new List { "parent->getHandle()" }; + argsFinal.AddRange(callArgs); + sbBody.Append(" auto _res = vk." + create.Name + "(" + string.Join(", ", argsFinal) + ");\n"); + sbBody.AppendLine(" if (_res != VK_SUCCESS) throw std::runtime_error(\"" + create.Name + " failed\");"); + + if (!multiOut) { + sbBody.Append(" return IntrusivePtr<" + className + ">(new " + className + "("); + for (int i = 0; i < primaryOwners.Count; ++i) { + if (i == 0) sbBody.Append("parent"); + else { + var t = primaryOwners[i]; + var pFound = ps.FirstOrDefault(pp => pp.IsHandle && pp.Type == t); + sbBody.Append(", " + (pFound != null ? (pFound.IsPointer && pFound.LenAttr != null ? "(" + ToClass(t) + ")nullptr /*from array*/" : pFound.Name) : "parent")); + } + } + if (primaryOwners.Count > 0) sbBody.Append(", "); + sbBody.Append(outName); + sbBody.Append(", true"); + foreach (var u in lifetimeUses) { + var cls = ToClass(u.Type); + sbBody.Append(", "); + sbBody.Append(u.IsMany ? "owners_" + cls : "owner_" + cls); + } + sbBody.AppendLine("));"); + } else { + sbBody.AppendLine(" std::vector> _ret; _ret.reserve(" + outName + ".size());"); + sbBody.AppendLine(" for (auto h : " + outName + ") {"); + sbBody.Append(" _ret.emplace_back(new " + className + "("); + for (int i = 0; i < primaryOwners.Count; ++i) { + if (i == 0) sbBody.Append("parent"); + else { + var t = primaryOwners[i]; + var pFound = ps.FirstOrDefault(pp => pp.IsHandle && pp.Type == t); + sbBody.Append(", " + (pFound != null ? (pFound.IsPointer && pFound.LenAttr != null ? "(" + ToClass(t) + ")nullptr /*from array*/" : pFound.Name) : "parent")); + } + } + if (primaryOwners.Count > 0) sbBody.Append(", "); + sbBody.Append("h"); + sbBody.Append(", true"); + foreach (var u in lifetimeUses) { + var cls = ToClass(u.Type); + sbBody.Append(", "); + sbBody.Append(u.IsMany ? "owners_" + cls : "owner_" + cls); + } + sbBody.AppendLine("));"); + sbBody.AppendLine(" }"); + sbBody.AppendLine(" return _ret;"); + } + + sbBody.AppendLine(" }"); + + return (sbDecl.ToString(), sbBody.ToString()); + } + + // ========================================= + // Producer methods (Get*/Enumerate*) → parent->getChild() + // ========================================= + void GenerateProducerMemberMethods() { + var headers = Directory.GetFiles(_outDir, "*.hpp") + .ToDictionary(f => Path.GetFileNameWithoutExtension(f), f => new StringBuilder(File.ReadAllText(f))); + + foreach (var pr in _spec.Producers) { + var parentClass = ToClass(pr.ParentType); + var childClass = ToClass(pr.HandleType); + if (!headers.TryGetValue(parentClass, out var h)) { + var minimal = new StringBuilder(); + minimal.AppendLine("#pragma once"); + minimal.AppendLine("#include "); + minimal.AppendLine("#include "); + minimal.AppendLine("#include \"VulkanFunctions.hpp\""); + minimal.AppendLine("class " + parentClass + " : public RefCounted {"); + minimal.AppendLine("public:"); + minimal.AppendLine(" " + pr.ParentType + " getHandle() const noexcept { return _handle; }"); + minimal.AppendLine("private:"); + minimal.AppendLine(" " + pr.ParentType + " _handle{VK_NULL_HANDLE};"); + minimal.AppendLine("};"); + h = minimal; + headers[parentClass] = h; + } + + var cmd = _spec.Commands[pr.CommandName]; + + var outParam = pr.OutHandleParam; + var lenName = outParam.LenAttr; + var sigParams = new List(); + var callArgs = new List(); + foreach (var p in cmd.Params) { + if (p == cmd.Params[0]) continue; + if (p.Name == outParam.Name) continue; + if (!string.IsNullOrEmpty(lenName) && p.Name == lenName) continue; + + if (p.IsHandle) { + var cls = ToClass(p.Type); + sigParams.Add("IntrusivePtr<" + cls + "> " + p.Name); + callArgs.Add(p.Name + "->getHandle()"); + } else { + sigParams.Add((p.IsConst ? "const " : "") + p.Type + (p.IsPointer ? "*" : "") + " " + p.Name); + callArgs.Add(p.Name); + } + } + + var retWrapper = pr.IsMany + ? "std::vector>" + : "IntrusivePtr<" + childClass + ">"; + + var methodName = "get" + childClass + (pr.IsMany ? "s" : ""); + var method = new StringBuilder(); + method.Append("public:\n "); + // <-- Comment with original vkGet*/vkEnumerate* + method.Append("// Original: " + cmd.Name + "\n "); + method.Append(retWrapper + " " + methodName + "(" + string.Join(", ", sigParams) + ") {\n"); + if (pr.IsMany) { + var countParam = cmd.Params.FirstOrDefault(p => p.Name == lenName)!; + var countType = countParam.Type; + method.Append(" " + countType + " _count = 0;\n"); + + var firstArgs = new List { "_handle" }; + foreach (var p in cmd.Params.Skip(1)) { + if (p.Name == lenName) firstArgs.Add("&_count"); + else if (p.Name == outParam.Name) firstArgs.Add("nullptr"); + else if (p.IsHandle) firstArgs.Add(p.Name + "->getHandle()"); + else firstArgs.Add(p.Name); + } + var call1 = "vk." + cmd.Name + "(" + string.Join(", ", firstArgs) + ")"; + if (cmd.ReturnType.Trim() == "VkResult") { + method.Append(" auto _r1 = " + call1 + ";\n"); + method.Append(" if (_r1 != VK_SUCCESS && _r1 != VK_INCOMPLETE) throw std::runtime_error(\"" + cmd.Name + " (count) failed\");\n"); + } else { + method.Append(" " + call1 + ";\n"); + } + + method.Append(" std::vector<" + pr.HandleType + "> _out;\n"); + method.Append(" _out.resize(_count);\n"); + + var secondArgs = new List { "_handle" }; + foreach (var p in cmd.Params.Skip(1)) { + if (p.Name == lenName) secondArgs.Add("&_count"); + else if (p.Name == outParam.Name) secondArgs.Add("_out.data()"); + else if (p.IsHandle) secondArgs.Add(p.Name + "->getHandle()"); + else secondArgs.Add(p.Name); + } + var call2 = "vk." + cmd.Name + "(" + string.Join(", ", secondArgs) + ")"; + if (cmd.ReturnType.Trim() == "VkResult") { + method.Append(" auto _r2 = " + call2 + ";\n"); + method.Append(" if (_r2 != VK_SUCCESS) throw std::runtime_error(\"" + cmd.Name + " (fill) failed\");\n"); + } else { + method.Append(" " + call2 + ";\n"); + } + + method.Append(" std::vector> _ret; _ret.reserve(_out.size());\n"); + method.Append(" IntrusivePtr<" + parentClass + "> _self(this);\n"); + method.Append(" for (auto h : _out) {\n"); + method.Append(" _ret.emplace_back(new " + childClass + "(_self, h, false));\n"); + method.Append(" }\n"); + method.Append(" return _ret;\n"); + } else { + method.Append(" " + pr.HandleType + " _h = VK_NULL_HANDLE;\n"); + var args = new List { "_handle" }; + foreach (var p in cmd.Params.Skip(1)) { + if (p.Name == outParam.Name) args.Add("&_h"); + else if (p.IsHandle) args.Add(p.Name + "->getHandle()"); + else args.Add(p.Name); + } + var call = "vk." + cmd.Name + "(" + string.Join(", ", args) + ")"; + if (cmd.ReturnType.Trim() == "VkResult") { + method.Append(" auto _r = " + call + ";\n"); + method.Append(" if (_r != VK_SUCCESS) throw std::runtime_error(\"" + cmd.Name + " failed\");\n"); + } else { + method.Append(" " + call + ";\n"); + } + method.Append(" return IntrusivePtr<" + childClass + ">(new " + childClass + "(IntrusivePtr<" + parentClass + ">(this), _h, false));\n"); + } + method.Append(" }\n"); + + var s = h.ToString(); + var idx = s.LastIndexOf("};", StringComparison.Ordinal); + s = s.Insert(idx, method.ToString()); + headers[parentClass] = new StringBuilder(s); + } + + foreach (var kv in headers) + File.WriteAllText(Path.Combine(_outDir, kv.Key + ".hpp"), kv.Value.ToString()); + } + + // ========================================= + // vk* → class methods (skip producer-out-handle) with pretty names + // ========================================= + void GenerateMemberMethods() { + var headers = Directory.GetFiles(_outDir, "*.hpp") + .ToDictionary(f => Path.GetFileNameWithoutExtension(f), f => new StringBuilder(File.ReadAllText(f))); + + foreach (var cmd in _spec.Commands.Values) { + if (cmd.Name.StartsWith("vkCreate") || cmd.Name.StartsWith("vkDestroy") || + cmd.Name.StartsWith("vkAllocate") || cmd.Name.StartsWith("vkFree")) + continue; + + bool isProducerOutHandle = + (cmd.Name.StartsWith("vkGet") || cmd.Name.StartsWith("vkEnumerate")) && + cmd.Params.Any(p => p.IsPointer && _spec.Handles.ContainsKey(p.Type)); + if (isProducerOutHandle) continue; + + if (cmd.Params.Count == 0) continue; + var first = cmd.Params[0]; + if (!_spec.Handles.ContainsKey(first.Type)) continue; + + var className = ToClass(first.Type); + if (!headers.TryGetValue(className, out var h)) { + var minimal = new StringBuilder(); + minimal.AppendLine("#pragma once"); + minimal.AppendLine("#include "); + minimal.AppendLine("#include "); + minimal.AppendLine("#include \"VulkanFunctions.hpp\""); + minimal.AppendLine("class " + className + " : public RefCounted {"); + minimal.AppendLine("public:"); + minimal.AppendLine(" " + first.Type + " getHandle() const noexcept { return _handle; }"); + minimal.AppendLine("private:"); + minimal.AppendLine(" " + first.Type + " _handle{VK_NULL_HANDLE};"); + minimal.AppendLine("};"); + h = minimal; + headers[className] = h; + } + + var methodName = MakeMethodName(className, cmd.Name); + var retType = string.IsNullOrWhiteSpace(cmd.ReturnType) ? "void" : cmd.ReturnType; + + var ps = cmd.Params.Skip(1).ToList(); + var sigParams = new List(); + var prep = new StringBuilder(); + var callArgs = new List(); + var lenNamesToSkip = new HashSet(); + + foreach (var p in ps) { + if (p.IsHandle) { + var cls = ToClass(p.Type); + if (p.IsPointer && p.LenAttr != null) { + sigParams.Add("const std::vector>& " + p.Name); + var raw = "_raw_" + p.Name; + prep.AppendLine(" std::vector<" + p.Type + "> " + raw + "; " + raw + ".reserve(" + p.Name + ".size());"); + prep.AppendLine(" for (auto const& it : " + p.Name + ") " + raw + ".push_back(it->getHandle());"); + lenNamesToSkip.Add(p.LenAttr); + callArgs.Add(p.Name + ".size()"); + callArgs.Add(raw + ".data()"); + } else { + sigParams.Add("IntrusivePtr<" + cls + "> " + p.Name); + callArgs.Add(p.Name + "->getHandle()"); + } + } else { + if (lenNamesToSkip.Contains(p.Name)) continue; + sigParams.Add((p.IsConst ? "const " : "") + p.Type + (p.IsPointer ? "*" : "") + " " + p.Name); + callArgs.Add(p.Name); + } + } + + var method = new StringBuilder(); + method.Append("public:\n "); + // <-- Comment with original vk* name + method.Append("// Original: " + cmd.Name + "\n "); + method.Append(retType + " " + methodName + "(" + string.Join(", ", sigParams) + ") {\n"); + method.Append(prep.ToString()); + method.Append(" "); + if (retType != "void") method.Append("return "); + var call = "vk." + cmd.Name + "(_handle"; + if (callArgs.Count > 0) call += ", " + string.Join(", ", callArgs); + call += ");"; + method.Append(call + "\n"); + method.Append(" }\n"); + + var s = h.ToString(); + var idx = s.LastIndexOf("};", StringComparison.Ordinal); + s = s.Insert(idx, method.ToString()); + headers[className] = new StringBuilder(s); + } + + foreach (var kv in headers) + File.WriteAllText(Path.Combine(_outDir, kv.Key + ".hpp"), kv.Value.ToString()); + } +} diff --git a/Tools/VulkanWrapperGenerator/Program.cs b/Tools/VulkanWrapperGenerator/Program.cs new file mode 100644 index 0000000..5b0f932 --- /dev/null +++ b/Tools/VulkanWrapperGenerator/Program.cs @@ -0,0 +1,141 @@ +using System; +using System.Collections.Generic; +using System.IO; +//using VulkanWrapperGenerator; + +class Program { + static string ResolveVkXmlPath(string[] args) { + // 1) Explicit path via arg + if (args.Length > 0 && File.Exists(args[0])) return args[0]; + + // 2) VULKAN_SDK based locations + var sdk = Environment.GetEnvironmentVariable("VULKAN_SDK"); + if (!string.IsNullOrWhiteSpace(sdk)) { + var p1 = Path.Combine(sdk, "share", "vulkan", "registry", "vk.xml"); + if (File.Exists(p1)) return p1; + var p2 = Path.Combine(sdk, "Include", "vulkan", "vk.xml"); + if (File.Exists(p2)) return p2; + } + + // 3) Typical Windows SDK layout (adjust if needed) + var winShare = Path.Combine(@"C:\VulkanSDK", "1.4.313.0", "share", "vulkan", "registry", "vk.xml"); + if (File.Exists(winShare)) return winShare; + + throw new FileNotFoundException( + "vk.xml not found. Pass a path as the first argument or set VULKAN_SDK.\n" + + "Searched: , " + + (sdk != null ? $"{sdk}\\share\\vulkan\\registry\\vk.xml, {sdk}\\Include\\vulkan\\vk.xml" : "") + + ", C:\\VulkanSDK\\1.4.313.0\\share\\vulkan\\registry\\vk.xml"); + } + + static void Main(string[] args) { + var xmlPath = ResolveVkXmlPath(args); + var spec = VkSpec.Load(xmlPath); + + // External "sugar" methods you want to inject into wrappers + var sugar = new List + { + new SugarClassSpec("Buffer") + { + Includes = { "" }, + Methods = + { + new SugarMethodSpec + { + Name = "uploadData", + ReturnType = "void", + Parameters = { + "const void* data", + "size_t size", + "IntrusivePtr memory", + "VkDeviceSize offset = 0" + }, + RequiresFunctions = { "vkMapMemory", "vkUnmapMemory", "vkBindBufferMemory" }, + BodyLines = { + "void* _p = nullptr;", + "auto _r = ${vk}.vkMapMemory(${owner0.handle}, ${memory.handle}, offset, size, 0, &_p);", + "if (_r != VK_SUCCESS) throw std::runtime_error(\"vkMapMemory failed\");", + "std::memcpy(_p, data, size);", + "${vk}.vkUnmapMemory(${owner0.handle});", + "(void)${vk}.vkBindBufferMemory(${owner0.handle}, ${handle}, ${memory.handle}, 0);" + } + } + } + }, + + new SugarClassSpec("Image") + { + Methods = + { + new SugarMethodSpec + { + Name = "transitionLayout", + ReturnType = "void", + Parameters = { + "IntrusivePtr cmd", + "VkImageLayout oldLayout", + "VkImageLayout newLayout", + "VkImageAspectFlags aspect = VK_IMAGE_ASPECT_COLOR_BIT" + }, + RequiresFunctions = { "vkCmdPipelineBarrier" }, + RequiresExtensions = { "VK_KHR_synchronization2" }, // example of extension guard + BodyLines = { + "VkImageMemoryBarrier barrier{ VK_STRUCTURE_TYPE_IMAGE_MEMORY_BARRIER };", + "barrier.oldLayout = oldLayout;", + "barrier.newLayout = newLayout;", + "barrier.srcQueueFamilyIndex = VK_QUEUE_FAMILY_IGNORED;", + "barrier.dstQueueFamilyIndex = VK_QUEUE_FAMILY_IGNORED;", + "barrier.image = ${handle};", + "barrier.subresourceRange.aspectMask = aspect;", + "barrier.subresourceRange.baseMipLevel = 0;", + "barrier.subresourceRange.levelCount = VK_REMAINING_MIP_LEVELS;", + "barrier.subresourceRange.baseArrayLayer = 0;", + "barrier.subresourceRange.layerCount = VK_REMAINING_ARRAY_LAYERS;", + "VkPipelineStageFlags srcStage = VK_PIPELINE_STAGE_ALL_COMMANDS_BIT;", + "VkPipelineStageFlags dstStage = VK_PIPELINE_STAGE_ALL_COMMANDS_BIT;", + "${vk}.vkCmdPipelineBarrier(${cmd.handle}, srcStage, dstStage, 0, 0, nullptr, 0, nullptr, 1, &barrier);" + } + } + } + }, + + new SugarClassSpec("Device") + { + Includes = { "", "" }, + Methods = + { + new SugarMethodSpec + { + Name = "createShaderModuleFromFile", + ReturnType = "IntrusivePtr", + Parameters = { "const char* path" }, + RequiresFunctions = { "vkCreateShaderModule" }, + BodyLines = { + "std::ifstream f(path, std::ios::binary | std::ios::ate);", + "if (!f) throw std::runtime_error(\"Can't open SPV file\");", + "auto sz = (size_t)f.tellg(); f.seekg(0);", + "std::vector code((sz + 3) / 4);", + "f.read(reinterpret_cast(code.data()), sz);", + "VkShaderModuleCreateInfo ci{ VK_STRUCTURE_TYPE_SHADER_MODULE_CREATE_INFO };", + "ci.codeSize = sz;", + "ci.pCode = code.data();", + "VkShaderModule h = VK_NULL_HANDLE;", + "auto r = ${vk}.vkCreateShaderModule(${handle}, &ci, nullptr, &h);", + "if (r != VK_SUCCESS) throw std::runtime_error(\"vkCreateShaderModule failed\");", + "return IntrusivePtr(new ShaderModule(${selfOwners}, h));" + } + } + } + } + }; + + var outDir = "Generated"; + Directory.CreateDirectory(outDir); + + var gen = new Generator(spec, outDir, sugar); + gen.GenerateAll(); + + Console.WriteLine("✅ Generated into ./Generated"); + Console.WriteLine("vk.xml used: " + xmlPath); + } +} diff --git a/Tools/VulkanWrapperGenerator/Sugar.cs b/Tools/VulkanWrapperGenerator/Sugar.cs new file mode 100644 index 0000000..349a7ed --- /dev/null +++ b/Tools/VulkanWrapperGenerator/Sugar.cs @@ -0,0 +1,29 @@ +using System.Collections.Generic; + +public sealed class SugarMethodSpec { + public string Name { get; init; } = ""; + public string ReturnType { get; init; } = "void"; + public List Parameters { get; init; } = new(); + public List BodyLines { get; init; } = new(); + + // Optional compile-time guard (combined with RequiresExtensions) + public string? Guard { get; init; } + + // Requirements (checked against vk.xml). Method is generated only if all match. + public List RequiresFunctions { get; init; } = new(); + public List RequiresExtensions { get; init; } = new(); + + // Method qualifiers + public bool IsConst { get; init; } = false; + public bool IsNoexcept { get; init; } = false; + public bool IsStatic { get; init; } = false; +} + +public sealed class SugarClassSpec { + public string TargetClass { get; init; } = ""; + public List Includes { get; init; } = new(); + public List Methods { get; init; } = new(); + + public SugarClassSpec() { } + public SugarClassSpec(string targetClass) { TargetClass = targetClass; } +} diff --git a/Tools/VulkanWrapperGenerator/VkSpec.cs b/Tools/VulkanWrapperGenerator/VkSpec.cs new file mode 100644 index 0000000..368de67 --- /dev/null +++ b/Tools/VulkanWrapperGenerator/VkSpec.cs @@ -0,0 +1,342 @@ +using System; +using System.Collections.Generic; +using System.Linq; +using System.Text.RegularExpressions; +using System.Xml.Linq; + +public sealed class VkSpec { + // ---------- Data model ---------- + + public sealed record Handle(string Name, bool Dispatchable); + public sealed record Param(string Raw, string Type, string Name, bool IsPointer, bool IsConst, string? LenAttr, bool IsHandle); + public sealed record Command(string Name, string ReturnType, List Params); + + public sealed record StructMember(string Type, string Name, bool IsPointer, string? LenAttr, bool IsHandle); + public sealed record StructDef(string Name, List Members); + + // Owned pair (created/destroyed or allocated/freed) + public sealed record Pair( + string HandleType, + string CreateOrAllocate, string DestroyOrFree, + string ParentType, + Param OutHandleParam, + List ExtraOwnersForDestroy + ); + + // Simple producer (Get*/Enumerate* with single parent handle and a single out handle param) + public sealed record SimpleProducer( + string HandleType, + string CommandName, + string ParentType, + Param OutHandleParam, + bool IsMany + ); + + public sealed record OwnerUse(string Type, bool IsMany, string SourceStruct, string SourceMember); + + public Dictionary Handles { get; private init; } = new(); + public Dictionary Commands { get; private init; } = new(); + public Dictionary Structs { get; private init; } = new(); + public HashSet Extensions { get; private init; } = new(); + public List Pairs { get; private init; } = new(); + public List Producers { get; private init; } = new(); + + // ---------- Load & parse ---------- + + public static VkSpec Load(string xmlPath) { + var doc = XDocument.Load(xmlPath); + + var handles = ParseHandles(doc); + var structs = ParseStructs(doc, handles); + var (real, aliasElems) = ParseCommandsPre(doc); + var cmds = BuildCommands(real, handles); + ApplyAliases(cmds, aliasElems, handles); + var exts = ParseExtensions(doc); + var pairs = BuildPairs(cmds, handles); + var producers = BuildSimpleProducers(cmds, handles); + + return new VkSpec { + Handles = handles, + Structs = structs, + Commands = cmds, + Extensions = exts, + Pairs = pairs, + Producers = producers + }; + } + + static Dictionary ParseHandles(XDocument doc) { + var dict = new Dictionary(); + var types = doc.Root!.Element("types")!.Elements("type") + .Where(t => (string?)t.Attribute("category") == "handle"); + + foreach (var t in types) { + var name = t.Element("name")?.Value; + if (string.IsNullOrEmpty(name)) { + var txt = (t.Value ?? ""); + var m = Regex.Match(txt, @"VK_DEFINE_(?:NON_)?DISPATCHABLE_HANDLE\((Vk\w+)\)"); + if (m.Success) name = m.Groups[1].Value; + } + if (string.IsNullOrEmpty(name)) continue; + + var body = t.Value ?? string.Empty; + bool dispatchable = + body.Contains("VK_DEFINE_HANDLE(") || + body.Contains("VK_DEFINE_DISPATCHABLE_HANDLE("); + dict[name!] = new Handle(name!, dispatchable); + } + return dict; + } + + static Dictionary ParseStructs(XDocument doc, Dictionary handles) { + var dict = new Dictionary(); + var xs = doc.Root!.Element("types")!.Elements("type") + .Where(t => (string?)t.Attribute("category") == "struct"); + + foreach (var t in xs) { + var name = t.Element("name")?.Value ?? (string?)t.Attribute("name"); + if (string.IsNullOrEmpty(name)) continue; + + var members = new List(); + foreach (var m in t.Elements("member")) { + var type = m.Element("type")?.Value ?? "void"; + var mname = m.Element("name")?.Value ?? ""; + var raw = Regex.Replace(m.Value ?? "", @"\s+", " "); + bool isPtr = raw.Contains("*"); + string? len = (string?)m.Attribute("len"); + bool isHandle = handles.ContainsKey(type); + members.Add(new StructMember(type, mname, isPtr, len, isHandle)); + } + dict[name!] = new StructDef(name!, members); + } + + return dict; + } + + static (List real, List aliases) ParseCommandsPre(XDocument doc) { + var all = doc.Root!.Element("commands")!.Elements("command").ToList(); + var real = new List(); + var aliases = new List(); + foreach (var c in all) { + if (c.Attribute("alias") != null) aliases.Add(c); + else if (c.Element("proto") != null) real.Add(c); + } + return (real, aliases); + } + + static string Cleanup(string s) => Regex.Replace(s.Trim(), @"\s+", " "); + + static Dictionary BuildCommands(List realCmds, Dictionary handles) { + bool IsPointer(string raw) => raw.Contains("*"); + bool IsConst(string raw) => raw.Contains("const "); + + var dict = new Dictionary(); + + foreach (var c in realCmds) { + var proto = c.Element("proto")!; + var name = proto.Element("name")!.Value; + + string retType = string.Concat(proto.Nodes().Select(n => { + if (n is XElement e) return e.Name == "name" ? "" : e.Value; + return ((XText)n).Value; + })); + retType = Cleanup(retType.Replace(name, "")); + + var prms = new List(); + foreach (var p in c.Elements("param")) { + var type = p.Element("type")?.Value ?? "void"; + var pname = p.Element("name")?.Value ?? ""; + var raw = Cleanup(p.Value); + var len = (string?)p.Attribute("len"); + bool isPtr = IsPointer(raw); + bool isConst = IsConst(raw); + bool isHandle = handles.ContainsKey(type); + prms.Add(new Param(raw, type, pname, isPtr, isConst, len, isHandle)); + } + + dict[name] = new Command(name, retType, prms); + } + + return dict; + } + + static void ApplyAliases(Dictionary dict, List aliases, Dictionary handles) { + foreach (var c in aliases) { + var name = (string?)c.Attribute("name"); + var aliasOf = (string?)c.Attribute("alias"); + if (string.IsNullOrWhiteSpace(name) || string.IsNullOrWhiteSpace(aliasOf)) continue; + + if (dict.TryGetValue(aliasOf!, out var baseCmd)) { + var clone = new Command(name!, baseCmd.ReturnType, + baseCmd.Params.Select(p => new Param(p.Raw, p.Type, p.Name, p.IsPointer, p.IsConst, p.LenAttr, p.IsHandle)).ToList()); + dict[name!] = clone; + } + } + } + + static HashSet ParseExtensions(XDocument doc) { + var set = new HashSet(StringComparer.Ordinal); + var exts = doc.Root!.Element("extensions")!.Elements("extension"); + foreach (var e in exts) { + var name = (string?)e.Attribute("name"); + if (!string.IsNullOrWhiteSpace(name)) set.Add(name!); + } + return set; + } + + static List BuildPairs(Dictionary cmds, Dictionary handles) { + var list = new List(); + + // vkCreateX / vkDestroyX + foreach (var create in cmds.Values.Where(x => x.Name.StartsWith("vkCreate"))) { + var obj = create.Name.Substring("vkCreate".Length); + var destroyName = "vkDestroy" + obj; + if (!cmds.TryGetValue(destroyName, out var destroy)) continue; + + var parent = create.Params.First().Type; + var outParam = create.Params.LastOrDefault(p => p.IsPointer && handles.ContainsKey(p.Type)); + if (outParam is null) continue; + + var handleType = outParam.Type; + var extraOwners = destroy.Params + .Where(p => p.IsHandle && p.Type != handleType) + .Select(p => p.Type).Distinct().ToList(); + + list.Add(new Pair(handleType, create.Name, destroyName, parent, outParam, extraOwners)); + } + + // vkAllocateX / vkFreeX + foreach (var alloc in cmds.Values.Where(x => x.Name.StartsWith("vkAllocate"))) { + var obj = alloc.Name.Substring("vkAllocate".Length); + var freeName = "vkFree" + obj; + if (!cmds.TryGetValue(freeName, out var free)) continue; + + var parent = alloc.Params.First().Type; + var outParam = alloc.Params.LastOrDefault(p => p.IsPointer && handles.ContainsKey(p.Type)); + if (outParam is null) continue; + + var handleType = outParam.Type; + var extraOwners = free.Params + .Where(p => p.IsHandle && p.Type != handleType) + .Select(p => p.Type).Distinct().ToList(); + + list.Add(new Pair(handleType, alloc.Name, freeName, parent, outParam, extraOwners)); + } + + return list + .GroupBy(p => p.HandleType) + .Select(g => g.First()) + .ToList(); + } + + // Build "simple producers": vkGet*/vkEnumerate* that output a handle with exactly one parent handle argument + static List BuildSimpleProducers(Dictionary cmds, Dictionary handles) { + var list = new List(); + + foreach (var c in cmds.Values) { + if (!(c.Name.StartsWith("vkGet") || c.Name.StartsWith("vkEnumerate"))) + continue; + + // pick first handle output param (pointer to handle) + int outIdx = c.Params.FindIndex(p => p.IsPointer && handles.ContainsKey(p.Type)); + if (outIdx < 0) continue; + var outParam = c.Params[outIdx]; + + // consider handle params BEFORE out param as potential parents + var parents = c.Params.Take(outIdx).Where(p => p.IsHandle).ToList(); + if (parents.Count != 1) continue; // only "simple" case supported: a single parent + + var parent = parents[0].Type; + bool isMany = outParam.LenAttr != null; + + list.Add(new SimpleProducer(outParam.Type, c.Name, parent, outParam, isMany)); + } + + return list; + } + + public enum LoadLevel { Global, Instance, Device } + public LoadLevel ClassifyLoadLevel(Command c) { + if (c.Params.Count == 0) return LoadLevel.Global; + var first = c.Params[0].Type; + if (first == "VkInstance" || first == "VkPhysicalDevice") return LoadLevel.Instance; + if (first == "VkDevice" || first == "VkQueue" || first == "VkCommandBuffer") return LoadLevel.Device; + return LoadLevel.Global; + } + + // ---------- Deterministic keep filter ---------- + public static class LifetimeFilter { + // Explicit "do not keep" cases (semantically not required after creation) + private static readonly HashSet<(string st, string mem, string ht)> Exact = new() + { + ("VkPipelineShaderStageCreateInfo", "module", "VkShaderModule"), + ("VkFramebufferCreateInfo", "renderPass", "VkRenderPass"), + ("VkGraphicsPipelineCreateInfo", "renderPass", "VkRenderPass"), + ("VkSwapchainCreateInfoKHR", "oldSwapchain", "VkSwapchainKHR"), + }; + + private static readonly string[] NameHints = { "old", "scratch", "staging" }; + + public static bool ShouldKeep(string structName, string memberName, string handleType) { + if (Exact.Contains((structName, memberName, handleType))) + return false; + + var mn = memberName.ToLowerInvariant(); + if (NameHints.Any(h => mn.Contains(h))) + return false; + + // Immutable samplers ARE kept + return true; + } + } + + // Traverse CreateInfo/AllocateInfo trees and collect keep-alive owners + public List AnalyzeCreateInfoOwners(Command create, string handleType) { + var owners = new Dictionary<(string Type, string Struct, string Member), OwnerUse>(); + + void Mark(string t, bool many, string srcStruct, string srcMember) { + if (t == handleType) return; // do not keep self + if (!LifetimeFilter.ShouldKeep(srcStruct, srcMember, t)) return; + + var key = (t, srcStruct, srcMember); + if (owners.TryGetValue(key, out var ex)) + owners[key] = new OwnerUse(t, ex.IsMany || many, srcStruct, srcMember); + else + owners[key] = new OwnerUse(t, many, srcStruct, srcMember); + } + + var stack = new Stack<(string type, bool many)>(); + + foreach (var p in create.Params.Where(p => p.Name.Contains("CreateInfo") || p.Name.Contains("AllocateInfo"))) { + if (Structs.ContainsKey(p.Type)) + stack.Push((p.Type, p.IsPointer && p.LenAttr != null)); + } + + var visited = new HashSet(); + while (stack.Count > 0) { + var (stype, parentMany) = stack.Pop(); + if (!visited.Add(stype)) { /* ok */ } + + if (!Structs.TryGetValue(stype, out var s)) continue; + foreach (var m in s.Members) { + bool manyHere = parentMany || (m.IsPointer && m.LenAttr != null); + + if (Handles.ContainsKey(m.Type)) { + Mark(m.Type, manyHere, stype, m.Name); + } else if (Structs.ContainsKey(m.Type)) { + stack.Push((m.Type, manyHere)); + } + } + } + + return owners.Values + .GroupBy(o => o.Type) + .Select(g => new OwnerUse( + g.Key, + g.Any(x => x.IsMany), + g.First().SourceStruct, + g.First().SourceMember + )) + .ToList(); + } +} diff --git a/Tools/VulkanWrapperGenerator/VulkanWrapperGenerator.csproj b/Tools/VulkanWrapperGenerator/VulkanWrapperGenerator.csproj new file mode 100644 index 0000000..fd4bd08 --- /dev/null +++ b/Tools/VulkanWrapperGenerator/VulkanWrapperGenerator.csproj @@ -0,0 +1,10 @@ + + + + Exe + net9.0 + enable + enable + + + diff --git a/Tools/VulkanWrapperGenerator/VulkanWrapperGenerator.sln b/Tools/VulkanWrapperGenerator/VulkanWrapperGenerator.sln new file mode 100644 index 0000000..0836ca2 --- /dev/null +++ b/Tools/VulkanWrapperGenerator/VulkanWrapperGenerator.sln @@ -0,0 +1,25 @@ + +Microsoft Visual Studio Solution File, Format Version 12.00 +# Visual Studio Version 17 +VisualStudioVersion = 17.14.36221.1 d17.14 +MinimumVisualStudioVersion = 10.0.40219.1 +Project("{FAE04EC0-301F-11D3-BF4B-00C04F79EFBC}") = "VulkanWrapperGenerator", "VulkanWrapperGenerator.csproj", "{8552C09F-5C09-CEAF-7342-DF426CEC77E3}" +EndProject +Global + GlobalSection(SolutionConfigurationPlatforms) = preSolution + Debug|Any CPU = Debug|Any CPU + Release|Any CPU = Release|Any CPU + EndGlobalSection + GlobalSection(ProjectConfigurationPlatforms) = postSolution + {8552C09F-5C09-CEAF-7342-DF426CEC77E3}.Debug|Any CPU.ActiveCfg = Debug|Any CPU + {8552C09F-5C09-CEAF-7342-DF426CEC77E3}.Debug|Any CPU.Build.0 = Debug|Any CPU + {8552C09F-5C09-CEAF-7342-DF426CEC77E3}.Release|Any CPU.ActiveCfg = Release|Any CPU + {8552C09F-5C09-CEAF-7342-DF426CEC77E3}.Release|Any CPU.Build.0 = Release|Any CPU + EndGlobalSection + GlobalSection(SolutionProperties) = preSolution + HideSolutionNode = FALSE + EndGlobalSection + GlobalSection(ExtensibilityGlobals) = postSolution + SolutionGuid = {B1300085-B1F5-4AB7-A22D-B9918A2528B1} + EndGlobalSection +EndGlobal