tidal-link-1.2.1: link/include/ableton/link_audio/UdpMessenger.hpp
/* Copyright 2025, Ableton AG, Berlin. All rights reserved.
*
* This program is free software: you can redistribute it and/or modify
* it under the terms of the GNU General Public License as published by
* the Free Software Foundation, either version 2 of the License, or
* (at your option) any later version.
*
* This program is distributed in the hope that it will be useful,
* but WITHOUT ANY WARRANTY; without even the implied warranty of
* MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
* GNU General Public License for more details.
*
* You should have received a copy of the GNU General Public License
* along with this program. If not, see <http://www.gnu.org/licenses/>.
*
* If you would like to incorporate Link into a proprietary software application,
* please contact <link-devs@ableton.com>.
*/
#pragma once
#include <ableton/discovery/AsioTypes.hpp>
#include <ableton/discovery/UdpMessenger.hpp>
#include <ableton/discovery/UnicastIpInterface.hpp>
#include <ableton/link/PayloadEntries.hpp>
#include <ableton/link_audio/ChannelRequests.hpp>
#include <ableton/link_audio/NetworkMetrics.hpp>
#include <ableton/link_audio/v1/Messages.hpp>
#include <ableton/util/Injected.hpp>
#include <ableton/util/SafeAsyncHandler.hpp>
#include <algorithm>
#include <chrono>
#include <memory>
#include <optional>
#include <stdexcept>
#include <vector>
namespace ableton
{
namespace link_audio
{
// Throws UdpSendException
template <typename Interface, typename NodeId, typename Payload>
void sendLinkAudioUdpMessage(Interface& iface,
NodeId from,
const uint8_t ttl,
const v1::MessageType messageType,
const Payload& payload,
const discovery::UdpEndpoint& to)
{
using namespace std;
v1::MessageBuffer buffer;
const auto messageBegin = begin(buffer);
try
{
const auto messageEnd =
v1::detail::encodeMessage(std::move(from), ttl, messageType, payload, messageBegin);
const auto numBytes = static_cast<size_t>(distance(messageBegin, messageEnd));
iface.send(buffer.data(), numBytes, to);
}
catch (const std::runtime_error& err)
{
throw discovery::UdpSendException{err, iface.endpoint().address()};
}
}
// UdpMessenger uses a "shared_ptr pImpl" pattern to make it movable
// and to support safe async handler callbacks when receiving messages
// on the given interface.
template <typename ChannelsMessageHandler,
typename Interface,
typename Observer,
typename IoContext>
class UdpMessenger
{
public:
using ObserverType = typename util::Injected<Observer>::type;
using Announcement = typename ObserverType::GatewayObserverAnnouncement;
using NodeId = typename ObserverType::GatewayObserverNodeId;
using Timer = typename util::Injected<IoContext>::type::Timer;
using TimerError = typename Timer::ErrorCode;
using TimePoint = typename Timer::TimePoint;
using SharedInterface = std::shared_ptr<Interface>;
struct ExtendedAnnouncement
{
link::NodeId ident() const { return announcement.ident(); }
Announcement announcement;
double networkQuality;
SharedInterface pInterface;
discovery::UdpEndpoint from;
int ttl;
};
UdpMessenger(util::Injected<ChannelsMessageHandler> handler,
SharedInterface pIface,
Announcement announcement,
util::Injected<IoContext> io,
const uint8_t ttl,
const uint8_t ttlRatio,
util::Injected<Observer> observer)
: mpImpl(std::make_shared<Impl>(std::move(handler),
std::move(pIface),
std::move(announcement),
std::move(io),
ttl,
ttlRatio,
std::move(observer)))
{
// We need to always listen for incoming traffic in order to
// respond to announcement broadcasts
mpImpl->listen();
mpImpl->broadcastAnnouncement();
}
UdpMessenger(const UdpMessenger&) = delete;
UdpMessenger& operator=(const UdpMessenger&) = delete;
UdpMessenger(UdpMessenger&& rhs)
: mpImpl(std::move(rhs.mpImpl))
{
}
~UdpMessenger()
{
if (mpImpl != nullptr)
{
mpImpl->mTimer.cancel();
mpImpl->sendAudioChannelByes({});
}
}
void updateAnnouncement(Announcement announcement)
{
mpImpl->updateAnnouncement(std::move(announcement));
}
discovery::UdpEndpoint endpoint() const { return mpImpl->mpInterface->endpoint(); }
template <typename It>
void pruneChannelsEndpoints(It peersBegin, It peersEnd)
{
auto it = std::remove_if(
mpImpl->mReceivers.begin(),
mpImpl->mReceivers.end(),
[&](const auto& r)
{
return std::none_of(
peersBegin, peersEnd, [&](const auto& peerId) { return peerId == r.id; });
});
mpImpl->mReceivers.erase(it, mpImpl->mReceivers.end());
}
void sawLinkAudioEndpoint(link::NodeId peerId,
std::optional<discovery::UdpEndpoint> endpoint)
{
if (endpoint)
{
if (endpoint->address().is_v6())
{
*endpoint = discovery::ipV6Endpoint(*mpImpl->mpInterface, *endpoint);
}
auto it =
std::find_if(mpImpl->mReceivers.begin(),
mpImpl->mReceivers.end(),
[&](const auto& receiver) { return receiver.endpoint == endpoint; });
if (it == mpImpl->mReceivers.end())
{
mpImpl->mReceivers.push_back({peerId, *endpoint, {}});
}
}
else
{
auto it = std::remove_if(mpImpl->mReceivers.begin(),
mpImpl->mReceivers.end(),
[&](const auto& r) { return r.id == peerId; });
mpImpl->mReceivers.erase(it, mpImpl->mReceivers.end());
}
}
private:
struct Receiver
{
link::NodeId id;
discovery::UdpEndpoint endpoint;
NetworkMetricsFilter metricsFilter;
};
struct Impl : std::enable_shared_from_this<Impl>
{
Impl(util::Injected<ChannelsMessageHandler> handler,
SharedInterface pIface,
Announcement announcement,
util::Injected<IoContext> io,
const uint8_t ttl,
const uint8_t ttlRatio,
util::Injected<Observer> observer)
: mIo(std::move(io))
, mChannelsMessageHandler(std::move(handler))
, mpInterface(std::move(pIface))
, mTimer(mIo->makeTimer())
, mLastBroadcastTime{}
, mTtl(ttl)
, mTtlRatio(ttlRatio)
, mObserver(std::move(observer))
{
updateAnnouncement(std::move(announcement));
}
void sendAudioChannelByes(const ChannelAnnouncements& newAnnouncements)
{
auto channelByes = ChannelByes{};
for (const auto& announcement : mAnnouncements)
{
for (const auto& channel : announcement.channels.channels)
{
if (std::none_of(newAnnouncements.channels.begin(),
newAnnouncements.channels.end(),
[&](const auto& c) { return c.id == channel.id; }))
{
channelByes.byes.push_back({channel.id});
}
}
}
if (channelByes.byes.size() > 0)
{
auto byesToSend = std::vector<ChannelByes>{{}};
for (const auto& bye : channelByes.byes)
{
auto addedSize = sizeInByteStream(bye);
if (sizeInByteStream(byesToSend.back()) + addedSize > v1::kMaxPayloadSize)
{
byesToSend.emplace_back();
}
byesToSend.back().byes.push_back(bye);
}
for (const auto& receiver : mReceivers)
{
for (const auto& byes : byesToSend)
{
try
{
sendLinkAudioUdpMessage(*mpInterface,
mAnnouncements.back().ident(),
mTtl,
v1::kChannelByes,
discovery::makePayload(byes),
receiver.endpoint);
}
catch (const discovery::UdpSendException&)
{
}
}
}
}
}
void updateAnnouncement(Announcement announcement)
{
sendAudioChannelByes(announcement.channels);
const auto pingPayload = discovery::makePayload(link::HostTime{});
const auto pingSize = sizeInByteStream(pingPayload);
mAnnouncements = {Announcement{
announcement.nodeId, announcement.sessionId, announcement.peerInfo, {}}};
for (const auto& channel : announcement.channels.channels)
{
const auto channelSize = sizeInByteStream(channel);
// A ping is sent along with the first announcement
auto addedSize =
mAnnouncements.size() == 1 ? channelSize + pingSize : channelSize;
if (sizeInByteStream(toPayload(mAnnouncements.back())) + addedSize
> v1::kMaxPayloadSize)
{
mAnnouncements.push_back(Announcement{
announcement.nodeId, announcement.sessionId, announcement.peerInfo, {}});
}
mAnnouncements.back().channels.channels.push_back(channel);
}
}
void broadcastAnnouncement()
{
using namespace std::chrono;
const auto minBroadcastPeriod = milliseconds{50};
const auto nominalBroadcastPeriod = milliseconds(mTtl * 1000 / mTtlRatio);
const auto timeSinceLastBroadcast =
duration_cast<milliseconds>(mTimer.now() - mLastBroadcastTime);
// The rate is limited to maxBroadcastRate to prevent flooding the network.
const auto delay = minBroadcastPeriod - timeSinceLastBroadcast;
// Schedule the next broadcast before we actually send the
// message so that if sending throws an exception we are still
// scheduled to try again. We want to keep trying at our
// interval as long as this instance is alive.
mTimer.expires_from_now(delay > milliseconds{0} ? delay : nominalBroadcastPeriod);
mTimer.async_wait(
[this](const TimerError e)
{
if (!e)
{
broadcastAnnouncement();
}
});
// // If we're not delaying, broadcast now
if (delay < milliseconds{1})
{
debug(mIo->log()) << "Broadcasting Announcement";
sendAnnouncement();
}
}
void sendAnnouncement()
{
const auto pingTime = std::chrono::duration_cast<std::chrono::microseconds>(
mTimer.now().time_since_epoch());
for (auto& receiver : mReceivers)
{
try
{
// Send one ping per receiver
auto shouldSendPing = true;
for (const auto& announcement : mAnnouncements)
{
if (shouldSendPing)
{
const auto hostTime = link::HostTime{pingTime};
sendLinkAudioUdpMessage(
*mpInterface,
announcement.ident(),
mTtl,
v1::kPeerAnnouncement,
toPayload(announcement) + discovery::makePayload(hostTime),
receiver.endpoint);
shouldSendPing = false;
}
else
{
sendLinkAudioUdpMessage(*mpInterface,
announcement.ident(),
mTtl,
v1::kPeerAnnouncement,
toPayload(announcement),
receiver.endpoint);
}
}
}
catch (const discovery::UdpSendException&)
{
}
}
mLastBroadcastTime = mTimer.now();
}
void listen() { mpInterface->receive(util::makeAsyncSafe(this->shared_from_this())); }
template <typename It>
void operator()(const discovery::UdpEndpoint& from,
const It messageBegin,
const It messageEnd)
{
auto result = v1::parseMessageHeader(messageBegin, messageEnd);
const auto& header = result.first;
// Ignore messages from self and other groups
if (header.ident != mAnnouncements.front().ident() && header.groupId == 0)
{
debug(mIo->log()) << "Received message type "
<< static_cast<int>(header.messageType) << " from peer "
<< header.ident;
switch (header.messageType)
{
case v1::kPeerAnnouncement:
receiveAnnouncement(std::move(result.first), result.second, messageEnd, from);
receivePing(result.second, messageEnd, from);
break;
case v1::kChannelByes:
receiveChannelByes(result.second, messageEnd);
break;
case v1::kPong:
receivePong(result.second, messageEnd, from);
break;
case v1::kChannelRequest:
receiveChannelRequest(std::move(result.first), result.second, messageEnd);
break;
case v1::kStopChannelRequest:
receiveChannelStopRequest(std::move(result.first), result.second, messageEnd);
break;
case v1::kAudioBuffer:
receiveAudioBuffer(std::move(result.first), result.second, messageEnd);
break;
default:
info(mIo->log()) << "Unknown message received of type: " << header.messageType;
}
}
listen();
}
template <typename It>
void receivePing(It payloadBegin, It payloadEnd, discovery::UdpEndpoint from)
{
std::optional<std::chrono::microseconds> oHostTime;
discovery::parsePayload<link::HostTime>(payloadBegin,
payloadEnd,
[&oHostTime](link::HostTime ht)
{ oHostTime = std::move(ht.time); });
if (oHostTime)
{
try
{
sendLinkAudioUdpMessage(*mpInterface,
mAnnouncements.front().ident(),
mTtl,
v1::kPong,
discovery::makePayload(link::HostTime{*oHostTime}),
from);
}
catch (const std::runtime_error& err)
{
return;
}
}
}
template <typename It>
void receivePong(It payloadBegin, It payloadEnd, discovery::UdpEndpoint from)
{
const auto receiveTime = std::chrono::duration_cast<std::chrono::microseconds>(
mTimer.now().time_since_epoch());
std::chrono::microseconds sendTime{0};
try
{
discovery::parsePayload<link::HostTime>(payloadBegin,
payloadEnd,
[&sendTime](link::HostTime ht)
{ sendTime = std::move(ht.time); });
auto it =
std::find_if(mReceivers.begin(),
mReceivers.end(),
[&](const auto& receiver) { return receiver.endpoint == from; });
if (it != mReceivers.end())
{
it->metricsFilter(receiveTime - sendTime);
}
}
catch (const std::runtime_error& err)
{
return;
}
}
template <typename It>
void receiveAnnouncement(v1::MessageHeader header,
It payloadBegin,
It payloadEnd,
discovery::UdpEndpoint from)
{
const auto it =
std::find_if(mReceivers.begin(),
mReceivers.end(),
[&](const auto& receiver) { return receiver.endpoint == from; });
if (it != mReceivers.end())
{
try
{
auto announcement = Announcement::fromPayload(
std::move(header.ident), std::move(payloadBegin), std::move(payloadEnd));
sawAnnouncement(*mObserver,
ExtendedAnnouncement{std::move(announcement),
it->metricsFilter.metrics().quality(),
mpInterface,
from,
mTtl});
}
catch (const std::runtime_error& err)
{
info(mIo->log()) << "Ignoring peer announcement message: " << err.what();
}
}
}
template <typename It>
void receiveChannelByes(It payloadBegin, It payloadEnd)
{
auto byes = ChannelByes{};
discovery::parsePayload<ChannelByes>(
payloadBegin, payloadEnd, [&](const auto& b) { byes = std::move(b); });
std::vector<NodeId> byesVector;
for (const auto& bye : byes.byes)
{
byesVector.push_back(bye.id);
}
channelsLeft(*mObserver, begin(byesVector), end(byesVector));
}
template <typename It>
void receiveChannelRequest(v1::MessageHeader header, It payloadBegin, It payloadEnd)
{
try
{
auto request = ChannelRequest::fromPayload(
std::move(header.ident), std::move(payloadBegin), std::move(payloadEnd));
mChannelsMessageHandler->receiveChannelRequest(request, header.ttl);
}
catch (const std::runtime_error& err)
{
info(mIo->log()) << "Ignoring AudioRequest message: " << err.what();
}
}
template <typename It>
void receiveChannelStopRequest(v1::MessageHeader header,
It payloadBegin,
It payloadEnd)
{
try
{
auto stopRequest = ChannelStopRequest::fromPayload(
std::move(header.ident), std::move(payloadBegin), std::move(payloadEnd));
mChannelsMessageHandler->receiveChannelRequest(stopRequest, header.ttl);
}
catch (const std::runtime_error& err)
{
info(mIo->log()) << "Ignoring ChannelStopRequest message: " << err.what();
}
}
template <typename It>
void receiveAudioBuffer(v1::MessageHeader header, It payloadBegin, It payloadEnd)
{
try
{
mChannelsMessageHandler->receiveAudioBuffer(payloadBegin, payloadEnd);
}
catch (const std::runtime_error& err)
{
info(mIo->log()) << "Ignoring AudioBuffer message: " << err.what();
}
}
util::Injected<IoContext> mIo;
util::Injected<ChannelsMessageHandler> mChannelsMessageHandler;
SharedInterface mpInterface;
std::vector<Announcement> mAnnouncements;
Timer mTimer;
TimePoint mLastBroadcastTime;
uint8_t mTtl;
uint8_t mTtlRatio;
util::Injected<Observer> mObserver;
std::vector<Receiver> mReceivers;
};
std::shared_ptr<Impl> mpImpl;
};
template <typename IoContext>
using MessengerInterface =
UnicastIpInterface<typename util::Injected<IoContext>::type&, v1::kMaxMessageSize>;
template <typename IoContext>
using MessengerInterface = MessengerInterface<IoContext>;
template <typename ChannelsMessageHandler, typename Observer, typename IoContext>
using Messenger = UdpMessenger<ChannelsMessageHandler,
MessengerInterface<IoContext>,
Observer,
IoContext>;
template <typename ChannelsMessageHandler, typename Observer, typename IoContext>
using MessengerPtr =
std::shared_ptr<Messenger<ChannelsMessageHandler, Observer, IoContext>>;
// Factory function
template <typename ChannelsMessageHandler,
typename Announcement,
typename IoContext,
typename Observer>
MessengerPtr<ChannelsMessageHandler, Observer, IoContext> makeMessengerPtr(
util::Injected<ChannelsMessageHandler> handler,
util::Injected<IoContext> io,
const discovery::IpAddress& addr,
util::Injected<Observer> observer,
Announcement announcement)
{
const uint8_t ttl = 5;
const uint8_t ttlRatio = 20;
auto pIface =
makeSharedUnicastIpInterface<v1::kMaxMessageSize>(util::injectRef(*io), addr);
return std::make_shared<Messenger<ChannelsMessageHandler, Observer, IoContext>>(
std::move(handler),
pIface,
std::move(announcement),
std::move(io),
ttl,
ttlRatio,
std::move(observer));
}
} // namespace link_audio
} // namespace ableton