#include "Extension.h" #include #include #include "CallbackHandler.h" #include "Callback.h" #include "Socket.h" using namespace boost::asio::ip; Extension extension; SMEXT_LINK(&extension); void GameFrame(bool simulating) { callbackHandler.ExecuteQueuedCallbacks(); } extern const sp_nativeinfo_t smsock_natives[]; bool Extension::SDK_OnLoad(char *error, size_t err_max, bool late) { smutils->AddGameFrameHook(&GameFrame); sharesys->AddNatives(myself, smsock_natives); socketHandleType = handlesys->CreateType("Socket", this, 0, NULL, NULL, myself->GetIdentity(), NULL); //if (_debug) smutils->LogError(myself, "[Debug] Extension loaded"); socketHandler.StartProcessing(); return true; } void Extension::SDK_OnUnload() { smutils->RemoveGameFrameHook(&GameFrame); handlesys->RemoveType(socketHandleType, NULL); socketHandler.Shutdown(); } void Extension::OnHandleDestroy(HandleType_t type, void *object) { if (type == socketHandleType && object != NULL) { socketHandler.DestroySocket((SocketWrapper*) object); } } SocketWrapper* Extension::GetSocketWrapperByHandle(Handle_t handle) { HandleSecurity sec; sec.pOwner = NULL; sec.pIdentity = myself->GetIdentity(); SocketWrapper* sw; if (handlesys->ReadHandle(handle, socketHandleType, &sec, (void**)&sw) != HandleError_None) return NULL; return sw; } // native bool:SocketIsConnected(Handle:socket); cell_t SocketIsConnected(IPluginContext *pContext, const cell_t *params) { SocketWrapper* sw = extension.GetSocketWrapperByHandle(static_cast(params[1])); if (sw == NULL) return pContext->ThrowNativeError("Invalid handle: %i", params[1]); switch (sw->socketType) { case SM_SocketType_Tcp: return ((Socket*) sw->socket)->IsOpen(); case SM_SocketType_Udp: return ((Socket*) sw->socket)->IsOpen(); default: return false; } } // native Handle:SocketCreate(SocketType:protocol=SOCKET_TCP, SocketErrorCB:efunc); cell_t SocketCreate(IPluginContext *pContext, const cell_t *params) { if (params[1] != SM_SocketType_Tcp && params[1] != SM_SocketType_Udp) return pContext->ThrowNativeError("Invalid protocol specified"); if (!pContext->GetFunctionById(params[2])) return pContext->ThrowNativeError("Invalid error callback specified"); cell_t handle = -1; switch (params[1]) { case SM_SocketType_Tcp: { Socket* socket = socketHandler.CreateSocket(SM_SocketType_Tcp); SocketWrapper* sw = socketHandler.GetSocketWrapper(socket); handle = handlesys->CreateHandle(extension.socketHandleType, sw, pContext->GetIdentity(), myself->GetIdentity(), NULL); socket->smHandle = handle; socket->errorCallback = pContext->GetFunctionById(params[2]); break; } case SM_SocketType_Udp: { Socket* socket = socketHandler.CreateSocket(SM_SocketType_Udp); SocketWrapper* sw = socketHandler.GetSocketWrapper(socket); handle = handlesys->CreateHandle(extension.socketHandleType, sw, pContext->GetIdentity(), myself->GetIdentity(), NULL); socket->smHandle = handle; socket->errorCallback = pContext->GetFunctionById(params[2]); break; } } return handle; } // native SocketBind(Handle:socket, String:hostname[], port); cell_t SocketBind(IPluginContext *pContext, const cell_t *params) { SocketWrapper* sw = extension.GetSocketWrapperByHandle(static_cast(params[1])); if (sw == NULL) return pContext->ThrowNativeError("Invalid handle: %i", params[1]); if (params[3] < 0 || params[3] > 65535) return pContext->ThrowNativeError("Invalid port specified"); char *hostname = NULL; pContext->LocalToString(params[2], &hostname); switch (sw->socketType) { case SM_SocketType_Tcp: return ((Socket*) sw->socket)->Bind(hostname, params[3], false); case SM_SocketType_Udp: return ((Socket*) sw->socket)->Bind(hostname, params[3], false); default: return false; } } // native SocketConnect(Handle:socket, SocketConnectCB:cfunc, SocketReceiveCB:rfunc, SocketDisconnectCB:dfunc, String:hostname[], port); cell_t SocketConnect(IPluginContext *pContext, const cell_t *params) { SocketWrapper* sw = extension.GetSocketWrapperByHandle(static_cast(params[1])); if (sw == NULL) return pContext->ThrowNativeError("Invalid handle: %i", params[1]); //if (socket->shouldListen()) return pContext->ThrowNativeError("You can't connect a listening socket"); if (!pContext->GetFunctionById(params[2])) return pContext->ThrowNativeError("Invalid connect callback specified"); if (!pContext->GetFunctionById(params[3])) return pContext->ThrowNativeError("Invalid receive callback specified"); if (!pContext->GetFunctionById(params[4])) return pContext->ThrowNativeError("Invalid disconnect callback specified"); if (params[6] < 0 || params[6] > 65535) return pContext->ThrowNativeError("Invalid port specified"); char *hostname = NULL; pContext->LocalToString(params[5], &hostname); switch (sw->socketType) { case SM_SocketType_Tcp: { Socket* socket = (Socket*) sw->socket; if (socket->IsOpen()) return pContext->ThrowNativeError("Socket is already connected"); socket->connectCallback = pContext->GetFunctionById(params[2]); socket->receiveCallback = pContext->GetFunctionById(params[3]); socket->disconnectCallback = pContext->GetFunctionById(params[4]); return socket->Connect(hostname, params[6]); } case SM_SocketType_Udp: { Socket* socket = (Socket*) sw->socket; if (socket->IsOpen()) return pContext->ThrowNativeError("Socket is already connected"); socket->connectCallback = pContext->GetFunctionById(params[2]); socket->receiveCallback = pContext->GetFunctionById(params[3]); socket->disconnectCallback = pContext->GetFunctionById(params[4]); return socket->Connect(hostname, params[6]); } default: return false; } } // native SocketDisconnect(Handle:socket); cell_t SocketDisconnect(IPluginContext *pContext, const cell_t *params) { SocketWrapper* sw = extension.GetSocketWrapperByHandle(static_cast(params[1])); if (sw == NULL) return pContext->ThrowNativeError("Invalid handle: %i", params[1]); switch (sw->socketType) { case SM_SocketType_Tcp: { Socket* socket = (Socket*) sw->socket; if (!socket->IsOpen()) return pContext->ThrowNativeError("Socket is not connected/listening"); return socket->Disconnect(); } case SM_SocketType_Udp: { Socket* socket = (Socket*) sw->socket; if (!socket->IsOpen()) return pContext->ThrowNativeError("Socket is not connected/listening"); return socket->Disconnect(); } default: return false; } } // native SocketListen(Handle:socket, SocketIncomingCB:ifunc); cell_t SocketListen(IPluginContext *pContext, const cell_t *params) { SocketWrapper* sw = extension.GetSocketWrapperByHandle(static_cast(params[1])); if (sw == NULL) return pContext->ThrowNativeError("Invalid handle: %i", params[1]); if (sw->socketType != SM_SocketType_Tcp) return pContext->ThrowNativeError("The socket must use the TCP/SOCK_STREAM protocol"); if (!pContext->GetFunctionById(params[2])) return pContext->ThrowNativeError("Invalid incoming callback specified"); switch (sw->socketType) { case SM_SocketType_Tcp: { Socket* socket = (Socket*) sw->socket; if (socket->IsOpen()) return pContext->ThrowNativeError("Socket is already open"); socket->incomingCallback = pContext->GetFunctionById(params[2]); return socket->Listen(); } case SM_SocketType_Udp: { Socket* socket = (Socket*) sw->socket; if (socket->IsOpen()) return pContext->ThrowNativeError("Socket is already open"); socket->incomingCallback = pContext->GetFunctionById(params[2]); return socket->Listen(); } default: return false; } } // native SocketSend(Handle:socket, String:command[], size); cell_t SocketSend(IPluginContext *pContext, const cell_t *params) { SocketWrapper* sw = extension.GetSocketWrapperByHandle(static_cast(params[1])); if (sw == NULL) return pContext->ThrowNativeError("Invalid handle: %i", params[1]); char* dataTmp = NULL; pContext->LocalToString(params[2], &dataTmp); std::string data; if (params[3] == -1) { data.assign(dataTmp); } else { data.assign(dataTmp, params[3]); } switch (sw->socketType) { case SM_SocketType_Tcp: { Socket* socket = (Socket*) sw->socket; if (!socket->IsOpen()) return pContext->ThrowNativeError("Can't send, socket is not connected"); return socket->Send(data); } case SM_SocketType_Udp: { Socket* socket = (Socket*) sw->socket; if (!socket->IsOpen()) return pContext->ThrowNativeError("Can't send, socket is not connected"); socket->incomingCallback = pContext->GetFunctionById(params[2]); return socket->Send(data); } default: return false; } } // native SocketSendTo(Handle:socket, const String:data[], size=-1, const String:hostname[], port); cell_t SocketSendTo(IPluginContext *pContext, const cell_t *params) { SocketWrapper* sw = extension.GetSocketWrapperByHandle(static_cast(params[1])); if (sw == NULL) return pContext->ThrowNativeError("Invalid handle: %i", params[1]); if (sw->socketType == SM_SocketType_Tcp) return pContext->ThrowNativeError("This native doesn't support connection orientated protocols"); char* dataTmp = NULL; pContext->LocalToString(params[2], &dataTmp); std::string data; if (params[3] == -1) { data.assign(dataTmp); } else { data.assign(dataTmp, params[3]); } char* hostname = NULL; pContext->LocalToString(params[4], &hostname); switch (sw->socketType) { case SM_SocketType_Udp: { Socket* socket = (Socket*) sw->socket; //if (!socket->IsOpen()) return pContext->ThrowNativeError("Can't send, socket is not connected"); socket->incomingCallback = pContext->GetFunctionById(params[2]); return socket->SendTo(data, hostname, params[5]); } default: return false; } } // native SocketSetOption(Handle:socket, SocketOption:option, value) cell_t SocketSetOption(IPluginContext *pContext, const cell_t *params) { SocketWrapper* sw = extension.GetSocketWrapperByHandle(static_cast(params[1])); if (params[2] != SM_SO_ConcatenateCallbacks && params[2] != SM_SO_ForceFrameLock && params[2] != SM_SO_CallbacksPerFrame && params[2] != SM_SO_DebugMode) { if (sw == NULL) return pContext->ThrowNativeError("Invalid handle: %i", params[1]); switch (sw->socketType) { case SM_SocketType_Tcp: { return ((Socket*) sw->socket)->SetOption((SM_SocketOption) params[2], params[3]); } case SM_SocketType_Udp: { return ((Socket*) sw->socket)->SetOption((SM_SocketOption) params[2], params[3]); } default: return false; } } else { return false; } #if 0 switch (params[2]) { case ConcatenateCallbacks: socket->setOption(ConcatenateCallbacks, value); return 1; case ForceFrameLock: callbacks->setOption(ForceFrameLock, value); return 1; case CallbacksPerFrame: if (value > 0) { callbacks->setOption(CallbacksPerFrame, value); return 1; } else { return 0; } case DebugMode: sockets._debug = value != 0; return 1; ... } #endif } // native SocketSetReceiveCallback(Handle:socket, SocketReceiveCB:rfunc); cell_t SocketSetReceiveCallback(IPluginContext *pContext, const cell_t *params) { SocketWrapper* sw = extension.GetSocketWrapperByHandle(static_cast(params[1])); if (sw == NULL) return pContext->ThrowNativeError("Invalid handle: %i", params[1]); switch (sw->socketType) { case SM_SocketType_Tcp: ((Socket*) sw->socket)->receiveCallback = pContext->GetFunctionById((params[2])); break; case SM_SocketType_Udp: ((Socket*) sw->socket)->receiveCallback = pContext->GetFunctionById((params[2])); default: return false; } return true; } // native SocketSetSendqueueEmptyCallback(Handle:socket, SocketSendqueueEmptyCB:sfunc); cell_t SocketSetSendqueueEmptyCallback(IPluginContext *pContext, const cell_t *params) { SocketWrapper* sw = extension.GetSocketWrapperByHandle(static_cast(params[1])); if (sw == NULL) return pContext->ThrowNativeError("Invalid handle: %i", params[1]); bool forceSendqueueEmptyCallback = false; switch (sw->socketType) { case SM_SocketType_Tcp: ((Socket*) sw->socket)->sendqueueEmptyCallback = pContext->GetFunctionById((params[2])); if (!((Socket*) sw->socket)->sendQueueLength) forceSendqueueEmptyCallback = true; break; case SM_SocketType_Udp: ((Socket*) sw->socket)->sendqueueEmptyCallback = pContext->GetFunctionById((params[2])); if (!((Socket*) sw->socket)->sendQueueLength) forceSendqueueEmptyCallback = true; default: return false; } if (forceSendqueueEmptyCallback) { callbackHandler.AddCallback(new Callback(CallbackEvent_SendQueueEmpty, sw->socket)); } return true; } // native SocketSetDisconnectCallback(Handle:socket, SocketDisconnectCB:dfunc); cell_t SocketSetDisconnectCallback(IPluginContext *pContext, const cell_t *params) { SocketWrapper* sw = extension.GetSocketWrapperByHandle(static_cast(params[1])); if (sw == NULL) return pContext->ThrowNativeError("Invalid handle: %i", params[1]); switch (sw->socketType) { case SM_SocketType_Tcp: ((Socket*) sw->socket)->disconnectCallback = pContext->GetFunctionById((params[2])); break; case SM_SocketType_Udp: ((Socket*) sw->socket)->disconnectCallback = pContext->GetFunctionById((params[2])); default: return false; } return true; } // native SocketSetErrorCallback(Handle:socket, SocketErrorCB:efunc); cell_t SocketSetErrorCallback(IPluginContext *pContext, const cell_t *params) { SocketWrapper* sw = extension.GetSocketWrapperByHandle(static_cast(params[1])); if (sw == NULL) return pContext->ThrowNativeError("Invalid handle: %i", params[1]); switch (sw->socketType) { case SM_SocketType_Tcp: ((Socket*) sw->socket)->errorCallback = pContext->GetFunctionById((params[2])); break; case SM_SocketType_Udp: ((Socket*) sw->socket)->errorCallback = pContext->GetFunctionById((params[2])); default: return false; } return true; } // native SocketSetArg(Handle:socket, any:arg); cell_t SocketSetArg(IPluginContext *pContext, const cell_t *params) { SocketWrapper* sw = extension.GetSocketWrapperByHandle(static_cast(params[1])); if (sw == NULL) return pContext->ThrowNativeError("Invalid handle: %i", params[1]); switch (sw->socketType) { case SM_SocketType_Tcp: ((Socket*) sw->socket)->smCallbackArg = params[2]; break; case SM_SocketType_Udp: ((Socket*) sw->socket)->smCallbackArg = params[2]; default: return false; } return true; } // native SocketGetHostName(String:dest[], destLen); cell_t SocketGetHostName(IPluginContext *pContext, const cell_t *params) { char* dest = NULL; pContext->LocalToString(params[1], &dest); boost::system::error_code errorCode; std::string hostName = host_name(errorCode); if (!errorCode) { size_t len = hostName.copy(dest, params[2]-1); dest[len] = '\0'; return true; } else { dest[0] = '\0'; return false; } } const sp_nativeinfo_t smsock_natives[] = { {"SocketCreate", SocketCreate}, {"SocketBind", SocketBind}, {"SocketConnect", SocketConnect}, {"SocketDisconnect", SocketDisconnect}, {"SocketListen", SocketListen}, {"SocketSend", SocketSend}, {"SocketSendTo", SocketSendTo}, {"SocketSetOption", SocketSetOption}, {"SocketSetReceiveCallback", SocketSetReceiveCallback}, {"SocketSetSendqueueEmptyCallback", SocketSetSendqueueEmptyCallback}, {"SocketSetDisconnectCallback", SocketSetDisconnectCallback}, {"SocketSetErrorCallback", SocketSetErrorCallback}, {"SocketSetArg", SocketSetArg}, {"SocketGetHostName", SocketGetHostName}, {"SocketIsConnected", SocketIsConnected}, // Transitional syntax support. {"Socket.Socket", SocketCreate}, {"Socket.Bind", SocketBind}, {"Socket.Connect", SocketConnect}, {"Socket.Disconnect", SocketDisconnect}, {"Socket.Listen", SocketListen}, {"Socket.Send", SocketSend}, {"Socket.SendTo", SocketSendTo}, {"Socket.SetOption", SocketSetOption}, {"Socket.SetReceiveCallback", SocketSetReceiveCallback}, {"Socket.SetSendqueueEmptyCallback",SocketSetSendqueueEmptyCallback}, {"Socket.SetDisconnectCallback", SocketSetDisconnectCallback}, {"Socket.SetErrorCallback", SocketSetErrorCallback}, {"Socket.SetArg", SocketSetArg}, {"Socket.GetHostName", SocketGetHostName}, {"Socket.Connected.get", SocketIsConnected}, {NULL, NULL}, };