Files
openvpn-manager/backend/internal/service/service.go
T
cnbugs 4c8b7b5188 Add per-instance/user access whitelist
Features:
- New Instance.AccessMode: "open" (default) or "whitelist"
- New Instance.AllowNetworks + VPNUser.AllowNetworks: list of CIDRs
- Effective whitelist = instance allow_networks ∪ user allow_networks (dedup)
- Auto-generates client-connect.sh / client-disconnect.sh for OpenVPN:
    * Reads ccd/<cn> to extract CIDRs
    * Pushes "route <ip> <mask>" to client (client side)
    * Inserts iptables ACCEPT rules in FORWARD chain (server side, defense in depth)
    * Cleans up rules on disconnect
- server.conf auto-includes client-connect / client-disconnect directives
  and push "redirect-gateway def1 bypass-dhcp" in whitelist mode
- ccd/<cn> file format: first line ifconfig-push (static IP), then one CIDR per line
- Editing instance allow_networks refreshes all users' ccd automatically
- New PUT /api/instances/:id/users/:uid endpoint
- CIDR format validation; reject malformed inputs with friendly errors
- Dashboard shows whitelist_instances count and per-instance allow_networks table

Docs:
- README: new section "三、访问控制(白名单模式)" with usage, validation, pitfalls
- docs/API.md: updated Instance / VPNUser model + create/update payloads
- Renumbered client usage section as 四
2026-08-09 20:47:21 +08:00

585 lines
15 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"
"os"
"os/exec"
"path/filepath"
"strings"
"time"
"github.com/gin-gonic/gin"
"github.com/google/uuid"
"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}
}
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
}
// 校验白名单 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 := 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)
}
}
u.ID = uuid.NewString()
u.Enabled = true
u.CreatedAt = time.Now()
// 签发证书(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
}
return &u, nil
}
// GenerateOVPN 下载/重新生成 .ovpnremoteHost 由前端传入。
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
}
return s.Ovm.GenerateClientOVPNFor(u, in, remoteHost)
}
// 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.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)
}
// 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)
}