package workerclient import ( "context" "errors" "fmt" "net" "net/url" "strconv" "sync" "time" altv1 "git.toki-labs.com/toki/alt/packages/contracts/gen/go/alt/v1" apiContracts "git.toki-labs.com/toki/alt/services/api/internal/contracts" protoSocket "git.toki-labs.com/toki/proto-socket/go" ) var ( ErrUnavailable = errors.New("worker is not available") ErrTimeout = errors.New("worker request timeout") ) type WorkerClient interface { Connect(ctx context.Context) error Close() error Hello(ctx context.Context, req *altv1.HelloRequest) (*altv1.HelloResponse, error) } type socketClient struct { socketURL string mu sync.RWMutex wsClient *protoSocket.WsClient } func New(socketURL string) WorkerClient { return &socketClient{ socketURL: socketURL, } } func (c *socketClient) Connect(ctx context.Context) error { c.mu.Lock() defer c.mu.Unlock() if c.wsClient != nil && c.wsClient.IsAlive() { return nil } u, err := url.Parse(c.socketURL) if err != nil { return fmt.Errorf("invalid worker socket URL %q: %w", c.socketURL, err) } host, portStr, err := net.SplitHostPort(u.Host) if err != nil { host = u.Host if u.Scheme == "wss" { portStr = "443" } else { portStr = "80" } } port, err := strconv.Atoi(portStr) if err != nil { return fmt.Errorf("invalid worker port %q: %w", portStr, err) } path := u.Path if path == "" { path = "/" } var wsClient *protoSocket.WsClient if u.Scheme == "wss" { wsClient, err = protoSocket.DialWssWithHeartbeat(ctx, host, port, path, nil, 30, 10, apiContracts.ParserMap()) } else { wsClient, err = protoSocket.DialWsWithHeartbeat(ctx, host, port, path, 30, 10, apiContracts.ParserMap()) } if err != nil { return fmt.Errorf("%w: %v", ErrUnavailable, err) } c.wsClient = wsClient return nil } func (c *socketClient) Close() error { c.mu.Lock() defer c.mu.Unlock() if c.wsClient != nil { err := c.wsClient.Close() c.wsClient = nil return err } return nil } func (c *socketClient) Hello(ctx context.Context, req *altv1.HelloRequest) (*altv1.HelloResponse, error) { c.mu.RLock() client := c.wsClient c.mu.RUnlock() if client == nil || !client.IsAlive() { return nil, ErrUnavailable } timeout := 5 * time.Second if dl, ok := ctx.Deadline(); ok { timeout = time.Until(dl) } res, err := protoSocket.SendRequestTyped[*altv1.HelloRequest, *altv1.HelloResponse](&client.Communicator, req, timeout) if err != nil { if errors.Is(err, protoSocket.ErrNotConnected) { return nil, ErrUnavailable } return nil, fmt.Errorf("%w: %v", ErrTimeout, err) } return res, nil }