Files
groot 47c523dd98
All checks were successful
release-tag / release-image (push) Successful in 3m51s
RC-14
2026-08-14 06:17:30 +02:00

574 lines
19 KiB
Go

package servicecontroller
import (
"bytes"
"context"
"crypto/subtle"
"encoding/json"
"errors"
"fmt"
"io"
"log"
"net/http"
"net/url"
"strings"
"sync"
"time"
"github.com/go-chi/chi/v5"
"neuralhunt/internal/customer"
)
type Config struct {
ID string
Name string
MasterURL string
AdvertiseURL string
SharedSecret string
ControlAddr string
AdminAddr string
AdminUser string
AdminPassword string
AdminCookieSecureMode string
MaxWorkers int
MaxRunning int
HeartbeatInterval time.Duration
DefaultWorkerImage string
WorkerNetworkOverride string
WorkerRegisterURLOverride string
GameURLOverride string
RegistryUsername string
RegistryPassword string
RegistryServer string
}
type Service struct {
docker *customer.DockerClient
cfg Config
hc *http.Client
adminSessions sync.Map
regMu sync.RWMutex
lastRegOK bool
lastRegAt time.Time
lastRegErr string
}
func New(docker *customer.DockerClient, cfg Config) *Service {
if cfg.MaxWorkers <= 0 {
cfg.MaxWorkers = 1000
}
if cfg.MaxRunning <= 0 {
cfg.MaxRunning = 100
}
if cfg.HeartbeatInterval < 5*time.Second {
cfg.HeartbeatInterval = 15 * time.Second
}
return &Service{docker: docker, cfg: cfg, hc: &http.Client{Timeout: 10 * time.Second}}
}
func jsonOut(w http.ResponseWriter, status int, v any) {
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(status)
_ = json.NewEncoder(w).Encode(v)
}
func decode(r *http.Request, v any, limit int64) error {
if limit <= 0 {
limit = 1 << 20
}
d := json.NewDecoder(io.LimitReader(r.Body, limit))
d.DisallowUnknownFields()
return d.Decode(v)
}
func requestIsHTTPS(r *http.Request) bool {
if r.TLS != nil {
return true
}
proto := strings.TrimSpace(strings.Split(r.Header.Get("X-Forwarded-Proto"), ",")[0])
return strings.EqualFold(proto, "https")
}
func (s *Service) security(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")
w.Header().Set("Permissions-Policy", "camera=(), microphone=(), geolocation=()")
w.Header().Set("Content-Security-Policy", "default-src 'self'; base-uri 'none'; object-src 'none'; frame-ancestors 'none'; script-src 'self'; style-src 'self'; img-src 'self' data:; connect-src 'self'; form-action 'self'")
if requestIsHTTPS(r) {
w.Header().Set("Strict-Transport-Security", "max-age=31536000")
}
next.ServeHTTP(w, r)
})
}
func (s *Service) authorized(r *http.Request) bool {
provided := strings.TrimSpace(strings.TrimPrefix(r.Header.Get("Authorization"), "Bearer "))
secret := strings.TrimSpace(s.cfg.SharedSecret)
return len(provided) == len(secret) && len(secret) >= 24 && subtle.ConstantTimeCompare([]byte(provided), []byte(secret)) == 1
}
func (s *Service) requireControl(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if !s.authorized(r) {
jsonOut(w, http.StatusUnauthorized, map[string]string{"error": "unauthorized"})
return
}
next.ServeHTTP(w, r)
})
}
func (s *Service) counts(ctx context.Context) (total, running int, err error) {
items, err := s.docker.ManagedContainers(ctx)
if err != nil {
return 0, 0, err
}
for _, item := range items {
if strings.TrimSpace(item.WorkerID) == "" {
continue
}
total++
if item.Running {
running++
}
}
return total, running, nil
}
func (s *Service) ControlRoutes() http.Handler {
r := chi.NewRouter()
r.Use(s.security)
r.Group(func(r chi.Router) {
r.Use(s.requireControl)
r.Get("/internal/health", s.health)
r.Post("/internal/image/ensure", s.ensureImage)
r.Post("/internal/image/pull", s.pullImage)
r.Post("/internal/volumes", s.createVolume)
r.Delete("/internal/volumes/{name}", s.removeVolume)
r.Post("/internal/workers", s.createWorker)
r.Post("/internal/workers/{id}/start", s.startWorker)
r.Post("/internal/workers/{id}/stop", s.stopWorker)
r.Delete("/internal/workers/{id}", s.removeWorker)
r.Get("/internal/workers/{id}/running", s.runningWorker)
r.Post("/internal/workers/{id}/file/get", s.getWorkerFile)
r.Post("/internal/workers/{id}/file/put", s.putWorkerFile)
})
return r
}
func (s *Service) health(w http.ResponseWriter, r *http.Request) {
total, running, err := s.counts(r.Context())
if err != nil {
jsonOut(w, http.StatusServiceUnavailable, map[string]any{"ok": false, "error": err.Error()})
return
}
jsonOut(w, http.StatusOK, map[string]any{"ok": true, "controller_id": s.cfg.ID, "name": s.cfg.Name, "workers": total, "running": running, "max_workers": s.cfg.MaxWorkers, "max_running": s.cfg.MaxRunning})
}
func (s *Service) ensureImage(w http.ResponseWriter, r *http.Request) {
var in struct {
Image string `json:"image"`
AutoPull bool `json:"auto_pull"`
RegistryAuth string `json:"registry_auth"`
}
if decode(r, &in, 256<<10) != nil || strings.TrimSpace(in.Image) == "" {
jsonOut(w, 400, map[string]string{"error": "image required"})
return
}
if err := s.docker.EnsureImage(r.Context(), in.Image, in.AutoPull, in.RegistryAuth); err != nil {
jsonOut(w, 502, map[string]string{"error": err.Error()})
return
}
jsonOut(w, 200, map[string]bool{"ok": true})
}
func (s *Service) pullImage(w http.ResponseWriter, r *http.Request) {
var in struct {
Image string `json:"image"`
RegistryAuth string `json:"registry_auth"`
}
if decode(r, &in, 256<<10) != nil || strings.TrimSpace(in.Image) == "" {
jsonOut(w, 400, map[string]string{"error": "image required"})
return
}
if err := s.docker.PullImage(r.Context(), in.Image, in.RegistryAuth); err != nil {
jsonOut(w, 502, map[string]string{"error": err.Error()})
return
}
jsonOut(w, 200, map[string]bool{"ok": true})
}
func (s *Service) createVolume(w http.ResponseWriter, r *http.Request) {
var in struct {
Name string `json:"name"`
}
if decode(r, &in, 64<<10) != nil || strings.TrimSpace(in.Name) == "" {
jsonOut(w, 400, map[string]string{"error": "name required"})
return
}
if err := s.docker.CreateVolume(r.Context(), in.Name); err != nil {
jsonOut(w, 502, map[string]string{"error": err.Error()})
return
}
jsonOut(w, 200, map[string]bool{"ok": true})
}
func (s *Service) removeVolume(w http.ResponseWriter, r *http.Request) {
if err := s.docker.RemoveVolume(r.Context(), chi.URLParam(r, "name")); err != nil {
jsonOut(w, 502, map[string]string{"error": err.Error()})
return
}
jsonOut(w, 200, map[string]bool{"ok": true})
}
func (s *Service) createWorker(w http.ResponseWriter, r *http.Request) {
var cfg customer.WorkerContainerConfig
if decode(r, &cfg, 1<<20) != nil || strings.TrimSpace(cfg.WorkerID) == "" || strings.TrimSpace(cfg.Image) == "" {
jsonOut(w, 400, map[string]string{"error": "invalid worker config"})
return
}
if v := strings.TrimSpace(s.cfg.WorkerNetworkOverride); v != "" {
cfg.Network = v
}
if v := strings.TrimSpace(s.cfg.WorkerRegisterURLOverride); v != "" {
cfg.RegisterURL = v
}
if v := strings.TrimSpace(s.cfg.GameURLOverride); v != "" {
cfg.GameURL = v
}
total, _, err := s.counts(r.Context())
if err != nil {
jsonOut(w, 502, map[string]string{"error": err.Error()})
return
}
if total >= s.cfg.MaxWorkers {
jsonOut(w, 503, map[string]string{"error": "service controller worker capacity reached"})
return
}
id, err := s.docker.CreateWorker(r.Context(), cfg)
if err != nil {
jsonOut(w, 502, map[string]string{"error": err.Error()})
return
}
jsonOut(w, 201, map[string]string{"id": id})
}
func (s *Service) startWorker(w http.ResponseWriter, r *http.Request) {
_, running, err := s.counts(r.Context())
if err != nil {
jsonOut(w, 502, map[string]string{"error": err.Error()})
return
}
if running >= s.cfg.MaxRunning {
jsonOut(w, 503, map[string]string{"error": "service controller running capacity reached"})
return
}
if err := s.docker.Start(r.Context(), chi.URLParam(r, "id")); err != nil {
jsonOut(w, 502, map[string]string{"error": err.Error()})
return
}
jsonOut(w, 200, map[string]bool{"ok": true})
}
func (s *Service) stopWorker(w http.ResponseWriter, r *http.Request) {
var in struct {
Seconds int `json:"seconds"`
}
_ = decode(r, &in, 64<<10)
if err := s.docker.Stop(r.Context(), chi.URLParam(r, "id"), in.Seconds); err != nil && !strings.Contains(err.Error(), "304") {
jsonOut(w, 502, map[string]string{"error": err.Error()})
return
}
jsonOut(w, 200, map[string]bool{"ok": true})
}
func (s *Service) removeWorker(w http.ResponseWriter, r *http.Request) {
if err := s.docker.Remove(r.Context(), chi.URLParam(r, "id")); err != nil {
jsonOut(w, 502, map[string]string{"error": err.Error()})
return
}
jsonOut(w, 200, map[string]bool{"ok": true})
}
func (s *Service) runningWorker(w http.ResponseWriter, r *http.Request) {
running, err := s.docker.Running(r.Context(), chi.URLParam(r, "id"))
if err != nil {
jsonOut(w, 502, map[string]string{"error": err.Error()})
return
}
jsonOut(w, 200, map[string]bool{"running": running})
}
func (s *Service) getWorkerFile(w http.ResponseWriter, r *http.Request) {
var in struct {
Path string `json:"path"`
}
if decode(r, &in, 64<<10) != nil || strings.TrimSpace(in.Path) == "" {
jsonOut(w, 400, map[string]string{"error": "path required"})
return
}
b, err := s.docker.GetFile(r.Context(), chi.URLParam(r, "id"), in.Path)
if err != nil {
jsonOut(w, 502, map[string]string{"error": err.Error()})
return
}
jsonOut(w, 200, map[string]any{"data": b})
}
func (s *Service) putWorkerFile(w http.ResponseWriter, r *http.Request) {
var in struct {
Dir string `json:"dir"`
Name string `json:"name"`
Data []byte `json:"data"`
}
if decode(r, &in, 2<<20) != nil || strings.TrimSpace(in.Dir) == "" || strings.TrimSpace(in.Name) == "" || len(in.Data) == 0 {
jsonOut(w, 400, map[string]string{"error": "dir/name/data required"})
return
}
if err := s.docker.PutFile(r.Context(), chi.URLParam(r, "id"), in.Dir, in.Name, in.Data); err != nil {
jsonOut(w, 502, map[string]string{"error": err.Error()})
return
}
jsonOut(w, 200, map[string]bool{"ok": true})
}
const adminCookie = "neuralhunt_service_controller_admin"
func (s *Service) adminCookieSecure(r *http.Request) bool {
switch strings.ToLower(strings.TrimSpace(s.cfg.AdminCookieSecureMode)) {
case "0", "false", "no", "off":
return false
case "1", "true", "yes", "on":
return true
case "", "auto":
return requestIsHTTPS(r)
default:
return true
}
}
func (s *Service) setAdminCookie(w http.ResponseWriter, r *http.Request, value string, ttl time.Duration) {
http.SetCookie(w, &http.Cookie{Name: adminCookie, Value: value, Path: "/", HttpOnly: true, Secure: s.adminCookieSecure(r), SameSite: http.SameSiteStrictMode, MaxAge: int(ttl.Seconds()), Expires: time.Now().UTC().Add(ttl)})
}
func (s *Service) clearAdminCookie(w http.ResponseWriter, r *http.Request) {
http.SetCookie(w, &http.Cookie{Name: adminCookie, Value: "", Path: "/", HttpOnly: true, Secure: s.adminCookieSecure(r), SameSite: http.SameSiteStrictMode, MaxAge: -1, Expires: time.Unix(1, 0).UTC()})
}
func (s *Service) requireAdmin(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
c, err := r.Cookie(adminCookie)
if err != nil {
jsonOut(w, 401, map[string]string{"error": "unauthorized"})
return
}
v, ok := s.adminSessions.Load(c.Value)
exp, _ := v.(time.Time)
if !ok || time.Now().After(exp) {
s.adminSessions.Delete(c.Value)
jsonOut(w, 401, map[string]string{"error": "unauthorized"})
return
}
next.ServeHTTP(w, r)
})
}
func (s *Service) AdminRoutes(ui http.Handler) http.Handler {
r := chi.NewRouter()
r.Use(s.security)
r.Post("/api/admin/login", s.adminLogin)
r.Post("/api/admin/logout", s.adminLogout)
r.Group(func(r chi.Router) {
r.Use(s.requireAdmin)
r.Get("/api/admin/overview", s.adminOverview)
r.Post("/api/admin/workers/{id}/start", s.adminStart)
r.Post("/api/admin/workers/{id}/stop", s.adminStop)
r.Post("/api/admin/workers/{id}/restart", s.adminRestart)
r.Post("/api/admin/workers/stop-all", s.adminStopAll)
r.Post("/api/admin/image/pull", s.adminPullImage)
})
r.Mount("/", ui)
return r
}
func (s *Service) adminLogin(w http.ResponseWriter, r *http.Request) {
var in struct{ Username, Password string }
if decode(r, &in, 64<<10) != nil {
jsonOut(w, 400, map[string]string{"error": "bad json"})
return
}
if subtle.ConstantTimeCompare([]byte(in.Username), []byte(s.cfg.AdminUser)) != 1 || subtle.ConstantTimeCompare([]byte(in.Password), []byte(s.cfg.AdminPassword)) != 1 {
time.Sleep(250 * time.Millisecond)
jsonOut(w, 401, map[string]string{"error": "invalid credentials"})
return
}
sid := customer.RandomToken(32)
s.adminSessions.Store(sid, time.Now().Add(12*time.Hour))
s.setAdminCookie(w, r, sid, 12*time.Hour)
jsonOut(w, 200, map[string]bool{"ok": true})
}
func (s *Service) adminLogout(w http.ResponseWriter, r *http.Request) {
if c, err := r.Cookie(adminCookie); err == nil {
s.adminSessions.Delete(c.Value)
}
s.clearAdminCookie(w, r)
jsonOut(w, 200, map[string]bool{"ok": true})
}
func (s *Service) adminOverview(w http.ResponseWriter, r *http.Request) {
items, err := s.docker.ManagedContainers(r.Context())
if err != nil {
jsonOut(w, 502, map[string]string{"error": err.Error()})
return
}
total, running := 0, 0
var workers []customer.ManagedContainer
for _, item := range items {
if item.WorkerID == "" {
continue
}
total++
if item.Running {
running++
}
workers = append(workers, item)
}
regOK, regAt, regErr := s.registrationStatus()
jsonOut(w, 200, map[string]any{"controller": map[string]any{"id": s.cfg.ID, "name": s.cfg.Name, "advertise_url": s.cfg.AdvertiseURL, "master_url": s.cfg.MasterURL, "max_workers": s.cfg.MaxWorkers, "max_running": s.cfg.MaxRunning, "worker_image": s.cfg.DefaultWorkerImage, "master_connected": regOK, "master_last_attempt": regAt, "master_last_error": regErr}, "workers": workers, "total": total, "running": running})
}
func (s *Service) adminStart(w http.ResponseWriter, r *http.Request) {
if err := s.docker.Start(r.Context(), chi.URLParam(r, "id")); err != nil {
jsonOut(w, 502, map[string]string{"error": err.Error()})
return
}
jsonOut(w, 200, map[string]bool{"ok": true})
}
func (s *Service) adminStop(w http.ResponseWriter, r *http.Request) {
if err := s.docker.Stop(r.Context(), chi.URLParam(r, "id"), 5); err != nil && !strings.Contains(err.Error(), "304") {
jsonOut(w, 502, map[string]string{"error": err.Error()})
return
}
jsonOut(w, 200, map[string]bool{"ok": true})
}
func (s *Service) adminRestart(w http.ResponseWriter, r *http.Request) {
id := chi.URLParam(r, "id")
_ = s.docker.Stop(r.Context(), id, 5)
if err := s.docker.Start(r.Context(), id); err != nil {
jsonOut(w, 502, map[string]string{"error": err.Error()})
return
}
jsonOut(w, 200, map[string]bool{"ok": true})
}
func (s *Service) adminStopAll(w http.ResponseWriter, r *http.Request) {
items, err := s.docker.ManagedContainers(r.Context())
if err != nil {
jsonOut(w, 502, map[string]string{"error": err.Error()})
return
}
stopped := 0
var warnings []string
for _, item := range items {
if item.WorkerID == "" || !item.Running {
continue
}
if err := s.docker.Stop(r.Context(), item.ID, 5); err != nil && !strings.Contains(err.Error(), "304") {
warnings = append(warnings, item.WorkerID+": "+err.Error())
continue
}
stopped++
}
jsonOut(w, 200, map[string]any{"ok": len(warnings) == 0, "stopped": stopped, "warnings": warnings})
}
func (s *Service) adminPullImage(w http.ResponseWriter, r *http.Request) {
image := strings.TrimSpace(s.cfg.DefaultWorkerImage)
if image == "" {
jsonOut(w, 400, map[string]string{"error": "SC_WORKER_IMAGE is empty"})
return
}
auth, err := customer.RegistryAuthHeader(s.cfg.RegistryUsername, s.cfg.RegistryPassword, s.cfg.RegistryServer)
if err != nil {
jsonOut(w, 400, map[string]string{"error": err.Error()})
return
}
if err := s.docker.PullImage(r.Context(), image, auth); err != nil {
jsonOut(w, 502, map[string]string{"error": err.Error()})
return
}
jsonOut(w, 200, map[string]any{"ok": true, "image": image})
}
func (s *Service) setRegistrationStatus(ok bool, err error) {
s.regMu.Lock()
defer s.regMu.Unlock()
s.lastRegOK = ok
s.lastRegAt = time.Now().UTC()
if err != nil {
s.lastRegErr = err.Error()
} else {
s.lastRegErr = ""
}
}
func (s *Service) registrationStatus() (bool, time.Time, string) {
s.regMu.RLock()
defer s.regMu.RUnlock()
return s.lastRegOK, s.lastRegAt, s.lastRegErr
}
func (s *Service) registerOnce(ctx context.Context) error {
total, running, err := s.counts(ctx)
if err != nil {
return err
}
payload := map[string]any{"id": s.cfg.ID, "name": s.cfg.Name, "base_url": strings.TrimRight(s.cfg.AdvertiseURL, "/"), "protocol_version": 1, "max_workers": s.cfg.MaxWorkers, "max_running": s.cfg.MaxRunning, "reported_workers": total, "reported_running": running}
b, _ := json.Marshal(payload)
endpoint := strings.TrimRight(s.cfg.MasterURL, "/") + "/internal/controllers/register"
req, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint, bytes.NewReader(b))
if err != nil {
return err
}
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", "Bearer "+s.cfg.SharedSecret)
resp, err := s.hc.Do(req)
if err != nil {
return err
}
defer resp.Body.Close()
body, _ := io.ReadAll(io.LimitReader(resp.Body, 16<<10))
if resp.StatusCode/100 != 2 {
return fmt.Errorf("master register HTTP %d: %s", resp.StatusCode, strings.TrimSpace(string(body)))
}
return nil
}
func (s *Service) RunRegistration(ctx context.Context) {
ticker := time.NewTicker(s.cfg.HeartbeatInterval)
defer ticker.Stop()
for {
callCtx, cancel := context.WithTimeout(ctx, 8*time.Second)
err := s.registerOnce(callCtx)
cancel()
s.setRegistrationStatus(err == nil, err)
if err != nil && !errors.Is(err, context.Canceled) {
log.Printf("service-controller registration: %v", err)
}
select {
case <-ctx.Done():
return
case <-ticker.C:
}
}
}
func ValidateConfig(cfg Config) error {
if strings.TrimSpace(cfg.ID) == "" {
return errors.New("SC_ID is required and must remain stable for this host")
}
if strings.TrimSpace(cfg.MasterURL) == "" {
return errors.New("SC_MASTER_URL is required")
}
if strings.TrimSpace(cfg.AdvertiseURL) == "" {
return errors.New("SC_ADVERTISE_URL is required")
}
for _, raw := range []string{cfg.MasterURL, cfg.AdvertiseURL} {
u, err := url.Parse(raw)
if err != nil || u.Host == "" || (u.Scheme != "http" && u.Scheme != "https") {
return fmt.Errorf("invalid controller URL %q", raw)
}
}
secret := strings.TrimSpace(cfg.SharedSecret)
if len(secret) < 32 || strings.Contains(strings.ToLower(secret), "replace-with") || strings.Contains(strings.ToLower(secret), "change-me") {
return errors.New("SERVICE_CONTROLLER_SHARED_SECRET must be a unique random value of at least 32 characters")
}
adminPass := strings.TrimSpace(cfg.AdminPassword)
if len(adminPass) < 16 || strings.Contains(strings.ToLower(adminPass), "replace-with") || strings.Contains(strings.ToLower(adminPass), "change-me") {
return errors.New("SC_ADMIN_PASSWORD must be a unique value of at least 16 characters")
}
return nil
}