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 }