Skip to content

Commit 35d8091

Browse files
committed
refactor port forward
1 parent adc0fec commit 35d8091

5 files changed

Lines changed: 100 additions & 37 deletions

File tree

pkg/forwarder/forwarder.go

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,5 @@
1+
package forwarder
2+
3+
type Forwarder interface {
4+
NextPort() string
5+
}

pkg/portforward/mock/mock.go

Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,22 @@
1+
package mock
2+
3+
import "github.com/revolyssup/k8sdebug/pkg/forwarder"
4+
5+
type MockForwarder struct {
6+
ports []string
7+
index int
8+
}
9+
10+
// NewMock creates a mock forwarder with predefined ports.
11+
func New(ports ...string) forwarder.Forwarder {
12+
return &MockForwarder{
13+
ports: ports,
14+
}
15+
}
16+
17+
// Port returns the next predefined port or an error if exhausted.
18+
func (m *MockForwarder) NextPort() string {
19+
port := m.ports[m.index]
20+
m.index = (m.index + 1) % len(m.ports)
21+
return port
22+
}

pkg/portforward/mock/mock_test.go

Lines changed: 19 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,19 @@
1+
package mock_test
2+
3+
import (
4+
"testing"
5+
6+
"github.com/revolyssup/k8sdebug/pkg/portforward/mock"
7+
"github.com/stretchr/testify/assert"
8+
)
9+
10+
func TestRoundRobin(t *testing.T) {
11+
mock := mock.New("8080", "8081", "8082")
12+
13+
assert.Equal(t, "8080", mock.NextPort())
14+
assert.Equal(t, "8081", mock.NextPort())
15+
assert.Equal(t, "8082", mock.NextPort())
16+
assert.Equal(t, "8080", mock.NextPort())
17+
assert.Equal(t, "8081", mock.NextPort())
18+
assert.Equal(t, "8082", mock.NextPort())
19+
}

pkg/portforward/portforward.go

Lines changed: 9 additions & 37 deletions
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,9 @@ import (
1414
"time"
1515

1616
"github.com/revolyssup/k8sdebug/pkg"
17+
"github.com/revolyssup/k8sdebug/pkg/forwarder"
18+
"github.com/revolyssup/k8sdebug/pkg/portforward/mock"
19+
"github.com/revolyssup/k8sdebug/pkg/portforward/roundrobin"
1720
"github.com/spf13/cobra"
1821
v1 "k8s.io/api/core/v1"
1922
metav1 "k8s.io/apimachinery/pkg/apis/meta/v1"
@@ -45,52 +48,21 @@ func forwardToPod(hostConn net.Conn, podCon net.Conn) {
4548
io.Copy(podCon, hostConn)
4649
}
4750

48-
func getPodConnection(fw forwarder) (net.Conn, error) {
49-
port := fw.Port()
51+
func getPodConnection(fw forwarder.Forwarder) (net.Conn, error) {
52+
port := fw.NextPort()
5053
podConn, err := net.Dial("tcp", fmt.Sprintf(":%s", port))
5154
if err != nil {
5255
return nil, err
5356
}
5457
return podConn, nil
5558
}
5659

57-
type forwarder interface {
58-
Port() string
59-
}
60-
61-
type roundRobin struct {
62-
connNumber int
63-
mx sync.Mutex
64-
}
65-
66-
func (rr *roundRobin) Port() string {
67-
rr.mx.Lock()
68-
defer rr.mx.Unlock()
69-
70-
initial := rr.connNumber
71-
for {
72-
rr.connNumber = (rr.connNumber + 1) % len(connPool)
73-
portNum := connPool[rr.connNumber]
74-
if portNum != "" {
75-
// Check if port is actually listening
76-
conn, err := net.DialTimeout("tcp", ":"+portNum, 50*time.Millisecond)
77-
if err == nil {
78-
conn.Close()
79-
fmt.Println("PORT RETURNED ", portNum)
80-
return portNum
81-
}
82-
}
83-
if rr.connNumber == initial {
84-
break // Avoid infinite loop
85-
}
86-
}
87-
return ""
88-
}
89-
90-
func getForwarder(policy string) forwarder {
60+
func getForwarder(policy string) forwarder.Forwarder {
9161
switch policy {
9262
case "round-robin":
93-
return &roundRobin{}
63+
return roundrobin.New(connPool)
64+
case "mock":
65+
return mock.New()
9466
}
9567
return nil
9668
}
Lines changed: 45 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,45 @@
1+
package roundrobin
2+
3+
import (
4+
"fmt"
5+
"net"
6+
"sync"
7+
"time"
8+
9+
"github.com/revolyssup/k8sdebug/pkg/forwarder"
10+
)
11+
12+
type RoundRobin struct {
13+
connNumber int
14+
mx sync.Mutex
15+
connPool []string
16+
}
17+
18+
func New(connPool []string) forwarder.Forwarder {
19+
return &RoundRobin{
20+
connPool: connPool,
21+
}
22+
}
23+
func (rr *RoundRobin) NextPort() string {
24+
rr.mx.Lock()
25+
defer rr.mx.Unlock()
26+
27+
initial := rr.connNumber
28+
for {
29+
rr.connNumber = (rr.connNumber + 1) % len(rr.connPool)
30+
portNum := rr.connPool[rr.connNumber]
31+
if portNum != "" {
32+
// Check if port is actually listening
33+
conn, err := net.DialTimeout("tcp", ":"+portNum, 50*time.Millisecond)
34+
if err == nil {
35+
conn.Close()
36+
fmt.Println("PORT RETURNED ", portNum)
37+
return portNum
38+
}
39+
}
40+
if rr.connNumber == initial {
41+
break // Avoid infinite loop
42+
}
43+
}
44+
return ""
45+
}

0 commit comments

Comments
 (0)