This commit is contained in:
chermack
2026-07-26 00:09:57 +08:00
commit 059be96536
134 changed files with 19678 additions and 0 deletions
+33
View File
@@ -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"`
}
+50
View File
@@ -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
}
+41
View File
@@ -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)
}
+412
View File
@@ -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, &quota, &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, &quota, &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
}
+122
View File
@@ -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, "")
}
+113
View File
@@ -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
}