Phase 3+4: 共享 HTTP 客户端 + admin API + audit log
Phase 3 (httpclient): - shared Client (Timeout + MaxResponseBytes 截断) - SSRF: 16 段私网 CIDR (含 169.254/::ffff:0:0/96 IPv4-mapped) + hostname allowlist (精确/后缀/*. 通配) + scheme 校验 (仅 http/https) - auth 适配器: none / simple (YARN user.name) / basic (SetBasicAuth) - DoWithRedirect: 跨主机跳转保留 Authorization, max 5 默认 - 19 个测试全绿 (httptest 模拟) Phase 4 (admin + audit): - internal/audit: Entry + Repo.Insert/List + MarshalDetails, snake_case JSON - internal/middleware: AdminAuth/AgentAuth (constant-time 比对) - internal/admin: 6 端点 (GET/POST/PUT/DELETE /admin/clusters, GET /admin/audit) + 写操作触发 audit_log (含 before/after diff) + AuthPassword 空=保留旧密码 + 校验: 必填字段 + auth_type 枚举 + rate_limit>=0 - main.go: 挂载 storage.Open + admin.Mount - L fix: audit.Entry 加 json tag (smoke test 发现大写 key 不规范) Co-Authored-By: Claude <noreply@anthropic.com>
This commit is contained in:
@@ -0,0 +1,394 @@
|
||||
package admin
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"spark-mcp-go/internal/audit"
|
||||
"spark-mcp-go/internal/storage"
|
||||
)
|
||||
|
||||
func init() {
|
||||
gin.SetMode(gin.TestMode)
|
||||
}
|
||||
|
||||
func newTestAdmin(t *testing.T) (*gin.Engine, *storage.ClusterRepo, *audit.Repo) {
|
||||
t.Helper()
|
||||
db, err := storage.Open(":memory:")
|
||||
if err != nil {
|
||||
t.Fatalf("open: %v", err)
|
||||
}
|
||||
t.Cleanup(func() { _ = db.Close() })
|
||||
r := gin.New()
|
||||
Mount(r, db.Clusters(), audit.NewRepo(db), []string{"good-token"})
|
||||
return r, db.Clusters(), audit.NewRepo(db)
|
||||
}
|
||||
|
||||
func doReq(t *testing.T, r *gin.Engine, method, path, token string, body any) *httptest.ResponseRecorder {
|
||||
t.Helper()
|
||||
var rdr io.Reader
|
||||
if body != nil {
|
||||
b, _ := json.Marshal(body)
|
||||
rdr = bytes.NewReader(b)
|
||||
}
|
||||
req := httptest.NewRequest(method, path, rdr)
|
||||
if token != "" {
|
||||
req.Header.Set("Authorization", "Bearer "+token)
|
||||
}
|
||||
if body != nil {
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
}
|
||||
w := httptest.NewRecorder()
|
||||
r.ServeHTTP(w, req)
|
||||
return w
|
||||
}
|
||||
|
||||
func fullClusterMap(id string) map[string]any {
|
||||
return map[string]any{
|
||||
"id": id,
|
||||
"name": id + "-name",
|
||||
"rm_url": "http://rm.example.com:8088",
|
||||
"shs_url": "http://shs.example.com:18080",
|
||||
"spark_submit_execute_bin": "/usr/bin/spark-submit",
|
||||
"is_active": true,
|
||||
"auth_type": "basic",
|
||||
"auth_username": "admin",
|
||||
"auth_password": "real-secret",
|
||||
"ssl_verify": false,
|
||||
"ssl_ca_bundle": "",
|
||||
"url_allowlist": []string{"*.example.com"},
|
||||
"default_submit_args": []string{"--master", "yarn"},
|
||||
"rate_limit_per_min": 10,
|
||||
}
|
||||
}
|
||||
|
||||
func TestMount_RequiresAuth(t *testing.T) {
|
||||
r, _, _ := newTestAdmin(t)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
method string
|
||||
path string
|
||||
token string
|
||||
want int
|
||||
}{
|
||||
{"no header", "GET", "/admin/clusters", "", http.StatusUnauthorized},
|
||||
{"wrong scheme", "GET", "/admin/clusters", "Basic xxx", http.StatusUnauthorized},
|
||||
{"wrong token", "GET", "/admin/clusters", "wrong", http.StatusUnauthorized},
|
||||
{"valid token", "GET", "/admin/clusters", "good-token", http.StatusOK},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
w := doReq(t, r, tt.method, tt.path, tt.token, nil)
|
||||
if w.Code != tt.want {
|
||||
t.Errorf("got status %d, want %d", w.Code, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestListClusters(t *testing.T) {
|
||||
r, _, _ := newTestAdmin(t)
|
||||
|
||||
doReq(t, r, "POST", "/admin/clusters", "good-token", fullClusterMap("c1"))
|
||||
doReq(t, r, "POST", "/admin/clusters", "good-token", fullClusterMap("c2"))
|
||||
|
||||
w := doReq(t, r, "GET", "/admin/clusters", "good-token", nil)
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("got status %d, want %d", w.Code, http.StatusOK)
|
||||
}
|
||||
|
||||
var list []map[string]any
|
||||
if err := json.Unmarshal(w.Body.Bytes(), &list); err != nil {
|
||||
t.Fatalf("unmarshal list: %v", err)
|
||||
}
|
||||
if len(list) != 2 {
|
||||
t.Errorf("got %d clusters, want 2", len(list))
|
||||
}
|
||||
|
||||
for _, cl := range list {
|
||||
if _, ok := cl["auth_password"]; ok {
|
||||
t.Errorf("response must not contain auth_password")
|
||||
}
|
||||
}
|
||||
|
||||
body := w.Body.String()
|
||||
if strings.Contains(body, "real-secret") {
|
||||
t.Errorf("response body leaks auth_password")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateCluster(t *testing.T) {
|
||||
r, _, auditRepo := newTestAdmin(t)
|
||||
|
||||
w := doReq(t, r, "POST", "/admin/clusters", "good-token", fullClusterMap("c1"))
|
||||
if w.Code != http.StatusCreated {
|
||||
t.Fatalf("got status %d, want %d", w.Code, http.StatusCreated)
|
||||
}
|
||||
|
||||
var cl map[string]any
|
||||
if err := json.Unmarshal(w.Body.Bytes(), &cl); err != nil {
|
||||
t.Fatalf("unmarshal cluster: %v", err)
|
||||
}
|
||||
if cl["id"] != "c1" || cl["name"] != "c1-name" {
|
||||
t.Errorf("unexpected cluster fields: id=%v name=%v", cl["id"], cl["name"])
|
||||
}
|
||||
|
||||
entries, err := auditRepo.List(context.Background(), 100)
|
||||
if err != nil {
|
||||
t.Fatalf("list audit: %v", err)
|
||||
}
|
||||
if len(entries) < 1 {
|
||||
t.Fatalf("got %d audit entries, want at least 1", len(entries))
|
||||
}
|
||||
found := false
|
||||
for _, e := range entries {
|
||||
if e.Action == audit.ActionClusterCreate && e.Actor == "admin:good-tok" {
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Errorf("missing create audit entry from admin:good-tok")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCreateCluster_ValidationErrors(t *testing.T) {
|
||||
r, _, _ := newTestAdmin(t)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
body map[string]any
|
||||
}{
|
||||
{
|
||||
"missing id",
|
||||
func() map[string]any {
|
||||
m := fullClusterMap("c1")
|
||||
delete(m, "id")
|
||||
return m
|
||||
}(),
|
||||
},
|
||||
{
|
||||
"missing name",
|
||||
func() map[string]any {
|
||||
m := fullClusterMap("c1")
|
||||
delete(m, "name")
|
||||
return m
|
||||
}(),
|
||||
},
|
||||
{
|
||||
"invalid auth_type",
|
||||
func() map[string]any {
|
||||
m := fullClusterMap("c1")
|
||||
m["auth_type"] = "kerberos"
|
||||
return m
|
||||
}(),
|
||||
},
|
||||
{
|
||||
"negative rate limit",
|
||||
func() map[string]any {
|
||||
m := fullClusterMap("c1")
|
||||
m["rate_limit_per_min"] = -1
|
||||
return m
|
||||
}(),
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
w := doReq(t, r, "POST", "/admin/clusters", "good-token", tt.body)
|
||||
if w.Code != http.StatusBadRequest {
|
||||
t.Errorf("got status %d, want %d", w.Code, http.StatusBadRequest)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetCluster(t *testing.T) {
|
||||
r, _, _ := newTestAdmin(t)
|
||||
doReq(t, r, "POST", "/admin/clusters", "good-token", fullClusterMap("c1"))
|
||||
|
||||
w := doReq(t, r, "GET", "/admin/clusters/c1", "good-token", nil)
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("got status %d, want %d", w.Code, http.StatusOK)
|
||||
}
|
||||
var cl map[string]any
|
||||
if err := json.Unmarshal(w.Body.Bytes(), &cl); err != nil {
|
||||
t.Fatalf("unmarshal cluster: %v", err)
|
||||
}
|
||||
if cl["id"] != "c1" || cl["name"] != "c1-name" {
|
||||
t.Errorf("unexpected cluster fields: id=%v name=%v", cl["id"], cl["name"])
|
||||
}
|
||||
|
||||
w = doReq(t, r, "GET", "/admin/clusters/not-found", "good-token", nil)
|
||||
if w.Code != http.StatusNotFound {
|
||||
t.Errorf("got status %d, want %d", w.Code, http.StatusNotFound)
|
||||
}
|
||||
var errBody map[string]any
|
||||
if err := json.Unmarshal(w.Body.Bytes(), &errBody); err != nil {
|
||||
t.Fatalf("unmarshal error body: %v", err)
|
||||
}
|
||||
if errBody["error"] != "cluster not found" {
|
||||
t.Errorf("unexpected error message: %v", errBody["error"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateCluster(t *testing.T) {
|
||||
r, _, auditRepo := newTestAdmin(t)
|
||||
doReq(t, r, "POST", "/admin/clusters", "good-token", fullClusterMap("c1"))
|
||||
|
||||
update := map[string]any{
|
||||
"name": "updated",
|
||||
"rate_limit_per_min": 99,
|
||||
}
|
||||
w := doReq(t, r, "PUT", "/admin/clusters/c1", "good-token", update)
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("got status %d, want %d: %s", w.Code, http.StatusOK, w.Body.String())
|
||||
}
|
||||
var cl map[string]any
|
||||
if err := json.Unmarshal(w.Body.Bytes(), &cl); err != nil {
|
||||
t.Fatalf("unmarshal cluster: %v", err)
|
||||
}
|
||||
if cl["name"] != "updated" || cl["rate_limit_per_min"] != float64(99) {
|
||||
t.Errorf("unexpected update response: name=%v rate_limit_per_min=%v", cl["name"], cl["rate_limit_per_min"])
|
||||
}
|
||||
|
||||
w = doReq(t, r, "GET", "/admin/clusters/c1", "good-token", nil)
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("get after update: got status %d", w.Code)
|
||||
}
|
||||
if err := json.Unmarshal(w.Body.Bytes(), &cl); err != nil {
|
||||
t.Fatalf("unmarshal cluster after update: %v", err)
|
||||
}
|
||||
if cl["name"] != "updated" || cl["rate_limit_per_min"] != float64(99) {
|
||||
t.Errorf("cluster not updated: name=%v rate_limit_per_min=%v", cl["name"], cl["rate_limit_per_min"])
|
||||
}
|
||||
|
||||
entries, err := auditRepo.List(context.Background(), 100)
|
||||
if err != nil {
|
||||
t.Fatalf("list audit: %v", err)
|
||||
}
|
||||
found := false
|
||||
for _, e := range entries {
|
||||
if e.Action == audit.ActionClusterUpdate {
|
||||
found = true
|
||||
var details map[string]any
|
||||
if err := json.Unmarshal([]byte(e.Details), &details); err != nil {
|
||||
t.Fatalf("unmarshal audit details: %v", err)
|
||||
}
|
||||
if _, ok := details["before"]; !ok {
|
||||
t.Errorf("audit details missing before")
|
||||
}
|
||||
if _, ok := details["after"]; !ok {
|
||||
t.Errorf("audit details missing after")
|
||||
}
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Errorf("missing update audit entry")
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateCluster_PreservesEmptyPassword(t *testing.T) {
|
||||
r, repo, auditRepo := newTestAdmin(t)
|
||||
doReq(t, r, "POST", "/admin/clusters", "good-token", fullClusterMap("c1"))
|
||||
|
||||
// AuthPassword has json:"-", so POST cannot receive it via JSON.
|
||||
// Seed the stored password directly through the repo.
|
||||
cl, err := repo.Get(context.Background(), "c1")
|
||||
if err != nil {
|
||||
t.Fatalf("get cluster from repo: %v", err)
|
||||
}
|
||||
cl.AuthPassword = "real-secret"
|
||||
if err := repo.Update(context.Background(), cl); err != nil {
|
||||
t.Fatalf("seed password: %v", err)
|
||||
}
|
||||
// Discard the audit row produced by the seed update so it does not confuse later checks.
|
||||
_, _ = auditRepo.List(context.Background(), 100)
|
||||
|
||||
update := map[string]any{
|
||||
"name": "renamed",
|
||||
}
|
||||
w := doReq(t, r, "PUT", "/admin/clusters/c1", "good-token", update)
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("got status %d, want %d: %s", w.Code, http.StatusOK, w.Body.String())
|
||||
}
|
||||
|
||||
if strings.Contains(w.Body.String(), "real-secret") {
|
||||
t.Errorf("update response leaks auth_password")
|
||||
}
|
||||
|
||||
cl, err = repo.Get(context.Background(), "c1")
|
||||
if err != nil {
|
||||
t.Fatalf("get cluster from repo: %v", err)
|
||||
}
|
||||
if cl.AuthPassword != "real-secret" {
|
||||
t.Errorf("password changed: got %q, want %q", cl.AuthPassword, "real-secret")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDeleteCluster(t *testing.T) {
|
||||
r, _, auditRepo := newTestAdmin(t)
|
||||
doReq(t, r, "POST", "/admin/clusters", "good-token", fullClusterMap("c1"))
|
||||
|
||||
w := doReq(t, r, "DELETE", "/admin/clusters/c1", "good-token", nil)
|
||||
if w.Code != http.StatusNoContent {
|
||||
t.Fatalf("got status %d, want %d", w.Code, http.StatusNoContent)
|
||||
}
|
||||
|
||||
w = doReq(t, r, "GET", "/admin/clusters/c1", "good-token", nil)
|
||||
if w.Code != http.StatusNotFound {
|
||||
t.Errorf("get after delete: got status %d, want %d", w.Code, http.StatusNotFound)
|
||||
}
|
||||
|
||||
entries, err := auditRepo.List(context.Background(), 100)
|
||||
if err != nil {
|
||||
t.Fatalf("list audit: %v", err)
|
||||
}
|
||||
found := false
|
||||
for _, e := range entries {
|
||||
if e.Action == audit.ActionClusterDelete {
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Errorf("missing delete audit entry")
|
||||
}
|
||||
}
|
||||
|
||||
func TestListAudit(t *testing.T) {
|
||||
r, _, _ := newTestAdmin(t)
|
||||
doReq(t, r, "POST", "/admin/clusters", "good-token", fullClusterMap("c1"))
|
||||
doReq(t, r, "PUT", "/admin/clusters/c1", "good-token", map[string]any{"name": "updated"})
|
||||
doReq(t, r, "DELETE", "/admin/clusters/c1", "good-token", nil)
|
||||
|
||||
w := doReq(t, r, "GET", "/admin/audit", "good-token", nil)
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("got status %d, want %d", w.Code, http.StatusOK)
|
||||
}
|
||||
|
||||
var entries []*audit.Entry
|
||||
if err := json.Unmarshal(w.Body.Bytes(), &entries); err != nil {
|
||||
t.Fatalf("unmarshal audit list: %v", err)
|
||||
}
|
||||
if len(entries) != 3 {
|
||||
t.Fatalf("got %d audit entries, want 3", len(entries))
|
||||
}
|
||||
|
||||
wantActions := []string{string(audit.ActionClusterDelete), string(audit.ActionClusterUpdate), string(audit.ActionClusterCreate)}
|
||||
for i, want := range wantActions {
|
||||
got := string(entries[i].Action)
|
||||
if got != want {
|
||||
t.Errorf("entry[%d].action=%s, want %s", i, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user