Add security limits for network to avoid too many messages processing

This commit is contained in:
2026-09-09 06:02:37 +02:00
parent 9ded33b717
commit 652fbf9a52
4 changed files with 117 additions and 5 deletions
+54 -4
View File
@@ -6,6 +6,7 @@
#include "NetworkInternal.h"
#include "NetworkStream.h"
#include "NetworkMessage.h"
#include "NetworkSettings.h"
#include "NetworkPeer.h"
#include "NetworkChannelType.h"
#include "NetworkEvent.h"
@@ -19,6 +20,7 @@
#include "Engine/Core/Types/DataContainer.h"
#include "Engine/Platform/CriticalSection.h"
#include "Engine/Engine/EngineService.h"
#include "Engine/Engine/Time.h"
#include "Engine/Level/Actor.h"
#include "Engine/Level/SceneObject.h"
#include "Engine/Level/Prefabs/Prefab.h"
@@ -200,6 +202,7 @@ struct PartsItem
uint16 PartsLeft;
uint32 OwnerFrame;
uint32 OwnerClientId;
float LastPartTime;
const void* Tag;
Array<byte> Data;
};
@@ -211,6 +214,7 @@ struct SpawnItem
bool HasOwnership = false;
bool HierarchicalOwnership = false;
uint32 OwnerClientId;
float LastPartTime;
NetworkObjectRole Role;
};
@@ -218,6 +222,7 @@ struct SpawnItemParts
{
NetworkMessageObjectSpawn MsgData;
uint32 OwnerClientId; // Duplicate of MsgData.OwnerClientId to reuse template code for other structs
float LastPartTime;
Guid PrefabId;
Array<NetworkMessageObjectSpawnItem> Items;
};
@@ -266,6 +271,7 @@ namespace
Array<Guid> DespawnedObjects;
uint32 SpawnId = 0;
uint32 RpcId = 0;
float NetworkTime = 0;
NetworkStream* GetWriteStream()
{
@@ -544,6 +550,35 @@ void RemoveOwnerFromItems(Array<T>& items, uint32 clientId)
}
}
template<typename T>
void CleanupOldParts(Array<T>& items, int32 maxParts, float ttl, const Char* hint)
{
// Limit amount of parts
if (items.Count() > maxParts && maxParts > 0)
{
LOG(Warning, "Too many {} network parts in-flight: {}, limit: {}", hint, items.Count(), maxParts);
items.Resize(maxParts);
}
// Remove expired parts
if (ttl > 0.0f)
{
int32 count = 0;
for (int32 i = items.Count() - 1; i >= 0; i--)
{
if (NetworkTime - items[i].LastPartTime > ttl)
{
items.RemoveAt(i);
count++;
}
}
if (count > 0)
{
LOG(Warning, "Removed {} expired {} network parts (TTL is {}s)", count, hint, ttl);
}
}
}
void SetupObjectSpawnMessageItem(SpawnItem* e, NetworkMessage& msg)
{
ScriptingObject* obj = e->Object.Get();
@@ -986,6 +1021,7 @@ void FindObjectsForSpawn(SpawnGroup& group, ChunkedArray<SpawnItem, 256>& spawnI
spawnItem.Object = obj;
spawnItem.Targets.Link(item.TargetClientIds);
spawnItem.OwnerClientId = item.OwnerClientId;
spawnItem.LastPartTime = NetworkTime;
spawnItem.Role = item.Role;
group.Items.Add(&spawnItem);
}
@@ -1039,9 +1075,10 @@ PartsItem* AddPartsItem(Array<PartsItem>& items, NetworkEvent& event, uint32 own
// Copy part data
item->PartsLeft--;
const void* partData = event.Message.SkipBytes(partSize);
if (!partData)
if (!partData || EnumHasAnyFlags(event.Message.Flags, NetworkMessageFlags::HasError))
return nullptr;
Platform::MemoryCopy(item->Data.Get() + partStart, partData, partSize);
item->LastPartTime = NetworkTime;
return item;
}
@@ -1654,6 +1691,8 @@ bool NetworkReplicator::HasObject(const ScriptingObject* obj)
const auto it = Objects.Find(obj->GetID());
if (it != Objects.End())
return true;
// Check in-flight spawn queue
for (const SpawnItem& item : SpawnQueue)
{
if (item.Object == obj)
@@ -1699,6 +1738,7 @@ uint32 NetworkReplicator::GetObjectOwnerClientId(const ScriptingObject* obj)
id = it->Item.OwnerClientId;
else
{
// Check in-flight spawn queue
for (const SpawnItem& item : SpawnQueue)
{
if (item.Object == obj)
@@ -1731,6 +1771,7 @@ NetworkObjectRole NetworkReplicator::GetObjectRole(const ScriptingObject* obj)
role = it->Item.Role;
else
{
// Check in-flight spawn queue
for (const SpawnItem& item : SpawnQueue)
{
if (item.Object == obj)
@@ -1790,6 +1831,7 @@ void NetworkReplicator::SetObjectOwnership(ScriptingObject* obj, uint32 ownerCli
item.HasOwnership = true;
item.HierarchicalOwnership = hierarchical;
item.OwnerClientId = ownerClientId;
item.LastPartTime = NetworkTime;
item.Role = localRole;
break;
}
@@ -2035,6 +2077,7 @@ void NetworkInternal::NetworkReplicatorUpdate()
if (Objects.Count() == 0)
return;
const bool isClient = NetworkManager::IsClient();
NetworkTime = Time::Update.UnscaledTime.GetTotalSeconds();
if (!isClient && NewClients.Count() != 0)
{
@@ -2055,6 +2098,7 @@ void NetworkInternal::NetworkReplicatorUpdate()
spawnItem.Object = obj;
spawnItem.Targets.Link(item.TargetClientIds);
spawnItem.OwnerClientId = item.OwnerClientId;
spawnItem.LastPartTime = NetworkTime;
spawnItem.Role = item.Role;
SetupObjectSpawnGroupItem(obj, spawnGroups, spawnItem);
@@ -2192,9 +2236,11 @@ void NetworkInternal::NetworkReplicatorUpdate()
}
}
// TODO: remove items from RpcParts after some TTL to reduce memory usage
// TODO: remove items from SpawnParts after some TTL to reduce memory usage
// TODO: remove items from ReplicationParts after some TTL to reduce memory usage
// Limit partial messages to avoid overallocating or plugging
const auto& settings = *NetworkSettings::Get();
CleanupOldParts(RpcParts, settings.MaxSyncParts, settings.MaxSyncPartTTL, TEXT("RPC"));
CleanupOldParts(SpawnParts, settings.MaxSyncParts, settings.MaxSyncPartTTL, TEXT("Spawn"));
CleanupOldParts(ReplicationParts, settings.MaxSyncParts, settings.MaxSyncPartTTL, TEXT("Replication"));
// Replicate all owned networked objects with other clients or server
if (!CachedReplicationResult)
@@ -2338,6 +2384,7 @@ void NetworkInternal::OnNetworkMessageObjectSpawn(NetworkEvent& event, NetworkCl
auto& parts = SpawnParts.AddOne();
parts.MsgData = msgData;
parts.OwnerClientId = msgData.OwnerClientId;
parts.LastPartTime = NetworkTime;
parts.PrefabId = prefabId;
parts.Items.Resize(msgData.ItemsCount);
for (auto& item : parts.Items)
@@ -2371,6 +2418,7 @@ void NetworkInternal::OnNetworkMessageObjectSpawnPart(NetworkEvent& event, Netwo
return;
}
auto& spawnParts = SpawnParts.Get()[spawnPartsIndex];
spawnParts.LastPartTime = NetworkTime;
// Read all items from this part
constexpr uint32 spawnItemMaxSize = sizeof(uint16) + sizeof(NetworkMessageObjectSpawnItem); // Index + Data
@@ -2379,6 +2427,8 @@ void NetworkInternal::OnNetworkMessageObjectSpawnPart(NetworkEvent& event, Netwo
const uint16 itemIndex = event.Message.ReadUInt16();
event.Message.ReadStructure(spawnParts.Items[itemIndex]);
}
if (EnumHasAnyFlags(event.Message.Flags, NetworkMessageFlags::HasError))
return;
// Invoke spawning if we've got all items
for (auto& e : spawnParts.Items)