summaryrefslogtreecommitdiff
path: root/Source/Core/Common/TraversalClient.h
blob: 759edd653a35ba815584b0a5298a299160094d7f (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
// SPDX-License-Identifier: CC0-1.0

#pragma once

#include <cstddef>
#include <list>
#include <memory>
#include <string>
#include <string_view>

#include <enet/enet.h>

#include "Common/CommonTypes.h"
#include "Common/ENet.h"
#include "Common/TraversalProto.h"

namespace Common
{
class TraversalClientClient
{
public:
  virtual ~TraversalClientClient() = default;
  virtual void OnTraversalStateChanged() = 0;
  virtual void OnConnectReady(ENetAddress addr) = 0;
  virtual void OnConnectFailed(TraversalConnectFailedReason reason) = 0;
  virtual void OnTtlDetermined(u8 ttl) = 0;
};

class TraversalClient
{
public:
  enum class State
  {
    Connecting,
    Connected,
    Failure,
  };
  enum class FailureReason
  {
    BadHost = 0x300,
    VersionTooOld,
    ServerForgotAboutUs,
    SocketSendError,
    ResendTimeout,
  };
  TraversalClient(ENetHost* netHost, const std::string& server, const u16 port,
                  const u16 port_alt = 0);
  ~TraversalClient();

  TraversalHostId GetHostID() const;
  TraversalInetAddress GetExternalAddress() const;
  State GetState() const;
  FailureReason GetFailureReason() const;

  bool HasFailed() const { return m_State == State::Failure; }
  bool IsConnecting() const { return m_State == State::Connecting; }
  bool IsConnected() const { return m_State == State::Connected; }

  void Reset();
  void ConnectToClient(std::string_view host);
  void ReconnectToServer();
  void Update();
  void HandleResends();

  TraversalClientClient* m_Client = nullptr;

private:
  struct OutgoingTraversalPacketInfo
  {
    TraversalPacket packet;
    int tries;
    u32 sendTime;
  };
  void HandleServerPacket(TraversalPacket* packet);
  // called from NetHost
  bool TestPacket(u8* data, size_t size, ENetAddress* from);
  void ResendPacket(OutgoingTraversalPacketInfo* info);
  TraversalRequestId SendTraversalPacket(const TraversalPacket& packet);
  void OnFailure(FailureReason reason);
  void HandlePing();
  static int ENET_CALLBACK InterceptCallback(ENetHost* host, ENetEvent* event);

  void NewTraversalTest();
  void HandleTraversalTest();

  ENetHost* m_NetHost;
  TraversalHostId m_HostId{};
  TraversalInetAddress m_external_address{};
  State m_State{};
  FailureReason m_FailureReason{};
  TraversalRequestId m_ConnectRequestId = 0;
  bool m_PendingConnect = false;
  std::list<OutgoingTraversalPacketInfo> m_OutgoingTraversalPackets;
  ENetAddress m_ServerAddress{};
  std::string m_Server;
  u16 m_port;
  u16 m_portAlt;
  u32 m_PingTime = 0;

  ENetSocket m_TestSocket = ENET_SOCKET_NULL;
  TraversalRequestId m_TestRequestId = 0;
  u8 m_ttl = 2;
  bool m_ttlReady = false;
};

extern std::unique_ptr<TraversalClient> g_TraversalClient;
// the NetHost connected to the TraversalClient.
extern ENet::ENetHostPtr g_MainNetHost;

// Create g_TraversalClient and g_MainNetHost if necessary.
bool EnsureTraversalClient(const std::string& server, u16 server_port, u16 server_port_alt = 0,
                           u16 listen_port = 0);
void ReleaseTraversalClient();
}  // namespace Common