Files
MontaukOS/kernel/src/Net/Udp.cpp
T

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();
}
}