init
This commit is contained in:
@@ -0,0 +1,33 @@
|
||||
package core
|
||||
|
||||
// Config 结构体
|
||||
type Config struct {
|
||||
RunType string `json:"run_type"`
|
||||
LocalAddr string `json:"local_addr"`
|
||||
LocalPort int `json:"local_port"`
|
||||
RemoteAddr string `json:"remote_addr"`
|
||||
RemotePort int `json:"remote_port"`
|
||||
Password []string `json:"password"`
|
||||
LogLevel int `json:"log_level"`
|
||||
}
|
||||
|
||||
// SSL 结构体
|
||||
type SSL struct {
|
||||
Cert string `json:"cert"`
|
||||
Cipher string `json:"cipher"`
|
||||
CipherTls13 string `json:"cipher_tls13"`
|
||||
Alpn []string `json:"alpn"`
|
||||
ReuseSession bool `json:"reuse_session"`
|
||||
SessionTicket bool `json:"session_ticket"`
|
||||
Curves string `json:"curves"`
|
||||
Sni string `json:"sni"`
|
||||
}
|
||||
|
||||
// TCP 结构体
|
||||
type TCP struct {
|
||||
NoDelay bool `json:"no_delay"`
|
||||
KeepAlive bool `json:"keep_alive"`
|
||||
ReusePort bool `json:"reuse_port"`
|
||||
FastOpen bool `json:"fast_open"`
|
||||
FastOpenQlen int `json:"fast_open_qlen"`
|
||||
}
|
||||
@@ -0,0 +1,50 @@
|
||||
package core
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"trojan/asset"
|
||||
)
|
||||
|
||||
// ClientConfig 结构体
|
||||
type ClientConfig struct {
|
||||
Config
|
||||
SSl ClientSSL `json:"ssl"`
|
||||
Tcp ClientTCP `json:"tcp"`
|
||||
}
|
||||
|
||||
// ClientSSL 结构体
|
||||
type ClientSSL struct {
|
||||
SSL
|
||||
Verify bool `json:"verify"`
|
||||
VerifyHostname bool `json:"verify_hostname"`
|
||||
}
|
||||
|
||||
// ClientTCP 结构体
|
||||
type ClientTCP struct {
|
||||
TCP
|
||||
}
|
||||
|
||||
// WriteClient 生成客户端json
|
||||
func WriteClient(port int, password, domain, writePath string) bool {
|
||||
data := asset.GetAsset("client.json")
|
||||
config := ClientConfig{}
|
||||
if err := json.Unmarshal(data, &config); err != nil {
|
||||
fmt.Println(err)
|
||||
return false
|
||||
}
|
||||
config.RemoteAddr = domain
|
||||
config.RemotePort = port
|
||||
config.Password = []string{password}
|
||||
outData, err := json.MarshalIndent(config, "", " ")
|
||||
if err != nil {
|
||||
fmt.Println(err)
|
||||
return false
|
||||
}
|
||||
if err = os.WriteFile(writePath, outData, 0644); err != nil {
|
||||
fmt.Println(err)
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
@@ -0,0 +1,41 @@
|
||||
package core
|
||||
|
||||
import (
|
||||
"github.com/syndtr/goleveldb/leveldb"
|
||||
)
|
||||
|
||||
var dbPath = "/var/lib/trojan-manager"
|
||||
|
||||
// GetValue 获取leveldb值
|
||||
func GetValue(key string) (string, error) {
|
||||
db, err := leveldb.OpenFile(dbPath, nil)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer db.Close()
|
||||
result, err := db.Get([]byte(key), nil)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return string(result), nil
|
||||
}
|
||||
|
||||
// SetValue 设置leveldb值
|
||||
func SetValue(key string, value string) error {
|
||||
db, err := leveldb.OpenFile(dbPath, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer db.Close()
|
||||
return db.Put([]byte(key), []byte(value), nil)
|
||||
}
|
||||
|
||||
// DelValue 删除值
|
||||
func DelValue(key string) error {
|
||||
db, err := leveldb.OpenFile(dbPath, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer db.Close()
|
||||
return db.Delete([]byte(key), nil)
|
||||
}
|
||||
@@ -0,0 +1,412 @@
|
||||
package core
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"fmt"
|
||||
mysqlDriver "github.com/go-sql-driver/mysql"
|
||||
"io"
|
||||
"log"
|
||||
"time"
|
||||
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
// mysql sql驱动
|
||||
_ "github.com/go-sql-driver/mysql"
|
||||
)
|
||||
|
||||
// Mysql 结构体
|
||||
type Mysql struct {
|
||||
Enabled bool `json:"enabled"`
|
||||
ServerAddr string `json:"server_addr"`
|
||||
ServerPort int `json:"server_port"`
|
||||
Database string `json:"database"`
|
||||
Username string `json:"username"`
|
||||
Password string `json:"password"`
|
||||
Cafile string `json:"cafile"`
|
||||
}
|
||||
|
||||
// User 用户表记录结构体
|
||||
type User struct {
|
||||
ID uint
|
||||
Username string
|
||||
Password string
|
||||
EncryptPass string
|
||||
Quota int64
|
||||
Download uint64
|
||||
Upload uint64
|
||||
UseDays uint
|
||||
ExpiryDate string
|
||||
}
|
||||
|
||||
// PageQuery 分页查询的结构体
|
||||
type PageQuery struct {
|
||||
PageNum int
|
||||
CurPage int
|
||||
Total int
|
||||
PageSize int
|
||||
DataList []*User
|
||||
}
|
||||
|
||||
// CreateTableSql 创表sql
|
||||
var CreateTableSql = `
|
||||
CREATE TABLE IF NOT EXISTS users (
|
||||
id INT UNSIGNED NOT NULL AUTO_INCREMENT,
|
||||
username VARCHAR(64) NOT NULL,
|
||||
password CHAR(56) NOT NULL,
|
||||
passwordShow VARCHAR(255) NOT NULL,
|
||||
quota BIGINT NOT NULL DEFAULT 0,
|
||||
download BIGINT UNSIGNED NOT NULL DEFAULT 0,
|
||||
upload BIGINT UNSIGNED NOT NULL DEFAULT 0,
|
||||
useDays int(10) DEFAULT 0,
|
||||
expiryDate char(10) DEFAULT '',
|
||||
PRIMARY KEY (id),
|
||||
INDEX (password)
|
||||
) DEFAULT CHARSET=utf8mb4;
|
||||
`
|
||||
|
||||
// GetDB 获取mysql数据库连接
|
||||
func (mysql *Mysql) GetDB() *sql.DB {
|
||||
// 屏蔽mysql驱动包的日志输出
|
||||
mysqlDriver.SetLogger(log.New(io.Discard, "", 0))
|
||||
conn := fmt.Sprintf("%s:%s@tcp(%s:%d)/%s", mysql.Username, mysql.Password, mysql.ServerAddr, mysql.ServerPort, mysql.Database)
|
||||
db, err := sql.Open("mysql", conn)
|
||||
if err != nil {
|
||||
fmt.Println(err.Error())
|
||||
return nil
|
||||
}
|
||||
return db
|
||||
}
|
||||
|
||||
// CreateTable 不存在trojan user表则自动创建
|
||||
func (mysql *Mysql) CreateTable() {
|
||||
db := mysql.GetDB()
|
||||
defer db.Close()
|
||||
if _, err := db.Exec(CreateTableSql); err != nil {
|
||||
fmt.Println(err)
|
||||
}
|
||||
}
|
||||
|
||||
func queryUserList(db *sql.DB, sql string) ([]*User, error) {
|
||||
var (
|
||||
username string
|
||||
encryptPass string
|
||||
passShow string
|
||||
download uint64
|
||||
upload uint64
|
||||
quota int64
|
||||
id uint
|
||||
useDays uint
|
||||
expiryDate string
|
||||
)
|
||||
var userList []*User
|
||||
rows, err := db.Query(sql)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
for rows.Next() {
|
||||
if err := rows.Scan(&id, &username, &encryptPass, &passShow, "a, &download, &upload, &useDays, &expiryDate); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
userList = append(userList, &User{
|
||||
ID: id,
|
||||
Username: username,
|
||||
Password: passShow,
|
||||
EncryptPass: encryptPass,
|
||||
Download: download,
|
||||
Upload: upload,
|
||||
Quota: quota,
|
||||
UseDays: useDays,
|
||||
ExpiryDate: expiryDate,
|
||||
})
|
||||
}
|
||||
return userList, nil
|
||||
}
|
||||
|
||||
func queryUser(db *sql.DB, sql string) (*User, error) {
|
||||
var (
|
||||
username string
|
||||
encryptPass string
|
||||
passShow string
|
||||
download uint64
|
||||
upload uint64
|
||||
quota int64
|
||||
id uint
|
||||
useDays uint
|
||||
expiryDate string
|
||||
)
|
||||
row := db.QueryRow(sql)
|
||||
if err := row.Scan(&id, &username, &encryptPass, &passShow, "a, &download, &upload, &useDays, &expiryDate); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &User{ID: id, Username: username, Password: passShow, EncryptPass: encryptPass, Download: download, Upload: upload, Quota: quota, UseDays: useDays, ExpiryDate: expiryDate}, nil
|
||||
}
|
||||
|
||||
// CreateUser 创建Trojan用户
|
||||
func (mysql *Mysql) CreateUser(username string, base64Pass string, originPass string) error {
|
||||
db := mysql.GetDB()
|
||||
if db == nil {
|
||||
return errors.New("can't connect mysql")
|
||||
}
|
||||
defer db.Close()
|
||||
encryPass := sha256.Sum224([]byte(originPass))
|
||||
if _, err := db.Exec(fmt.Sprintf("INSERT INTO users(username, password, passwordShow, quota) VALUES ('%s', '%x', '%s', -1);", username, encryPass, base64Pass)); err != nil {
|
||||
fmt.Println(err)
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// UpdateUser 更新Trojan用户名和密码
|
||||
func (mysql *Mysql) UpdateUser(id uint, username string, base64Pass string, originPass string) error {
|
||||
db := mysql.GetDB()
|
||||
if db == nil {
|
||||
return errors.New("can't connect mysql")
|
||||
}
|
||||
defer db.Close()
|
||||
encryPass := sha256.Sum224([]byte(originPass))
|
||||
if _, err := db.Exec(fmt.Sprintf("UPDATE users SET username='%s', password='%x', passwordShow='%s' WHERE id=%d;", username, encryPass, base64Pass, id)); err != nil {
|
||||
fmt.Println(err)
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// DeleteUser 删除用户
|
||||
func (mysql *Mysql) DeleteUser(id uint) error {
|
||||
db := mysql.GetDB()
|
||||
if db == nil {
|
||||
return errors.New("can't connect mysql")
|
||||
}
|
||||
defer db.Close()
|
||||
if userList, err := mysql.GetData(strconv.Itoa(int(id))); err != nil {
|
||||
return err
|
||||
} else if userList != nil && len(userList) == 0 {
|
||||
return fmt.Errorf("不存在id为%d的用户", id)
|
||||
}
|
||||
if _, err := db.Exec(fmt.Sprintf("DELETE FROM users WHERE id=%d;", id)); err != nil {
|
||||
fmt.Println(err)
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// MonthlyResetData 设置了过期时间的用户,每月定时清空使用流量
|
||||
func (mysql *Mysql) MonthlyResetData() error {
|
||||
db := mysql.GetDB()
|
||||
if db == nil {
|
||||
return errors.New("can't connect mysql")
|
||||
}
|
||||
defer db.Close()
|
||||
userList, err := queryUserList(db, "SELECT * FROM users WHERE useDays != 0 AND quota != 0")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
for _, user := range userList {
|
||||
if _, err := db.Exec(fmt.Sprintf("UPDATE users SET download=0, upload=0 WHERE id=%d;", user.ID)); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// DailyCheckExpire 检查是否有过期,过期了设置流量上限为0
|
||||
func (mysql *Mysql) DailyCheckExpire() (bool, error) {
|
||||
needRestart := false
|
||||
now := time.Now()
|
||||
utc, err := time.LoadLocation("Asia/Shanghai")
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
addDay, _ := time.ParseDuration("-24h")
|
||||
yesterdayStr := now.Add(addDay).In(utc).Format("2006-01-02")
|
||||
yesterday, _ := time.Parse("2006-01-02", yesterdayStr)
|
||||
db := mysql.GetDB()
|
||||
if db == nil {
|
||||
return false, errors.New("can't connect mysql")
|
||||
}
|
||||
defer db.Close()
|
||||
userList, err := queryUserList(db, "SELECT * FROM users WHERE quota != 0")
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
for _, user := range userList {
|
||||
if expireDate, err := time.Parse("2006-01-02", user.ExpiryDate); err == nil {
|
||||
if yesterday.Sub(expireDate).Seconds() >= 0 {
|
||||
if _, err := db.Exec(fmt.Sprintf("UPDATE users SET quota=0 WHERE id=%d;", user.ID)); err != nil {
|
||||
return false, err
|
||||
}
|
||||
if !needRestart {
|
||||
needRestart = true
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
return needRestart, nil
|
||||
}
|
||||
|
||||
// CancelExpire 取消过期时间
|
||||
func (mysql *Mysql) CancelExpire(id uint) error {
|
||||
db := mysql.GetDB()
|
||||
if db == nil {
|
||||
return errors.New("can't connect mysql")
|
||||
}
|
||||
defer db.Close()
|
||||
if _, err := db.Exec(fmt.Sprintf("UPDATE users SET useDays=0, expiryDate='' WHERE id=%d;", id)); err != nil {
|
||||
fmt.Println(err)
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// SetExpire 设置过期时间
|
||||
func (mysql *Mysql) SetExpire(id uint, useDays uint) error {
|
||||
now := time.Now()
|
||||
utc, err := time.LoadLocation("Asia/Shanghai")
|
||||
if err != nil {
|
||||
fmt.Println(err)
|
||||
return err
|
||||
}
|
||||
addDay, _ := time.ParseDuration(strconv.Itoa(int(24*useDays)) + "h")
|
||||
expiryDate := now.Add(addDay).In(utc).Format("2006-01-02")
|
||||
|
||||
db := mysql.GetDB()
|
||||
if db == nil {
|
||||
return errors.New("can't connect mysql")
|
||||
}
|
||||
defer db.Close()
|
||||
if _, err := db.Exec(fmt.Sprintf("UPDATE users SET useDays=%d, expiryDate='%s' WHERE id=%d;", useDays, expiryDate, id)); err != nil {
|
||||
fmt.Println(err)
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// SetQuota 限制流量
|
||||
func (mysql *Mysql) SetQuota(id uint, quota int) error {
|
||||
db := mysql.GetDB()
|
||||
if db == nil {
|
||||
return errors.New("can't connect mysql")
|
||||
}
|
||||
defer db.Close()
|
||||
if _, err := db.Exec(fmt.Sprintf("UPDATE users SET quota=%d WHERE id=%d;", quota, id)); err != nil {
|
||||
fmt.Println(err)
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// CleanData 清空流量统计
|
||||
func (mysql *Mysql) CleanData(id uint) error {
|
||||
db := mysql.GetDB()
|
||||
if db == nil {
|
||||
return errors.New("can't connect mysql")
|
||||
}
|
||||
defer db.Close()
|
||||
if _, err := db.Exec(fmt.Sprintf("UPDATE users SET download=0, upload=0 WHERE id=%d;", id)); err != nil {
|
||||
fmt.Println(err)
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// CleanDataByName 清空指定用户名流量统计数据
|
||||
func (mysql *Mysql) CleanDataByName(usernames []string) error {
|
||||
db := mysql.GetDB()
|
||||
if db == nil {
|
||||
return errors.New("can't connect mysql")
|
||||
}
|
||||
defer db.Close()
|
||||
runSql := "UPDATE users SET download=0, upload=0 WHERE BINARY username in ("
|
||||
for i, name := range usernames {
|
||||
runSql = runSql + "'" + name + "'"
|
||||
if i == len(usernames)-1 {
|
||||
runSql = runSql + ")"
|
||||
} else {
|
||||
runSql = runSql + ","
|
||||
}
|
||||
}
|
||||
if _, err := db.Exec(runSql); err != nil {
|
||||
fmt.Println(err)
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetUserByName 通过用户名来获取用户
|
||||
func (mysql *Mysql) GetUserByName(name string) *User {
|
||||
db := mysql.GetDB()
|
||||
if db == nil {
|
||||
return nil
|
||||
}
|
||||
defer db.Close()
|
||||
user, err := queryUser(db, fmt.Sprintf("SELECT * FROM users WHERE BINARY username='%s'", name))
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
return user
|
||||
}
|
||||
|
||||
// GetUserByPass 通过密码来获取用户
|
||||
func (mysql *Mysql) GetUserByPass(pass string) *User {
|
||||
db := mysql.GetDB()
|
||||
if db == nil {
|
||||
return nil
|
||||
}
|
||||
defer db.Close()
|
||||
user, err := queryUser(db, fmt.Sprintf("SELECT * FROM users WHERE BINARY passwordShow='%s'", pass))
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
return user
|
||||
}
|
||||
|
||||
// PageList 通过分页获取用户记录
|
||||
func (mysql *Mysql) PageList(curPage int, pageSize int) (*PageQuery, error) {
|
||||
var (
|
||||
total int
|
||||
)
|
||||
|
||||
db := mysql.GetDB()
|
||||
if db == nil {
|
||||
return nil, errors.New("连接mysql失败")
|
||||
}
|
||||
defer db.Close()
|
||||
offset := (curPage - 1) * pageSize
|
||||
querySQL := fmt.Sprintf("SELECT * FROM users LIMIT %d, %d", offset, pageSize)
|
||||
userList, err := queryUserList(db, querySQL)
|
||||
if err != nil {
|
||||
fmt.Println(err)
|
||||
return nil, err
|
||||
}
|
||||
db.QueryRow("SELECT COUNT(id) FROM users").Scan(&total)
|
||||
return &PageQuery{
|
||||
CurPage: curPage,
|
||||
PageSize: pageSize,
|
||||
Total: total,
|
||||
DataList: userList,
|
||||
PageNum: (total + pageSize - 1) / pageSize,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// GetData 获取用户记录
|
||||
func (mysql *Mysql) GetData(ids ...string) ([]*User, error) {
|
||||
querySQL := "SELECT * FROM users"
|
||||
db := mysql.GetDB()
|
||||
if db == nil {
|
||||
return nil, errors.New("连接mysql失败")
|
||||
}
|
||||
defer db.Close()
|
||||
if len(ids) > 0 {
|
||||
querySQL = querySQL + " WHERE id in (" + strings.Join(ids, ",") + ")"
|
||||
}
|
||||
userList, err := queryUserList(db, querySQL)
|
||||
if err != nil {
|
||||
fmt.Println(err)
|
||||
return nil, err
|
||||
}
|
||||
return userList, nil
|
||||
}
|
||||
@@ -0,0 +1,122 @@
|
||||
package core
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"github.com/tidwall/pretty"
|
||||
"github.com/tidwall/sjson"
|
||||
"os"
|
||||
)
|
||||
|
||||
var configPath = "/usr/local/etc/trojan/config.json"
|
||||
|
||||
// ServerConfig 结构体
|
||||
type ServerConfig struct {
|
||||
Config
|
||||
SSl ServerSSL `json:"ssl"`
|
||||
Tcp ServerTCP `json:"tcp"`
|
||||
Mysql Mysql `json:"mysql"`
|
||||
}
|
||||
|
||||
// ServerSSL 结构体
|
||||
type ServerSSL struct {
|
||||
SSL
|
||||
Key string `json:"key"`
|
||||
KeyPassword string `json:"key_password"`
|
||||
PreferServerCipher bool `json:"prefer_server_cipher"`
|
||||
SessionTimeout int `json:"session_timeout"`
|
||||
PlainHttpResponse string `json:"plain_http_response"`
|
||||
Dhparam string `json:"dhparam"`
|
||||
}
|
||||
|
||||
// ServerTCP 结构体
|
||||
type ServerTCP struct {
|
||||
TCP
|
||||
PreferIPv4 bool `json:"prefer_ipv4"`
|
||||
}
|
||||
|
||||
// Load 加载服务端配置文件
|
||||
func Load(path string) []byte {
|
||||
if path == "" {
|
||||
path = configPath
|
||||
}
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
fmt.Println(err)
|
||||
return nil
|
||||
}
|
||||
return data
|
||||
}
|
||||
|
||||
// Save 保存服务端配置文件
|
||||
func Save(data []byte, path string) bool {
|
||||
if path == "" {
|
||||
path = configPath
|
||||
}
|
||||
if err := os.WriteFile(path, pretty.Pretty(data), 0644); err != nil {
|
||||
fmt.Println(err)
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// GetConfig 获取config配置
|
||||
func GetConfig() *ServerConfig {
|
||||
data := Load("")
|
||||
config := ServerConfig{}
|
||||
if err := json.Unmarshal(data, &config); err != nil {
|
||||
fmt.Println(err)
|
||||
return nil
|
||||
}
|
||||
return &config
|
||||
}
|
||||
|
||||
// GetMysql 获取mysql连接
|
||||
func GetMysql() *Mysql {
|
||||
return &GetConfig().Mysql
|
||||
}
|
||||
|
||||
// WriteMysql 写mysql配置
|
||||
func WriteMysql(mysql *Mysql) bool {
|
||||
mysql.Enabled = true
|
||||
data := Load("")
|
||||
result, _ := sjson.SetBytes(data, "mysql", mysql)
|
||||
return Save(result, "")
|
||||
}
|
||||
|
||||
// WriteTls 写tls配置
|
||||
func WriteTls(cert, key, domain string) bool {
|
||||
data := Load("")
|
||||
data, _ = sjson.SetBytes(data, "ssl.cert", cert)
|
||||
data, _ = sjson.SetBytes(data, "ssl.key", key)
|
||||
data, _ = sjson.SetBytes(data, "ssl.sni", domain)
|
||||
return Save(data, "")
|
||||
}
|
||||
|
||||
// WriteDomain 写域名
|
||||
func WriteDomain(domain string) bool {
|
||||
data := Load("")
|
||||
data, _ = sjson.SetBytes(data, "ssl.sni", domain)
|
||||
return Save(data, "")
|
||||
}
|
||||
|
||||
// WritePassword 写密码
|
||||
func WritePassword(pass []string) bool {
|
||||
data := Load("")
|
||||
data, _ = sjson.SetBytes(data, "password", pass)
|
||||
return Save(data, "")
|
||||
}
|
||||
|
||||
// WritePort 写trojan端口
|
||||
func WritePort(port int) bool {
|
||||
data := Load("")
|
||||
data, _ = sjson.SetBytes(data, "local_port", port)
|
||||
return Save(data, "")
|
||||
}
|
||||
|
||||
// WriteLogLevel 写日志等级
|
||||
func WriteLogLevel(level int) bool {
|
||||
data := Load("")
|
||||
data, _ = sjson.SetBytes(data, "log_level", level)
|
||||
return Save(data, "")
|
||||
}
|
||||
@@ -0,0 +1,113 @@
|
||||
package core
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"database/sql"
|
||||
"encoding/base64"
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"strings"
|
||||
"trojan/util"
|
||||
)
|
||||
|
||||
// UpgradeDB 升级数据库表结构以及迁移数据
|
||||
func (mysql *Mysql) UpgradeDB() error {
|
||||
db := mysql.GetDB()
|
||||
if db == nil {
|
||||
return errors.New("can't connect mysql")
|
||||
}
|
||||
var field string
|
||||
err := db.QueryRow("SHOW COLUMNS FROM users LIKE 'passwordShow';").Scan(&field)
|
||||
if err == sql.ErrNoRows {
|
||||
fmt.Println(util.Yellow("正在进行数据库升级, 请稍等.."))
|
||||
if _, err := db.Exec("ALTER TABLE users ADD COLUMN passwordShow VARCHAR(255) NOT NULL AFTER password;"); err != nil {
|
||||
fmt.Println(err)
|
||||
return err
|
||||
}
|
||||
userList, err := mysql.GetData()
|
||||
if err != nil {
|
||||
fmt.Println(err)
|
||||
return err
|
||||
}
|
||||
for _, user := range userList {
|
||||
pass, _ := GetValue(fmt.Sprintf("%s_pass", user.Username))
|
||||
if pass != "" {
|
||||
base64Pass := base64.StdEncoding.EncodeToString([]byte(pass))
|
||||
if _, err := db.Exec(fmt.Sprintf("UPDATE users SET passwordShow='%s' WHERE id=%d;", base64Pass, user.ID)); err != nil {
|
||||
fmt.Println(err)
|
||||
return err
|
||||
}
|
||||
DelValue(fmt.Sprintf("%s_pass", user.Username))
|
||||
}
|
||||
}
|
||||
}
|
||||
err = db.QueryRow("SHOW COLUMNS FROM users LIKE 'useDays';").Scan(&field)
|
||||
if err == sql.ErrNoRows {
|
||||
fmt.Println(util.Yellow("正在进行数据库升级, 请稍等.."))
|
||||
if _, err := db.Exec(`
|
||||
ALTER TABLE users
|
||||
ADD COLUMN useDays int(10) DEFAULT 0,
|
||||
ADD COLUMN expiryDate char(10) DEFAULT '';
|
||||
`); err != nil {
|
||||
fmt.Println(err)
|
||||
return err
|
||||
}
|
||||
}
|
||||
var tableName string
|
||||
err = db.QueryRow(fmt.Sprintf(
|
||||
"SELECT * FROM information_schema.TABLES WHERE TABLE_NAME = 'users' AND TABLE_SCHEMA = '%s' ",
|
||||
mysql.Database) + " AND TABLE_COLLATION LIKE 'utf8%';").Scan(&tableName)
|
||||
if err == sql.ErrNoRows {
|
||||
tempFile := "temp.sql"
|
||||
mysql.DumpSql(tempFile)
|
||||
mysql.ExecSql(tempFile)
|
||||
os.Remove(tempFile)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// DumpSql 导出sql
|
||||
func (mysql *Mysql) DumpSql(filePath string) error {
|
||||
file, err := os.OpenFile(filePath, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, 0644)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer file.Close()
|
||||
writer := bufio.NewWriter(file)
|
||||
writer.WriteString("DROP TABLE IF EXISTS users;")
|
||||
writer.WriteString(CreateTableSql)
|
||||
db := mysql.GetDB()
|
||||
userList, err := queryUserList(db, "SELECT * FROM users;")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
for _, user := range userList {
|
||||
writer.WriteString(fmt.Sprintf(`
|
||||
INSERT INTO users(username, password, passwordShow, quota, download, upload, useDays, expiryDate) VALUES ('%s','%s','%s', %d, %d, %d, %d, '%s');`,
|
||||
user.Username, user.EncryptPass, user.Password, user.Quota, user.Download, user.Upload, user.UseDays, user.ExpiryDate))
|
||||
}
|
||||
writer.WriteString("\n")
|
||||
writer.Flush()
|
||||
return nil
|
||||
}
|
||||
|
||||
// ExecSql 执行sql
|
||||
func (mysql *Mysql) ExecSql(filePath string) error {
|
||||
db := mysql.GetDB()
|
||||
fileByte, err := os.ReadFile(filePath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
sqlStr := string(fileByte)
|
||||
sqls := strings.Split(strings.Replace(sqlStr, "\r\n", "\n", -1), ";\n")
|
||||
for _, s := range sqls {
|
||||
s = strings.TrimSpace(s)
|
||||
if s != "" {
|
||||
if _, err = db.Exec(s); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
Reference in New Issue
Block a user