// 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) }