Files
2026-07-20 21:03:05 +02:00

248 lines
6.9 KiB
Go

package licensing
import (
"crypto/ed25519"
"encoding/base64"
"encoding/json"
"errors"
"fmt"
"net/url"
"sort"
"strings"
"time"
)
const (
FeatureCustomText = "custom_text"
FeatureCustomBadge = "custom_badge"
FeatureWhiteLabel = "white_label"
)
type Claims struct {
Version int `json:"version"`
Customer string `json:"customer"`
Plan string `json:"plan"`
Features []string `json:"features"`
Domains []string `json:"domains,omitempty"`
IssuedAt int64 `json:"issuedAt"`
ExpiresAt int64 `json:"expiresAt"`
}
type Status struct {
Edition string `json:"edition"`
Licensed bool `json:"licensed"`
Customer string `json:"customer,omitempty"`
Plan string `json:"plan,omitempty"`
Features []string `json:"features"`
ExpiresAt string `json:"expiresAt,omitempty"`
Reason string `json:"reason,omitempty"`
}
type Manager struct {
status Status
features map[string]bool
}
func Community() *Manager {
return &Manager{status: Status{Edition: "community", Features: []string{}}, features: map[string]bool{}}
}
func New(publicKeyEncoded, token, baseURL string, allowInsecurePro bool, now time.Time) *Manager {
if allowInsecurePro {
features := []string{FeatureCustomBadge, FeatureCustomText, FeatureWhiteLabel}
return &Manager{
status: Status{Edition: "pro", Licensed: true, Customer: "development", Plan: "pro-dev", Features: features, Reason: "insecure development override"},
features: featureSet(features),
}
}
if strings.TrimSpace(publicKeyEncoded) == "" || strings.TrimSpace(token) == "" {
return Community()
}
claims, err := Verify(publicKeyEncoded, token, now)
if err != nil {
m := Community()
m.status.Reason = err.Error()
return m
}
if err := validateDomain(claims.Domains, baseURL); err != nil {
m := Community()
m.status.Reason = err.Error()
return m
}
features := uniqueSorted(claims.Features)
return &Manager{
status: Status{
Edition: "pro", Licensed: true, Customer: claims.Customer, Plan: claims.Plan,
Features: features, ExpiresAt: time.Unix(claims.ExpiresAt, 0).UTC().Format(time.RFC3339),
},
features: featureSet(features),
}
}
func (m *Manager) Has(feature string) bool { return m != nil && m.features[feature] }
func (m *Manager) Status() Status {
if m == nil {
return Community().status
}
out := m.status
out.Features = make([]string, len(m.status.Features))
copy(out.Features, m.status.Features)
return out
}
func Sign(privateKey ed25519.PrivateKey, claims Claims) (string, error) {
if len(privateKey) != ed25519.PrivateKeySize {
return "", errors.New("invalid Ed25519 private key")
}
if err := validateClaims(claims, time.Unix(claims.IssuedAt, 0)); err != nil {
return "", err
}
payload, err := json.Marshal(claims)
if err != nil {
return "", fmt.Errorf("marshal claims: %w", err)
}
payloadPart := base64.RawURLEncoding.EncodeToString(payload)
sig := ed25519.Sign(privateKey, []byte(payloadPart))
return payloadPart + "." + base64.RawURLEncoding.EncodeToString(sig), nil
}
func Verify(publicKeyEncoded, token string, now time.Time) (Claims, error) {
publicKeyBytes, err := decodeKey(publicKeyEncoded)
if err != nil {
return Claims{}, fmt.Errorf("decode public key: %w", err)
}
if len(publicKeyBytes) != ed25519.PublicKeySize {
return Claims{}, errors.New("public key must be an Ed25519 public key")
}
parts := strings.Split(token, ".")
if len(parts) != 2 {
return Claims{}, errors.New("license token has invalid format")
}
sig, err := base64.RawURLEncoding.DecodeString(parts[1])
if err != nil {
return Claims{}, errors.New("license signature is not valid base64url")
}
if !ed25519.Verify(ed25519.PublicKey(publicKeyBytes), []byte(parts[0]), sig) {
return Claims{}, errors.New("license signature verification failed")
}
payload, err := base64.RawURLEncoding.DecodeString(parts[0])
if err != nil {
return Claims{}, errors.New("license payload is not valid base64url")
}
var claims Claims
dec := json.NewDecoder(strings.NewReader(string(payload)))
dec.DisallowUnknownFields()
if err := dec.Decode(&claims); err != nil {
return Claims{}, fmt.Errorf("decode license payload: %w", err)
}
if err := validateClaims(claims, now); err != nil {
return Claims{}, err
}
return claims, nil
}
func validateClaims(c Claims, now time.Time) error {
if c.Version != 1 {
return errors.New("unsupported license version")
}
if strings.TrimSpace(c.Customer) == "" {
return errors.New("license customer is required")
}
if strings.TrimSpace(c.Plan) == "" {
return errors.New("license plan is required")
}
if c.IssuedAt <= 0 || c.ExpiresAt <= 0 || c.ExpiresAt <= c.IssuedAt {
return errors.New("license timestamps are invalid")
}
if now.Unix() < c.IssuedAt-300 {
return errors.New("license is not active yet")
}
if now.Unix() >= c.ExpiresAt {
return errors.New("license has expired")
}
allowed := map[string]bool{FeatureCustomText: true, FeatureCustomBadge: true, FeatureWhiteLabel: true}
for _, feature := range c.Features {
if !allowed[feature] {
return fmt.Errorf("unknown license feature %q", feature)
}
}
return nil
}
func validateDomain(domains []string, baseURL string) error {
if len(domains) == 0 {
return nil
}
u, err := url.Parse(baseURL)
if err != nil || u.Hostname() == "" {
return errors.New("BASE_URL has no valid host for licensed domain validation")
}
host := strings.ToLower(u.Hostname())
for _, allowed := range domains {
allowed = strings.ToLower(strings.TrimSpace(allowed))
if allowed == "*" {
return nil
}
if host == allowed {
return nil
}
if strings.HasPrefix(allowed, "*.") {
suffix := strings.TrimPrefix(allowed, "*")
if strings.HasSuffix(host, suffix) && host != strings.TrimPrefix(suffix, ".") {
return nil
}
}
}
return fmt.Errorf("host %q is not covered by the license", host)
}
func DecodePrivateKey(encoded string) (ed25519.PrivateKey, error) {
b, err := decodeKey(encoded)
if err != nil {
return nil, err
}
if len(b) == ed25519.SeedSize {
return ed25519.NewKeyFromSeed(b), nil
}
if len(b) != ed25519.PrivateKeySize {
return nil, errors.New("private key must contain an Ed25519 seed or private key")
}
return ed25519.PrivateKey(b), nil
}
func EncodeKey(key []byte) string { return base64.RawURLEncoding.EncodeToString(key) }
func decodeKey(value string) ([]byte, error) {
value = strings.TrimSpace(value)
if b, err := base64.RawURLEncoding.DecodeString(value); err == nil {
return b, nil
}
if b, err := base64.StdEncoding.DecodeString(value); err == nil {
return b, nil
}
return nil, errors.New("key is not valid base64")
}
func featureSet(features []string) map[string]bool {
out := make(map[string]bool, len(features))
for _, f := range features {
out[f] = true
}
return out
}
func uniqueSorted(values []string) []string {
seen := map[string]bool{}
out := make([]string, 0, len(values))
for _, v := range values {
v = strings.TrimSpace(v)
if v != "" && !seen[v] {
seen[v] = true
out = append(out, v)
}
}
sort.Strings(out)
return out
}