feat: add integration tests and refactor connection request handling in `handlers.go` for improved modularity.

lemon-mint committed Nov 19, 2025 at 09:34 UTC b321894341d46dbe21aef6d8c0e5dd1fc4262ab7
2 files changed +204 -206
portal/handlers.go
+94 -206
@@ -164,276 +164,164 @@ func (g *RelayServer) handleConnectionRequest(ctx *StreamContext, packet *rdverb
164 Int64("conn_id", ctx.ConnectionID).
165 Msg("[RelayServer] Handling connection request")
166
167 - var resp rdverb.ConnectionResponse
168 -
169 - // Check if lease exists and get lease connection using LeaseId
167 + // Check if lease exists and get lease connection
168 leaseEntry, exists := g.leaseManager.GetLeaseByID(req.LeaseId)
169 if !exists {
172 - log.Warn().Str("lease_id", req.LeaseId).Msg("[RelayServer] Lease not found")
173 - resp.Code = rdverb.ResponseCode_RESPONSE_CODE_INVALID_IDENTITY
174 -
175 - response, err := resp.MarshalVT()
176 - if err != nil {
177 - log.Error().Err(err).Msg("[RelayServer] Failed to marshal connection response")
178 - return err
179 - }
180 -
181 - return writePacket(ctx.Stream, &rdverb.Packet{
182 - Type: rdverb.PacketType_PACKET_TYPE_CONNECTION_RESPONSE,
183 - Payload: response,
184 - })
170 + return g.sendConnectionResponse(ctx.Stream, rdverb.ResponseCode_RESPONSE_CODE_INVALID_IDENTITY)
171 }
172
187 - log.Debug().
188 - Str("lease_id", req.LeaseId).
189 - Int64("lease_conn_id", leaseEntry.ConnectionID).
190 - Msg("[RelayServer] Lease found, forwarding to lease holder")
191 -
192 - // Get the lease connection using connection ID
173 + // Get the lease connection
174 g.connectionsLock.RLock()
175 leaseConn, leaseExists := g.connections[leaseEntry.ConnectionID]
176 g.connectionsLock.RUnlock()
177
178 if !leaseExists {
198 - log.Warn().
199 - Str("lease_id", req.LeaseId).
200 - Int64("lease_conn_id", leaseEntry.ConnectionID).
201 - Msg("[RelayServer] Lease connection no longer active")
202 - resp.Code = rdverb.ResponseCode_RESPONSE_CODE_INVALID_IDENTITY
179 + return g.sendConnectionResponse(ctx.Stream, rdverb.ResponseCode_RESPONSE_CODE_INVALID_IDENTITY)
180 + }
181
204 - response, err := resp.MarshalVT()
205 - if err != nil {
206 - log.Error().Err(err).Msg("[RelayServer] Failed to marshal connection response")
207 - return err
182 + // Forward request to lease holder
183 + leaseStream, respCode, err := g.forwardConnectionRequest(leaseConn, &req)
184 + if err != nil {
185 + // If forwarding failed, we might need to close the stream if it was opened
186 + if leaseStream != nil {
187 + leaseStream.Close()
188 }
209 -
210 - return writePacket(ctx.Stream, &rdverb.Packet{
211 - Type: rdverb.PacketType_PACKET_TYPE_CONNECTION_RESPONSE,
212 - Payload: response,
213 - })
189 + return g.sendConnectionResponse(ctx.Stream, respCode)
190 }
191
216 - // Open a stream to the lease holder
217 - log.Debug().Str("lease_id", req.LeaseId).Msg("[RelayServer] Opening stream to lease holder")
218 - leaseStream, err := leaseConn.sess.OpenStream()
219 - if err != nil {
220 - log.Error().Err(err).Str("lease_id", req.LeaseId).Msg("[RelayServer] Failed to open stream to lease holder")
221 - resp.Code = rdverb.ResponseCode_RESPONSE_CODE_REJECTED
192 + // Enforce relayed connection limits
193 + if respCode == rdverb.ResponseCode_RESPONSE_CODE_ACCEPTED {
194 + leaseID := string(leaseEntry.Lease.Identity.Id)
195 + g.limitsLock.Lock()
196 + overPerLease := g.maxRelayedPerLease > 0 && g.relayedPerLeaseCount[leaseID] >= g.maxRelayedPerLease
197 + g.limitsLock.Unlock()
198
223 - response, err := resp.MarshalVT()
224 - if err != nil {
225 - return err
199 + if overPerLease {
200 + log.Warn().Str("lease_id", leaseID).Msg("[RelayServer] Relayed connection per-lease limit reached")
201 + respCode = rdverb.ResponseCode_RESPONSE_CODE_REJECTED
202 + leaseStream.Close()
203 }
204 + }
205
228 - return writePacket(ctx.Stream, &rdverb.Packet{
229 - Type: rdverb.PacketType_PACKET_TYPE_CONNECTION_RESPONSE,
230 - Payload: response,
231 - })
206 + // Send response to client
207 + if err := g.sendConnectionResponse(ctx.Stream, respCode); err != nil {
208 + leaseStream.Close()
209 + return err
210 }
211
234 - // Forward the connection request
235 - requestPayload, err := req.MarshalVT()
236 - if err != nil {
237 - log.Error().Err(err).Msg("[RelayServer] Failed to marshal forward request")
212 + // If accepted, set up bidirectional forwarding
213 + if respCode == rdverb.ResponseCode_RESPONSE_CODE_ACCEPTED {
214 + ctx.Hijack()
215 + go g.establishRelayedConnection(ctx.Stream, leaseStream, string(leaseEntry.Lease.Identity.Id))
216 + } else {
217 leaseStream.Close()
239 - resp.Code = rdverb.ResponseCode_RESPONSE_CODE_REJECTED
218 + }
219
241 - response, err := resp.MarshalVT()
242 - if err != nil {
243 - return err
244 - }
220 + return nil
221 +}
222
246 - return writePacket(ctx.Stream, &rdverb.Packet{
247 - Type: rdverb.PacketType_PACKET_TYPE_CONNECTION_RESPONSE,
248 - Payload: response,
249 - })
223 +// forwardConnectionRequest opens a stream to the lease holder and forwards the request
224 +func (g *RelayServer) forwardConnectionRequest(leaseConn *Connection, req *rdverb.ConnectionRequest) (*yamux.Stream, rdverb.ResponseCode, error) {
225 + leaseStream, err := leaseConn.sess.OpenStream()
226 + if err != nil {
227 + return nil, rdverb.ResponseCode_RESPONSE_CODE_REJECTED, err
228 + }
229 +
230 + reqPayload, err := req.MarshalVT()
231 + if err != nil {
232 + leaseStream.Close()
233 + return nil, rdverb.ResponseCode_RESPONSE_CODE_REJECTED, err
234 }
235
252 - log.Debug().Str("lease_id", req.LeaseId).Msg("[RelayServer] Sending connection request to lease holder")
236 err = writePacket(leaseStream, &rdverb.Packet{
237 Type: rdverb.PacketType_PACKET_TYPE_CONNECTION_REQUEST,
255 - Payload: requestPayload,
238 + Payload: reqPayload,
239 })
240 if err != nil {
258 - log.Error().Err(err).Msg("[RelayServer] Failed to write forward request")
241 leaseStream.Close()
260 - resp.Code = rdverb.ResponseCode_RESPONSE_CODE_REJECTED
261 -
262 - response, err := resp.MarshalVT()
263 - if err != nil {
264 - return err
265 - }
266 -
267 - return writePacket(ctx.Stream, &rdverb.Packet{
268 - Type: rdverb.PacketType_PACKET_TYPE_CONNECTION_RESPONSE,
269 - Payload: response,
270 - })
242 + return nil, rdverb.ResponseCode_RESPONSE_CODE_REJECTED, err
243 }
244
273 - // Read the response
274 - log.Debug().Str("lease_id", req.LeaseId).Msg("[RelayServer] Waiting for response from lease holder")
245 respPacket, err := readPacket(leaseStream)
246 if err != nil {
277 - log.Error().Str("lease_id", req.LeaseId).Err(err).Msg("[RelayServer] Failed to read forward response")
247 leaseStream.Close()
279 - resp.Code = rdverb.ResponseCode_RESPONSE_CODE_REJECTED
280 -
281 - response, err := resp.MarshalVT()
282 - if err != nil {
283 - return err
284 - }
285 -
286 - return writePacket(ctx.Stream, &rdverb.Packet{
287 - Type: rdverb.PacketType_PACKET_TYPE_CONNECTION_RESPONSE,
288 - Payload: response,
289 - })
248 + return nil, rdverb.ResponseCode_RESPONSE_CODE_REJECTED, err
249 }
250
251 if respPacket.Type != rdverb.PacketType_PACKET_TYPE_CONNECTION_RESPONSE {
293 - log.Warn().Str("packet_type", respPacket.Type.String()).Msg("[RelayServer] Unexpected response packet type")
252 leaseStream.Close()
295 - resp.Code = rdverb.ResponseCode_RESPONSE_CODE_REJECTED
296 -
297 - response, err := resp.MarshalVT()
298 - if err != nil {
299 - return err
300 - }
301 -
302 - return writePacket(ctx.Stream, &rdverb.Packet{
303 - Type: rdverb.PacketType_PACKET_TYPE_CONNECTION_RESPONSE,
304 - Payload: response,
305 - })
253 + return nil, rdverb.ResponseCode_RESPONSE_CODE_REJECTED, nil
254 }
255
308 - err = resp.UnmarshalVT(respPacket.Payload)
309 - if err != nil {
310 - log.Error().Err(err).Msg("[RelayServer] Failed to unmarshal forward response")
256 + var resp rdverb.ConnectionResponse
257 + if err := resp.UnmarshalVT(respPacket.Payload); err != nil {
258 leaseStream.Close()
312 - resp.Code = rdverb.ResponseCode_RESPONSE_CODE_REJECTED
313 -
314 - response, err := resp.MarshalVT()
315 - if err != nil {
316 - return err
317 - }
318 -
319 - return writePacket(ctx.Stream, &rdverb.Packet{
320 - Type: rdverb.PacketType_PACKET_TYPE_CONNECTION_RESPONSE,
321 - Payload: response,
322 - })
259 + return nil, rdverb.ResponseCode_RESPONSE_CODE_REJECTED, err
260 }
261
325 - log.Debug().
326 - Str("lease_id", req.LeaseId).
327 - Str("response_code", resp.Code.String()).
328 - Msg("[RelayServer] Received response from lease holder, sending to client")
329 -
330 - // Enforce relayed connection limits if currently accepted
331 - if resp.Code == rdverb.ResponseCode_RESPONSE_CODE_ACCEPTED {
332 - leaseID := string(leaseEntry.Lease.Identity.Id)
333 - g.limitsLock.Lock()
334 - overPerLease := g.maxRelayedPerLease > 0 && g.relayedPerLeaseCount[leaseID] >= g.maxRelayedPerLease
335 - if overPerLease {
336 - log.Warn().
337 - Str("lease_id", leaseID).
338 - Bool("over_per_lease", overPerLease).
339 - Msg("[RelayServer] Relayed connection per-lease limit reached, rejecting")
340 - resp.Code = rdverb.ResponseCode_RESPONSE_CODE_REJECTED
341 - }
342 - g.limitsLock.Unlock()
343 - }
262 + return leaseStream, resp.Code, nil
263 +}
264
345 - // Send response to client
346 - response, err := resp.MarshalVT()
265 +func (g *RelayServer) sendConnectionResponse(stream *yamux.Stream, code rdverb.ResponseCode) error {
266 + resp := rdverb.ConnectionResponse{Code: code}
267 + payload, err := resp.MarshalVT()
268 if err != nil {
348 - log.Error().Err(err).Msg("[RelayServer] Failed to marshal connection response")
349 - leaseStream.Close()
269 return err
270 }
352 -
353 - err = writePacket(ctx.Stream, &rdverb.Packet{
271 + return writePacket(stream, &rdverb.Packet{
272 Type: rdverb.PacketType_PACKET_TYPE_CONNECTION_RESPONSE,
355 - Payload: response,
273 + Payload: payload,
274 })
357 - if err != nil {
358 - log.Error().Err(err).Msg("[RelayServer] Failed to write connection response")
359 - leaseStream.Close()
360 - return err
361 - }
275 +}
276
363 - // If accepted, hijack both streams and set up bidirectional forwarding
364 - if resp.Code == rdverb.ResponseCode_RESPONSE_CODE_ACCEPTED {
365 - log.Debug().Str("lease_id", req.LeaseId).Msg("[RelayServer] Connection accepted, setting up bidirectional forwarding")
366 - ctx.Hijack()
277 +func (g *RelayServer) establishRelayedConnection(clientStream, leaseStream *yamux.Stream, leaseID string) {
278 + // Register connection
279 + g.limitsLock.Lock()
280 + g.relayedPerLeaseCount[leaseID]++
281 + g.limitsLock.Unlock()
282
368 - leaseID := string(leaseEntry.Lease.Identity.Id)
369 - // Increment counters for active relayed connections
370 - g.limitsLock.Lock()
371 - g.relayedPerLeaseCount[leaseID] = g.relayedPerLeaseCount[leaseID] + 1
372 - g.limitsLock.Unlock()
373 - g.relayedConnectionsLock.Lock()
374 - g.relayedConnections[leaseID] = append(g.relayedConnections[leaseID], ctx.Stream)
375 - g.relayedConnectionsLock.Unlock()
283 + g.relayedConnectionsLock.Lock()
284 + g.relayedConnections[leaseID] = append(g.relayedConnections[leaseID], clientStream)
285 + g.relayedConnectionsLock.Unlock()
286
377 - // Set up bidirectional copying
378 - var wg sync.WaitGroup
379 - wg.Add(2)
380 -
381 - // Copy from client to lease holder (with optional per-lease BPS limit)
382 - go func() {
383 - defer wg.Done()
384 - n, err := ratelimit.Copy(leaseStream, ctx.Stream, g.getLeaseBPSBucket(leaseID))
385 - log.Debug().
386 - Str("lease_id", leaseID).
387 - Int64("bytes", n).
388 - Err(err).
389 - Msg("[RelayServer] Client -> Lease copy finished")
390 - leaseStream.Close()
391 - }()
392 -
393 - // Copy from lease holder to client (with optional per-lease BPS limit)
394 - go func() {
395 - defer wg.Done()
396 - n, err := ratelimit.Copy(ctx.Stream, leaseStream, g.getLeaseBPSBucket(leaseID))
397 - log.Debug().
398 - Str("lease_id", leaseID).
399 - Int64("bytes", n).
400 - Err(err).
401 - Msg("[RelayServer] Lease -> Client copy finished")
402 - ctx.Stream.Close()
403 - }()
404 -
405 - wg.Wait()
406 - log.Debug().Str("lease_id", leaseID).Msg("[RelayServer] Connection forwarding completed successfully")
407 -
408 - // Decrement counters after forwarding completes
287 + // Cleanup function
288 + defer func() {
289 g.limitsLock.Lock()
410 - if v := g.relayedPerLeaseCount[leaseID]; v > 1 {
411 - g.relayedPerLeaseCount[leaseID] = v - 1
412 - } else {
413 - delete(g.relayedPerLeaseCount, leaseID)
290 + if g.relayedPerLeaseCount[leaseID] > 0 {
291 + g.relayedPerLeaseCount[leaseID]--
292 }
293 g.limitsLock.Unlock()
294
417 - // Clean up relayed connection tracking
295 g.relayedConnectionsLock.Lock()
419 - if streams, exists := g.relayedConnections[leaseID]; exists {
420 - for i, stream := range streams {
421 - if stream == ctx.Stream {
296 + if streams, ok := g.relayedConnections[leaseID]; ok {
297 + for i, s := range streams {
298 + if s == clientStream {
299 g.relayedConnections[leaseID] = append(streams[:i], streams[i+1:]...)
300 break
301 }
302 }
426 - if len(g.relayedConnections[leaseID]) == 0 {
427 - delete(g.relayedConnections, leaseID)
428 - }
303 }
304 g.relayedConnectionsLock.Unlock()
431 - } else {
432 - // Connection rejected, close lease stream
305 + }()
306 +
307 + var wg sync.WaitGroup
308 + wg.Add(2)
309 +
310 + // Client -> Lease
311 + go func() {
312 + defer wg.Done()
313 + ratelimit.Copy(leaseStream, clientStream, g.getLeaseBPSBucket(leaseID))
314 leaseStream.Close()
434 - }
315 + }()
316
436 - return nil
317 + // Lease -> Client
318 + go func() {
319 + defer wg.Done()
320 + ratelimit.Copy(clientStream, leaseStream, g.getLeaseBPSBucket(leaseID))
321 + clientStream.Close()
322 + }()
323 +
324 + wg.Wait()
325 }
326
327 // Helper function to read packet from stream
portal/integration_test.go new
+110
@@ -0,0 +1,110 @@
1 +package portal
2 +
3 +import (
4 + "io"
5 + "net"
6 + "testing"
7 + "time"
8 +
9 + "github.com/stretchr/testify/assert"
10 + "github.com/stretchr/testify/require"
11 + "gosuda.org/portal/portal/core/cryptoops"
12 + "gosuda.org/portal/portal/core/proto/rdverb"
13 +)
14 +
15 +// generateTestCredential creates a new credential for testing
16 +func generateTestCredential(t *testing.T) *cryptoops.Credential {
17 + cred, err := cryptoops.NewCredential()
18 + require.NoError(t, err)
19 + return cred
20 +}
21 +
22 +func TestIntegration_FullFlow(t *testing.T) {
23 + // 1. Setup Relay Server
24 + serverCred := generateTestCredential(t)
25 + server := NewRelayServer(serverCred, []string{"localhost:8080"})
26 + server.Start()
27 + defer server.Stop()
28 +
29 + // Create a listener for the server
30 + listener, err := net.Listen("tcp", "127.0.0.1:0")
31 + require.NoError(t, err)
32 + defer listener.Close()
33 +
34 + go func() {
35 + for {
36 + conn, err := listener.Accept()
37 + if err != nil {
38 + return
39 + }
40 + go server.HandleConnection(conn)
41 + }
42 + }()
43 +
44 + serverAddr := listener.Addr().String()
45 +
46 + // 2. Setup Host Client (Service Provider)
47 + hostCred := generateTestCredential(t)
48 + hostConn, err := net.Dial("tcp", serverAddr)
49 + require.NoError(t, err)
50 +
51 + hostClient := NewRelayClient(hostConn)
52 + require.NotNil(t, hostClient)
53 + defer hostClient.Close()
54 +
55 + // Register Lease
56 + lease := &rdverb.Lease{
57 + Name: "test-service",
58 + Alpn: []string{"test-proto"},
59 + }
60 + err = hostClient.RegisterLease(hostCred, lease)
61 + require.NoError(t, err)
62 +
63 + // Handle incoming connections on Host
64 + go func() {
65 + for conn := range hostClient.IncomingConnection() {
66 + go func(c *IncomingConn) {
67 + defer c.Close()
68 + // Echo server
69 + io.Copy(c, c)
70 + }(conn)
71 + }
72 + }()
73 +
74 + // 3. Setup Peer Client (Consumer)
75 + peerCred := generateTestCredential(t)
76 + peerConn, err := net.Dial("tcp", serverAddr)
77 + require.NoError(t, err)
78 +
79 + peerClient := NewRelayClient(peerConn)
80 + require.NotNil(t, peerClient)
81 + defer peerClient.Close()
82 +
83 + // 4. Peer connects to Host
84 + code, conn, err := peerClient.RequestConnection(hostCred.ID(), "test-proto", peerCred)
85 + require.NoError(t, err)
86 + assert.Equal(t, rdverb.ResponseCode_RESPONSE_CODE_ACCEPTED, code)
87 + require.NotNil(t, conn)
88 + defer conn.Close()
89 +
90 + // 5. Verify Data Transfer
91 + message := []byte("Hello, Portal!")
92 + _, err = conn.Write(message)
93 + require.NoError(t, err)
94 +
95 + buffer := make([]byte, len(message))
96 + _, err = io.ReadFull(conn, buffer)
97 + require.NoError(t, err)
98 + assert.Equal(t, message, buffer)
99 +
100 + // 6. Verify Lease Cleanup
101 + err = hostClient.DeregisterLease(hostCred)
102 + require.NoError(t, err)
103 +
104 + // Wait a bit for propagation
105 + time.Sleep(100 * time.Millisecond)
106 +
107 + // Connection should fail now
108 + code, _, err = peerClient.RequestConnection(hostCred.ID(), "test-proto", peerCred)
109 + assert.Equal(t, rdverb.ResponseCode_RESPONSE_CODE_INVALID_IDENTITY, code)
110 +}