50 lines
1.4 KiB
Go
50 lines
1.4 KiB
Go
package agent
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"net/http"
|
|
|
|
"github.com/coder/websocket"
|
|
)
|
|
|
|
var ErrUnexpectedMessage = errors.New("agent transport received a non-binary WebSocket message")
|
|
|
|
type WebSocketTransport struct{ connection *websocket.Conn }
|
|
|
|
func DialWebSocket(ctx context.Context, address string, client *http.Client) (*WebSocketTransport, error) {
|
|
connection, _, err := websocket.Dial(ctx, address, &websocket.DialOptions{HTTPClient: client, CompressionMode: websocket.CompressionDisabled})
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return &WebSocketTransport{connection: connection}, nil
|
|
}
|
|
|
|
func (transport *WebSocketTransport) Write(ctx context.Context, payload []byte) error {
|
|
if transport == nil || transport.connection == nil {
|
|
return ErrSessionClosed
|
|
}
|
|
return transport.connection.Write(ctx, websocket.MessageBinary, payload)
|
|
}
|
|
|
|
func (transport *WebSocketTransport) Read(ctx context.Context) ([]byte, error) {
|
|
if transport == nil || transport.connection == nil {
|
|
return nil, ErrSessionClosed
|
|
}
|
|
kind, payload, err := transport.connection.Read(ctx)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
if kind != websocket.MessageBinary {
|
|
return nil, ErrUnexpectedMessage
|
|
}
|
|
return payload, nil
|
|
}
|
|
|
|
func (transport *WebSocketTransport) Close() error {
|
|
if transport == nil || transport.connection == nil {
|
|
return nil
|
|
}
|
|
return transport.connection.Close(websocket.StatusNormalClosure, "")
|
|
}
|