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

272 lines
8.6 KiB
Go

package licenseserver
import (
"crypto/ed25519"
"crypto/rand"
"crypto/subtle"
"encoding/hex"
"encoding/json"
"errors"
"fmt"
"io"
"log/slog"
"net/http"
"os"
"strings"
"time"
"github.com/b1tsblog/ai-disclosure-standard/pkg/licensekit"
)
type Config struct {
TrustStore licensekit.TrustStore
LeasePrivateKey ed25519.PrivateKey
LeaseKeyID string
AdminToken string
DefaultLeaseTTL time.Duration
MaxLeaseTTL time.Duration
}
type Registry interface {
Get(id string) (Record, bool)
List() []Record
Put(record Record) error
SetRevoked(id string, revoked bool, reason string) error
}
type Server struct {
cfg Config
store Registry
logger *slog.Logger
mux *http.ServeMux
}
type introspectRequest struct {
Token string `json:"token"`
Product string `json:"product"`
BaseURL string `json:"baseUrl"`
Host string `json:"host"`
InstanceID string `json:"instanceId,omitempty"`
ClientVersion string `json:"clientVersion,omitempty"`
}
type introspectResponse struct {
Valid bool `json:"valid"`
LeaseToken string `json:"leaseToken,omitempty"`
ExpiresAt string `json:"expiresAt,omitempty"`
Reason string `json:"reason,omitempty"`
}
type registerRequest struct {
Token string `json:"token"`
}
type revokeRequest struct {
Reason string `json:"reason"`
}
func New(cfg Config, store Registry, logger *slog.Logger) (*Server, error) {
if store == nil {
return nil, errors.New("license registry is required")
}
if len(cfg.LeasePrivateKey) != ed25519.PrivateKeySize {
return nil, errors.New("a valid Ed25519 lease signing private key is required")
}
if strings.TrimSpace(cfg.LeaseKeyID) == "" {
return nil, errors.New("lease key id is required")
}
if _, ok := cfg.TrustStore.LeaseKeys[cfg.LeaseKeyID]; !ok {
return nil, fmt.Errorf("lease public key %q is not present in the trust store", cfg.LeaseKeyID)
}
if cfg.DefaultLeaseTTL <= 0 {
cfg.DefaultLeaseTTL = time.Hour
}
if cfg.MaxLeaseTTL <= 0 {
cfg.MaxLeaseTTL = 24 * time.Hour
}
if logger == nil {
logger = slog.Default()
}
s := &Server{cfg: cfg, store: store, logger: logger, mux: http.NewServeMux()}
s.routes()
return s, nil
}
func (s *Server) Handler() http.Handler { return s.securityHeaders(s.mux) }
func (s *Server) routes() {
s.mux.HandleFunc("GET /healthz", s.health)
s.mux.HandleFunc("POST /v1/introspect", s.introspect)
s.mux.HandleFunc("GET /v1/admin/licenses", s.admin(s.list))
s.mux.HandleFunc("POST /v1/admin/licenses", s.admin(s.register))
s.mux.HandleFunc("POST /v1/admin/licenses/{id}/revoke", s.admin(s.revoke))
s.mux.HandleFunc("POST /v1/admin/licenses/{id}/restore", s.admin(s.restore))
}
func (s *Server) health(w http.ResponseWriter, _ *http.Request) {
s.writeJSON(w, http.StatusOK, map[string]any{"status": "ok"})
}
func (s *Server) introspect(w http.ResponseWriter, r *http.Request) {
var request introspectRequest
if err := decodeBody(r, &request); err != nil {
s.writeJSON(w, http.StatusBadRequest, introspectResponse{Reason: err.Error()})
return
}
now := time.Now().UTC()
verified, err := licensekit.VerifyLicense(s.cfg.TrustStore, request.Token, now)
if err != nil {
s.writeJSON(w, http.StatusForbidden, introspectResponse{Reason: err.Error()})
return
}
claims := verified.Claims
if err := licensekit.ValidateLicenseContext(claims, request.Product, request.BaseURL, request.InstanceID); err != nil {
s.writeJSON(w, http.StatusForbidden, introspectResponse{Reason: err.Error()})
return
}
record, ok := s.store.Get(claims.LicenseID)
if !ok {
s.writeJSON(w, http.StatusForbidden, introspectResponse{Reason: "license is not registered"})
return
}
if record.TokenHash != licensekit.TokenHash(request.Token) {
s.writeJSON(w, http.StatusForbidden, introspectResponse{Reason: "registered token does not match"})
return
}
if record.Revoked {
reason := "license is revoked"
if record.Reason != "" {
reason += ": " + record.Reason
}
s.writeJSON(w, http.StatusForbidden, introspectResponse{Reason: reason})
return
}
ttl := s.cfg.DefaultLeaseTTL
if claims.Verification.LeaseTTLSeconds > 0 {
ttl = time.Duration(claims.Verification.LeaseTTLSeconds) * time.Second
}
if ttl > s.cfg.MaxLeaseTTL {
ttl = s.cfg.MaxLeaseTTL
}
if remaining := time.Until(time.Unix(claims.ExpiresAt, 0)); ttl > remaining {
ttl = remaining
}
if ttl <= 0 {
s.writeJSON(w, http.StatusForbidden, introspectResponse{Reason: "license has expired"})
return
}
host, _ := licensekit.HostFromBaseURL(request.BaseURL)
leaseID := randomID("lease")
leaseClaims := licensekit.LeaseClaims{
Version: 1, LeaseID: leaseID, LicenseID: claims.LicenseID, Product: claims.Product,
Customer: claims.Customer, Edition: claims.Edition, Features: claims.Features,
Host: host, InstanceID: request.InstanceID, IssuedAt: now.Unix(), ExpiresAt: now.Add(ttl).Unix(),
}
token, err := licensekit.SignLease(s.cfg.LeasePrivateKey, s.cfg.LeaseKeyID, leaseClaims)
if err != nil {
s.logger.Error("lease signing failed", "error", err, "license_id", claims.LicenseID)
s.writeJSON(w, http.StatusInternalServerError, introspectResponse{Reason: "lease signing failed"})
return
}
s.writeJSON(w, http.StatusOK, introspectResponse{Valid: true, LeaseToken: token, ExpiresAt: time.Unix(leaseClaims.ExpiresAt, 0).UTC().Format(time.RFC3339)})
}
func (s *Server) list(w http.ResponseWriter, _ *http.Request) {
s.writeJSON(w, http.StatusOK, map[string]any{"licenses": s.store.List()})
}
func (s *Server) register(w http.ResponseWriter, r *http.Request) {
var request registerRequest
if err := decodeBody(r, &request); err != nil {
s.writeJSON(w, http.StatusBadRequest, map[string]any{"error": err.Error()})
return
}
verified, err := licensekit.VerifyLicense(s.cfg.TrustStore, request.Token, time.Now().UTC())
if err != nil {
s.writeJSON(w, http.StatusBadRequest, map[string]any{"error": err.Error()})
return
}
claims := verified.Claims
record := Record{LicenseID: claims.LicenseID, TokenHash: licensekit.TokenHash(request.Token), Product: claims.Product, Customer: claims.Customer, Edition: claims.Edition, ExpiresAt: claims.ExpiresAt}
if err := s.store.Put(record); err != nil {
s.writeJSON(w, http.StatusInternalServerError, map[string]any{"error": err.Error()})
return
}
s.writeJSON(w, http.StatusCreated, record)
}
func (s *Server) revoke(w http.ResponseWriter, r *http.Request) {
var request revokeRequest
_ = decodeBodyAllowEmpty(r, &request)
if err := s.store.SetRevoked(r.PathValue("id"), true, strings.TrimSpace(request.Reason)); err != nil {
if errors.Is(err, os.ErrNotExist) {
s.writeJSON(w, http.StatusNotFound, map[string]any{"error": "license not found"})
return
}
s.writeJSON(w, http.StatusNotFound, map[string]any{"error": "license not found"})
return
}
s.writeJSON(w, http.StatusOK, map[string]any{"status": "revoked"})
}
func (s *Server) restore(w http.ResponseWriter, r *http.Request) {
if err := s.store.SetRevoked(r.PathValue("id"), false, ""); err != nil {
s.writeJSON(w, http.StatusNotFound, map[string]any{"error": "license not found"})
return
}
s.writeJSON(w, http.StatusOK, map[string]any{"status": "active"})
}
func (s *Server) admin(next http.HandlerFunc) http.HandlerFunc {
return func(w http.ResponseWriter, r *http.Request) {
expected := strings.TrimSpace(s.cfg.AdminToken)
actual := strings.TrimPrefix(r.Header.Get("Authorization"), "Bearer ")
if expected == "" || subtle.ConstantTimeCompare([]byte(expected), []byte(actual)) != 1 {
w.Header().Set("WWW-Authenticate", "Bearer")
s.writeJSON(w, http.StatusUnauthorized, map[string]any{"error": "unauthorized"})
return
}
next(w, r)
}
}
func (s *Server) securityHeaders(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("X-Content-Type-Options", "nosniff")
w.Header().Set("X-Frame-Options", "DENY")
w.Header().Set("Referrer-Policy", "no-referrer")
next.ServeHTTP(w, r)
})
}
func (s *Server) writeJSON(w http.ResponseWriter, status int, value any) {
w.Header().Set("Content-Type", "application/json; charset=utf-8")
w.WriteHeader(status)
_ = json.NewEncoder(w).Encode(value)
}
func decodeBody(r *http.Request, target any) error {
dec := json.NewDecoder(io.LimitReader(r.Body, 1<<20))
dec.DisallowUnknownFields()
if err := dec.Decode(target); err != nil {
return err
}
return nil
}
func decodeBodyAllowEmpty(r *http.Request, target any) error {
err := decodeBody(r, target)
if errors.Is(err, io.EOF) {
return nil
}
return err
}
func randomID(prefix string) string {
var raw [16]byte
if _, err := rand.Read(raw[:]); err != nil {
return prefix + "_" + fmt.Sprint(time.Now().UnixNano())
}
return prefix + "_" + hex.EncodeToString(raw[:])
}