first commit
This commit is contained in:
@@ -0,0 +1,73 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"server/models"
|
||||
"server/pkg/passwordutil"
|
||||
)
|
||||
|
||||
func NormalizeAccount(s string) string {
|
||||
return strings.TrimSpace(s)
|
||||
}
|
||||
|
||||
func CreateAdminUser(account, password string, name, phone, email, qq, avatar *string, sex uint8, roleID uint64, status uint8) (uint64, error) {
|
||||
hashed, err := passwordutil.Hash(password)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
u := &models.AdminUser{
|
||||
Account: NormalizeAccount(account),
|
||||
Password: hashed,
|
||||
Name: name,
|
||||
Phone: phone,
|
||||
Email: email,
|
||||
Qq: qq,
|
||||
Avatar: avatar,
|
||||
Sex: sex,
|
||||
RoleID: roleID,
|
||||
Status: status,
|
||||
}
|
||||
id, err := models.Orm.Insert(u)
|
||||
return uint64(id), err
|
||||
}
|
||||
|
||||
func GetAdminUserByID(id uint64) (*models.AdminUser, error) {
|
||||
u := &models.AdminUser{ID: id}
|
||||
if err := models.Orm.Read(u); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return u, nil
|
||||
}
|
||||
|
||||
func UpdateAdminUser(id uint64, fields map[string]interface{}) error {
|
||||
_, err := models.Orm.QueryTable(new(models.AdminUser)).Filter("id", id).Update(fields)
|
||||
return err
|
||||
}
|
||||
|
||||
func DeleteAdminUser(id uint64) error {
|
||||
_, err := models.Orm.QueryTable(new(models.AdminUser)).Filter("id", id).Delete()
|
||||
return err
|
||||
}
|
||||
|
||||
func ChangeAdminUserPassword(id uint64, newPassword string) error {
|
||||
hashed, err := passwordutil.Hash(newPassword)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
_, err = models.Orm.QueryTable(new(models.AdminUser)).Filter("id", id).Update(map[string]interface{}{
|
||||
"password": hashed,
|
||||
})
|
||||
return err
|
||||
}
|
||||
|
||||
func ListAdminUsers() ([]models.AdminUser, int64, error) {
|
||||
var rows []models.AdminUser
|
||||
total, err := models.Orm.QueryTable(new(models.AdminUser)).Count()
|
||||
if err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
_, err = models.Orm.QueryTable(new(models.AdminUser)).OrderBy("-id").All(&rows)
|
||||
return rows, total, err
|
||||
}
|
||||
|
||||
@@ -0,0 +1,246 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"math/rand"
|
||||
"net/http"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"server/models"
|
||||
)
|
||||
|
||||
type loginCodeItem struct {
|
||||
Code string
|
||||
Channel string
|
||||
ExpiredAt time.Time
|
||||
}
|
||||
|
||||
var loginCodeStore sync.Map
|
||||
|
||||
func codeKey(account, channel string) string {
|
||||
return strings.ToLower(strings.TrimSpace(account)) + "|" + strings.TrimSpace(channel)
|
||||
}
|
||||
|
||||
func SendPlatformLoginCode(account, channel string) error {
|
||||
account = strings.TrimSpace(account)
|
||||
channel = strings.TrimSpace(channel)
|
||||
if account == "" {
|
||||
return errors.New("账号不能为空")
|
||||
}
|
||||
if channel != "sms" && channel != "email" {
|
||||
return errors.New("仅支持短信或邮箱验证码")
|
||||
}
|
||||
|
||||
var u models.AdminUser
|
||||
if err := models.Orm.QueryTable(new(models.AdminUser)).Filter("account", account).One(&u); err != nil {
|
||||
return errors.New("用户不存在")
|
||||
}
|
||||
if u.Status == 0 {
|
||||
return errors.New("账号已禁用")
|
||||
}
|
||||
if channel == "sms" && (u.Phone == nil || strings.TrimSpace(*u.Phone) == "") {
|
||||
return errors.New("该账号未绑定手机号")
|
||||
}
|
||||
if channel == "email" && (u.Email == nil || strings.TrimSpace(*u.Email) == "") {
|
||||
return errors.New("该账号未绑定邮箱")
|
||||
}
|
||||
|
||||
rand.Seed(time.Now().UnixNano())
|
||||
code := fmt.Sprintf("%06d", rand.Intn(1000000))
|
||||
loginCodeStore.Store(codeKey(account, channel), loginCodeItem{
|
||||
Code: code,
|
||||
Channel: channel,
|
||||
ExpiredAt: time.Now().Add(5 * time.Minute),
|
||||
})
|
||||
return nil
|
||||
}
|
||||
|
||||
func VerifyPlatformLoginCode(account, channel, code string) error {
|
||||
account = strings.TrimSpace(account)
|
||||
channel = strings.TrimSpace(channel)
|
||||
code = strings.TrimSpace(code)
|
||||
if account == "" || code == "" {
|
||||
return errors.New("验证码不能为空")
|
||||
}
|
||||
val, ok := loginCodeStore.Load(codeKey(account, channel))
|
||||
if !ok {
|
||||
return errors.New("验证码不存在或已失效")
|
||||
}
|
||||
item, ok := val.(loginCodeItem)
|
||||
if !ok {
|
||||
return errors.New("验证码状态异常")
|
||||
}
|
||||
if time.Now().After(item.ExpiredAt) {
|
||||
loginCodeStore.Delete(codeKey(account, channel))
|
||||
return errors.New("验证码已过期")
|
||||
}
|
||||
if item.Code != code {
|
||||
return errors.New("验证码错误")
|
||||
}
|
||||
loginCodeStore.Delete(codeKey(account, channel))
|
||||
return nil
|
||||
}
|
||||
|
||||
func SendBackendLoginCode(tenantName, account, channel string) error {
|
||||
tenantName = strings.TrimSpace(tenantName)
|
||||
account = strings.TrimSpace(account)
|
||||
channel = strings.TrimSpace(channel)
|
||||
if tenantName == "" || account == "" {
|
||||
return errors.New("租户名称和账号不能为空")
|
||||
}
|
||||
if channel != "sms" && channel != "email" {
|
||||
return errors.New("仅支持短信或邮箱验证码")
|
||||
}
|
||||
|
||||
var tenant models.SystemTenant
|
||||
if err := models.Orm.QueryTable(new(models.SystemTenant)).Filter("tenant_name", tenantName).One(&tenant); err != nil {
|
||||
return errors.New("租户不存在")
|
||||
}
|
||||
|
||||
rand.Seed(time.Now().UnixNano())
|
||||
code := fmt.Sprintf("%06d", rand.Intn(1000000))
|
||||
|
||||
switch channel {
|
||||
case "sms":
|
||||
phone := account
|
||||
var user models.SystemTenantUser
|
||||
if err := models.Orm.QueryTable(new(models.SystemTenantUser)).
|
||||
Filter("tid", tenant.ID).
|
||||
Filter("phone", phone).
|
||||
One(&user); err != nil {
|
||||
return errors.New("该手机号非当前企业绑定号码,请重试")
|
||||
}
|
||||
if user.Status == 0 {
|
||||
return errors.New("账号已禁用")
|
||||
}
|
||||
if user.Phone == nil || strings.TrimSpace(*user.Phone) == "" {
|
||||
return errors.New("该手机号非当前企业绑定号码,请重试")
|
||||
}
|
||||
|
||||
content := "短信验证码:" + code
|
||||
if err := enqueueSMSTaskForLogin(tenant.ID, phone, content, code); err != nil {
|
||||
return errors.New("短信发送失败,请重试")
|
||||
}
|
||||
case "email":
|
||||
email := account
|
||||
var user models.SystemTenantUser
|
||||
if err := models.Orm.QueryTable(new(models.SystemTenantUser)).
|
||||
Filter("tid", tenant.ID).
|
||||
Filter("email", email).
|
||||
One(&user); err != nil {
|
||||
return errors.New("该账号未绑定邮箱")
|
||||
}
|
||||
if user.Status == 0 {
|
||||
return errors.New("账号已禁用")
|
||||
}
|
||||
if user.Email == nil || strings.TrimSpace(*user.Email) == "" {
|
||||
return errors.New("该账号未绑定邮箱")
|
||||
}
|
||||
}
|
||||
|
||||
loginCodeStore.Store(codeKey(tenantName+"#"+account, channel), loginCodeItem{
|
||||
Code: code,
|
||||
Channel: channel,
|
||||
ExpiredAt: time.Now().Add(5 * time.Minute),
|
||||
})
|
||||
return nil
|
||||
}
|
||||
|
||||
func VerifyBackendLoginCode(tenantName, account, channel, code string) error {
|
||||
return VerifyPlatformLoginCode(tenantName+"#"+account, channel, code)
|
||||
}
|
||||
|
||||
func getDefaultSystemSMSConfig() (backendURL string, apiKey string, err error) {
|
||||
var row models.SystemSMS
|
||||
err = models.Orm.QueryTable(new(models.SystemSMS)).
|
||||
Filter("is_default", 1).
|
||||
Filter("status", 1).
|
||||
OrderBy("-weight", "-id").
|
||||
Limit(1).
|
||||
One(&row)
|
||||
if err != nil {
|
||||
err2 := models.Orm.QueryTable(new(models.SystemSMS)).
|
||||
Filter("config_code", "custom").
|
||||
OrderBy("-id").
|
||||
Limit(1).
|
||||
One(&row)
|
||||
if err2 != nil {
|
||||
return "", "", err2
|
||||
}
|
||||
}
|
||||
backendURL = strings.TrimSpace(row.ApiURL)
|
||||
apiKey = strings.TrimSpace(row.ApiKey)
|
||||
return backendURL, apiKey, nil
|
||||
}
|
||||
|
||||
// enqueueSMSTaskForLogin 入队短信任务到网关,并写入 yz_system_sms_tasks
|
||||
func enqueueSMSTaskForLogin(tid uint64, phone, content, code string) error {
|
||||
backendURL, apiKey, err := getDefaultSystemSMSConfig()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if backendURL == "" || apiKey == "" {
|
||||
return errors.New("短信网关未配置")
|
||||
}
|
||||
|
||||
enqueueURL := strings.TrimRight(backendURL, "/") + "/api/v1/business/outbound-tasks"
|
||||
payload := map[string]interface{}{
|
||||
"phone": phone,
|
||||
"content": content,
|
||||
}
|
||||
bs, _ := json.Marshal(payload)
|
||||
|
||||
client := &http.Client{Timeout: 10 * time.Second}
|
||||
req, err := http.NewRequest("POST", enqueueURL, bytes.NewReader(bs))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
req.Header.Set("X-Api-Key", apiKey)
|
||||
req.Header.Set("Content-Type", "application/json; charset=utf-8")
|
||||
req.Header.Set("Accept", "application/json")
|
||||
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
bodyBytes, _ := io.ReadAll(resp.Body)
|
||||
bodyStr := strings.TrimSpace(string(bodyBytes))
|
||||
|
||||
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||
return fmt.Errorf("gateway http status: %d, body: %s", resp.StatusCode, bodyStr)
|
||||
}
|
||||
|
||||
now := time.Now()
|
||||
tidCopy := tid
|
||||
contentPtr := content
|
||||
var reportPtr *string
|
||||
if bodyStr != "" {
|
||||
reportPtr = &bodyStr
|
||||
}
|
||||
|
||||
task := &models.SystemSMSTask{
|
||||
Tid: &tidCopy,
|
||||
ApiKey: apiKey,
|
||||
Phone: phone,
|
||||
Content: &contentPtr,
|
||||
Status: 3,
|
||||
Code: code,
|
||||
ReportRaw: reportPtr,
|
||||
CreateTime: &now,
|
||||
UpdateTime: &now,
|
||||
}
|
||||
|
||||
_, insertErr := models.Orm.Insert(task)
|
||||
if insertErr != nil {
|
||||
return nil
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,51 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"strings"
|
||||
|
||||
"server/models"
|
||||
)
|
||||
|
||||
// CheckUserPermission 校验用户是否拥有指定权限标识。
|
||||
// 兼容 rights 为 JSON 数组 / 逗号分隔字符串;解析失败时默认放行,避免历史数据阻断请求。
|
||||
func CheckUserPermission(userID int, permission string) (bool, error) {
|
||||
if permission == "" || userID <= 0 {
|
||||
return true, nil
|
||||
}
|
||||
|
||||
var user models.AdminUser
|
||||
if err := models.Orm.QueryTable(new(models.AdminUser)).Filter("id", userID).One(&user); err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
var role models.AdminRole
|
||||
if err := models.Orm.QueryTable(new(models.AdminRole)).Filter("id", user.RoleID).One(&role); err != nil {
|
||||
return false, err
|
||||
}
|
||||
if role.Rights == nil || strings.TrimSpace(*role.Rights) == "" {
|
||||
return true, nil
|
||||
}
|
||||
rightsRaw := strings.TrimSpace(*role.Rights)
|
||||
|
||||
// 1) JSON 数组格式
|
||||
var arr []string
|
||||
if err := json.Unmarshal([]byte(rightsRaw), &arr); err == nil {
|
||||
for _, p := range arr {
|
||||
if strings.TrimSpace(p) == permission {
|
||||
return true, nil
|
||||
}
|
||||
}
|
||||
return false, nil
|
||||
}
|
||||
|
||||
// 2) 逗号分隔字符串
|
||||
for _, p := range strings.Split(rightsRaw, ",") {
|
||||
if strings.TrimSpace(p) == permission {
|
||||
return true, nil
|
||||
}
|
||||
}
|
||||
|
||||
return false, nil
|
||||
}
|
||||
|
||||
@@ -0,0 +1,158 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"strings"
|
||||
|
||||
"server/models"
|
||||
"server/pkg/jwtutil"
|
||||
"server/pkg/passwordutil"
|
||||
)
|
||||
|
||||
type PlatformLoginUser struct {
|
||||
ID uint64
|
||||
Account string
|
||||
Name string
|
||||
Tid uint64
|
||||
Rid uint64
|
||||
Avatar string
|
||||
RoleName string
|
||||
}
|
||||
|
||||
func adminRoleNameByID(roleID uint64) string {
|
||||
if roleID == 0 {
|
||||
return ""
|
||||
}
|
||||
var role models.AdminRole
|
||||
err := models.Orm.QueryTable(new(models.AdminRole)).Filter("id", roleID).One(&role)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return role.Name
|
||||
}
|
||||
|
||||
func toPlatformLoginUser(user *models.AdminUser) *PlatformLoginUser {
|
||||
name := ""
|
||||
if user.Name != nil {
|
||||
name = *user.Name
|
||||
}
|
||||
avatar := ""
|
||||
if user.Avatar != nil {
|
||||
avatar = *user.Avatar
|
||||
}
|
||||
return &PlatformLoginUser{
|
||||
ID: user.ID,
|
||||
Account: user.Account,
|
||||
Name: name,
|
||||
Tid: 0,
|
||||
Rid: user.RoleID,
|
||||
Avatar: avatar,
|
||||
RoleName: adminRoleNameByID(user.RoleID),
|
||||
}
|
||||
}
|
||||
|
||||
// PlatformAdminLogin 平台端登录:仅校验 yz_system_admin_user,不需要租户。
|
||||
func PlatformAdminLogin(account, password string) (string, *PlatformLoginUser, error) {
|
||||
account = strings.TrimSpace(account)
|
||||
password = strings.TrimSpace(password)
|
||||
if account == "" || password == "" {
|
||||
return "", nil, errors.New("用户名或密码不能为空")
|
||||
}
|
||||
|
||||
var user models.AdminUser
|
||||
err := models.Orm.QueryTable(new(models.AdminUser)).
|
||||
Filter("account", account).
|
||||
One(&user)
|
||||
if err != nil {
|
||||
return "", nil, errors.New("用户名或密码错误")
|
||||
}
|
||||
if user.Status == 0 {
|
||||
return "", nil, errors.New("账号已禁用")
|
||||
}
|
||||
if !passwordutil.Verify(user.Password, password) {
|
||||
return "", nil, errors.New("用户名或密码错误")
|
||||
}
|
||||
|
||||
const tenantID = 0
|
||||
const userType = "platform"
|
||||
token, err := jwtutil.GenerateToken(int(user.ID), user.Account, tenantID, userType)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
|
||||
loginUser := toPlatformLoginUser(&user)
|
||||
return token, loginUser, nil
|
||||
}
|
||||
|
||||
// BackendLogin backend 登录:先校验租户,再校验租户下用户账号和密码。
|
||||
func BackendLogin(tenantName, account, password string) (string, *PlatformLoginUser, error) {
|
||||
tenantName = strings.TrimSpace(tenantName)
|
||||
account = strings.TrimSpace(account)
|
||||
password = strings.TrimSpace(password)
|
||||
if tenantName == "" || account == "" || password == "" {
|
||||
return "", nil, errors.New("租户名称、用户名或密码不能为空")
|
||||
}
|
||||
|
||||
var tenant models.SystemTenant
|
||||
err := models.Orm.QueryTable(new(models.SystemTenant)).
|
||||
Filter("tenant_name", tenantName).
|
||||
One(&tenant)
|
||||
if err != nil {
|
||||
return "", nil, errors.New("租户不存在")
|
||||
}
|
||||
if tenant.Status != 1 {
|
||||
return "", nil, errors.New("租户已停用")
|
||||
}
|
||||
|
||||
var tenantUser models.SystemTenantUser
|
||||
err = models.Orm.QueryTable(new(models.SystemTenantUser)).
|
||||
Filter("tid", tenant.ID).
|
||||
Filter("account", account).
|
||||
One(&tenantUser)
|
||||
if err != nil {
|
||||
return "", nil, errors.New("用户名或密码错误")
|
||||
}
|
||||
if tenantUser.Status == 0 {
|
||||
return "", nil, errors.New("账号已禁用")
|
||||
}
|
||||
if tenantUser.Password == nil || !passwordutil.Verify(*tenantUser.Password, password) {
|
||||
return "", nil, errors.New("用户名或密码错误")
|
||||
}
|
||||
|
||||
tenantID := int(tenant.ID)
|
||||
const userType = "backend"
|
||||
token, err := jwtutil.GenerateToken(int(tenantUser.Uid), account, tenantID, userType)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
|
||||
loginUser := &PlatformLoginUser{
|
||||
ID: tenantUser.Uid,
|
||||
Account: account,
|
||||
Name: "",
|
||||
Tid: tenant.ID,
|
||||
Rid: 0,
|
||||
Avatar: "",
|
||||
RoleName: "",
|
||||
}
|
||||
if tenantUser.Account != nil && strings.TrimSpace(*tenantUser.Account) != "" {
|
||||
loginUser.Account = strings.TrimSpace(*tenantUser.Account)
|
||||
}
|
||||
if tenantUser.Name != nil {
|
||||
loginUser.Name = strings.TrimSpace(*tenantUser.Name)
|
||||
}
|
||||
|
||||
return token, loginUser, nil
|
||||
}
|
||||
|
||||
// PlatformGetCurrentUser 根据平台管理员用户 ID 返回登录用户信息(含角色名称)。
|
||||
func PlatformGetCurrentUser(uid uint64) (*PlatformLoginUser, error) {
|
||||
u, err := GetAdminUserByID(uid)
|
||||
if err != nil {
|
||||
return nil, errors.New("用户不存在")
|
||||
}
|
||||
if u.Status == 0 {
|
||||
return nil, errors.New("账号已禁用")
|
||||
}
|
||||
return toPlatformLoginUser(u), nil
|
||||
}
|
||||
@@ -0,0 +1,60 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"server/models"
|
||||
|
||||
beego "github.com/beego/beego/v2/server/web"
|
||||
)
|
||||
|
||||
// PublicRequestBaseURL 根据请求拼出对外访问的根(用于把 /uploads/... 拼成完整下载地址)
|
||||
func PublicRequestBaseURL(c *beego.Controller) (scheme, host string) {
|
||||
scheme = "http"
|
||||
if c.Ctx.Request.TLS != nil {
|
||||
scheme = "https"
|
||||
}
|
||||
if strings.EqualFold(strings.TrimSpace(c.Ctx.Request.Header.Get("X-Forwarded-Proto")), "https") {
|
||||
scheme = "https"
|
||||
}
|
||||
host = strings.TrimSpace(c.Ctx.Request.Host)
|
||||
return scheme, host
|
||||
}
|
||||
|
||||
// ResolveSoftwareDownloadURL 优先使用自定义 download_url;否则根据 file_id 读附件 src 拼完整 URL
|
||||
func ResolveSoftwareDownloadURL(scheme, host string, downloadURL *string, fileID *uint64) string {
|
||||
if downloadURL != nil {
|
||||
u := strings.TrimSpace(*downloadURL)
|
||||
if u != "" {
|
||||
return u
|
||||
}
|
||||
}
|
||||
if fileID == nil || *fileID == 0 {
|
||||
return ""
|
||||
}
|
||||
var f models.SystemFile
|
||||
err := models.Orm.QueryTable(new(models.SystemFile)).
|
||||
Filter("id", *fileID).
|
||||
Filter("delete_time__isnull", true).
|
||||
One(&f)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
src := strings.TrimSpace(f.Src)
|
||||
if src == "" {
|
||||
return ""
|
||||
}
|
||||
if strings.HasPrefix(strings.ToLower(src), "http://") || strings.HasPrefix(strings.ToLower(src), "https://") {
|
||||
return src
|
||||
}
|
||||
if host == "" {
|
||||
return src
|
||||
}
|
||||
if !strings.HasPrefix(src, "/") {
|
||||
src = "/" + src
|
||||
}
|
||||
if scheme == "" {
|
||||
scheme = "http"
|
||||
}
|
||||
return scheme + "://" + host + src
|
||||
}
|
||||
@@ -0,0 +1,191 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"mime/multipart"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"server/models"
|
||||
)
|
||||
|
||||
// MigrationProgress 迁移进度
|
||||
type MigrationProgress struct {
|
||||
Total int
|
||||
Success int
|
||||
Failed int
|
||||
Current string
|
||||
Errors []string
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
// AddSuccess 增加成功计数
|
||||
func (p *MigrationProgress) AddSuccess() {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
p.Success++
|
||||
}
|
||||
|
||||
// AddFailed 增加失败计数
|
||||
func (p *MigrationProgress) AddFailed(err string) {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
p.Failed++
|
||||
p.Errors = append(p.Errors, err)
|
||||
}
|
||||
|
||||
// SetCurrent 设置当前处理的文件
|
||||
func (p *MigrationProgress) SetCurrent(filename string) {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
p.Current = filename
|
||||
}
|
||||
|
||||
// GetProgress 获取进度信息
|
||||
func (p *MigrationProgress) GetProgress() (int, int, int, string) {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
return p.Total, p.Success, p.Failed, p.Current
|
||||
}
|
||||
|
||||
// StorageMigration 存储迁移服务
|
||||
type StorageMigration struct {
|
||||
fromService StorageService
|
||||
toService StorageService
|
||||
progress *MigrationProgress
|
||||
}
|
||||
|
||||
// NewStorageMigration 创建存储迁移服务
|
||||
func NewStorageMigration(from, to StorageService) *StorageMigration {
|
||||
return &StorageMigration{
|
||||
fromService: from,
|
||||
toService: to,
|
||||
progress: &MigrationProgress{
|
||||
Errors: make([]string, 0),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
// MigrateFile 迁移单个文件
|
||||
func (m *StorageMigration) MigrateFile(file *models.SystemFile) error {
|
||||
m.progress.SetCurrent(file.Name)
|
||||
|
||||
// 如果是本地存储,从本地读取文件
|
||||
if localFrom, ok := m.fromService.(*LocalStorage); ok {
|
||||
// 从本地文件系统读取
|
||||
localPath := strings.TrimPrefix(file.Src, "/")
|
||||
filePath := filepath.Join(localFrom.BaseDir, localPath)
|
||||
|
||||
f, err := os.Open(filePath)
|
||||
if err != nil {
|
||||
return fmt.Errorf("打开本地文件失败: %w", err)
|
||||
}
|
||||
defer f.Close()
|
||||
|
||||
// 获取文件信息
|
||||
stat, err := f.Stat()
|
||||
if err != nil {
|
||||
return fmt.Errorf("获取文件信息失败: %w", err)
|
||||
}
|
||||
|
||||
// 创建 multipart.FileHeader
|
||||
header := &multipart.FileHeader{
|
||||
Filename: file.Name,
|
||||
Size: stat.Size(),
|
||||
}
|
||||
|
||||
// 上传到目标存储
|
||||
result, err := m.toService.Upload(f, header)
|
||||
if err != nil {
|
||||
return fmt.Errorf("上传到目标存储失败: %w", err)
|
||||
}
|
||||
|
||||
// 更新数据库记录
|
||||
_, err = models.Orm.QueryTable(new(models.SystemFile)).
|
||||
Filter("id", file.ID).
|
||||
Update(map[string]interface{}{
|
||||
"src": result.URL,
|
||||
})
|
||||
if err != nil {
|
||||
// 上传成功但更新数据库失败,尝试删除已上传的文件
|
||||
_ = m.toService.Delete(result.Key)
|
||||
return fmt.Errorf("更新数据库失败: %w", err)
|
||||
}
|
||||
|
||||
m.progress.AddSuccess()
|
||||
return nil
|
||||
}
|
||||
|
||||
// 如果是七牛云存储,需要先下载再上传(这里简化处理)
|
||||
return fmt.Errorf("暂不支持从七牛云迁移到本地")
|
||||
}
|
||||
|
||||
// MigrateAll 迁移所有文件
|
||||
func (m *StorageMigration) MigrateAll(tid uint64) error {
|
||||
// 获取所有文件
|
||||
var files []models.SystemFile
|
||||
_, err := models.Orm.QueryTable(new(models.SystemFile)).
|
||||
Filter("tid", tid).
|
||||
Filter("delete_time__isnull", true).
|
||||
All(&files)
|
||||
if err != nil {
|
||||
return fmt.Errorf("获取文件列表失败: %w", err)
|
||||
}
|
||||
|
||||
m.progress.Total = len(files)
|
||||
|
||||
// 并发迁移(限制并发数)
|
||||
concurrency := 5
|
||||
sem := make(chan struct{}, concurrency)
|
||||
var wg sync.WaitGroup
|
||||
|
||||
for i := range files {
|
||||
wg.Add(1)
|
||||
go func(file *models.SystemFile) {
|
||||
defer wg.Done()
|
||||
sem <- struct{}{} // 获取信号量
|
||||
defer func() { <-sem }() // 释放信号量
|
||||
|
||||
if err := m.MigrateFile(file); err != nil {
|
||||
m.progress.AddFailed(fmt.Sprintf("%s: %v", file.Name, err))
|
||||
}
|
||||
}(&files[i])
|
||||
}
|
||||
|
||||
wg.Wait()
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetProgress 获取迁移进度
|
||||
func (m *StorageMigration) GetProgress() *MigrationProgress {
|
||||
return m.progress
|
||||
}
|
||||
|
||||
// MigrateLocalToQiniu 从本地存储迁移到七牛云
|
||||
func MigrateLocalToQiniu(tid uint64) (*MigrationProgress, error) {
|
||||
// 获取存储配置
|
||||
cfg, err := models.GetStorageConfig()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("获取存储配置失败: %w", err)
|
||||
}
|
||||
|
||||
if cfg.StorageType != "qiniu" {
|
||||
return nil, fmt.Errorf("当前存储类型不是七牛云")
|
||||
}
|
||||
|
||||
// 创建存储服务
|
||||
localStorage := NewLocalStorage()
|
||||
qiniuStorage := NewQiniuStorage(cfg)
|
||||
|
||||
// 创建迁移服务
|
||||
migration := NewStorageMigration(localStorage, qiniuStorage)
|
||||
|
||||
// 执行迁移
|
||||
if err := migration.MigrateAll(tid); err != nil {
|
||||
return migration.GetProgress(), err
|
||||
}
|
||||
|
||||
return migration.GetProgress(), nil
|
||||
}
|
||||
@@ -0,0 +1,252 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/md5"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"io"
|
||||
"mime/multipart"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"server/models"
|
||||
|
||||
"github.com/qiniu/go-sdk/v7/auth/qbox"
|
||||
"github.com/qiniu/go-sdk/v7/storage"
|
||||
)
|
||||
|
||||
// StorageService 存储服务接口
|
||||
type StorageService interface {
|
||||
Upload(file multipart.File, header *multipart.FileHeader) (*UploadResult, error)
|
||||
GetPublicURL(key string) string
|
||||
Delete(key string) error
|
||||
}
|
||||
|
||||
// UploadResult 上传结果
|
||||
type UploadResult struct {
|
||||
URL string // 完整访问URL
|
||||
Key string // 存储key/路径
|
||||
Size int64 // 文件大小
|
||||
MD5 string // 文件MD5
|
||||
MimeType string // 文件类型
|
||||
}
|
||||
|
||||
// LocalStorage 本地存储实现
|
||||
type LocalStorage struct {
|
||||
BaseDir string // 基础目录,默认 "uploads"
|
||||
BaseURL string // 基础URL,默认 "/"
|
||||
}
|
||||
|
||||
// NewLocalStorage 创建本地存储服务
|
||||
func NewLocalStorage() *LocalStorage {
|
||||
return &LocalStorage{
|
||||
BaseDir: "uploads",
|
||||
BaseURL: "/",
|
||||
}
|
||||
}
|
||||
|
||||
// Upload 上传文件到本地
|
||||
func (s *LocalStorage) Upload(file multipart.File, header *multipart.FileHeader) (*UploadResult, error) {
|
||||
// 生成存储路径
|
||||
ext := filepath.Ext(header.Filename)
|
||||
datePath := time.Now().Format("2006/01/02")
|
||||
fileName := fmt.Sprintf("%d%s", time.Now().UnixNano(), ext)
|
||||
savePath := filepath.Join(datePath, fileName)
|
||||
|
||||
// 创建目录
|
||||
destDir := filepath.Join(s.BaseDir, filepath.FromSlash(datePath))
|
||||
if err := os.MkdirAll(destDir, 0755); err != nil {
|
||||
return nil, fmt.Errorf("创建目录失败: %w", err)
|
||||
}
|
||||
|
||||
// 保存文件
|
||||
destPath := filepath.Join(s.BaseDir, filepath.FromSlash(savePath))
|
||||
dst, err := os.Create(destPath)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("创建文件失败: %w", err)
|
||||
}
|
||||
defer dst.Close()
|
||||
|
||||
// 计算MD5并复制文件
|
||||
hash := md5.New()
|
||||
size, err := io.Copy(io.MultiWriter(dst, hash), file)
|
||||
if err != nil {
|
||||
_ = os.Remove(destPath)
|
||||
return nil, fmt.Errorf("保存文件失败: %w", err)
|
||||
}
|
||||
|
||||
md5Sum := hex.EncodeToString(hash.Sum(nil))
|
||||
webURL := s.BaseURL + strings.ReplaceAll(filepath.ToSlash(destPath), "\\", "/")
|
||||
|
||||
return &UploadResult{
|
||||
URL: webURL,
|
||||
Key: savePath,
|
||||
Size: size,
|
||||
MD5: md5Sum,
|
||||
MimeType: header.Header.Get("Content-Type"),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// GetPublicURL 获取公开访问URL
|
||||
func (s *LocalStorage) GetPublicURL(key string) string {
|
||||
return s.BaseURL + filepath.ToSlash(filepath.Join(s.BaseDir, key))
|
||||
}
|
||||
|
||||
// Delete 删除本地文件
|
||||
func (s *LocalStorage) Delete(key string) error {
|
||||
filePath := filepath.Join(s.BaseDir, filepath.FromSlash(key))
|
||||
return os.Remove(filePath)
|
||||
}
|
||||
|
||||
// QiniuStorage 七牛云存储实现
|
||||
type QiniuStorage struct {
|
||||
AccessKey string
|
||||
SecretKey string
|
||||
Bucket string
|
||||
Domain string
|
||||
Region string
|
||||
}
|
||||
|
||||
// NewQiniuStorage 创建七牛云存储服务
|
||||
func NewQiniuStorage(cfg *models.StorageConfig) *QiniuStorage {
|
||||
return &QiniuStorage{
|
||||
AccessKey: cfg.QiniuAccessKey,
|
||||
SecretKey: cfg.QiniuSecretKey,
|
||||
Bucket: cfg.QiniuBucket,
|
||||
Domain: cfg.QiniuDomain,
|
||||
Region: cfg.QiniuRegion,
|
||||
}
|
||||
}
|
||||
|
||||
// getZone 根据区域代码获取存储区域
|
||||
func (s *QiniuStorage) getZone() *storage.Region {
|
||||
switch s.Region {
|
||||
case "z0":
|
||||
return &storage.ZoneHuadong
|
||||
case "z1":
|
||||
return &storage.ZoneHuabei
|
||||
case "z2":
|
||||
return &storage.ZoneHuanan
|
||||
case "na0":
|
||||
return &storage.ZoneBeimei
|
||||
case "as0":
|
||||
return &storage.ZoneXinjiapo
|
||||
case "cn-east-2":
|
||||
return &storage.ZoneHuadongZheJiang2
|
||||
default:
|
||||
return &storage.ZoneHuadong // 默认华东
|
||||
}
|
||||
}
|
||||
|
||||
// Upload 上传文件到七牛云
|
||||
func (s *QiniuStorage) Upload(file multipart.File, header *multipart.FileHeader) (*UploadResult, error) {
|
||||
// 生成存储key
|
||||
ext := filepath.Ext(header.Filename)
|
||||
datePath := time.Now().Format("2006/01/02")
|
||||
fileName := fmt.Sprintf("%d%s", time.Now().UnixNano(), ext)
|
||||
key := filepath.ToSlash(filepath.Join(datePath, fileName))
|
||||
|
||||
// 创建上传凭证
|
||||
mac := qbox.NewMac(s.AccessKey, s.SecretKey)
|
||||
putPolicy := storage.PutPolicy{
|
||||
Scope: s.Bucket,
|
||||
}
|
||||
upToken := putPolicy.UploadToken(mac)
|
||||
|
||||
// 配置上传参数
|
||||
cfg := storage.Config{
|
||||
Region: s.getZone(),
|
||||
UseHTTPS: true,
|
||||
UseCdnDomains: false,
|
||||
}
|
||||
|
||||
// 创建表单上传器
|
||||
formUploader := storage.NewFormUploader(&cfg)
|
||||
ret := storage.PutRet{}
|
||||
putExtra := storage.PutExtra{}
|
||||
|
||||
// 计算文件大小和MD5
|
||||
tmpFile, err := os.CreateTemp("", "qiniu_upload_*")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("创建临时文件失败: %w", err)
|
||||
}
|
||||
defer os.Remove(tmpFile.Name())
|
||||
defer tmpFile.Close()
|
||||
|
||||
hash := md5.New()
|
||||
size, err := io.Copy(io.MultiWriter(tmpFile, hash), file)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("读取文件失败: %w", err)
|
||||
}
|
||||
md5Sum := hex.EncodeToString(hash.Sum(nil))
|
||||
|
||||
// 重置文件指针
|
||||
if _, err := tmpFile.Seek(0, 0); err != nil {
|
||||
return nil, fmt.Errorf("重置文件指针失败: %w", err)
|
||||
}
|
||||
|
||||
// 执行上传
|
||||
err = formUploader.Put(context.Background(), &ret, upToken, key, tmpFile, size, &putExtra)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("上传到七牛云失败: %w", err)
|
||||
}
|
||||
|
||||
// 构建完整URL
|
||||
domain := strings.TrimRight(s.Domain, "/")
|
||||
url := fmt.Sprintf("%s/%s", domain, ret.Key)
|
||||
|
||||
return &UploadResult{
|
||||
URL: url,
|
||||
Key: ret.Key,
|
||||
Size: size,
|
||||
MD5: md5Sum,
|
||||
MimeType: header.Header.Get("Content-Type"),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// GetPublicURL 获取七牛云公开访问URL
|
||||
func (s *QiniuStorage) GetPublicURL(key string) string {
|
||||
domain := strings.TrimRight(s.Domain, "/")
|
||||
return fmt.Sprintf("%s/%s", domain, key)
|
||||
}
|
||||
|
||||
// Delete 删除七牛云文件
|
||||
func (s *QiniuStorage) Delete(key string) error {
|
||||
mac := qbox.NewMac(s.AccessKey, s.SecretKey)
|
||||
cfg := storage.Config{
|
||||
Region: s.getZone(),
|
||||
UseHTTPS: true,
|
||||
}
|
||||
|
||||
bucketManager := storage.NewBucketManager(mac, &cfg)
|
||||
err := bucketManager.Delete(s.Bucket, key)
|
||||
if err != nil {
|
||||
return fmt.Errorf("删除七牛云文件失败: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetStorageService 根据配置获取存储服务
|
||||
func GetStorageService() (StorageService, error) {
|
||||
cfg, err := models.GetStorageConfig()
|
||||
if err != nil {
|
||||
// 默认使用本地存储
|
||||
return NewLocalStorage(), nil
|
||||
}
|
||||
|
||||
switch cfg.StorageType {
|
||||
case "qiniu":
|
||||
if cfg.QiniuAccessKey == "" || cfg.QiniuSecretKey == "" ||
|
||||
cfg.QiniuBucket == "" || cfg.QiniuDomain == "" {
|
||||
return nil, fmt.Errorf("七牛云配置不完整")
|
||||
}
|
||||
return NewQiniuStorage(cfg), nil
|
||||
case "local":
|
||||
return NewLocalStorage(), nil
|
||||
default:
|
||||
return NewLocalStorage(), nil
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,126 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"crypto/tls"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/smtp"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// SMTPConfig 发送邮件所需参数(与 yz_system_email 字段对应)
|
||||
type SMTPConfig struct {
|
||||
FromAddress string
|
||||
FromName string
|
||||
Host string
|
||||
Port uint
|
||||
Password string
|
||||
Encryption string // ssl / tls / none
|
||||
Timeout uint // 秒
|
||||
}
|
||||
|
||||
// SendTestEmailSMTP 发送一封简单测试邮件(纯文本 UTF-8)
|
||||
func SendTestEmailSMTP(cfg SMTPConfig, to string) error {
|
||||
to = strings.TrimSpace(to)
|
||||
if to == "" {
|
||||
return fmt.Errorf("收件人不能为空")
|
||||
}
|
||||
if cfg.Host == "" || cfg.FromAddress == "" {
|
||||
return fmt.Errorf("SMTP 主机或发件人不能为空")
|
||||
}
|
||||
if cfg.Port == 0 {
|
||||
cfg.Port = 465
|
||||
}
|
||||
timeout := cfg.Timeout
|
||||
if timeout == 0 {
|
||||
timeout = 30
|
||||
}
|
||||
d := net.Dialer{Timeout: time.Duration(timeout) * time.Second}
|
||||
addr := net.JoinHostPort(cfg.Host, strconv.FormatUint(uint64(cfg.Port), 10))
|
||||
enc := strings.ToLower(strings.TrimSpace(cfg.Encryption))
|
||||
if enc == "" {
|
||||
enc = "ssl"
|
||||
}
|
||||
|
||||
var client *smtp.Client
|
||||
var err error
|
||||
|
||||
switch enc {
|
||||
case "ssl":
|
||||
conn, derr := tls.DialWithDialer(&d, "tcp", addr, &tls.Config{ServerName: cfg.Host, MinVersion: tls.VersionTLS12})
|
||||
if derr != nil {
|
||||
return fmt.Errorf("连接 SMTP 失败: %w", derr)
|
||||
}
|
||||
defer conn.Close()
|
||||
client, err = smtp.NewClient(conn, cfg.Host)
|
||||
if err != nil {
|
||||
return fmt.Errorf("SMTP 握手失败: %w", err)
|
||||
}
|
||||
case "tls":
|
||||
conn, derr := d.Dial("tcp", addr)
|
||||
if derr != nil {
|
||||
return fmt.Errorf("连接 SMTP 失败: %w", derr)
|
||||
}
|
||||
defer conn.Close()
|
||||
client, err = smtp.NewClient(conn, cfg.Host)
|
||||
if err != nil {
|
||||
return fmt.Errorf("SMTP 握手失败: %w", err)
|
||||
}
|
||||
if ok, _ := client.Extension("STARTTLS"); ok {
|
||||
if err = client.StartTLS(&tls.Config{ServerName: cfg.Host, MinVersion: tls.VersionTLS12}); err != nil {
|
||||
_ = client.Close()
|
||||
return fmt.Errorf("STARTTLS 失败: %w", err)
|
||||
}
|
||||
}
|
||||
case "none":
|
||||
conn, derr := d.Dial("tcp", addr)
|
||||
if derr != nil {
|
||||
return fmt.Errorf("连接 SMTP 失败: %w", derr)
|
||||
}
|
||||
defer conn.Close()
|
||||
client, err = smtp.NewClient(conn, cfg.Host)
|
||||
if err != nil {
|
||||
return fmt.Errorf("SMTP 握手失败: %w", err)
|
||||
}
|
||||
default:
|
||||
return fmt.Errorf("不支持的加密方式: %s", cfg.Encryption)
|
||||
}
|
||||
defer func() { _ = client.Close() }()
|
||||
|
||||
auth := smtp.PlainAuth("", cfg.FromAddress, cfg.Password, cfg.Host)
|
||||
if err = client.Auth(auth); err != nil {
|
||||
return fmt.Errorf("SMTP 认证失败: %w", err)
|
||||
}
|
||||
if err = client.Mail(cfg.FromAddress); err != nil {
|
||||
return fmt.Errorf("MAIL FROM 失败: %w", err)
|
||||
}
|
||||
if err = client.Rcpt(to); err != nil {
|
||||
return fmt.Errorf("RCPT TO 失败: %w", err)
|
||||
}
|
||||
wc, err := client.Data()
|
||||
if err != nil {
|
||||
return fmt.Errorf("DATA 失败: %w", err)
|
||||
}
|
||||
fromName := strings.TrimSpace(cfg.FromName)
|
||||
subject := "平台邮箱测试"
|
||||
body := "这是一封来自管理后台「邮箱管理」的测试邮件。\r\nThis is a test email from the platform email settings.\r\n"
|
||||
headers := fmt.Sprintf("From: %s\r\nTo: %s\r\nSubject: %s\r\nMIME-Version: 1.0\r\nContent-Type: text/plain; charset=UTF-8\r\nContent-Transfer-Encoding: 8bit\r\n\r\n",
|
||||
formatFromHeader(fromName, cfg.FromAddress), to, subject)
|
||||
if _, err = wc.Write([]byte(headers + body)); err != nil {
|
||||
return fmt.Errorf("写入邮件内容失败: %w", err)
|
||||
}
|
||||
if err = wc.Close(); err != nil {
|
||||
return fmt.Errorf("结束 DATA 失败: %w", err)
|
||||
}
|
||||
return client.Quit()
|
||||
}
|
||||
|
||||
func formatFromHeader(name, addr string) string {
|
||||
name = strings.TrimSpace(name)
|
||||
if name == "" {
|
||||
return addr
|
||||
}
|
||||
return fmt.Sprintf("%s <%s>", name, addr)
|
||||
}
|
||||
@@ -0,0 +1,77 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"server/models"
|
||||
)
|
||||
|
||||
// ListSystemEmails 返回全部邮箱配置(按 id 升序,通常仅一条)
|
||||
func ListSystemEmails() ([]models.SystemEmail, error) {
|
||||
var rows []models.SystemEmail
|
||||
_, err := models.Orm.QueryTable(new(models.SystemEmail)).OrderBy("id").All(&rows)
|
||||
return rows, err
|
||||
}
|
||||
|
||||
// UpsertFirstSystemEmail 若已有记录则更新第一条,否则插入
|
||||
func UpsertFirstSystemEmail(fromAddress string, fromName *string, host string, port uint, password string, encryption string, timeout uint, status int8, remark *string) error {
|
||||
if encryption == "" {
|
||||
encryption = "ssl"
|
||||
}
|
||||
if port == 0 {
|
||||
port = 465
|
||||
}
|
||||
if timeout == 0 {
|
||||
timeout = 30
|
||||
}
|
||||
if status == 0 {
|
||||
status = 1
|
||||
}
|
||||
fromAddress = strings.TrimSpace(fromAddress)
|
||||
host = strings.TrimSpace(host)
|
||||
|
||||
cnt, err := models.Orm.QueryTable(new(models.SystemEmail)).Count()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if cnt == 0 {
|
||||
if strings.TrimSpace(password) == "" {
|
||||
return fmt.Errorf("首次保存必须填写授权码/密码")
|
||||
}
|
||||
row := &models.SystemEmail{
|
||||
FromAddress: fromAddress,
|
||||
FromName: fromName,
|
||||
Host: host,
|
||||
Port: port,
|
||||
Password: strings.TrimSpace(password),
|
||||
Encryption: encryption,
|
||||
Timeout: timeout,
|
||||
Status: status,
|
||||
Remark: remark,
|
||||
}
|
||||
_, err = models.Orm.Insert(row)
|
||||
return err
|
||||
}
|
||||
|
||||
var first models.SystemEmail
|
||||
if err := models.Orm.QueryTable(new(models.SystemEmail)).OrderBy("id").Limit(1).One(&first); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
up := map[string]interface{}{
|
||||
"from_address": fromAddress,
|
||||
"from_name": fromName,
|
||||
"host": host,
|
||||
"port": port,
|
||||
"encryption": encryption,
|
||||
"timeout": timeout,
|
||||
"status": status,
|
||||
"remark": remark,
|
||||
}
|
||||
if strings.TrimSpace(password) != "" {
|
||||
up["password"] = strings.TrimSpace(password)
|
||||
}
|
||||
_, err = models.Orm.QueryTable(new(models.SystemEmail)).Filter("id", first.ID).Update(up)
|
||||
return err
|
||||
}
|
||||
@@ -0,0 +1,144 @@
|
||||
package services
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"server/models"
|
||||
)
|
||||
|
||||
// BindTenantUser 绑定用户到租户(若已存在则更新状态/默认值)
|
||||
func BindTenantUser(tid, uid uint64, account, name, phone, email *string, sex *uint8, birth *string, password *string, isDefault, status int8, remark *string) (uint64, error) {
|
||||
var existed models.SystemTenantUser
|
||||
err := models.Orm.QueryTable(new(models.SystemTenantUser)).
|
||||
Filter("tid", tid).
|
||||
Filter("uid", uid).
|
||||
One(&existed)
|
||||
if err == nil {
|
||||
update := map[string]interface{}{
|
||||
"account": account,
|
||||
"name": name,
|
||||
"phone": phone,
|
||||
"email": email,
|
||||
"password": password,
|
||||
"status": status,
|
||||
"is_default": isDefault,
|
||||
"remark": remark,
|
||||
}
|
||||
if sex != nil {
|
||||
update["sex"] = *sex
|
||||
}
|
||||
if birth != nil {
|
||||
trimmedBirth := strings.TrimSpace(*birth)
|
||||
if trimmedBirth == "" {
|
||||
update["birth"] = nil
|
||||
} else {
|
||||
update["birth"] = trimmedBirth
|
||||
}
|
||||
}
|
||||
_, uErr := models.Orm.QueryTable(new(models.SystemTenantUser)).Filter("id", existed.ID).Update(update)
|
||||
return existed.ID, uErr
|
||||
}
|
||||
|
||||
m := &models.SystemTenantUser{
|
||||
Tid: tid,
|
||||
Uid: uid,
|
||||
Account: account,
|
||||
Name: name,
|
||||
Phone: phone,
|
||||
Email: email,
|
||||
Password: password,
|
||||
IsDefault: isDefault,
|
||||
Status: status,
|
||||
Remark: remark,
|
||||
}
|
||||
if sex != nil {
|
||||
m.Sex = *sex
|
||||
}
|
||||
if birth != nil {
|
||||
trimmedBirth := strings.TrimSpace(*birth)
|
||||
if trimmedBirth != "" {
|
||||
m.Birth = &trimmedBirth
|
||||
}
|
||||
}
|
||||
id, iErr := models.Orm.Insert(m)
|
||||
return uint64(id), iErr
|
||||
}
|
||||
|
||||
// UnbindTenantUser 删除绑定关系
|
||||
func UnbindTenantUser(id uint64) error {
|
||||
_, err := models.Orm.QueryTable(new(models.SystemTenantUser)).Filter("id", id).Delete()
|
||||
return err
|
||||
}
|
||||
|
||||
// ListTenantUsersByTid 根据租户ID查询绑定关系
|
||||
func ListTenantUsersByTid(tid uint64) ([]models.SystemTenantUser, error) {
|
||||
var rows []models.SystemTenantUser
|
||||
_, err := models.Orm.QueryTable(new(models.SystemTenantUser)).
|
||||
Filter("tid", tid).
|
||||
OrderBy("-is_default", "-id").
|
||||
All(&rows)
|
||||
return rows, err
|
||||
}
|
||||
|
||||
// ListTenantBindingsByUid 根据用户ID查询绑定关系
|
||||
func ListTenantBindingsByUid(uid uint64) ([]models.SystemTenantUser, error) {
|
||||
var rows []models.SystemTenantUser
|
||||
_, err := models.Orm.QueryTable(new(models.SystemTenantUser)).
|
||||
Filter("uid", uid).
|
||||
OrderBy("-is_default", "-id").
|
||||
All(&rows)
|
||||
return rows, err
|
||||
}
|
||||
|
||||
// GetTenantUserByUidAndTid 根据用户ID和租户ID查询租户用户绑定关系
|
||||
func GetTenantUserByUidAndTid(uid, tid uint64) (*models.SystemTenantUser, error) {
|
||||
var row models.SystemTenantUser
|
||||
err := models.Orm.QueryTable(new(models.SystemTenantUser)).
|
||||
Filter("uid", uid).
|
||||
Filter("tid", tid).
|
||||
One(&row)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &row, nil
|
||||
}
|
||||
|
||||
// GetTenantUserByUid 根据用户ID查询默认/最新租户用户绑定关系
|
||||
func GetTenantUserByUid(uid uint64) (*models.SystemTenantUser, error) {
|
||||
var row models.SystemTenantUser
|
||||
err := models.Orm.QueryTable(new(models.SystemTenantUser)).
|
||||
Filter("uid", uid).
|
||||
OrderBy("-is_default", "-id").
|
||||
One(&row)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &row, nil
|
||||
}
|
||||
|
||||
// GetTenantByID 根据租户ID查询租户信息
|
||||
func GetTenantByID(id uint64) (*models.SystemTenant, error) {
|
||||
var row models.SystemTenant
|
||||
err := models.Orm.QueryTable(new(models.SystemTenant)).
|
||||
Filter("id", id).
|
||||
One(&row)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &row, nil
|
||||
}
|
||||
|
||||
// SetDefaultTenant 设置用户默认租户(同一用户仅一个默认)
|
||||
func SetDefaultTenant(uid, tid uint64) error {
|
||||
_, err := models.Orm.QueryTable(new(models.SystemTenantUser)).Filter("uid", uid).Update(map[string]interface{}{
|
||||
"is_default": 0,
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
_, err = models.Orm.QueryTable(new(models.SystemTenantUser)).
|
||||
Filter("uid", uid).
|
||||
Filter("tid", tid).
|
||||
Update(map[string]interface{}{"is_default": 1})
|
||||
return err
|
||||
}
|
||||
Reference in New Issue
Block a user