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