mirror of
https://github.com/xxnuo/MTranServer.git
synced 2026-09-03 06:35:20 +08:00
feat: update
This commit is contained in:
18
Makefile
18
Makefile
@@ -26,16 +26,20 @@ download-core:
|
||||
@echo "Detecting platform: $(GOOS)-$(GOARCH)"
|
||||
@echo "Downloading $(WORKER_BINARY) from $(DOWNLOAD_URL)..."
|
||||
@mkdir -p bin
|
||||
@rm -f bin/worker$(SUFFIX)
|
||||
@rm -f bin/worker
|
||||
@curl -L -o bin/worker$(SUFFIX) $(DOWNLOAD_URL) || (echo "Failed to download worker binary" && exit 1)
|
||||
@chmod +x bin/worker$(SUFFIX)
|
||||
@echo "Downloaded successfully to bin/worker$(SUFFIX)"
|
||||
@chmod +x bin/worker
|
||||
@echo "Downloaded successfully to bin/worker"
|
||||
@go generate ./bin
|
||||
@echo "Generated successfully to bin/bin.go"
|
||||
|
||||
build-core:
|
||||
@echo "Building core..."
|
||||
@mkdir -p bin
|
||||
@rm -f bin/worker$(SUFFIX)
|
||||
@rm -f bin/worker
|
||||
@cd ../../MTranCore && make build-worker
|
||||
@cp ../../MTranCore/build/worker bin/worker$(SUFFIX)
|
||||
@chmod +x bin/worker$(SUFFIX)
|
||||
@echo "Built successfully to bin/worker$(SUFFIX)"
|
||||
@cp ../../MTranCore/build/worker bin/worker
|
||||
@chmod +x bin/worker
|
||||
@echo "Built successfully to bin/worker"
|
||||
@go generate ./bin
|
||||
@echo "Generated successfully to bin/bin.go"
|
||||
|
||||
@@ -5,4 +5,4 @@ import (
|
||||
)
|
||||
|
||||
//go:embed worker
|
||||
var workerBin []byte
|
||||
var WorkerBinary []byte
|
||||
|
||||
4
bin/bin_hash.go
Normal file
4
bin/bin_hash.go
Normal file
@@ -0,0 +1,4 @@
|
||||
// Code generated by go generate; DO NOT EDIT.
|
||||
package bin
|
||||
|
||||
const WorkerHash = "5ba53860f605a731b07d77a8a94455a359ef73f798d2882829aad8db395de799"
|
||||
42
bin/gen_hash.go
Normal file
42
bin/gen_hash.go
Normal file
@@ -0,0 +1,42 @@
|
||||
//go:build ignore
|
||||
|
||||
package main
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"log"
|
||||
"os"
|
||||
"text/template"
|
||||
)
|
||||
|
||||
const OutputPath = "bin_hash.go"
|
||||
const Template = `// Code generated by go generate; DO NOT EDIT.
|
||||
package bin
|
||||
|
||||
const WorkerHash = "{{.Hash}}"
|
||||
`
|
||||
|
||||
func main() {
|
||||
data, err := os.ReadFile("worker")
|
||||
if err != nil {
|
||||
log.Fatalf("Failed to read worker file: %v", err)
|
||||
}
|
||||
|
||||
hashBytes := sha256.Sum256(data)
|
||||
hashString := hex.EncodeToString(hashBytes[:])
|
||||
|
||||
f, err := os.Create(OutputPath)
|
||||
if err != nil {
|
||||
log.Fatalf("Failed to create output file: %v", err)
|
||||
}
|
||||
defer f.Close()
|
||||
|
||||
t := template.Must(template.New("hash").Parse(Template))
|
||||
err = t.Execute(f, struct{ Hash string }{Hash: hashString})
|
||||
if err != nil {
|
||||
log.Fatalf("Failed to execute template: %v", err)
|
||||
}
|
||||
|
||||
log.Printf("Successfully generated %s, WorkerHash: %s", OutputPath, hashString)
|
||||
}
|
||||
15
bin/utils.go
Normal file
15
bin/utils.go
Normal file
@@ -0,0 +1,15 @@
|
||||
package bin
|
||||
|
||||
import "crypto/sha256"
|
||||
|
||||
//go:generate go run gen_hash.go
|
||||
|
||||
// GetWorkerInfo returns information about the embedded worker binary
|
||||
func GetWorkerInfo() (hash string, size int) {
|
||||
return WorkerHash, len(WorkerBinary)
|
||||
}
|
||||
|
||||
// ComputeHash computes the SHA256 hash of the given data
|
||||
func ComputeHash(data []byte) [32]byte {
|
||||
return sha256.Sum256(data)
|
||||
}
|
||||
7
go.mod
7
go.mod
@@ -1,3 +1,8 @@
|
||||
module github.com/xxnuo/MTranServer/v3
|
||||
module github.com/xxnuo/MTranServer
|
||||
|
||||
go 1.25.3
|
||||
|
||||
require (
|
||||
github.com/ShinyTrinkets/meta-logger v0.2.0 // indirect
|
||||
github.com/ShinyTrinkets/overseer v0.6.0 // indirect
|
||||
)
|
||||
|
||||
4
go.sum
Normal file
4
go.sum
Normal file
@@ -0,0 +1,4 @@
|
||||
github.com/ShinyTrinkets/meta-logger v0.2.0 h1:oR533+wuhSJ+vLsnSq1CBSGQygNv8nDsvuRUVcOls0g=
|
||||
github.com/ShinyTrinkets/meta-logger v0.2.0/go.mod h1:cY1KnpPfpLIopR+arZXHYVrVGO6AETrhi3HmRGFjU+U=
|
||||
github.com/ShinyTrinkets/overseer v0.6.0 h1:SxzF7nACfr9Ln9oR4yyRfHSYkfXIu81UAwBw5A6LsPc=
|
||||
github.com/ShinyTrinkets/overseer v0.6.0/go.mod h1:NB2kg7uXESqJYBC7Y+5cdCFTpCZBPwEsSdYnabfNOOw=
|
||||
61
internal/config.go
Normal file
61
internal/config.go
Normal file
@@ -0,0 +1,61 @@
|
||||
package internal
|
||||
|
||||
import (
|
||||
"flag"
|
||||
|
||||
"github.com/xxnuo/MTranServer/internal/utils"
|
||||
)
|
||||
|
||||
// Config 包含服务器配置
|
||||
type Config struct {
|
||||
LogLevel string
|
||||
ConfigDir string
|
||||
ModelDir string
|
||||
|
||||
// 服务器配置
|
||||
Host string
|
||||
Port string
|
||||
EnableWebUI bool
|
||||
EnableOfflineMode bool
|
||||
}
|
||||
|
||||
// LoadConfig 加载配置,优先级:命令行参数 > 环境变量 > 默认值
|
||||
func LoadConfig() *Config {
|
||||
// Define command line flags
|
||||
logLevel := flag.String("log-level", "", "Log level (debug, info, warn, error)")
|
||||
configDir := flag.String("config-dir", "", "Config directory")
|
||||
modelDir := flag.String("model-dir", "", "Model directory")
|
||||
host := flag.String("host", "", "Server host address")
|
||||
port := flag.String("port", "", "Server port")
|
||||
enableWebUI := flag.String("ui", "", "Enable web UI (true/false)")
|
||||
enableOfflineMode := flag.String("offline", "", "Enable offline mode (true/false)")
|
||||
|
||||
flag.Parse()
|
||||
return &Config{
|
||||
LogLevel: getConfigValue(*logLevel, "LOG_LEVEL", "info"),
|
||||
ConfigDir: getConfigValue(*configDir, "CONFIG_DIR", "./"),
|
||||
ModelDir: getConfigValue(*modelDir, "MODEL_DIR", "./"),
|
||||
|
||||
// 服务器配置
|
||||
Host: getConfigValue(*host, "HOST", "0.0.0.0"),
|
||||
Port: getConfigValue(*port, "PORT", "8989"),
|
||||
EnableWebUI: getBoolConfigValue(*enableWebUI, "UI", "true"),
|
||||
EnableOfflineMode: getBoolConfigValue(*enableOfflineMode, "OFFLINE", "true"),
|
||||
}
|
||||
}
|
||||
|
||||
// getConfigValue 获取字符串值,优先级:命令行参数 > 环境变量 > 默认值
|
||||
func getConfigValue(flagValue, envKey, defaultValue string) string {
|
||||
if flagValue != "" {
|
||||
return flagValue
|
||||
}
|
||||
return utils.GetEnv(envKey, defaultValue)
|
||||
}
|
||||
|
||||
// getBoolConfigValue 获取布尔值,优先级:命令行参数 > 环境变量 > 默认值
|
||||
func getBoolConfigValue(flagValue, envKey, defaultValue string) bool {
|
||||
if flagValue != "" {
|
||||
return utils.ParseBoolEnv(flagValue)
|
||||
}
|
||||
return utils.ParseBoolEnv(utils.GetEnv(envKey, defaultValue))
|
||||
}
|
||||
382
internal/process/daemon.go
Normal file
382
internal/process/daemon.go
Normal file
@@ -0,0 +1,382 @@
|
||||
package process
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"sync"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
"github.com/ShinyTrinkets/overseer"
|
||||
"github.com/xxnuo/MTranServer/bin"
|
||||
)
|
||||
|
||||
const (
|
||||
maxLogLines = 1000
|
||||
)
|
||||
|
||||
// WorkerArgs 包含工作进程的配置
|
||||
type WorkerArgs struct {
|
||||
Host string
|
||||
Port int
|
||||
WorkDir string
|
||||
EnableGRPC bool
|
||||
EnableHTTP bool
|
||||
EnableWebSocket bool
|
||||
GRPCUnixSocket string
|
||||
LogLevel string
|
||||
BinaryPath string // 写入工作程序二进制文件的路径,如果为空则使用 /tmp
|
||||
}
|
||||
|
||||
// NewWorkerArgs 创建一个新的 WorkerArgs 实例,使用默认值
|
||||
func NewWorkerArgs() *WorkerArgs {
|
||||
return &WorkerArgs{
|
||||
Host: "0.0.0.0",
|
||||
Port: 8988,
|
||||
WorkDir: ".",
|
||||
EnableGRPC: false,
|
||||
EnableHTTP: false,
|
||||
EnableWebSocket: true,
|
||||
GRPCUnixSocket: "",
|
||||
LogLevel: "info",
|
||||
}
|
||||
}
|
||||
|
||||
// Worker 管理使用 overseer 的工作进程
|
||||
type Worker struct {
|
||||
args *WorkerArgs
|
||||
overseer *overseer.Overseer
|
||||
id string
|
||||
binaryPath string // 实际写入二进制文件的路径
|
||||
mu sync.RWMutex
|
||||
logChan chan *overseer.LogMsg
|
||||
stateChan chan *overseer.ProcessJSON
|
||||
logs []string
|
||||
maxLogs int
|
||||
}
|
||||
|
||||
// NewWorker 创建一个新的 Worker 实例
|
||||
func NewWorker(args *WorkerArgs) *Worker {
|
||||
// 确定二进制文件路径
|
||||
binaryPath := args.BinaryPath
|
||||
if binaryPath == "" {
|
||||
binaryPath = "/tmp/mtran-worker"
|
||||
}
|
||||
|
||||
// 根据二进制文件路径和端口生成唯一的 worker ID
|
||||
workerID := fmt.Sprintf("mtran-worker-%d", args.Port)
|
||||
|
||||
w := &Worker{
|
||||
args: args,
|
||||
overseer: overseer.NewOverseer(),
|
||||
id: workerID,
|
||||
binaryPath: binaryPath,
|
||||
logChan: make(chan *overseer.LogMsg, 100),
|
||||
stateChan: make(chan *overseer.ProcessJSON, 10),
|
||||
logs: make([]string, 0, maxLogLines),
|
||||
maxLogs: maxLogLines,
|
||||
}
|
||||
|
||||
// 订阅日志和状态变化
|
||||
w.overseer.WatchLogs(w.logChan)
|
||||
w.overseer.WatchState(w.stateChan)
|
||||
|
||||
// 启动日志收集器
|
||||
go w.collectLogs()
|
||||
|
||||
return w
|
||||
}
|
||||
|
||||
// ensureWorkerBinary 提取嵌入的工作程序二进制文件到指定路径
|
||||
func (w *Worker) ensureWorkerBinary() error {
|
||||
// 检查二进制文件是否已存在并且匹配哈希
|
||||
if data, err := os.ReadFile(w.binaryPath); err == nil {
|
||||
// 二进制文件存在,计算其哈希并比较
|
||||
existingHash := fmt.Sprintf("%x", bin.ComputeHash(data))
|
||||
if existingHash == bin.WorkerHash {
|
||||
// 哈希匹配,二进制文件是最新的
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
// 确保父目录存在
|
||||
if err := os.MkdirAll(filepath.Dir(w.binaryPath), 0755); err != nil {
|
||||
return fmt.Errorf("failed to create directory for worker binary: %w", err)
|
||||
}
|
||||
|
||||
// 写入嵌入的二进制文件
|
||||
if err := os.WriteFile(w.binaryPath, bin.WorkerBinary, 0755); err != nil {
|
||||
return fmt.Errorf("failed to write worker binary: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// buildArgs 构建工作程序的命令行参数
|
||||
func (w *Worker) buildArgs() []string {
|
||||
args := []string{
|
||||
"--host", w.args.Host,
|
||||
"--port", strconv.Itoa(w.args.Port),
|
||||
"--log-level", w.args.LogLevel,
|
||||
}
|
||||
|
||||
if w.args.WorkDir != "" {
|
||||
absWorkDir, err := filepath.Abs(w.args.WorkDir)
|
||||
if err == nil {
|
||||
args = append(args, "--work-dir", absWorkDir)
|
||||
} else {
|
||||
args = append(args, "--work-dir", w.args.WorkDir)
|
||||
}
|
||||
}
|
||||
if w.args.EnableGRPC {
|
||||
args = append(args, "--enable-grpc", "true")
|
||||
} else {
|
||||
args = append(args, "--enable-grpc", "false")
|
||||
}
|
||||
|
||||
if w.args.EnableHTTP {
|
||||
args = append(args, "--enable-http", "true")
|
||||
} else {
|
||||
args = append(args, "--enable-http", "false")
|
||||
}
|
||||
|
||||
if w.args.EnableWebSocket {
|
||||
args = append(args, "--enable-websocket", "true")
|
||||
} else {
|
||||
args = append(args, "--enable-websocket", "false")
|
||||
}
|
||||
|
||||
if w.args.GRPCUnixSocket != "" {
|
||||
args = append(args, "--grpc-unix-socket", w.args.GRPCUnixSocket)
|
||||
}
|
||||
|
||||
return args
|
||||
}
|
||||
|
||||
// Start 启动工作进程
|
||||
func (w *Worker) Start() error {
|
||||
w.mu.Lock()
|
||||
defer w.mu.Unlock()
|
||||
|
||||
// 检查是否已经运行
|
||||
if w.overseer.HasProc(w.id) {
|
||||
status := w.overseer.Status(w.id)
|
||||
if status != nil && status.State == "running" {
|
||||
return fmt.Errorf("worker already running")
|
||||
}
|
||||
// 如果存在但未运行,则移除旧进程
|
||||
w.overseer.Remove(w.id)
|
||||
}
|
||||
|
||||
// 确保工作程序二进制文件可用
|
||||
if err := w.ensureWorkerBinary(); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// 构建命令行参数
|
||||
args := w.buildArgs()
|
||||
|
||||
// 将工作进程添加到 overseer
|
||||
// 注意: overseer.Add 接受 []string 作为单个参数,而不是可变参数字符串
|
||||
cmd := w.overseer.Add(w.id, w.binaryPath, args)
|
||||
if cmd == nil {
|
||||
return fmt.Errorf("failed to add worker to overseer")
|
||||
}
|
||||
|
||||
// 配置进程
|
||||
cmd.Dir = w.args.WorkDir
|
||||
cmd.DelayStart = 0
|
||||
cmd.RetryTimes = 0 // 默认不自动重启
|
||||
|
||||
// 在 goroutine 中启动监督
|
||||
go w.overseer.Supervise(w.id)
|
||||
|
||||
// 等待一段时间让进程启动
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Stop 优雅地停止工作进程
|
||||
func (w *Worker) Stop() error {
|
||||
w.mu.Lock()
|
||||
defer w.mu.Unlock()
|
||||
|
||||
if !w.overseer.HasProc(w.id) {
|
||||
return fmt.Errorf("worker not found")
|
||||
}
|
||||
|
||||
status := w.overseer.Status(w.id)
|
||||
if status == nil || status.State != "running" {
|
||||
return fmt.Errorf("worker not running")
|
||||
}
|
||||
|
||||
// 优雅地停止进程
|
||||
if err := w.overseer.Stop(w.id); err != nil {
|
||||
return fmt.Errorf("failed to stop worker: %w", err)
|
||||
}
|
||||
|
||||
// 等待进程停止
|
||||
timeout := time.After(10 * time.Second)
|
||||
ticker := time.NewTicker(100 * time.Millisecond)
|
||||
defer ticker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-timeout:
|
||||
// 如果优雅停止失败,则强制杀死进程
|
||||
w.overseer.Signal(w.id, syscall.SIGKILL)
|
||||
return fmt.Errorf("worker stop timeout, forced kill")
|
||||
case <-ticker.C:
|
||||
status := w.overseer.Status(w.id)
|
||||
if status != nil && status.State != "running" {
|
||||
return nil
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Restart 重启工作进程
|
||||
func (w *Worker) Restart() error {
|
||||
// 如果正在运行,则停止
|
||||
if w.overseer.HasProc(w.id) {
|
||||
status := w.overseer.Status(w.id)
|
||||
if status != nil && status.State == "running" {
|
||||
if err := w.Stop(); err != nil {
|
||||
return fmt.Errorf("failed to stop worker: %w", err)
|
||||
}
|
||||
}
|
||||
// 移除旧进程
|
||||
w.overseer.Remove(w.id)
|
||||
}
|
||||
|
||||
// 等待一段时间再重启
|
||||
time.Sleep(500 * time.Millisecond)
|
||||
|
||||
// 再次启动
|
||||
return w.Start()
|
||||
}
|
||||
|
||||
// Status 返回工作进程的当前状态
|
||||
func (w *Worker) Status() string {
|
||||
w.mu.RLock()
|
||||
defer w.mu.RUnlock()
|
||||
|
||||
if !w.overseer.HasProc(w.id) {
|
||||
return "not_started"
|
||||
}
|
||||
|
||||
status := w.overseer.Status(w.id)
|
||||
if status == nil {
|
||||
return "unknown"
|
||||
}
|
||||
|
||||
return status.State
|
||||
}
|
||||
|
||||
// GetDetailedStatus 返回详细的状态信息
|
||||
func (w *Worker) GetDetailedStatus() *overseer.ProcessJSON {
|
||||
w.mu.RLock()
|
||||
defer w.mu.RUnlock()
|
||||
|
||||
if !w.overseer.HasProc(w.id) {
|
||||
return nil
|
||||
}
|
||||
|
||||
return w.overseer.Status(w.id)
|
||||
}
|
||||
|
||||
// Logs 返回最近的工作日志行
|
||||
func (w *Worker) Logs() []string {
|
||||
w.mu.RLock()
|
||||
defer w.mu.RUnlock()
|
||||
|
||||
// 返回一个副本以避免竞争条件
|
||||
logsCopy := make([]string, len(w.logs))
|
||||
copy(logsCopy, w.logs)
|
||||
return logsCopy
|
||||
}
|
||||
|
||||
// collectLogs 收集工作进程的日志
|
||||
func (w *Worker) collectLogs() {
|
||||
for {
|
||||
select {
|
||||
case msg, ok := <-w.logChan:
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
w.mu.Lock()
|
||||
// 格式化日志消息
|
||||
logType := "INFO"
|
||||
if msg.Type == 1 {
|
||||
logType = "ERROR"
|
||||
}
|
||||
logLine := fmt.Sprintf("[%s] [%s] %s",
|
||||
time.Now().Format("2006-01-02 15:04:05"), logType, msg.Text)
|
||||
w.logs = append(w.logs, logLine)
|
||||
|
||||
// 只保留最近的日志
|
||||
if len(w.logs) > w.maxLogs {
|
||||
w.logs = w.logs[len(w.logs)-w.maxLogs:]
|
||||
}
|
||||
w.mu.Unlock()
|
||||
|
||||
case state, ok := <-w.stateChan:
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
w.mu.Lock()
|
||||
// 记录状态变化
|
||||
stateLog := fmt.Sprintf("[%s] Worker state: %s (PID: %d)",
|
||||
time.Now().Format("2006-01-02 15:04:05"), state.State, state.PID)
|
||||
w.logs = append(w.logs, stateLog)
|
||||
|
||||
if len(w.logs) > w.maxLogs {
|
||||
w.logs = w.logs[len(w.logs)-w.maxLogs:]
|
||||
}
|
||||
w.mu.Unlock()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// IsRunning 检查工作进程是否正在运行
|
||||
func (w *Worker) IsRunning() bool {
|
||||
return w.Status() == "running"
|
||||
}
|
||||
|
||||
// Signal 发送信号到工作进程
|
||||
func (w *Worker) Signal(sig syscall.Signal) error {
|
||||
w.mu.Lock()
|
||||
defer w.mu.Unlock()
|
||||
|
||||
if !w.overseer.HasProc(w.id) {
|
||||
return fmt.Errorf("worker not found")
|
||||
}
|
||||
|
||||
return w.overseer.Signal(w.id, sig)
|
||||
}
|
||||
|
||||
// Cleanup 清理资源
|
||||
func (w *Worker) Cleanup() error {
|
||||
// 如果正在运行,则停止
|
||||
if w.IsRunning() {
|
||||
if err := w.Stop(); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
// 取消订阅通道
|
||||
w.overseer.UnWatchLogs(w.logChan)
|
||||
w.overseer.UnWatchState(w.stateChan)
|
||||
|
||||
// 关闭通道
|
||||
close(w.logChan)
|
||||
close(w.stateChan)
|
||||
|
||||
// 注意: 我们不在这里删除工作程序二进制文件,因为它可能被共享
|
||||
// 或者用户有意放置在特定位置
|
||||
|
||||
return nil
|
||||
}
|
||||
246
internal/process/daemon_test.go
Normal file
246
internal/process/daemon_test.go
Normal file
@@ -0,0 +1,246 @@
|
||||
package process_test
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/xxnuo/MTranServer/bin"
|
||||
"github.com/xxnuo/MTranServer/internal/process"
|
||||
"github.com/xxnuo/MTranServer/internal/utils"
|
||||
)
|
||||
|
||||
// Example demonstrates basic usage of the Worker
|
||||
func TestBasicUsage(t *testing.T) {
|
||||
if testing.Short() {
|
||||
t.Skip("Skipping integration test in short mode")
|
||||
}
|
||||
|
||||
// Create worker arguments with custom configuration
|
||||
args := process.NewWorkerArgs()
|
||||
args.Host = "127.0.0.1"
|
||||
port, err := utils.GetFreePort()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get free port: %v", err)
|
||||
}
|
||||
args.Port = port
|
||||
args.EnableWebSocket = true
|
||||
args.EnableHTTP = true
|
||||
args.LogLevel = "debug"
|
||||
args.WorkDir = "/tmp/mtran"
|
||||
// Binary will be written to /tmp by default (BinaryPath is empty)
|
||||
|
||||
// Create a new worker
|
||||
worker := process.NewWorker(args)
|
||||
defer worker.Cleanup()
|
||||
|
||||
// Start the worker
|
||||
if err := worker.Start(); err != nil {
|
||||
t.Logf("Failed to start worker: %v\n", err)
|
||||
t.Fatalf("Failed to start worker: %v", err)
|
||||
}
|
||||
|
||||
// Wait for worker to be running
|
||||
time.Sleep(2 * time.Second)
|
||||
|
||||
// Check status
|
||||
status := worker.Status()
|
||||
t.Logf("Worker status: %s\n", status)
|
||||
|
||||
// Get detailed status
|
||||
detailedStatus := worker.GetDetailedStatus()
|
||||
if detailedStatus != nil {
|
||||
t.Logf("Worker PID: %d\n", detailedStatus.PID)
|
||||
t.Logf("Worker state: %s\n", detailedStatus.State)
|
||||
}
|
||||
|
||||
// Check if running
|
||||
if worker.IsRunning() {
|
||||
t.Log("Worker is running")
|
||||
}
|
||||
|
||||
// Get logs
|
||||
logs := worker.Logs()
|
||||
t.Logf("Collected %d log lines\n", len(logs))
|
||||
for _, log := range logs {
|
||||
t.Log(log)
|
||||
}
|
||||
|
||||
// Restart the worker
|
||||
if err := worker.Restart(); err != nil {
|
||||
t.Logf("Failed to restart worker: %v\n", err)
|
||||
t.Fatalf("Failed to restart worker: %v", err)
|
||||
}
|
||||
|
||||
time.Sleep(2 * time.Second)
|
||||
|
||||
// Stop the worker
|
||||
if err := worker.Stop(); err != nil {
|
||||
t.Logf("Failed to stop worker: %v\n", err)
|
||||
t.Fatalf("Failed to stop worker: %v", err)
|
||||
}
|
||||
|
||||
t.Log("Worker stopped successfully")
|
||||
}
|
||||
|
||||
// Example demonstrates worker lifecycle
|
||||
func TestLifecycle(t *testing.T) {
|
||||
if testing.Short() {
|
||||
t.Skip("Skipping integration test in short mode")
|
||||
}
|
||||
|
||||
args := process.NewWorkerArgs()
|
||||
port, err := utils.GetFreePort()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get free port: %v", err)
|
||||
}
|
||||
args.Port = port
|
||||
// Binary will be written to /tmp by default
|
||||
worker := process.NewWorker(args)
|
||||
defer worker.Cleanup()
|
||||
|
||||
// Start
|
||||
worker.Start()
|
||||
time.Sleep(1 * time.Second)
|
||||
|
||||
// Check status
|
||||
t.Log("Status:", worker.Status())
|
||||
|
||||
// Stop
|
||||
worker.Stop()
|
||||
time.Sleep(500 * time.Millisecond)
|
||||
|
||||
// Check status again
|
||||
t.Log("Status after stop:", worker.Status())
|
||||
}
|
||||
|
||||
// TestWorkerHash verifies that WorkerHash is properly computed
|
||||
func TestWorkerHash(t *testing.T) {
|
||||
// Check that WorkerHash is not empty
|
||||
if bin.WorkerHash == "" {
|
||||
t.Fatal("WorkerHash should not be empty")
|
||||
}
|
||||
|
||||
// Check that hash is a valid hex string (64 chars for SHA256)
|
||||
if len(bin.WorkerHash) != 64 {
|
||||
t.Fatalf("WorkerHash should be 64 characters (SHA256), got %d", len(bin.WorkerHash))
|
||||
}
|
||||
|
||||
t.Logf("Worker binary hash: %s", bin.WorkerHash)
|
||||
t.Logf("Worker binary size: %d bytes", len(bin.WorkerBinary))
|
||||
|
||||
// Verify the worker starts successfully with the hash check
|
||||
args := process.NewWorkerArgs()
|
||||
port, err := utils.GetFreePort()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get free port: %v", err)
|
||||
}
|
||||
args.Port = port
|
||||
worker := process.NewWorker(args)
|
||||
defer worker.Cleanup()
|
||||
|
||||
// First start - should write binary and hash file
|
||||
if err := worker.Start(); err != nil {
|
||||
t.Fatalf("Failed to start worker on first attempt: %v", err)
|
||||
}
|
||||
time.Sleep(1 * time.Second)
|
||||
worker.Stop()
|
||||
time.Sleep(500 * time.Millisecond)
|
||||
|
||||
// Second start - should detect existing hash and skip writing
|
||||
if err := worker.Start(); err != nil {
|
||||
t.Fatalf("Failed to start worker on second attempt: %v", err)
|
||||
}
|
||||
time.Sleep(1 * time.Second)
|
||||
worker.Stop()
|
||||
time.Sleep(500 * time.Millisecond)
|
||||
|
||||
t.Log("Worker hash verification successful")
|
||||
}
|
||||
|
||||
// TestCustomBinaryPath demonstrates using a custom binary path
|
||||
func TestCustomBinaryPath(t *testing.T) {
|
||||
if testing.Short() {
|
||||
t.Skip("Skipping integration test in short mode")
|
||||
}
|
||||
|
||||
// Test with custom binary path
|
||||
args := process.NewWorkerArgs()
|
||||
port, err := utils.GetFreePort()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get free port: %v", err)
|
||||
}
|
||||
args.Port = port
|
||||
args.BinaryPath = "/tmp/custom-mtran-worker" // Custom path
|
||||
worker := process.NewWorker(args)
|
||||
defer worker.Cleanup()
|
||||
|
||||
if err := worker.Start(); err != nil {
|
||||
t.Fatalf("Failed to start worker with custom binary path: %v", err)
|
||||
}
|
||||
time.Sleep(1 * time.Second)
|
||||
|
||||
if !worker.IsRunning() {
|
||||
t.Fatal("Worker should be running")
|
||||
}
|
||||
|
||||
worker.Stop()
|
||||
time.Sleep(500 * time.Millisecond)
|
||||
t.Log("Worker with custom binary path stopped successfully")
|
||||
}
|
||||
|
||||
// TestMultipleWorkers demonstrates running multiple workers concurrently
|
||||
func TestMultipleWorkers(t *testing.T) {
|
||||
if testing.Short() {
|
||||
t.Skip("Skipping integration test in short mode")
|
||||
}
|
||||
|
||||
// Create multiple workers with different ports
|
||||
workers := make([]*process.Worker, 0, 3)
|
||||
|
||||
for i := 0; i < 3; i++ {
|
||||
port, err := utils.GetFreePort()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get free port: %v", err)
|
||||
}
|
||||
|
||||
args := process.NewWorkerArgs()
|
||||
args.Port = port
|
||||
args.Host = "127.0.0.1"
|
||||
args.EnableWebSocket = true
|
||||
// Each worker writes to /tmp by default but has unique ID based on port
|
||||
|
||||
worker := process.NewWorker(args)
|
||||
workers = append(workers, worker)
|
||||
|
||||
if err := worker.Start(); err != nil {
|
||||
t.Fatalf("Failed to start worker %d: %v", i, err)
|
||||
}
|
||||
t.Logf("Worker %d started on port %d", i, port)
|
||||
}
|
||||
|
||||
// Let all workers run for a bit
|
||||
time.Sleep(2 * time.Second)
|
||||
|
||||
// Verify all workers are running
|
||||
for i, worker := range workers {
|
||||
if !worker.IsRunning() {
|
||||
t.Errorf("Worker %d should be running", i)
|
||||
}
|
||||
status := worker.GetDetailedStatus()
|
||||
if status != nil {
|
||||
t.Logf("Worker %d: PID=%d, State=%s", i, status.PID, status.State)
|
||||
}
|
||||
}
|
||||
|
||||
// Stop all workers
|
||||
for i, worker := range workers {
|
||||
if err := worker.Stop(); err != nil {
|
||||
t.Errorf("Failed to stop worker %d: %v", i, err)
|
||||
}
|
||||
worker.Cleanup()
|
||||
t.Logf("Worker %d stopped", i)
|
||||
}
|
||||
|
||||
time.Sleep(500 * time.Millisecond)
|
||||
t.Log("All workers stopped successfully")
|
||||
}
|
||||
31
internal/utils/env.go
Normal file
31
internal/utils/env.go
Normal file
@@ -0,0 +1,31 @@
|
||||
package utils
|
||||
|
||||
import (
|
||||
"os"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// GetEnv gets the environment variable value or returns the default value
|
||||
func GetEnv(key, defaultValue string) string {
|
||||
if value := os.Getenv(key); value != "" {
|
||||
return value
|
||||
}
|
||||
return defaultValue
|
||||
}
|
||||
|
||||
// ParseBoolEnv parses boolean values from environment variables
|
||||
// Supports: true/false, 1/0, yes/no (case insensitive)
|
||||
func ParseBoolEnv(value string) bool {
|
||||
value = strings.ToLower(strings.TrimSpace(value))
|
||||
switch value {
|
||||
case "true", "1", "yes":
|
||||
return true
|
||||
case "false", "0", "no":
|
||||
return false
|
||||
default:
|
||||
// Try standard parsing as fallback
|
||||
result, _ := strconv.ParseBool(value)
|
||||
return result
|
||||
}
|
||||
}
|
||||
112
internal/utils/env_test.go
Normal file
112
internal/utils/env_test.go
Normal file
@@ -0,0 +1,112 @@
|
||||
package utils
|
||||
|
||||
import (
|
||||
"os"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestGetEnv(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
key string
|
||||
defaultValue string
|
||||
envValue string
|
||||
setEnv bool
|
||||
want string
|
||||
}{
|
||||
{
|
||||
name: "environment variable exists",
|
||||
key: "TEST_VAR",
|
||||
defaultValue: "default",
|
||||
envValue: "custom",
|
||||
setEnv: true,
|
||||
want: "custom",
|
||||
},
|
||||
{
|
||||
name: "environment variable not set",
|
||||
key: "TEST_VAR_NOT_SET",
|
||||
defaultValue: "default",
|
||||
envValue: "",
|
||||
setEnv: false,
|
||||
want: "default",
|
||||
},
|
||||
{
|
||||
name: "environment variable is empty string",
|
||||
key: "TEST_VAR_EMPTY",
|
||||
defaultValue: "default",
|
||||
envValue: "",
|
||||
setEnv: true,
|
||||
want: "default",
|
||||
},
|
||||
{
|
||||
name: "default value is empty",
|
||||
key: "TEST_VAR_DEFAULT_EMPTY",
|
||||
defaultValue: "",
|
||||
envValue: "",
|
||||
setEnv: false,
|
||||
want: "",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
// Setup
|
||||
if tt.setEnv {
|
||||
os.Setenv(tt.key, tt.envValue)
|
||||
defer os.Unsetenv(tt.key)
|
||||
}
|
||||
|
||||
// Test
|
||||
got := GetEnv(tt.key, tt.defaultValue)
|
||||
if got != tt.want {
|
||||
t.Errorf("GetEnv() = %v, want %v", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseBoolEnv(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
value string
|
||||
want bool
|
||||
}{
|
||||
// true cases
|
||||
{name: "true lowercase", value: "true", want: true},
|
||||
{name: "true uppercase", value: "TRUE", want: true},
|
||||
{name: "true mixed case", value: "True", want: true},
|
||||
{name: "1", value: "1", want: true},
|
||||
{name: "yes lowercase", value: "yes", want: true},
|
||||
{name: "yes uppercase", value: "YES", want: true},
|
||||
{name: "yes mixed case", value: "Yes", want: true},
|
||||
{name: "true with spaces", value: " true ", want: true},
|
||||
{name: "yes with spaces", value: " yes ", want: true},
|
||||
|
||||
// false cases
|
||||
{name: "false lowercase", value: "false", want: false},
|
||||
{name: "false uppercase", value: "FALSE", want: false},
|
||||
{name: "false mixed case", value: "False", want: false},
|
||||
{name: "0", value: "0", want: false},
|
||||
{name: "no lowercase", value: "no", want: false},
|
||||
{name: "no uppercase", value: "NO", want: false},
|
||||
{name: "no mixed case", value: "No", want: false},
|
||||
{name: "false with spaces", value: " false ", want: false},
|
||||
{name: "no with spaces", value: " no ", want: false},
|
||||
|
||||
// invalid/edge cases
|
||||
{name: "empty string", value: "", want: false},
|
||||
{name: "invalid string", value: "invalid", want: false},
|
||||
{name: "random text", value: "xyz", want: false},
|
||||
{name: "number 2", value: "2", want: false},
|
||||
{name: "spaces only", value: " ", want: false},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got := ParseBoolEnv(tt.value)
|
||||
if got != tt.want {
|
||||
t.Errorf("ParseBoolEnv(%q) = %v, want %v", tt.value, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
18
internal/utils/port.go
Normal file
18
internal/utils/port.go
Normal file
@@ -0,0 +1,18 @@
|
||||
package utils
|
||||
|
||||
import "net"
|
||||
|
||||
// GetFreePort 获取一个未被占用的随机端口
|
||||
func GetFreePort() (int, error) {
|
||||
addr, err := net.ResolveTCPAddr("tcp", "localhost:0")
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
|
||||
l, err := net.ListenTCP("tcp", addr)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
defer l.Close()
|
||||
return l.Addr().(*net.TCPAddr).Port, nil
|
||||
}
|
||||
15
internal/utils/port_test.go
Normal file
15
internal/utils/port_test.go
Normal file
@@ -0,0 +1,15 @@
|
||||
package utils_test
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/xxnuo/MTranServer/internal/utils"
|
||||
)
|
||||
|
||||
func TestGetFreePort(t *testing.T) {
|
||||
port, err := utils.GetFreePort()
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to get free port: %v", err)
|
||||
}
|
||||
t.Logf("Free port: %d", port)
|
||||
}
|
||||
Reference in New Issue
Block a user