FreeRDP
Loading...
Searching...
No Matches
pf_channel.c
1
18#include <winpr/assert.h>
19#include <winpr/cast.h>
20
21#include <freerdp/freerdp.h>
22#include <freerdp/server/proxy/proxy_log.h>
23
24#include "proxy_modules.h"
25#include "pf_channel.h"
26#include "pf_client.h"
27#include "pf_server.h"
28
29#define TAG PROXY_TAG("channel")
30
32struct sChannelStateTracker
33{
34 pServerStaticChannelContext* channel;
35 ChannelTrackerMode mode;
36 wStream* currentPacket;
37 size_t currentPacketReceived;
38 size_t currentPacketSize;
39 size_t currentPacketFragments;
40
41 ChannelTrackerPeekFn peekFn;
42 void* trackerData;
43 proxyData* pdata;
44};
45
46WINPR_ATTR_NODISCARD
47static BOOL channelTracker_resetCurrentPacket(ChannelStateTracker* tracker)
48{
49 WINPR_ASSERT(tracker);
50
51 BOOL create = TRUE;
52 if (tracker->currentPacket)
53 {
54 const size_t cap = Stream_Capacity(tracker->currentPacket);
55 if (cap < 1ULL * 1000ULL * 1000ULL)
56 create = FALSE;
57 else
58 Stream_Free(tracker->currentPacket, TRUE);
59 }
60
61 if (create)
62 tracker->currentPacket = Stream_New(nullptr, 10ULL * 1024ULL);
63 if (!tracker->currentPacket)
64 return FALSE;
65 Stream_ResetPosition(tracker->currentPacket);
66 return TRUE;
67}
68
69ChannelStateTracker* channelTracker_new(pServerStaticChannelContext* channel,
70 ChannelTrackerPeekFn fn, void* data)
71{
72 ChannelStateTracker* ret = calloc(1, sizeof(ChannelStateTracker));
73 if (!ret)
74 return ret;
75
76 WINPR_ASSERT(fn);
77
78 ret->channel = channel;
79 ret->peekFn = fn;
80
81 if (!channelTracker_setCustomData(ret, data))
82 goto fail;
83
84 if (!channelTracker_resetCurrentPacket(ret))
85 goto fail;
86
87 return ret;
88
89fail:
90 WINPR_PRAGMA_DIAG_PUSH
91 WINPR_PRAGMA_DIAG_IGNORED_MISMATCHED_DEALLOC
92 channelTracker_free(ret);
93 WINPR_PRAGMA_DIAG_POP
94 return nullptr;
95}
96
97PfChannelResult channelTracker_update(ChannelStateTracker* tracker, const BYTE* xdata, size_t xsize,
98 UINT32 flags, size_t totalSize)
99{
100 PfChannelResult result = PF_CHANNEL_RESULT_ERROR;
101 BOOL firstPacket = (flags & CHANNEL_FLAG_FIRST) != 0;
102 BOOL lastPacket = (flags & CHANNEL_FLAG_LAST) != 0;
103
104 WINPR_ASSERT(tracker);
105
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)
109 {
110 if (!channelTracker_resetCurrentPacket(tracker))
111 return PF_CHANNEL_RESULT_ERROR;
112 channelTracker_setCurrentPacketSize(tracker, totalSize);
113 tracker->currentPacketReceived = 0;
114 tracker->currentPacketFragments = 0;
115 }
116
117 {
118 const size_t currentPacketSize = channelTracker_getCurrentPacketSize(tracker);
119 if ((xsize > currentPacketSize) || (currentPacketSize != totalSize))
120 {
121 WLog_WARN(TAG,
122 "current fragment size is bigger (%" PRIuz ") than total size (%" PRIuz ")",
123 xsize, currentPacketSize);
124 return PF_CHANNEL_RESULT_ERROR;
125 }
126
127 if (xsize > SIZE_MAX - tracker->currentPacketReceived)
128 {
129 WLog_WARN(TAG,
130 "current fragment size is overflowing size_t (%" PRIuz
131 ") when added to (%" PRIuz ")",
132 xsize, tracker->currentPacketReceived);
133 return PF_CHANNEL_RESULT_ERROR;
134 }
135
136 tracker->currentPacketReceived += xsize;
137 tracker->currentPacketFragments++;
138 if (tracker->currentPacketReceived > currentPacketSize)
139 {
140 WLog_WARN(TAG, "cumulated size is bigger (%" PRIuz ") than total size (%" PRIuz ")",
141 tracker->currentPacketReceived + xsize, currentPacketSize);
142 return PF_CHANNEL_RESULT_ERROR;
143 }
144 else if ((tracker->currentPacketReceived == currentPacketSize) && !lastPacket)
145 return PF_CHANNEL_RESULT_ERROR;
146 }
147
148 switch (channelTracker_getMode(tracker))
149 {
150 case CHANNEL_TRACKER_PEEK:
151 {
152 wStream* currentPacket = channelTracker_getCurrentPacket(tracker);
153 if (!Stream_EnsureRemainingCapacity(currentPacket, xsize))
154 return PF_CHANNEL_RESULT_ERROR;
155
156 Stream_Write(currentPacket, xdata, xsize);
157
158 WINPR_ASSERT(tracker->peekFn);
159 result = tracker->peekFn(tracker, firstPacket, lastPacket);
160 }
161 break;
162 case CHANNEL_TRACKER_PASS:
163 result = PF_CHANNEL_RESULT_PASS;
164 break;
165 case CHANNEL_TRACKER_DROP:
166 result = PF_CHANNEL_RESULT_DROP;
167 break;
168 default:
169 break;
170 }
171
172 if (lastPacket)
173 {
174 const size_t currentPacketSize = channelTracker_getCurrentPacketSize(tracker);
175 channelTracker_setMode(tracker, CHANNEL_TRACKER_PEEK);
176
177 if (tracker->currentPacketReceived != currentPacketSize)
178 WLog_INFO(TAG, "cumulated size(%" PRIuz ") does not match total size (%" PRIuz ")",
179 tracker->currentPacketReceived, currentPacketSize);
180 }
181
182 return result;
183}
184
185void channelTracker_free(ChannelStateTracker* t)
186{
187 if (!t)
188 return;
189
190 Stream_Free(t->currentPacket, TRUE);
191 free(t);
192}
193
200PfChannelResult channelTracker_flushCurrent(ChannelStateTracker* t, BOOL first, BOOL last,
201 BOOL toBack)
202{
203 UINT32 flags = CHANNEL_FLAG_FIRST;
204 BOOL r = 0;
205 const char* direction = toBack ? "F->B" : "B->F";
206 const size_t currentPacketSize = channelTracker_getCurrentPacketSize(t);
207 wStream* currentPacket = channelTracker_getCurrentPacket(t);
208
209 WINPR_ASSERT(t);
210
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);
213
214 if (first)
215 return PF_CHANNEL_RESULT_PASS;
216
217 proxyData* pdata = t->pdata;
218 pServerStaticChannelContext* channel = t->channel;
219 if (last)
220 flags |= CHANNEL_FLAG_LAST;
221
222 if (toBack)
223 {
224 proxyChannelDataEventInfo ev = WINPR_C_ARRAY_INIT;
225
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);
230 ev.flags = flags;
231 ev.total_size = currentPacketSize;
232
233 pClientContext* pc = proxy_data_get_client_context(pdata);
234 if (!pc->sendChannelData)
235 return PF_CHANNEL_RESULT_ERROR;
236
237 return pc->sendChannelData(pc, &ev) ? PF_CHANNEL_RESULT_DROP : PF_CHANNEL_RESULT_ERROR;
238 }
239
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));
244
245 return r ? PF_CHANNEL_RESULT_DROP : PF_CHANNEL_RESULT_ERROR;
246}
247
248WINPR_ATTR_NODISCARD
249static PfChannelResult pf_channel_generic_back_data(proxyData* pdata,
250 const pServerStaticChannelContext* channel,
251 const BYTE* xdata, size_t xsize, UINT32 flags,
252 size_t totalSize)
253{
254 proxyChannelDataEventInfo ev = WINPR_C_ARRAY_INIT;
255
256 WINPR_ASSERT(pdata);
257 WINPR_ASSERT(channel);
258
259 switch (channel->channelMode)
260 {
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;
264 ev.data = xdata;
265 ev.data_len = xsize;
266 ev.flags = flags;
267 ev.total_size = totalSize;
268
269 if (!pf_modules_run_filter(pdata->module, FILTER_TYPE_CLIENT_PASSTHROUGH_CHANNEL_DATA,
270 pdata, &ev))
271 return PF_CHANNEL_RESULT_DROP; /* Silently drop */
272
273 return PF_CHANNEL_RESULT_PASS;
274
275 case PF_UTILS_CHANNEL_INTERCEPT:
276 /* TODO */
277 case PF_UTILS_CHANNEL_BLOCK:
278 default:
279 return PF_CHANNEL_RESULT_DROP;
280 }
281}
282
283WINPR_ATTR_NODISCARD
284static PfChannelResult pf_channel_generic_front_data(proxyData* pdata,
285 const pServerStaticChannelContext* channel,
286 const BYTE* xdata, size_t xsize, UINT32 flags,
287 size_t totalSize)
288{
289 proxyChannelDataEventInfo ev = WINPR_C_ARRAY_INIT;
290
291 WINPR_ASSERT(pdata);
292 WINPR_ASSERT(channel);
293
294 switch (channel->channelMode)
295 {
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;
299 ev.data = xdata;
300 ev.data_len = xsize;
301 ev.flags = flags;
302 ev.total_size = totalSize;
303
304 if (!pf_modules_run_filter(pdata->module, FILTER_TYPE_SERVER_PASSTHROUGH_CHANNEL_DATA,
305 pdata, &ev))
306 return PF_CHANNEL_RESULT_DROP; /* Silently drop */
307
308 return PF_CHANNEL_RESULT_PASS;
309
310 case PF_UTILS_CHANNEL_INTERCEPT:
311 /* TODO */
312 case PF_UTILS_CHANNEL_BLOCK:
313 default:
314 return PF_CHANNEL_RESULT_DROP;
315 }
316}
317
318BOOL pf_channel_setup_generic(pServerStaticChannelContext* channel)
319{
320 WINPR_ASSERT(channel);
321 channel->onBackData = pf_channel_generic_back_data;
322 channel->onFrontData = pf_channel_generic_front_data;
323 return TRUE;
324}
325
326BOOL channelTracker_setMode(ChannelStateTracker* tracker, ChannelTrackerMode mode)
327{
328 WINPR_ASSERT(tracker);
329 tracker->mode = mode;
330 return TRUE;
331}
332
333ChannelTrackerMode channelTracker_getMode(ChannelStateTracker* tracker)
334{
335 WINPR_ASSERT(tracker);
336 return tracker->mode;
337}
338
339BOOL channelTracker_setPData(ChannelStateTracker* tracker, proxyData* pdata)
340{
341 WINPR_ASSERT(tracker);
342 tracker->pdata = pdata;
343 return TRUE;
344}
345
346proxyData* channelTracker_getPData(ChannelStateTracker* tracker)
347{
348 WINPR_ASSERT(tracker);
349 return tracker->pdata;
350}
351
352wStream* channelTracker_getCurrentPacket(ChannelStateTracker* tracker)
353{
354 WINPR_ASSERT(tracker);
355 return tracker->currentPacket;
356}
357
358BOOL channelTracker_setCustomData(ChannelStateTracker* tracker, void* data)
359{
360 WINPR_ASSERT(tracker);
361 tracker->trackerData = data;
362 return TRUE;
363}
364
365void* channelTracker_getCustomData(ChannelStateTracker* tracker)
366{
367 WINPR_ASSERT(tracker);
368 return tracker->trackerData;
369}
370
371size_t channelTracker_getCurrentPacketSize(ChannelStateTracker* tracker)
372{
373 WINPR_ASSERT(tracker);
374 return tracker->currentPacketSize;
375}
376
377BOOL channelTracker_setCurrentPacketSize(ChannelStateTracker* tracker, size_t size)
378{
379 WINPR_ASSERT(tracker);
380 tracker->currentPacketSize = size;
381 return TRUE;
382}