Files
openvpn-manager/backend/internal/service/service.go
T
cnbugs a5519c2f6c Feature: per-instance public_host for .ovpn
When creating/editing an instance, add '公网地址' (public_host) field.
This is the domain/IP clients will use to reach the OpenVPN server.
When downloading a .ovpn, this field is used first, falling back to:
  1. ?host= query parameter (one-time override)
  2. Request Host header
  3. 'vpn.example.com' default

This eliminates the need to manually edit .ovpn files after download
when the server's public address differs from its internal IP.

Backend changes:
  model.Instance: add PublicHost string field
  service.GenerateOVPN: priority PublicHost > remoteHost > default
  service.UpdateInstance: PublicHost persisted (already via JSON tag)

Frontend changes:
  Instances.vue: new '公网地址' input in create/edit dialog
  Instances.vue: new column in list table (shows '未配置' when empty)
  Users.vue: hint text updated to '可选: 覆盖实例公网地址'

E2E verified:
  public_host='vpn.yunwei.blog'  → .ovpn has 'remote vpn.yunwei.blog 1194'
  public_host='' + host=X        → .ovpn has 'remote X 1194'
  public_host set + no host      → public_host wins
2026-08-10 00:35:30 +08:00

832 lines
22 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package service
import (
"archive/tar"
"compress/gzip"
"context"
"crypto/rand"
"encoding/hex"
"fmt"
"io"
"log"
"os"
"os/exec"
"path/filepath"
"strings"
"time"
"github.com/gin-gonic/gin"
"github.com/google/uuid"
"golang.org/x/crypto/bcrypt"
"openvpn-manager/internal/config"
"openvpn-manager/internal/model"
"openvpn-manager/internal/store"
"openvpn-manager/pkg/openvpn"
)
// Service 业务逻辑聚合,供 API 层调用。
// 任何对实例/用户/证书/备份的变更都应经过这里,从而写入审计日志。
type Service struct {
Cfg *config.Config
Store *store.Store
Ovm *openvpn.Manager
}
func New(cfg *config.Config, st *store.Store, ovm *openvpn.Manager) *Service {
return &Service{Cfg: cfg, Store: st, Ovm: ovm}
}
// StartStatusSync 启动后台定时任务,每 10 秒读取所有实例的 status.log
// 解析 CLIENT_LIST 并同步到 ConnLogs。应在 main.go 中以 goroutine 启动。
func (s *Service) StartStatusSync() {
go func() {
ticker := time.NewTicker(10 * time.Second)
defer ticker.Stop()
for range ticker.C {
s.syncAllInstanceStatus()
}
}()
}
func (s *Service) syncAllInstanceStatus() {
instances := s.Store.ListInstances()
for _, in := range instances {
// 不检查 Status 字段,直接尝试读 status.log。
// OpenVPN 运行时会自动写 status.log,读到就同步。
statusPath := filepath.Join(s.Cfg.InstanceDir(in.Name), "status.log")
entries, err := s.Ovm.ParseStatus(statusPath)
if err != nil {
// 文件不存在或解析失败,跳过
continue
}
log.Printf("[status-sync] instance=%s entries=%d", in.Name, len(entries))
// 构建当前在线 CN 集合
onlineCNs := map[string]bool{}
for _, e := range entries {
onlineCNs[e.CommonName] = true
// 如果没有活跃连接记录,则新建
if s.Store.FindActiveConn(in.ID, e.CommonName) == nil {
_ = s.Store.AppendConnLog(model.ConnectionLog{
InstanceID: in.ID,
CommonName: e.CommonName,
RealIP: e.RealAddress,
VPNIP: e.VPNAddress,
BytesIn: e.BytesRecv,
BytesOut: e.BytesSent,
ConnectedAt: e.ConnectedAt,
})
log.Printf("[status-sync] new conn: %s %s %s", e.CommonName, e.RealAddress, e.VPNAddress)
} else {
// 更新流量统计
_ = s.Store.UpdateActiveConnStats(in.ID, e.CommonName, e.BytesRecv, e.BytesSent)
}
}
// 标记已断开的连接
for _, c := range s.Store.ListActiveConns(in.ID) {
if !onlineCNs[c.CommonName] {
now := time.Now()
_ = s.Store.CloseActiveConn(in.ID, c.CommonName, now)
log.Printf("[status-sync] disconnected: %s", c.CommonName)
}
}
}
}
func (s *Service) audit(c context.Context, action, target, detail, result, ip string) {
username, _ := c.Value("user").(string)
if username == "" {
username = "system"
}
_ = s.Store.AppendAudit(model.AuditLog{
ID: uuid.NewString(),
Time: time.Now(),
User: username,
Action: action,
Target: target,
Result: result,
Detail: detail,
IP: ip,
})
}
// AuditForAPI 在 API 层被调用时手动写入(因为 gin context 转为 context.Context)。
func (s *Service) AuditForAPI(c *gin.Context, action, target, detail, result string) {
username, _ := c.Get("user")
un, _ := username.(string)
if un == "" {
un = "system"
}
_ = s.Store.AppendAudit(model.AuditLog{
ID: uuid.NewString(),
Time: time.Now(),
User: un,
Action: action,
Target: target,
Result: result,
Detail: detail,
IP: c.ClientIP(),
})
}
// CreateInstance 新建一个 OpenVPN 实例,并签发服务端证书、生成 server.conf。
func (s *Service) CreateInstance(in model.Instance) (*model.Instance, error) {
if in.Name == "" {
return nil, fmt.Errorf("name required")
}
if in.Port == 0 {
return nil, fmt.Errorf("port required")
}
if in.Proto == "" {
in.Proto = "udp"
}
if in.Dev == "" {
in.Dev = "tun"
}
if in.Subnet == "" {
in.Subnet = "10.8.0.0/24"
}
if in.AccessMode == "" {
in.AccessMode = model.AccessOpen
}
if in.AuthMode == "" {
in.AuthMode = model.AuthCertPassword // 默认双因素(更安全)
}
// 校验白名单 CIDR
for _, c := range in.AllowNetworks {
if err := openvpn.ValidateCIDR(c); err != nil {
return nil, fmt.Errorf("instance allow_networks: %w", err)
}
}
// 名称查重
if _, err := s.Store.GetInstanceByName(in.Name); err == nil {
return nil, fmt.Errorf("instance %s already exists", in.Name)
}
// 确保 CA
if err := s.Ovm.EnsureCA(); err != nil {
return nil, err
}
in.ID = uuid.NewString()
in.Status = "stopped"
in.CreatedAt = time.Now()
in.UpdatedAt = time.Now()
// 创建实例目录
if err := os.MkdirAll(s.Cfg.InstanceDir(in.Name), 0o755); err != nil {
return nil, err
}
if err := os.MkdirAll(filepath.Join(s.Cfg.InstanceDir(in.Name), "logs"), 0o755); err != nil {
return nil, err
}
if err := os.MkdirAll(filepath.Join(s.Cfg.InstanceDir(in.Name), "ccd"), 0o755); err != nil {
return nil, err
}
// 签发服务端证书
if err := s.Ovm.IssueServerCert(in.Name); err != nil {
return nil, fmt.Errorf("issue server cert: %w", err)
}
// 写 server.conf + 白名单脚本
ccdDir := filepath.Join(s.Cfg.InstanceDir(in.Name), "ccd")
if err := s.Ovm.WriteServerConf(&in, ccdDir); err != nil {
return nil, fmt.Errorf("write conf: %w", err)
}
if in.AccessMode == model.AccessWhitelist {
if _, err := s.Ovm.AllowNetworksScript(in.Name, in.AllowNetworks); err != nil {
return nil, fmt.Errorf("gen whitelist scripts: %w", err)
}
}
if err := s.Store.UpsertInstance(in); err != nil {
return nil, err
}
return &in, nil
}
// UpdateInstance 仅更新可热改字段(端口/子网需要重启生效)。
func (s *Service) UpdateInstance(in model.Instance) error {
old, err := s.Store.GetInstance(in.ID)
if err != nil {
return err
}
if in.AccessMode == "" {
in.AccessMode = old.AccessMode
}
if in.Cipher == "" {
in.Cipher = old.Cipher
}
if in.AuthDigest == "" {
in.AuthDigest = old.AuthDigest
}
if in.Dev == "" {
in.Dev = old.Dev
}
if in.Subnet == "" {
in.Subnet = old.Subnet
}
for _, c := range in.AllowNetworks {
if err := openvpn.ValidateCIDR(c); err != nil {
return fmt.Errorf("instance allow_networks: %w", err)
}
}
in.CreatedAt = old.CreatedAt
in.UpdatedAt = time.Now()
in.Status = old.Status
in.PID = old.PID
ccdDir := filepath.Join(s.Cfg.InstanceDir(in.Name), "ccd")
if err := s.Ovm.WriteServerConf(&in, ccdDir); err != nil {
return err
}
if in.AccessMode == model.AccessWhitelist {
if _, err := s.Ovm.AllowNetworksScript(in.Name, in.AllowNetworks); err != nil {
return err
}
}
if err := s.Store.UpsertInstance(in); err != nil {
return err
}
// 实例级 allow_networks 变更后,刷新所有用户的 ccd
inst := in
for _, u := range s.Store.ListUsers(inst.ID) {
u := u
merged := openvpn.MergeAllowNetworks(inst.AllowNetworks, u.AllowNetworks)
if err := s.writeUserCCD(&inst, &u, merged); err != nil {
return err
}
}
return nil
}
// DeleteInstance 移除实例及其 PKI/配置。
func (s *Service) DeleteInstance(id string) error {
in, err := s.Store.GetInstance(id)
if err != nil {
return err
}
if in.Status == "running" {
_ = s.StopInstance(id)
}
// 清理用户记录与目录
for _, u := range s.Store.ListUsers(id) {
_ = s.Store.DeleteUser(u.ID)
}
_ = os.RemoveAll(s.Cfg.InstanceDir(in.Name))
_ = os.RemoveAll(filepath.Join(s.Cfg.ClientsDir(), in.Name))
return s.Store.DeleteInstance(id)
}
// StartInstance 在前台启动 openvpn。
// 注意:本服务应以 root 运行;非 root 场景下应通过 systemd 单元托管。
func (s *Service) StartInstance(id string) error {
in, err := s.Store.GetInstance(id)
if err != nil {
return err
}
if in.Status == "running" {
return fmt.Errorf("already running")
}
conf := s.Ovm.InstanceConf(in.Name)
if _, err := os.Stat(conf); err != nil {
return fmt.Errorf("conf missing: %w", err)
}
cmd := exec.Command(s.Cfg.OpenVPNBin,
"--cd", s.Cfg.InstanceDir(in.Name),
"--config", conf)
logf, _ := os.OpenFile(filepath.Join(s.Cfg.InstanceDir(in.Name), "logs", "openvpn-stdout.log"),
os.O_CREATE|os.O_APPEND|os.O_WRONLY, 0o644)
if logf != nil {
cmd.Stdout = logf
cmd.Stderr = logf
}
if err := cmd.Start(); err != nil {
return err
}
in.Status = "running"
in.PID = cmd.Process.Pid
in.UpdatedAt = time.Now()
_ = s.Store.UpsertInstance(*in)
// 后台释放
go func() {
_ = cmd.Wait()
// 进程退出时回写状态(简单模型)
cur, err := s.Store.GetInstance(id)
if err == nil && cur.PID == cmd.Process.Pid {
cur.Status = "stopped"
cur.PID = 0
cur.UpdatedAt = time.Now()
_ = s.Store.UpsertInstance(*cur)
}
}()
return nil
}
// StopInstance 通过 SIGTERM 停止实例。
func (s *Service) StopInstance(id string) error {
in, err := s.Store.GetInstance(id)
if err != nil {
return err
}
if in.PID == 0 {
in.Status = "stopped"
return s.Store.UpsertInstance(*in)
}
proc, err := os.FindProcess(in.PID)
if err != nil {
return err
}
if err := proc.Signal(os.Interrupt); err != nil {
// 兜底:直接 Kill
_ = proc.Signal(os.Kill)
}
in.Status = "stopped"
in.PID = 0
in.UpdatedAt = time.Now()
return s.Store.UpsertInstance(*in)
}
// CreateUser 新建客户端用户并签发证书。
func (s *Service) CreateUser(u model.VPNUser) (*model.VPNUser, error) {
if u.Username == "" {
return nil, fmt.Errorf("username required")
}
in, err := s.Store.GetInstance(u.InstanceID)
if err != nil {
return nil, err
}
if _, err := s.Store.GetUserByCN(u.InstanceID, u.Username); err == nil {
return nil, fmt.Errorf("user %s already exists", u.Username)
}
// 校验用户级 CIDR
for _, c := range u.AllowNetworks {
if err := openvpn.ValidateCIDR(c); err != nil {
return nil, fmt.Errorf("user allow_networks: %w", err)
}
}
// 双因素模式必须有密码
if in.AuthMode == model.AuthCertPassword {
if u.Password == "" {
return nil, fmt.Errorf("password required for cert+password auth mode")
}
}
// 哈希密码(bcrypt),即使 cert 模式也存储(便于后续切换)
if u.Password != "" {
if len(u.Password) < 4 {
return nil, fmt.Errorf("password too short (min 4)")
}
hash, err := bcrypt.GenerateFromPassword([]byte(u.Password), bcrypt.DefaultCost)
if err != nil {
return nil, fmt.Errorf("hash password: %w", err)
}
u.PasswordHash = string(hash)
}
u.ID = uuid.NewString()
u.Enabled = true
u.CreatedAt = time.Now()
u.Password = "" // 不持久化明文
// 签发证书(Manager 按实例名索引 PKI
if _, _, err := s.Ovm.IssueCert(in.Name, u.Username); err != nil {
return nil, err
}
// CCD: 合并实例级与用户级白名单(白名单模式才生效)
merged := openvpn.MergeAllowNetworks(in.AllowNetworks, u.AllowNetworks)
if err := s.writeUserCCD(in, &u, merged); err != nil {
return nil, err
}
// 预生成 ovpn(以空 host 生成占位,用户在 UI 上下载)
if _, err := s.Ovm.GenerateClientOVPNFor(&u, in, "vpn.example.com"); err != nil {
return nil, err
}
if err := s.Store.UpsertUser(u); err != nil {
return nil, err
}
u.PasswordHash = "" // 不向外暴露
return &u, nil
}
// GenerateOVPN 下载/重新生成 .ovpn。
// remoteHost 优先级:Instance.PublicHost > 传入的 remoteHost > "vpn.example.com"。
func (s *Service) GenerateOVPN(userID, remoteHost string) (string, error) {
u, err := s.Store.GetUser(userID)
if err != nil {
return "", err
}
in, err := s.Store.GetInstance(u.InstanceID)
if err != nil {
return "", err
}
host := in.PublicHost
if host == "" {
host = remoteHost
}
return s.Ovm.GenerateClientOVPNFor(u, in, host)
}
// RevokeUser 吊销用户:禁用 + 标记。
func (s *Service) RevokeUser(userID string) error {
u, err := s.Store.GetUser(userID)
if err != nil {
return err
}
u.Enabled = false
now := time.Now()
u.RevokedAt = &now
if err := s.Store.UpsertUser(*u); err != nil {
return err
}
// 在 ccd 写入禁用标记
in, err2 := s.Store.GetInstance(u.InstanceID)
if err2 == nil {
body := "# revoked by manager\n"
_ = s.Ovm.WriteCCD(in.Name, u.Username, body)
}
return nil
}
// DeleteUser 删除用户及证书。
func (s *Service) DeleteUser(userID string) error {
u, err := s.Store.GetUser(userID)
if err != nil {
return err
}
in, _ := s.Store.GetInstance(u.InstanceID)
if in != nil {
_ = s.Ovm.DeleteCCD(in.Name, u.Username)
_ = os.Remove(filepath.Join(s.Ovm.PKIPath(in.Name), "issued", u.Username+".crt"))
_ = os.Remove(filepath.Join(s.Ovm.PKIPath(in.Name), "private", u.Username+".key"))
_ = os.Remove(filepath.Join(s.Cfg.ClientsDir(), in.Name, u.Username+".ovpn"))
}
return s.Store.DeleteUser(userID)
}
// UpdateUser 修改用户(主要用于改 allow_networks / static_ip 等)。
func (s *Service) UpdateUser(u model.VPNUser) error {
old, err := s.Store.GetUser(u.ID)
if err != nil {
return err
}
for _, c := range u.AllowNetworks {
if err := openvpn.ValidateCIDR(c); err != nil {
return fmt.Errorf("user allow_networks: %w", err)
}
}
u.CreatedAt = old.CreatedAt
u.Enabled = old.Enabled
u.PasswordHash = old.PasswordHash // 保留原密码哈希,不被空覆盖
u.Password = "" // 不持久化明文
u.RevokedAt = old.RevokedAt
in, err := s.Store.GetInstance(u.InstanceID)
if err != nil {
return err
}
merged := openvpn.MergeAllowNetworks(in.AllowNetworks, u.AllowNetworks)
if err := s.writeUserCCD(in, &u, merged); err != nil {
return err
}
return s.Store.UpsertUser(u)
}
// ResetVPNPassword 由管理员重置 VPN 用户的登录密码。
func (s *Service) ResetVPNPassword(userID, newPassword string) error {
if len(newPassword) < 4 {
return fmt.Errorf("password too short (min 4)")
}
u, err := s.Store.GetUser(userID)
if err != nil {
return err
}
hash, err := bcrypt.GenerateFromPassword([]byte(newPassword), bcrypt.DefaultCost)
if err != nil {
return err
}
u.PasswordHash = string(hash)
return s.Store.UpsertUser(*u)
}
// writeUserCCD 把允许网段写进 ccd/<cn>,client-connect 脚本读取后推送 route。
// 文件内容:
// - 第 1 行 ifconfig-push (固定 IP)
// - 之后每行一个 CIDR (允许的网段)
// 这样设计既兼容现有 static_ip 场景,又能让 client-connect.sh 简单 grep。
func (s *Service) writeUserCCD(in *model.Instance, u *model.VPNUser, allowNets []string) error {
var b strings.Builder
if u.StaticIP != "" {
b.WriteString("ifconfig-push " + u.StaticIP + " 255.255.255.0\n")
}
for _, c := range allowNets {
b.WriteString(c + "\n")
}
return s.Ovm.WriteCCD(in.Name, u.Username, b.String())
}
// CertInfos 汇总所有用户证书的过期时间。
func (s *Service) CertInfos() ([]model.CertInfo, error) {
var out []model.CertInfo
now := time.Now()
for _, in := range s.Store.ListInstances() {
pki := s.Ovm.PKIPath(in.Name)
issuedDir := filepath.Join(pki, "issued")
entries, err := os.ReadDir(issuedDir)
if err != nil {
continue
}
for _, e := range entries {
if e.IsDir() || !strings.HasSuffix(e.Name(), ".crt") {
continue
}
cn := strings.TrimSuffix(e.Name(), ".crt")
if cn == "server" {
continue
}
notAfter, err := s.Ovm.CertNotAfter(filepath.Join(issuedDir, e.Name()))
if err != nil {
continue
}
out = append(out, model.CertInfo{
InstanceID: in.ID,
Username: cn,
NotAfter: notAfter,
DaysLeft: int(notAfter.Sub(now).Hours() / 24),
})
}
}
return out, nil
}
// ListOnline 解析 status 文件获取在线客户端。
func (s *Service) ListOnline(instanceID string) ([]openvpn.StatusEntry, error) {
in, err := s.Store.GetInstance(instanceID)
if err != nil {
return nil, err
}
statusPath := filepath.Join(s.Cfg.InstanceDir(in.Name), "status.log")
return s.Ovm.ParseStatus(statusPath)
}
// Backup 创建 tar.gz 备份。
func (s *Service) Backup(note string) (*model.Backup, error) {
if err := os.MkdirAll(s.Cfg.BackupsDir(), 0o755); err != nil {
return nil, err
}
id := uuid.NewString()
ts := time.Now().Format("20060102-150405")
fp := filepath.Join(s.Cfg.BackupsDir(), "backup-"+ts+"-"+id[:8]+".tar.gz")
f, err := os.Create(fp)
if err != nil {
return nil, err
}
defer f.Close()
gz := gzip.NewWriter(f)
defer gz.Close()
tw := tar.NewWriter(gz)
defer tw.Close()
add := func(rel string) error {
abs := filepath.Join(s.Cfg.DataDir, rel)
return filepath.Walk(abs, func(path string, info os.FileInfo, err error) error {
if err != nil {
return nil
}
if info.IsDir() {
return nil
}
hdr, err := tar.FileInfoHeader(info, "")
if err != nil {
return nil
}
hdr.Name = filepath.ToSlash(filepath.Join(rel, strings.TrimPrefix(path, abs)))
if err := tw.WriteHeader(hdr); err != nil {
return nil
}
data, err := os.ReadFile(path)
if err != nil {
return nil
}
_, _ = tw.Write(data)
return nil
})
}
for _, sub := range []string{"pki", "instances", "clients"} {
_ = add(sub)
}
b := &model.Backup{
ID: id,
CreatedAt: time.Now(),
Filename: filepath.Base(fp),
Note: note,
Includes: []string{"pki", "instances", "clients"},
}
fi, _ := os.Stat(fp)
if fi != nil {
b.Size = fi.Size()
}
_ = s.Store.AddBackup(*b)
return b, nil
}
// Restore 从备份恢复。会覆盖现有数据。
func (s *Service) Restore(backupID string) error {
var bk *model.Backup
for _, b := range s.Store.ListBackups() {
if b.ID == backupID {
b := b
bk = &b
break
}
}
if bk == nil {
return fmt.Errorf("backup not found")
}
src := filepath.Join(s.Cfg.BackupsDir(), bk.Filename)
f, err := os.Open(src)
if err != nil {
return err
}
defer f.Close()
var r io.Reader = f
if strings.HasSuffix(src, ".gz") {
gz, err := gzip.NewReader(f)
if err != nil {
return err
}
defer gz.Close()
r = gz
}
tr := tar.NewReader(r)
for {
hdr, err := tr.Next()
if err == io.EOF {
break
}
if err != nil {
return err
}
target := filepath.Join(s.Cfg.DataDir, hdr.Name)
if hdr.FileInfo().IsDir() {
_ = os.MkdirAll(target, 0o755)
continue
}
_ = os.MkdirAll(filepath.Dir(target), 0o755)
out, err := os.OpenFile(target, os.O_CREATE|os.O_WRONLY|os.O_TRUNC, 0o600)
if err != nil {
return err
}
_, _ = io.Copy(out, tr)
_ = out.Close()
}
return nil
}
// DeleteBackup 删除备份文件与索引。
func (s *Service) DeleteBackup(id string) error {
for _, b := range s.Store.ListBackups() {
if b.ID == id {
_ = os.Remove(filepath.Join(s.Cfg.BackupsDir(), b.Filename))
return s.Store.DeleteBackup(id)
}
}
return fmt.Errorf("not found")
}
// RandomToken 生成短随机串。
func RandomToken(n int) string {
b := make([]byte, n)
_, _ = rand.Read(b)
return hex.EncodeToString(b)
}
// ---------- Admin 账号管理 ----------
// SeedDefaultAdminIfEmpty 当 db.json 中无任何 admin 账号时,用 env 里的默认账号创建。
// MustChangePassword=true,提示用户首次登录后必须改密。
func (s *Service) SeedDefaultAdminIfEmpty() error {
if s.Store.AdminCount() > 0 {
return nil
}
hash, err := bcrypt.GenerateFromPassword([]byte(s.Cfg.AdminPass), bcrypt.DefaultCost)
if err != nil {
return err
}
now := time.Now()
return s.Store.UpsertAdmin(model.AdminUser{
ID: uuid.NewString(),
Username: s.Cfg.AdminUser,
PasswordHash: string(hash),
Role: "admin",
Status: "active",
CreatedAt: now,
UpdatedAt: now,
MustChangePassword: true,
})
}
// CreateAdmin 新建管理员账号。密码由调用方提供,内部 bcrypt 哈希后存。
// 只允许 role=admin 的现有账号调用。
func (s *Service) CreateAdmin(username, password, role string) (*model.AdminUser, error) {
if len(username) < 3 {
return nil, fmt.Errorf("用户名至少 3 个字符")
}
if len(password) < 6 {
return nil, fmt.Errorf("密码至少 6 个字符")
}
if _, err := s.Store.GetAdminByUsername(username); err == nil {
return nil, fmt.Errorf("用户名已存在")
}
if role != "admin" && role != "operator" {
role = "operator"
}
hash, err := bcrypt.GenerateFromPassword([]byte(password), bcrypt.DefaultCost)
if err != nil {
return nil, err
}
now := time.Now()
a := model.AdminUser{
ID: uuid.NewString(),
Username: username,
PasswordHash: string(hash),
Role: role,
Status: "active",
CreatedAt: now,
UpdatedAt: now,
}
if err := s.Store.UpsertAdmin(a); err != nil {
return nil, err
}
return &a, nil
}
// ChangePassword 修改自身密码。校验旧密码后写新密码。
// actorID 是当前登录账号的 id。
func (s *Service) ChangePassword(actorID, oldPwd, newPwd string) error {
if len(newPwd) < 6 {
return fmt.Errorf("新密码至少 6 个字符")
}
a, err := s.Store.GetAdmin(actorID)
if err != nil {
return err
}
if err := bcrypt.CompareHashAndPassword([]byte(a.PasswordHash), []byte(oldPwd)); err != nil {
return fmt.Errorf("原密码错误")
}
hash, err := bcrypt.GenerateFromPassword([]byte(newPwd), bcrypt.DefaultCost)
if err != nil {
return err
}
a.PasswordHash = string(hash)
a.MustChangePassword = false
a.UpdatedAt = time.Now()
return s.Store.UpsertAdmin(*a)
}
// ResetPassword 由管理员重置他人密码(无需知道旧密码)。
func (s *Service) ResetPassword(targetID, newPwd string) error {
if len(newPwd) < 6 {
return fmt.Errorf("新密码至少 6 个字符")
}
a, err := s.Store.GetAdmin(targetID)
if err != nil {
return err
}
hash, err := bcrypt.GenerateFromPassword([]byte(newPwd), bcrypt.DefaultCost)
if err != nil {
return err
}
a.PasswordHash = string(hash)
a.MustChangePassword = false
a.UpdatedAt = time.Now()
return s.Store.UpsertAdmin(*a)
}
// SetAdminStatus 启用/禁用账号。
func (s *Service) SetAdminStatus(targetID, status string) error {
if status != "active" && status != "disabled" {
return fmt.Errorf("status 必须是 active 或 disabled")
}
a, err := s.Store.GetAdmin(targetID)
if err != nil {
return err
}
a.Status = status
a.UpdatedAt = time.Now()
return s.Store.UpsertAdmin(*a)
}
// DeleteAdmin 删除账号。保护:不能删除自己、不能删除最后一个 admin。
func (s *Service) DeleteAdmin(actorID, targetID string) error {
if actorID == targetID {
return fmt.Errorf("不能删除自己")
}
target, err := s.Store.GetAdmin(targetID)
if err != nil {
return err
}
if target.Role == "admin" {
// 统计其他 admin 数量
others := 0
for _, a := range s.Store.ListAdmins() {
if a.Role == "admin" && a.ID != targetID {
others++
}
}
if others == 0 {
return fmt.Errorf("不能删除最后一个 admin 账号")
}
}
return s.Store.DeleteAdmin(targetID)
}