Improve Harmony Hook

1. Hook.HookMethod can patch overloaded methods based on parameter types
2. Hook.HookMethod can add more patches to the method
3. Find all hook methods by the address of the function (origin method), not by the method path (string).
4. Fixed unable to add postfix patch to method (__params changed to __args)
5. Fix patching method will cause duplicate patches to be added
This commit is contained in:
zhurengong
2022-01-25 21:54:56 +08:00
parent 740e51ce48
commit 6daf410e50
2 changed files with 184 additions and 88 deletions
+79 -30
View File
@@ -1,36 +1,85 @@
Hook.HookMethod("Barotrauma.Item", "TryInteract", function (instance, p) Hook.HookMethod(
if Hook.Call("itemInteract", instance, p.picker, p.ignoreRequiredItems, p.forceSelectKey, p.forceActionKey) == true then "Barotrauma.Item", "TryInteract",
return false {
end "Barotrauma.Character",
end, Hook.HookMethodType.Before) "System.Boolean",
"System.Boolean",
"System.Boolean"
},
function (instance, p)
if Hook.Call("itemInteract", instance, p.picker, p.ignoreRequiredItems, p.forceSelectKey, p.forceActionKey) == true then
return false
end
end,
Hook.HookMethodType.Before
)
Hook.HookMethod("Barotrauma.Item", "ApplyTreatment", function (instance, p) Hook.HookMethod(
if Hook.Call("itemApplyTreatment", instance, p.user, p.character, p.targetLimb) then "Barotrauma.Item", "ApplyTreatment",
return false {
end "Barotrauma.Character",
end, Hook.HookMethodType.Before) "Barotrauma.Character",
"Barotrauma.Limb"
},
function (instance, p)
if Hook.Call("itemApplyTreatment", instance, p.user, p.character, p.targetLimb) then
return false
end
end,
Hook.HookMethodType.Before
)
Hook.HookMethod("Barotrauma.Item", "Combine", function (instance, p) Hook.HookMethod(
if Hook.Call("itemCombine", instance, p.item, p.user) == true then "Barotrauma.Item", "Combine",
return false {
end "Barotrauma.Item",
end, Hook.HookMethodType.Before) "Barotrauma.Character"
},
function (instance, p)
if Hook.Call("itemCombine", instance, p.item, p.user) == true then
return false
end
end,
Hook.HookMethodType.Before
)
Hook.HookMethod("Barotrauma.Item", "Drop", function (instance, p) Hook.HookMethod(
if Hook.Call("itemDrop", instance, p.dropper) == true then "Barotrauma.Item", "Drop",
return false {
end "Barotrauma.Character",
end, Hook.HookMethodType.Before) "System.Boolean"
},
function (instance, p)
if Hook.Call("itemDrop", instance, p.dropper) == true then
return false
end
end,
Hook.HookMethodType.Before
)
Hook.HookMethod("Barotrauma.Item", "Equip", function (instance, p) Hook.HookMethod(
if Hook.Call("itemEquip", instance, p.character) == true then "Barotrauma.Item", "Equip",
return false {
end "Barotrauma.Character"
end, Hook.HookMethodType.Before) },
function (instance, p)
if Hook.Call("itemEquip", instance, p.character) == true then
return false
end
end,
Hook.HookMethodType.Before
)
Hook.HookMethod("Barotrauma.Item", "Unequip", function (instance, p) Hook.HookMethod(
if Hook.Call("itemUnequip", instance, p.character) == true then "Barotrauma.Item", "Unequip",
return false {
end "Barotrauma.Character"
end, Hook.HookMethodType.Before) },
function (instance, p)
if Hook.Call("itemUnequip", instance, p.character) == true then
return false
end
end,
Hook.HookMethodType.Before
)
@@ -799,7 +799,8 @@ namespace Barotrauma
public LuaHook(LuaSetup e) public LuaHook(LuaSetup e)
{ {
env = e; env = e;
_hookMethods = new Dictionary<string, object>(); _hookPrefixMethods = new Dictionary<long, HashSet<object>>();
_hookPostfixMethods = new Dictionary<long, HashSet<object>>();
} }
public class HookFunction public class HookFunction
@@ -818,7 +819,8 @@ namespace Barotrauma
private Dictionary<string, Dictionary<string, HookFunction>> hookFunctions = new Dictionary<string, Dictionary<string, HookFunction>>(); private Dictionary<string, Dictionary<string, HookFunction>> hookFunctions = new Dictionary<string, Dictionary<string, HookFunction>>();
private static Dictionary<string, object> _hookMethods; private static Dictionary<long, HashSet<object>> _hookPrefixMethods;
private static Dictionary<long, HashSet<object>> _hookPostfixMethods;
private Queue<Tuple<object, object[]>> queuedFunctionCalls = new Queue<Tuple<object, object[]>>(); private Queue<Tuple<object, object[]>> queuedFunctionCalls = new Queue<Tuple<object, object[]>>();
@@ -829,29 +831,36 @@ namespace Barotrauma
static void _hookLuaPatch(MethodBase __originalMethod, object[] __args, object __instance, out LuaResult result, HookMethodType hookMethodType) static void _hookLuaPatch(MethodBase __originalMethod, object[] __args, object __instance, out LuaResult result, HookMethodType hookMethodType)
{ {
// Although it works correctly, the performance is low
result = new LuaResult(null); result = new LuaResult(null);
try try
{ {
var classType = __originalMethod.DeclaringType;
var methodPath = $"{hookMethodType}:{classType.Namespace}.{classType.Name}.{__originalMethod.Name}";
var @params = __originalMethod.GetParameters(); var @params = __originalMethod.GetParameters();
var ptable = new Dictionary<string, object>(); var ptable = new Dictionary<string, object>();
for (int i = 0; i < @params.Length; i++) for (int i = 0; i < @params.Length; i++)
{ {
ptable.Add(@params[i].Name, __args[i]); ptable.Add(@params[i].Name, __args[i]);
} }
if (_hookMethods.TryGetValue(methodPath, out object hookMethod)) var funcAddr = ((long)__originalMethod.MethodHandle.GetFunctionPointer());
{ HashSet<object> methodSet = null;
result = new LuaResult(luaSetup.hook.env.lua.Call(hookMethod, __instance, ptable)); switch (hookMethodType)
}
else
{ {
luaSetup.PrintError($"No hook method found in _hookMethods[{methodPath}]"); case HookMethodType.Before:
methodSet = _hookPrefixMethods[funcAddr];
break;
case HookMethodType.After:
methodSet = _hookPostfixMethods[funcAddr];
break;
default:
break;
}
if (methodSet != null)
{
foreach (var hookMethod in methodSet)
{
result = new LuaResult(luaSetup.lua.Call(hookMethod, __instance, ptable));
}
} }
} }
@@ -881,14 +890,14 @@ namespace Barotrauma
return true; return true;
} }
static void HookLuaPatchPostfix(MethodBase __originalMethod, object[] __params, object __instance) static void HookLuaPatchPostfix(MethodBase __originalMethod, object[] __args, object __instance)
{ {
_hookLuaPatch(__originalMethod, __params, __instance, out LuaResult result, HookMethodType.After); _hookLuaPatch(__originalMethod, __args, __instance, out LuaResult result, HookMethodType.After);
} }
static void HookLuaPatchRetPostfix(MethodBase __originalMethod, object[] __params, ref object __result, object __instance) static void HookLuaPatchRetPostfix(MethodBase __originalMethod, object[] __args, ref object __result, object __instance)
{ {
_hookLuaPatch(__originalMethod, __params, __instance, out LuaResult result, HookMethodType.After); _hookLuaPatch(__originalMethod, __args, __instance, out LuaResult result, HookMethodType.After);
if (!result.IsNull()) if (!result.IsNull())
__result = result.Object(); __result = result.Object();
@@ -898,61 +907,99 @@ namespace Barotrauma
private static MethodInfo _miHookLuaPatchRetPrefix = typeof(LuaHook).GetMethod("HookLuaPatchRetPrefix", BindingFlags.NonPublic | BindingFlags.Static); private static MethodInfo _miHookLuaPatchRetPrefix = typeof(LuaHook).GetMethod("HookLuaPatchRetPrefix", BindingFlags.NonPublic | BindingFlags.Static);
private static MethodInfo _miHookLuaPatchPostfix = typeof(LuaHook).GetMethod("HookLuaPatchPostfix", BindingFlags.NonPublic | BindingFlags.Static); private static MethodInfo _miHookLuaPatchPostfix = typeof(LuaHook).GetMethod("HookLuaPatchPostfix", BindingFlags.NonPublic | BindingFlags.Static);
private static MethodInfo _miHookLuaPatchRetPostfix = typeof(LuaHook).GetMethod("HookLuaPatchRetPostfix", BindingFlags.NonPublic | BindingFlags.Static); private static MethodInfo _miHookLuaPatchRetPostfix = typeof(LuaHook).GetMethod("HookLuaPatchRetPostfix", BindingFlags.NonPublic | BindingFlags.Static);
public void HookMethod(string className, string methodName, object hookMethod, HookMethodType hookMethodType = HookMethodType.Before) public void HookMethod(string className, string methodName, string[] parameterNames, object hookMethod, HookMethodType hookMethodType = HookMethodType.Before)
{ {
if (hookMethod == null)
{
env.PrintError("hookMethod cannot be null");
return;
}
var classType = Type.GetType(className); var classType = Type.GetType(className);
var methodInfos = classType.GetMethods(); MethodInfo methodInfo = null;
HarmonyMethod harmonyMethod = new HarmonyMethod();
HarmonyMethod harmonyMethodRet = new HarmonyMethod(); if (parameterNames.Length > 0)
{
Type[] parameterTypes = parameterNames.Select(x => AccessTools.TypeByName(x)).ToArray();
methodInfo = classType.GetMethod(methodName, parameterTypes);
}
else
{
methodInfo = classType.GetMethod(methodName);
}
if (methodInfo == null)
{
env.PrintError($"can't find method({className}.{methodName}) with these parameters' types({string.Join(", ", parameterNames)})");
return;
}
var funcAddr = ((long)methodInfo.MethodHandle.GetFunctionPointer());
var patches = Harmony.GetPatchInfo(methodInfo);
if (hookMethodType == HookMethodType.Before) if (hookMethodType == HookMethodType.Before)
{ {
harmonyMethod = new HarmonyMethod(_miHookLuaPatchPrefix); if (methodInfo.ReturnType == typeof(void))
harmonyMethodRet = new HarmonyMethod(_miHookLuaPatchRetPrefix); {
if (patches == null || patches.Prefixes == null || patches.Prefixes.Find(patch => patch.PatchMethod == _miHookLuaPatchPrefix) == null)
{
env.harmony.Patch(methodInfo, prefix: new HarmonyMethod(_miHookLuaPatchPrefix));
}
}
else
{
if (patches == null || patches.Prefixes == null || patches.Prefixes.Find(patch => patch.PatchMethod == _miHookLuaPatchRetPrefix) == null)
{
env.harmony.Patch(methodInfo, prefix: new HarmonyMethod(_miHookLuaPatchRetPrefix));
}
}
if (_hookPrefixMethods.TryGetValue(funcAddr, out HashSet<object> methodSet))
{
methodSet.Add(hookMethod);
}
else
{
_hookPrefixMethods.Add(funcAddr, new HashSet<object>() { hookMethod });
}
} }
else if (hookMethodType == HookMethodType.After) else if (hookMethodType == HookMethodType.After)
{ {
harmonyMethod = new HarmonyMethod(_miHookLuaPatchPostfix); if (methodInfo.ReturnType == typeof(void))
harmonyMethodRet = new HarmonyMethod(_miHookLuaPatchRetPostfix);
}
foreach (var methodInfo in methodInfos)
{
if (methodInfo.Name == methodName)
{ {
if (hookMethodType == HookMethodType.Before) if (patches == null || patches.Postfixes == null || patches.Postfixes.Find(patch => patch.PatchMethod == _miHookLuaPatchPostfix) == null)
{
if (methodInfo.ReturnType == typeof(void))
env.harmony.Patch(methodInfo, prefix: harmonyMethod);
else
env.harmony.Patch(methodInfo, prefix: harmonyMethodRet);
}
else if (hookMethodType == HookMethodType.After)
{
if (methodInfo.ReturnType == typeof(void))
env.harmony.Patch(methodInfo, postfix: harmonyMethod);
else
env.harmony.Patch(methodInfo, postfix: harmonyMethodRet);
}
// build an unique method path by patch type, class, method self
var methodPath = $"{hookMethodType}:{classType.Namespace}.{classType.Name}.{methodInfo.Name}";
if (hookMethod != null)
{ {
if (!_hookMethods.TryAdd(methodPath, hookMethod)) env.harmony.Patch(methodInfo, postfix: new HarmonyMethod(_miHookLuaPatchPostfix));
env.PrintError($"Failed to add key-value in {nameof(_hookMethods)}\n[{methodPath}, {hookMethod.ToString()}]");
#if DEBUG
else
env.PrintMessage($"Sucessfully added key-value in {nameof(_hookMethods)}\n[{methodPath}, {hookMethod.ToString()}]");
#endif
} }
break;
} }
else
{
if (patches == null || patches.Postfixes == null || patches.Postfixes.Find(patch => patch.PatchMethod == _miHookLuaPatchRetPostfix) == null)
{
env.harmony.Patch(methodInfo, postfix: new HarmonyMethod(_miHookLuaPatchRetPostfix));
}
}
if (_hookPostfixMethods.TryGetValue(funcAddr, out HashSet<object> methodSet))
{
methodSet.Add(hookMethod);
}
else
{
_hookPostfixMethods.Add(funcAddr, new HashSet<object>() { hookMethod });
}
} }
} }
public void HookMethod(string className, string methodName, object hookMethod, HookMethodType hookMethodType = HookMethodType.Before)
{
HookMethod(className, methodName, new string[] {}, hookMethod, hookMethodType);
}
public void EnqueueFunction(object function, params object[] args) public void EnqueueFunction(object function, params object[] args)
{ {
queuedFunctionCalls.Enqueue(new Tuple<object, object[]>(function, args)); queuedFunctionCalls.Enqueue(new Tuple<object, object[]>(function, args));