233 lines
6.1 KiB
Go
233 lines
6.1 KiB
Go
package api
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"fmt"
|
|
"net/http"
|
|
"net/url"
|
|
"strings"
|
|
"sync"
|
|
"sync/atomic"
|
|
"time"
|
|
|
|
"git.ipao.vip/rogee/creator-hub/internal/creator"
|
|
"git.ipao.vip/rogee/creator-hub/internal/environment"
|
|
"github.com/gorilla/websocket"
|
|
)
|
|
|
|
type channelMessage struct {
|
|
Type string `json:"type"`
|
|
Subscription string `json:"subscription"`
|
|
Error string `json:"error"`
|
|
Deliveries []creator.ListenerDelivery `json:"deliveries"`
|
|
}
|
|
type eventSubscription struct {
|
|
hub *gatewayEventChannel
|
|
id, alias string
|
|
messages chan channelMessage
|
|
}
|
|
type gatewayEventChannel struct {
|
|
target environment.Gateway
|
|
once sync.Once
|
|
connectErr error
|
|
conn *websocket.Conn
|
|
write sync.Mutex
|
|
mu sync.Mutex
|
|
subs map[string]*eventSubscription
|
|
done chan struct{}
|
|
finish sync.Once
|
|
}
|
|
|
|
var eventChannels = struct {
|
|
sync.Mutex
|
|
items map[string]*gatewayEventChannel
|
|
}{items: map[string]*gatewayEventChannel{}}
|
|
var subscriptionSequence atomic.Uint64
|
|
|
|
func acquireEventSubscription(ctx context.Context, target environment.Gateway, params map[string]any) (*eventSubscription, error) {
|
|
key := target.Endpoint + "\n" + target.Token
|
|
eventChannels.Lock()
|
|
hub := eventChannels.items[key]
|
|
if hub != nil {
|
|
select {
|
|
case <-hub.done:
|
|
hub = nil
|
|
default:
|
|
}
|
|
}
|
|
if hub == nil {
|
|
hub = &gatewayEventChannel{target: target, subs: map[string]*eventSubscription{}, done: make(chan struct{})}
|
|
eventChannels.items[key] = hub
|
|
}
|
|
alias, _ := params["alias"].(string)
|
|
sub := &eventSubscription{hub: hub, id: fmt.Sprint(subscriptionSequence.Add(1)), alias: alias, messages: make(chan channelMessage, 2)}
|
|
// Reserve the subscription before releasing the pool lock, so closing the
|
|
// previous last account cannot retire a channel another account is joining.
|
|
hub.mu.Lock()
|
|
hub.subs[sub.id] = sub
|
|
hub.mu.Unlock()
|
|
eventChannels.Unlock()
|
|
hub.once.Do(func() {
|
|
u, err := url.Parse(strings.TrimRight(target.Endpoint, "/") + "/v1/channel")
|
|
if err != nil {
|
|
hub.connectErr = err
|
|
hub.fail(err)
|
|
return
|
|
}
|
|
switch u.Scheme {
|
|
case "http":
|
|
u.Scheme = "ws"
|
|
case "https":
|
|
u.Scheme = "wss"
|
|
default:
|
|
hub.connectErr = errors.New("invalid gateway scheme")
|
|
hub.fail(hub.connectErr)
|
|
return
|
|
}
|
|
headers := http.Header{"Authorization": []string{"Bearer " + target.Token}}
|
|
dialer := websocket.Dialer{HandshakeTimeout: 10 * time.Second}
|
|
hub.conn, _, hub.connectErr = dialer.DialContext(ctx, u.String(), headers)
|
|
if hub.connectErr != nil {
|
|
hub.fail(hub.connectErr)
|
|
return
|
|
}
|
|
hub.conn.SetReadLimit(16 << 20)
|
|
hub.conn.SetPingHandler(func(data string) error {
|
|
_ = hub.conn.SetReadDeadline(time.Now().Add(45 * time.Second))
|
|
return hub.conn.WriteControl(websocket.PongMessage, []byte(data), time.Now().Add(5*time.Second))
|
|
})
|
|
go hub.receive()
|
|
})
|
|
if hub.connectErr != nil {
|
|
sub.close()
|
|
return nil, fmt.Errorf("gateway WS connection: %w", hub.connectErr)
|
|
}
|
|
value := make(map[string]any, len(params)+2)
|
|
for k, v := range params {
|
|
value[k] = v
|
|
}
|
|
value["type"] = "subscribe"
|
|
value["subscription"] = sub.id
|
|
if err := hub.send(value); err != nil {
|
|
sub.close()
|
|
return nil, err
|
|
}
|
|
select {
|
|
case message := <-sub.messages:
|
|
if message.Type != "subscribed" {
|
|
sub.close()
|
|
return nil, fmt.Errorf("gateway WS subscription: %s", message.Error)
|
|
}
|
|
return sub, nil
|
|
case <-hub.done:
|
|
sub.close()
|
|
return nil, errors.New("gateway WS disconnected")
|
|
case <-ctx.Done():
|
|
sub.close()
|
|
return nil, ctx.Err()
|
|
}
|
|
}
|
|
func (h *gatewayEventChannel) send(value any) error {
|
|
h.write.Lock()
|
|
defer h.write.Unlock()
|
|
select {
|
|
case <-h.done:
|
|
return errors.New("gateway WS disconnected")
|
|
default:
|
|
}
|
|
if err := h.conn.SetWriteDeadline(time.Now().Add(5 * time.Second)); err != nil {
|
|
return err
|
|
}
|
|
if err := h.conn.WriteJSON(value); err != nil {
|
|
h.fail(err)
|
|
return err
|
|
}
|
|
return nil
|
|
}
|
|
func (h *gatewayEventChannel) receive() {
|
|
for {
|
|
_ = h.conn.SetReadDeadline(time.Now().Add(45 * time.Second))
|
|
var message channelMessage
|
|
if err := h.conn.ReadJSON(&message); err != nil {
|
|
h.fail(err)
|
|
return
|
|
}
|
|
if message.Type != "subscribed" && message.Type != "deliveries" && message.Type != "error" {
|
|
h.fail(errors.New("invalid gateway WS message"))
|
|
return
|
|
}
|
|
if message.Type == "error" && message.Subscription == "" {
|
|
h.fail(errors.New(message.Error))
|
|
return
|
|
}
|
|
h.mu.Lock()
|
|
sub := h.subs[message.Subscription]
|
|
h.mu.Unlock()
|
|
if sub == nil {
|
|
continue
|
|
} // A closed generation cannot receive delayed replies.
|
|
select {
|
|
case sub.messages <- message:
|
|
default:
|
|
h.fail(errors.New("gateway WS subscription queue full"))
|
|
return
|
|
}
|
|
}
|
|
}
|
|
func (h *gatewayEventChannel) fail(err error) {
|
|
h.finish.Do(func() {
|
|
close(h.done)
|
|
if h.conn != nil {
|
|
_ = h.conn.Close()
|
|
}
|
|
h.mu.Lock()
|
|
defer h.mu.Unlock()
|
|
for _, sub := range h.subs {
|
|
select {
|
|
case sub.messages <- channelMessage{Type: "error", Error: err.Error()}:
|
|
default:
|
|
}
|
|
}
|
|
})
|
|
}
|
|
func (s *eventSubscription) poll(ctx context.Context, acks []string) ([]creator.ListenerDelivery, error) {
|
|
if len(acks) > 0 {
|
|
if err := s.hub.send(map[string]any{"type": "ack", "alias": s.alias, "subscription": s.id, "delivery_ids": acks}); err != nil {
|
|
return nil, err
|
|
}
|
|
}
|
|
select {
|
|
case message := <-s.messages:
|
|
if message.Type != "deliveries" || len(message.Deliveries) == 0 {
|
|
return nil, fmt.Errorf("gateway event delivery: %s", message.Error)
|
|
}
|
|
for _, d := range message.Deliveries {
|
|
if err := d.Validate(); err != nil {
|
|
return nil, err
|
|
}
|
|
}
|
|
return message.Deliveries, nil
|
|
case <-s.hub.done:
|
|
return nil, errors.New("gateway WS disconnected")
|
|
case <-ctx.Done():
|
|
return nil, ctx.Err()
|
|
}
|
|
}
|
|
func (s *eventSubscription) close() {
|
|
_ = s.hub.send(map[string]any{"type": "unsubscribe", "alias": s.alias, "subscription": s.id})
|
|
key := s.hub.target.Endpoint + "\n" + s.hub.target.Token
|
|
eventChannels.Lock()
|
|
s.hub.mu.Lock()
|
|
delete(s.hub.subs, s.id)
|
|
empty := len(s.hub.subs) == 0
|
|
if empty && eventChannels.items[key] == s.hub {
|
|
delete(eventChannels.items, key)
|
|
}
|
|
s.hub.mu.Unlock()
|
|
eventChannels.Unlock()
|
|
if empty {
|
|
s.hub.fail(errors.New("gateway event channel stopped"))
|
|
}
|
|
}
|