| 1 | package provider |
| 2 | |
| 3 | import ( |
| 4 | "context" |
| 5 | "crypto/tls" |
| 6 | "errors" |
| 7 | "io" |
| 8 | "net" |
| 9 | "net/http" |
| 10 | "net/http/httptest" |
| 11 | "strings" |
| 12 | "sync/atomic" |
| 13 | "testing" |
| 14 | |
| 15 | "reasonix/internal/netclient" |
| 16 | ) |
| 17 | |
| 18 | // deadUpstreamTunnel forwards bytes to a TLS backend. Connections accepted up |
| 19 | // to killThrough are dropped the moment the client writes on them, the way a |
| 20 | // local proxy whose upstream half died while idle only finds out on the next write. |
| 21 | type deadUpstreamTunnel struct { |
| 22 | ln net.Listener |
| 23 | backend string |
| 24 | accepted atomic.Int32 |
| 25 | killThrough atomic.Int32 |
| 26 | } |
| 27 | |
| 28 | // killExisting marks every connection accepted so far as dead upstream; new |
| 29 | // connections still reach the backend. |
| 30 | func (p *deadUpstreamTunnel) killExisting() { p.killThrough.Store(p.accepted.Load()) } |
| 31 | |
| 32 | func (p *deadUpstreamTunnel) serve() { |
| 33 | for { |
| 34 | c, err := p.ln.Accept() |
| 35 | if err != nil { |
| 36 | return |
| 37 | } |
| 38 | go p.pipe(c, p.accepted.Add(1)) |
| 39 | } |
| 40 | } |
| 41 | |
| 42 | func (p *deadUpstreamTunnel) pipe(c net.Conn, index int32) { |
| 43 | b, err := net.Dial("tcp", p.backend) |
| 44 | if err != nil { |
| 45 | c.Close() |
| 46 | return |
| 47 | } |
| 48 | go func() { _, _ = io.Copy(c, b); c.Close() }() |
| 49 | buf := make([]byte, 32<<10) |
| 50 | for { |
| 51 | n, err := c.Read(buf) |
| 52 | if n > 0 && index <= p.killThrough.Load() { |
| 53 | c.Close() |
| 54 | b.Close() |
| 55 | return |
| 56 | } |
| 57 | if n > 0 { |
| 58 | if _, werr := b.Write(buf[:n]); werr != nil { |
| 59 | c.Close() |
| 60 | return |
| 61 | } |
| 62 | } |
| 63 | if err != nil { |
| 64 | b.Close() |
| 65 | return |
| 66 | } |
| 67 | } |
| 68 | } |
| 69 | |
| 70 | type staleConnFixture struct { |
| 71 | tunnel *deadUpstreamTunnel |
| 72 | client *http.Client |
| 73 | url string |
| 74 | hits *atomic.Int32 |
| 75 | } |
| 76 | |
| 77 | func newStaleConnFixture(t *testing.T, http2 bool) *staleConnFixture { |
| 78 | t.Helper() |
| 79 | var hits atomic.Int32 |
| 80 | srv := httptest.NewUnstartedServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { |
| 81 | _, _ = io.Copy(io.Discard, r.Body) |
| 82 | hits.Add(1) |
| 83 | _, _ = io.WriteString(w, "data: [DONE]\n\n") |
| 84 | })) |
| 85 | srv.EnableHTTP2 = http2 |
| 86 | srv.StartTLS() |
| 87 | t.Cleanup(srv.Close) |
| 88 | ln, err := net.Listen("tcp", "127.0.0.1:0") |
| 89 | if err != nil { |
| 90 | t.Fatal(err) |
| 91 | } |
| 92 | t.Cleanup(func() { ln.Close() }) |
| 93 | tunnel := &deadUpstreamTunnel{ln: ln, backend: srv.Listener.Addr().String()} |
| 94 | go tunnel.serve() |
| 95 | tr, err := netclient.NewTransport(netclient.ProxySpec{Mode: netclient.ModeOff}, netclient.TransportOptions{}) |
| 96 | if err != nil { |
| 97 | t.Fatal(err) |
| 98 | } |
| 99 | tr.TLSClientConfig = &tls.Config{RootCAs: srv.Client().Transport.(*http.Transport).TLSClientConfig.RootCAs} |
| 100 | t.Cleanup(tr.CloseIdleConnections) |
| 101 | return &staleConnFixture{tunnel: tunnel, client: &http.Client{Transport: tr}, url: "https://" + ln.Addr().String() + "/chat/completions", hits: &hits} |
| 102 | } |
| 103 | |
| 104 | func (f *staleConnFixture) send(ctx context.Context) (*http.Response, error) { |
| 105 | return SendWithRetry(ctx, f.client, SendOptions{Provider: "deepseek", Protocol: "openai"}, func(ctx context.Context) (*http.Request, error) { |
| 106 | return http.NewRequestWithContext(ctx, http.MethodPost, f.url, strings.NewReader(`{"stream":true}`)) |
| 107 | }) |
| 108 | } |
| 109 | |
| 110 | func (f *staleConnFixture) warm(t *testing.T) { |
| 111 | t.Helper() |
| 112 | resp, err := f.send(context.Background()) |
| 113 | if err != nil { |
| 114 | t.Fatal(err) |
| 115 | } |
| 116 | _, _ = io.Copy(io.Discard, resp.Body) |
| 117 | resp.Body.Close() |
| 118 | } |
| 119 | |
| 120 | // #11254: after an idle pause the pooled connection is dead, and the next turn |
| 121 | // failed with "unexpected EOF" although the request never reached the provider. |
| 122 | func TestSendWithRetryResendsOnceWhenIdlePooledConnectionIsDead(t *testing.T) { |
| 123 | for _, tc := range []struct { |
| 124 | name string |
| 125 | http2 bool |
| 126 | }{{"http2", true}, {"http1", false}} { |
| 127 | t.Run(tc.name, func(t *testing.T) { |
| 128 | f := newStaleConnFixture(t, tc.http2) |
| 129 | f.warm(t) |
| 130 | f.tunnel.killExisting() |
| 131 | ctx := WithRequestAttemptCounter(context.Background()) |
| 132 | resp, err := f.send(ctx) |
| 133 | if err != nil { |
| 134 | t.Fatalf("send on dead pooled connection: %v", err) |
| 135 | } |
| 136 | _, _ = io.Copy(io.Discard, resp.Body) |
| 137 | resp.Body.Close() |
| 138 | if got := f.hits.Load(); got != 2 { |
| 139 | t.Fatalf("provider saw %d requests, want 2 (warm-up + one delivery)", got) |
| 140 | } |
| 141 | if got := RequestAttemptCount(ctx); got != 2 { |
| 142 | t.Fatalf("attempts = %d, want the failed write and the resend counted", got) |
| 143 | } |
| 144 | }) |
| 145 | } |
| 146 | } |
| 147 | |
| 148 | // A connection that was never pooled carries no staleness: its failure is the |
| 149 | // network's, and the user decides whether to send again. |
| 150 | func TestSendWithRetryDoesNotResendOnFreshConnection(t *testing.T) { |
| 151 | f := newStaleConnFixture(t, true) |
| 152 | f.tunnel.killThrough.Store(1 << 30) |
| 153 | ctx := WithRequestAttemptCounter(context.Background()) |
| 154 | _, err := f.send(ctx) |
| 155 | if err == nil { |
| 156 | t.Fatal("fresh connection failure was absorbed") |
| 157 | } |
| 158 | var failure *RequestFailure |
| 159 | if !errors.As(err, &failure) { |
| 160 | t.Fatalf("error = %T, want *RequestFailure", err) |
| 161 | } |
| 162 | if got := RequestAttemptCount(ctx); got != 1 { |
| 163 | t.Fatalf("attempts = %d, want 1", got) |
| 164 | } |
| 165 | } |
| 166 | |
| 167 | func TestStaleIdleConnectionNeedsSilenceAndIdleReuse(t *testing.T) { |
| 168 | cases := []struct { |
| 169 | name string |
| 170 | idleReuse bool |
| 171 | responded bool |
| 172 | err error |
| 173 | want bool |
| 174 | }{ |
| 175 | {"idle reuse, silent, EOF", true, false, io.ErrUnexpectedEOF, true}, |
| 176 | {"not reused", false, false, io.ErrUnexpectedEOF, false}, |
| 177 | {"response bytes arrived", true, true, io.ErrUnexpectedEOF, false}, |
| 178 | {"not a closed connection", true, false, errors.New("tls: bad certificate"), false}, |
| 179 | {"deadline", true, false, context.DeadlineExceeded, false}, |
| 180 | } |
| 181 | for _, tc := range cases { |
| 182 | t.Run(tc.name, func(t *testing.T) { |
| 183 | p := &connectionProbe{} |
| 184 | p.idleReuse.Store(tc.idleReuse) |
| 185 | p.responded.Store(tc.responded) |
| 186 | if got := p.staleIdleConnection(context.Background(), tc.err); got != tc.want { |
| 187 | t.Fatalf("staleIdleConnection = %v, want %v", got, tc.want) |
| 188 | } |
| 189 | }) |
| 190 | } |
| 191 | canceled, cancel := context.WithCancel(context.Background()) |
| 192 | cancel() |
| 193 | p := &connectionProbe{} |
| 194 | p.idleReuse.Store(true) |
| 195 | if p.staleIdleConnection(canceled, io.ErrUnexpectedEOF) { |
| 196 | t.Fatal("a canceled turn must not be resent") |
| 197 | } |
| 198 | } |
| 199 |