Initial commit: TaskPool React panel
- React frontend with route-level code splitting - Backend rebranded from Baihu to TaskPool - DB brand migration script and local compatibility
This commit is contained in:
@@ -0,0 +1,225 @@
|
||||
package utils
|
||||
|
||||
import (
|
||||
"archive/tar"
|
||||
"archive/zip"
|
||||
"compress/gzip"
|
||||
"io"
|
||||
"io/fs"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func CreateZip(dst io.Writer, basePaths []string) (err error) {
|
||||
w := zip.NewWriter(dst)
|
||||
defer func() {
|
||||
if closeErr := w.Close(); err == nil {
|
||||
err = closeErr
|
||||
}
|
||||
}()
|
||||
|
||||
for _, basePath := range basePaths {
|
||||
info, err := os.Lstat(basePath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if info.Mode()&os.ModeSymlink != 0 {
|
||||
continue
|
||||
}
|
||||
|
||||
if info.Mode().IsRegular() {
|
||||
if err := addZipFile(w, basePath, filepath.Base(basePath), info); err != nil {
|
||||
return err
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
if info.IsDir() {
|
||||
baseName := filepath.Base(basePath)
|
||||
if err := filepath.WalkDir(basePath, func(path string, d fs.DirEntry, err error) error {
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if d.Type()&fs.ModeSymlink != 0 {
|
||||
return nil
|
||||
}
|
||||
|
||||
rel, err := filepath.Rel(basePath, path)
|
||||
if err != nil || strings.HasPrefix(rel, ".."+string(filepath.Separator)) || rel == ".." {
|
||||
return nil
|
||||
}
|
||||
|
||||
name := baseName
|
||||
if rel != "." {
|
||||
name = filepath.Join(baseName, rel)
|
||||
}
|
||||
name = filepath.ToSlash(name)
|
||||
|
||||
info, err := d.Info()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if d.IsDir() {
|
||||
header, err := zip.FileInfoHeader(info)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
header.Name = name + "/"
|
||||
_, err = w.CreateHeader(header)
|
||||
return err
|
||||
}
|
||||
|
||||
if info.Mode().IsRegular() {
|
||||
return addZipFile(w, path, name, info)
|
||||
}
|
||||
return nil
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func addZipFile(w *zip.Writer, path, name string, info os.FileInfo) error {
|
||||
header, err := zip.FileInfoHeader(info)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
header.Name = filepath.ToSlash(name)
|
||||
header.Method = zip.Deflate
|
||||
|
||||
writer, err := w.CreateHeader(header)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
file, err := os.Open(path)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer file.Close()
|
||||
|
||||
_, err = io.Copy(writer, file)
|
||||
return err
|
||||
}
|
||||
|
||||
func ExtractZip(src, dest string) error {
|
||||
r, err := zip.OpenReader(src)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer r.Close()
|
||||
|
||||
for _, f := range r.File {
|
||||
fpath := filepath.Join(dest, f.Name)
|
||||
|
||||
// 安全检查:防止路径遍历 (ZipSlip)
|
||||
rel, err := filepath.Rel(dest, fpath)
|
||||
if err != nil || strings.HasPrefix(rel, ".."+string(filepath.Separator)) || rel == ".." {
|
||||
continue
|
||||
}
|
||||
|
||||
if f.FileInfo().IsDir() {
|
||||
os.MkdirAll(fpath, 0755)
|
||||
continue
|
||||
}
|
||||
|
||||
if err := os.MkdirAll(filepath.Dir(fpath), 0755); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
outFile, err := os.OpenFile(fpath, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, f.Mode())
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
rc, err := f.Open()
|
||||
if err != nil {
|
||||
outFile.Close()
|
||||
return err
|
||||
}
|
||||
|
||||
_, err = io.Copy(outFile, rc)
|
||||
outFile.Close()
|
||||
rc.Close()
|
||||
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func ExtractTar(src, dest string) error {
|
||||
file, err := os.Open(src)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer file.Close()
|
||||
|
||||
return extractTarReader(tar.NewReader(file), dest)
|
||||
}
|
||||
|
||||
func ExtractTarGz(src, dest string) error {
|
||||
file, err := os.Open(src)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer file.Close()
|
||||
|
||||
gzr, err := gzip.NewReader(file)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer gzr.Close()
|
||||
|
||||
return extractTarReader(tar.NewReader(gzr), dest)
|
||||
}
|
||||
|
||||
func extractTarReader(tr *tar.Reader, dest string) error {
|
||||
for {
|
||||
header, err := tr.Next()
|
||||
if err == io.EOF {
|
||||
break
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
fpath := filepath.Join(dest, header.Name)
|
||||
|
||||
// 安全检查:防止路径遍历
|
||||
rel, err := filepath.Rel(dest, fpath)
|
||||
if err != nil || strings.HasPrefix(rel, ".."+string(filepath.Separator)) || rel == ".." {
|
||||
continue
|
||||
}
|
||||
|
||||
switch header.Typeflag {
|
||||
case tar.TypeDir:
|
||||
os.MkdirAll(fpath, 0755)
|
||||
case tar.TypeReg:
|
||||
if err := os.MkdirAll(filepath.Dir(fpath), 0755); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
outFile, err := os.Create(fpath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if _, err := io.Copy(outFile, tr); err != nil {
|
||||
outFile.Close()
|
||||
return err
|
||||
}
|
||||
outFile.Close()
|
||||
|
||||
os.Chmod(fpath, os.FileMode(header.Mode))
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,79 @@
|
||||
package utils
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sync"
|
||||
)
|
||||
|
||||
// TailBuffer 是一个只保留最后 N 字节数据的缓冲区
|
||||
type TailBuffer struct {
|
||||
mu sync.Mutex
|
||||
limit int
|
||||
data []byte
|
||||
size int // 当前实际存储的大小
|
||||
pos int // 下一个写入位置 (针对环形缓冲区逻辑,但这里为了简单使用切片重排)
|
||||
}
|
||||
|
||||
// NewTailBuffer 创建一个限制大小为 limit 的尾部缓冲区
|
||||
func NewTailBuffer(limit int) *TailBuffer {
|
||||
return &TailBuffer{
|
||||
limit: limit,
|
||||
data: make([]byte, 0, limit),
|
||||
}
|
||||
}
|
||||
|
||||
// Write 实现 io.Writer 接口
|
||||
func (b *TailBuffer) Write(p []byte) (n int, err error) {
|
||||
b.mu.Lock()
|
||||
defer b.mu.Unlock()
|
||||
|
||||
n = len(p)
|
||||
if n >= b.limit {
|
||||
// 如果单次写入就超过了限制,直接取最后 limit 字节
|
||||
b.data = append(b.data[:0], p[n-b.limit:]...)
|
||||
return
|
||||
}
|
||||
|
||||
available := b.limit - len(b.data)
|
||||
if n <= available {
|
||||
// 空间足够,直接追加
|
||||
b.data = append(b.data, p...)
|
||||
} else {
|
||||
// 空间不足,需要移除旧数据
|
||||
toRemove := n - available
|
||||
b.data = append(b.data[toRemove:], p...)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
// Bytes 返回缓冲区内的所有数据
|
||||
func (b *TailBuffer) Bytes() []byte {
|
||||
b.mu.Lock()
|
||||
defer b.mu.Unlock()
|
||||
res := make([]byte, len(b.data))
|
||||
copy(res, b.data)
|
||||
return res
|
||||
}
|
||||
|
||||
// String 返回缓冲区内的字符串表示
|
||||
func (b *TailBuffer) String() string {
|
||||
b.mu.Lock()
|
||||
defer b.mu.Unlock()
|
||||
return string(b.data)
|
||||
}
|
||||
|
||||
// Len 返回当前存储的数据长度
|
||||
func (b *TailBuffer) Len() int {
|
||||
b.mu.Lock()
|
||||
defer b.mu.Unlock()
|
||||
return len(b.data)
|
||||
}
|
||||
|
||||
// TrimLog 裁剪日志,保留末尾指定大小
|
||||
func TrimLog(content string, limit int) string {
|
||||
if len(content) <= limit {
|
||||
return content
|
||||
}
|
||||
// 简单裁剪,不考虑字符完整性,因为这是针对大文本的保护
|
||||
return fmt.Sprintf("\n\n[System] 日志过长,已自动截断,仅保留末尾 %d MB...\n\n", limit/1024/1024) + content[len(content)-limit:]
|
||||
}
|
||||
@@ -0,0 +1,134 @@
|
||||
package utils
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"compress/zlib"
|
||||
"encoding/base64"
|
||||
"io"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/klauspost/compress/zstd"
|
||||
)
|
||||
|
||||
const (
|
||||
zstdPrefix = "zstd:"
|
||||
rawPrefix = "raw:"
|
||||
MinCompressSize = 128
|
||||
)
|
||||
|
||||
var zlibWriterPool = sync.Pool{
|
||||
New: func() interface{} {
|
||||
return zlib.NewWriter(io.Discard)
|
||||
},
|
||||
}
|
||||
|
||||
// GetZlibWriter 从对象池中获取 zlib 写入器
|
||||
func GetZlibWriter(w io.Writer) *zlib.Writer {
|
||||
zw := zlibWriterPool.Get().(*zlib.Writer)
|
||||
zw.Reset(w)
|
||||
return zw
|
||||
}
|
||||
|
||||
// PutZlibWriter 将 zlib 写入器还回对象池
|
||||
func PutZlibWriter(zw *zlib.Writer) {
|
||||
zlibWriterPool.Put(zw)
|
||||
}
|
||||
|
||||
var zstdEncoderPool = sync.Pool{
|
||||
New: func() interface{} {
|
||||
zw, _ := zstd.NewWriter(nil)
|
||||
return zw
|
||||
},
|
||||
}
|
||||
|
||||
// GetZstdWriter 从对象池中获取 zstd 写入器并定向到 w
|
||||
func GetZstdWriter(w io.Writer) *zstd.Encoder {
|
||||
zw := zstdEncoderPool.Get().(*zstd.Encoder)
|
||||
zw.Reset(w)
|
||||
return zw
|
||||
}
|
||||
|
||||
// PutZstdWriter 将 zstd 写入器还回对象池
|
||||
func PutZstdWriter(zw *zstd.Encoder) {
|
||||
zstdEncoderPool.Put(zw)
|
||||
}
|
||||
|
||||
var (
|
||||
zstdEncoder *zstd.Encoder
|
||||
zstdDecoder *zstd.Decoder
|
||||
zstdOnce sync.Once
|
||||
)
|
||||
|
||||
func initZstd() {
|
||||
zstdOnce.Do(func() {
|
||||
var err error
|
||||
// 默认级别适合常规压缩
|
||||
zstdEncoder, err = zstd.NewWriter(nil)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
zstdDecoder, err = zstd.NewReader(nil)
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// CompressToBase64 compresses data using zstd and encodes to base64 with a prefix (uses raw for data under MinCompressSize bytes)
|
||||
func CompressToBase64(data string) (string, error) {
|
||||
if data == "" {
|
||||
return "", nil
|
||||
}
|
||||
|
||||
// 小于等于阈值,不需要压缩,直接前缀明文保存
|
||||
if len(data) <= MinCompressSize {
|
||||
return rawPrefix + data, nil
|
||||
}
|
||||
|
||||
initZstd()
|
||||
compressed := zstdEncoder.EncodeAll([]byte(data), nil)
|
||||
return zstdPrefix + base64.StdEncoding.EncodeToString(compressed), nil
|
||||
}
|
||||
|
||||
// DecompressFromBase64 decodes base64 and decompresses data (supports raw/zstd prefix and falls back to zlib)
|
||||
func DecompressFromBase64(data string) (string, error) {
|
||||
if data == "" {
|
||||
return "", nil
|
||||
}
|
||||
|
||||
if strings.HasPrefix(data, rawPrefix) {
|
||||
return data[len(rawPrefix):], nil
|
||||
}
|
||||
|
||||
if strings.HasPrefix(data, zstdPrefix) {
|
||||
initZstd()
|
||||
encoded := data[len(zstdPrefix):]
|
||||
decoded, err := base64.StdEncoding.DecodeString(encoded)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
decompressed, err := zstdDecoder.DecodeAll(decoded, nil)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return string(decompressed), nil
|
||||
}
|
||||
|
||||
// Fallback to legacy zlib
|
||||
decoded, err := base64.StdEncoding.DecodeString(data)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
zr, err := zlib.NewReader(bytes.NewReader(decoded))
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer zr.Close()
|
||||
result, err := io.ReadAll(zr)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return string(result), nil
|
||||
}
|
||||
|
||||
@@ -0,0 +1,121 @@
|
||||
package utils
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"compress/zlib"
|
||||
"encoding/base64"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// legacyZlibCompress 以前的 zlib 压缩逻辑,用于构造测试样本
|
||||
func legacyZlibCompress(data string) (string, error) {
|
||||
if data == "" {
|
||||
return "", nil
|
||||
}
|
||||
var buf bytes.Buffer
|
||||
zw := zlib.NewWriter(&buf)
|
||||
if _, err := zw.Write([]byte(data)); err != nil {
|
||||
return "", err
|
||||
}
|
||||
if err := zw.Close(); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return base64.StdEncoding.EncodeToString(buf.Bytes()), nil
|
||||
}
|
||||
|
||||
func TestCompressAndDecompressZstd(t *testing.T) {
|
||||
originalText := "Hello, this is a test log message for ZSTD compression in TaskPool! Repeat: Hello, this is a test log message for ZSTD compression in TaskPool!"
|
||||
|
||||
// 1. 测试 ZSTD 压缩
|
||||
compressed, err := CompressToBase64(originalText)
|
||||
if err != nil {
|
||||
t.Fatalf("CompressToBase64 failed: %v", err)
|
||||
}
|
||||
|
||||
// 验证前缀是否正确
|
||||
if !strings.HasPrefix(compressed, "zstd:") {
|
||||
t.Errorf("Expected compressed output to have 'zstd:' prefix, got: %s", compressed)
|
||||
}
|
||||
|
||||
// 2. 测试 ZSTD 解密解压
|
||||
decompressed, err := DecompressFromBase64(compressed)
|
||||
if err != nil {
|
||||
t.Fatalf("DecompressFromBase64 failed: %v", err)
|
||||
}
|
||||
|
||||
if decompressed != originalText {
|
||||
t.Errorf("Decompressed text mismatch.\nExpected: %s\nGot: %s", originalText, decompressed)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecompressLegacyZlibCompatibility(t *testing.T) {
|
||||
originalText := "This is a legacy log message compressed using zlib. It should be decompressed successfully."
|
||||
|
||||
// 1. 用老逻辑压缩生成旧数据
|
||||
legacyCompressed, err := legacyZlibCompress(originalText)
|
||||
if err != nil {
|
||||
t.Fatalf("legacyZlibCompress failed: %v", err)
|
||||
}
|
||||
|
||||
// 验证没有 zstd 前缀
|
||||
if strings.HasPrefix(legacyCompressed, "zstd:") {
|
||||
t.Fatalf("Legacy compressed string shouldn't have 'zstd:' prefix")
|
||||
}
|
||||
|
||||
// 2. 使用新版的 DecompressFromBase64 解压,验证其对旧格式的兼容性
|
||||
decompressed, err := DecompressFromBase64(legacyCompressed)
|
||||
if err != nil {
|
||||
t.Fatalf("DecompressFromBase64 failed to decompress legacy data: %v", err)
|
||||
}
|
||||
|
||||
if decompressed != originalText {
|
||||
t.Errorf("Decompressed legacy text mismatch.\nExpected: %s\nGot: %s", originalText, decompressed)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEmptyString(t *testing.T) {
|
||||
compressed, err := CompressToBase64("")
|
||||
if err != nil {
|
||||
t.Fatalf("CompressToBase64 for empty string failed: %v", err)
|
||||
}
|
||||
if compressed != "" {
|
||||
t.Errorf("Expected empty string for empty input, got: %q", compressed)
|
||||
}
|
||||
|
||||
decompressed, err := DecompressFromBase64("")
|
||||
if err != nil {
|
||||
t.Fatalf("DecompressFromBase64 for empty string failed: %v", err)
|
||||
}
|
||||
if decompressed != "" {
|
||||
t.Errorf("Expected empty string for empty input, got: %q", decompressed)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCompressShortTextRaw(t *testing.T) {
|
||||
shortText := "python secret.py" // 16 bytes
|
||||
|
||||
// 1. 压缩(应转为 raw 前缀)
|
||||
output, err := CompressToBase64(shortText)
|
||||
if err != nil {
|
||||
t.Fatalf("CompressToBase64 short text failed: %v", err)
|
||||
}
|
||||
|
||||
if !strings.HasPrefix(output, "raw:") {
|
||||
t.Errorf("Expected short text output to have 'raw:' prefix, got: %q", output)
|
||||
}
|
||||
if output != "raw:python secret.py" {
|
||||
t.Errorf("Expected output raw:python secret.py, got: %q", output)
|
||||
}
|
||||
|
||||
// 2. 解码解密
|
||||
decompressed, err := DecompressFromBase64(output)
|
||||
if err != nil {
|
||||
t.Fatalf("DecompressFromBase64 short text failed: %v", err)
|
||||
}
|
||||
|
||||
if decompressed != shortText {
|
||||
t.Errorf("Decompressed short text mismatch.\nExpected: %q\nGot: %q", shortText, decompressed)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,134 @@
|
||||
package utils
|
||||
|
||||
import (
|
||||
"crypto/aes"
|
||||
"crypto/cipher"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"errors"
|
||||
"io"
|
||||
"os"
|
||||
"strings"
|
||||
)
|
||||
|
||||
var (
|
||||
masterSecretKey []byte
|
||||
ErrKeyNotSet = errors.New("加密秘钥未配置,请按照文档使用 TASKPOOL_SECRET_KEY 环境变量启动服务配置秘钥")
|
||||
)
|
||||
|
||||
// InitSecretKey initialized the master secret key from the environment and unsets it
|
||||
func InitSecretKey() {
|
||||
// 优先新变量;兼容旧 BAIHU_SECRET_KEY
|
||||
key := os.Getenv("TASKPOOL_SECRET_KEY")
|
||||
if key == "" {
|
||||
key = os.Getenv("BAIHU_SECRET_KEY")
|
||||
}
|
||||
if key != "" {
|
||||
hash := sha256.Sum256([]byte(key))
|
||||
masterSecretKey = hash[:]
|
||||
// Ensure it's only in memory by unsetting the environment variable
|
||||
os.Unsetenv("TASKPOOL_SECRET_KEY")
|
||||
os.Unsetenv("BAIHU_SECRET_KEY")
|
||||
}
|
||||
}
|
||||
|
||||
// IsSecretKeySet returns true if the master secret key is configured
|
||||
func IsSecretKeySet() bool {
|
||||
return len(masterSecretKey) > 0
|
||||
}
|
||||
|
||||
// Encrypt encrypts a plaintext string using AES-GCM
|
||||
func Encrypt(plaintext string) (string, error) {
|
||||
if plaintext == "" {
|
||||
return "", nil
|
||||
}
|
||||
if !IsSecretKeySet() {
|
||||
return "", ErrKeyNotSet
|
||||
}
|
||||
|
||||
block, err := aes.NewCipher(masterSecretKey)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
aesGCM, err := cipher.NewGCM(block)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
nonce := make([]byte, aesGCM.NonceSize())
|
||||
if _, err = io.ReadFull(rand.Reader, nonce); err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
ciphertext := aesGCM.Seal(nonce, nonce, []byte(plaintext), nil)
|
||||
return base64.StdEncoding.EncodeToString(ciphertext), nil
|
||||
}
|
||||
|
||||
// Decrypt decrypts a ciphertext string using AES-GCM
|
||||
// Returns the original string if decryption fails or if it wasn't encrypted
|
||||
func Decrypt(ciphertext string) (string, error) {
|
||||
if ciphertext == "" {
|
||||
return "", nil
|
||||
}
|
||||
if !IsSecretKeySet() {
|
||||
return ciphertext, ErrKeyNotSet
|
||||
}
|
||||
|
||||
data, err := base64.StdEncoding.DecodeString(ciphertext)
|
||||
if err != nil {
|
||||
return ciphertext, err
|
||||
}
|
||||
|
||||
block, err := aes.NewCipher(masterSecretKey)
|
||||
if err != nil {
|
||||
return ciphertext, err
|
||||
}
|
||||
|
||||
aesGCM, err := cipher.NewGCM(block)
|
||||
if err != nil {
|
||||
return ciphertext, err
|
||||
}
|
||||
|
||||
nonceSize := aesGCM.NonceSize()
|
||||
if len(data) < nonceSize {
|
||||
return ciphertext, errors.New("ciphertext too short")
|
||||
}
|
||||
|
||||
nonce, ciphertextBytes := data[:nonceSize], data[nonceSize:]
|
||||
plaintext, err := aesGCM.Open(nil, nonce, ciphertextBytes, nil)
|
||||
if err != nil {
|
||||
return ciphertext, err
|
||||
}
|
||||
|
||||
return string(plaintext), nil
|
||||
}
|
||||
|
||||
// MaskSecrets 将文本中的所有敏感机密值替换为脱敏字符串 "********"
|
||||
func MaskSecrets(text string, secrets []string) string {
|
||||
if len(secrets) == 0 || text == "" {
|
||||
return text
|
||||
}
|
||||
for _, mask := range secrets {
|
||||
if mask != "" {
|
||||
text = strings.ReplaceAll(text, mask, "********")
|
||||
}
|
||||
}
|
||||
return text
|
||||
}
|
||||
|
||||
// MaskString 对字符串进行脱敏处理,保留首尾,中间用星号遮掩
|
||||
func MaskString(s string) string {
|
||||
if s == "" {
|
||||
return ""
|
||||
}
|
||||
n := len(s)
|
||||
if n <= 3 {
|
||||
return "***"
|
||||
}
|
||||
if n <= 6 {
|
||||
return s[:1] + "***" + s[n-1:]
|
||||
}
|
||||
return s[:2] + "****" + s[n-2:]
|
||||
}
|
||||
@@ -0,0 +1,92 @@
|
||||
package utils
|
||||
|
||||
import (
|
||||
"bufio"
|
||||
"io"
|
||||
"unicode/utf8"
|
||||
|
||||
"golang.org/x/text/encoding/simplifiedchinese"
|
||||
"golang.org/x/text/transform"
|
||||
)
|
||||
|
||||
// ToUTF8 converts potentially non-UTF8 data (like GBK on Windows) to UTF-8
|
||||
func ToUTF8(data []byte) string {
|
||||
if utf8.Valid(data) {
|
||||
return string(data)
|
||||
}
|
||||
// Try GBK (common on Windows)
|
||||
reader := transform.NewReader(
|
||||
bufio.NewReader(
|
||||
&byteReader{data: data},
|
||||
),
|
||||
simplifiedchinese.GBK.NewDecoder(),
|
||||
)
|
||||
result, err := io.ReadAll(reader)
|
||||
if err != nil {
|
||||
return string(data)
|
||||
}
|
||||
return string(result)
|
||||
}
|
||||
|
||||
type byteReader struct {
|
||||
data []byte
|
||||
pos int
|
||||
}
|
||||
|
||||
func (r *byteReader) Read(p []byte) (n int, err error) {
|
||||
if r.pos >= len(r.data) {
|
||||
return 0, io.EOF
|
||||
}
|
||||
n = copy(p, r.data[r.pos:])
|
||||
r.pos += n
|
||||
return n, nil
|
||||
}
|
||||
|
||||
// TrimLastRunes 从字符串尾部保留最多 maxRunes 个字符(不仅限于 ASCII,支持中英文混排的真实字符数量)
|
||||
func TrimLastRunes(s string, maxRunes int) string {
|
||||
// 如果字符串的总字节数小于等于 maxRunes,那么它的字符数一定也小于等于 maxRunes
|
||||
if len(s) <= maxRunes {
|
||||
return s
|
||||
}
|
||||
|
||||
count := 0
|
||||
for i := len(s); i > 0; {
|
||||
_, size := utf8.DecodeLastRuneInString(s[:i])
|
||||
i -= size
|
||||
count++
|
||||
if count == maxRunes {
|
||||
return s[i:]
|
||||
}
|
||||
}
|
||||
return s
|
||||
}
|
||||
|
||||
// IntToStr 将整数转换为字符串
|
||||
func IntToStr(n int) string {
|
||||
if n == 0 {
|
||||
return "0"
|
||||
}
|
||||
|
||||
var negative bool
|
||||
if n < 0 {
|
||||
negative = true
|
||||
n = -n
|
||||
}
|
||||
|
||||
var digits []byte
|
||||
for n > 0 {
|
||||
digits = append(digits, byte('0'+n%10))
|
||||
n /= 10
|
||||
}
|
||||
|
||||
if negative {
|
||||
digits = append(digits, '-')
|
||||
}
|
||||
|
||||
// Reverse
|
||||
for i, j := 0, len(digits)-1; i < j; i, j = i+1, j-1 {
|
||||
digits[i], digits[j] = digits[j], digits[i]
|
||||
}
|
||||
|
||||
return string(digits)
|
||||
}
|
||||
@@ -0,0 +1,85 @@
|
||||
package utils
|
||||
|
||||
import (
|
||||
"io"
|
||||
"os"
|
||||
"path/filepath"
|
||||
)
|
||||
|
||||
// CopyPath copies a file or directory from src to dest
|
||||
func CopyPath(src, dest string) error {
|
||||
info, err := os.Stat(src)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if info.IsDir() {
|
||||
return copyDir(src, dest)
|
||||
}
|
||||
return CopyFile(src, dest)
|
||||
}
|
||||
|
||||
// CopyFile copies a single file from src to dest
|
||||
func CopyFile(src, dest string) error {
|
||||
srcFile, err := os.Open(src)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer srcFile.Close()
|
||||
|
||||
if err := os.MkdirAll(filepath.Dir(dest), 0755); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
destFile, err := os.Create(dest)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer destFile.Close()
|
||||
|
||||
if _, err := io.Copy(destFile, srcFile); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
info, err := os.Stat(src)
|
||||
if err == nil {
|
||||
os.Chmod(dest, info.Mode())
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func copyDir(src, dest string) error {
|
||||
info, err := os.Stat(src)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if err := os.MkdirAll(dest, info.Mode()); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
entries, err := os.ReadDir(src)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
for _, entry := range entries {
|
||||
srcPath := filepath.Join(src, entry.Name())
|
||||
destPath := filepath.Join(dest, entry.Name())
|
||||
|
||||
if err := CopyPath(srcPath, destPath); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// IsInDocker 判断程序是否运行在 Docker 容器中
|
||||
func IsInDocker() bool {
|
||||
if _, err := os.Stat("/.dockerenv"); err == nil {
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
@@ -0,0 +1,20 @@
|
||||
package utils
|
||||
|
||||
import (
|
||||
"runtime"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// GetGoroutineID 获取当前 Goroutine ID
|
||||
// 注意:这只是为了调试和日志目的,不应该用于业务逻辑
|
||||
func GetGoroutineID() int64 {
|
||||
var buf [64]byte
|
||||
n := runtime.Stack(buf[:], false)
|
||||
idField := strings.Fields(strings.TrimPrefix(string(buf[:n]), "goroutine "))[0]
|
||||
id, err := strconv.ParseInt(idField, 10, 64)
|
||||
if err != nil {
|
||||
return -1
|
||||
}
|
||||
return id
|
||||
}
|
||||
@@ -0,0 +1,20 @@
|
||||
package utils
|
||||
|
||||
import (
|
||||
"github.com/rs/xid"
|
||||
)
|
||||
|
||||
// GenerateID 生成一个新的 ID (使用 xid,20位字符)
|
||||
func GenerateID() string {
|
||||
return xid.New().String()
|
||||
}
|
||||
|
||||
// IsNumeric 检查字符串是否全为数字
|
||||
func IsNumeric(s string) bool {
|
||||
for _, c := range s {
|
||||
if c < '0' || c > '9' {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return s != ""
|
||||
}
|
||||
@@ -0,0 +1,51 @@
|
||||
package utils
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"net"
|
||||
"os"
|
||||
"runtime"
|
||||
"sort"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// GenerateMachineID 生成机器识别码
|
||||
func GenerateMachineID() string {
|
||||
var parts []string
|
||||
|
||||
// 主机名
|
||||
if hostname, err := os.Hostname(); err == nil {
|
||||
parts = append(parts, hostname)
|
||||
}
|
||||
|
||||
// 获取所有非回环网卡的 MAC 地址,排序后取第一个(最稳定)
|
||||
if interfaces, err := net.Interfaces(); err == nil {
|
||||
var macs []string
|
||||
for _, iface := range interfaces {
|
||||
// 跳过回环接口、没有 MAC 地址的接口、虚拟接口
|
||||
if iface.Flags&net.FlagLoopback != 0 || len(iface.HardwareAddr) == 0 {
|
||||
continue
|
||||
}
|
||||
// 跳过 docker/veth 等虚拟网卡
|
||||
name := strings.ToLower(iface.Name)
|
||||
if strings.HasPrefix(name, "docker") || strings.HasPrefix(name, "veth") ||
|
||||
strings.HasPrefix(name, "br-") || strings.HasPrefix(name, "virbr") {
|
||||
continue
|
||||
}
|
||||
macs = append(macs, iface.HardwareAddr.String())
|
||||
}
|
||||
sort.Strings(macs)
|
||||
// 只使用第一个 MAC 地址(最稳定)
|
||||
if len(macs) > 0 {
|
||||
parts = append(parts, macs[0])
|
||||
}
|
||||
}
|
||||
|
||||
// 操作系统和架构
|
||||
parts = append(parts, runtime.GOOS, runtime.GOARCH)
|
||||
|
||||
data := strings.Join(parts, "|")
|
||||
hash := sha256.Sum256([]byte(data))
|
||||
return hex.EncodeToString(hash[:])
|
||||
}
|
||||
@@ -0,0 +1,167 @@
|
||||
package utils
|
||||
|
||||
import (
|
||||
"os/exec"
|
||||
"strings"
|
||||
"sync"
|
||||
)
|
||||
|
||||
var nodePathCache sync.Map
|
||||
|
||||
// GetMiseNodePath 获取指定版本的 node 全局包路径,使用内存缓存避免重复获取
|
||||
func GetMiseNodePath(version string) string {
|
||||
if version == "" {
|
||||
version = "latest"
|
||||
}
|
||||
|
||||
if val, ok := nodePathCache.Load(version); ok {
|
||||
return val.(string)
|
||||
}
|
||||
|
||||
cmd := exec.Command("mise", "where", "node@"+version)
|
||||
out, err := cmd.CombinedOutput()
|
||||
if err == nil {
|
||||
nodeDir := strings.TrimSpace(string(out))
|
||||
if nodeDir != "" {
|
||||
// 采用双路径策略:lib/node_modules 是标准路径,lib 是某些环境(如 mise Docker)下的特殊路径
|
||||
// 通过冒号分隔,让 Node.js 按顺序搜索,保证最大兼容性
|
||||
nodePath := nodeDir + "/lib/node_modules:" + nodeDir + "/lib"
|
||||
nodePathCache.Store(version, nodePath)
|
||||
return nodePath
|
||||
}
|
||||
}
|
||||
|
||||
return ""
|
||||
}
|
||||
|
||||
// InjectNodePath 检查语言环境中是否有 node,如果有则自动获取并注入 NODE_PATH 到环境变量切片中
|
||||
func InjectNodePath(envs *[]string, languages []map[string]string) {
|
||||
for _, lang := range languages {
|
||||
if lang["name"] == "node" {
|
||||
if nodePath := GetMiseNodePath(lang["version"]); nodePath != "" {
|
||||
*envs = append(*envs, "NODE_PATH="+nodePath)
|
||||
}
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// BuildMiseCommand 构建多语言 mise 执行命令 (字符串形式)
|
||||
func BuildMiseCommand(command string, languages []map[string]string) string {
|
||||
if len(languages) == 0 {
|
||||
return command
|
||||
}
|
||||
|
||||
var builder strings.Builder
|
||||
builder.WriteString("mise exec")
|
||||
|
||||
for _, lang := range languages {
|
||||
name := lang["name"]
|
||||
version := lang["version"]
|
||||
if name == "" {
|
||||
continue
|
||||
}
|
||||
if version == "" {
|
||||
version = "latest"
|
||||
}
|
||||
builder.WriteString(" " + name + "@" + version)
|
||||
}
|
||||
|
||||
builder.WriteString(" -- " + command)
|
||||
return builder.String()
|
||||
}
|
||||
|
||||
// BuildMiseCommandArgs 构建多语言 mise 执行命令 (参数列表形式)
|
||||
func BuildMiseCommandArgs(cmdArgs []string, languages []map[string]string) []string {
|
||||
if len(languages) == 0 {
|
||||
return cmdArgs
|
||||
}
|
||||
|
||||
args := []string{"mise", "exec"}
|
||||
for _, lang := range languages {
|
||||
name := lang["name"]
|
||||
version := lang["version"]
|
||||
if name == "" {
|
||||
continue
|
||||
}
|
||||
if version == "" {
|
||||
version = "latest"
|
||||
}
|
||||
args = append(args, name+"@"+version)
|
||||
}
|
||||
args = append(args, "--")
|
||||
args = append(args, cmdArgs...)
|
||||
return args
|
||||
}
|
||||
|
||||
// BuildMiseCommandSimple 构建单个语言的 mise 执行命令
|
||||
func BuildMiseCommandSimple(command string, language, version string) string {
|
||||
if language == "" {
|
||||
return command
|
||||
}
|
||||
spec := language
|
||||
if version != "" {
|
||||
spec += "@" + version
|
||||
}
|
||||
return "mise exec " + spec + " -- " + command
|
||||
}
|
||||
|
||||
// BuildMiseCommandArgsSimple 构建单个语言的 mise 执行命令 (参数列表形式)
|
||||
func BuildMiseCommandArgsSimple(cmdArgs []string, language, version string) []string {
|
||||
if language == "" {
|
||||
return cmdArgs
|
||||
}
|
||||
spec := language
|
||||
if version != "" {
|
||||
spec += "@" + version
|
||||
}
|
||||
return append([]string{"mise", "exec", spec, "--"}, cmdArgs...)
|
||||
}
|
||||
|
||||
// ListMiseInstalledVersions 获取指定语言已安装的所有版本列表
|
||||
func ListMiseInstalledVersions(language string) ([]string, error) {
|
||||
// 执行 mise ls <language> 命令
|
||||
cmd := exec.Command("mise", "ls", language)
|
||||
out, err := cmd.CombinedOutput()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var versions []string
|
||||
lines := strings.Split(string(out), "\n")
|
||||
for _, line := range lines {
|
||||
v := strings.TrimSpace(line)
|
||||
if v == "" {
|
||||
continue
|
||||
}
|
||||
// mise ls 的输出可能包含状态标识或插件名,例:
|
||||
// * 20.10.0 (active)
|
||||
// node 18.17.0
|
||||
fields := strings.Fields(v)
|
||||
if len(fields) == 0 {
|
||||
continue
|
||||
}
|
||||
|
||||
startIdx := 0
|
||||
// 跳过状态标识符
|
||||
if fields[startIdx] == "*" || fields[startIdx] == "->" || fields[startIdx] == ">" {
|
||||
startIdx++
|
||||
}
|
||||
|
||||
if len(fields) <= startIdx {
|
||||
continue
|
||||
}
|
||||
|
||||
vstr := fields[startIdx]
|
||||
// 如果第一个有效字段是插件名,则版本号在第二个字段
|
||||
if vstr == language && len(fields) > startIdx+1 {
|
||||
vstr = fields[startIdx+1]
|
||||
}
|
||||
|
||||
// 确保解析出来的不是插件名
|
||||
if vstr != "" && vstr != language {
|
||||
versions = append(versions, vstr)
|
||||
}
|
||||
}
|
||||
return versions, nil
|
||||
}
|
||||
@@ -0,0 +1,72 @@
|
||||
package utils
|
||||
|
||||
import (
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"os"
|
||||
"strings"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// CheckWSOrigin 校验 WebSocket 的 Origin 来源是否安全。
|
||||
// 默认仅允许同源请求,可通过环境变量 BH_ALLOWED_ORIGINS 配置额外的允许列表(逗号分隔)。
|
||||
func CheckWSOrigin(r *http.Request) bool {
|
||||
origin := r.Header.Get("Origin")
|
||||
if origin == "" {
|
||||
// 非浏览器发起的请求(如直接用脚本连接)通常不带 Origin,默认放行。
|
||||
return true
|
||||
}
|
||||
|
||||
u, err := url.Parse(origin)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
|
||||
// 0. 开发环境校验:如果是非 Release 模式,默认放行
|
||||
if gin.Mode() != gin.ReleaseMode {
|
||||
return true
|
||||
}
|
||||
|
||||
// 1. 同源校验:Origin 的 Host 与请求头中的 Host 一致
|
||||
if strings.EqualFold(u.Host, r.Host) {
|
||||
return true
|
||||
}
|
||||
|
||||
// 容错处理:如果因为 Nginx 配置了 $host 而丢掉了端口,导致一方带端口一方不带端口时,尝试忽略端口进行域名比对
|
||||
uHostOnly := u.Hostname()
|
||||
rHostOnly := r.Host
|
||||
if h, _, err := net.SplitHostPort(r.Host); err == nil {
|
||||
rHostOnly = h
|
||||
}
|
||||
if strings.EqualFold(uHostOnly, rHostOnly) {
|
||||
return true
|
||||
}
|
||||
|
||||
// 2. 环境变量配置的允许列表校验
|
||||
allowedOrigins := os.Getenv("BH_ALLOWED_ORIGINS")
|
||||
if allowedOrigins != "" {
|
||||
origins := strings.Split(allowedOrigins, ",")
|
||||
for _, o := range origins {
|
||||
o = strings.TrimSpace(o)
|
||||
if o == "*" {
|
||||
return true
|
||||
}
|
||||
// 匹配完整 Origin (如 http://localhost:5173) 或仅 Host 部分
|
||||
if strings.EqualFold(o, origin) || strings.EqualFold(o, u.Host) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 3. 允许来自 localhost、127.0.0.1 以及局域网内网 IP 的请求 (方便本地开发、虚拟机调试和局域网部署)
|
||||
hostname := u.Hostname()
|
||||
if strings.HasPrefix(hostname, "localhost") || strings.HasPrefix(hostname, "127.0.0.1") ||
|
||||
strings.HasPrefix(hostname, "192.168.") || strings.HasPrefix(hostname, "10.") ||
|
||||
strings.HasPrefix(hostname, "172.") {
|
||||
return true
|
||||
}
|
||||
|
||||
return false
|
||||
}
|
||||
@@ -0,0 +1,79 @@
|
||||
package utils
|
||||
|
||||
import (
|
||||
"strconv"
|
||||
|
||||
"github.com/engigu/taskpool/internal/cache"
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// ToInt 解析字符串为整数,如果解析失败则返回默认值
|
||||
func ToInt(s string, defaultVal int) int {
|
||||
val, err := strconv.Atoi(s)
|
||||
if err != nil {
|
||||
return defaultVal
|
||||
}
|
||||
return val
|
||||
}
|
||||
|
||||
// ParseInt 解析字符串为整数
|
||||
func ParseInt(s string) (int, error) {
|
||||
return strconv.Atoi(s)
|
||||
}
|
||||
|
||||
// Pagination 分页参数
|
||||
type Pagination struct {
|
||||
Page int
|
||||
PageSize int
|
||||
}
|
||||
|
||||
// getDefaultPageSize 从缓存获取默认分页大小
|
||||
func getDefaultPageSize() int {
|
||||
pageSizeStr := cache.GetSiteCache(constant.KeyPageSize)
|
||||
pageSize, err := strconv.Atoi(pageSizeStr)
|
||||
if err != nil || pageSize < 1 {
|
||||
return 10
|
||||
}
|
||||
return pageSize
|
||||
}
|
||||
|
||||
// ParsePagination 从请求中解析分页参数
|
||||
func ParsePagination(c *gin.Context) Pagination {
|
||||
defaultPageSize := getDefaultPageSize()
|
||||
page, _ := strconv.Atoi(c.DefaultQuery("page", "1"))
|
||||
pageSize, _ := strconv.Atoi(c.DefaultQuery("page_size", strconv.Itoa(defaultPageSize)))
|
||||
|
||||
if page < 1 {
|
||||
page = 1
|
||||
}
|
||||
if pageSize < 1 || pageSize > 10000 {
|
||||
pageSize = defaultPageSize
|
||||
}
|
||||
|
||||
return Pagination{Page: page, PageSize: pageSize}
|
||||
}
|
||||
|
||||
// Offset 计算偏移量
|
||||
func (p Pagination) Offset() int {
|
||||
return (p.Page - 1) * p.PageSize
|
||||
}
|
||||
|
||||
// PaginationData 分页数据
|
||||
type PaginationData struct {
|
||||
Data interface{} `json:"data"`
|
||||
Total int64 `json:"total"`
|
||||
Page int `json:"page"`
|
||||
PageSize int `json:"page_size"`
|
||||
}
|
||||
|
||||
// PaginatedResponse 分页响应
|
||||
func PaginatedResponse(c *gin.Context, data interface{}, total int64, p Pagination) {
|
||||
Success(c, PaginationData{
|
||||
Data: data,
|
||||
Total: total,
|
||||
Page: p.Page,
|
||||
PageSize: p.PageSize,
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,14 @@
|
||||
package utils
|
||||
|
||||
// BoolPtr returns a pointer to the bool value
|
||||
func BoolPtr(b bool) *bool {
|
||||
return &b
|
||||
}
|
||||
|
||||
// DerefBool returns the value of the bool pointer or default if nil
|
||||
func DerefBool(b *bool, defaultVal bool) bool {
|
||||
if b == nil {
|
||||
return defaultVal
|
||||
}
|
||||
return *b
|
||||
}
|
||||
@@ -0,0 +1,18 @@
|
||||
package utils
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"math/big"
|
||||
)
|
||||
|
||||
const charset = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789"
|
||||
|
||||
// RandomString 生成指定长度的随机字符串
|
||||
func RandomString(n int) string {
|
||||
b := make([]byte, n)
|
||||
for i := range b {
|
||||
num, _ := rand.Int(rand.Reader, big.NewInt(int64(len(charset))))
|
||||
b[i] = charset[num.Int64()]
|
||||
}
|
||||
return string(b)
|
||||
}
|
||||
@@ -0,0 +1,68 @@
|
||||
package utils
|
||||
|
||||
import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// GetRepoIdentifier 返回根据仓库URL和分支生成的作者_仓库名标识符
|
||||
func GetRepoIdentifier(url string, branch string) string {
|
||||
url = strings.TrimSuffix(url, ".git")
|
||||
url = strings.TrimSuffix(url, "/")
|
||||
|
||||
repoName := url[strings.LastIndex(url, "/")+1:]
|
||||
|
||||
author := ""
|
||||
lastSlash := strings.LastIndex(url, "/")
|
||||
if lastSlash != -1 {
|
||||
prefix := url[:lastSlash]
|
||||
if strings.Contains(prefix, ":") {
|
||||
parts := strings.Split(prefix, ":")
|
||||
prefix = parts[len(parts)-1]
|
||||
}
|
||||
lastSlashPrefix := strings.LastIndex(prefix, "/")
|
||||
if lastSlashPrefix != -1 {
|
||||
author = prefix[lastSlashPrefix+1:]
|
||||
} else {
|
||||
author = prefix
|
||||
}
|
||||
}
|
||||
|
||||
if dotIdx := strings.LastIndex(author, "."); dotIdx != -1 {
|
||||
author = author[dotIdx+1:]
|
||||
}
|
||||
|
||||
identifier := ""
|
||||
if author != "" {
|
||||
identifier = author + "_" + repoName
|
||||
} else {
|
||||
identifier = repoName
|
||||
}
|
||||
|
||||
if branch != "" && branch != "master" && branch != "main" {
|
||||
identifier = identifier + "_" + branch
|
||||
}
|
||||
|
||||
// Replace any invalid characters for tags or paths
|
||||
identifier = strings.ReplaceAll(identifier, "/", "_")
|
||||
identifier = strings.ReplaceAll(identifier, ".", "_")
|
||||
return identifier
|
||||
}
|
||||
|
||||
// GetActualRepoDir 返回仓库真实的物理目录
|
||||
func GetActualRepoDir(targetPath, sourceURL, branch, sourceType string) string {
|
||||
repoDir := targetPath
|
||||
if sourceType == "git" && sourceURL != "" {
|
||||
repoName := GetRepoIdentifier(sourceURL, branch)
|
||||
// 检查 targetPath 是否已存在且是 Git 仓库
|
||||
gitDir := filepath.Join(repoDir, ".git")
|
||||
if info, err := os.Stat(repoDir); err == nil && info.IsDir() {
|
||||
if _, err := os.Stat(gitDir); os.IsNotExist(err) {
|
||||
// 只有当目标目录存在但不是 Git 仓库时,才追加仓库名
|
||||
repoDir = filepath.Join(repoDir, repoName)
|
||||
}
|
||||
}
|
||||
}
|
||||
return repoDir
|
||||
}
|
||||
@@ -0,0 +1,59 @@
|
||||
package utils
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
type Response struct {
|
||||
Code int `json:"code"`
|
||||
Msg string `json:"msg"`
|
||||
Data interface{} `json:"data,omitempty"`
|
||||
}
|
||||
|
||||
func Success(c *gin.Context, data interface{}) {
|
||||
c.JSON(http.StatusOK, Response{
|
||||
Code: 200,
|
||||
Msg: "success",
|
||||
Data: data,
|
||||
})
|
||||
}
|
||||
|
||||
func SuccessMsg(c *gin.Context, msg string) {
|
||||
c.JSON(http.StatusOK, Response{
|
||||
Code: 200,
|
||||
Msg: msg,
|
||||
})
|
||||
}
|
||||
|
||||
func Error(c *gin.Context, code int, msg string) {
|
||||
c.JSON(http.StatusOK, Response{
|
||||
Code: code,
|
||||
Msg: msg,
|
||||
})
|
||||
}
|
||||
|
||||
func BadRequest(c *gin.Context, msg string) {
|
||||
Error(c, 400, msg)
|
||||
}
|
||||
|
||||
func Unauthorized(c *gin.Context, msg string) {
|
||||
Error(c, 401, msg)
|
||||
}
|
||||
|
||||
func Forbidden(c *gin.Context, msg string) {
|
||||
Error(c, 403, msg)
|
||||
}
|
||||
|
||||
func NotFound(c *gin.Context, msg string) {
|
||||
Error(c, 404, msg)
|
||||
}
|
||||
|
||||
func TooManyRequests(c *gin.Context, msg string) {
|
||||
Error(c, 429, msg)
|
||||
}
|
||||
|
||||
func ServerError(c *gin.Context, msg string) {
|
||||
Error(c, 500, msg)
|
||||
}
|
||||
@@ -0,0 +1,41 @@
|
||||
package utils
|
||||
|
||||
import (
|
||||
"os"
|
||||
"runtime"
|
||||
"runtime/debug"
|
||||
"strconv"
|
||||
|
||||
"github.com/engigu/taskpool/internal/logger"
|
||||
)
|
||||
|
||||
// InitRuntime 设置运行时内存和性能优化参数
|
||||
func InitRuntime() {
|
||||
// 从环境变量或配置读取 GOGC
|
||||
if gogc := os.Getenv("GOGC"); gogc == "" {
|
||||
// 默认设为 80,比 100 更激进一点,减少 RSS 峰值
|
||||
debug.SetGCPercent(80)
|
||||
}
|
||||
|
||||
// 从环境变量读取 GOMEMLIMIT (Go 1.19+)
|
||||
// 建议用户在系统层面设置,例如 BH_MEM_LIMIT=256MiB
|
||||
if memLimit := os.Getenv("BH_MEM_LIMIT"); memLimit != "" {
|
||||
// 这里可以解析一下单位,但简单处理可以直接读取字节
|
||||
if limit, err := strconv.ParseInt(memLimit, 10, 64); err == nil {
|
||||
debug.SetMemoryLimit(limit)
|
||||
logger.Infof("[Runtime] 已设置内存上限: %d 字节", limit)
|
||||
}
|
||||
}
|
||||
|
||||
// 打印当前运行时信息
|
||||
logger.Infof("[Runtime] CPU 核心数: %d, Goroutine 数量: %d", runtime.NumCPU(), runtime.NumGoroutine())
|
||||
}
|
||||
|
||||
// FreeMemory 显式触发内存回收,释放物理资源给 OS
|
||||
// 仅建议在执行了超大批量任务或处理了大型文件后调用
|
||||
func FreeMemory() {
|
||||
// 触发 GC
|
||||
runtime.GC()
|
||||
// 尽可能将内存归还 OS
|
||||
debug.FreeOSMemory()
|
||||
}
|
||||
@@ -0,0 +1,109 @@
|
||||
package utils
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
|
||||
"github.com/engigu/taskpool/internal/constant"
|
||||
)
|
||||
|
||||
// BuildRuntimeProcessEnv 构造 TaskPool 内部可信子进程需要继承的运行时环境变量。
|
||||
// 仅包含 TaskPool 自己的路径/数据库配置,不包含用户任务环境变量。
|
||||
func BuildRuntimeProcessEnv() []string {
|
||||
envs := make([]string, 0, 11)
|
||||
|
||||
configPath := constant.ConfigPath
|
||||
if absConfig, err := filepath.Abs(constant.ConfigPath); err == nil {
|
||||
configPath = absConfig
|
||||
}
|
||||
envs = append(envs, formatEnvVar("BH_CONFIG_PATH", configPath))
|
||||
|
||||
if scriptsDir := ResolveAbsScriptsDir(); strings.TrimSpace(scriptsDir) != "" {
|
||||
envs = append(envs, formatEnvVar("BH_SCRIPTS_DIR", scriptsDir))
|
||||
}
|
||||
|
||||
appendEnvIfSet(&envs, "BH_DB_TYPE", constant.RuntimeDBType)
|
||||
appendEnvIfSet(&envs, "BH_DB_HOST", constant.RuntimeDBHost)
|
||||
if constant.RuntimeDBPort > 0 {
|
||||
envs = append(envs, formatEnvVar("BH_DB_PORT", fmt.Sprintf("%d", constant.RuntimeDBPort)))
|
||||
}
|
||||
appendEnvIfSet(&envs, "BH_DB_USER", constant.RuntimeDBUser)
|
||||
appendEnvIfSet(&envs, "BH_DB_PASSWORD", constant.RuntimeDBPassword)
|
||||
appendEnvIfSet(&envs, "BH_DB_NAME", constant.RuntimeDBName)
|
||||
appendEnvIfSet(&envs, "BH_DB_PATH", constant.RuntimeDBPath)
|
||||
appendEnvIfSet(&envs, "BH_DB_DSN", constant.RuntimeDBDSN)
|
||||
appendEnvIfSet(&envs, "BH_DB_TABLE_PREFIX", constant.RuntimeDBTablePrefix)
|
||||
appendEnvIfSet(&envs, "BH_DB_SSL_MODE", constant.RuntimeDBSSLMode)
|
||||
|
||||
return envs
|
||||
}
|
||||
|
||||
// GetSystemSecrets 返回当前运行时的所有系统级敏感机密(如数据库账号密码等)
|
||||
func GetSystemSecrets() []string {
|
||||
secrets := make([]string, 0, 8)
|
||||
addIfNotEmpty := func(s string) {
|
||||
if s != "" {
|
||||
secrets = append(secrets, s)
|
||||
}
|
||||
}
|
||||
|
||||
addIfNotEmpty(constant.RuntimeDBPassword)
|
||||
addIfNotEmpty(constant.RuntimeDBUser)
|
||||
addIfNotEmpty(constant.RuntimeDBHost)
|
||||
addIfNotEmpty(constant.RuntimeDBName)
|
||||
addIfNotEmpty(constant.RuntimeDBPath)
|
||||
addIfNotEmpty(constant.RuntimeDBDSN)
|
||||
|
||||
return secrets
|
||||
}
|
||||
|
||||
// BuildShellEnvPrefix 将 KEY=VALUE 环境变量切片转换为 shell 前缀。
|
||||
func BuildShellEnvPrefix(envs []string) string {
|
||||
parts := make([]string, 0, len(envs))
|
||||
for _, env := range envs {
|
||||
key, value, ok := strings.Cut(env, "=")
|
||||
if !ok || strings.TrimSpace(key) == "" {
|
||||
continue
|
||||
}
|
||||
parts = append(parts, ShellEnvAssignment(key, value))
|
||||
}
|
||||
|
||||
if len(parts) == 0 {
|
||||
return ""
|
||||
}
|
||||
|
||||
return strings.Join(parts, " ") + " "
|
||||
}
|
||||
|
||||
// ShellEnvAssignment 生成 shell 可安全使用的 KEY='VALUE' 赋值片段。
|
||||
func ShellEnvAssignment(key, value string) string {
|
||||
return key + "='" + strings.ReplaceAll(value, "'", "'\\''") + "'"
|
||||
}
|
||||
|
||||
// ResolveAbsScriptsDir 解析 TaskPool 运行时脚本目录的绝对路径。
|
||||
func ResolveAbsScriptsDir() string {
|
||||
if scriptsDir := os.Getenv("BH_SCRIPTS_DIR"); scriptsDir != "" {
|
||||
if filepath.IsAbs(scriptsDir) {
|
||||
return filepath.Clean(scriptsDir)
|
||||
}
|
||||
if absScriptsDir, err := filepath.Abs(scriptsDir); err == nil {
|
||||
return absScriptsDir
|
||||
}
|
||||
return filepath.Clean(scriptsDir)
|
||||
}
|
||||
|
||||
return constant.ScriptsWorkDir
|
||||
}
|
||||
|
||||
func appendEnvIfSet(envs *[]string, key, value string) {
|
||||
if strings.TrimSpace(value) == "" {
|
||||
return
|
||||
}
|
||||
*envs = append(*envs, formatEnvVar(key, value))
|
||||
}
|
||||
|
||||
func formatEnvVar(key, value string) string {
|
||||
return fmt.Sprintf("%s=%s", key, value)
|
||||
}
|
||||
@@ -0,0 +1,101 @@
|
||||
package utils
|
||||
|
||||
import (
|
||||
"os"
|
||||
"os/exec"
|
||||
"runtime"
|
||||
"strings"
|
||||
"sync"
|
||||
)
|
||||
|
||||
var (
|
||||
defaultShell string
|
||||
defaultArgs []string
|
||||
shellOnce sync.Once
|
||||
)
|
||||
|
||||
// GetShell 返回当前操作系统的 shell 和参数
|
||||
func GetShell() (shell string, args []string) {
|
||||
shellOnce.Do(func() {
|
||||
if runtime.GOOS == "windows" {
|
||||
defaultShell = "cmd"
|
||||
defaultArgs = []string{}
|
||||
return
|
||||
}
|
||||
|
||||
// 1. 优先在 PATH 中查找 bash
|
||||
if path, err := exec.LookPath("bash"); err == nil {
|
||||
defaultShell = path
|
||||
defaultArgs = []string{}
|
||||
return
|
||||
}
|
||||
|
||||
// 2. 其次使用环境变量中的 SHELL
|
||||
if envShell := os.Getenv("SHELL"); envShell != "" {
|
||||
if _, err := os.Stat(envShell); err == nil {
|
||||
defaultShell = envShell
|
||||
defaultArgs = []string{}
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// 3. 尝试在 PATH 中查找 zsh 或 sh
|
||||
for _, s := range []string{"zsh", "sh"} {
|
||||
if path, err := exec.LookPath(s); err == nil {
|
||||
defaultShell = path
|
||||
defaultArgs = []string{}
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// 4. 最后回退到硬编码路径
|
||||
shells := []string{"/usr/bin/bash", "/bin/bash", "/usr/bin/sh", "/bin/sh"}
|
||||
for _, sh := range shells {
|
||||
if _, err := os.Stat(sh); err == nil {
|
||||
defaultShell = sh
|
||||
defaultArgs = []string{}
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
defaultShell = "sh"
|
||||
defaultArgs = []string{}
|
||||
})
|
||||
|
||||
return defaultShell, defaultArgs
|
||||
}
|
||||
|
||||
// GetShellCommand 返回执行命令的 shell 和参数
|
||||
func GetShellCommand(command string) (shell string, args []string) {
|
||||
shell, _ = GetShell()
|
||||
if runtime.GOOS == "windows" {
|
||||
return shell, []string{"/c", command}
|
||||
}
|
||||
return shell, []string{"-c", command}
|
||||
}
|
||||
|
||||
// NewShellCmd 创建一个交互式 shell 命令
|
||||
func NewShellCmd() *exec.Cmd {
|
||||
shell, _ := GetShell()
|
||||
if runtime.GOOS == "windows" {
|
||||
return exec.Command(shell)
|
||||
}
|
||||
// Unix 系统使用 -i 启用交互模式
|
||||
return exec.Command(shell, "-i")
|
||||
}
|
||||
|
||||
// NewShellCommandCmd 创建一个执行指定命令的 shell 命令
|
||||
func NewShellCommandCmd(command string) *exec.Cmd {
|
||||
shell, args := GetShellCommand(command)
|
||||
return exec.Command(shell, args...)
|
||||
}
|
||||
|
||||
// QuotePath 转义并包裹路径,防止 Shell 注入
|
||||
func QuotePath(path string) string {
|
||||
if path == "" {
|
||||
return "''"
|
||||
}
|
||||
// 在 Unix-like 系统中,单引号包裹是最安全的
|
||||
// 需要将路径中的 ' 替换为 '\'' (结束当前引号,转义一个单引号,重新开启引号)
|
||||
return "'" + strings.ReplaceAll(path, "'", "'\\''") + "'"
|
||||
}
|
||||
@@ -0,0 +1,52 @@
|
||||
package utils
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"time"
|
||||
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
)
|
||||
|
||||
type Claims struct {
|
||||
UserID string `json:"user_id"`
|
||||
Username string `json:"username"`
|
||||
TokenVersion int `json:"version"`
|
||||
jwt.RegisteredClaims
|
||||
}
|
||||
|
||||
// GenerateToken 生成 JWT token
|
||||
func GenerateToken(userID string, username string, version int, expireDays int, secret string) (string, error) {
|
||||
claims := Claims{
|
||||
UserID: userID,
|
||||
Username: username,
|
||||
TokenVersion: version,
|
||||
RegisteredClaims: jwt.RegisteredClaims{
|
||||
ExpiresAt: jwt.NewNumericDate(time.Now().Add(time.Duration(expireDays) * 24 * time.Hour)),
|
||||
IssuedAt: jwt.NewNumericDate(time.Now()),
|
||||
},
|
||||
}
|
||||
|
||||
token := jwt.NewWithClaims(jwt.SigningMethodHS256, claims)
|
||||
return token.SignedString([]byte(secret))
|
||||
}
|
||||
|
||||
// ParseToken 解析 JWT token
|
||||
func ParseToken(tokenString string, secret string) (string, string, int, error) {
|
||||
token, err := jwt.ParseWithClaims(tokenString, &Claims{}, func(token *jwt.Token) (any, error) {
|
||||
// 校验算法
|
||||
if _, ok := token.Method.(*jwt.SigningMethodHMAC); !ok {
|
||||
return nil, errors.New("unexpected signing method")
|
||||
}
|
||||
return []byte(secret), nil
|
||||
})
|
||||
|
||||
if err != nil {
|
||||
return "", "", 0, err
|
||||
}
|
||||
|
||||
if claims, ok := token.Claims.(*Claims); ok && token.Valid {
|
||||
return claims.UserID, claims.Username, claims.TokenVersion, nil
|
||||
}
|
||||
|
||||
return "", "", 0, errors.New("invalid token")
|
||||
}
|
||||
Reference in New Issue
Block a user