This repository has no description
1package lexutil
2
3import (
4 "cmp"
5 "context"
6 "fmt"
7 "log/slog"
8 "net/http"
9 "net/url"
10 "time"
11
12 indigoxrpc "github.com/bluesky-social/indigo/xrpc"
13 "github.com/carlmjohnson/versioninfo"
14 "github.com/gorilla/websocket"
15 cbg "github.com/whyrusleeping/cbor-gen"
16)
17
18const minHealthyConn = 30 * time.Second
19
20type Client struct {
21 indigoxrpc.Client
22 Dialer websocket.Dialer
23 Logger *slog.Logger
24}
25
26var _ LexClient = (*Client)(nil)
27
28func makeParams(p map[string]any) url.Values {
29 params := url.Values{}
30 for k, v := range p {
31 if s, ok := v.([]string); ok {
32 for _, v := range s {
33 params.Add(k, v)
34 }
35 } else {
36 params.Add(k, fmt.Sprint(v))
37 }
38 }
39 return params
40}
41
42type processFn func(ctx context.Context, cr *cbg.CborReader) error
43
44func (c *Client) LexDo(ctx context.Context, method string, inputEncoding string, endpoint string, params map[string]any, bodyData any, out any) error {
45 switch method {
46 case Subscription:
47 if process, ok := out.(processFn); ok {
48 return c.LexSubscribe(ctx, endpoint, params, process)
49 } else if redialer, ok := out.(Redialer); ok {
50 return c.LexSubscribeWithRedialer(ctx, endpoint, params, redialer)
51 } else {
52 return fmt.Errorf("unknown output type: %T", out)
53 }
54 default:
55 return c.Client.LexDo(ctx, method, inputEncoding, endpoint, params, bodyData, out)
56 }
57}
58
59func (c *Client) getHeader() http.Header {
60 header := http.Header{}
61 if c.UserAgent != nil {
62 header.Set("User-Agent", *c.UserAgent)
63 } else {
64 header.Set("User-Agent", "extlexutil/"+versioninfo.Short())
65 }
66 if c.Headers != nil {
67 for k, v := range c.Headers {
68 header.Set(k, v)
69 }
70 }
71 return header
72}
73
74func (c *Client) LexSubscribe(ctx context.Context, endpoint string, params map[string]any, process func(ctx context.Context, cr *cbg.CborReader) error) error {
75 logger := cmp.Or(c.Logger, slog.Default().With("system", "events"))
76 rurl, err := url.Parse(c.Host)
77 if err != nil {
78 return err
79 }
80 if rurl.Scheme == "http" {
81 rurl.Scheme = "ws"
82 } else {
83 rurl.Scheme = "wss"
84 }
85 surl := rurl.JoinPath("/xrpc", endpoint)
86 surl.RawQuery = makeParams(params).Encode()
87
88 header := c.getHeader()
89
90 u := surl.String()
91 conn, resp, err := c.Dialer.DialContext(ctx, u, header)
92 if err != nil {
93 return fmt.Errorf("%w: %w", ErrDialFailure, err)
94 }
95
96 logger.Debug("event subscription response", "code", resp.StatusCode, "url", u)
97
98 return c.handleConn(ctx, conn, process)
99}
100
101func (c *Client) LexSubscribeWithRedialer(ctx context.Context, endpoint string, params map[string]any, redialer Redialer) error {
102 logger := cmp.Or(c.Logger, slog.Default().With("system", "events"))
103 rurl, err := url.Parse(c.Host)
104 if err != nil {
105 return err
106 }
107 if rurl.Scheme == "http" {
108 rurl.Scheme = "ws"
109 } else {
110 rurl.Scheme = "wss"
111 }
112 surl := rurl.JoinPath("/xrpc", endpoint)
113
114 header := c.getHeader()
115
116 var backoff int
117 // returns false if the retry budget is exhausted
118 sleepBackoff := func() bool {
119 select {
120 case <-ctx.Done():
121 case <-time.After(time.Duration(5+backoff) * time.Second):
122 }
123 backoff++
124 return backoff <= 15
125 }
126
127 for {
128 select {
129 case <-ctx.Done():
130 return ctx.Err()
131 default:
132 }
133
134 surl.RawQuery = makeParams(params).Encode()
135
136 u := surl.String()
137 conn, resp, err := c.Dialer.DialContext(ctx, u, header)
138 if err != nil {
139 logger.Warn("dialing failed", "err", err, "backoff", backoff)
140 if !sleepBackoff() {
141 return fmt.Errorf("%w: %w", ErrDialFailure, err)
142 }
143 continue
144 }
145
146 logger.Debug("event subscription response", "code", resp.StatusCode, "url", u)
147
148 connectedAt := time.Now()
149 connErr := c.handleConn(ctx, conn, redialer.Process)
150 if connErr != nil {
151 logger.Warn("host connection failed", "err", connErr, "backoff", backoff)
152 }
153
154 // updates cursor
155 updated := redialer.UpdateParams(ctx, params)
156
157 // a connection that drops immediately shouldnt reset backoff
158 // this to avoid reconnect storms
159 if updated || time.Since(connectedAt) >= minHealthyConn {
160 backoff = 0
161 continue
162 }
163 if !sleepBackoff() {
164 return fmt.Errorf("%w: %w", ErrConnFailure, connErr)
165 }
166 }
167}
168
169func (c *Client) handleConn(ctx context.Context, conn *websocket.Conn, process func(ctx context.Context, cr *cbg.CborReader) error) error {
170 logger := cmp.Or(c.Logger, slog.Default().With("system", "events"))
171 ctx, cancel := context.WithCancel(ctx)
172 defer cancel()
173
174 go func() {
175 t := time.NewTicker(time.Second * 30)
176 defer t.Stop()
177 failcount := 0
178
179 for {
180
181 select {
182 case <-t.C:
183 if err := conn.WriteControl(websocket.PingMessage, []byte{}, time.Now().Add(time.Second*10)); err != nil {
184 logger.Warn("failed to ping", "err", err)
185 failcount++
186 if failcount >= 4 {
187 logger.Error("too many ping fails", "count", failcount)
188 conn.Close()
189 return
190 }
191 } else {
192 failcount = 0 // ok ping
193 }
194 case <-ctx.Done():
195 conn.Close()
196 return
197 }
198 }
199 }()
200
201 conn.SetPingHandler(func(message string) error {
202 err := conn.WriteControl(websocket.PongMessage, []byte(message), time.Now().Add(time.Second*60))
203 if err == websocket.ErrCloseSent {
204 return nil
205 }
206 return err
207 })
208
209 conn.SetPongHandler(func(_ string) error {
210 if err := conn.SetReadDeadline(time.Now().Add(time.Minute)); err != nil {
211 logger.Error("failed to set read deadline", "err", err)
212 }
213
214 return nil
215 })
216
217 cr := new(cbg.CborReader)
218
219 for {
220 select {
221 case <-ctx.Done():
222 return ctx.Err()
223 default:
224 }
225
226 mt, rawReader, err := conn.NextReader()
227 if err != nil {
228 return fmt.Errorf("conn err at read: %w", err)
229 }
230
231 if mt != websocket.BinaryMessage {
232 return fmt.Errorf("expected binary message from subscription endpoint")
233 }
234
235 cr.SetReader(rawReader)
236
237 if err := process(ctx, cr); err != nil {
238 return err
239 }
240 }
241}