net: better protocol headers
Juan Batiz-Benet committed
Dec 16, 2014 at 06:52 UTC
061e1ab861596096216af93a5dba6ef165dab7e5
4 files changed
+68
-33
net/interface.go
+3
-2
@@ -46,7 +46,7 @@ type Conn interface {
46
conn.PeerConn
47
48
// NewStream constructs a new Stream directly connected to p.
49
- NewStream(p peer.Peer) (Stream, error)
49
+ NewStream(pr ProtocolID, p peer.Peer) (Stream, error)
50
}
51
52
// Network is the interface IPFS uses for connecting to the world.
@@ -63,7 +63,8 @@ type Network interface {
63
64
// NewStream returns a new stream to given peer p.
65
// If there is no connection to p, attempts to create one.
66
- NewStream(p peer.Peer) (Stream, error)
66
+ // If ProtocolID is "", writes no header.
67
+ NewStream(ProtocolID, peer.Peer) (Stream, error)
68
69
// Swarm returns the connection Swarm
70
Swarm() *swarm.Swarm
net/mux.go
+25
-27
@@ -37,31 +37,10 @@ type Mux struct {
37
sync.RWMutex
38
}
39
40
-// NextName reads the stream and returns the next protocol name
40
+// ReadProtocolHeader reads the stream and returns the next Handler function
41
// according to the muxer encoding.
42
-func (m *Mux) NextName(s io.Reader) (string, error) {
43
-
44
- // c-string identifier
45
- // the first byte is our length
46
- l := make([]byte, 1)
47
- if _, err := io.ReadFull(s, l); err != nil {
48
- return "", err
49
- }
50
- length := int(l[0])
51
-
52
- // the next are our identifier
53
- name := make([]byte, length)
54
- if _, err := io.ReadFull(s, name); err != nil {
55
- return "", err
56
- }
57
-
58
- return string(name), nil
59
-}
60
-
61
-// NextHandler reads the stream and returns the next Handler function
62
-// according to the muxer encoding.
63
-func (m *Mux) NextHandler(s io.Reader) (string, StreamHandler, error) {
64
- name, err := m.NextName(s)
42
+func (m *Mux) ReadProtocolHeader(s io.Reader) (string, StreamHandler, error) {
43
+ name, err := ReadLengthPrefix(s)
44
if err != nil {
45
return "", nil, err
46
}
@@ -92,7 +71,7 @@ func (m *Mux) SetHandler(p ProtocolID, h StreamHandler) {
71
func (m *Mux) Handle(s Stream) {
72
ctx := context.Background()
73
95
- name, handler, err := m.NextHandler(s)
74
+ name, handler, err := m.ReadProtocolHeader(s)
75
if err != nil {
76
err = fmt.Errorf("protocol mux error: %s", err)
77
log.Error(err)
@@ -105,8 +84,27 @@ func (m *Mux) Handle(s Stream) {
84
handler(s)
85
}
86
108
-// Write writes the name into Writer with a length-byte-prefix.
109
-func Write(w io.Writer, name string) error {
87
+// ReadLengthPrefix reads the name from Reader with a length-byte-prefix.
88
+func ReadLengthPrefix(r io.Reader) (string, error) {
89
+ // c-string identifier
90
+ // the first byte is our length
91
+ l := make([]byte, 1)
92
+ if _, err := io.ReadFull(r, l); err != nil {
93
+ return "", err
94
+ }
95
+ length := int(l[0])
96
+
97
+ // the next are our identifier
98
+ name := make([]byte, length)
99
+ if _, err := io.ReadFull(r, name); err != nil {
100
+ return "", err
101
+ }
102
+
103
+ return string(name), nil
104
+}
105
+
106
+// WriteLengthPrefix writes the name into Writer with a length-byte-prefix.
107
+func WriteLengthPrefix(w io.Writer, name string) error {
108
s := make([]byte, len(name)+1)
109
s[0] = byte(len(name))
110
copy(s[1:], []byte(name))
net/mux_test.go
+2
-2
@@ -15,7 +15,7 @@ var testCases = map[string]string{
15
func TestWrite(t *testing.T) {
16
for k, v := range testCases {
17
var buf bytes.Buffer
18
- Write(&buf, k)
18
+ WriteLengthPrefix(&buf, k)
19
20
v2 := buf.Bytes()
21
if !bytes.Equal(v2, []byte(v)) {
@@ -48,7 +48,7 @@ func TestHandler(t *testing.T) {
48
continue
49
}
50
51
- name, _, err := m.NextHandler(&buf)
51
+ name, _, err := m.ReadProtocolHeader(&buf)
52
if err != nil {
53
t.Error(err)
54
continue
net/net.go
+38
-2
@@ -44,12 +44,20 @@ func (c *conn_) SwarmConn() *swarm.Conn {
44
return (*swarm.Conn)(c)
45
}
46
47
-func (c *conn_) NewStream(p peer.Peer) (Stream, error) {
47
+func (c *conn_) NewStream(pr ProtocolID, p peer.Peer) (Stream, error) {
48
s, err := (*swarm.Conn)(c).NewStream()
49
if err != nil {
50
return nil, err
51
}
52
- return (*stream)(s), nil
52
+
53
+ ss := (*stream)(s)
54
+
55
+ if err := writeProtocolHeader(pr, ss); err != nil {
56
+ ss.Close()
57
+ return nil, err
58
+ }
59
+
60
+ return ss, nil
61
}
62
63
// LocalMultiaddr is the Multiaddr on this side
@@ -154,8 +162,36 @@ func (n *network) Connectedness(p peer.Peer) Connectedness {
162
return NotConnected
163
}
164
165
+// NewStream returns a new stream to given peer p.
166
+// If there is no connection to p, attempts to create one.
167
+// If ProtocolID is "", writes no header.
168
+func (c *network) NewStreamWithPeer(pr ProtocolID, p peer.Peer) (Stream, error) {
169
+ s, err := c.swarm.NewStreamWithPeer(p)
170
+ if err != nil {
171
+ return nil, err
172
+ }
173
+
174
+ ss := (*stream)(s)
175
+
176
+ if err := writeProtocolHeader(pr, ss); err != nil {
177
+ ss.Close()
178
+ return nil, err
179
+ }
180
+
181
+ return ss, nil
182
+}
183
+
184
// SetHandler sets the protocol handler on the Network's Muxer.
185
// This operation is threadsafe.
186
func (n *network) SetHandler(p ProtocolID, h StreamHandler) {
187
n.mux.SetHandler(p, h)
188
}
189
+
190
+func writeProtocolHeader(pr ProtocolID, s Stream) error {
191
+ if pr != "" { // only write proper protocol headers
192
+ if err := WriteLengthPrefix(s, string(pr)); err != nil {
193
+ return err
194
+ }
195
+ }
196
+ return nil
197
+}