This commit is contained in:
+27
-2
@@ -27,9 +27,23 @@ export WXAGENT_NODE_TOKENS='{"node-a":"token-a","node-b":"token-b"}'
|
||||
export WXAGENT_WEB_USERS='{"admin":"change-me","auditor":"read-only-password"}'
|
||||
```
|
||||
|
||||
可选:`WXAGENT_CONTROL_PLANE_ADDR`(默认 `127.0.0.1:8090`)、`WXAGENT_CONTROL_PLANE_DATA`(默认 `control-plane-data.json`)。Token 只从环境变量读取,不写入控制面数据文件和日志。节点 `service.json` 可设置 `remoteConfigurationFile` 指向 CLI 管理的 `remote.json`;修改上报配置会在下一轮生效,修改远程地址/Token 后需重启 Agent。
|
||||
可选:`WXAGENT_CONTROL_PLANE_ADDR`(默认 `127.0.0.1:8090`)、`WXAGENT_CONTROL_PLANE_DATA`(默认 `control-plane-data.json`)。凭据可通过环境变量或 `*_FILE` Secret 文件读取,不写入控制面数据文件和日志。节点 `service.json` 可设置 `remoteConfigurationFile` 指向 CLI 管理的 `remote.json`;修改上报配置会在下一轮生效,修改远程地址/Token 后需重启 Agent。
|
||||
|
||||
浏览器访问 `http://127.0.0.1:8090/`,登录后管理节点、任务、白名单事件和审计记录。生产部署必须使用 HTTPS 和外部密钥管理;本地 HTTP 仅用于 loopback 集成测试。
|
||||
生产控制面可直接启用 TLS 和节点 mTLS:
|
||||
|
||||
```text
|
||||
WXAGENT_CONTROL_PLANE_TLS_CERT_FILE
|
||||
WXAGENT_CONTROL_PLANE_TLS_KEY_FILE
|
||||
WXAGENT_MTLS_CLIENT_CA_FILE
|
||||
WXAGENT_MTLS_REQUIRE_NODE_CERT=true
|
||||
WXAGENT_MTLS_REVOKED_CERTS_FILE
|
||||
WXAGENT_NODE_TOKENS_FILE=/run/secrets/node_tokens
|
||||
WXAGENT_WEB_USERS_FILE=/run/secrets/web_users
|
||||
```
|
||||
|
||||
节点证书按证书原始 DER 的 SHA-256 指纹逐行写入撤销文件;文件读取失败时节点认证拒绝。控制面还提供 `/readyz`,并对数据文件使用单写入者锁、自动轮转备份和任务/事件/审计保留期限。生产示例见 [`compose.production.yml`](./compose.production.yml)。
|
||||
|
||||
浏览器访问 `http://127.0.0.1:8090/`,登录后管理节点、任务、白名单事件和审计记录。生产部署必须使用 HTTPS、节点 mTLS 和外部 Secret 文件;本地 HTTP 仅用于显式的私有网络/loopback 集成测试。
|
||||
|
||||
远程只读查询使用同一持久化任务队列,返回 `task_id` 后通过 `GET /v1/tasks/{task_id}` 取结果:
|
||||
|
||||
@@ -55,6 +69,15 @@ cd ../..
|
||||
|
||||
`remote-control-smoke.sh` 会启动一次临时控制面,运行 .NET 节点协议客户端,验证注册、心跳、任务租约、幂等、结果回传和白名单拒绝;不会操作真实微信联系人。
|
||||
|
||||
控制面并发/读取任务基线(使用已注册的测试节点,不执行写任务):
|
||||
|
||||
```bash
|
||||
WXAGENT_SCALE_WEB_PASSWORD='<test-password>' \
|
||||
scripts/control-plane-scale-smoke.py --base-url http://127.0.0.1:8090 \
|
||||
--node-id node-a --account-id account-a --requests 100 --workers 10
|
||||
```
|
||||
|
||||
`control-plane-data.sh backup`/`restore` 用于停机前后的手工备份和原子恢复;自动备份由 Store 按 `WXAGENT_CONTROL_PLANE_BACKUP_INTERVAL`(默认 5 分钟)产生。当前 JSON 存储使用 active/passive 单写锁,不支持 active-active 多实例共享写入。
|
||||
## Windows 手工验收
|
||||
|
||||
1. 在已登录且未锁定的微信桌面会话中运行 `WxAgent.Host doctor` 和 `WxAgent.Host inspect-ui --output artifacts/ui-tree.json`,确认账号绑定使用稳定 `accountId`,会话使用稳定 `AutomationId`,不使用昵称/PID/窗口句柄猜测。
|
||||
@@ -78,3 +101,5 @@ cd ../..
|
||||
## 容器镜像
|
||||
|
||||
`.gitea/workflows/build-web-image.yml` 会构建并推送 `git.ipao.vip/<owner>/<repo>` 镜像:主分支额外更新 `latest`,`v*` 标签额外更新对应版本标签。请在 Gitea 账号级 Secrets 配置 `REGISTRY_TOKEN`;登录用户名使用仓库所有者。
|
||||
|
||||
生产环境不要提交 `compose.production.yml` 所引用的 `secrets/` 文件。将节点 Token map、Web 用户 map、TLS 私钥、客户端 CA 和撤销指纹列表放入外部 Secret 管理或受 ACL 保护的部署目录,再运行 `docker compose -f compose.production.yml up -d`。证书泄露时可用 `scripts/revoke-client-certificate.sh` 更新指纹列表,并重建 Docker Secret 挂载的容器。
|
||||
|
||||
@@ -3,12 +3,13 @@ package main
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log"
|
||||
"os"
|
||||
"os/signal"
|
||||
"strconv"
|
||||
"strings"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
controlplane "git.ipao.vip/rogee/wx-win-agent/control-plane"
|
||||
)
|
||||
@@ -21,22 +22,34 @@ func main() {
|
||||
if value := os.Getenv("WXAGENT_CONTROL_PLANE_DATA"); value != "" {
|
||||
config.DataFile = value
|
||||
}
|
||||
config.NodeTokens = readMapEnv("WXAGENT_NODE_TOKENS")
|
||||
config.BackupDir = os.Getenv("WXAGENT_CONTROL_PLANE_BACKUP_DIR")
|
||||
config.BackupCount = readIntEnv("WXAGENT_CONTROL_PLANE_BACKUP_COUNT", config.BackupCount)
|
||||
config.BackupInterval = readDurationEnv("WXAGENT_CONTROL_PLANE_BACKUP_INTERVAL", config.BackupInterval)
|
||||
config.TaskRetention = readDurationEnv("WXAGENT_CONTROL_PLANE_TASK_RETENTION", config.TaskRetention)
|
||||
config.EventRetention = readDurationEnv("WXAGENT_CONTROL_PLANE_EVENT_RETENTION", config.EventRetention)
|
||||
config.AuditRetention = readDurationEnv("WXAGENT_CONTROL_PLANE_AUDIT_RETENTION", config.AuditRetention)
|
||||
config.TLSCertFile = os.Getenv("WXAGENT_CONTROL_PLANE_TLS_CERT_FILE")
|
||||
config.TLSKeyFile = os.Getenv("WXAGENT_CONTROL_PLANE_TLS_KEY_FILE")
|
||||
config.MTLSClientCAFile = os.Getenv("WXAGENT_MTLS_CLIENT_CA_FILE")
|
||||
config.MTLSRequireNodeCert = readBoolEnv("WXAGENT_MTLS_REQUIRE_NODE_CERT", false)
|
||||
config.MTLSRevokedCertsFile = os.Getenv("WXAGENT_MTLS_REVOKED_CERTS_FILE")
|
||||
|
||||
config.NodeTokens = readMapSecret("WXAGENT_NODE_TOKENS", "WXAGENT_NODE_TOKENS_FILE")
|
||||
if len(config.NodeTokens) == 0 {
|
||||
nodeID, nodeToken := os.Getenv("WXAGENT_NODE_ID"), os.Getenv("WXAGENT_NODE_TOKEN")
|
||||
nodeID, nodeToken := os.Getenv("WXAGENT_NODE_ID"), readSecret("WXAGENT_NODE_TOKEN", "WXAGENT_NODE_TOKEN_FILE")
|
||||
if strings.TrimSpace(nodeID) != "" && nodeToken != "" {
|
||||
config.NodeTokens = map[string]string{nodeID: nodeToken}
|
||||
}
|
||||
}
|
||||
config.WebUsers = readMapEnv("WXAGENT_WEB_USERS")
|
||||
config.WebUsers = readMapSecret("WXAGENT_WEB_USERS", "WXAGENT_WEB_USERS_FILE")
|
||||
if len(config.WebUsers) == 0 {
|
||||
webUser, webPassword := os.Getenv("WXAGENT_WEB_USER"), os.Getenv("WXAGENT_WEB_PASSWORD")
|
||||
webUser, webPassword := os.Getenv("WXAGENT_WEB_USER"), readSecret("WXAGENT_WEB_PASSWORD", "WXAGENT_WEB_PASSWORD_FILE")
|
||||
if strings.TrimSpace(webUser) != "" && webPassword != "" {
|
||||
config.WebUsers = map[string]string{webUser: webPassword}
|
||||
}
|
||||
}
|
||||
if len(config.NodeTokens) == 0 || len(config.WebUsers) == 0 {
|
||||
log.Fatal("configure WXAGENT_NODE_TOKENS/WXAGENT_WEB_USERS as JSON maps, or the single-node fallback environment variables")
|
||||
log.Fatal("configure node and web credentials with environment variables or *_FILE secret files")
|
||||
}
|
||||
|
||||
server, err := controlplane.NewServer(config)
|
||||
@@ -47,19 +60,66 @@ func main() {
|
||||
defer stop()
|
||||
log.Printf("wxagent control plane listening on %s; data=%s", config.ListenAddr, config.DataFile)
|
||||
if err := server.ListenAndServe(ctx); err != nil {
|
||||
fmt.Fprintln(os.Stderr, err)
|
||||
log.Printf("control plane stopped: %v", err)
|
||||
os.Exit(1)
|
||||
}
|
||||
}
|
||||
|
||||
func readMapEnv(name string) map[string]string {
|
||||
value := strings.TrimSpace(os.Getenv(name))
|
||||
func readSecret(valueName, fileName string) string {
|
||||
if path := strings.TrimSpace(os.Getenv(fileName)); path != "" {
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
log.Fatalf("read %s: %v", fileName, err)
|
||||
}
|
||||
return strings.TrimRight(string(data), "\r\n")
|
||||
}
|
||||
return os.Getenv(valueName)
|
||||
}
|
||||
|
||||
func readMapSecret(valueName, fileName string) map[string]string {
|
||||
value := readSecret(valueName, fileName)
|
||||
if value == "" {
|
||||
return nil
|
||||
}
|
||||
var result map[string]string
|
||||
result := make(map[string]string)
|
||||
if err := json.Unmarshal([]byte(value), &result); err != nil {
|
||||
log.Fatalf("%s must be a JSON object", name)
|
||||
log.Fatalf("%s must contain a JSON object", fileName)
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
func readBoolEnv(name string, fallback bool) bool {
|
||||
value := strings.TrimSpace(os.Getenv(name))
|
||||
if value == "" {
|
||||
return fallback
|
||||
}
|
||||
parsed, err := strconv.ParseBool(value)
|
||||
if err != nil {
|
||||
log.Fatalf("%s must be true or false", name)
|
||||
}
|
||||
return parsed
|
||||
}
|
||||
|
||||
func readIntEnv(name string, fallback int) int {
|
||||
value := strings.TrimSpace(os.Getenv(name))
|
||||
if value == "" {
|
||||
return fallback
|
||||
}
|
||||
parsed, err := strconv.Atoi(value)
|
||||
if err != nil || parsed < 1 {
|
||||
log.Fatalf("%s must be a positive integer", name)
|
||||
}
|
||||
return parsed
|
||||
}
|
||||
|
||||
func readDurationEnv(name string, fallback time.Duration) time.Duration {
|
||||
value := strings.TrimSpace(os.Getenv(name))
|
||||
if value == "" {
|
||||
return fallback
|
||||
}
|
||||
parsed, err := time.ParseDuration(value)
|
||||
if err != nil || parsed < 0 {
|
||||
log.Fatalf("%s must be a non-negative duration", name)
|
||||
}
|
||||
return parsed
|
||||
}
|
||||
|
||||
@@ -0,0 +1,34 @@
|
||||
package main
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestReadSecretPrefersFileAndRemovesOnlyLineEndings(t *testing.T) {
|
||||
directory := t.TempDir()
|
||||
path := filepath.Join(directory, "token")
|
||||
if err := os.WriteFile(path, []byte("file-token\r\n"), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Setenv("TEST_TOKEN", "environment-token")
|
||||
t.Setenv("TEST_TOKEN_FILE", path)
|
||||
if got := readSecret("TEST_TOKEN", "TEST_TOKEN_FILE"); got != "file-token" {
|
||||
t.Fatalf("readSecret() = %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadMapSecretFromFile(t *testing.T) {
|
||||
directory := t.TempDir()
|
||||
path := filepath.Join(directory, "users.json")
|
||||
if err := os.WriteFile(path, []byte("{\"admin\":\"password\"}\n"), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Setenv("TEST_USERS", "")
|
||||
t.Setenv("TEST_USERS_FILE", path)
|
||||
users := readMapSecret("TEST_USERS", "TEST_USERS_FILE")
|
||||
if users["admin"] != "password" || len(users) != 1 {
|
||||
t.Fatalf("readMapSecret() = %#v", users)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,56 @@
|
||||
services:
|
||||
control-plane:
|
||||
image: ${WXAGENT_CONTROL_PLANE_IMAGE:?set image tag}
|
||||
restart: unless-stopped
|
||||
read_only: true
|
||||
user: "65532:65532"
|
||||
security_opt:
|
||||
- no-new-privileges:true
|
||||
cap_drop:
|
||||
- ALL
|
||||
tmpfs:
|
||||
- /tmp:rw,noexec,nosuid,size=16m
|
||||
ports:
|
||||
- "${CONTROL_PLANE_BIND:-127.0.0.1:18090}:8090"
|
||||
volumes:
|
||||
- control-plane-data:/data
|
||||
secrets:
|
||||
- node_tokens
|
||||
- web_users
|
||||
- tls_cert
|
||||
- tls_key
|
||||
- client_ca
|
||||
- revoked_client_certificates
|
||||
environment:
|
||||
WXAGENT_CONTROL_PLANE_ADDR: 0.0.0.0:8090
|
||||
WXAGENT_CONTROL_PLANE_DATA: /data/control-plane-data.json
|
||||
WXAGENT_CONTROL_PLANE_BACKUP_DIR: /data/backups
|
||||
WXAGENT_CONTROL_PLANE_BACKUP_COUNT: "7"
|
||||
WXAGENT_CONTROL_PLANE_BACKUP_INTERVAL: 5m
|
||||
WXAGENT_CONTROL_PLANE_TASK_RETENTION: 720h
|
||||
WXAGENT_CONTROL_PLANE_EVENT_RETENTION: 720h
|
||||
WXAGENT_CONTROL_PLANE_AUDIT_RETENTION: 2160h
|
||||
WXAGENT_CONTROL_PLANE_TLS_CERT_FILE: /run/secrets/tls_cert
|
||||
WXAGENT_CONTROL_PLANE_TLS_KEY_FILE: /run/secrets/tls_key
|
||||
WXAGENT_MTLS_CLIENT_CA_FILE: /run/secrets/client_ca
|
||||
WXAGENT_MTLS_REQUIRE_NODE_CERT: "true"
|
||||
WXAGENT_MTLS_REVOKED_CERTS_FILE: /run/secrets/revoked_client_certificates
|
||||
WXAGENT_NODE_TOKENS_FILE: /run/secrets/node_tokens
|
||||
WXAGENT_WEB_USERS_FILE: /run/secrets/web_users
|
||||
|
||||
secrets:
|
||||
node_tokens:
|
||||
file: ./secrets/node-tokens.json
|
||||
web_users:
|
||||
file: ./secrets/web-users.json
|
||||
tls_cert:
|
||||
file: ./secrets/server.crt
|
||||
tls_key:
|
||||
file: ./secrets/server.key
|
||||
client_ca:
|
||||
file: ./secrets/client-ca.crt
|
||||
revoked_client_certificates:
|
||||
file: ./secrets/revoked-client-certificates.txt
|
||||
|
||||
volumes:
|
||||
control-plane-data:
|
||||
@@ -0,0 +1,197 @@
|
||||
package controlplane
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"crypto/rsa"
|
||||
"crypto/sha256"
|
||||
"crypto/tls"
|
||||
"crypto/x509"
|
||||
"crypto/x509/pkix"
|
||||
"encoding/hex"
|
||||
"encoding/pem"
|
||||
"math/big"
|
||||
"net/http"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestStoreLocksDataAndKeepsBoundedBackups(t *testing.T) {
|
||||
directory := t.TempDir()
|
||||
path := filepath.Join(directory, "state.json")
|
||||
store, err := OpenStore(path, StoreOptions{BackupDir: filepath.Join(directory, "backups"), BackupCount: 2})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := store.Mutate(func(state *PersistedState) error {
|
||||
state.Nodes["node-1"] = Node{NodeID: "node-1"}
|
||||
return nil
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := store.Mutate(func(state *PersistedState) error {
|
||||
state.Nodes["node-1"] = Node{NodeID: "node-1", AgentVersion: "2"}
|
||||
return nil
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := store.Mutate(func(state *PersistedState) error {
|
||||
state.Nodes["node-1"] = Node{NodeID: "node-1", AgentVersion: "3"}
|
||||
return nil
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := OpenStore(path, StoreOptions{BackupDir: filepath.Join(directory, "backups"), BackupCount: 2}); err == nil || !strings.Contains(err.Error(), "already in use") {
|
||||
t.Fatalf("second writer was not rejected: %v", err)
|
||||
}
|
||||
if err := store.Close(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
entries, err := os.ReadDir(filepath.Join(directory, "backups"))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(entries) != 2 {
|
||||
t.Fatalf("expected two backups, got %d", len(entries))
|
||||
}
|
||||
for _, entry := range entries {
|
||||
if entry.Type().Perm() != 0o600 {
|
||||
info, infoErr := entry.Info()
|
||||
if infoErr != nil {
|
||||
t.Fatal(infoErr)
|
||||
}
|
||||
if info.Mode().Perm() != 0o600 {
|
||||
t.Fatalf("backup permissions = %o", info.Mode().Perm())
|
||||
}
|
||||
}
|
||||
}
|
||||
second, err := OpenStore(path, StoreOptions{BackupDir: filepath.Join(directory, "backups"), BackupCount: 2})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer second.Close()
|
||||
}
|
||||
|
||||
func TestStorePrunesOnlyTerminalRecordsByRetention(t *testing.T) {
|
||||
directory := t.TempDir()
|
||||
path := filepath.Join(directory, "state.json")
|
||||
seed, err := OpenStore(path, StoreOptions{BackupDir: "", BackupCount: -1})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
old := time.Now().UTC().Add(-2 * time.Hour)
|
||||
if err := seed.Mutate(func(state *PersistedState) error {
|
||||
state.Tasks["old-terminal"] = Task{TaskID: "old-terminal", Status: TaskSucceeded, UpdatedAt: old}
|
||||
state.Tasks["old-running"] = Task{TaskID: "old-running", Status: TaskRunning, UpdatedAt: old}
|
||||
state.Events = append(state.Events, StoredEvent{ReceivedAt: old})
|
||||
state.Audit = append(state.Audit, AuditEntry{At: old})
|
||||
return nil
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := seed.Close(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
store, err := OpenStore(path, StoreOptions{
|
||||
BackupDir: filepath.Join(directory, "backups"), BackupCount: -1,
|
||||
TaskRetention: time.Hour, EventRetention: time.Hour, AuditRetention: time.Hour,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer store.Close()
|
||||
if err := store.Mutate(func(state *PersistedState) error {
|
||||
state.Nodes["node-1"] = Node{NodeID: "node-1"}
|
||||
return nil
|
||||
}); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
snapshot := store.Snapshot()
|
||||
if _, ok := snapshot.Tasks["old-terminal"]; ok {
|
||||
t.Fatal("old terminal task was retained")
|
||||
}
|
||||
if _, ok := snapshot.Tasks["old-running"]; !ok {
|
||||
t.Fatal("old running task was pruned")
|
||||
}
|
||||
if len(snapshot.Events) != 0 || len(snapshot.Audit) != 0 {
|
||||
t.Fatalf("old events/audit were retained: events=%d audit=%d", len(snapshot.Events), len(snapshot.Audit))
|
||||
}
|
||||
}
|
||||
|
||||
func TestTLSConfigSupportsVerifiedNodeClientCertificates(t *testing.T) {
|
||||
directory := t.TempDir()
|
||||
certFile, keyFile, caFile := writeTestCertificateMaterial(t, directory)
|
||||
config := ServerConfig{
|
||||
TLSCertFile: certFile,
|
||||
TLSKeyFile: keyFile,
|
||||
MTLSClientCAFile: caFile,
|
||||
MTLSRequireNodeCert: true,
|
||||
MTLSRevokedCertsFile: filepath.Join(directory, "revoked.txt"),
|
||||
}
|
||||
if err := os.WriteFile(config.MTLSRevokedCertsFile, []byte("# initially empty\n"), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
tlsConfig, err := newTLSConfig(config)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if tlsConfig.MinVersion != tls.VersionTLS13 || tlsConfig.ClientAuth != tls.VerifyClientCertIfGiven || len(tlsConfig.Certificates) != 1 {
|
||||
t.Fatalf("unexpected TLS config: min=%d auth=%d certs=%d", tlsConfig.MinVersion, tlsConfig.ClientAuth, len(tlsConfig.Certificates))
|
||||
}
|
||||
|
||||
rawCertificate := []byte("client-cert")
|
||||
digest := sha256.Sum256(rawCertificate)
|
||||
fingerprint := hex.EncodeToString(digest[:])
|
||||
if err := os.WriteFile(config.MTLSRevokedCertsFile, []byte(fingerprint+"\n"), 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
server := &Server{config: config}
|
||||
request := &http.Request{TLS: &tls.ConnectionState{PeerCertificates: []*x509.Certificate{{Raw: rawCertificate}}}}
|
||||
if server.clientCertificateAllowed(request) {
|
||||
t.Fatal("revoked certificate was accepted")
|
||||
}
|
||||
}
|
||||
|
||||
func TestTLSConfigRejectsIncompleteSecuritySettings(t *testing.T) {
|
||||
if _, err := newTLSConfig(ServerConfig{TLSCertFile: "cert.pem"}); err == nil {
|
||||
t.Fatal("incomplete TLS certificate settings were accepted")
|
||||
}
|
||||
if _, err := newTLSConfig(ServerConfig{MTLSRequireNodeCert: true}); err == nil {
|
||||
t.Fatal("mTLS without TLS material was accepted")
|
||||
}
|
||||
}
|
||||
|
||||
func writeTestCertificateMaterial(t *testing.T, directory string) (string, string, string) {
|
||||
t.Helper()
|
||||
key, err := rsa.GenerateKey(rand.Reader, 2048)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
template := &x509.Certificate{
|
||||
SerialNumber: big.NewInt(1),
|
||||
Subject: pkix.Name{CommonName: "wxagent-test"},
|
||||
NotBefore: time.Now().Add(-time.Minute),
|
||||
NotAfter: time.Now().Add(time.Hour),
|
||||
KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageKeyEncipherment | x509.KeyUsageCertSign,
|
||||
IsCA: true,
|
||||
BasicConstraintsValid: true,
|
||||
}
|
||||
der, err := x509.CreateCertificate(rand.Reader, template, template, &key.PublicKey, key)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
certFile := filepath.Join(directory, "server.pem")
|
||||
keyFile := filepath.Join(directory, "server-key.pem")
|
||||
caFile := filepath.Join(directory, "client-ca.pem")
|
||||
certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: der})
|
||||
keyPEM := pem.EncodeToMemory(&pem.Block{Type: "RSA PRIVATE KEY", Bytes: x509.MarshalPKCS1PrivateKey(key)})
|
||||
for path, data := range map[string][]byte{certFile: certPEM, keyFile: keyPEM, caFile: certPEM} {
|
||||
if err := os.WriteFile(path, data, 0o600); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
return certFile, keyFile, caFile
|
||||
}
|
||||
+79
-9
@@ -5,6 +5,7 @@ import (
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"crypto/subtle"
|
||||
"crypto/tls"
|
||||
"encoding/base64"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
@@ -26,13 +27,24 @@ import (
|
||||
)
|
||||
|
||||
type ServerConfig struct {
|
||||
ListenAddr string
|
||||
DataFile string
|
||||
NodeTokens map[string]string
|
||||
WebUsers map[string]string
|
||||
LeaseTTL time.Duration
|
||||
HeartbeatTimeout time.Duration
|
||||
SessionTTL time.Duration
|
||||
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 {
|
||||
@@ -44,12 +56,18 @@ func DefaultServerConfig() ServerConfig {
|
||||
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
|
||||
@@ -77,6 +95,24 @@ func NewServer(config ServerConfig) (*Server, error) {
|
||||
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
|
||||
}
|
||||
@@ -92,13 +128,24 @@ func NewServer(config ServerConfig) (*Server, error) {
|
||||
if config.WebUsers == nil {
|
||||
config.WebUsers = map[string]string{}
|
||||
}
|
||||
store, err := OpenStore(config.DataFile)
|
||||
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{},
|
||||
@@ -126,13 +173,23 @@ func (s *Server) ListenAndServe(ctx context.Context) error {
|
||||
defer cancel()
|
||||
_ = server.Shutdown(shutdownCtx)
|
||||
}()
|
||||
err := server.ListenAndServe()
|
||||
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) {
|
||||
@@ -158,6 +215,8 @@ func (s *Server) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
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:
|
||||
@@ -182,6 +241,14 @@ func (s *Server) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
}
|
||||
}
|
||||
|
||||
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"`
|
||||
@@ -1071,6 +1138,9 @@ func (s *Server) auditRoute(w http.ResponseWriter, r *http.Request, _ string) er
|
||||
}
|
||||
|
||||
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."}
|
||||
|
||||
@@ -19,6 +19,7 @@ func TestReactFrontendIsEmbedded(t *testing.T) {
|
||||
}
|
||||
httpServer := httptest.NewServer(server.Handler())
|
||||
defer httpServer.Close()
|
||||
defer server.Close()
|
||||
response, err := http.Get(httpServer.URL + "/")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
@@ -70,6 +71,7 @@ func TestNodeWebTaskAndEventFlow(t *testing.T) {
|
||||
}
|
||||
httpServer := httptest.NewServer(server.Handler())
|
||||
defer httpServer.Close()
|
||||
defer server.Close()
|
||||
client := httpServer.Client()
|
||||
|
||||
if response := doJSON(t, client, http.MethodGet, httpServer.URL+"/v1/nodes", "", nil); response.Code != http.StatusUnauthorized {
|
||||
@@ -195,6 +197,9 @@ func TestNodeWebTaskAndEventFlow(t *testing.T) {
|
||||
t.Fatalf("expected two account-scoped events, got %d", len(eventList.Events))
|
||||
}
|
||||
|
||||
if err := server.Close(); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
restarted, err := NewServer(ServerConfig{DataFile: dataFile, NodeTokens: map[string]string{"node-1": "node-secret"}, WebUsers: map[string]string{"admin": "web-secret"}})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
@@ -226,6 +231,7 @@ func TestRemoteReadRoutesUseTheTaskQueueAndRetainReadContent(t *testing.T) {
|
||||
}
|
||||
httpServer := httptest.NewServer(server.Handler())
|
||||
defer httpServer.Close()
|
||||
defer server.Close()
|
||||
client := httpServer.Client()
|
||||
|
||||
login := doJSON(t, client, http.MethodPost, httpServer.URL+"/v1/auth/login", "", map[string]string{"username": "admin", "password": "web-secret"})
|
||||
|
||||
+231
-8
@@ -6,33 +6,88 @@ import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
"strings"
|
||||
"sync"
|
||||
"syscall"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Store is a small durable JSON store for the single-process control-plane MVP.
|
||||
// The file is an implementation detail; callers only observe transactional methods.
|
||||
type Store struct {
|
||||
path string
|
||||
mu sync.Mutex
|
||||
state PersistedState
|
||||
// StoreOptions controls durability and bounded retention for the JSON store.
|
||||
// A Store still has one active writer; the lock file enables active/passive failover.
|
||||
type StoreOptions struct {
|
||||
BackupDir string
|
||||
BackupCount int
|
||||
BackupInterval time.Duration
|
||||
TaskRetention time.Duration
|
||||
EventRetention time.Duration
|
||||
AuditRetention time.Duration
|
||||
}
|
||||
|
||||
func OpenStore(path string) (*Store, error) {
|
||||
// Store is a small durable JSON store for the control-plane MVP.
|
||||
// The file is an implementation detail; callers only observe transactional methods.
|
||||
type Store struct {
|
||||
path string
|
||||
lockFile *os.File
|
||||
options StoreOptions
|
||||
lastBackup time.Time
|
||||
mu sync.Mutex
|
||||
state PersistedState
|
||||
closed bool
|
||||
}
|
||||
|
||||
func OpenStore(path string, options ...StoreOptions) (*Store, error) {
|
||||
if path == "" {
|
||||
return nil, errors.New("data file is required")
|
||||
}
|
||||
store := &Store{path: filepath.Clean(path), state: PersistedState{Nodes: map[string]Node{}, Tasks: map[string]Task{}, Events: []StoredEvent{}, Audit: []AuditEntry{}}}
|
||||
cleanPath := filepath.Clean(path)
|
||||
storeOptions := StoreOptions{}
|
||||
if len(options) > 0 {
|
||||
storeOptions = options[0]
|
||||
}
|
||||
if storeOptions.BackupDir == "" {
|
||||
storeOptions.BackupDir = cleanPath + ".backups"
|
||||
}
|
||||
if storeOptions.BackupCount == 0 {
|
||||
storeOptions.BackupCount = 7
|
||||
}
|
||||
if storeOptions.BackupCount < -1 {
|
||||
return nil, errors.New("backup count must be -1 or non-negative")
|
||||
}
|
||||
if err := os.MkdirAll(filepath.Dir(cleanPath), 0o700); err != nil {
|
||||
return nil, fmt.Errorf("create data directory: %w", err)
|
||||
}
|
||||
lockFile, err := os.OpenFile(cleanPath+".lock", os.O_CREATE|os.O_RDWR, 0o600)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("open data lock: %w", err)
|
||||
}
|
||||
if err := syscall.Flock(int(lockFile.Fd()), syscall.LOCK_EX|syscall.LOCK_NB); err != nil {
|
||||
_ = lockFile.Close()
|
||||
if errors.Is(err, syscall.EWOULDBLOCK) {
|
||||
return nil, fmt.Errorf("data file is already in use: %s", cleanPath)
|
||||
}
|
||||
return nil, fmt.Errorf("lock data file: %w", err)
|
||||
}
|
||||
|
||||
store := &Store{
|
||||
path: cleanPath,
|
||||
lockFile: lockFile,
|
||||
options: storeOptions,
|
||||
state: PersistedState{Nodes: map[string]Node{}, Tasks: map[string]Task{}, Events: []StoredEvent{}, Audit: []AuditEntry{}},
|
||||
}
|
||||
data, err := os.ReadFile(store.path)
|
||||
if errors.Is(err, os.ErrNotExist) {
|
||||
return store, nil
|
||||
}
|
||||
if err != nil {
|
||||
_ = store.Close()
|
||||
return nil, fmt.Errorf("read data file: %w", err)
|
||||
}
|
||||
if len(data) == 0 {
|
||||
return store, nil
|
||||
}
|
||||
if err := json.Unmarshal(data, &store.state); err != nil {
|
||||
_ = store.Close()
|
||||
return nil, fmt.Errorf("decode data file: %w", err)
|
||||
}
|
||||
store.ensureMaps()
|
||||
@@ -48,6 +103,9 @@ func (s *Store) Snapshot() PersistedState {
|
||||
func (s *Store) Mutate(fn func(*PersistedState) error) error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
if s.closed {
|
||||
return errors.New("store is closed")
|
||||
}
|
||||
s.ensureMaps()
|
||||
if err := fn(&s.state); err != nil {
|
||||
return err
|
||||
@@ -58,9 +116,35 @@ func (s *Store) Mutate(fn func(*PersistedState) error) error {
|
||||
func (s *Store) Read(fn func(PersistedState) error) error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
if s.closed {
|
||||
return errors.New("store is closed")
|
||||
}
|
||||
return fn(cloneState(s.state))
|
||||
}
|
||||
|
||||
// Close releases the single-writer lock. A second control-plane process can then
|
||||
// be promoted by the supervisor using the same data directory.
|
||||
func (s *Store) Close() error {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
if s.closed {
|
||||
return nil
|
||||
}
|
||||
s.closed = true
|
||||
if s.lockFile == nil {
|
||||
return nil
|
||||
}
|
||||
unlockErr := syscall.Flock(int(s.lockFile.Fd()), syscall.LOCK_UN)
|
||||
closeErr := s.lockFile.Close()
|
||||
if unlockErr != nil {
|
||||
return fmt.Errorf("unlock data file: %w", unlockErr)
|
||||
}
|
||||
if closeErr != nil {
|
||||
return fmt.Errorf("close data lock: %w", closeErr)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Store) ensureMaps() {
|
||||
if s.state.Nodes == nil {
|
||||
s.state.Nodes = map[string]Node{}
|
||||
@@ -80,6 +164,10 @@ func (s *Store) saveLocked() error {
|
||||
if err := os.MkdirAll(filepath.Dir(s.path), 0o700); err != nil {
|
||||
return fmt.Errorf("create data directory: %w", err)
|
||||
}
|
||||
s.pruneLocked(time.Now().UTC())
|
||||
if err := s.backupLocked(); err != nil {
|
||||
return err
|
||||
}
|
||||
data, err := json.MarshalIndent(s.state, "", " ")
|
||||
if err != nil {
|
||||
return fmt.Errorf("encode data file: %w", err)
|
||||
@@ -108,6 +196,141 @@ func (s *Store) saveLocked() error {
|
||||
if err := os.Rename(temporaryName, s.path); err != nil {
|
||||
return fmt.Errorf("replace data file: %w", err)
|
||||
}
|
||||
return syncDirectory(filepath.Dir(s.path))
|
||||
}
|
||||
|
||||
func (s *Store) pruneLocked(now time.Time) {
|
||||
if s.options.TaskRetention > 0 {
|
||||
cutoff := now.Add(-s.options.TaskRetention)
|
||||
for id, task := range s.state.Tasks {
|
||||
if terminal(task.Status) && !task.UpdatedAt.IsZero() && task.UpdatedAt.Before(cutoff) {
|
||||
delete(s.state.Tasks, id)
|
||||
}
|
||||
}
|
||||
}
|
||||
if s.options.EventRetention > 0 {
|
||||
cutoff := now.Add(-s.options.EventRetention)
|
||||
kept := s.state.Events[:0]
|
||||
for _, event := range s.state.Events {
|
||||
if event.ReceivedAt.IsZero() || !event.ReceivedAt.Before(cutoff) {
|
||||
kept = append(kept, event)
|
||||
}
|
||||
}
|
||||
s.state.Events = kept
|
||||
}
|
||||
if s.options.AuditRetention > 0 {
|
||||
cutoff := now.Add(-s.options.AuditRetention)
|
||||
kept := s.state.Audit[:0]
|
||||
for _, entry := range s.state.Audit {
|
||||
if entry.At.IsZero() || !entry.At.Before(cutoff) {
|
||||
kept = append(kept, entry)
|
||||
}
|
||||
}
|
||||
s.state.Audit = kept
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Store) backupLocked() error {
|
||||
if s.options.BackupCount < 1 {
|
||||
return nil
|
||||
}
|
||||
now := time.Now().UTC()
|
||||
if s.options.BackupInterval > 0 && !s.lastBackup.IsZero() && now.Sub(s.lastBackup) < s.options.BackupInterval {
|
||||
return nil
|
||||
}
|
||||
data, err := os.ReadFile(s.path)
|
||||
if errors.Is(err, os.ErrNotExist) {
|
||||
return nil
|
||||
}
|
||||
if err != nil {
|
||||
return fmt.Errorf("read current data for backup: %w", err)
|
||||
}
|
||||
if err := os.MkdirAll(s.options.BackupDir, 0o700); err != nil {
|
||||
return fmt.Errorf("create backup directory: %w", err)
|
||||
}
|
||||
name := fmt.Sprintf("%s.%d.json", filepath.Base(s.path), time.Now().UTC().UnixNano())
|
||||
backupPath := filepath.Join(s.options.BackupDir, name)
|
||||
temporary, err := os.CreateTemp(s.options.BackupDir, ".wxagent-backup-*")
|
||||
if err != nil {
|
||||
return fmt.Errorf("create backup file: %w", err)
|
||||
}
|
||||
temporaryName := temporary.Name()
|
||||
defer os.Remove(temporaryName)
|
||||
if err := temporary.Chmod(0o600); err != nil {
|
||||
_ = temporary.Close()
|
||||
return fmt.Errorf("protect backup file: %w", err)
|
||||
}
|
||||
if _, err := temporary.Write(data); err != nil {
|
||||
_ = temporary.Close()
|
||||
return fmt.Errorf("write backup file: %w", err)
|
||||
}
|
||||
if err := temporary.Sync(); err != nil {
|
||||
_ = temporary.Close()
|
||||
return fmt.Errorf("sync backup file: %w", err)
|
||||
}
|
||||
if err := temporary.Close(); err != nil {
|
||||
return fmt.Errorf("close backup file: %w", err)
|
||||
}
|
||||
if err := os.Rename(temporaryName, backupPath); err != nil {
|
||||
return fmt.Errorf("publish backup file: %w", err)
|
||||
}
|
||||
if err := pruneBackups(s.options.BackupDir, filepath.Base(s.path), s.options.BackupCount); err != nil {
|
||||
return fmt.Errorf("prune backups: %w", err)
|
||||
}
|
||||
if err := syncDirectory(s.options.BackupDir); err != nil {
|
||||
return err
|
||||
}
|
||||
s.lastBackup = now
|
||||
return nil
|
||||
}
|
||||
|
||||
func pruneBackups(directory, base string, keep int) error {
|
||||
entries, err := os.ReadDir(directory)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
type backup struct {
|
||||
name string
|
||||
when time.Time
|
||||
}
|
||||
backups := make([]backup, 0, len(entries))
|
||||
prefix := base + "."
|
||||
for _, entry := range entries {
|
||||
if entry.IsDir() || !strings.HasPrefix(entry.Name(), prefix) || !strings.HasSuffix(entry.Name(), ".json") {
|
||||
continue
|
||||
}
|
||||
info, err := entry.Info()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
backups = append(backups, backup{name: entry.Name(), when: info.ModTime()})
|
||||
}
|
||||
sort.Slice(backups, func(i, j int) bool {
|
||||
if backups[i].when.Equal(backups[j].when) {
|
||||
return backups[i].name > backups[j].name
|
||||
}
|
||||
return backups[i].when.After(backups[j].when)
|
||||
})
|
||||
if len(backups) <= keep {
|
||||
return nil
|
||||
}
|
||||
for _, item := range backups[keep:] {
|
||||
if err := os.Remove(filepath.Join(directory, item.name)); err != nil && !errors.Is(err, os.ErrNotExist) {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func syncDirectory(path string) error {
|
||||
directory, err := os.Open(path)
|
||||
if err != nil {
|
||||
return fmt.Errorf("open directory for sync: %w", err)
|
||||
}
|
||||
defer directory.Close()
|
||||
if err := directory.Sync(); err != nil && !errors.Is(err, syscall.EINVAL) {
|
||||
return fmt.Errorf("sync directory: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,94 @@
|
||||
package controlplane
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"crypto/tls"
|
||||
"crypto/x509"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"os"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func newTLSConfig(config ServerConfig) (*tls.Config, error) {
|
||||
if (config.TLSCertFile == "") != (config.TLSKeyFile == "") {
|
||||
return nil, errors.New("TLS cert and key must be configured together")
|
||||
}
|
||||
if config.TLSCertFile == "" {
|
||||
if config.MTLSClientCAFile != "" || config.MTLSRequireNodeCert || config.MTLSRevokedCertsFile != "" {
|
||||
return nil, errors.New("mTLS settings require TLS cert and key")
|
||||
}
|
||||
return nil, nil
|
||||
}
|
||||
certificate, err := tls.LoadX509KeyPair(config.TLSCertFile, config.TLSKeyFile)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("load TLS certificate: %w", err)
|
||||
}
|
||||
result := &tls.Config{
|
||||
MinVersion: tls.VersionTLS13,
|
||||
Certificates: []tls.Certificate{certificate},
|
||||
}
|
||||
if config.MTLSClientCAFile == "" {
|
||||
if config.MTLSRequireNodeCert || config.MTLSRevokedCertsFile != "" {
|
||||
return nil, errors.New("mTLS client CA is required for node certificates or revocation")
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
caBytes, err := os.ReadFile(config.MTLSClientCAFile)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("read mTLS client CA: %w", err)
|
||||
}
|
||||
pool := x509.NewCertPool()
|
||||
if !pool.AppendCertsFromPEM(caBytes) {
|
||||
return nil, errors.New("mTLS client CA does not contain a PEM certificate")
|
||||
}
|
||||
result.ClientCAs = pool
|
||||
// Web users may use HTTPS without a client certificate. Node routes enforce
|
||||
// a certificate separately, so one TLS listener serves both audiences.
|
||||
result.ClientAuth = tls.VerifyClientCertIfGiven
|
||||
if config.MTLSRevokedCertsFile != "" {
|
||||
if _, err := os.Stat(config.MTLSRevokedCertsFile); err != nil {
|
||||
return nil, fmt.Errorf("check revoked certificate file: %w", err)
|
||||
}
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (s *Server) clientCertificateAllowed(r *http.Request) bool {
|
||||
if r.TLS == nil || len(r.TLS.PeerCertificates) == 0 {
|
||||
return !s.config.MTLSRequireNodeCert
|
||||
}
|
||||
if s.config.MTLSRevokedCertsFile == "" {
|
||||
return true
|
||||
}
|
||||
revoked, err := readRevokedFingerprints(s.config.MTLSRevokedCertsFile)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
fingerprint := sha256.Sum256(r.TLS.PeerCertificates[0].Raw)
|
||||
_, isRevoked := revoked[hex.EncodeToString(fingerprint[:])]
|
||||
return !isRevoked
|
||||
}
|
||||
|
||||
func readRevokedFingerprints(path string) (map[string]struct{}, error) {
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result := make(map[string]struct{})
|
||||
for _, line := range strings.Split(string(data), "\n") {
|
||||
line = strings.TrimSpace(strings.SplitN(line, "#", 2)[0])
|
||||
if line == "" {
|
||||
continue
|
||||
}
|
||||
line = strings.ReplaceAll(strings.ToLower(line), ":", "")
|
||||
decoded, err := hex.DecodeString(line)
|
||||
if err != nil || len(decoded) != sha256.Size {
|
||||
return nil, errors.New("revoked certificate list contains an invalid SHA-256 fingerprint")
|
||||
}
|
||||
result[line] = struct{}{}
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
Reference in New Issue
Block a user