Files
2026-09-01 09:57:33 +01:00

135 lines
4.5 KiB
Go

// Package workload adapts projected JWT workload credentials to the strict
// domain policy. JWT signature/key trust stays injected at this boundary.
package workload
import (
"encoding/base64"
"encoding/json"
"fmt"
"strings"
"time"
"github.com/cosmic-clash/cosmic-clash/server/domain"
)
type SignatureVerifier func(signingInput, signature []byte) bool
// ParseAndValidate parses a compact JWT, verifies its signature before domain
// validation, and returns only the exact one-allocation binding accepted by
// the policy. It intentionally does not fetch keys or trust an alg claim.
func ParseAndValidate(token string, expected domain.WorkloadBinding, verify SignatureVerifier, now time.Time) (domain.WorkloadBinding, error) {
header, claims, signingInput, signature, err := parse(token)
if err != nil || header.Alg == "" || strings.EqualFold(header.Alg, "none") || verify == nil || !verify(signingInput, signature) {
return domain.WorkloadBinding{}, domain.ErrWorkloadCredential
}
credential, err := claims.credential(signature)
if err != nil {
return domain.WorkloadBinding{}, domain.ErrWorkloadCredential
}
policy, err := domain.NewWorkloadCredentialPolicy(expected, func(candidate domain.WorkloadCredential) bool {
return verify(signingInput, candidate.Signature)
})
if err != nil {
return domain.WorkloadBinding{}, domain.ErrWorkloadCredential
}
return policy.Validate(credential, now)
}
type tokenHeader struct {
Alg string `json:"alg"`
}
type tokenClaims map[string]json.RawMessage
func parse(token string) (tokenHeader, tokenClaims, []byte, []byte, error) {
parts := strings.Split(token, ".")
if len(parts) != 3 || parts[0] == "" || parts[1] == "" || parts[2] == "" {
return tokenHeader{}, nil, nil, nil, fmt.Errorf("invalid compact token")
}
headerBytes, err := decode(parts[0])
if err != nil {
return tokenHeader{}, nil, nil, nil, err
}
claimsBytes, err := decode(parts[1])
if err != nil {
return tokenHeader{}, nil, nil, nil, err
}
signature, err := decode(parts[2])
if err != nil || len(signature) == 0 {
return tokenHeader{}, nil, nil, nil, fmt.Errorf("invalid token signature")
}
var header tokenHeader
if err := json.Unmarshal(headerBytes, &header); err != nil {
return tokenHeader{}, nil, nil, nil, err
}
var claims tokenClaims
if err := json.Unmarshal(claimsBytes, &claims); err != nil {
return tokenHeader{}, nil, nil, nil, err
}
return header, claims, []byte(parts[0] + "." + parts[1]), signature, nil
}
func (c tokenClaims) credential(signature []byte) (domain.WorkloadCredential, error) {
issuer, err := c.string("iss")
if err != nil {
return domain.WorkloadCredential{}, err
}
audience, err := c.audience()
if err != nil {
return domain.WorkloadCredential{}, err
}
issuedAt, err := c.time("iat")
if err != nil {
return domain.WorkloadCredential{}, err
}
expiresAt, err := c.time("exp")
if err != nil {
return domain.WorkloadCredential{}, err
}
values := make([]string, 7)
for i, name := range []string{"namespace", "service_account", "pod_uid", "gameserver_uid", "allocation_id", "match_id", "server_id"} {
values[i], err = c.string(name)
if err != nil {
return domain.WorkloadCredential{}, err
}
}
return domain.WorkloadCredential{Issuer: issuer, Audience: audience, IssuedAt: issuedAt, ExpiresAt: expiresAt, Namespace: values[0], ServiceAcct: values[1], PodUID: values[2], GameServerUID: values[3], AllocationID: values[4], MatchID: values[5], ServerID: values[6], Signature: signature}, nil
}
func (c tokenClaims) string(name string) (string, error) {
var value string
raw, ok := c[name]
if !ok || json.Unmarshal(raw, &value) != nil || value == "" {
return "", fmt.Errorf("missing %s", name)
}
return value, nil
}
func (c tokenClaims) time(name string) (time.Time, error) {
var seconds float64
raw, ok := c[name]
if !ok || json.Unmarshal(raw, &seconds) != nil || seconds <= 0 || seconds != float64(int64(seconds)) {
return time.Time{}, fmt.Errorf("invalid %s", name)
}
return time.Unix(int64(seconds), 0).UTC(), nil
}
func (c tokenClaims) audience() (string, error) {
if raw, ok := c["aud"]; ok {
var single string
if json.Unmarshal(raw, &single) == nil && single != "" {
return single, nil
}
var many []string
if json.Unmarshal(raw, &many) == nil && len(many) == 1 && many[0] != "" {
return many[0], nil
}
}
return "", fmt.Errorf("missing aud")
}
func decode(value string) ([]byte, error) {
decoded, err := base64.RawURLEncoding.DecodeString(value)
if err == nil {
return decoded, nil
}
return base64.URLEncoding.DecodeString(value)
}