517 lines
18 KiB
Go
517 lines
18 KiB
Go
package companion
|
|
|
|
import (
|
|
"bytes"
|
|
"context"
|
|
"crypto/hmac"
|
|
"crypto/rand"
|
|
"crypto/sha256"
|
|
"crypto/tls"
|
|
"encoding/base64"
|
|
"encoding/hex"
|
|
"encoding/json"
|
|
"errors"
|
|
"fmt"
|
|
"io"
|
|
"net"
|
|
"net/http"
|
|
"net/url"
|
|
"os"
|
|
"sort"
|
|
"strconv"
|
|
"strings"
|
|
"sync"
|
|
"time"
|
|
)
|
|
|
|
const (
|
|
registerPath = "/api/v1/client-managers/register"
|
|
heartbeatPath = "/api/v1/client-managers/heartbeat"
|
|
claimPath = "/api/v1/game-client-bridge/companion/commands/claim"
|
|
snapshotPath = "/api/v1/game-client-bridge/companion/snapshots"
|
|
|
|
maxRequestBytes = 128 * 1024
|
|
maxResponseBytes = 4 * 1024 * 1024
|
|
)
|
|
|
|
type Options struct {
|
|
Now func() time.Time
|
|
Nonce func() (string, error)
|
|
}
|
|
|
|
type Client struct {
|
|
config Config
|
|
proof string
|
|
httpClient *http.Client
|
|
now func() time.Time
|
|
nonce func() (string, error)
|
|
|
|
mu sync.Mutex
|
|
sessionToken string
|
|
sessionExpiresAt time.Time
|
|
heartbeatSequence uint64
|
|
snapshotSequences map[string]uint64
|
|
}
|
|
|
|
type Registration struct {
|
|
ExpiresAt time.Time
|
|
HeartbeatEverySeconds int
|
|
ServerTime time.Time
|
|
}
|
|
|
|
type HealthReport struct {
|
|
Status string
|
|
Reason string
|
|
}
|
|
|
|
type HeartbeatResult struct {
|
|
Status string
|
|
Health string
|
|
NextHeartbeatSeconds int
|
|
SessionExpiresAt time.Time
|
|
ServerTime time.Time
|
|
}
|
|
|
|
type ClaimedCommand struct {
|
|
ID string `json:"id"`
|
|
ProfileKey string `json:"profileKey"`
|
|
CommandType string `json:"commandType"`
|
|
Payload map[string]any `json:"payload"`
|
|
Priority int `json:"priority"`
|
|
FencingToken uint64 `json:"fencingToken"`
|
|
ClaimedAt time.Time `json:"claimedAt"`
|
|
LeaseExpiresAt time.Time `json:"leaseExpiresAt"`
|
|
ExpiresAt time.Time `json:"expiresAt"`
|
|
}
|
|
|
|
type CommandAck struct {
|
|
CommandID string `json:"commandId"`
|
|
State string `json:"state"`
|
|
FencingToken uint64 `json:"fencingToken"`
|
|
AcknowledgedAt time.Time `json:"acknowledgedAt"`
|
|
}
|
|
|
|
type CommandResult struct {
|
|
Status string
|
|
Summary string
|
|
Payload map[string]any
|
|
}
|
|
|
|
type CompletedCommand struct {
|
|
CommandID string `json:"commandId"`
|
|
State string `json:"state"`
|
|
UpdatedAt time.Time `json:"updatedAt"`
|
|
CompletedAt time.Time `json:"completedAt"`
|
|
}
|
|
|
|
type Snapshot struct {
|
|
Type string
|
|
SchemaVersion string
|
|
StreamKey string
|
|
Sequence uint64
|
|
ObservedAt time.Time
|
|
Payload map[string]any
|
|
KeepForSeconds int
|
|
MaxRecords int
|
|
}
|
|
|
|
type AcceptedSnapshot struct {
|
|
SnapshotID string `json:"snapshotId"`
|
|
ProfileKey string `json:"profileKey"`
|
|
Type string `json:"type"`
|
|
SchemaVersion string `json:"schemaVersion"`
|
|
StreamKey string `json:"streamKey"`
|
|
Sequence uint64 `json:"sequence"`
|
|
AcceptedAt time.Time `json:"acceptedAt"`
|
|
ExpiresAt time.Time `json:"expiresAt"`
|
|
}
|
|
|
|
type HTTPError struct {
|
|
StatusCode int
|
|
ExpectedStatus int
|
|
}
|
|
|
|
func (err HTTPError) Error() string {
|
|
return fmt.Sprintf("platform request returned HTTP %d; expected %d", err.StatusCode, err.ExpectedStatus)
|
|
}
|
|
|
|
func NewClient(config Config, options Options) (*Client, error) {
|
|
if err := config.Validate(); err != nil {
|
|
return nil, err
|
|
}
|
|
proof, configured := os.LookupEnv(config.Proof.MaterialEnv)
|
|
if !configured || proof == "" {
|
|
return nil, fmt.Errorf("component proof environment variable is not configured")
|
|
}
|
|
if len(proof) > 4096 {
|
|
return nil, fmt.Errorf("component proof environment variable is invalid")
|
|
}
|
|
|
|
baseURL, err := canonicalPlatformOrigin(config.Platform.BaseURL)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
transport := secureDefaultTransport()
|
|
timeout := time.Duration(config.Timing.RequestTimeoutSeconds) * time.Second
|
|
httpClient := &http.Client{
|
|
Transport: transport,
|
|
Timeout: timeout,
|
|
CheckRedirect: func(_ *http.Request, _ []*http.Request) error {
|
|
return http.ErrUseLastResponse
|
|
},
|
|
}
|
|
now := options.Now
|
|
if now == nil {
|
|
now = time.Now
|
|
}
|
|
nonce := options.Nonce
|
|
if nonce == nil {
|
|
nonce = randomNonce
|
|
}
|
|
config.Capabilities = append([]string(nil), config.Capabilities...)
|
|
config.Platform.BaseURL = baseURL
|
|
return &Client{config: config, proof: proof, httpClient: httpClient, now: now, nonce: nonce, snapshotSequences: make(map[string]uint64)}, nil
|
|
}
|
|
|
|
func secureDefaultTransport() *http.Transport {
|
|
return &http.Transport{
|
|
Proxy: http.ProxyFromEnvironment,
|
|
DialContext: (&net.Dialer{Timeout: 30 * time.Second, KeepAlive: 30 * time.Second}).DialContext,
|
|
ForceAttemptHTTP2: true,
|
|
MaxIdleConns: 100,
|
|
IdleConnTimeout: 90 * time.Second,
|
|
TLSHandshakeTimeout: 10 * time.Second,
|
|
ExpectContinueTimeout: time.Second,
|
|
TLSClientConfig: &tls.Config{MinVersion: tls.VersionTLS12},
|
|
}
|
|
}
|
|
|
|
func (client *Client) Register(ctx context.Context) (Registration, error) {
|
|
nonce, err := client.nonce()
|
|
if err != nil {
|
|
return Registration{}, fmt.Errorf("create component registration nonce: %w", err)
|
|
}
|
|
if !validNonce(nonce) {
|
|
return Registration{}, fmt.Errorf("component registration nonce is invalid")
|
|
}
|
|
component := client.config.Component
|
|
request := registerRequest{
|
|
InstallationID: component.InstallationID,
|
|
ServerInstanceID: component.ServerInstanceID,
|
|
ProfileKey: component.ProfileKey,
|
|
ArtifactID: component.ArtifactID,
|
|
Version: component.Version,
|
|
SourceRevision: component.SourceRevision,
|
|
TargetOS: component.TargetOS,
|
|
TargetArch: component.TargetArch,
|
|
KeyGeneration: component.KeyGeneration,
|
|
DeploymentGeneration: component.DeploymentGeneration,
|
|
Capabilities: append([]string(nil), client.config.Capabilities...),
|
|
Timestamp: client.now().UTC(),
|
|
Nonce: nonce,
|
|
}
|
|
request.Signature = registrationSignature(client.proof, request)
|
|
var response registerResponse
|
|
if err := client.postJSON(ctx, registerPath, http.StatusOK, request, &response); err != nil {
|
|
return Registration{}, err
|
|
}
|
|
if !response.Accepted || response.InstallationID != component.InstallationID || response.SessionToken == "" || response.ExpiresAt.IsZero() || !client.now().Before(response.ExpiresAt) {
|
|
return Registration{}, fmt.Errorf("platform returned an invalid component registration")
|
|
}
|
|
client.mu.Lock()
|
|
client.sessionToken = response.SessionToken
|
|
client.sessionExpiresAt = response.ExpiresAt
|
|
client.heartbeatSequence = 0
|
|
client.mu.Unlock()
|
|
return Registration{ExpiresAt: response.ExpiresAt, HeartbeatEverySeconds: response.HeartbeatEverySeconds, ServerTime: response.ServerTime}, nil
|
|
}
|
|
|
|
func (client *Client) Heartbeat(ctx context.Context, report HealthReport) (HeartbeatResult, error) {
|
|
if report.Status != "healthy" && report.Status != "degraded" && report.Status != "unhealthy" && report.Status != "offline" {
|
|
return HeartbeatResult{}, fmt.Errorf("component health status is invalid")
|
|
}
|
|
token, sessionExpiresAt, sequence, err := client.nextHeartbeat()
|
|
if err != nil {
|
|
return HeartbeatResult{}, err
|
|
}
|
|
request := heartbeatRequest{
|
|
InstallationID: client.config.Component.InstallationID,
|
|
SessionToken: token,
|
|
Sequence: sequence,
|
|
Health: report.Status,
|
|
HealthReason: report.Reason,
|
|
Capabilities: append([]string(nil), client.config.Capabilities...),
|
|
SentAt: client.now().UTC(),
|
|
}
|
|
var response heartbeatResponse
|
|
if err := client.postJSON(ctx, heartbeatPath, http.StatusOK, request, &response); err != nil {
|
|
return HeartbeatResult{}, err
|
|
}
|
|
if !response.Accepted || response.InstallationID != client.config.Component.InstallationID || !response.SessionExpiresAt.Equal(sessionExpiresAt) {
|
|
return HeartbeatResult{}, fmt.Errorf("platform returned an invalid component heartbeat")
|
|
}
|
|
return HeartbeatResult{Status: response.Status, Health: response.Health, NextHeartbeatSeconds: response.NextHeartbeatSeconds, SessionExpiresAt: response.SessionExpiresAt, ServerTime: response.ServerTime}, nil
|
|
}
|
|
|
|
func (client *Client) ClaimCommands(ctx context.Context, limit int) ([]ClaimedCommand, error) {
|
|
if limit < 0 || limit > 50 {
|
|
return nil, fmt.Errorf("claim limit must be between 0 and 50")
|
|
}
|
|
token, err := client.currentSession()
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
var response claimResponse
|
|
if err := client.postJSON(ctx, claimPath, http.StatusOK, claimRequest{SessionToken: token, Limit: limit}, &response); err != nil {
|
|
return nil, err
|
|
}
|
|
if response.Count != len(response.Items) {
|
|
return nil, fmt.Errorf("platform returned an invalid command claim batch")
|
|
}
|
|
return append([]ClaimedCommand(nil), response.Items...), nil
|
|
}
|
|
|
|
func (client *Client) AckCommand(ctx context.Context, commandID string, fencingToken uint64) (CommandAck, error) {
|
|
if commandID == "" || fencingToken == 0 {
|
|
return CommandAck{}, fmt.Errorf("command ID and fencing token are required")
|
|
}
|
|
token, err := client.currentSession()
|
|
if err != nil {
|
|
return CommandAck{}, err
|
|
}
|
|
path := "/api/v1/game-client-bridge/companion/commands/" + url.PathEscape(commandID) + "/ack"
|
|
var response CommandAck
|
|
if err := client.postJSON(ctx, path, http.StatusOK, ackRequest{SessionToken: token, FencingToken: fencingToken}, &response); err != nil {
|
|
return CommandAck{}, err
|
|
}
|
|
return response, nil
|
|
}
|
|
|
|
func (client *Client) CompleteCommand(ctx context.Context, commandID string, fencingToken uint64, result CommandResult) (CompletedCommand, error) {
|
|
if commandID == "" || fencingToken == 0 {
|
|
return CompletedCommand{}, fmt.Errorf("command ID and fencing token are required")
|
|
}
|
|
if result.Status != "succeeded" && result.Status != "failed" && result.Status != "cancelled" {
|
|
return CompletedCommand{}, fmt.Errorf("command result status is invalid")
|
|
}
|
|
token, err := client.currentSession()
|
|
if err != nil {
|
|
return CompletedCommand{}, err
|
|
}
|
|
path := "/api/v1/game-client-bridge/companion/commands/" + url.PathEscape(commandID) + "/result"
|
|
request := resultRequest{SessionToken: token, FencingToken: fencingToken, Status: result.Status, Summary: result.Summary, Payload: result.Payload}
|
|
var response CompletedCommand
|
|
if err := client.postJSON(ctx, path, http.StatusOK, request, &response); err != nil {
|
|
return CompletedCommand{}, err
|
|
}
|
|
return response, nil
|
|
}
|
|
|
|
func (client *Client) UploadSnapshot(ctx context.Context, snapshot Snapshot) (AcceptedSnapshot, error) {
|
|
if snapshot.Type == "" || snapshot.SchemaVersion == "" || snapshot.StreamKey == "" || snapshot.Sequence == 0 || snapshot.ObservedAt.IsZero() || snapshot.Payload == nil || snapshot.KeepForSeconds <= 0 {
|
|
return AcceptedSnapshot{}, fmt.Errorf("typed snapshot is incomplete")
|
|
}
|
|
if err := client.reserveSnapshotSequence(snapshot); err != nil {
|
|
return AcceptedSnapshot{}, err
|
|
}
|
|
token, err := client.currentSession()
|
|
if err != nil {
|
|
return AcceptedSnapshot{}, err
|
|
}
|
|
request := snapshotRequest{
|
|
SessionToken: token, Type: snapshot.Type, SchemaVersion: snapshot.SchemaVersion, StreamKey: snapshot.StreamKey,
|
|
Sequence: snapshot.Sequence, ObservedAt: snapshot.ObservedAt.UTC(), Payload: snapshot.Payload,
|
|
KeepForSeconds: snapshot.KeepForSeconds, MaxRecords: snapshot.MaxRecords,
|
|
}
|
|
var response AcceptedSnapshot
|
|
if err := client.postJSON(ctx, snapshotPath, http.StatusAccepted, request, &response); err != nil {
|
|
return AcceptedSnapshot{}, err
|
|
}
|
|
return response, nil
|
|
}
|
|
|
|
func (client *Client) currentSession() (string, error) {
|
|
client.mu.Lock()
|
|
defer client.mu.Unlock()
|
|
if client.sessionToken == "" || client.sessionExpiresAt.IsZero() || !client.now().Before(client.sessionExpiresAt) {
|
|
return "", errors.New("component session is unavailable or expired")
|
|
}
|
|
return client.sessionToken, nil
|
|
}
|
|
|
|
func (client *Client) nextHeartbeat() (string, time.Time, uint64, error) {
|
|
client.mu.Lock()
|
|
defer client.mu.Unlock()
|
|
if client.sessionToken == "" || client.sessionExpiresAt.IsZero() || !client.now().Before(client.sessionExpiresAt) {
|
|
return "", time.Time{}, 0, errors.New("component session is unavailable or expired")
|
|
}
|
|
client.heartbeatSequence++
|
|
return client.sessionToken, client.sessionExpiresAt, client.heartbeatSequence, nil
|
|
}
|
|
|
|
func (client *Client) reserveSnapshotSequence(snapshot Snapshot) error {
|
|
client.mu.Lock()
|
|
defer client.mu.Unlock()
|
|
key := snapshot.Type + "\x00" + snapshot.StreamKey
|
|
if snapshot.Sequence <= client.snapshotSequences[key] {
|
|
return fmt.Errorf("snapshot sequence must increase for its typed stream")
|
|
}
|
|
client.snapshotSequences[key] = snapshot.Sequence
|
|
return nil
|
|
}
|
|
|
|
func (client *Client) postJSON(ctx context.Context, path string, expectedStatus int, payload any, target any) error {
|
|
encoded, err := json.Marshal(payload)
|
|
if err != nil {
|
|
return fmt.Errorf("encode platform request: %w", err)
|
|
}
|
|
if len(encoded) > maxRequestBytes {
|
|
return fmt.Errorf("platform request exceeds the bounded payload size")
|
|
}
|
|
request, err := http.NewRequestWithContext(ctx, http.MethodPost, client.config.Platform.BaseURL+path, bytes.NewReader(encoded))
|
|
if err != nil {
|
|
return fmt.Errorf("create platform request: %w", err)
|
|
}
|
|
request.Header.Set("Accept", "application/json")
|
|
request.Header.Set("Content-Type", "application/json")
|
|
response, err := client.httpClient.Do(request)
|
|
if err != nil {
|
|
return fmt.Errorf("send platform request: %w", err)
|
|
}
|
|
defer response.Body.Close()
|
|
if response.StatusCode != expectedStatus {
|
|
_, _ = io.Copy(io.Discard, io.LimitReader(response.Body, 4096))
|
|
return HTTPError{StatusCode: response.StatusCode, ExpectedStatus: expectedStatus}
|
|
}
|
|
decoder := json.NewDecoder(io.LimitReader(response.Body, maxResponseBytes))
|
|
if err := decoder.Decode(target); err != nil {
|
|
return fmt.Errorf("decode platform response: %w", err)
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func registrationSignature(proof string, request registerRequest) string {
|
|
capabilities := append([]string(nil), request.Capabilities...)
|
|
sort.Strings(capabilities)
|
|
canonical := strings.Join([]string{
|
|
request.InstallationID,
|
|
request.ServerInstanceID,
|
|
request.ProfileKey,
|
|
request.ArtifactID,
|
|
request.Version,
|
|
request.SourceRevision,
|
|
request.TargetOS,
|
|
request.TargetArch,
|
|
strconv.Itoa(request.KeyGeneration),
|
|
strconv.Itoa(request.DeploymentGeneration),
|
|
request.Timestamp.UTC().Format(time.RFC3339Nano),
|
|
request.Nonce,
|
|
strings.Join(capabilities, ","),
|
|
}, "\n")
|
|
mac := hmac.New(sha256.New, []byte(proof))
|
|
_, _ = mac.Write([]byte(canonical))
|
|
return "sha256:" + hex.EncodeToString(mac.Sum(nil))
|
|
}
|
|
|
|
func randomNonce() (string, error) {
|
|
value := make([]byte, 24)
|
|
if _, err := rand.Read(value); err != nil {
|
|
return "", err
|
|
}
|
|
return base64.RawURLEncoding.EncodeToString(value), nil
|
|
}
|
|
|
|
func validNonce(value string) bool {
|
|
if len(value) < 16 || len(value) > 128 {
|
|
return false
|
|
}
|
|
for _, character := range value {
|
|
if character >= 'A' && character <= 'Z' || character >= 'a' && character <= 'z' || character >= '0' && character <= '9' || character == '_' || character == '-' {
|
|
continue
|
|
}
|
|
return false
|
|
}
|
|
return true
|
|
}
|
|
|
|
type registerRequest struct {
|
|
InstallationID string `json:"installationId"`
|
|
ServerInstanceID string `json:"serverInstanceId"`
|
|
ProfileKey string `json:"profileKey"`
|
|
ArtifactID string `json:"artifactId"`
|
|
Version string `json:"version"`
|
|
SourceRevision string `json:"sourceRevision"`
|
|
TargetOS string `json:"targetOs"`
|
|
TargetArch string `json:"targetArch"`
|
|
KeyGeneration int `json:"keyGeneration"`
|
|
DeploymentGeneration int `json:"deploymentGeneration"`
|
|
Capabilities []string `json:"capabilities"`
|
|
Timestamp time.Time `json:"timestamp"`
|
|
Nonce string `json:"nonce"`
|
|
Signature string `json:"signature"`
|
|
}
|
|
|
|
type registerResponse struct {
|
|
Accepted bool `json:"accepted"`
|
|
InstallationID string `json:"installationId"`
|
|
SessionToken string `json:"sessionToken"`
|
|
ExpiresAt time.Time `json:"expiresAt"`
|
|
HeartbeatEverySeconds int `json:"heartbeatEverySeconds"`
|
|
ServerTime time.Time `json:"serverTime"`
|
|
}
|
|
|
|
type heartbeatRequest struct {
|
|
InstallationID string `json:"installationId"`
|
|
SessionToken string `json:"sessionToken"`
|
|
Sequence uint64 `json:"sequence"`
|
|
Health string `json:"health"`
|
|
HealthReason string `json:"healthReason,omitempty"`
|
|
Capabilities []string `json:"capabilities"`
|
|
SentAt time.Time `json:"sentAt"`
|
|
}
|
|
|
|
type heartbeatResponse struct {
|
|
Accepted bool `json:"accepted"`
|
|
InstallationID string `json:"installationId"`
|
|
Status string `json:"status"`
|
|
Health string `json:"health"`
|
|
NextHeartbeatSeconds int `json:"nextHeartbeatSeconds"`
|
|
SessionExpiresAt time.Time `json:"sessionExpiresAt"`
|
|
ServerTime time.Time `json:"serverTime"`
|
|
}
|
|
|
|
type claimRequest struct {
|
|
SessionToken string `json:"sessionToken"`
|
|
Limit int `json:"limit,omitempty"`
|
|
}
|
|
|
|
type claimResponse struct {
|
|
Items []ClaimedCommand `json:"items"`
|
|
Count int `json:"count"`
|
|
}
|
|
|
|
type ackRequest struct {
|
|
SessionToken string `json:"sessionToken"`
|
|
FencingToken uint64 `json:"fencingToken"`
|
|
}
|
|
|
|
type resultRequest struct {
|
|
SessionToken string `json:"sessionToken"`
|
|
FencingToken uint64 `json:"fencingToken"`
|
|
Status string `json:"status"`
|
|
Summary string `json:"summary,omitempty"`
|
|
Payload map[string]any `json:"payload,omitempty"`
|
|
}
|
|
|
|
type snapshotRequest struct {
|
|
SessionToken string `json:"sessionToken"`
|
|
Type string `json:"type"`
|
|
SchemaVersion string `json:"schemaVersion"`
|
|
StreamKey string `json:"streamKey"`
|
|
Sequence uint64 `json:"sequence"`
|
|
ObservedAt time.Time `json:"observedAt"`
|
|
Payload map[string]any `json:"payload"`
|
|
KeepForSeconds int `json:"keepForSeconds"`
|
|
MaxRecords int `json:"maxRecords,omitempty"`
|
|
}
|