package controlplane import ( "context" "crypto/rand" "crypto/sha256" "crypto/subtle" "crypto/tls" "encoding/base64" "encoding/hex" "encoding/json" "errors" "fmt" "io" "io/fs" "mime" "net/http" "net/url" "path" "sort" "strconv" "strings" "sync" "time" webassets "git.ipao.vip/rogee/wx-win-agent/control-plane/web" ) type ServerConfig struct { ListenAddr string DataFile string NodeTokens map[string]string WebUsers map[string]string LeaseTTL time.Duration HeartbeatTimeout time.Duration SessionTTL time.Duration TLSCertFile string TLSKeyFile string MTLSClientCAFile string MTLSRequireNodeCert bool MTLSRevokedCertsFile string BackupDir string BackupCount int BackupInterval time.Duration TaskRetention time.Duration EventRetention time.Duration AuditRetention time.Duration } func DefaultServerConfig() ServerConfig { return ServerConfig{ ListenAddr: "127.0.0.1:8090", DataFile: "control-plane-data.json", NodeTokens: map[string]string{}, WebUsers: map[string]string{}, LeaseTTL: 30 * time.Second, HeartbeatTimeout: 45 * time.Second, SessionTTL: 8 * time.Hour, BackupCount: 7, BackupInterval: 5 * time.Minute, TaskRetention: 30 * 24 * time.Hour, EventRetention: 30 * 24 * time.Hour, AuditRetention: 90 * 24 * time.Hour, } } type Server struct { config ServerConfig store *Store tlsConfig *tls.Config nodeTokenHashes map[string][32]byte userPasswords map[string][32]byte sessionMu sync.Mutex sessions map[string]session } type session struct { Username string Expires time.Time } type requestError struct { status int code string message string } func (e requestError) Error() string { return e.message } func NewServer(config ServerConfig) (*Server, error) { defaults := DefaultServerConfig() if config.ListenAddr == "" { config.ListenAddr = defaults.ListenAddr } if config.DataFile == "" { config.DataFile = defaults.DataFile } if config.BackupDir == "" { config.BackupDir = config.DataFile + ".backups" } if config.BackupCount == 0 { config.BackupCount = defaults.BackupCount } if config.BackupInterval == 0 { config.BackupInterval = defaults.BackupInterval } if config.TaskRetention == 0 { config.TaskRetention = defaults.TaskRetention } if config.EventRetention == 0 { config.EventRetention = defaults.EventRetention } if config.AuditRetention == 0 { config.AuditRetention = defaults.AuditRetention } if config.LeaseTTL <= 0 { config.LeaseTTL = defaults.LeaseTTL } if config.HeartbeatTimeout <= 0 { config.HeartbeatTimeout = defaults.HeartbeatTimeout } if config.SessionTTL <= 0 { config.SessionTTL = defaults.SessionTTL } if config.NodeTokens == nil { config.NodeTokens = map[string]string{} } if config.WebUsers == nil { config.WebUsers = map[string]string{} } if config.BackupCount < 0 || config.BackupInterval < 0 || config.TaskRetention < 0 || config.EventRetention < 0 || config.AuditRetention < 0 { return nil, errors.New("backup count and retention settings must be non-negative") } tlsConfig, err := newTLSConfig(config) if err != nil { return nil, err } store, err := OpenStore(config.DataFile, StoreOptions{ BackupDir: config.BackupDir, BackupCount: config.BackupCount, BackupInterval: config.BackupInterval, TaskRetention: config.TaskRetention, EventRetention: config.EventRetention, AuditRetention: config.AuditRetention, }) if err != nil { return nil, err } s := &Server{ config: config, store: store, tlsConfig: tlsConfig, nodeTokenHashes: map[string][32]byte{}, userPasswords: map[string][32]byte{}, sessions: map[string]session{}, } for nodeID, token := range config.NodeTokens { if validIdentifier(nodeID, 200) && token != "" { s.nodeTokenHashes[nodeID] = sha256.Sum256([]byte(token)) } } for username, password := range config.WebUsers { if username != "" && password != "" { s.userPasswords[username] = sha256.Sum256([]byte(password)) } } return s, nil } func (s *Server) Handler() http.Handler { return s } func (s *Server) ListenAndServe(ctx context.Context) error { server := &http.Server{Addr: s.config.ListenAddr, Handler: s, ReadHeaderTimeout: 5 * time.Second, ReadTimeout: 15 * time.Second, WriteTimeout: 15 * time.Second, IdleTimeout: 60 * time.Second} go func() { <-ctx.Done() shutdownCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() _ = server.Shutdown(shutdownCtx) }() defer func() { _ = s.Close() }() var err error if s.tlsConfig != nil { server.TLSConfig = s.tlsConfig err = server.ListenAndServeTLS("", "") } else { err = server.ListenAndServe() } if errors.Is(err, http.ErrServerClosed) { return nil } return err } // Close releases the active/passive store lock. func (s *Server) Close() error { return s.store.Close() } func (s *Server) ServeHTTP(w http.ResponseWriter, r *http.Request) { correlationID := r.Header.Get("X-Correlation-Id") if !validIdentifier(correlationID, 128) { correlationID = randomID() } w.Header().Set("X-Correlation-Id", correlationID) w.Header().Set("Cache-Control", "no-store") w.Header().Set("X-Content-Type-Options", "nosniff") w.Header().Set("Referrer-Policy", "no-referrer") if r.Method == http.MethodOptions { w.WriteHeader(http.StatusNoContent) return } var err error switch { case r.URL.Path == "/" && r.Method == http.MethodGet: s.serveFrontend(w, "dist/index.html") return case strings.HasPrefix(r.URL.Path, "/assets/") && r.Method == http.MethodGet: s.serveFrontend(w, path.Join("dist", strings.TrimPrefix(r.URL.Path, "/"))) return case r.URL.Path == "/healthz" && r.Method == http.MethodGet: writeJSON(w, http.StatusOK, map[string]any{"status": "ok", "protocol_version": ProtocolVersion, "correlation_id": correlationID}) return case r.URL.Path == "/readyz" && r.Method == http.MethodGet: err = s.ready(w, correlationID) case r.URL.Path == "/v1/auth/login" && r.Method == http.MethodPost: err = s.login(w, r, correlationID) case r.URL.Path == "/v1/nodes" && r.Method == http.MethodGet: err = s.listNodes(w, r) case r.URL.Path == "/v1/nodes/register" && r.Method == http.MethodPost: err = s.registerNode(w, r, correlationID) case strings.HasPrefix(r.URL.Path, "/v1/nodes/"): err = s.nodeRoute(w, r, correlationID) case strings.HasPrefix(r.URL.Path, "/v1/reads/"): err = s.readRoute(w, r, correlationID) case r.URL.Path == "/v1/tasks" || strings.HasPrefix(r.URL.Path, "/v1/tasks/"): err = s.taskRoute(w, r, correlationID) case r.URL.Path == "/v1/events": err = s.eventRoute(w, r, correlationID) case r.URL.Path == "/v1/audit" && r.Method == http.MethodGet: err = s.auditRoute(w, r, correlationID) default: err = requestError{status: http.StatusNotFound, code: "NotFound", message: "Resource was not found."} } if err != nil { writeError(w, err, correlationID) } } func (s *Server) ready(w http.ResponseWriter, correlationID string) error { if err := s.store.Read(func(PersistedState) error { return nil }); err != nil { return requestError{status: http.StatusServiceUnavailable, code: "NotReady", message: "The control plane store is not ready."} } writeJSON(w, http.StatusOK, map[string]any{"status": "ready", "protocol_version": ProtocolVersion, "correlation_id": correlationID}) return nil } func (s *Server) login(w http.ResponseWriter, r *http.Request, correlationID string) error { var request struct { Username string `json:"username"` Password string `json:"password"` } if err := decodeJSON(r, &request, 8*1024); err != nil { return err } password, exists := s.userPasswords[request.Username] provided := sha256.Sum256([]byte(request.Password)) if !exists || subtle.ConstantTimeCompare(password[:], provided[:]) != 1 { return requestError{status: http.StatusUnauthorized, code: "Unauthorized", message: "Authentication failed."} } token := randomID() s.sessionMu.Lock() s.sessions[token] = session{Username: request.Username, Expires: time.Now().UTC().Add(s.config.SessionTTL)} s.sessionMu.Unlock() _ = s.appendAudit(request.Username, "login", "session", correlationID, "success") writeJSON(w, http.StatusOK, map[string]any{"access_token": token, "token_type": "Bearer", "expires_in": int(s.config.SessionTTL.Seconds()), "correlation_id": correlationID}) return nil } func (s *Server) registerNode(w http.ResponseWriter, r *http.Request, correlationID string) error { nodeID, err := s.authenticateNode(r) if err != nil { return err } var request NodeRegistration if err := decodeJSON(r, &request, 128*1024); err != nil { return err } if request.NodeID != nodeID { return requestError{status: http.StatusForbidden, code: "NodeIdentityMismatch", message: "The token is not assigned to this node."} } if !validIdentifier(request.AgentVersion, 80) || request.ProtocolVersion != ProtocolVersion || len(request.Capabilities) > 100 || !validCapabilities(request.Capabilities) || !validAccountSummaries(request.Accounts) { return requestError{status: http.StatusBadRequest, code: "InvalidRegistration", message: "Node registration is invalid."} } now := time.Now().UTC() if err := s.store.Mutate(func(state *PersistedState) error { node := state.Nodes[nodeID] node.NodeID = nodeID node.AgentVersion = request.AgentVersion node.ProtocolVersion = request.ProtocolVersion node.Capabilities = append([]string(nil), request.Capabilities...) node.Status = NodeOnline node.LastHeartbeatAt = &now node.ReportingConfigVersion = request.ReportingConfigVersion node.Accounts = append([]AccountSummary(nil), request.Accounts...) node.LastCorrelationID = correlationID state.Nodes[nodeID] = node state.Audit = appendAudit(state.Audit, "node:"+nodeID, "node.register", nodeID, correlationID, "success", now) return nil }); err != nil { return requestError{status: http.StatusInternalServerError, code: "PersistenceFailed", message: "The control plane could not persist the node registration."} } writeJSON(w, http.StatusOK, NodeResponse{NodeID: nodeID, Status: NodeOnline, Authenticated: true, CorrelationID: correlationID}) return nil } func (s *Server) nodeRoute(w http.ResponseWriter, r *http.Request, correlationID string) error { parts := pathParts(r.URL.Path) if len(parts) < 3 || parts[0] != "v1" || parts[1] != "nodes" { return requestError{status: http.StatusNotFound, code: "NotFound", message: "Resource was not found."} } nodeID, err := s.authenticateNode(r) if err != nil { return err } if nodeID != parts[2] { return requestError{status: http.StatusForbidden, code: "NodeIdentityMismatch", message: "The token is not assigned to this node."} } if len(parts) == 3 && r.Method == http.MethodGet { return s.nodeDetails(w, nodeID, correlationID) } if len(parts) < 4 { return requestError{status: http.StatusNotFound, code: "NotFound", message: "Resource was not found."} } switch parts[3] { case "heartbeat": if r.Method != http.MethodPost || len(parts) != 4 { return requestError{status: http.StatusMethodNotAllowed, code: "MethodNotAllowed", message: "Method is not allowed."} } return s.heartbeat(w, r, nodeID, correlationID) case "tasks": return s.nodeTaskRoute(w, r, nodeID, parts[4:], correlationID) case "events": if r.Method != http.MethodPost || len(parts) != 4 { return requestError{status: http.StatusMethodNotAllowed, code: "MethodNotAllowed", message: "Method is not allowed."} } return s.ingestEvent(w, r, nodeID, correlationID) default: return requestError{status: http.StatusNotFound, code: "NotFound", message: "Resource was not found."} } } func (s *Server) listNodes(w http.ResponseWriter, r *http.Request) error { if _, err := s.authenticateWeb(r); err != nil { return err } var nodes []Node if err := s.store.Mutate(func(state *PersistedState) error { now := time.Now().UTC() for nodeID, value := range state.Nodes { node := value nodeChanged := false if node.LastHeartbeatAt == nil || now.Sub(*node.LastHeartbeatAt) > s.config.HeartbeatTimeout { if node.Status != NodeOffline { node.Status = NodeOffline nodeChanged = true } } if nodeChanged { state.Nodes[nodeID] = node } nodes = append(nodes, node) } return nil }); err != nil { return err } sort.Slice(nodes, func(i, j int) bool { return nodes[i].NodeID < nodes[j].NodeID }) writeJSON(w, http.StatusOK, map[string]any{"nodes": nodes}) return nil } func (s *Server) nodeDetails(w http.ResponseWriter, nodeID, _ string) error { var node Node if err := s.store.Read(func(state PersistedState) error { value, ok := state.Nodes[nodeID] if !ok { return requestError{status: http.StatusNotFound, code: "NodeNotFound", message: "Node was not registered."} } node = value return nil }); err != nil { return err } writeJSON(w, http.StatusOK, node) return nil } func (s *Server) heartbeat(w http.ResponseWriter, r *http.Request, nodeID, correlationID string) error { var request Heartbeat if err := decodeJSON(r, &request, 32*1024); err != nil { return err } if request.NodeID != nodeID || request.ProtocolVersion != ProtocolVersion || !validIdentifier(request.AgentVersion, 80) || !validNodeStatus(request.NodeStatus) || request.QueueLength < 0 || request.QueueLength > 100000 { return requestError{status: http.StatusBadRequest, code: "InvalidHeartbeat", message: "Heartbeat is invalid."} } now := time.Now().UTC() if err := s.store.Mutate(func(state *PersistedState) error { node := state.Nodes[nodeID] if node.NodeID == "" { return requestError{status: http.StatusConflict, code: "NodeNotRegistered", message: "Register the node before sending a heartbeat."} } node.AgentVersion = request.AgentVersion node.ProtocolVersion = request.ProtocolVersion node.Status = request.NodeStatus node.WechatRunning = request.WechatRunning node.WechatLoggedIn = request.WechatLoggedIn node.SessionLocked = request.SessionLocked node.ActiveAccountID = request.ActiveAccountID node.QueueLength = request.QueueLength node.ReportingConfigVersion = request.ReportingConfigVersion node.LastErrorCode = request.LastErrorCode node.LastHeartbeatAt = &now node.LastCorrelationID = correlationID state.Nodes[nodeID] = node state.Audit = appendAudit(state.Audit, "node:"+nodeID, "node.heartbeat", nodeID, correlationID, "success", now) return nil }); err != nil { return err } writeJSON(w, http.StatusOK, NodeResponse{NodeID: nodeID, Status: request.NodeStatus, LastHeartbeatAt: &now, CorrelationID: correlationID}) return nil } func (s *Server) nodeTaskRoute(w http.ResponseWriter, r *http.Request, nodeID string, parts []string, correlationID string) error { if len(parts) == 0 && r.Method == http.MethodGet { return s.pollTasks(w, r, nodeID, correlationID) } if len(parts) != 2 || r.Method != http.MethodPost { return requestError{status: http.StatusNotFound, code: "NotFound", message: "Resource was not found."} } taskID := parts[0] switch parts[1] { case "ack": return s.ackTask(w, r, nodeID, taskID, correlationID) case "start": return s.startTask(w, r, nodeID, taskID, correlationID) case "renew": return s.renewTask(w, r, nodeID, taskID, correlationID) case "result": return s.recordTaskResult(w, r, nodeID, taskID, correlationID) default: return requestError{status: http.StatusNotFound, code: "NotFound", message: "Resource was not found."} } } func (s *Server) pollTasks(w http.ResponseWriter, r *http.Request, nodeID, correlationID string) error { waitSeconds := queryWaitSeconds(r.URL.Query().Get("wait_seconds")) deadline := time.Now().Add(time.Duration(waitSeconds) * time.Second) for { tasks, err := s.claimTasks(r, nodeID, correlationID) if err != nil { return err } if len(tasks) > 0 || waitSeconds == 0 || !time.Now().Before(deadline) { writeJSON(w, http.StatusOK, TaskBatch{Tasks: tasks}) return nil } select { case <-r.Context().Done(): return nil case <-time.After(250 * time.Millisecond): } } } func (s *Server) claimTasks(r *http.Request, nodeID, correlationID string) ([]Task, error) { accountID := r.URL.Query().Get("account_id") if accountID != "" && !validIdentifier(accountID, 200) { return nil, requestError{status: http.StatusBadRequest, code: "InvalidAccount", message: "account_id is invalid."} } now := time.Now().UTC() var tasks []Task if err := s.store.Mutate(func(state *PersistedState) error { for taskID, value := range state.Tasks { task := value if task.NodeID != nodeID || (accountID != "" && task.AccountID != accountID) || terminal(task.Status) { continue } changed := false if task.NotAfter != nil && !now.Before(*task.NotAfter) { if task.Status == TaskPending { task.Status = TaskExpired task.StateVersion++ task.UpdatedAt = now task.LeaseOwner = "" task.LeaseExpiresAt = nil changed = true } else if task.Status == TaskAccepted || task.Status == TaskRunning { markUnconfirmed(&task, "TaskDeadlineReached", now) changed = true } } if terminal(task.Status) { if changed { state.Tasks[taskID] = task } continue } if task.CancelRequestedAt != nil && task.Status == TaskPending && !leaseActive(task, now) { task.Status = TaskCancelled task.StateVersion++ task.UpdatedAt = now task.LeaseOwner = "" task.LeaseExpiresAt = nil changed = true } if terminal(task.Status) { if changed { state.Tasks[taskID] = task } continue } if task.Status == TaskPending { if leaseActive(task, now) && task.LeaseOwner != nodeID { continue } if !leaseActive(task, now) { task.LeaseGeneration++ expires := now.Add(s.config.LeaseTTL) task.LeaseExpiresAt = &expires task.LeaseOwner = nodeID task.UpdatedAt = now task.LastCorrelationID = correlationID changed = true } } if task.Status == TaskAccepted || task.Status == TaskRunning { if task.LeaseOwner != nodeID || !leaseActive(task, now) { if task.LeaseOwner == nodeID && !leaseActive(task, now) { markUnconfirmed(&task, "LeaseExpired", now) changed = true } if changed { state.Tasks[taskID] = task } continue } } if changed { state.Tasks[taskID] = task } if task.Status == TaskPending || task.Status == TaskAccepted || task.Status == TaskRunning { tasks = append(tasks, task) } } state.Audit = appendAudit(state.Audit, "node:"+nodeID, "task.poll", nodeID, correlationID, "success", now) return nil }); err != nil { return nil, err } sort.Slice(tasks, func(i, j int) bool { return tasks[i].CreatedAt.Before(tasks[j].CreatedAt) }) return tasks, nil } func (s *Server) ackTask(w http.ResponseWriter, r *http.Request, nodeID, taskID, correlationID string) error { var request TaskAck if err := decodeJSON(r, &request, 8*1024); err != nil { return err } if request.TaskID != taskID { return requestError{status: http.StatusBadRequest, code: "TaskIdentityMismatch", message: "Task identity does not match the path."} } var task Task err := s.store.Mutate(func(state *PersistedState) error { value, ok := state.Tasks[taskID] if !ok || value.NodeID != nodeID { return requestError{status: http.StatusNotFound, code: "TaskNotFound", message: "Task was not found."} } task = value if task.AccountID != request.AccountID { return requestError{status: http.StatusForbidden, code: "AccountMismatch", message: "Task account does not match."} } if terminal(task.Status) { return nil } if task.LeaseGeneration != request.LeaseGeneration || task.LeaseOwner != nodeID || !leaseActive(task, time.Now().UTC()) { return requestError{status: http.StatusConflict, code: "LeaseMismatch", message: "The task lease is no longer valid."} } if task.Status != TaskPending { if task.Status == TaskAccepted { return nil } return requestError{status: http.StatusConflict, code: "InvalidTaskState", message: "The task cannot be accepted in its current state."} } if task.CancelRequestedAt != nil { task.Status = TaskCancelled task.StateVersion++ task.LeaseOwner = "" task.LeaseExpiresAt = nil } else { task.Status = TaskAccepted task.StateVersion++ } task.UpdatedAt = time.Now().UTC() task.LastCorrelationID = correlationID state.Tasks[taskID] = task state.Audit = appendAudit(state.Audit, "node:"+nodeID, "task.accept", taskID, correlationID, "success", task.UpdatedAt) return nil }) if err != nil { return err } writeJSON(w, http.StatusOK, task) return nil } func (s *Server) startTask(w http.ResponseWriter, r *http.Request, nodeID, taskID, correlationID string) error { var request TaskAck if err := decodeJSON(r, &request, 8*1024); err != nil { return err } if request.TaskID != taskID { return requestError{status: http.StatusBadRequest, code: "TaskIdentityMismatch", message: "Task identity does not match the path."} } var task Task err := s.store.Mutate(func(state *PersistedState) error { value, ok := state.Tasks[taskID] if !ok || value.NodeID != nodeID { return requestError{status: http.StatusNotFound, code: "TaskNotFound", message: "Task was not found."} } task = value if task.AccountID != request.AccountID || task.LeaseGeneration != request.LeaseGeneration || task.LeaseOwner != nodeID { return requestError{status: http.StatusConflict, code: "LeaseMismatch", message: "The task lease is no longer valid."} } now := time.Now().UTC() if terminal(task.Status) { return nil } if task.Status != TaskAccepted { return requestError{status: http.StatusConflict, code: "InvalidTaskState", message: "The task must be accepted before execution."} } if task.CancelRequestedAt != nil { task.Status = TaskCancelled task.StateVersion++ task.LeaseOwner = "" task.LeaseExpiresAt = nil } else if task.NotAfter != nil && !now.Before(*task.NotAfter) { task.Status = TaskExpired task.StateVersion++ task.LeaseOwner = "" task.LeaseExpiresAt = nil } else if !leaseActive(task, now) { return requestError{status: http.StatusConflict, code: "LeaseExpired", message: "The task lease expired before execution started."} } else { task.Status = TaskRunning task.StateVersion++ } task.UpdatedAt = now task.LastCorrelationID = correlationID state.Tasks[taskID] = task state.Audit = appendAudit(state.Audit, "node:"+nodeID, "task.start", taskID, correlationID, "success", now) return nil }) if err != nil { return err } writeJSON(w, http.StatusOK, task) return nil } func (s *Server) renewTask(w http.ResponseWriter, r *http.Request, nodeID, taskID, correlationID string) error { var request TaskAck if err := decodeJSON(r, &request, 8*1024); err != nil { return err } if request.TaskID != taskID { return requestError{status: http.StatusBadRequest, code: "TaskIdentityMismatch", message: "Task identity does not match the path."} } var task Task if err := s.store.Mutate(func(state *PersistedState) error { value, ok := state.Tasks[taskID] if !ok || value.NodeID != nodeID { return requestError{status: http.StatusNotFound, code: "TaskNotFound", message: "Task was not found."} } task = value if task.AccountID != request.AccountID || task.LeaseGeneration != request.LeaseGeneration || task.LeaseOwner != nodeID { return requestError{status: http.StatusConflict, code: "LeaseMismatch", message: "The task lease is no longer valid."} } if task.CancelRequestedAt != nil { return requestError{status: http.StatusConflict, code: "CancelRequested", message: "The task lease cannot be renewed after cancellation was requested."} } now := time.Now().UTC() if task.Status != TaskAccepted && task.Status != TaskRunning || !leaseActive(task, now) { return requestError{status: http.StatusConflict, code: "LeaseExpired", message: "The task lease is no longer renewable."} } expires := now.Add(s.config.LeaseTTL) task.LeaseExpiresAt = &expires task.StateVersion++ task.UpdatedAt = now task.LastCorrelationID = correlationID state.Tasks[taskID] = task state.Audit = appendAudit(state.Audit, "node:"+nodeID, "task.renew", taskID, correlationID, "success", now) return nil }); err != nil { return err } writeJSON(w, http.StatusOK, task) return nil } func (s *Server) recordTaskResult(w http.ResponseWriter, r *http.Request, nodeID, taskID, correlationID string) error { var result TaskResult if err := decodeJSON(r, &result, MaxTaskResultBytes+32*1024); err != nil { return err } if result.TaskID != taskID || !terminal(result.Status) { return requestError{status: http.StatusBadRequest, code: "InvalidTaskResult", message: "Task result or terminal status is invalid."} } if len(result.Content) > MaxTaskResultBytes { return requestError{status: http.StatusRequestEntityTooLarge, code: "ContentTooLarge", message: "The task result content is too large."} } var task Task err := s.store.Mutate(func(state *PersistedState) error { value, ok := state.Tasks[taskID] if !ok || value.NodeID != nodeID { return requestError{status: http.StatusNotFound, code: "TaskNotFound", message: "Task was not found."} } task = value if task.AccountID != result.AccountID { return requestError{status: http.StatusForbidden, code: "AccountMismatch", message: "Task account does not match."} } if hasTaskContent(result.Content) && !isReadTaskKind(task.Kind) { return requestError{status: http.StatusForbidden, code: "ContentNotAllowed", message: "Only read task results may contain content."} } if terminal(task.Status) { if task.Result != nil && task.Result.Status == result.Status && task.Result.ErrorCode == result.ErrorCode { return nil } return requestError{status: http.StatusConflict, code: "TerminalState", message: "The task already has a different terminal result."} } if task.LeaseGeneration != result.LeaseGeneration || task.LeaseOwner != nodeID { return requestError{status: http.StatusConflict, code: "LeaseMismatch", message: "The result belongs to an expired task lease."} } if task.Status != TaskAccepted && task.Status != TaskRunning { return requestError{status: http.StatusConflict, code: "InvalidTaskState", message: "The task is not awaiting a final result."} } result.CorrelationID = firstNonEmpty(result.CorrelationID, correlationID) task.Result = &result task.Status = result.Status task.StateVersion++ task.UpdatedAt = time.Now().UTC() task.LeaseOwner = "" task.LeaseExpiresAt = nil task.LastCorrelationID = correlationID state.Tasks[taskID] = task state.Audit = appendAudit(state.Audit, "node:"+nodeID, "task.result", taskID, correlationID, "success", task.UpdatedAt) return nil }) if err != nil { return err } writeJSON(w, http.StatusOK, task) return nil } func (s *Server) readRoute(w http.ResponseWriter, r *http.Request, correlationID string) error { username, err := s.authenticateWeb(r) if err != nil { return err } parts := pathParts(r.URL.Path) if len(parts) != 3 || parts[0] != "v1" || parts[1] != "reads" || r.Method != http.MethodPost { return requestError{status: http.StatusNotFound, code: "NotFound", message: "Resource was not found."} } kind := "read-" + parts[2] if !isReadTaskKind(kind) { return requestError{status: http.StatusNotFound, code: "NotFound", message: "Resource was not found."} } var request ReadTaskSubmission if err := decodeJSON(r, &request, 64*1024); err != nil { return err } payload, err := readTaskPayload(kind, request) if err != nil { return err } return s.createTaskSubmission(w, TaskSubmission{ NodeID: request.NodeID, AccountID: request.AccountID, Kind: kind, IdempotencyKey: request.IdempotencyKey, Payload: payload, NotAfter: request.NotAfter, }, username, correlationID) } type readSessionsPayload struct { Limit int `json:"limit"` Offset int `json:"offset"` } type readContactsPayload struct { Limit int `json:"limit"` Offset int `json:"offset"` Contains string `json:"contains,omitempty"` GroupsOnly bool `json:"groups_only"` } type readMessagesPayload struct { Limit int `json:"limit"` Offset int `json:"offset"` ChatID string `json:"chat_id"` IncludeContent bool `json:"include_content"` } func isReadTaskKind(kind string) bool { switch kind { case "read-sessions", "read-contacts", "read-messages": return true default: return false } } func readTaskPayload(kind string, request ReadTaskSubmission) (jsonRaw, error) { limit := request.Limit if limit == 0 { limit = 50 } if limit < 1 || limit > 200 || request.Offset < 0 { return nil, requestError{status: http.StatusBadRequest, code: "InvalidPagination", message: "limit must be 1..200 and offset must be non-negative."} } if request.Contains != "" && !validIdentifier(request.Contains, 200) { return nil, requestError{status: http.StatusBadRequest, code: "InvalidReadRequest", message: "contains is invalid."} } switch kind { case "read-sessions": if request.Contains != "" || request.GroupsOnly != nil || request.ChatID != "" || request.IncludeContent { return nil, requestError{status: http.StatusBadRequest, code: "InvalidReadRequest", message: "The sessions read request contains unsupported fields."} } return json.Marshal(readSessionsPayload{Limit: limit, Offset: request.Offset}) case "read-contacts": if request.GroupsOnly == nil || request.ChatID != "" || request.IncludeContent { return nil, requestError{status: http.StatusBadRequest, code: "InvalidReadRequest", message: "contacts requires groups_only and does not accept chat_id or include_content."} } return json.Marshal(readContactsPayload{Limit: limit, Offset: request.Offset, Contains: request.Contains, GroupsOnly: *request.GroupsOnly}) case "read-messages": if !validIdentifier(request.ChatID, 512) || request.GroupsOnly != nil || request.Contains != "" { return nil, requestError{status: http.StatusBadRequest, code: "InvalidReadRequest", message: "messages requires a stable chat_id and does not accept contains or groups_only."} } return json.Marshal(readMessagesPayload{Limit: limit, Offset: request.Offset, ChatID: request.ChatID, IncludeContent: request.IncludeContent}) default: return nil, requestError{status: http.StatusNotFound, code: "NotFound", message: "Resource was not found."} } } func validTaskPayload(kind string, payload jsonRaw) bool { if kind == "send-text" { return validSendTextPayload(payload) } if !isReadTaskKind(kind) || len(payload) == 0 || len(payload) > 64*1024 || !json.Valid(payload) { return false } switch kind { case "read-sessions": var value readSessionsPayload return decodeRaw(payload, &value) && validPagination(value.Limit, value.Offset) case "read-contacts": var value readContactsPayload return decodeRaw(payload, &value) && validPagination(value.Limit, value.Offset) && (value.Contains == "" || validIdentifier(value.Contains, 200)) case "read-messages": var value readMessagesPayload return decodeRaw(payload, &value) && validPagination(value.Limit, value.Offset) && validIdentifier(value.ChatID, 512) default: return false } } func decodeRaw(payload jsonRaw, target any) bool { decoder := json.NewDecoder(strings.NewReader(string(payload))) decoder.DisallowUnknownFields() if decoder.Decode(target) != nil { return false } var extra any return decoder.Decode(&extra) == io.EOF } func validPagination(limit, offset int) bool { return limit >= 1 && limit <= 200 && offset >= 0 } func hasTaskContent(content jsonRaw) bool { return len(content) != 0 && string(content) != "null" } func (s *Server) taskRoute(w http.ResponseWriter, r *http.Request, correlationID string) error { username, err := s.authenticateWeb(r) if err != nil { return err } parts := pathParts(r.URL.Path) if len(parts) == 2 && parts[0] == "v1" && parts[1] == "tasks" { switch r.Method { case http.MethodGet: return s.listTasks(w, r) case http.MethodPost: return s.createTask(w, r, username, correlationID) default: return requestError{status: http.StatusMethodNotAllowed, code: "MethodNotAllowed", message: "Method is not allowed."} } } if len(parts) == 3 && r.Method == http.MethodGet { return s.getTask(w, parts[2]) } if len(parts) == 4 && parts[3] == "cancel" && r.Method == http.MethodPost { return s.cancelTask(w, parts[2], username, correlationID) } return requestError{status: http.StatusNotFound, code: "NotFound", message: "Resource was not found."} } func (s *Server) createTask(w http.ResponseWriter, r *http.Request, username, correlationID string) error { var request TaskSubmission if err := decodeJSON(r, &request, 128*1024); err != nil { return err } return s.createTaskSubmission(w, request, username, correlationID) } func (s *Server) createTaskSubmission(w http.ResponseWriter, request TaskSubmission, username, correlationID string) error { if !validIdentifier(request.NodeID, 200) || !validIdentifier(request.AccountID, 200) || !validIdentifier(request.Kind, 80) || !validIdentifier(request.IdempotencyKey, 128) { return requestError{status: http.StatusBadRequest, code: "InvalidTask", message: "Task identity fields are invalid."} } if !validTaskPayload(request.Kind, request.Payload) { return requestError{status: http.StatusBadRequest, code: "UnsupportedTask", message: "The task kind or payload is not supported."} } if request.NotAfter != nil && !request.NotAfter.After(time.Now().UTC()) { return requestError{status: http.StatusBadRequest, code: "InvalidDeadline", message: "not_after must be in the future."} } var response TaskSubmissionResponse err := s.store.Mutate(func(state *PersistedState) error { node, registered := state.Nodes[request.NodeID] if !registered { return requestError{status: http.StatusConflict, code: "NodeNotReady", message: "The target node is not registered."} } if !nodeHasAccount(node, request.AccountID) { return requestError{status: http.StatusConflict, code: "AccountNotReady", message: "The target account is not currently registered on the node."} } payloadKey := string(request.Payload) for _, existing := range state.Tasks { if existing.NodeID != request.NodeID || existing.AccountID != request.AccountID || existing.IdempotencyKey != request.IdempotencyKey { continue } if existing.Kind != request.Kind || string(existing.Payload) != payloadKey { return requestError{status: http.StatusConflict, code: "IdempotencyConflict", message: "The idempotency key is already bound to different task parameters."} } response = TaskSubmissionResponse{TaskID: existing.TaskID, Status: existing.Status, Duplicate: true, StateVersion: existing.StateVersion} return nil } now := time.Now().UTC() task := Task{ TaskID: randomID(), NodeID: request.NodeID, AccountID: request.AccountID, Kind: request.Kind, IdempotencyKey: request.IdempotencyKey, Payload: append(jsonRaw(nil), request.Payload...), NotAfter: request.NotAfter, Status: TaskPending, StateVersion: 1, CreatedAt: now, UpdatedAt: now, LastCorrelationID: correlationID, } state.Tasks[task.TaskID] = task state.Audit = appendAudit(state.Audit, "user:"+username, "task.create", task.TaskID, correlationID, "success", now) response = TaskSubmissionResponse{TaskID: task.TaskID, Status: task.Status, Duplicate: false, StateVersion: task.StateVersion} return nil }) if err != nil { return err } writeJSON(w, http.StatusAccepted, response) return nil } func (s *Server) listTasks(w http.ResponseWriter, r *http.Request) error { nodeID := r.URL.Query().Get("node_id") accountID := r.URL.Query().Get("account_id") limit := queryLimit(r.URL.Query().Get("limit")) var tasks []Task if err := s.store.Read(func(state PersistedState) error { for _, task := range state.Tasks { if nodeID != "" && task.NodeID != nodeID || accountID != "" && task.AccountID != accountID { continue } tasks = append(tasks, task) } return nil }); err != nil { return err } sort.Slice(tasks, func(i, j int) bool { return tasks[i].CreatedAt.After(tasks[j].CreatedAt) }) if len(tasks) > limit { tasks = tasks[:limit] } writeJSON(w, http.StatusOK, map[string]any{"tasks": tasks}) return nil } func (s *Server) getTask(w http.ResponseWriter, taskID string) error { if !validIdentifier(taskID, 200) { return requestError{status: http.StatusBadRequest, code: "InvalidTask", message: "Task ID is invalid."} } var task Task if err := s.store.Read(func(state PersistedState) error { value, ok := state.Tasks[taskID] if !ok { return requestError{status: http.StatusNotFound, code: "TaskNotFound", message: "Task was not found."} } task = value return nil }); err != nil { return err } writeJSON(w, http.StatusOK, task) return nil } func (s *Server) cancelTask(w http.ResponseWriter, taskID, username, correlationID string) error { var task Task err := s.store.Mutate(func(state *PersistedState) error { value, ok := state.Tasks[taskID] if !ok { return requestError{status: http.StatusNotFound, code: "TaskNotFound", message: "Task was not found."} } task = value if terminal(task.Status) { return nil } now := time.Now().UTC() if task.CancelRequestedAt == nil { task.CancelRequestedAt = &now task.StateVersion++ if task.Status == TaskPending && !leaseActive(task, now) { task.Status = TaskCancelled task.LeaseOwner = "" task.LeaseExpiresAt = nil } task.UpdatedAt = now state.Audit = appendAudit(state.Audit, "user:"+username, "task.cancel-request", taskID, correlationID, "success", now) state.Tasks[taskID] = task } return nil }) if err != nil { return err } writeJSON(w, http.StatusOK, task) return nil } func (s *Server) eventRoute(w http.ResponseWriter, r *http.Request, _ string) error { if r.Method != http.MethodGet { return requestError{status: http.StatusMethodNotAllowed, code: "MethodNotAllowed", message: "Method is not allowed."} } if _, err := s.authenticateWeb(r); err != nil { return err } return s.listEvents(w, r) } func (s *Server) ingestEvent(w http.ResponseWriter, r *http.Request, nodeID, correlationID string) error { var event MessageEvent if err := decodeJSON(r, &event, 64*1024); err != nil { return err } if event.NodeID != nodeID || !validIdentifier(event.AccountID, 200) || !validIdentifier(event.ChatID, 512) || event.EventSeq <= 0 || !validIdentifier(event.EventType, 80) || event.OccurredAt.IsZero() || event.ConfigVersion <= 0 || event.AuthorizationVersion <= 0 { return requestError{status: http.StatusBadRequest, code: "InvalidEvent", message: "Event identity or authorization metadata is invalid."} } if event.ChatType != ChatGroup && event.ChatType != ChatPrivate { return requestError{status: http.StatusBadRequest, code: "InvalidChatType", message: "Only group and private chat events are supported."} } if event.EventType != "message" { return requestError{status: http.StatusBadRequest, code: "UnsupportedEvent", message: "Only message events are supported in v1."} } if len(event.Content) > 16*1024 { return requestError{status: http.StatusBadRequest, code: "EventTooLarge", message: "Event content is too large."} } if !event.Authorized { return requestError{status: http.StatusForbidden, code: "ReportingNotAuthorized", message: "The node did not authorize this event."} } contentHash := eventHash(event) var receipt EventReceipt err := s.store.Mutate(func(state *PersistedState) error { for _, existing := range state.Events { if existing.NodeID != event.NodeID || existing.AccountID != event.AccountID || existing.ChatID != event.ChatID || existing.EventSeq != event.EventSeq { continue } if existing.ContentHash != contentHash { return requestError{status: http.StatusConflict, code: "EventIdempotencyConflict", message: "The event sequence is already bound to different content."} } receipt = EventReceipt{Accepted: true, Duplicate: true, EventID: existing.EventID} return nil } now := time.Now().UTC() stored := StoredEvent{MessageEvent: event, EventID: randomID(), ContentHash: contentHash, ReceivedAt: now} state.Events = append(state.Events, stored) if len(state.Events) > 10000 { state.Events = state.Events[len(state.Events)-10000:] } state.Audit = appendAudit(state.Audit, "node:"+nodeID, "event.accept", event.AccountID+":"+event.ChatID, correlationID, "success", now) receipt = EventReceipt{Accepted: true, EventID: stored.EventID} return nil }) if err != nil { return err } writeJSON(w, http.StatusAccepted, receipt) return nil } func (s *Server) listEvents(w http.ResponseWriter, r *http.Request) error { query := r.URL.Query() nodeID, accountID, chatID := query.Get("node_id"), query.Get("account_id"), query.Get("chat_id") limit := queryLimit(query.Get("limit")) var events []StoredEvent if err := s.store.Read(func(state PersistedState) error { for index := len(state.Events) - 1; index >= 0 && len(events) < limit; index-- { event := state.Events[index] if nodeID != "" && event.NodeID != nodeID || accountID != "" && event.AccountID != accountID || chatID != "" && event.ChatID != chatID { continue } events = append(events, event) } return nil }); err != nil { return err } writeJSON(w, http.StatusOK, map[string]any{"events": events}) return nil } func (s *Server) auditRoute(w http.ResponseWriter, r *http.Request, _ string) error { if _, err := s.authenticateWeb(r); err != nil { return err } limit := queryLimit(r.URL.Query().Get("limit")) var audit []AuditEntry if err := s.store.Read(func(state PersistedState) error { for index := len(state.Audit) - 1; index >= 0 && len(audit) < limit; index-- { audit = append(audit, state.Audit[index]) } return nil }); err != nil { return err } writeJSON(w, http.StatusOK, map[string]any{"audit": audit}) return nil } func (s *Server) authenticateNode(r *http.Request) (string, error) { if !s.clientCertificateAllowed(r) { return "", requestError{status: http.StatusUnauthorized, code: "Unauthorized", message: "Node certificate authentication failed."} } token := bearerToken(r) if token == "" { return "", requestError{status: http.StatusUnauthorized, code: "Unauthorized", message: "Node authentication is required."} } provided := sha256.Sum256([]byte(token)) for nodeID, expected := range s.nodeTokenHashes { if subtle.ConstantTimeCompare(provided[:], expected[:]) == 1 { return nodeID, nil } } return "", requestError{status: http.StatusUnauthorized, code: "Unauthorized", message: "Node authentication failed."} } func (s *Server) authenticateWeb(r *http.Request) (string, error) { token := bearerToken(r) if token == "" { return "", requestError{status: http.StatusUnauthorized, code: "Unauthorized", message: "Web authentication is required."} } now := time.Now().UTC() s.sessionMu.Lock() value, ok := s.sessions[token] if ok && !now.Before(value.Expires) { delete(s.sessions, token) ok = false } s.sessionMu.Unlock() if !ok { return "", requestError{status: http.StatusUnauthorized, code: "Unauthorized", message: "Web session is missing or expired."} } return value.Username, nil } func (s *Server) appendAudit(principal, action, resource, correlationID, outcome string) error { return s.store.Mutate(func(state *PersistedState) error { state.Audit = appendAudit(state.Audit, principal, action, resource, correlationID, outcome, time.Now().UTC()) return nil }) } func (s *Server) serveFrontend(w http.ResponseWriter, name string) { if name == "" || name == "." || strings.Contains(name, "..") { writeError(w, requestError{status: http.StatusNotFound, code: "NotFound", message: "Resource was not found."}, randomID()) return } content, err := fs.ReadFile(webassets.Dist, name) if err != nil { writeError(w, requestError{status: http.StatusNotFound, code: "NotFound", message: "Resource was not found."}, randomID()) return } contentType := mime.TypeByExtension(path.Ext(name)) if contentType == "" { contentType = "application/octet-stream" } w.Header().Set("Content-Type", contentType+"; charset=utf-8") w.Header().Set("Content-Security-Policy", "default-src 'self'; script-src 'self'; style-src 'self'; connect-src 'self'; frame-ancestors 'none'; base-uri 'none'; form-action 'self'") w.WriteHeader(http.StatusOK) _, _ = w.Write(content) } func decodeJSON(r *http.Request, target any, maxBytes int64) error { defer r.Body.Close() decoder := json.NewDecoder(io.LimitReader(r.Body, maxBytes)) decoder.DisallowUnknownFields() if err := decoder.Decode(target); err != nil { return requestError{status: http.StatusBadRequest, code: "InvalidRequest", message: "Request JSON is invalid."} } var extra any if decoder.Decode(&extra) != io.EOF { return requestError{status: http.StatusBadRequest, code: "InvalidRequest", message: "Request must contain one JSON value."} } return nil } func 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 writeError(w http.ResponseWriter, err error, correlationID string) { requestErr, ok := err.(requestError) if !ok { requestErr = requestError{status: http.StatusInternalServerError, code: "InternalError", message: "Request failed; use the correlation ID for diagnosis."} } writeJSON(w, requestErr.status, map[string]any{ "error": map[string]string{"code": requestErr.code, "message": requestErr.message}, "correlation_id": correlationID, }) } func pathParts(path string) []string { trimmed := strings.Trim(path, "/") if trimmed == "" { return nil } parts := strings.Split(trimmed, "/") for index, part := range parts { decoded, err := url.PathUnescape(part) if err == nil { parts[index] = decoded } } return parts } func bearerToken(r *http.Request) string { value := r.Header.Get("Authorization") if len(value) < 8 || !strings.EqualFold(value[:7], "Bearer ") { return "" } return strings.TrimSpace(value[7:]) } func validIdentifier(value string, max int) bool { return value != "" && len(value) <= max && !strings.ContainsAny(value, "\r\n\x00") } func terminal(status TaskStatus) bool { switch status { case TaskSucceeded, TaskFailed, TaskCancelled, TaskExpired, TaskResultUnconfirmed: return true default: return false } } func leaseActive(task Task, now time.Time) bool { return task.LeaseOwner != "" && task.LeaseExpiresAt != nil && now.Before(*task.LeaseExpiresAt) } func markUnconfirmed(task *Task, reason string, now time.Time) { task.Status = TaskResultUnconfirmed task.StateVersion++ task.UpdatedAt = now task.LeaseOwner = "" task.LeaseExpiresAt = nil task.Result = &TaskResult{TaskID: task.TaskID, AccountID: task.AccountID, LeaseGeneration: task.LeaseGeneration, Status: TaskResultUnconfirmed, ErrorCode: reason, Message: "The control plane could not prove that the previous executor stopped before the lease expired.", CorrelationID: task.LastCorrelationID} } func validCapabilities(capabilities []string) bool { for _, capability := range capabilities { if !validIdentifier(capability, 80) { return false } } return true } func validAccountSummaries(accounts []AccountSummary) bool { seen := make(map[string]struct{}, len(accounts)) for _, account := range accounts { if !validIdentifier(account.AccountID, 200) || account.AllowedGroupCount < 0 || account.AllowedPrivateCount < 0 { return false } if _, exists := seen[account.AccountID]; exists { return false } seen[account.AccountID] = struct{}{} } return len(accounts) <= 100 } func validNodeStatus(status NodeStatus) bool { switch status { case NodeRegistered, NodeOnline, NodeDegraded, NodeOffline, NodeSessionLocked, NodeWechatNotRunning, NodeWechatNotLogged: return true default: return false } } func nodeHasAccount(node Node, accountID string) bool { for _, account := range node.Accounts { if account.AccountID == accountID && account.Active && account.Verified { return true } } return false } func validSendTextPayload(payload jsonRaw) bool { if len(payload) == 0 || len(payload) > 64*1024 || !json.Valid(payload) { return false } var value struct { TargetID string `json:"target_id"` Text string `json:"text"` Confirmed bool `json:"confirmed"` } if json.Unmarshal(payload, &value) != nil { return false } return validIdentifier(value.TargetID, 512) && len(value.Text) > 0 && len(value.Text) <= 4000 && value.Confirmed } func eventHash(event MessageEvent) string { value := struct { NodeID string `json:"node_id"` AccountID string `json:"account_id"` ChatID string `json:"chat_id"` ChatType ChatType `json:"chat_type"` EventSeq int64 `json:"event_seq"` EventType string `json:"event_type"` OccurredAt time.Time `json:"occurred_at"` Content string `json:"content"` ConfigVersion int64 `json:"config_version"` }{event.NodeID, event.AccountID, event.ChatID, event.ChatType, event.EventSeq, event.EventType, event.OccurredAt, event.Content, event.ConfigVersion} data, _ := json.Marshal(value) digest := sha256.Sum256(data) return hex.EncodeToString(digest[:]) } func appendAudit(entries []AuditEntry, principal, action, resource, correlationID, outcome string, at time.Time) []AuditEntry { entries = append(entries, AuditEntry{ID: randomID(), At: at, Principal: principal, Action: action, Resource: resource, CorrelationID: correlationID, Outcome: outcome}) if len(entries) > 10000 { entries = entries[len(entries)-10000:] } return entries } func queryLimit(raw string) int { if value, err := strconv.Atoi(raw); err == nil && value >= 1 && value <= 200 { return value } return 50 } func queryWaitSeconds(raw string) int { if value, err := strconv.Atoi(raw); err == nil && value >= 1 && value <= 30 { return value } return 0 } func firstNonEmpty(value, fallback string) string { if value != "" { return value } return fallback } func randomID() string { buffer := make([]byte, 16) if _, err := rand.Read(buffer); err != nil { return fmt.Sprintf("fallback-%d", time.Now().UnixNano()) } return base64.RawURLEncoding.EncodeToString(buffer) }