#[cfg(test)]
#[path = "tests/damage_hooks.rs"]
mod tests;
use crate::MetamodApi;
use crate::hook::{
Handler, HookAction, HookCall, HookError, HookId, HookTarget, HookTiming, VirtualFunction,
};
use source_sdk_2013::entities::Entity;
use source_sdk_2013::raw::tf2::damage::{
ON_TAKE_DAMAGE_ALIVE_SLOT, ON_TAKE_DAMAGE_SLOT, TakeDamageFn as TakeDamage,
};
use source_sdk_2013::raw::util::vtable::vtable_pointer;
use source_sdk_2013::tf2::damage::DamageEvent;
use source_sdk_2013::{Game, Server, ServerBinding, sys};
use std::cell::Cell;
use std::ffi::{c_int, c_void};
use std::ptr::NonNull;
pub type DamageFn = for<'s> fn(Server<'s>, DamageStage, &mut DamageEvent<'s>) -> DamageAction;
static ROUTES: [DamageRoute; 32] = [const { DamageRoute::new() }; 32];
#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)]
pub enum DamageAction {
#[default]
Continue,
Apply,
Block,
}
#[derive(Debug, thiserror::Error)]
pub enum DamageHookError {
#[error("damage hooks require a TF2 server and a CTFPlayer entity")]
NotTfPlayer,
#[error(transparent)]
Hook(#[from] HookError),
}
struct DamageRoute {
state: Cell<Option<RoutedDamage>>,
}
impl DamageRoute {
const fn new() -> Self {
Self {
state: Cell::new(None),
}
}
}
impl Handler<TakeDamage> for DamageRoute {
fn call(&self, call: &HookCall<'_, TakeDamage>) -> HookAction<c_int> {
if call.superseded() == Some(true) {
return HookAction::Ignore;
}
let Some(route) = self.state.get() else {
return HookAction::Ignore;
};
let (info,) = call.args();
let (Some(victim), Some(info)) = (NonNull::new(call.this()), NonNull::new(info.cast_mut()))
else {
return HookAction::Ignore;
};
let scope = ();
let server = unsafe { route.binding.server(&scope) };
unsafe {
dispatch(
server,
route.stage,
route.original,
route.callback,
victim,
info,
)
}
}
}
unsafe impl Sync for DamageRoute {}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum DamageStage {
Incoming,
Alive,
}
impl DamageStage {
const fn function(self) -> VirtualFunction<TakeDamage> {
VirtualFunction::new(match self {
Self::Incoming => ON_TAKE_DAMAGE_SLOT,
Self::Alive => ON_TAKE_DAMAGE_ALIVE_SLOT,
})
}
}
#[derive(Clone, Copy)]
struct RoutedDamage {
binding: ServerBinding,
callback: DamageFn,
original: TakeDamage,
hook: HookId,
stage: DamageStage,
vtable: usize,
}
impl MetamodApi<'_> {
pub fn hook_player_damage(
self,
player: Entity<'_>,
binding: ServerBinding,
stage: DamageStage,
callback: DamageFn,
) -> Result<HookId, DamageHookError> {
if binding.game() != Game::TeamFortress2
|| !player
.server_class()
.is_some_and(|class| class.name() == c"CTFPlayer")
{
return Err(DamageHookError::NotTfPlayer);
}
let vtable = unsafe { vtable_pointer::<c_void>(player.as_ptr()) }.addr();
if ROUTES.iter().any(|route| {
route.state.get().is_some_and(|state| {
state.vtable == vtable && state.stage == stage && self.has_hook(state.hook)
})
}) {
return Err(HookError::AlreadyInstalled.into());
}
let route = ROUTES
.iter()
.find(|route| {
route
.state
.get()
.is_none_or(|state| !self.has_hook(state.hook))
})
.ok_or(HookError::TooManyFunctions)?;
let target = HookTarget::class_of(NonNull::new(player.as_ptr()).unwrap());
let original = unsafe { self.original_function(stage.function(), target) }?;
let hook = unsafe { self.add_hook(stage.function(), target, HookTiming::Pre, route) }?;
route.state.set(Some(RoutedDamage {
binding,
callback,
original,
hook,
stage,
vtable,
}));
Ok(hook)
}
}
unsafe fn dispatch(
server: Server<'_>,
stage: DamageStage,
original: TakeDamage,
callback: DamageFn,
victim: NonNull<sys::CBaseEntity>,
info: NonNull<sys::CTakeDamageInfo>,
) -> HookAction<c_int> {
let mut event = unsafe { DamageEvent::from_raw(server, victim, info) };
match callback(server, stage, &mut event) {
DamageAction::Continue => HookAction::Ignore,
DamageAction::Block => HookAction::Supersede(0),
DamageAction::Apply => {
let result = unsafe { original(victim.as_ptr(), event.info.as_ptr()) };
HookAction::Supersede(result)
}
}
}