#[cfg(test)]
#[path = "tests/server_hooks.rs"]
mod tests;
use crate::MetamodApi;
use crate::hook::{
Handler, HookAction, HookCall, HookError, HookId, HookTarget, HookTiming, VirtualFunction,
};
use crate::sys::plugin::{self as raw, HookStatus};
use source_sdk_2013::interfaces::ServerGameDll;
use source_sdk_2013::net::incoming::{
HookTargetError, IncomingHandler, IncomingKind, Verdict, hook_target, route_incoming,
};
use source_sdk_2013::raw::interfaces::server_game_dll::{
GAME_FRAME_SLOT, GameFrameFn as GameFrame,
};
use source_sdk_2013::raw::net::incoming::ProcessMessageFn as ProcessMessage;
use source_sdk_2013::{Server, ServerBinding};
use std::cell::Cell;
use std::ffi::{CStr, c_char, c_int, c_void};
use std::panic::{AssertUnwindSafe, catch_unwind};
use std::ptr::{self, NonNull};
pub type GameFrameFn = fn(server: Server<'_>, simulating: bool);
const GAME_FRAME: VirtualFunction<GameFrame> = VirtualFunction::new(GAME_FRAME_SLOT);
static GAME_FRAMES: Route<GameFrameFn> = Route::new();
static GAME_FRAMES_POST: Route<GameFrameFn> = Route::new();
static LEVELS: LevelRoute = LevelRoute(Cell::new(None));
static NET_MESSAGE_KINDS: [NetMessageKind; IncomingKind::ALL.len()] = {
let mut kinds = [const { NetMessageKind(0) }; IncomingKind::ALL.len()];
let mut kind = 0;
while kind < kinds.len() {
kinds[kind] = NetMessageKind(kind as c_int);
kind += 1;
}
kinds
};
static NET_MESSAGES: Route<&'static dyn IncomingHandler, { IncomingKind::ALL.len() }> =
Route::new();
#[derive(Debug, Clone, Copy, Default)]
pub struct LevelEvents {
pub init: Option<fn(server: Server<'_>, map: &CStr)>,
pub shutdown: Option<fn(server: Server<'_>)>,
}
struct LevelRoute(Cell<Option<(ServerBinding, LevelEvents)>>);
impl LevelRoute {
unsafe fn from_context(context: *mut c_void) -> Option<(ServerBinding, LevelEvents)> {
unsafe { &*context.cast::<Self>() }.0.get()
}
fn context(&'static self) -> *mut c_void {
ptr::from_ref(self).cast_mut().cast()
}
}
unsafe impl Sync for LevelRoute {}
#[derive(Debug, thiserror::Error)]
pub enum NetMessageHookError {
#[error(transparent)]
Target(#[from] HookTargetError),
#[error(transparent)]
Hook(#[from] HookError),
}
struct NetMessageKind(c_int);
impl Handler<ProcessMessage> for NetMessageKind {
fn call(&self, call: &HookCall<'_, ProcessMessage>) -> HookAction<bool> {
if call.superseded() == Some(true) {
return HookAction::Ignore;
}
let (message,) = call.args();
let (Some(routed), Some(handler), Some(message)) = (
NET_MESSAGES.get(),
NonNull::new(call.this()),
NonNull::new(message),
) else {
return HookAction::Ignore;
};
match unsafe { route_incoming(&routed.binding, routed.target, self.0, handler, message) } {
Verdict::Block => HookAction::Supersede(true),
Verdict::Continue => HookAction::Ignore,
}
}
}
struct Route<T, const HOOKS: usize = 1>(Cell<Option<Routed<T, HOOKS>>>);
impl<T: Copy, const HOOKS: usize> Route<T, HOOKS> {
const fn new() -> Self {
Self(Cell::new(None))
}
fn get(&self) -> Option<Routed<T, HOOKS>> {
self.0.get()
}
fn installed(&self, api: MetamodApi<'_>) -> bool {
self.get().is_some_and(|routed| {
routed
.hooks
.into_iter()
.flatten()
.any(|hook| api.has_hook(hook))
})
}
fn set(&self, hooks: [Option<HookId>; HOOKS], binding: ServerBinding, target: T) {
self.0.set(Some(Routed {
hooks,
binding,
target,
}));
}
}
impl Handler<GameFrame> for Route<GameFrameFn> {
fn call(&self, call: &HookCall<'_, GameFrame>) -> HookAction<()> {
let (simulating,) = call.args();
if let Some(routed) = self.get() {
with_server(routed.binding, |server| (routed.target)(server, simulating));
}
HookAction::Ignore
}
}
unsafe impl<T, const HOOKS: usize> Sync for Route<T, HOOKS> {}
#[derive(Clone, Copy)]
struct Routed<T, const HOOKS: usize> {
hooks: [Option<HookId>; HOOKS],
binding: ServerBinding,
target: T,
}
impl MetamodApi<'_> {
pub fn hook_game_frame(
self,
game_dll: ServerGameDll<'_>,
binding: ServerBinding,
callback: GameFrameFn,
) -> Result<(), HookError> {
hook_game_frames(
self,
&GAME_FRAMES,
HookTiming::Pre,
game_dll,
binding,
callback,
)
}
pub fn hook_game_frame_post(
self,
game_dll: ServerGameDll<'_>,
binding: ServerBinding,
callback: GameFrameFn,
) -> Result<(), HookError> {
hook_game_frames(
self,
&GAME_FRAMES_POST,
HookTiming::Post,
game_dll,
binding,
callback,
)
}
pub fn hook_net_messages(
self,
server: Server<'_>,
binding: ServerBinding,
handler: &'static dyn IncomingHandler,
) -> Result<(), NetMessageHookError> {
if NET_MESSAGES.installed(self) {
return Err(HookError::AlreadyInstalled.into());
}
let target = hook_target(server)?;
let mut hooks = [None; IncomingKind::ALL.len()];
for (kind, &slot) in target.slots.iter().enumerate() {
let hooked = usize::try_from(slot)
.map_err(|_| HookError::InvalidArgument)
.and_then(|slot| {
unsafe {
self.add_hook(
VirtualFunction::<ProcessMessage>::new(slot),
HookTarget::class_of(target.handler),
HookTiming::Pre,
&NET_MESSAGE_KINDS[kind],
)
}
});
match hooked {
Ok(hook) => hooks[kind] = Some(hook),
Err(error) => {
for hook in hooks.into_iter().flatten() {
self.remove_hook(hook);
}
return Err(error.into());
}
}
}
NET_MESSAGES.set(hooks, binding, handler);
Ok(())
}
pub fn listen_level_events(
self,
binding: ServerBinding,
events: LevelEvents,
) -> Result<(), HookError> {
LEVELS.0.set(Some((binding, events)));
let status = unsafe {
raw::cpp_metamod_listen_levels(
self.version().plugin_api_version(),
events.init.map(|_| level_init as raw::LevelInitCallback),
events
.shutdown
.map(|_| level_shutdown as raw::LevelShutdownCallback),
LEVELS.context(),
)
};
match status {
HookStatus::INSTALLED => Ok(()),
HookStatus::NOT_BOUND => Err(HookError::NotBound),
HookStatus::ALREADY_INSTALLED => Err(HookError::AlreadyInstalled),
HookStatus::INVALID_ARGUMENT => Err(HookError::InvalidArgument),
_ => Err(HookError::Unsupported),
}
}
}
fn hook_game_frames(
api: MetamodApi<'_>,
route: &'static Route<GameFrameFn>,
timing: HookTiming,
game_dll: ServerGameDll<'_>,
binding: ServerBinding,
callback: GameFrameFn,
) -> Result<(), HookError> {
if route.installed(api) {
return Err(HookError::AlreadyInstalled);
}
let game_dll = NonNull::new(game_dll.as_ptr()).ok_or(HookError::InvalidArgument)?;
let hook = unsafe { api.add_hook(GAME_FRAME, HookTarget::instance(game_dll), timing, route) }?;
route.set([Some(hook)], binding, callback);
Ok(())
}
unsafe extern "C" fn level_init(context: *mut c_void, map: *const c_char) {
let Some((binding, events)) = (unsafe { LevelRoute::from_context(context) }) else {
return;
};
let (Some(init), false) = (events.init, map.is_null()) else {
return;
};
let map = unsafe { CStr::from_ptr(map) };
with_server(binding, |server| init(server, map));
}
unsafe extern "C" fn level_shutdown(context: *mut c_void) {
if let Some((
binding,
LevelEvents {
shutdown: Some(shutdown),
..
},
)) = unsafe { LevelRoute::from_context(context) }
{
with_server(binding, shutdown);
}
}
fn with_server(binding: ServerBinding, f: impl FnOnce(Server<'_>)) {
let scope = ();
let server = unsafe { binding.server(&scope) };
catch_unwind(AssertUnwindSafe(|| f(server))).ok();
}