144 lines
4.3 KiB
C++
144 lines
4.3 KiB
C++
/*
|
|
* Udp.cpp
|
|
* User Datagram Protocol
|
|
* Copyright (c) 2025 Daniel Hammer
|
|
*/
|
|
|
|
#include "Udp.hpp"
|
|
#include <Net/Ipv4.hpp>
|
|
#include <Net/ByteOrder.hpp>
|
|
#include <Net/NetConfig.hpp>
|
|
#include <Libraries/Memory.hpp>
|
|
#include <Terminal/Terminal.hpp>
|
|
#include <CppLib/Stream.hpp>
|
|
#include <CppLib/Spinlock.hpp>
|
|
|
|
using namespace Kt;
|
|
|
|
namespace Net::Udp {
|
|
|
|
struct PortBinding {
|
|
uint16_t Port;
|
|
RecvCallback Callback;
|
|
bool Active;
|
|
};
|
|
|
|
static constexpr uint32_t MAX_BINDINGS = 64;
|
|
static PortBinding g_bindings[MAX_BINDINGS] = {};
|
|
static kcp::Spinlock g_bindingsLock;
|
|
|
|
void Initialize() {
|
|
for (uint32_t i = 0; i < MAX_BINDINGS; i++) {
|
|
g_bindings[i].Active = false;
|
|
}
|
|
KernelLogStream(OK, "Net") << "UDP initialized";
|
|
}
|
|
|
|
void OnPacketReceived(uint32_t srcIp, uint32_t dstIp, const uint8_t* data, uint16_t length) {
|
|
if (length < HEADER_SIZE) {
|
|
return;
|
|
}
|
|
|
|
const Header* hdr = (const Header*)data;
|
|
uint16_t srcPort = Ntohs(hdr->SrcPort);
|
|
uint16_t dstPort = Ntohs(hdr->DstPort);
|
|
uint16_t udpLen = Ntohs(hdr->Length);
|
|
|
|
if (udpLen < HEADER_SIZE || udpLen > length) {
|
|
return;
|
|
}
|
|
|
|
// Verify checksum if present
|
|
if (hdr->Checksum != 0) {
|
|
uint16_t check = Ipv4::PseudoHeaderChecksum(srcIp, dstIp, Ipv4::PROTO_UDP,
|
|
udpLen, data, udpLen);
|
|
if (check != 0) {
|
|
return;
|
|
}
|
|
}
|
|
|
|
const uint8_t* payload = data + HEADER_SIZE;
|
|
uint16_t payloadLen = udpLen - HEADER_SIZE;
|
|
|
|
// Snapshot the callback under the binding lock, then invoke it after
|
|
// releasing the lock so callbacks may safely bind/unbind other ports.
|
|
RecvCallback callback = nullptr;
|
|
g_bindingsLock.Acquire();
|
|
for (uint32_t i = 0; i < MAX_BINDINGS; i++) {
|
|
if (g_bindings[i].Active && g_bindings[i].Port == dstPort) {
|
|
callback = g_bindings[i].Callback;
|
|
break;
|
|
}
|
|
}
|
|
g_bindingsLock.Release();
|
|
if (callback) callback(srcIp, srcPort, dstPort, payload, payloadLen);
|
|
}
|
|
|
|
bool Send(uint32_t destIp, uint16_t srcPort, uint16_t destPort,
|
|
const uint8_t* payload, uint16_t payloadLen) {
|
|
uint16_t udpLen = HEADER_SIZE + payloadLen;
|
|
uint8_t packet[1500];
|
|
|
|
if (udpLen > sizeof(packet)) {
|
|
return false;
|
|
}
|
|
|
|
Header* hdr = (Header*)packet;
|
|
hdr->SrcPort = Htons(srcPort);
|
|
hdr->DstPort = Htons(destPort);
|
|
hdr->Length = Htons(udpLen);
|
|
hdr->Checksum = 0;
|
|
|
|
memcpy(packet + HEADER_SIZE, payload, payloadLen);
|
|
|
|
// Calculate checksum with pseudo-header
|
|
hdr->Checksum = Ipv4::PseudoHeaderChecksum(
|
|
Net::GetIpAddress(), destIp, Ipv4::PROTO_UDP,
|
|
udpLen, packet, udpLen);
|
|
if (hdr->Checksum == 0) {
|
|
hdr->Checksum = 0xFFFF; // RFC 768: zero checksum transmitted as all ones
|
|
}
|
|
|
|
return Ipv4::Send(destIp, Ipv4::PROTO_UDP, packet, udpLen);
|
|
}
|
|
|
|
bool Bind(uint16_t port, RecvCallback callback) {
|
|
if (port == 0 || callback == nullptr) return false;
|
|
g_bindingsLock.Acquire();
|
|
// Check for duplicate
|
|
for (uint32_t i = 0; i < MAX_BINDINGS; i++) {
|
|
if (g_bindings[i].Active && g_bindings[i].Port == port) {
|
|
g_bindingsLock.Release();
|
|
return false;
|
|
}
|
|
}
|
|
|
|
// Find empty slot
|
|
for (uint32_t i = 0; i < MAX_BINDINGS; i++) {
|
|
if (!g_bindings[i].Active) {
|
|
g_bindings[i].Port = port;
|
|
g_bindings[i].Callback = callback;
|
|
g_bindings[i].Active = true;
|
|
g_bindingsLock.Release();
|
|
return true;
|
|
}
|
|
}
|
|
g_bindingsLock.Release();
|
|
return false;
|
|
}
|
|
|
|
void Unbind(uint16_t port) {
|
|
g_bindingsLock.Acquire();
|
|
for (uint32_t i = 0; i < MAX_BINDINGS; i++) {
|
|
if (g_bindings[i].Active && g_bindings[i].Port == port) {
|
|
g_bindings[i].Active = false;
|
|
g_bindings[i].Callback = nullptr;
|
|
g_bindingsLock.Release();
|
|
return;
|
|
}
|
|
}
|
|
g_bindingsLock.Release();
|
|
}
|
|
|
|
}
|