18#include <winpr/assert.h>
19#include <winpr/cast.h>
21#include <freerdp/freerdp.h>
22#include <freerdp/server/proxy/proxy_log.h>
24#include "proxy_modules.h"
25#include "pf_channel.h"
29#define TAG PROXY_TAG("channel")
32struct sChannelStateTracker
34 pServerStaticChannelContext* channel;
35 ChannelTrackerMode mode;
37 size_t currentPacketReceived;
38 size_t currentPacketSize;
39 size_t currentPacketFragments;
41 ChannelTrackerPeekFn peekFn;
47static BOOL channelTracker_resetCurrentPacket(ChannelStateTracker* tracker)
49 WINPR_ASSERT(tracker);
52 if (tracker->currentPacket)
54 const size_t cap = Stream_Capacity(tracker->currentPacket);
55 if (cap < 1ULL * 1000ULL * 1000ULL)
58 Stream_Free(tracker->currentPacket, TRUE);
62 tracker->currentPacket = Stream_New(
nullptr, 10ULL * 1024ULL);
63 if (!tracker->currentPacket)
65 Stream_ResetPosition(tracker->currentPacket);
69ChannelStateTracker* channelTracker_new(pServerStaticChannelContext* channel,
70 ChannelTrackerPeekFn fn,
void* data)
72 ChannelStateTracker* ret = calloc(1,
sizeof(ChannelStateTracker));
78 ret->channel = channel;
81 if (!channelTracker_setCustomData(ret, data))
84 if (!channelTracker_resetCurrentPacket(ret))
90 WINPR_PRAGMA_DIAG_PUSH
91 WINPR_PRAGMA_DIAG_IGNORED_MISMATCHED_DEALLOC
92 channelTracker_free(ret);
97PfChannelResult channelTracker_update(ChannelStateTracker* tracker,
const BYTE* xdata,
size_t xsize,
98 UINT32 flags,
size_t totalSize)
100 PfChannelResult result = PF_CHANNEL_RESULT_ERROR;
101 BOOL firstPacket = (flags & CHANNEL_FLAG_FIRST) != 0;
102 BOOL lastPacket = (flags & CHANNEL_FLAG_LAST) != 0;
104 WINPR_ASSERT(tracker);
106 WLog_VRB(TAG,
"channelTracker_update(%s): sz=%" PRIuz
" first=%d last=%d",
107 tracker->channel->channel_name, xsize, firstPacket, lastPacket);
108 if (flags & CHANNEL_FLAG_FIRST)
110 if (!channelTracker_resetCurrentPacket(tracker))
111 return PF_CHANNEL_RESULT_ERROR;
112 channelTracker_setCurrentPacketSize(tracker, totalSize);
113 tracker->currentPacketReceived = 0;
114 tracker->currentPacketFragments = 0;
118 const size_t currentPacketSize = channelTracker_getCurrentPacketSize(tracker);
119 if ((xsize > currentPacketSize) || (currentPacketSize != totalSize))
122 "current fragment size is bigger (%" PRIuz
") than total size (%" PRIuz
")",
123 xsize, currentPacketSize);
124 return PF_CHANNEL_RESULT_ERROR;
127 if (xsize > SIZE_MAX - tracker->currentPacketReceived)
130 "current fragment size is overflowing size_t (%" PRIuz
131 ") when added to (%" PRIuz
")",
132 xsize, tracker->currentPacketReceived);
133 return PF_CHANNEL_RESULT_ERROR;
136 tracker->currentPacketReceived += xsize;
137 tracker->currentPacketFragments++;
138 if (tracker->currentPacketReceived > currentPacketSize)
140 WLog_WARN(TAG,
"cumulated size is bigger (%" PRIuz
") than total size (%" PRIuz
")",
141 tracker->currentPacketReceived + xsize, currentPacketSize);
142 return PF_CHANNEL_RESULT_ERROR;
144 else if ((tracker->currentPacketReceived == currentPacketSize) && !lastPacket)
145 return PF_CHANNEL_RESULT_ERROR;
148 switch (channelTracker_getMode(tracker))
150 case CHANNEL_TRACKER_PEEK:
152 wStream* currentPacket = channelTracker_getCurrentPacket(tracker);
153 if (!Stream_EnsureRemainingCapacity(currentPacket, xsize))
154 return PF_CHANNEL_RESULT_ERROR;
156 Stream_Write(currentPacket, xdata, xsize);
158 WINPR_ASSERT(tracker->peekFn);
159 result = tracker->peekFn(tracker, firstPacket, lastPacket);
162 case CHANNEL_TRACKER_PASS:
163 result = PF_CHANNEL_RESULT_PASS;
165 case CHANNEL_TRACKER_DROP:
166 result = PF_CHANNEL_RESULT_DROP;
174 const size_t currentPacketSize = channelTracker_getCurrentPacketSize(tracker);
175 channelTracker_setMode(tracker, CHANNEL_TRACKER_PEEK);
177 if (tracker->currentPacketReceived != currentPacketSize)
178 WLog_INFO(TAG,
"cumulated size(%" PRIuz
") does not match total size (%" PRIuz
")",
179 tracker->currentPacketReceived, currentPacketSize);
185void channelTracker_free(ChannelStateTracker* t)
190 Stream_Free(t->currentPacket, TRUE);
200PfChannelResult channelTracker_flushCurrent(ChannelStateTracker* t, BOOL first, BOOL last,
203 UINT32 flags = CHANNEL_FLAG_FIRST;
205 const char* direction = toBack ?
"F->B" :
"B->F";
206 const size_t currentPacketSize = channelTracker_getCurrentPacketSize(t);
207 wStream* currentPacket = channelTracker_getCurrentPacket(t);
211 WLog_VRB(TAG,
"channelTracker_flushCurrent(%s): %s sz=%" PRIuz
" first=%d last=%d",
212 t->channel->channel_name, direction, Stream_GetPosition(currentPacket), first, last);
215 return PF_CHANNEL_RESULT_PASS;
217 proxyData* pdata = t->pdata;
218 pServerStaticChannelContext* channel = t->channel;
220 flags |= CHANNEL_FLAG_LAST;
226 ev.channel_id = WINPR_ASSERTING_INT_CAST(UINT16, channel->front_channel_id);
227 ev.channel_name = channel->channel_name;
228 ev.data = Stream_Buffer(currentPacket);
229 ev.data_len = Stream_GetPosition(currentPacket);
231 ev.total_size = currentPacketSize;
233 pClientContext* pc = proxy_data_get_client_context(pdata);
234 if (!pc->sendChannelData)
235 return PF_CHANNEL_RESULT_ERROR;
237 return pc->sendChannelData(pc, &ev) ? PF_CHANNEL_RESULT_DROP : PF_CHANNEL_RESULT_ERROR;
240 pServerContext* ps = proxy_data_get_server_context(pdata);
241 r = ps->context.peer->SendChannelPacket(
242 ps->context.peer, WINPR_ASSERTING_INT_CAST(UINT16, channel->front_channel_id),
243 currentPacketSize, flags, Stream_Buffer(currentPacket), Stream_GetPosition(currentPacket));
245 return r ? PF_CHANNEL_RESULT_DROP : PF_CHANNEL_RESULT_ERROR;
249static PfChannelResult pf_channel_generic_back_data(proxyData* pdata,
250 const pServerStaticChannelContext* channel,
251 const BYTE* xdata,
size_t xsize, UINT32 flags,
257 WINPR_ASSERT(channel);
259 switch (channel->channelMode)
261 case PF_UTILS_CHANNEL_PASSTHROUGH:
262 ev.channel_id = WINPR_ASSERTING_INT_CAST(UINT16, channel->back_channel_id);
263 ev.channel_name = channel->channel_name;
267 ev.total_size = totalSize;
269 if (!pf_modules_run_filter(pdata->module, FILTER_TYPE_CLIENT_PASSTHROUGH_CHANNEL_DATA,
271 return PF_CHANNEL_RESULT_DROP;
273 return PF_CHANNEL_RESULT_PASS;
275 case PF_UTILS_CHANNEL_INTERCEPT:
277 case PF_UTILS_CHANNEL_BLOCK:
279 return PF_CHANNEL_RESULT_DROP;
284static PfChannelResult pf_channel_generic_front_data(proxyData* pdata,
285 const pServerStaticChannelContext* channel,
286 const BYTE* xdata,
size_t xsize, UINT32 flags,
292 WINPR_ASSERT(channel);
294 switch (channel->channelMode)
296 case PF_UTILS_CHANNEL_PASSTHROUGH:
297 ev.channel_id = WINPR_ASSERTING_INT_CAST(UINT16, channel->front_channel_id);
298 ev.channel_name = channel->channel_name;
302 ev.total_size = totalSize;
304 if (!pf_modules_run_filter(pdata->module, FILTER_TYPE_SERVER_PASSTHROUGH_CHANNEL_DATA,
306 return PF_CHANNEL_RESULT_DROP;
308 return PF_CHANNEL_RESULT_PASS;
310 case PF_UTILS_CHANNEL_INTERCEPT:
312 case PF_UTILS_CHANNEL_BLOCK:
314 return PF_CHANNEL_RESULT_DROP;
318BOOL pf_channel_setup_generic(pServerStaticChannelContext* channel)
320 WINPR_ASSERT(channel);
321 channel->onBackData = pf_channel_generic_back_data;
322 channel->onFrontData = pf_channel_generic_front_data;
326BOOL channelTracker_setMode(ChannelStateTracker* tracker, ChannelTrackerMode mode)
328 WINPR_ASSERT(tracker);
329 tracker->mode = mode;
333ChannelTrackerMode channelTracker_getMode(ChannelStateTracker* tracker)
335 WINPR_ASSERT(tracker);
336 return tracker->mode;
339BOOL channelTracker_setPData(ChannelStateTracker* tracker, proxyData* pdata)
341 WINPR_ASSERT(tracker);
342 tracker->pdata = pdata;
346proxyData* channelTracker_getPData(ChannelStateTracker* tracker)
348 WINPR_ASSERT(tracker);
349 return tracker->pdata;
352wStream* channelTracker_getCurrentPacket(ChannelStateTracker* tracker)
354 WINPR_ASSERT(tracker);
355 return tracker->currentPacket;
358BOOL channelTracker_setCustomData(ChannelStateTracker* tracker,
void* data)
360 WINPR_ASSERT(tracker);
361 tracker->trackerData = data;
365void* channelTracker_getCustomData(ChannelStateTracker* tracker)
367 WINPR_ASSERT(tracker);
368 return tracker->trackerData;
371size_t channelTracker_getCurrentPacketSize(ChannelStateTracker* tracker)
373 WINPR_ASSERT(tracker);
374 return tracker->currentPacketSize;
377BOOL channelTracker_setCurrentPacketSize(ChannelStateTracker* tracker,
size_t size)
379 WINPR_ASSERT(tracker);
380 tracker->currentPacketSize = size;