@cryptotaxi247 / kubo / commits / a4021eb43

ctxio -- io with a context.

This commit introduces a reader and writer that respect contexts. Warning: careful how you use them. Returning leaves a goroutine reading until the read finishes.

Juan Batiz-Benet committed Dec 24, 2014 at 00:49 UTC a4021eb43352100e72e21ef0ee4c566c58751314
2 files changed +383
util/ctx/ctxio.go new
+110
@@ -0,0 +1,110 @@
1 +package ctxutil
2 +
3 +import (
4 + "io"
5 +
6 + context "github.com/jbenet/go-ipfs/Godeps/_workspace/src/code.google.com/p/go.net/context"
7 +)
8 +
9 +type ioret struct {
10 + n int
11 + err error
12 +}
13 +
14 +type Writer interface {
15 + io.Writer
16 +}
17 +
18 +type ctxWriter struct {
19 + w io.Writer
20 + ctx context.Context
21 +}
22 +
23 +// NewWriter wraps a writer to make it respect given Context.
24 +// If there is a blocking write, the returned Writer will return
25 +// whenever the context is cancelled (the return values are n=0
26 +// and err=ctx.Err().)
27 +//
28 +// Note well: this wrapper DOES NOT ACTUALLY cancel the underlying
29 +// write-- there is no way to do that with the standard go io
30 +// interface. So the read and write _will_ happen or hang. So, use
31 +// this sparingly, make sure to cancel the read or write as necesary
32 +// (e.g. closing a connection whose context is up, etc.)
33 +//
34 +// Furthermore, in order to protect your memory from being read
35 +// _after_ you've cancelled the context, this io.Writer will
36 +// first make a **copy** of the buffer.
37 +func NewWriter(ctx context.Context, w io.Writer) *ctxWriter {
38 + if ctx == nil {
39 + ctx = context.Background()
40 + }
41 + return &ctxWriter{ctx: ctx, w: w}
42 +}
43 +
44 +func (w *ctxWriter) Write(buf []byte) (int, error) {
45 + buf2 := make([]byte, len(buf))
46 + copy(buf2, buf)
47 +
48 + c := make(chan ioret)
49 +
50 + go func() {
51 + n, err := w.w.Write(buf2)
52 + c <- ioret{n, err}
53 + close(c)
54 + }()
55 +
56 + select {
57 + case r := <-c:
58 + return r.n, r.err
59 + case <-w.ctx.Done():
60 + return 0, w.ctx.Err()
61 + }
62 +}
63 +
64 +type Reader interface {
65 + io.Reader
66 +}
67 +
68 +type ctxReader struct {
69 + r io.Reader
70 + ctx context.Context
71 +}
72 +
73 +// NewReader wraps a reader to make it respect given Context.
74 +// If there is a blocking read, the returned Reader will return
75 +// whenever the context is cancelled (the return values are n=0
76 +// and err=ctx.Err().)
77 +//
78 +// Note well: this wrapper DOES NOT ACTUALLY cancel the underlying
79 +// write-- there is no way to do that with the standard go io
80 +// interface. So the read and write _will_ happen or hang. So, use
81 +// this sparingly, make sure to cancel the read or write as necesary
82 +// (e.g. closing a connection whose context is up, etc.)
83 +//
84 +// Furthermore, in order to protect your memory from being read
85 +// _before_ you've cancelled the context, this io.Reader will
86 +// allocate a buffer of the same size, and **copy** into the client's
87 +// if the read succeeds in time.
88 +func NewReader(ctx context.Context, r io.Reader) *ctxReader {
89 + return &ctxReader{ctx: ctx, r: r}
90 +}
91 +
92 +func (r *ctxReader) Read(buf []byte) (int, error) {
93 + buf2 := make([]byte, len(buf))
94 +
95 + c := make(chan ioret)
96 +
97 + go func() {
98 + n, err := r.r.Read(buf2)
99 + c <- ioret{n, err}
100 + close(c)
101 + }()
102 +
103 + select {
104 + case ret := <-c:
105 + copy(buf, buf2)
106 + return ret.n, ret.err
107 + case <-r.ctx.Done():
108 + return 0, r.ctx.Err()
109 + }
110 +}
util/ctx/ctxio_test.go new
+273
@@ -0,0 +1,273 @@
1 +package ctxutil
2 +
3 +import (
4 + "bytes"
5 + "io"
6 + "testing"
7 + "time"
8 +
9 + context "github.com/jbenet/go-ipfs/Godeps/_workspace/src/code.google.com/p/go.net/context"
10 +)
11 +
12 +func TestReader(t *testing.T) {
13 + buf := []byte("abcdef")
14 + buf2 := make([]byte, 3)
15 + r := NewReader(context.Background(), bytes.NewReader(buf))
16 +
17 + // read first half
18 + n, err := r.Read(buf2)
19 + if n != 3 {
20 + t.Error("n should be 3")
21 + }
22 + if err != nil {
23 + t.Error("should have no error")
24 + }
25 + if string(buf2) != string(buf[:3]) {
26 + t.Error("incorrect contents")
27 + }
28 +
29 + // read second half
30 + n, err = r.Read(buf2)
31 + if n != 3 {
32 + t.Error("n should be 3")
33 + }
34 + if err != nil {
35 + t.Error("should have no error")
36 + }
37 + if string(buf2) != string(buf[3:6]) {
38 + t.Error("incorrect contents")
39 + }
40 +
41 + // read more.
42 + n, err = r.Read(buf2)
43 + if n != 0 {
44 + t.Error("n should be 0", n)
45 + }
46 + if err != io.EOF {
47 + t.Error("should be EOF", err)
48 + }
49 +}
50 +
51 +func TestWriter(t *testing.T) {
52 + var buf bytes.Buffer
53 + w := NewWriter(context.Background(), &buf)
54 +
55 + // write three
56 + n, err := w.Write([]byte("abc"))
57 + if n != 3 {
58 + t.Error("n should be 3")
59 + }
60 + if err != nil {
61 + t.Error("should have no error")
62 + }
63 + if string(buf.Bytes()) != string("abc") {
64 + t.Error("incorrect contents")
65 + }
66 +
67 + // write three more
68 + n, err = w.Write([]byte("def"))
69 + if n != 3 {
70 + t.Error("n should be 3")
71 + }
72 + if err != nil {
73 + t.Error("should have no error")
74 + }
75 + if string(buf.Bytes()) != string("abcdef") {
76 + t.Error("incorrect contents")
77 + }
78 +}
79 +
80 +func TestReaderCancel(t *testing.T) {
81 + ctx, cancel := context.WithCancel(context.Background())
82 + piper, pipew := io.Pipe()
83 + r := NewReader(ctx, piper)
84 +
85 + buf := make([]byte, 10)
86 + done := make(chan ioret)
87 +
88 + go func() {
89 + n, err := r.Read(buf)
90 + done <- ioret{n, err}
91 + }()
92 +
93 + pipew.Write([]byte("abcdefghij"))
94 +
95 + select {
96 + case ret := <-done:
97 + if ret.n != 10 {
98 + t.Error("ret.n should be 10", ret.n)
99 + }
100 + if ret.err != nil {
101 + t.Error("ret.err should be nil", ret.err)
102 + }
103 + if string(buf) != "abcdefghij" {
104 + t.Error("read contents differ")
105 + }
106 + case <-time.After(20 * time.Millisecond):
107 + t.Fatal("failed to read")
108 + }
109 +
110 + go func() {
111 + n, err := r.Read(buf)
112 + done <- ioret{n, err}
113 + }()
114 +
115 + cancel()
116 +
117 + select {
118 + case ret := <-done:
119 + if ret.n != 0 {
120 + t.Error("ret.n should be 0", ret.n)
121 + }
122 + if ret.err == nil {
123 + t.Error("ret.err should be ctx error", ret.err)
124 + }
125 + case <-time.After(20 * time.Millisecond):
126 + t.Fatal("failed to stop reading after cancel")
127 + }
128 +}
129 +
130 +func TestWriterCancel(t *testing.T) {
131 + ctx, cancel := context.WithCancel(context.Background())
132 + piper, pipew := io.Pipe()
133 + w := NewWriter(ctx, pipew)
134 +
135 + buf := make([]byte, 10)
136 + done := make(chan ioret)
137 +
138 + go func() {
139 + n, err := w.Write([]byte("abcdefghij"))
140 + done <- ioret{n, err}
141 + }()
142 +
143 + piper.Read(buf)
144 +
145 + select {
146 + case ret := <-done:
147 + if ret.n != 10 {
148 + t.Error("ret.n should be 10", ret.n)
149 + }
150 + if ret.err != nil {
151 + t.Error("ret.err should be nil", ret.err)
152 + }
153 + if string(buf) != "abcdefghij" {
154 + t.Error("write contents differ")
155 + }
156 + case <-time.After(20 * time.Millisecond):
157 + t.Fatal("failed to write")
158 + }
159 +
160 + go func() {
161 + n, err := w.Write([]byte("abcdefghij"))
162 + done <- ioret{n, err}
163 + }()
164 +
165 + cancel()
166 +
167 + select {
168 + case ret := <-done:
169 + if ret.n != 0 {
170 + t.Error("ret.n should be 0", ret.n)
171 + }
172 + if ret.err == nil {
173 + t.Error("ret.err should be ctx error", ret.err)
174 + }
175 + case <-time.After(20 * time.Millisecond):
176 + t.Fatal("failed to stop writing after cancel")
177 + }
178 +}
179 +
180 +func TestReadPostCancel(t *testing.T) {
181 + ctx, cancel := context.WithCancel(context.Background())
182 + piper, pipew := io.Pipe()
183 + r := NewReader(ctx, piper)
184 +
185 + buf := make([]byte, 10)
186 + done := make(chan ioret)
187 +
188 + go func() {
189 + n, err := r.Read(buf)
190 + done <- ioret{n, err}
191 + }()
192 +
193 + cancel()
194 +
195 + select {
196 + case ret := <-done:
197 + if ret.n != 0 {
198 + t.Error("ret.n should be 0", ret.n)
199 + }
200 + if ret.err == nil {
201 + t.Error("ret.err should be ctx error", ret.err)
202 + }
203 + case <-time.After(20 * time.Millisecond):
204 + t.Fatal("failed to stop reading after cancel")
205 + }
206 +
207 + pipew.Write([]byte("abcdefghij"))
208 +
209 + if !bytes.Equal(buf, make([]byte, len(buf))) {
210 + t.Fatal("buffer should have not been written to")
211 + }
212 +}
213 +
214 +func TestWritePostCancel(t *testing.T) {
215 + ctx, cancel := context.WithCancel(context.Background())
216 + piper, pipew := io.Pipe()
217 + w := NewWriter(ctx, pipew)
218 +
219 + buf := []byte("abcdefghij")
220 + buf2 := make([]byte, 10)
221 + done := make(chan ioret)
222 +
223 + go func() {
224 + n, err := w.Write(buf)
225 + done <- ioret{n, err}
226 + }()
227 +
228 + piper.Read(buf2)
229 +
230 + select {
231 + case ret := <-done:
232 + if ret.n != 10 {
233 + t.Error("ret.n should be 10", ret.n)
234 + }
235 + if ret.err != nil {
236 + t.Error("ret.err should be nil", ret.err)
237 + }
238 + if string(buf2) != "abcdefghij" {
239 + t.Error("write contents differ")
240 + }
241 + case <-time.After(20 * time.Millisecond):
242 + t.Fatal("failed to write")
243 + }
244 +
245 + go func() {
246 + n, err := w.Write(buf)
247 + done <- ioret{n, err}
248 + }()
249 +
250 + cancel()
251 +
252 + select {
253 + case ret := <-done:
254 + if ret.n != 0 {
255 + t.Error("ret.n should be 0", ret.n)
256 + }
257 + if ret.err == nil {
258 + t.Error("ret.err should be ctx error", ret.err)
259 + }
260 + case <-time.After(20 * time.Millisecond):
261 + t.Fatal("failed to stop writing after cancel")
262 + }
263 +
264 + copy(buf, []byte("aaaaaaaaaa"))
265 +
266 + piper.Read(buf2)
267 +
268 + if string(buf2) == "aaaaaaaaaa" {
269 + t.Error("buffer was read from after ctx cancel")
270 + } else if string(buf2) != "abcdefghij" {
271 + t.Error("write contents differ from expected")
272 + }
273 +}