package admin_test import ( "context" "errors" "net/http" "net/http/cookiejar" "net/http/httptest" "path/filepath" "sync" "testing" "git.asio.asia/nixevol/NixMsg/internal/admin" "git.asio.asia/nixevol/NixMsg/internal/app/identity" "git.asio.asia/nixevol/NixMsg/internal/auth" "git.asio.asia/nixevol/NixMsg/internal/store" ) type failDisableIdentity struct { identity.Stub err error } func (f *failDisableIdentity) Disable(context.Context, string) error { return f.err } func (f *failDisableIdentity) Enable(context.Context, string) error { return nil } type trackIdentity struct { identity.Stub mu sync.Mutex disableN int enableN int } func (t *trackIdentity) Disable(context.Context, string) error { t.mu.Lock() defer t.mu.Unlock() t.disableN++ return nil } func (t *trackIdentity) Enable(context.Context, string) error { t.mu.Lock() defer t.mu.Unlock() t.enableN++ return nil } func setupEndpointsWithIdentity(t *testing.T, ident identity.Service) (*store.DB, *httptest.Server, *http.Client) { t.Helper() dir := t.TempDir() db, err := store.Open(filepath.Join(dir, "data"), "FULL") if err != nil { t.Fatal(err) } t.Cleanup(func() { _ = db.Close() }) hash := auth.NewStubHashPool() if seedErr := admin.SeedAdminPassword(context.Background(), db, hash, testPassword); seedErr != nil { t.Fatal(seedErr) } h := admin.New(admin.Deps{ DB: db, Hash: hash, Tokens: admin.NewRandomAPITokens(), Locks: admin.NewMemoryLoginLocks(), Identity: ident, }) srv := httptest.NewServer(h) t.Cleanup(srv.Close) jar, err := cookiejar.New(nil) if err != nil { t.Fatal(err) } client := &http.Client{Jar: jar} res := postJSON(t, client, srv.URL+"/api/admin/login", `{"username":"admin","password":"`+testPassword+`"}`, nil) env := decodeEnv(t, res) if res.StatusCode != http.StatusOK || !env.OK { t.Fatalf("login failed: %d %+v", res.StatusCode, env) } return db, srv, client } func TestPatchEnabledFalseKeepsEnabledWhenIdentityFails(t *testing.T) { ident := &failDisableIdentity{err: errors.New("disable failed")} db, srv, client := setupEndpointsWithIdentity(t, ident) base := srv.URL res := postJSON(t, client, base+"/api/admin/endpoints", `{"id":"keep-on","name":"仍启用","login_password":"password1"}`, csrfHeaders()) env := decodeEnv(t, res) if res.StatusCode != 200 || !env.OK { t.Fatalf("create: %d %+v", res.StatusCode, env) } res = doReq(t, client, http.MethodPatch, base+"/api/admin/endpoints/keep-on", `{"enabled":false}`, csrfHeaders()) env = decodeEnv(t, res) if res.StatusCode != http.StatusInternalServerError { t.Fatalf("want 500 got %d %+v", res.StatusCode, env) } var enabled int if err := db.Read.QueryRow(`SELECT enabled FROM endpoints WHERE id='keep-on'`).Scan(&enabled); err != nil { t.Fatal(err) } if enabled != 1 { t.Fatalf("want still enabled, got %d", enabled) } } func TestPatchNameOnlyDoesNotTouchEnabled(t *testing.T) { ident := &trackIdentity{} db, srv, client := setupEndpointsWithIdentity(t, ident) base := srv.URL res := postJSON(t, client, base+"/api/admin/endpoints", `{"id":"name-only","name":"旧名","login_password":"password1"}`, csrfHeaders()) env := decodeEnv(t, res) if res.StatusCode != 200 || !env.OK { t.Fatalf("create: %d %+v", res.StatusCode, env) } res = doReq(t, client, http.MethodPatch, base+"/api/admin/endpoints/name-only", `{"name":"新名"}`, csrfHeaders()) env = decodeEnv(t, res) if res.StatusCode != 200 || !env.OK { t.Fatalf("patch: %d %+v", res.StatusCode, env) } ident.mu.Lock() d, e := ident.disableN, ident.enableN ident.mu.Unlock() if d != 0 || e != 0 { t.Fatalf("identity enable/disable should not run, disable=%d enable=%d", d, e) } var enabled int var name string if err := db.Read.QueryRow(`SELECT enabled, name FROM endpoints WHERE id='name-only'`).Scan(&enabled, &name); err != nil { t.Fatal(err) } if enabled != 1 || name != "新名" { t.Fatalf("enabled=%d name=%q", enabled, name) } }