diff --git a/src/core/common/thread.cpp b/src/core/common/thread.cpp index 698fb5ee..b8e1b4c2 100644 --- a/src/core/common/thread.cpp +++ b/src/core/common/thread.cpp @@ -47,7 +47,6 @@ private: }; Thread::Thread() -: m_userReqTerminateLock(m_shouldTerminateMutex), m_threadStartBarrier(2) { } @@ -58,13 +57,18 @@ Thread::~Thread() void Thread::Run() { - // Create the boost thread object. boost::mutex::scoped_lock threadLock(m_threadObjMutex); + // Create the boost thread object. if (!m_threadObj.get()) { + // Initialise data structures within the context of the thread + // who runs/terminates this thread. + m_userReqTerminateLock.reset(new boost::timed_mutex::scoped_try_lock(m_shouldTerminateMutex)); + m_threadStartBarrier.reset(new boost::barrier(2)); + m_threadObj.reset(new boost::thread(ThreadStarter(*this))); - m_threadStartBarrier.wait(); + m_threadStartBarrier->wait(); } } @@ -72,7 +76,8 @@ void Thread::SignalTermination() { // Unlock the shouldTerminateMutex. - m_userReqTerminateLock.unlock(); + if (m_userReqTerminateLock.get()) // cannot signal before calling Run + m_userReqTerminateLock->unlock(); } bool @@ -118,7 +123,8 @@ void Thread::MainWrapper() { boost::timed_mutex::scoped_lock lock(m_isTerminatedMutex); - m_threadStartBarrier.wait(); + assert(m_threadStartBarrier.get()); + m_threadStartBarrier->wait(); this->Main(); } diff --git a/src/core/thread.h b/src/core/thread.h index 3e808085..eba21d37 100644 --- a/src/core/thread.h +++ b/src/core/thread.h @@ -72,13 +72,13 @@ private: // Flag specifying whether the thread should be terminated. mutable boost::timed_mutex m_shouldTerminateMutex; - mutable boost::timed_mutex::scoped_try_lock m_userReqTerminateLock; + mutable boost::shared_ptr m_userReqTerminateLock; // The boost thread object. boost::shared_ptr m_threadObj; mutable boost::mutex m_threadObjMutex; - boost::barrier m_threadStartBarrier; + mutable boost::shared_ptr m_threadStartBarrier; friend class ThreadStarter; }; diff --git a/src/net/common/clientstate.cpp b/src/net/common/clientstate.cpp index cd0dad33..fd02e372 100644 --- a/src/net/common/clientstate.cpp +++ b/src/net/common/clientstate.cpp @@ -159,7 +159,6 @@ void ClientStateResolving::SetResolver(ResolverThread *resolver) { Cleanup(); - m_resolver = resolver; } diff --git a/src/net/common/clientthread.cpp b/src/net/common/clientthread.cpp index 4cce357f..35e4176c 100644 --- a/src/net/common/clientthread.cpp +++ b/src/net/common/clientthread.cpp @@ -55,6 +55,8 @@ ClientThread::ClientThread(ClientCallback &cb) { m_context.reset(new ClientContext); m_senderCallback.reset(new ClientSenderCallback(*this)); + m_sender.reset(new SenderThread(GetSenderCallback())); + m_receiver.reset(new ReceiverHelper); } ClientThread::~ClientThread() @@ -86,8 +88,6 @@ ClientThread::Main() { SetState(CLIENT_INITIAL_STATE::Instance()); - m_sender.reset(new SenderThread(GetSenderCallback())); - m_receiver.reset(new ReceiverHelper); GetSender().Run(); try diff --git a/src/net/common/netpacket.cpp b/src/net/common/netpacket.cpp index 3d2016a2..748f398d 100644 --- a/src/net/common/netpacket.cpp +++ b/src/net/common/netpacket.cpp @@ -68,6 +68,20 @@ NetPacketInit::Init() m_data.test = htonl(0); } +boost::shared_ptr +NetPacketInit::Clone() const +{ + boost::shared_ptr newPacket(new NetPacketInit); + try + { + newPacket->SetData(GetData()); + } catch (const NetException &) + { + // Need to return the new packet anyway. + } + return newPacket; +} + const NetPacketHeader * NetPacketInit::GetData() const { @@ -118,6 +132,20 @@ NetPacketInitAck::Init() m_data.test = htonl(0); } +boost::shared_ptr +NetPacketInitAck::Clone() const +{ + boost::shared_ptr newPacket(new NetPacketInitAck); + try + { + newPacket->SetData(GetData()); + } catch (const NetException &) + { + // Need to return the new packet anyway. + } + return newPacket; +} + const NetPacketHeader * NetPacketInitAck::GetData() const { @@ -168,6 +196,20 @@ NetPacketGameStart::Init() m_data.test = htonl(0); } +boost::shared_ptr +NetPacketGameStart::Clone() const +{ + boost::shared_ptr newPacket(new NetPacketGameStart); + try + { + newPacket->SetData(GetData()); + } catch (const NetException &) + { + // Need to return the new packet anyway. + } + return newPacket; +} + const NetPacketHeader * NetPacketGameStart::GetData() const { diff --git a/src/net/common/senderthread.cpp b/src/net/common/senderthread.cpp index 3fd34e2d..aa1ddd61 100644 --- a/src/net/common/senderthread.cpp +++ b/src/net/common/senderthread.cpp @@ -24,7 +24,6 @@ using namespace std; -#define SEND_TIMEOUT_MSEC 50 SenderThread::SenderThread(SenderCallback &cb) : m_curSocket(INVALID_SOCKET), m_tmpOutBufSize(0), m_callback(cb) @@ -41,7 +40,9 @@ SenderThread::Send(boost::shared_ptr packet, SOCKET sock) if (packet.get() && IS_VALID_SOCKET(sock)) { boost::mutex::scoped_lock lock(m_outBufMutex); - m_outBuf.push_back(std::make_pair(packet, sock)); + if (m_outBuf.size() < SEND_QUEUE_SIZE) // Queue is limited in size. + m_outBuf.push_back(std::make_pair(packet, sock)); + // TODO: Throw exception if failed. } } diff --git a/src/net/common/serverrecvstate.cpp b/src/net/common/serverrecvstate.cpp index 4def430c..c13d712e 100644 --- a/src/net/common/serverrecvstate.cpp +++ b/src/net/common/serverrecvstate.cpp @@ -121,9 +121,10 @@ ServerRecvStateStartGame::Process(ServerRecvThread &server) { boost::shared_ptr answer(new NetPacketGameStart); - server.SendToAllClients(answer); + server.SendToAllPlayers(answer); + Thread::Msleep(100); - return MSG_SOCK_INIT_DONE; + return MSG_SOCK_INTERNAL_PENDING; } //----------------------------------------------------------------------------- diff --git a/src/net/common/serverrecvthread.cpp b/src/net/common/serverrecvthread.cpp index 8bad3b96..2f1fae18 100644 --- a/src/net/common/serverrecvthread.cpp +++ b/src/net/common/serverrecvthread.cpp @@ -44,10 +44,14 @@ private: ServerRecvThread::ServerRecvThread() { m_senderCallback.reset(new ServerSenderCallback(*this)); + m_sender.reset(new SenderThread(GetSenderCallback())); + m_receiver.reset(new ReceiverHelper); } ServerRecvThread::~ServerRecvThread() { + CleanupConnectQueue(); + CleanupSessionMap(); } void @@ -58,16 +62,18 @@ ServerRecvThread::StartGame() } void -ServerRecvThread::SendToAllClients(boost::shared_ptr packet) +ServerRecvThread::SendToAllPlayers(boost::shared_ptr packet) { - // TODO: possible race condition if used in multithreading - SocketSessionMap::iterator i = m_sessions.begin(); - SocketSessionMap::iterator end = m_sessions.end(); + // This function needs to be thread safe. + boost::mutex::scoped_lock lock(m_sessionMapMutex); + + SocketSessionMap::iterator i = m_sessionMap.begin(); + SocketSessionMap::iterator end = m_sessionMap.end(); while (i != end) { - // TODO: desparately need clone here, this is very dangerous. - GetSender().Send(packet, i->first); + // Send each client a copy of the packet. + GetSender().Send(boost::shared_ptr(packet->Clone()), i->first); ++i; } } @@ -83,9 +89,6 @@ void ServerRecvThread::Main() { SetState(SERVER_INITIAL_STATE::Instance()); - - m_sender.reset(new SenderThread(GetSenderCallback())); - m_receiver.reset(new ReceiverHelper); GetSender().Run(); try @@ -114,7 +117,8 @@ ServerRecvThread::Main() GetSender().SignalTermination(); GetSender().Join(SENDER_THREAD_TERMINATE_TIMEOUT); - // TODO: clear connection queue + CleanupConnectQueue(); + CleanupSessionMap(); } SOCKET @@ -122,31 +126,32 @@ ServerRecvThread::Select() { SOCKET retSock = INVALID_SOCKET; - if (m_sessions.empty()) + SOCKET maxSock = INVALID_SOCKET; + fd_set rdset; + FD_ZERO(&rdset); + + { + boost::mutex::scoped_lock lock(m_sessionMapMutex); + SocketSessionMap::iterator i = m_sessionMap.begin(); + SocketSessionMap::iterator end = m_sessionMap.end(); + + while (i != end) + { + SOCKET tmpSock = i->first; + FD_SET(tmpSock, &rdset); + if (tmpSock > maxSock) + maxSock = tmpSock; + ++i; + } + } + + if (maxSock == INVALID_SOCKET) { Msleep(RECV_TIMEOUT_MSEC); // just sleep if there is no session } else { // wait for data - SOCKET maxSock = 0; - fd_set rdset; - FD_ZERO(&rdset); - - { - SocketSessionMap::iterator i = m_sessions.begin(); - SocketSessionMap::iterator end = m_sessions.end(); - - while (i != end) - { - SOCKET tmpSock = i->first; - FD_SET(tmpSock, &rdset); - if (tmpSock > maxSock) - maxSock = tmpSock; - ++i; - } - } - struct timeval timeout; timeout.tv_sec = 0; timeout.tv_usec = RECV_TIMEOUT_MSEC * 1000; @@ -158,8 +163,9 @@ ServerRecvThread::Select() if (selectResult > 0) // one (or more) of the sockets is readable { // Check which socket is readable, return the first. - SocketSessionMap::iterator i = m_sessions.begin(); - SocketSessionMap::iterator end = m_sessions.end(); + boost::mutex::scoped_lock lock(m_sessionMapMutex); + SocketSessionMap::iterator i = m_sessionMap.begin(); + SocketSessionMap::iterator end = m_sessionMap.end(); while (i != end) { @@ -176,6 +182,34 @@ ServerRecvThread::Select() return retSock; } +void +ServerRecvThread::CleanupConnectQueue() +{ + boost::mutex::scoped_lock lock(m_connectQueueMutex); + + // Sockets will be closed automatically. + m_connectQueue.clear(); +} + +void +ServerRecvThread::CleanupSessionMap() +{ + boost::mutex::scoped_lock lock(m_sessionMapMutex); + + // We need to manually close all sockets for the sessions. + // This is "not great", but there are some issues when + // automatically closing them. + SocketSessionMap::iterator i = m_sessionMap.begin(); + SocketSessionMap::iterator end = m_sessionMap.end(); + + while (i != end) + { + CLOSESOCKET(i->first); + ++i; + } + m_sessionMap.clear(); +} + ServerRecvState & ServerRecvThread::GetState() { @@ -193,9 +227,10 @@ boost::shared_ptr ServerRecvThread::GetSession(SOCKET sock) { boost::shared_ptr tmpSession; - - SocketSessionMap::iterator pos = m_sessions.find(sock); - if (pos != m_sessions.end()) + boost::mutex::scoped_lock lock(m_sessionMapMutex); + + SocketSessionMap::iterator pos = m_sessionMap.find(sock); + if (pos != m_sessionMap.end()) { tmpSession = pos->second; } @@ -205,15 +240,17 @@ ServerRecvThread::GetSession(SOCKET sock) void ServerRecvThread::AddSession(boost::shared_ptr connData, boost::shared_ptr sessionData) { - SocketSessionMap::iterator pos = m_sessions.lower_bound(connData->GetSocket()); + boost::mutex::scoped_lock lock(m_sessionMapMutex); + + SocketSessionMap::iterator pos = m_sessionMap.lower_bound(connData->GetSocket()); // If pos points to a pair whose key is equivalent to the socket, this handle // already exists within the list. - if (pos != m_sessions.end() && connData->GetSocket() == pos->first) + if (pos != m_sessionMap.end() && connData->GetSocket() == pos->first) { throw ServerException(ERR_SOCK_CONN_EXISTS, 0); } - m_sessions.insert(pos, SocketSessionMap::value_type(connData->ReleaseSocket(), sessionData)); + m_sessionMap.insert(pos, SocketSessionMap::value_type(connData->ReleaseSocket(), sessionData)); } SenderThread & diff --git a/src/net/common/serverthread.cpp b/src/net/common/serverthread.cpp index 7c810587..7e479818 100644 --- a/src/net/common/serverthread.cpp +++ b/src/net/common/serverthread.cpp @@ -33,6 +33,7 @@ ServerThread::ServerThread(ServerCallback &cb) : m_callback(cb) { m_context.reset(new ServerContext); + m_recvThread.reset(new ServerRecvThread); } ServerThread::~ServerThread() @@ -71,7 +72,6 @@ ServerThread::GetCallback() void ServerThread::Main() { - m_recvThread.reset(new ServerRecvThread); try { Listen(); @@ -86,6 +86,8 @@ ServerThread::Main() { GetCallback().SignalNetServerError(e.GetErrorId(), e.GetOsErrorCode()); } + GetRecvThread().SignalTermination(); + GetRecvThread().Join(RECEIVER_THREAD_TERMINATE_TIMEOUT); } void diff --git a/src/net/netpacket.h b/src/net/netpacket.h index 17a6a800..638e5267 100644 --- a/src/net/netpacket.h +++ b/src/net/netpacket.h @@ -22,6 +22,7 @@ #define _NETPACKET_H_ #include +#include #include #define MAX_PACKET_SIZE 256 @@ -75,6 +76,8 @@ class NetPacket public: virtual ~NetPacket(); + virtual boost::shared_ptr Clone() const = 0; + virtual void SetData(const NetPacketHeader *p) = 0; virtual const NetPacketHeader *GetData() const = 0; @@ -90,6 +93,8 @@ public: NetPacketInit(u_int32_t value); virtual ~NetPacketInit(); + virtual boost::shared_ptr Clone() const; + virtual const NetPacketHeader *GetData() const; virtual void SetData(const NetPacketHeader *p); @@ -109,6 +114,8 @@ public: NetPacketInitAck(u_int32_t value); virtual ~NetPacketInitAck(); + virtual boost::shared_ptr Clone() const; + virtual const NetPacketHeader *GetData() const; virtual void SetData(const NetPacketHeader *p); @@ -128,6 +135,8 @@ public: NetPacketGameStart(u_int32_t value); virtual ~NetPacketGameStart(); + virtual boost::shared_ptr Clone() const; + virtual const NetPacketHeader *GetData() const; virtual void SetData(const NetPacketHeader *p); diff --git a/src/net/senderthread.h b/src/net/senderthread.h index 65002a00..d049b933 100644 --- a/src/net/senderthread.h +++ b/src/net/senderthread.h @@ -29,7 +29,9 @@ #include #include -#define SENDER_THREAD_TERMINATE_TIMEOUT 100 +#define SENDER_THREAD_TERMINATE_TIMEOUT 200 +#define SEND_TIMEOUT_MSEC 50 +#define SEND_QUEUE_SIZE 200 class SenderThread : public Thread { diff --git a/src/net/serverrecvthread.h b/src/net/serverrecvthread.h index bbf78b3f..b3dce6c4 100644 --- a/src/net/serverrecvthread.h +++ b/src/net/serverrecvthread.h @@ -29,6 +29,8 @@ #include #include +#define RECEIVER_THREAD_TERMINATE_TIMEOUT 200 + class ServerRecvState; class SenderThread; class ReceiverHelper; @@ -42,7 +44,7 @@ public: virtual ~ServerRecvThread(); void StartGame(); - void SendToAllClients(boost::shared_ptr packet); + void SendToAllPlayers(boost::shared_ptr packet); void AddConnection(boost::shared_ptr data); protected: @@ -54,6 +56,9 @@ protected: SOCKET Select(); + void CleanupConnectQueue(); + void CleanupSessionMap(); + ServerRecvState &GetState(); void SetState(ServerRecvState &newState); @@ -71,7 +76,8 @@ private: mutable boost::mutex m_connectQueueMutex; ServerRecvState *m_curState; - SocketSessionMap m_sessions; + SocketSessionMap m_sessionMap; + mutable boost::mutex m_sessionMapMutex; std::auto_ptr m_receiver; std::auto_ptr m_sender;