Files
NixMsg/test/load/provision.go

171 lines
4.5 KiB
Go

package load
import (
"bytes"
"encoding/json"
"fmt"
"io"
"net/http"
"strings"
"time"
"git.asio.asia/nixevol/NixMsg/test/harness"
)
const maxImportRows = 1000
type apiEnvelope struct {
OK bool `json:"ok"`
Data map[string]any `json:"data"`
Error *struct {
Code string `json:"code"`
Message string `json:"message"`
} `json:"error"`
}
// Provision 按配置准备端:管理 API 批量开通、自助注册,或复用已开通编号。
func Provision(cfg Config, ids []string) error {
if cfg.Reuse && cfg.AdminUser == "" && cfg.AdminPass == "" && cfg.APIToken == "" && cfg.RegisterCode == "" {
return nil
}
if cfg.RegisterCode != "" && cfg.AdminPass == "" && cfg.APIToken == "" {
return registerAll(cfg, ids)
}
if cfg.AdminPass == "" && cfg.APIToken == "" {
return fmt.Errorf("开通端需要 -admin-pass / -admin-token,或 -register-code;复用已开通端请加 -reuse")
}
ac, err := adminLogin(cfg)
if err != nil {
return err
}
if cfg.Reuse {
return createOneByOne(ac, ids, cfg.Password, true)
}
return importCSV(ac, ids, cfg.Password)
}
func adminLogin(cfg Config) (*harness.AdminClient, error) {
base := cfg.adminBase()
ac, err := harness.NewAdminClient(base)
if err != nil {
return nil, err
}
ac.HTTP.Timeout = 10 * time.Minute
if cfg.APIToken != "" {
ac.APIToken = cfg.APIToken
return ac, nil
}
user := cfg.AdminUser
if user == "" {
user = "admin"
}
body, _ := json.Marshal(map[string]string{
"username": user,
"password": cfg.AdminPass,
})
resp, err := ac.PostJSON("/api/admin/login", body)
if err != nil {
return nil, fmt.Errorf("管理员登录: %w", err)
}
raw, _ := io.ReadAll(resp.Body)
_ = resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return nil, fmt.Errorf("管理员登录失败: %d %s", resp.StatusCode, raw)
}
return ac, nil
}
func importCSV(ac *harness.AdminClient, ids []string, password string) error {
for start := 0; start < len(ids); start += maxImportRows {
end := start + maxImportRows
if end > len(ids) {
end = len(ids)
}
var b strings.Builder
b.WriteString("id,name,login_password,talk_password,default_delay_seconds,remark\n")
for _, id := range ids[start:end] {
fmt.Fprintf(&b, "%s,%s,%s,,,\n", id, id, password)
}
resp, err := ac.Do(http.MethodPost, "/api/admin/endpoints/import", []byte(b.String()), "text/csv")
if err != nil {
return fmt.Errorf("批量开通: %w", err)
}
raw, _ := io.ReadAll(resp.Body)
_ = resp.Body.Close()
if resp.StatusCode != http.StatusOK {
return fmt.Errorf("批量开通失败: %d %s", resp.StatusCode, raw)
}
}
return nil
}
func createOneByOne(ac *harness.AdminClient, ids []string, password string, reuse bool) error {
for _, id := range ids {
body, _ := json.Marshal(map[string]string{
"id": id,
"login_password": password,
})
resp, err := ac.PostJSON("/api/admin/endpoints", body)
if err != nil {
return fmt.Errorf("开通 %s: %w", id, err)
}
raw, _ := io.ReadAll(resp.Body)
_ = resp.Body.Close()
if resp.StatusCode == http.StatusOK || resp.StatusCode == http.StatusCreated {
continue
}
if reuse && isTaken(resp.StatusCode, raw) {
continue
}
return fmt.Errorf("开通 %s 失败: %d %s", id, resp.StatusCode, raw)
}
return nil
}
func registerAll(cfg Config, ids []string) error {
base := strings.TrimRight(cfg.HTTPBase, "/")
client := &http.Client{Timeout: 30 * time.Second}
for _, id := range ids {
payload, _ := json.Marshal(map[string]string{
"registration_code": cfg.RegisterCode,
"id": id,
"login_password": cfg.Password,
})
resp, err := client.Post(base+"/api/client/register", "application/json", bytes.NewReader(payload))
if err != nil {
return fmt.Errorf("注册 %s: %w", id, err)
}
raw, _ := io.ReadAll(resp.Body)
_ = resp.Body.Close()
if resp.StatusCode == http.StatusOK || resp.StatusCode == http.StatusCreated {
continue
}
if cfg.Reuse && isTaken(resp.StatusCode, raw) {
continue
}
return fmt.Errorf("注册 %s 失败: %d %s", id, resp.StatusCode, raw)
}
return nil
}
func isTaken(status int, raw []byte) bool {
if status == http.StatusConflict {
return true
}
var env apiEnvelope
if json.Unmarshal(raw, &env) != nil {
return false
}
if env.Error != nil && (env.Error.Code == "id_taken" || env.Error.Code == "conflict") {
return true
}
return false
}
func (c Config) adminBase() string {
if c.AdminHTTP != "" {
return strings.TrimRight(c.AdminHTTP, "/")
}
return strings.TrimRight(c.HTTPBase, "/")
}