mirror of
https://github.com/xxnuo/MTranServer.git
synced 2026-09-03 06:35:20 +08:00
chore: clean code
This commit is contained in:
@@ -18,7 +18,7 @@ const (
|
||||
)
|
||||
|
||||
func main() {
|
||||
// Detect platform
|
||||
|
||||
goos := runtime.GOOS
|
||||
goarch := runtime.GOARCH
|
||||
|
||||
@@ -38,7 +38,6 @@ func main() {
|
||||
log.Printf("Detecting platform: %s-%s", goos, goarch)
|
||||
log.Printf("Downloading %s from %s...", workerBinary, downloadURL)
|
||||
|
||||
// Ensure bin directory exists
|
||||
if err := os.MkdirAll(".", 0755); err != nil {
|
||||
log.Fatalf("Failed to create bin directory: %v", err)
|
||||
}
|
||||
@@ -46,7 +45,6 @@ func main() {
|
||||
targetFile := "worker"
|
||||
os.Remove(targetFile)
|
||||
|
||||
// Download
|
||||
d := downloader.New(".")
|
||||
err := d.Download(downloadURL, targetFile, &downloader.DownloadOptions{
|
||||
Context: context.Background(),
|
||||
@@ -56,7 +54,6 @@ func main() {
|
||||
log.Fatalf("Failed to download worker binary: %v", err)
|
||||
}
|
||||
|
||||
// Set executable permission
|
||||
if err := os.Chmod(targetFile, 0755); err != nil {
|
||||
log.Printf("Warning: Failed to set executable permission: %v", err)
|
||||
}
|
||||
|
||||
@@ -5,12 +5,10 @@ import "crypto/sha256"
|
||||
//go:generate go run gen_update.go
|
||||
//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)
|
||||
}
|
||||
|
||||
@@ -19,9 +19,8 @@ import (
|
||||
|
||||
// @contact.name API Support
|
||||
// @contact.url https://github.com/xxnuo/MTranServer/issues
|
||||
// @contact.email support@example.com
|
||||
|
||||
// @license.name MIT
|
||||
// @license.name Apache 2.0
|
||||
// @license.url https://github.com/xxnuo/MTranServer/blob/main/LICENSE
|
||||
|
||||
// @host localhost:8989
|
||||
@@ -36,11 +35,10 @@ import (
|
||||
// @name token
|
||||
|
||||
func main() {
|
||||
// 定义 version 和 help 标志
|
||||
|
||||
versionFlag := flag.Bool("version", false, "Show version information")
|
||||
versionShortFlag := flag.Bool("v", false, "Show version information (shorthand)")
|
||||
|
||||
// 自定义 Usage 函数
|
||||
flag.Usage = func() {
|
||||
fmt.Fprintf(os.Stderr, "MTranServer %s - Ultra-low resource consumption, ultra-fast offline translation server\n\n", version.GetVersion())
|
||||
fmt.Fprintf(os.Stderr, "Usage:\n")
|
||||
@@ -64,22 +62,17 @@ func main() {
|
||||
fmt.Fprintf(os.Stderr, "\nMore information: https://github.com/xxnuo/MTranServer\n")
|
||||
}
|
||||
|
||||
// 加载配置(会注册其他标志)
|
||||
cfg := config.GetConfig()
|
||||
|
||||
// 解析命令行参数
|
||||
flag.Parse()
|
||||
|
||||
// 设置日志级别
|
||||
logger.SetLevel(cfg.LogLevel)
|
||||
|
||||
// 处理 version 标志
|
||||
if *versionFlag || *versionShortFlag {
|
||||
fmt.Printf("MTranServer %s\n", version.GetVersion())
|
||||
os.Exit(0)
|
||||
}
|
||||
|
||||
// 启动服务器
|
||||
if err := server.Run(); err != nil {
|
||||
logger.Fatal("Server error: %v", err)
|
||||
}
|
||||
|
||||
@@ -13,7 +13,6 @@ import (
|
||||
func main() {
|
||||
log.Printf("Downloading records.json from %s...", models.RecordsUrl)
|
||||
|
||||
// Download using downloader
|
||||
d := downloader.New(".")
|
||||
err := d.Download(models.RecordsUrl, models.RecordsFileName, &downloader.DownloadOptions{
|
||||
Context: context.Background(),
|
||||
|
||||
@@ -1,3 +1,3 @@
|
||||
package data
|
||||
|
||||
//go:generate env GOOS= GOARCH= go run gen_records.go
|
||||
//go:generate go run gen_records.go
|
||||
|
||||
@@ -8,21 +8,18 @@ import (
|
||||
"github.com/xxnuo/MTranServer/internal/utils"
|
||||
)
|
||||
|
||||
// Config 包含服务器配置
|
||||
type Config struct {
|
||||
// 内部配置
|
||||
LogLevel string
|
||||
HomeDir string
|
||||
ConfigDir string
|
||||
ModelDir string
|
||||
|
||||
// 服务器配置
|
||||
Host string
|
||||
Port string
|
||||
EnableWebUI bool
|
||||
EnableOfflineMode bool
|
||||
WorkerIdleTimeout int // Worker 空闲超时时间(秒)
|
||||
APIToken string // API 访问令牌
|
||||
WorkerIdleTimeout int
|
||||
APIToken string
|
||||
}
|
||||
|
||||
var (
|
||||
|
||||
@@ -15,37 +15,30 @@ import (
|
||||
"github.com/xxnuo/MTranServer/internal/utils"
|
||||
)
|
||||
|
||||
// Downloader 下载器结构
|
||||
type Downloader struct {
|
||||
// 下载目录
|
||||
DestDir string
|
||||
// 进度回调函数
|
||||
|
||||
ProgressFunc getter.ProgressTracker
|
||||
}
|
||||
|
||||
// DownloadOptions 下载选项
|
||||
type DownloadOptions struct {
|
||||
// SHA256 校验和
|
||||
SHA256 string
|
||||
// 是否覆盖已存在的文件
|
||||
|
||||
Overwrite bool
|
||||
// Context 用于取消下载
|
||||
|
||||
Context context.Context
|
||||
}
|
||||
|
||||
// New 创建新的下载器
|
||||
func New(destDir string) *Downloader {
|
||||
return &Downloader{
|
||||
DestDir: destDir,
|
||||
}
|
||||
}
|
||||
|
||||
// SetProgressFunc 设置进度回调函数
|
||||
func (d *Downloader) SetProgressFunc(fn getter.ProgressTracker) {
|
||||
d.ProgressFunc = fn
|
||||
}
|
||||
|
||||
// Download 下载文件到指定目录
|
||||
func (d *Downloader) Download(urlStr, filename string, opts *DownloadOptions) error {
|
||||
if opts == nil {
|
||||
opts = &DownloadOptions{
|
||||
@@ -56,21 +49,18 @@ func (d *Downloader) Download(urlStr, filename string, opts *DownloadOptions) er
|
||||
opts.Context = context.Background()
|
||||
}
|
||||
|
||||
// 确保目标目录存在
|
||||
if err := os.MkdirAll(d.DestDir, 0755); err != nil {
|
||||
return fmt.Errorf("Failed to create directory: %w", err)
|
||||
}
|
||||
|
||||
// 目标文件路径
|
||||
dst := filepath.Join(d.DestDir, filename)
|
||||
|
||||
// 检查文件是否已存在
|
||||
if !opts.Overwrite {
|
||||
if _, err := os.Stat(dst); err == nil {
|
||||
// 文件存在,检查 SHA256
|
||||
|
||||
if opts.SHA256 != "" {
|
||||
if err := utils.VerifySHA256(dst, opts.SHA256); err == nil {
|
||||
// 文件已存在且校验通过,跳过下载
|
||||
|
||||
logger.Debug("File %s already exists and verified, skipping download", filename)
|
||||
return nil
|
||||
}
|
||||
@@ -80,15 +70,13 @@ func (d *Downloader) Download(urlStr, filename string, opts *DownloadOptions) er
|
||||
|
||||
logger.Info("Downloading %s from %s", filename, urlStr)
|
||||
|
||||
// 创建临时文件
|
||||
tmpFile := dst + ".tmp"
|
||||
defer os.Remove(tmpFile)
|
||||
|
||||
// 创建 HTTP 客户端,支持代理和重定向
|
||||
httpClient := &http.Client{
|
||||
Timeout: 30 * time.Minute,
|
||||
CheckRedirect: func(req *http.Request, via []*http.Request) error {
|
||||
// 允许最多 10 次重定向
|
||||
|
||||
if len(via) >= 10 {
|
||||
return fmt.Errorf("stopped after 10 redirects")
|
||||
}
|
||||
@@ -96,12 +84,10 @@ func (d *Downloader) Download(urlStr, filename string, opts *DownloadOptions) er
|
||||
},
|
||||
}
|
||||
|
||||
// 配置代理
|
||||
transport := &http.Transport{
|
||||
TLSClientConfig: &tls.Config{InsecureSkipVerify: false},
|
||||
}
|
||||
|
||||
// 从环境变量读取代理设置
|
||||
if proxyURL := os.Getenv("HTTP_PROXY"); proxyURL != "" {
|
||||
if parsedURL, err := url.Parse(proxyURL); err == nil {
|
||||
transport.Proxy = http.ProxyURL(parsedURL)
|
||||
@@ -112,7 +98,6 @@ func (d *Downloader) Download(urlStr, filename string, opts *DownloadOptions) er
|
||||
}
|
||||
}
|
||||
|
||||
// 支持 HTTPS_PROXY
|
||||
if proxyURL := os.Getenv("HTTPS_PROXY"); proxyURL != "" {
|
||||
if parsedURL, err := url.Parse(proxyURL); err == nil {
|
||||
transport.Proxy = http.ProxyURL(parsedURL)
|
||||
@@ -125,12 +110,10 @@ func (d *Downloader) Download(urlStr, filename string, opts *DownloadOptions) er
|
||||
|
||||
httpClient.Transport = transport
|
||||
|
||||
// 配置 HttpGetter
|
||||
httpGetter := &getter.HttpGetter{
|
||||
Client: httpClient,
|
||||
}
|
||||
|
||||
// 配置 getter 客户端选项
|
||||
clientOpts := []getter.ClientOption{
|
||||
getter.WithContext(opts.Context),
|
||||
getter.WithGetters(map[string]getter.Getter{
|
||||
@@ -143,26 +126,22 @@ func (d *Downloader) Download(urlStr, filename string, opts *DownloadOptions) er
|
||||
clientOpts = append(clientOpts, getter.WithProgress(d.ProgressFunc))
|
||||
}
|
||||
|
||||
// 创建 getter 客户端
|
||||
client := &getter.Client{
|
||||
Src: urlStr,
|
||||
Dst: tmpFile,
|
||||
Mode: getter.ClientModeFile,
|
||||
}
|
||||
|
||||
// 应用客户端选项
|
||||
if err := client.Configure(clientOpts...); err != nil {
|
||||
return fmt.Errorf("Failed to configure downloader: %w", err)
|
||||
}
|
||||
|
||||
// 执行下载
|
||||
if err := client.Get(); err != nil {
|
||||
return fmt.Errorf("Failed to download: %w", err)
|
||||
}
|
||||
|
||||
logger.Debug("Download completed: %s", filename)
|
||||
|
||||
// 校验 SHA256
|
||||
if opts.SHA256 != "" {
|
||||
logger.Debug("Verifying SHA256 for %s", filename)
|
||||
if err := utils.VerifySHA256(tmpFile, opts.SHA256); err != nil {
|
||||
@@ -171,7 +150,6 @@ func (d *Downloader) Download(urlStr, filename string, opts *DownloadOptions) er
|
||||
logger.Debug("SHA256 verification passed for %s", filename)
|
||||
}
|
||||
|
||||
// 移动临时文件到目标位置
|
||||
if err := os.Rename(tmpFile, dst); err != nil {
|
||||
return fmt.Errorf("Failed to move file: %w", err)
|
||||
}
|
||||
@@ -180,7 +158,6 @@ func (d *Downloader) Download(urlStr, filename string, opts *DownloadOptions) er
|
||||
return nil
|
||||
}
|
||||
|
||||
// DownloadFile 快捷方法:下载文件
|
||||
func DownloadFile(url, destPath, sha256sum string) error {
|
||||
dir := filepath.Dir(destPath)
|
||||
filename := filepath.Base(destPath)
|
||||
|
||||
@@ -11,7 +11,7 @@ import (
|
||||
)
|
||||
|
||||
func TestDownload(t *testing.T) {
|
||||
// 创建测试 HTTP 服务器
|
||||
|
||||
testContent := []byte("Hello, World!")
|
||||
expectedSHA256 := "dffd6021bb2bd5b0af676290809ec3a53191dd81c7f70a4b28688a362182986f"
|
||||
|
||||
@@ -20,14 +20,12 @@ func TestDownload(t *testing.T) {
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
// 创建临时目录
|
||||
tempDir, err := os.MkdirTemp("", "downloader-test-*")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer os.RemoveAll(tempDir)
|
||||
|
||||
// 测试下载
|
||||
d := New(tempDir)
|
||||
err = d.Download(server.URL, "test.txt", &DownloadOptions{
|
||||
SHA256: expectedSHA256,
|
||||
@@ -38,13 +36,11 @@ func TestDownload(t *testing.T) {
|
||||
t.Fatalf("下载失败: %v", err)
|
||||
}
|
||||
|
||||
// 验证文件存在
|
||||
filePath := filepath.Join(tempDir, "test.txt")
|
||||
if _, err := os.Stat(filePath); os.IsNotExist(err) {
|
||||
t.Fatal("文件不存在")
|
||||
}
|
||||
|
||||
// 验证文件内容
|
||||
content, err := os.ReadFile(filePath)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
@@ -56,7 +52,7 @@ func TestDownload(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestDownloadWithWrongSHA256(t *testing.T) {
|
||||
// 创建测试 HTTP 服务器
|
||||
|
||||
testContent := []byte("Hello, World!")
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
@@ -64,14 +60,12 @@ func TestDownloadWithWrongSHA256(t *testing.T) {
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
// 创建临时目录
|
||||
tempDir, err := os.MkdirTemp("", "downloader-test-*")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer os.RemoveAll(tempDir)
|
||||
|
||||
// 测试下载(使用错误的 SHA256)
|
||||
d := New(tempDir)
|
||||
err = d.Download(server.URL, "test.txt", &DownloadOptions{
|
||||
SHA256: "0000000000000000000000000000000000000000000000000000000000000000",
|
||||
@@ -84,7 +78,7 @@ func TestDownloadWithWrongSHA256(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestDownloadSkipExisting(t *testing.T) {
|
||||
// 创建测试 HTTP 服务器
|
||||
|
||||
testContent := []byte("Hello, World!")
|
||||
expectedSHA256 := "dffd6021bb2bd5b0af676290809ec3a53191dd81c7f70a4b28688a362182986f"
|
||||
requestCount := 0
|
||||
@@ -95,14 +89,12 @@ func TestDownloadSkipExisting(t *testing.T) {
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
// 创建临时目录
|
||||
tempDir, err := os.MkdirTemp("", "downloader-test-*")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer os.RemoveAll(tempDir)
|
||||
|
||||
// 第一次下载
|
||||
d := New(tempDir)
|
||||
err = d.Download(server.URL, "test.txt", &DownloadOptions{
|
||||
SHA256: expectedSHA256,
|
||||
@@ -115,7 +107,6 @@ func TestDownloadSkipExisting(t *testing.T) {
|
||||
|
||||
firstRequestCount := requestCount
|
||||
|
||||
// 第二次下载(应该跳过)
|
||||
err = d.Download(server.URL, "test.txt", &DownloadOptions{
|
||||
SHA256: expectedSHA256,
|
||||
Context: context.Background(),
|
||||
@@ -131,7 +122,7 @@ func TestDownloadSkipExisting(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestDownloadWithOverwrite(t *testing.T) {
|
||||
// 创建测试 HTTP 服务器
|
||||
|
||||
testContent := []byte("Hello, World!")
|
||||
expectedSHA256 := "dffd6021bb2bd5b0af676290809ec3a53191dd81c7f70a4b28688a362182986f"
|
||||
requestCount := 0
|
||||
@@ -142,14 +133,12 @@ func TestDownloadWithOverwrite(t *testing.T) {
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
// 创建临时目录
|
||||
tempDir, err := os.MkdirTemp("", "downloader-test-*")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer os.RemoveAll(tempDir)
|
||||
|
||||
// 第一次下载
|
||||
d := New(tempDir)
|
||||
err = d.Download(server.URL, "test.txt", &DownloadOptions{
|
||||
SHA256: expectedSHA256,
|
||||
@@ -162,7 +151,6 @@ func TestDownloadWithOverwrite(t *testing.T) {
|
||||
|
||||
firstRequestCount := requestCount
|
||||
|
||||
// 第二次下载(强制覆盖)
|
||||
err = d.Download(server.URL, "test.txt", &DownloadOptions{
|
||||
SHA256: expectedSHA256,
|
||||
Overwrite: true,
|
||||
@@ -179,25 +167,22 @@ func TestDownloadWithOverwrite(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestDownloadWithContext(t *testing.T) {
|
||||
// 创建测试 HTTP 服务器(延迟响应)
|
||||
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
time.Sleep(2 * time.Second)
|
||||
w.Write([]byte("Hello, World!"))
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
// 创建临时目录
|
||||
tempDir, err := os.MkdirTemp("", "downloader-test-*")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer os.RemoveAll(tempDir)
|
||||
|
||||
// 创建带超时的 context
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 500*time.Millisecond)
|
||||
defer cancel()
|
||||
|
||||
// 测试下载(应该超时)
|
||||
d := New(tempDir)
|
||||
err = d.Download(server.URL, "test.txt", &DownloadOptions{
|
||||
Context: ctx,
|
||||
@@ -209,7 +194,7 @@ func TestDownloadWithContext(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestDownloadFile(t *testing.T) {
|
||||
// 创建测试 HTTP 服务器
|
||||
|
||||
testContent := []byte("Hello, World!")
|
||||
expectedSHA256 := "dffd6021bb2bd5b0af676290809ec3a53191dd81c7f70a4b28688a362182986f"
|
||||
|
||||
@@ -218,14 +203,12 @@ func TestDownloadFile(t *testing.T) {
|
||||
}))
|
||||
defer server.Close()
|
||||
|
||||
// 创建临时目录
|
||||
tempDir, err := os.MkdirTemp("", "downloader-test-*")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer os.RemoveAll(tempDir)
|
||||
|
||||
// 测试快捷方法
|
||||
destPath := filepath.Join(tempDir, "test.txt")
|
||||
err = DownloadFile(server.URL, destPath, expectedSHA256)
|
||||
|
||||
@@ -233,12 +216,10 @@ func TestDownloadFile(t *testing.T) {
|
||||
t.Fatalf("下载失败: %v", err)
|
||||
}
|
||||
|
||||
// 验证文件存在
|
||||
if _, err := os.Stat(destPath); os.IsNotExist(err) {
|
||||
t.Fatal("文件不存在")
|
||||
}
|
||||
|
||||
// 验证文件内容
|
||||
content, err := os.ReadFile(destPath)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
|
||||
@@ -11,8 +11,6 @@ import (
|
||||
"github.com/xxnuo/MTranServer/internal/services"
|
||||
)
|
||||
|
||||
// deeplLangToBCP47 DeepL 语言代码转 BCP47
|
||||
// 只列出与 BCP47 不同的映射(DeepL 使用大写和特殊代码)
|
||||
var deeplLangToBCP47 = map[string]string{
|
||||
"NB": "no",
|
||||
"ZH": "zh-Hans",
|
||||
@@ -20,8 +18,6 @@ var deeplLangToBCP47 = map[string]string{
|
||||
"ZH-TW": "zh-Hant",
|
||||
}
|
||||
|
||||
// bcp47ToDeeplLang BCP47 转 DeepL 语言代码
|
||||
// 只列出与 BCP47 不同的映射(DeepL 使用大写和特殊代码)
|
||||
var bcp47ToDeeplLang = map[string]string{
|
||||
"no": "NB",
|
||||
"zh-Hans": "ZH",
|
||||
@@ -30,27 +26,24 @@ var bcp47ToDeeplLang = map[string]string{
|
||||
"zh-TW": "ZH-TW",
|
||||
}
|
||||
|
||||
// convertDeeplLangToBCP47 将 DeepL 语言代码转换为 BCP47
|
||||
func convertDeeplLangToBCP47(deeplLang string) string {
|
||||
// 转换为大写进行匹配
|
||||
|
||||
upperLang := strings.ToUpper(deeplLang)
|
||||
if bcp47, ok := deeplLangToBCP47[upperLang]; ok {
|
||||
return bcp47
|
||||
}
|
||||
// 如果不在映射表中,返回小写版本
|
||||
|
||||
return strings.ToLower(deeplLang)
|
||||
}
|
||||
|
||||
// convertBCP47ToDeeplLang 将 BCP47 语言代码转换为 DeepL
|
||||
func convertBCP47ToDeeplLang(bcp47Lang string) string {
|
||||
if deeplLang, ok := bcp47ToDeeplLang[bcp47Lang]; ok {
|
||||
return deeplLang
|
||||
}
|
||||
// 如果不在映射表中,返回大写版本
|
||||
|
||||
return strings.ToUpper(bcp47Lang)
|
||||
}
|
||||
|
||||
// DeeplTranslateRequest DeepL 翻译请求
|
||||
type DeeplTranslateRequest struct {
|
||||
Text []string `json:"text" binding:"required" example:"Hello, world!"`
|
||||
SourceLang string `json:"source_lang,omitempty" example:"EN"`
|
||||
@@ -69,13 +62,11 @@ type DeeplTranslateRequest struct {
|
||||
EnableBetaLanguages bool `json:"enable_beta_languages,omitempty"`
|
||||
}
|
||||
|
||||
// DeeplTranslation 翻译结果
|
||||
type DeeplTranslation struct {
|
||||
DetectedSourceLanguage string `json:"detected_source_language" example:"EN"`
|
||||
Text string `json:"text" example:"Hallo, Welt!"`
|
||||
}
|
||||
|
||||
// DeeplTranslateResponse DeepL 翻译响应
|
||||
type DeeplTranslateResponse struct {
|
||||
Translations []DeeplTranslation `json:"translations"`
|
||||
}
|
||||
@@ -95,19 +86,19 @@ type DeeplTranslateResponse struct {
|
||||
// @Router /deepl [post]
|
||||
func HandleDeeplTranslate(apiToken string) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
// 检查 token - 兼容 DeepL API 认证方式
|
||||
|
||||
if apiToken != "" {
|
||||
// 支持 DeepL 标准认证: Authorization: DeepL-Auth-Key [key]
|
||||
|
||||
authHeader := c.GetHeader("Authorization")
|
||||
token := ""
|
||||
|
||||
if strings.HasPrefix(authHeader, "DeepL-Auth-Key ") {
|
||||
token = strings.TrimPrefix(authHeader, "DeepL-Auth-Key ")
|
||||
} else if authHeader != "" {
|
||||
// 也支持标准 Bearer token
|
||||
|
||||
token = strings.TrimPrefix(authHeader, "Bearer ")
|
||||
} else {
|
||||
// 兼容 query 参数方式
|
||||
|
||||
token = c.Query("token")
|
||||
}
|
||||
|
||||
@@ -127,19 +118,16 @@ func HandleDeeplTranslate(apiToken string) gin.HandlerFunc {
|
||||
return
|
||||
}
|
||||
|
||||
// 转换语言代码:DeepL -> BCP47
|
||||
sourceLang := "auto"
|
||||
if req.SourceLang != "" {
|
||||
sourceLang = convertDeeplLangToBCP47(req.SourceLang)
|
||||
}
|
||||
targetLang := convertDeeplLangToBCP47(req.TargetLang)
|
||||
|
||||
// 批量翻译
|
||||
translations := make([]DeeplTranslation, len(req.Text))
|
||||
ctx, cancel := context.WithTimeout(c.Request.Context(), 120*time.Second)
|
||||
defer cancel()
|
||||
|
||||
// 确定是否需要 HTML 处理
|
||||
isHTML := req.TagHandling == "html" || req.TagHandling == "xml"
|
||||
|
||||
for i, text := range req.Text {
|
||||
@@ -151,7 +139,6 @@ func HandleDeeplTranslate(apiToken string) gin.HandlerFunc {
|
||||
return
|
||||
}
|
||||
|
||||
// 返回 DeepL 格式的语言代码
|
||||
detectedLang := req.SourceLang
|
||||
if detectedLang == "" {
|
||||
detectedLang = convertBCP47ToDeeplLang(sourceLang)
|
||||
|
||||
@@ -11,8 +11,6 @@ import (
|
||||
"github.com/xxnuo/MTranServer/internal/services"
|
||||
)
|
||||
|
||||
// googleLangToBCP47 Google 语言代码转 BCP47
|
||||
// 只列出与 BCP47 不同的映射,其他直接返回原值
|
||||
var googleLangToBCP47 = map[string]string{
|
||||
"zh-CN": "zh-Hans",
|
||||
"zh-TW": "zh-Hant",
|
||||
@@ -20,32 +18,27 @@ var googleLangToBCP47 = map[string]string{
|
||||
"zh-SG": "zh-Hans",
|
||||
}
|
||||
|
||||
// bcp47ToGoogleLang BCP47 转 Google 语言代码
|
||||
// 只列出与 BCP47 不同的映射,其他直接返回原值
|
||||
var bcp47ToGoogleLang = map[string]string{
|
||||
"zh-Hans": "zh-CN",
|
||||
"zh-Hant": "zh-TW",
|
||||
}
|
||||
|
||||
// convertGoogleLangToBCP47 将 Google 语言代码转换为 BCP47
|
||||
func convertGoogleLangToBCP47(googleLang string) string {
|
||||
if bcp47, ok := googleLangToBCP47[googleLang]; ok {
|
||||
return bcp47
|
||||
}
|
||||
// 如果不在映射表中,直接返回(大部分语言代码相同)
|
||||
|
||||
return googleLang
|
||||
}
|
||||
|
||||
// convertBCP47ToGoogleLang 将 BCP47 语言代码转换为 Google
|
||||
func convertBCP47ToGoogleLang(bcp47Lang string) string {
|
||||
if googleLang, ok := bcp47ToGoogleLang[bcp47Lang]; ok {
|
||||
return googleLang
|
||||
}
|
||||
// 如果不在映射表中,直接返回
|
||||
|
||||
return bcp47Lang
|
||||
}
|
||||
|
||||
// GoogleTranslateRequest Google 翻译兼容请求
|
||||
type GoogleTranslateRequest struct {
|
||||
Q string `json:"q" binding:"required" example:"The Great Pyramid of Giza"`
|
||||
Source string `json:"source" binding:"required" example:"en"`
|
||||
@@ -53,7 +46,6 @@ type GoogleTranslateRequest struct {
|
||||
Format string `json:"format" example:"text"`
|
||||
}
|
||||
|
||||
// GoogleTranslateResponse Google 翻译兼容响应
|
||||
type GoogleTranslateResponse struct {
|
||||
Data struct {
|
||||
Translations []struct {
|
||||
@@ -77,12 +69,11 @@ type GoogleTranslateResponse struct {
|
||||
// @Router /google/language/translate/v2 [post]
|
||||
func HandleGoogleCompatTranslate(apiToken string) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
// 检查 token - 兼容 Google API 认证方式
|
||||
|
||||
if apiToken != "" {
|
||||
// 支持 Google API 标准认证: ?key=xxx
|
||||
|
||||
token := c.Query("key")
|
||||
|
||||
// 也支持标准 Authorization header
|
||||
if token == "" {
|
||||
authHeader := c.GetHeader("Authorization")
|
||||
if strings.HasPrefix(authHeader, "Bearer ") {
|
||||
@@ -92,7 +83,6 @@ func HandleGoogleCompatTranslate(apiToken string) gin.HandlerFunc {
|
||||
}
|
||||
}
|
||||
|
||||
// 兼容通用 token 参数
|
||||
if token == "" {
|
||||
token = c.Query("token")
|
||||
}
|
||||
@@ -114,11 +104,9 @@ func HandleGoogleCompatTranslate(apiToken string) gin.HandlerFunc {
|
||||
return
|
||||
}
|
||||
|
||||
// 转换 Google 语言代码到 BCP47
|
||||
sourceBCP47 := convertGoogleLangToBCP47(req.Source)
|
||||
targetBCP47 := convertGoogleLangToBCP47(req.Target)
|
||||
|
||||
// 翻译
|
||||
ctx, cancel := context.WithTimeout(c.Request.Context(), 60*time.Second)
|
||||
defer cancel()
|
||||
|
||||
@@ -162,7 +150,7 @@ func HandleGoogleCompatTranslate(apiToken string) gin.HandlerFunc {
|
||||
// @Router /google/translate_a/single [get]
|
||||
func HandleGoogleTranslateSingle(apiToken string) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
// 检查 token
|
||||
|
||||
if apiToken != "" {
|
||||
token := c.Query("key")
|
||||
|
||||
@@ -187,7 +175,6 @@ func HandleGoogleTranslateSingle(apiToken string) gin.HandlerFunc {
|
||||
}
|
||||
}
|
||||
|
||||
// 获取参数
|
||||
sl := c.Query("sl")
|
||||
tl := c.Query("tl")
|
||||
q := c.Query("q")
|
||||
@@ -199,19 +186,15 @@ func HandleGoogleTranslateSingle(apiToken string) gin.HandlerFunc {
|
||||
return
|
||||
}
|
||||
|
||||
// 支持 auto 自动检测源语言
|
||||
if sl == "" {
|
||||
sl = "auto"
|
||||
}
|
||||
|
||||
// q 参数已经由 Gin 自动进行了 URL 解码
|
||||
text := q
|
||||
|
||||
// 转换 Google 语言代码到 BCP47
|
||||
sourceBCP47 := convertGoogleLangToBCP47(sl)
|
||||
targetBCP47 := convertGoogleLangToBCP47(tl)
|
||||
|
||||
// 翻译
|
||||
ctx, cancel := context.WithTimeout(c.Request.Context(), 60*time.Second)
|
||||
defer cancel()
|
||||
|
||||
@@ -223,10 +206,6 @@ func HandleGoogleTranslateSingle(apiToken string) gin.HandlerFunc {
|
||||
return
|
||||
}
|
||||
|
||||
// 返回 translate_a/single 格式的响应
|
||||
// 格式: [[["翻译结果","原文",null,null,1]],null,"检测到的源语言",null,null,null,null,[]]
|
||||
// response[0][0][0] 是翻译结果
|
||||
// response[2] 是检测到的源语言(返回 Google 格式)
|
||||
detectedLang := convertBCP47ToGoogleLang(sourceBCP47)
|
||||
response := []interface{}{
|
||||
[]interface{}{
|
||||
|
||||
@@ -11,7 +11,6 @@ import (
|
||||
"github.com/xxnuo/MTranServer/internal/services"
|
||||
)
|
||||
|
||||
// 划词翻译语言名称到 BCP47 的映射
|
||||
var hcfyLangToBCP47 = map[string]string{
|
||||
"中文(简体)": "zh-Hans",
|
||||
"中文(繁体)": "zh-Hant",
|
||||
@@ -119,7 +118,6 @@ var hcfyLangToBCP47 = map[string]string{
|
||||
"博多语": "brx",
|
||||
}
|
||||
|
||||
// bcp47ToHcfyLang BCP47 到划词翻译语言名称的映射
|
||||
var bcp47ToHcfyLang = map[string]string{
|
||||
"zh-Hans": "中文(简体)",
|
||||
"zh-CN": "中文(简体)",
|
||||
@@ -229,25 +227,22 @@ var bcp47ToHcfyLang = map[string]string{
|
||||
"brx": "博多语",
|
||||
}
|
||||
|
||||
// convertHcfyLangToBCP47 将划词翻译语言名称转换为 BCP47
|
||||
func convertHcfyLangToBCP47(hcfyLang string) string {
|
||||
if bcp47, ok := hcfyLangToBCP47[hcfyLang]; ok {
|
||||
return bcp47
|
||||
}
|
||||
// 如果不在映射表中,尝试直接返回(可能已经是 BCP47)
|
||||
|
||||
return hcfyLang
|
||||
}
|
||||
|
||||
// convertBCP47ToHcfyLang 将 BCP47 语言代码转换为划词翻译语言名称
|
||||
func convertBCP47ToHcfyLang(bcp47Lang string) string {
|
||||
if hcfyLang, ok := bcp47ToHcfyLang[bcp47Lang]; ok {
|
||||
return hcfyLang
|
||||
}
|
||||
// 如果不在映射表中,返回原值
|
||||
|
||||
return bcp47Lang
|
||||
}
|
||||
|
||||
// HcfyTranslateRequest 划词翻译请求
|
||||
type HcfyTranslateRequest struct {
|
||||
Name string `json:"name" binding:"required" example:"翻译一"`
|
||||
Text string `json:"text" binding:"required" example:"Hello, word translation."`
|
||||
@@ -255,20 +250,17 @@ type HcfyTranslateRequest struct {
|
||||
Source string `json:"source" example:"英语"`
|
||||
}
|
||||
|
||||
// HcfyPhonetic 音标
|
||||
type HcfyPhonetic struct {
|
||||
Name string `json:"name,omitempty" example:"美"`
|
||||
TtsURI string `json:"ttsURI,omitempty" example:"https://..."`
|
||||
Value string `json:"value,omitempty" example:"həˈloʊ"`
|
||||
}
|
||||
|
||||
// HcfyDict 词典释义
|
||||
type HcfyDict struct {
|
||||
Pos string `json:"pos,omitempty" example:"n."`
|
||||
Terms []string `json:"terms" example:"你好,问候"`
|
||||
}
|
||||
|
||||
// HcfyTranslateResponse 划词翻译响应
|
||||
type HcfyTranslateResponse struct {
|
||||
Text string `json:"text" example:"Hello, word translation."`
|
||||
From string `json:"from" example:"英语"`
|
||||
@@ -295,7 +287,7 @@ type HcfyTranslateResponse struct {
|
||||
// @Router /hcfy [post]
|
||||
func HandleHcfyTranslate(apiToken string) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
// 检查 token
|
||||
|
||||
if apiToken != "" {
|
||||
token := c.Query("token")
|
||||
if token == "" {
|
||||
@@ -324,15 +316,11 @@ func HandleHcfyTranslate(apiToken string) gin.HandlerFunc {
|
||||
return
|
||||
}
|
||||
|
||||
// 转换语言代码:划词翻译语言名称 -> BCP47
|
||||
sourceLang := "auto"
|
||||
if req.Source != "" {
|
||||
sourceLang = convertHcfyLangToBCP47(req.Source)
|
||||
}
|
||||
|
||||
// 确定目标语种
|
||||
// destination 是一个数组,首要目标语种是第一个元素
|
||||
// 如果源语种与首要目标语种相同,则使用次要目标语种
|
||||
if len(req.Destination) == 0 {
|
||||
c.JSON(http.StatusBadRequest, gin.H{
|
||||
"error": "destination is required",
|
||||
@@ -343,10 +331,9 @@ func HandleHcfyTranslate(apiToken string) gin.HandlerFunc {
|
||||
targetLangName := req.Destination[0]
|
||||
targetLang := convertHcfyLangToBCP47(targetLangName)
|
||||
|
||||
// 检测源文本语种(简单判断)
|
||||
detectedSourceLang := sourceLang
|
||||
if sourceLang == "auto" {
|
||||
// 简单的语种检测逻辑
|
||||
|
||||
if containsChinese(req.Text) {
|
||||
detectedSourceLang = "zh-Hans"
|
||||
} else if containsJapanese(req.Text) {
|
||||
@@ -358,17 +345,14 @@ func HandleHcfyTranslate(apiToken string) gin.HandlerFunc {
|
||||
}
|
||||
}
|
||||
|
||||
// 如果检测到的源语种与目标语种相同,且有次要目标语种,则使用次要目标语种
|
||||
if detectedSourceLang == targetLang && len(req.Destination) > 1 {
|
||||
targetLangName = req.Destination[1]
|
||||
targetLang = convertHcfyLangToBCP47(targetLangName)
|
||||
}
|
||||
|
||||
// 翻译
|
||||
ctx, cancel := context.WithTimeout(c.Request.Context(), 60*time.Second)
|
||||
defer cancel()
|
||||
|
||||
// 将文本按段落分割
|
||||
paragraphs := strings.Split(req.Text, "\n")
|
||||
results := make([]string, len(paragraphs))
|
||||
|
||||
@@ -388,7 +372,6 @@ func HandleHcfyTranslate(apiToken string) gin.HandlerFunc {
|
||||
results[i] = result
|
||||
}
|
||||
|
||||
// 构造响应
|
||||
response := HcfyTranslateResponse{
|
||||
Text: req.Text,
|
||||
From: convertBCP47ToHcfyLang(detectedSourceLang),
|
||||
@@ -400,7 +383,6 @@ func HandleHcfyTranslate(apiToken string) gin.HandlerFunc {
|
||||
}
|
||||
}
|
||||
|
||||
// containsChinese 检查文本是否包含中文字符
|
||||
func containsChinese(text string) bool {
|
||||
for _, r := range text {
|
||||
if r >= 0x4E00 && r <= 0x9FFF {
|
||||
@@ -410,18 +392,16 @@ func containsChinese(text string) bool {
|
||||
return false
|
||||
}
|
||||
|
||||
// containsJapanese 检查文本是否包含日文字符
|
||||
func containsJapanese(text string) bool {
|
||||
for _, r := range text {
|
||||
if (r >= 0x3040 && r <= 0x309F) || // 平假名
|
||||
(r >= 0x30A0 && r <= 0x30FF) { // 片假名
|
||||
if (r >= 0x3040 && r <= 0x309F) ||
|
||||
(r >= 0x30A0 && r <= 0x30FF) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// containsKorean 检查文本是否包含韩文字符
|
||||
func containsKorean(text string) bool {
|
||||
for _, r := range text {
|
||||
if r >= 0xAC00 && r <= 0xD7AF {
|
||||
|
||||
@@ -10,38 +10,31 @@ import (
|
||||
"github.com/xxnuo/MTranServer/internal/services"
|
||||
)
|
||||
|
||||
// immeLangToBCP47 沉浸式翻译语言代码转 BCP47
|
||||
// 只列出与 BCP47 不同的映射,其他直接返回原值
|
||||
var immeLangToBCP47 = map[string]string{
|
||||
"zh-CN": "zh-Hans",
|
||||
"zh-TW": "zh-Hant",
|
||||
}
|
||||
|
||||
// convertImmeLangToBCP47 将沉浸式翻译语言代码转换为 BCP47
|
||||
func convertImmeLangToBCP47(immeLang string) string {
|
||||
// 直接从映射表转换
|
||||
|
||||
if bcp47, ok := immeLangToBCP47[immeLang]; ok {
|
||||
return bcp47
|
||||
}
|
||||
|
||||
// 如果不在映射表中,返回原值
|
||||
return immeLang
|
||||
}
|
||||
|
||||
// ImmeTranslateRequest 沉浸式翻译请求
|
||||
type ImmeTranslateRequest struct {
|
||||
SourceLang string `json:"source_lang" binding:"required" example:"en"`
|
||||
TargetLang string `json:"target_lang" binding:"required" example:"zh-CN"`
|
||||
TextList []string `json:"text_list" binding:"required" example:"Hello, world!,Good morning!"`
|
||||
}
|
||||
|
||||
// ImmeTranslation 翻译结果
|
||||
type ImmeTranslation struct {
|
||||
DetectedSourceLang string `json:"detected_source_lang" example:"en"`
|
||||
Text string `json:"text" example:"你好,世界!"`
|
||||
}
|
||||
|
||||
// ImmeTranslateResponse 沉浸式翻译响应
|
||||
type ImmeTranslateResponse struct {
|
||||
Translations []ImmeTranslation `json:"translations"`
|
||||
}
|
||||
@@ -61,7 +54,7 @@ type ImmeTranslateResponse struct {
|
||||
// @Router /imme [post]
|
||||
func HandleImmeTranslate(apiToken string) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
// 检查 token
|
||||
|
||||
if apiToken != "" {
|
||||
token := c.Query("token")
|
||||
if token != apiToken {
|
||||
@@ -81,11 +74,9 @@ func HandleImmeTranslate(apiToken string) gin.HandlerFunc {
|
||||
return
|
||||
}
|
||||
|
||||
// 转换语言代码:沉浸式翻译 -> BCP47
|
||||
sourceLang := convertImmeLangToBCP47(req.SourceLang)
|
||||
targetLang := convertImmeLangToBCP47(req.TargetLang)
|
||||
|
||||
// 批量翻译
|
||||
translations := make([]ImmeTranslation, len(req.TextList))
|
||||
ctx, cancel := context.WithTimeout(c.Request.Context(), 120*time.Second)
|
||||
defer cancel()
|
||||
@@ -98,7 +89,7 @@ func HandleImmeTranslate(apiToken string) gin.HandlerFunc {
|
||||
})
|
||||
return
|
||||
}
|
||||
// 返回沉浸式翻译格式的语言代码
|
||||
|
||||
translations[i] = ImmeTranslation{
|
||||
DetectedSourceLang: req.SourceLang,
|
||||
Text: result,
|
||||
|
||||
@@ -10,14 +10,11 @@ import (
|
||||
"github.com/xxnuo/MTranServer/internal/services"
|
||||
)
|
||||
|
||||
// kissToBCP47 将 Kiss Translator 语言代码转换为 BCP 47 标准
|
||||
// 只列出与 BCP47 不同的映射,其他直接返回原值
|
||||
var kissToBCP47 = map[string]string{
|
||||
"zh-CN": "zh-Hans",
|
||||
"zh-TW": "zh-Hant",
|
||||
}
|
||||
|
||||
// convertKissToBCP47 转换 Kiss 语言代码到 BCP 47
|
||||
func convertKissToBCP47(kissLang string) string {
|
||||
if bcp47, ok := kissToBCP47[kissLang]; ok {
|
||||
return bcp47
|
||||
@@ -25,33 +22,28 @@ func convertKissToBCP47(kissLang string) string {
|
||||
return kissLang
|
||||
}
|
||||
|
||||
// KissTranslateRequest 简约翻译请求(非聚合)
|
||||
type KissTranslateRequest struct {
|
||||
From string `json:"from" binding:"required" example:"en"`
|
||||
To string `json:"to" binding:"required" example:"zh-CN"`
|
||||
Text string `json:"text" binding:"required" example:"Hello, world!"`
|
||||
}
|
||||
|
||||
// KissTranslateResponse 简约翻译响应(非聚合)
|
||||
type KissTranslateResponse struct {
|
||||
Text string `json:"text" example:"你好,世界!"`
|
||||
Src string `json:"src" example:"en"`
|
||||
}
|
||||
|
||||
// KissBatchTranslateRequest 简约翻译请求(聚合)
|
||||
type KissBatchTranslateRequest struct {
|
||||
From string `json:"from" binding:"required" example:"auto"`
|
||||
To string `json:"to" binding:"required" example:"zh-CN"`
|
||||
Texts []string `json:"texts" binding:"required" example:"Hello,World"`
|
||||
}
|
||||
|
||||
// KissBatchTranslateItem 聚合翻译单项响应
|
||||
type KissBatchTranslateItem struct {
|
||||
Text string `json:"text" example:"你好"`
|
||||
Src string `json:"src" example:"en"`
|
||||
}
|
||||
|
||||
// KissBatchTranslateResponse 简约翻译响应(聚合,v2.0.4+格式)
|
||||
type KissBatchTranslateResponse struct {
|
||||
Translations []KissBatchTranslateItem `json:"translations"`
|
||||
}
|
||||
@@ -71,7 +63,7 @@ type KissBatchTranslateResponse struct {
|
||||
// @Router /kiss [post]
|
||||
func HandleKissTranslate(apiToken string) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
// 检查 token
|
||||
|
||||
if apiToken != "" {
|
||||
token := c.GetHeader("KEY")
|
||||
if token != apiToken {
|
||||
@@ -82,7 +74,6 @@ func HandleKissTranslate(apiToken string) gin.HandlerFunc {
|
||||
}
|
||||
}
|
||||
|
||||
// 解析请求体为通用 map 以判断类型
|
||||
var rawReq map[string]interface{}
|
||||
if err := c.ShouldBindJSON(&rawReq); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{
|
||||
@@ -91,9 +82,8 @@ func HandleKissTranslate(apiToken string) gin.HandlerFunc {
|
||||
return
|
||||
}
|
||||
|
||||
// 判断是批量请求还是单个请求
|
||||
if texts, ok := rawReq["texts"].([]interface{}); ok && len(texts) > 0 {
|
||||
// 批量请求
|
||||
|
||||
var batchReq KissBatchTranslateRequest
|
||||
batchReq.From, _ = rawReq["from"].(string)
|
||||
batchReq.To, _ = rawReq["to"].(string)
|
||||
@@ -112,7 +102,6 @@ func HandleKissTranslate(apiToken string) gin.HandlerFunc {
|
||||
return
|
||||
}
|
||||
|
||||
// 单个请求
|
||||
var req KissTranslateRequest
|
||||
req.From, _ = rawReq["from"].(string)
|
||||
req.To, _ = rawReq["to"].(string)
|
||||
@@ -125,11 +114,9 @@ func HandleKissTranslate(apiToken string) gin.HandlerFunc {
|
||||
return
|
||||
}
|
||||
|
||||
// 转换 Kiss 语言代码到 BCP 47
|
||||
fromLang := convertKissToBCP47(req.From)
|
||||
toLang := convertKissToBCP47(req.To)
|
||||
|
||||
// 翻译
|
||||
ctx, cancel := context.WithTimeout(c.Request.Context(), 60*time.Second)
|
||||
defer cancel()
|
||||
|
||||
@@ -148,13 +135,11 @@ func HandleKissTranslate(apiToken string) gin.HandlerFunc {
|
||||
}
|
||||
}
|
||||
|
||||
// handleBatchTranslate 处理批量翻译请求
|
||||
func handleBatchTranslate(c *gin.Context, req KissBatchTranslateRequest) {
|
||||
// 转换 Kiss 语言代码到 BCP 47
|
||||
|
||||
fromLang := convertKissToBCP47(req.From)
|
||||
toLang := convertKissToBCP47(req.To)
|
||||
|
||||
// 翻译所有文本
|
||||
ctx, cancel := context.WithTimeout(c.Request.Context(), 120*time.Second)
|
||||
defer cancel()
|
||||
|
||||
|
||||
@@ -25,7 +25,6 @@ func HandleLanguages(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
// 从 records 中提取所有支持的语言
|
||||
langMap := make(map[string]bool)
|
||||
for _, record := range models.GlobalRecords.Data {
|
||||
langMap[record.FromLang] = true
|
||||
|
||||
@@ -13,7 +13,6 @@ import (
|
||||
func TestHandleLanguages(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
// 初始化测试数据
|
||||
models.GlobalRecords = &models.RecordsData{
|
||||
Data: []models.RecordItem{
|
||||
{FromLang: "en", ToLang: "zh-Hans"},
|
||||
@@ -37,7 +36,6 @@ func TestHandleLanguages(t *testing.T) {
|
||||
func TestHandleLanguagesNotInitialized(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
// 设置为 nil 模拟未初始化
|
||||
originalRecords := models.GlobalRecords
|
||||
models.GlobalRecords = nil
|
||||
defer func() {
|
||||
|
||||
@@ -47,7 +47,6 @@ func HandleTranslate(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
// 使用 TranslateWithPivot 处理可能需要中转的翻译(支持 auto 模式)
|
||||
logger.Debug("Translation request: %s -> %s, text length: %d", req.From, req.To, len(req.Text))
|
||||
ctx, cancel := context.WithTimeout(c.Request.Context(), 60*time.Second)
|
||||
defer cancel()
|
||||
@@ -67,7 +66,6 @@ func HandleTranslate(c *gin.Context) {
|
||||
})
|
||||
}
|
||||
|
||||
// TranslateBatchRequest 批量翻译请求
|
||||
type TranslateBatchRequest struct {
|
||||
From string `json:"from" binding:"required" example:"en"`
|
||||
To string `json:"to" binding:"required" example:"zh-Hans"`
|
||||
@@ -75,7 +73,6 @@ type TranslateBatchRequest struct {
|
||||
HTML bool `json:"html" example:"false"`
|
||||
}
|
||||
|
||||
// TranslateBatchResponse 批量翻译响应
|
||||
type TranslateBatchResponse struct {
|
||||
Results []string `json:"results" example:"你好,世界!,早上好!"`
|
||||
}
|
||||
@@ -103,7 +100,6 @@ func HandleTranslateBatch(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
|
||||
// 批量翻译,使用 TranslateWithPivot 处理可能需要中转的翻译(支持 auto 模式)
|
||||
logger.Debug("Batch translation request: %s -> %s, count: %d", req.From, req.To, len(req.Texts))
|
||||
results := make([]string, len(req.Texts))
|
||||
ctx, cancel := context.WithTimeout(c.Request.Context(), 120*time.Second)
|
||||
|
||||
@@ -7,7 +7,6 @@ import (
|
||||
"strings"
|
||||
)
|
||||
|
||||
// LogLevel 日志级别
|
||||
type LogLevel int
|
||||
|
||||
const (
|
||||
@@ -26,14 +25,13 @@ var (
|
||||
)
|
||||
|
||||
func init() {
|
||||
// 初始化各级别日志器
|
||||
|
||||
debugLogger = log.New(os.Stdout, "[DEBUG] ", log.Ldate|log.Ltime|log.Lshortfile)
|
||||
infoLogger = log.New(os.Stdout, "[INFO] ", log.Ldate|log.Ltime)
|
||||
warnLogger = log.New(os.Stdout, "[WARN] ", log.Ldate|log.Ltime)
|
||||
errorLogger = log.New(os.Stderr, "[ERROR] ", log.Ldate|log.Ltime|log.Lshortfile)
|
||||
}
|
||||
|
||||
// SetLevel 设置日志级别
|
||||
func SetLevel(level string) {
|
||||
switch strings.ToLower(level) {
|
||||
case "debug":
|
||||
@@ -49,7 +47,6 @@ func SetLevel(level string) {
|
||||
}
|
||||
}
|
||||
|
||||
// GetLevel 获取当前日志级别
|
||||
func GetLevel() string {
|
||||
switch currentLevel {
|
||||
case DEBUG:
|
||||
@@ -65,61 +62,51 @@ func GetLevel() string {
|
||||
}
|
||||
}
|
||||
|
||||
// Debug 输出调试日志
|
||||
func Debug(format string, v ...interface{}) {
|
||||
if currentLevel <= DEBUG {
|
||||
debugLogger.Output(2, fmt.Sprintf(format, v...))
|
||||
}
|
||||
}
|
||||
|
||||
// Info 输出信息日志
|
||||
func Info(format string, v ...interface{}) {
|
||||
if currentLevel <= INFO {
|
||||
infoLogger.Output(2, fmt.Sprintf(format, v...))
|
||||
}
|
||||
}
|
||||
|
||||
// Warn 输出警告日志
|
||||
func Warn(format string, v ...interface{}) {
|
||||
if currentLevel <= WARN {
|
||||
warnLogger.Output(2, fmt.Sprintf(format, v...))
|
||||
}
|
||||
}
|
||||
|
||||
// Error 输出错误日志
|
||||
func Error(format string, v ...interface{}) {
|
||||
if currentLevel <= ERROR {
|
||||
errorLogger.Output(2, fmt.Sprintf(format, v...))
|
||||
}
|
||||
}
|
||||
|
||||
// Fatal 输出致命错误日志并退出程序
|
||||
func Fatal(format string, v ...interface{}) {
|
||||
errorLogger.Output(2, fmt.Sprintf(format, v...))
|
||||
os.Exit(1)
|
||||
}
|
||||
|
||||
// Debugf Debug 的别名
|
||||
func Debugf(format string, v ...interface{}) {
|
||||
Debug(format, v...)
|
||||
}
|
||||
|
||||
// Infof Info 的别名
|
||||
func Infof(format string, v ...interface{}) {
|
||||
Info(format, v...)
|
||||
}
|
||||
|
||||
// Warnf Warn 的别名
|
||||
func Warnf(format string, v ...interface{}) {
|
||||
Warn(format, v...)
|
||||
}
|
||||
|
||||
// Errorf Error 的别名
|
||||
func Errorf(format string, v ...interface{}) {
|
||||
Error(format, v...)
|
||||
}
|
||||
|
||||
// Fatalf Fatal 的别名
|
||||
func Fatalf(format string, v ...interface{}) {
|
||||
Fatal(format, v...)
|
||||
}
|
||||
|
||||
@@ -10,13 +10,11 @@ import (
|
||||
"github.com/gorilla/websocket"
|
||||
)
|
||||
|
||||
// WSMessage WebSocket 消息
|
||||
type WSMessage struct {
|
||||
Type string `json:"type"`
|
||||
Data json.RawMessage `json:"data"`
|
||||
}
|
||||
|
||||
// WSResponse WebSocket 响应
|
||||
type WSResponse struct {
|
||||
Type string `json:"type"`
|
||||
Code int `json:"code"`
|
||||
@@ -24,7 +22,6 @@ type WSResponse struct {
|
||||
Data json.RawMessage `json:"data,omitempty"`
|
||||
}
|
||||
|
||||
// PoweronRequest poweron 请求参数
|
||||
type PoweronRequest struct {
|
||||
Path string `json:"path,omitempty"`
|
||||
ModelPath string `json:"model_path,omitempty"`
|
||||
@@ -33,50 +30,41 @@ type PoweronRequest struct {
|
||||
VocabularyPaths []string `json:"vocabulary_paths,omitempty"`
|
||||
}
|
||||
|
||||
// PoweroffRequest poweroff 请求参数
|
||||
type PoweroffRequest struct {
|
||||
Time int `json:"time"`
|
||||
Force bool `json:"force"`
|
||||
}
|
||||
|
||||
// RebootRequest reboot 请求参数
|
||||
type RebootRequest struct {
|
||||
Time int `json:"time"`
|
||||
Force bool `json:"force"`
|
||||
}
|
||||
|
||||
// ComputeRequest compute 请求参数
|
||||
type ComputeRequest struct {
|
||||
Text string `json:"text"`
|
||||
HTML bool `json:"html"`
|
||||
}
|
||||
|
||||
// ReadyResponse ready 响应数据
|
||||
type ReadyResponse struct {
|
||||
Ready bool `json:"ready"`
|
||||
}
|
||||
|
||||
// ComputeResponse compute 响应数据
|
||||
type ComputeResponse struct {
|
||||
TranslatedText string `json:"translated_text"`
|
||||
}
|
||||
|
||||
// PoweronResponse poweron 响应数据
|
||||
type PoweronResponse struct {
|
||||
Message string `json:"message"`
|
||||
}
|
||||
|
||||
// PoweroffResponse poweroff 响应数据
|
||||
type PoweroffResponse struct {
|
||||
Message string `json:"message"`
|
||||
}
|
||||
|
||||
// RebootResponse reboot 响应数据
|
||||
type RebootResponse struct {
|
||||
Message string `json:"message"`
|
||||
}
|
||||
|
||||
// Client WebSocket 客户端
|
||||
type Client struct {
|
||||
url string
|
||||
conn *websocket.Conn
|
||||
@@ -88,24 +76,20 @@ type Client struct {
|
||||
closeOnce sync.Once
|
||||
}
|
||||
|
||||
// ClientOption 客户端配置选项
|
||||
type ClientOption func(*Client)
|
||||
|
||||
// WithTimeout 设置请求超时时间
|
||||
func WithTimeout(timeout time.Duration) ClientOption {
|
||||
return func(c *Client) {
|
||||
c.timeout = timeout
|
||||
}
|
||||
}
|
||||
|
||||
// WithReconnect 设置是否自动重连
|
||||
func WithReconnect(reconnect bool) ClientOption {
|
||||
return func(c *Client) {
|
||||
c.reconnect = reconnect
|
||||
}
|
||||
}
|
||||
|
||||
// NewClient 创建新的 WebSocket 客户端
|
||||
func NewClient(url string, opts ...ClientOption) *Client {
|
||||
c := &Client{
|
||||
url: url,
|
||||
@@ -121,7 +105,6 @@ func NewClient(url string, opts ...ClientOption) *Client {
|
||||
return c
|
||||
}
|
||||
|
||||
// Connect 连接到 WebSocket 服务器
|
||||
func (c *Client) Connect() error {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
@@ -145,7 +128,6 @@ func (c *Client) Connect() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// Close 关闭连接
|
||||
func (c *Client) Close() error {
|
||||
var err error
|
||||
c.closeOnce.Do(func() {
|
||||
@@ -161,14 +143,12 @@ func (c *Client) Close() error {
|
||||
return err
|
||||
}
|
||||
|
||||
// IsConnected 检查是否已连接
|
||||
func (c *Client) IsConnected() bool {
|
||||
c.mu.RLock()
|
||||
defer c.mu.RUnlock()
|
||||
return c.connected
|
||||
}
|
||||
|
||||
// sendRequest 发送请求并接收响应
|
||||
func (c *Client) sendRequest(ctx context.Context, msgType string, data interface{}) (*WSResponse, error) {
|
||||
c.mu.Lock()
|
||||
if !c.connected {
|
||||
@@ -177,7 +157,6 @@ func (c *Client) sendRequest(ctx context.Context, msgType string, data interface
|
||||
}
|
||||
c.mu.Unlock()
|
||||
|
||||
// 序列化数据
|
||||
dataBytes, err := json.Marshal(data)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to marshal data: %w", err)
|
||||
@@ -188,11 +167,9 @@ func (c *Client) sendRequest(ctx context.Context, msgType string, data interface
|
||||
Data: dataBytes,
|
||||
}
|
||||
|
||||
// 创建带超时的 context
|
||||
reqCtx, cancel := context.WithTimeout(ctx, c.timeout)
|
||||
defer cancel()
|
||||
|
||||
// 发送消息
|
||||
c.mu.Lock()
|
||||
if err := c.conn.WriteJSON(msg); err != nil {
|
||||
c.mu.Unlock()
|
||||
@@ -201,7 +178,6 @@ func (c *Client) sendRequest(ctx context.Context, msgType string, data interface
|
||||
}
|
||||
c.mu.Unlock()
|
||||
|
||||
// 接收响应
|
||||
responseChan := make(chan *WSResponse, 1)
|
||||
errChan := make(chan error, 1)
|
||||
|
||||
@@ -228,7 +204,6 @@ func (c *Client) sendRequest(ctx context.Context, msgType string, data interface
|
||||
}
|
||||
}
|
||||
|
||||
// Poweron 加载翻译引擎
|
||||
func (c *Client) Poweron(ctx context.Context, req PoweronRequest) (*PoweronResponse, error) {
|
||||
resp, err := c.sendRequest(ctx, "poweron", req)
|
||||
if err != nil {
|
||||
@@ -249,14 +224,12 @@ func (c *Client) Poweron(ctx context.Context, req PoweronRequest) (*PoweronRespo
|
||||
return &result, nil
|
||||
}
|
||||
|
||||
// Poweroff 关闭服务器
|
||||
func (c *Client) Poweroff(ctx context.Context, req PoweroffRequest) (*PoweroffResponse, error) {
|
||||
resp, err := c.sendRequest(ctx, "poweroff", req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// poweroff 可能返回 1101 (等待任务完成),这也是成功的
|
||||
if resp.Code != 200 && resp.Code != 1101 {
|
||||
return nil, fmt.Errorf("poweroff failed (code %d): %s", resp.Code, resp.Msg)
|
||||
}
|
||||
@@ -273,7 +246,6 @@ func (c *Client) Poweroff(ctx context.Context, req PoweroffRequest) (*PoweroffRe
|
||||
return &result, nil
|
||||
}
|
||||
|
||||
// Reboot 重启引擎
|
||||
func (c *Client) Reboot(ctx context.Context, req RebootRequest) (*RebootResponse, error) {
|
||||
resp, err := c.sendRequest(ctx, "reboot", req)
|
||||
if err != nil {
|
||||
@@ -294,7 +266,6 @@ func (c *Client) Reboot(ctx context.Context, req RebootRequest) (*RebootResponse
|
||||
return &result, nil
|
||||
}
|
||||
|
||||
// Ready 检查引擎是否就绪
|
||||
func (c *Client) Ready(ctx context.Context) (bool, error) {
|
||||
resp, err := c.sendRequest(ctx, "ready", struct{}{})
|
||||
if err != nil {
|
||||
@@ -315,7 +286,6 @@ func (c *Client) Ready(ctx context.Context) (bool, error) {
|
||||
return result.Ready, nil
|
||||
}
|
||||
|
||||
// Compute 翻译文本
|
||||
func (c *Client) Compute(ctx context.Context, req ComputeRequest) (string, error) {
|
||||
resp, err := c.sendRequest(ctx, "compute", req)
|
||||
if err != nil {
|
||||
|
||||
@@ -20,7 +20,6 @@ var upgrader = websocket.Upgrader{
|
||||
},
|
||||
}
|
||||
|
||||
// mockWSServer 创建一个模拟的 WebSocket 服务器
|
||||
func mockWSServer(t *testing.T, handler func(*websocket.Conn)) *httptest.Server {
|
||||
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
conn, err := upgrader.Upgrade(w, r, nil)
|
||||
@@ -31,7 +30,6 @@ func mockWSServer(t *testing.T, handler func(*websocket.Conn)) *httptest.Server
|
||||
return server
|
||||
}
|
||||
|
||||
// handleEcho 回显处理器,用于测试基本连接
|
||||
func handleEcho(conn *websocket.Conn) {
|
||||
for {
|
||||
var msg manager.WSMessage
|
||||
@@ -52,7 +50,6 @@ func handleEcho(conn *websocket.Conn) {
|
||||
}
|
||||
}
|
||||
|
||||
// handlePoweron 模拟 poweron 处理
|
||||
func handlePoweron(conn *websocket.Conn) {
|
||||
var msg manager.WSMessage
|
||||
if err := conn.ReadJSON(&msg); err != nil {
|
||||
@@ -68,7 +65,6 @@ func handlePoweron(conn *websocket.Conn) {
|
||||
Msg: "success",
|
||||
}
|
||||
|
||||
// 检查参数
|
||||
if req.Path == "" && req.ModelPath == "" {
|
||||
resp.Code = 1000
|
||||
resp.Msg = "path is required"
|
||||
@@ -80,7 +76,6 @@ func handlePoweron(conn *websocket.Conn) {
|
||||
conn.WriteJSON(resp)
|
||||
}
|
||||
|
||||
// handleReady 模拟 ready 处理
|
||||
func handleReady(conn *websocket.Conn, ready bool) {
|
||||
var msg manager.WSMessage
|
||||
if err := conn.ReadJSON(&msg); err != nil {
|
||||
@@ -100,7 +95,6 @@ func handleReady(conn *websocket.Conn, ready bool) {
|
||||
conn.WriteJSON(resp)
|
||||
}
|
||||
|
||||
// handleCompute 模拟 compute 处理
|
||||
func handleCompute(conn *websocket.Conn) {
|
||||
var msg manager.WSMessage
|
||||
if err := conn.ReadJSON(&msg); err != nil {
|
||||
@@ -131,7 +125,7 @@ func TestClient_Connect(t *testing.T) {
|
||||
server := mockWSServer(t, handleEcho)
|
||||
defer server.Close()
|
||||
|
||||
wsURL := "ws" + server.URL[4:] // 将 http:// 替换为 ws://
|
||||
wsURL := "ws" + server.URL[4:]
|
||||
|
||||
client := manager.NewClient(wsURL)
|
||||
defer client.Close()
|
||||
@@ -153,7 +147,6 @@ func TestClient_ConnectTwice(t *testing.T) {
|
||||
err := client.Connect()
|
||||
require.NoError(t, err)
|
||||
|
||||
// 第二次连接应该直接返回
|
||||
err = client.Connect()
|
||||
assert.NoError(t, err)
|
||||
assert.True(t, client.IsConnected())
|
||||
@@ -287,7 +280,7 @@ func TestClient_Compute_EmptyText(t *testing.T) {
|
||||
|
||||
func TestClient_Timeout(t *testing.T) {
|
||||
server := mockWSServer(t, func(conn *websocket.Conn) {
|
||||
// 不响应,让客户端超时
|
||||
|
||||
time.Sleep(5 * time.Second)
|
||||
})
|
||||
defer server.Close()
|
||||
@@ -400,7 +393,7 @@ func TestClient_Reboot(t *testing.T) {
|
||||
|
||||
func TestClient_MultipleRequests(t *testing.T) {
|
||||
server := mockWSServer(t, func(conn *websocket.Conn) {
|
||||
// 处理多个请求
|
||||
|
||||
for i := 0; i < 3; i++ {
|
||||
var msg manager.WSMessage
|
||||
if err := conn.ReadJSON(&msg); err != nil {
|
||||
@@ -437,7 +430,6 @@ func TestClient_MultipleRequests(t *testing.T) {
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// 发送多个请求
|
||||
for i := 1; i <= 3; i++ {
|
||||
result, err := client.Compute(ctx, manager.ComputeRequest{
|
||||
Text: "Test " + string(rune('0'+i)),
|
||||
|
||||
@@ -25,7 +25,6 @@ var (
|
||||
workerBinaryMu sync.Mutex
|
||||
)
|
||||
|
||||
// WorkerArgs 包含工作进程的配置
|
||||
type WorkerArgs struct {
|
||||
Host string
|
||||
Port int
|
||||
@@ -35,10 +34,9 @@ type WorkerArgs struct {
|
||||
EnableWebSocket bool
|
||||
GRPCUnixSocket string
|
||||
LogLevel string
|
||||
BinaryPath string // 写入工作程序二进制文件的路径,如果为空则使用 ConfigDir/bin/mtrancore
|
||||
BinaryPath string
|
||||
}
|
||||
|
||||
// NewWorkerArgs 创建一个新的 WorkerArgs 实例,使用默认值
|
||||
func NewWorkerArgs() *WorkerArgs {
|
||||
return &WorkerArgs{
|
||||
Host: "127.0.0.1",
|
||||
@@ -52,22 +50,20 @@ func NewWorkerArgs() *WorkerArgs {
|
||||
}
|
||||
}
|
||||
|
||||
// Worker 管理使用 overseer 的工作进程
|
||||
type Worker struct {
|
||||
args *WorkerArgs
|
||||
overseer *overseer.Overseer
|
||||
id string
|
||||
binaryPath string // 实际写入二进制文件的路径
|
||||
binaryPath string
|
||||
mu sync.RWMutex
|
||||
logChan chan *overseer.LogMsg
|
||||
stateChan chan *overseer.ProcessJSON
|
||||
logs []string
|
||||
maxLogs int
|
||||
done chan struct{} // 用于通知 goroutine 退出
|
||||
done chan struct{}
|
||||
wg sync.WaitGroup
|
||||
}
|
||||
|
||||
// NewWorker 创建一个新的 Worker 实例
|
||||
func NewWorker(args *WorkerArgs) *Worker {
|
||||
binaryPath := args.BinaryPath
|
||||
if binaryPath == "" {
|
||||
@@ -79,7 +75,6 @@ func NewWorker(args *WorkerArgs) *Worker {
|
||||
binaryPath = filepath.Join(cfg.ConfigDir, "bin", binaryName)
|
||||
}
|
||||
|
||||
// 根据二进制文件路径和端口生成唯一的 worker ID
|
||||
workerID := fmt.Sprintf("mtran-worker-%d", args.Port)
|
||||
|
||||
w := &Worker{
|
||||
@@ -94,23 +89,19 @@ func NewWorker(args *WorkerArgs) *Worker {
|
||||
done: make(chan struct{}),
|
||||
}
|
||||
|
||||
// 订阅日志和状态变化
|
||||
w.overseer.WatchLogs(w.logChan)
|
||||
w.overseer.WatchState(w.stateChan)
|
||||
|
||||
// 启动日志收集器
|
||||
w.wg.Add(1)
|
||||
go w.collectLogs()
|
||||
|
||||
return w
|
||||
}
|
||||
|
||||
// EnsureWorkerBinary 提取嵌入的工作程序二进制文件到指定路径
|
||||
func EnsureWorkerBinary(cfg *config.Config) error {
|
||||
workerBinaryMu.Lock()
|
||||
defer workerBinaryMu.Unlock()
|
||||
|
||||
// 如果已经初始化过,直接返回
|
||||
if workerBinaryInitialized {
|
||||
return nil
|
||||
}
|
||||
@@ -121,12 +112,11 @@ func EnsureWorkerBinary(cfg *config.Config) error {
|
||||
}
|
||||
binaryPath := filepath.Join(cfg.ConfigDir, "bin", binaryName)
|
||||
|
||||
// 检查二进制文件是否已存在并且匹配哈希
|
||||
if data, err := os.ReadFile(binaryPath); err == nil {
|
||||
// 二进制文件存在,计算其哈希并比较
|
||||
|
||||
existingHash := fmt.Sprintf("%x", bin.ComputeHash(data))
|
||||
if existingHash == bin.WorkerHash {
|
||||
// 哈希匹配,二进制文件是最新的
|
||||
|
||||
logger.Debug("Worker binary already exists and is up to date")
|
||||
workerBinaryInitialized = true
|
||||
return nil
|
||||
@@ -134,13 +124,12 @@ func EnsureWorkerBinary(cfg *config.Config) error {
|
||||
logger.Info("Worker binary hash mismatch, updating...")
|
||||
}
|
||||
|
||||
// 确保父目录存在
|
||||
if err := os.MkdirAll(filepath.Dir(binaryPath), 0755); err != nil {
|
||||
return fmt.Errorf("failed to create directory for worker binary: %w", err)
|
||||
}
|
||||
|
||||
logger.Info("Extracting worker binary to %s", binaryPath)
|
||||
// 写入嵌入的二进制文件
|
||||
|
||||
if err := os.WriteFile(binaryPath, bin.WorkerBinary, 0755); err != nil {
|
||||
return fmt.Errorf("failed to write worker binary: %w", err)
|
||||
}
|
||||
@@ -150,7 +139,6 @@ func EnsureWorkerBinary(cfg *config.Config) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// buildArgs 构建工作程序的命令行参数
|
||||
func (w *Worker) buildArgs() []string {
|
||||
args := []string{
|
||||
"--host", w.args.Host,
|
||||
@@ -191,53 +179,43 @@ func (w *Worker) buildArgs() []string {
|
||||
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 := os.Stat(w.binaryPath); err != nil {
|
||||
return fmt.Errorf("worker binary not found at %s: %w", w.binaryPath, err)
|
||||
}
|
||||
|
||||
// 构建命令行参数
|
||||
args := w.buildArgs()
|
||||
|
||||
// 将工作进程添加到 overseer
|
||||
// 注意: overseer.Add 接受 []string 作为单个参数,而不是可变参数字符串
|
||||
logger.Debug("Starting worker %s on port %d", w.id, w.args.Port)
|
||||
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 // 默认不自动重启
|
||||
cmd.RetryTimes = 0
|
||||
|
||||
// 在 goroutine 中启动监督
|
||||
go w.overseer.Supervise(w.id)
|
||||
|
||||
// 等待一段时间让进程启动
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
|
||||
logger.Debug("Worker %s started", w.id)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Stop 优雅地停止工作进程
|
||||
func (w *Worker) Stop() error {
|
||||
w.mu.Lock()
|
||||
defer w.mu.Unlock()
|
||||
@@ -251,13 +229,11 @@ func (w *Worker) Stop() error {
|
||||
return fmt.Errorf("worker not running")
|
||||
}
|
||||
|
||||
// 优雅地停止进程
|
||||
logger.Debug("Stopping worker %s", w.id)
|
||||
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()
|
||||
@@ -265,7 +241,7 @@ func (w *Worker) Stop() error {
|
||||
for {
|
||||
select {
|
||||
case <-timeout:
|
||||
// 如果优雅停止失败,则强制杀死进程
|
||||
|
||||
logger.Warn("Worker %s stop timeout, forcing kill", w.id)
|
||||
w.overseer.Signal(w.id, syscall.SIGKILL)
|
||||
return fmt.Errorf("worker stop timeout, forced kill")
|
||||
@@ -279,9 +255,8 @@ func (w *Worker) Stop() error {
|
||||
}
|
||||
}
|
||||
|
||||
// Restart 重启工作进程
|
||||
func (w *Worker) Restart() error {
|
||||
// 如果正在运行,则停止
|
||||
|
||||
if w.overseer.HasProc(w.id) {
|
||||
status := w.overseer.Status(w.id)
|
||||
if status != nil && status.State == "running" {
|
||||
@@ -289,18 +264,15 @@ func (w *Worker) Restart() error {
|
||||
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()
|
||||
@@ -317,7 +289,6 @@ func (w *Worker) Status() string {
|
||||
return status.State
|
||||
}
|
||||
|
||||
// GetDetailedStatus 返回详细的状态信息
|
||||
func (w *Worker) GetDetailedStatus() *overseer.ProcessJSON {
|
||||
w.mu.RLock()
|
||||
defer w.mu.RUnlock()
|
||||
@@ -329,25 +300,22 @@ func (w *Worker) GetDetailedStatus() *overseer.ProcessJSON {
|
||||
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() {
|
||||
defer w.wg.Done()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-w.done:
|
||||
// 收到退出信号,清空剩余的日志消息后退出
|
||||
|
||||
for {
|
||||
select {
|
||||
case msg, ok := <-w.logChan:
|
||||
@@ -388,7 +356,7 @@ func (w *Worker) collectLogs() {
|
||||
return
|
||||
}
|
||||
w.mu.Lock()
|
||||
// 格式化日志消息
|
||||
|
||||
logType := "INFO"
|
||||
if msg.Type == 1 {
|
||||
logType = "ERROR"
|
||||
@@ -397,7 +365,6 @@ func (w *Worker) collectLogs() {
|
||||
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:]
|
||||
}
|
||||
@@ -408,7 +375,7 @@ func (w *Worker) collectLogs() {
|
||||
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)
|
||||
@@ -421,12 +388,10 @@ func (w *Worker) collectLogs() {
|
||||
}
|
||||
}
|
||||
|
||||
// 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()
|
||||
@@ -438,22 +403,19 @@ func (w *Worker) Signal(sig syscall.Signal) error {
|
||||
return w.overseer.Signal(w.id, sig)
|
||||
}
|
||||
|
||||
// Cleanup 清理资源
|
||||
func (w *Worker) Cleanup() error {
|
||||
w.mu.Lock()
|
||||
|
||||
var errs []error
|
||||
|
||||
// 如果正在运行,则停止
|
||||
if w.overseer.HasProc(w.id) {
|
||||
status := w.overseer.Status(w.id)
|
||||
if status != nil && status.State == "running" {
|
||||
// 先尝试优雅停止
|
||||
|
||||
if err := w.overseer.Stop(w.id); err != nil {
|
||||
errs = append(errs, fmt.Errorf("failed to stop worker gracefully: %w", err))
|
||||
}
|
||||
|
||||
// 等待进程停止,超时后强制杀死
|
||||
timeout := time.After(5 * time.Second)
|
||||
ticker := time.NewTicker(100 * time.Millisecond)
|
||||
defer ticker.Stop()
|
||||
@@ -462,7 +424,7 @@ func (w *Worker) Cleanup() error {
|
||||
for {
|
||||
select {
|
||||
case <-timeout:
|
||||
// 强制杀死进程
|
||||
|
||||
if err := w.overseer.Signal(w.id, syscall.SIGKILL); err != nil {
|
||||
errs = append(errs, fmt.Errorf("failed to kill worker: %w", err))
|
||||
}
|
||||
@@ -477,30 +439,23 @@ func (w *Worker) Cleanup() error {
|
||||
}
|
||||
}
|
||||
|
||||
// 从 overseer 中移除进程
|
||||
w.overseer.Remove(w.id)
|
||||
}
|
||||
|
||||
// 取消订阅通道(如果还未取消)
|
||||
w.overseer.UnWatchLogs(w.logChan)
|
||||
w.overseer.UnWatchState(w.stateChan)
|
||||
|
||||
// 通知 collectLogs goroutine 退出
|
||||
select {
|
||||
case <-w.done:
|
||||
// 已经关闭
|
||||
|
||||
default:
|
||||
close(w.done)
|
||||
}
|
||||
|
||||
w.mu.Unlock()
|
||||
|
||||
// 等待 collectLogs goroutine 退出
|
||||
w.wg.Wait()
|
||||
|
||||
// 注意: 我们不在这里删除工作程序二进制文件,因为它可能被共享
|
||||
// 或者用户有意放置在特定位置
|
||||
|
||||
if len(errs) > 0 {
|
||||
return fmt.Errorf("cleanup errors: %v", errs)
|
||||
}
|
||||
|
||||
@@ -9,13 +9,11 @@ import (
|
||||
"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 := manager.NewWorkerArgs()
|
||||
args.Host = "127.0.0.1"
|
||||
port, err := utils.GetFreePort()
|
||||
@@ -27,43 +25,36 @@ func TestBasicUsage(t *testing.T) {
|
||||
args.EnableHTTP = true
|
||||
args.LogLevel = "debug"
|
||||
args.WorkDir = "."
|
||||
// Create a new worker
|
||||
|
||||
worker := manager.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)
|
||||
@@ -71,7 +62,6 @@ func TestBasicUsage(t *testing.T) {
|
||||
|
||||
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)
|
||||
@@ -80,7 +70,6 @@ func TestBasicUsage(t *testing.T) {
|
||||
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")
|
||||
@@ -92,33 +81,27 @@ func TestLifecycle(t *testing.T) {
|
||||
t.Fatalf("Failed to get free port: %v", err)
|
||||
}
|
||||
args.Port = port
|
||||
// Binary will be written to /tmp by default
|
||||
|
||||
worker := manager.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))
|
||||
}
|
||||
@@ -126,7 +109,6 @@ func TestWorkerHash(t *testing.T) {
|
||||
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 := manager.NewWorkerArgs()
|
||||
port, err := utils.GetFreePort()
|
||||
if err != nil {
|
||||
@@ -136,7 +118,6 @@ func TestWorkerHash(t *testing.T) {
|
||||
worker := manager.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)
|
||||
}
|
||||
@@ -144,7 +125,6 @@ func TestWorkerHash(t *testing.T) {
|
||||
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)
|
||||
}
|
||||
@@ -155,20 +135,18 @@ func TestWorkerHash(t *testing.T) {
|
||||
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 := manager.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
|
||||
args.BinaryPath = "/tmp/custom-mtran-worker"
|
||||
worker := manager.NewWorker(args)
|
||||
defer worker.Cleanup()
|
||||
|
||||
@@ -186,13 +164,11 @@ func TestCustomBinaryPath(t *testing.T) {
|
||||
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([]*manager.Worker, 0, 3)
|
||||
|
||||
for i := 0; i < 3; i++ {
|
||||
@@ -205,7 +181,6 @@ func TestMultipleWorkers(t *testing.T) {
|
||||
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 := manager.NewWorker(args)
|
||||
workers = append(workers, worker)
|
||||
@@ -216,10 +191,8 @@ func TestMultipleWorkers(t *testing.T) {
|
||||
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)
|
||||
@@ -230,7 +203,6 @@ func TestMultipleWorkers(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// Stop all workers
|
||||
for i, worker := range workers {
|
||||
if err := worker.Stop(); err != nil {
|
||||
t.Errorf("Failed to stop worker %d: %v", i, err)
|
||||
|
||||
@@ -7,7 +7,6 @@ import (
|
||||
"time"
|
||||
)
|
||||
|
||||
// Manager 管理 Worker 和 Client,提供统一的翻译服务接口
|
||||
type Manager struct {
|
||||
worker *Worker
|
||||
client *Client
|
||||
@@ -15,12 +14,10 @@ type Manager struct {
|
||||
url string
|
||||
}
|
||||
|
||||
// ManagerOption 管理器配置选项
|
||||
type ManagerOption func(*Manager)
|
||||
|
||||
// NewManager 创建新的 Manager 实例
|
||||
func NewManager(args *WorkerArgs, opts ...ManagerOption) *Manager {
|
||||
// 构建 WebSocket URL
|
||||
|
||||
url := fmt.Sprintf("ws://%s:%d/ws", args.Host, args.Port)
|
||||
|
||||
m := &Manager{
|
||||
@@ -35,17 +32,14 @@ func NewManager(args *WorkerArgs, opts ...ManagerOption) *Manager {
|
||||
return m
|
||||
}
|
||||
|
||||
// Start 启动 Worker 并连接 Client
|
||||
func (m *Manager) Start() error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
// 启动 Worker
|
||||
if err := m.worker.Start(); err != nil {
|
||||
return fmt.Errorf("failed to start worker: %w", err)
|
||||
}
|
||||
|
||||
// 等待 Worker 启动
|
||||
timeout := time.After(10 * time.Second)
|
||||
ticker := time.NewTicker(100 * time.Millisecond)
|
||||
defer ticker.Stop()
|
||||
@@ -57,7 +51,7 @@ func (m *Manager) Start() error {
|
||||
return fmt.Errorf("worker start timeout")
|
||||
case <-ticker.C:
|
||||
if m.worker.IsRunning() {
|
||||
// Worker 已运行,创建并连接 Client
|
||||
|
||||
m.client = NewClient(m.url)
|
||||
if err := m.client.Connect(); err != nil {
|
||||
m.worker.Stop()
|
||||
@@ -69,14 +63,12 @@ func (m *Manager) Start() error {
|
||||
}
|
||||
}
|
||||
|
||||
// Stop 停止 Manager
|
||||
func (m *Manager) Stop() error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
var errs []error
|
||||
|
||||
// 关闭 Client
|
||||
if m.client != nil {
|
||||
if err := m.client.Close(); err != nil {
|
||||
errs = append(errs, fmt.Errorf("failed to close client: %w", err))
|
||||
@@ -84,7 +76,6 @@ func (m *Manager) Stop() error {
|
||||
m.client = nil
|
||||
}
|
||||
|
||||
// 停止 Worker
|
||||
if m.worker != nil {
|
||||
if err := m.worker.Stop(); err != nil {
|
||||
errs = append(errs, fmt.Errorf("failed to stop worker: %w", err))
|
||||
@@ -98,7 +89,6 @@ func (m *Manager) Stop() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// Restart 重启 Manager
|
||||
func (m *Manager) Restart() error {
|
||||
if err := m.Stop(); err != nil {
|
||||
return fmt.Errorf("failed to stop: %w", err)
|
||||
@@ -109,14 +99,12 @@ func (m *Manager) Restart() error {
|
||||
return m.Start()
|
||||
}
|
||||
|
||||
// Cleanup 清理资源
|
||||
func (m *Manager) Cleanup() error {
|
||||
m.mu.Lock()
|
||||
defer m.mu.Unlock()
|
||||
|
||||
var errs []error
|
||||
|
||||
// 关闭 Client(忽略错误,继续清理)
|
||||
if m.client != nil {
|
||||
if err := m.client.Close(); err != nil {
|
||||
errs = append(errs, fmt.Errorf("failed to close client: %w", err))
|
||||
@@ -124,12 +112,11 @@ func (m *Manager) Cleanup() error {
|
||||
m.client = nil
|
||||
}
|
||||
|
||||
// 清理 Worker(即使 client 关闭失败也要清理)
|
||||
if m.worker != nil {
|
||||
if err := m.worker.Cleanup(); err != nil {
|
||||
errs = append(errs, fmt.Errorf("failed to cleanup worker: %w", err))
|
||||
}
|
||||
// 不设置为 nil,因为 worker 结构体可能还需要保留
|
||||
|
||||
}
|
||||
|
||||
if len(errs) > 0 {
|
||||
@@ -139,7 +126,6 @@ func (m *Manager) Cleanup() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// IsRunning 检查 Manager 是否正在运行
|
||||
func (m *Manager) IsRunning() bool {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
@@ -147,7 +133,6 @@ func (m *Manager) IsRunning() bool {
|
||||
return m.worker != nil && m.worker.IsRunning() && m.client != nil && m.client.IsConnected()
|
||||
}
|
||||
|
||||
// Status 返回 Worker 状态
|
||||
func (m *Manager) Status() string {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
@@ -159,7 +144,6 @@ func (m *Manager) Status() string {
|
||||
return m.worker.Status()
|
||||
}
|
||||
|
||||
// Logs 返回 Worker 日志
|
||||
func (m *Manager) Logs() []string {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
@@ -171,7 +155,6 @@ func (m *Manager) Logs() []string {
|
||||
return m.worker.Logs()
|
||||
}
|
||||
|
||||
// Poweron 加载翻译引擎
|
||||
func (m *Manager) Poweron(ctx context.Context, req PoweronRequest) (*PoweronResponse, error) {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
@@ -183,7 +166,6 @@ func (m *Manager) Poweron(ctx context.Context, req PoweronRequest) (*PoweronResp
|
||||
return m.client.Poweron(ctx, req)
|
||||
}
|
||||
|
||||
// Poweroff 关闭翻译引擎
|
||||
func (m *Manager) Poweroff(ctx context.Context, req PoweroffRequest) (*PoweroffResponse, error) {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
@@ -195,7 +177,6 @@ func (m *Manager) Poweroff(ctx context.Context, req PoweroffRequest) (*PoweroffR
|
||||
return m.client.Poweroff(ctx, req)
|
||||
}
|
||||
|
||||
// Reboot 重启翻译引擎
|
||||
func (m *Manager) Reboot(ctx context.Context, req RebootRequest) (*RebootResponse, error) {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
@@ -207,7 +188,6 @@ func (m *Manager) Reboot(ctx context.Context, req RebootRequest) (*RebootRespons
|
||||
return m.client.Reboot(ctx, req)
|
||||
}
|
||||
|
||||
// Ready 检查翻译引擎是否就绪
|
||||
func (m *Manager) Ready(ctx context.Context) (bool, error) {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
@@ -219,7 +199,6 @@ func (m *Manager) Ready(ctx context.Context) (bool, error) {
|
||||
return m.client.Ready(ctx)
|
||||
}
|
||||
|
||||
// Compute 翻译文本
|
||||
func (m *Manager) Compute(ctx context.Context, req ComputeRequest) (string, error) {
|
||||
m.mu.RLock()
|
||||
defer m.mu.RUnlock()
|
||||
@@ -231,7 +210,6 @@ func (m *Manager) Compute(ctx context.Context, req ComputeRequest) (string, erro
|
||||
return m.client.Compute(ctx, req)
|
||||
}
|
||||
|
||||
// Translate 翻译文本(简化接口)
|
||||
func (m *Manager) Translate(ctx context.Context, text string) (string, error) {
|
||||
return m.Compute(ctx, ComputeRequest{
|
||||
Text: text,
|
||||
@@ -239,7 +217,6 @@ func (m *Manager) Translate(ctx context.Context, text string) (string, error) {
|
||||
})
|
||||
}
|
||||
|
||||
// TranslateHTML 翻译 HTML 文本
|
||||
func (m *Manager) TranslateHTML(ctx context.Context, html string) (string, error) {
|
||||
return m.Compute(ctx, ComputeRequest{
|
||||
Text: html,
|
||||
|
||||
@@ -27,15 +27,12 @@ func TestManager_StartStop(t *testing.T) {
|
||||
mgr := manager.NewManager(args)
|
||||
defer mgr.Cleanup()
|
||||
|
||||
// 启动 Manager
|
||||
err = mgr.Start()
|
||||
require.NoError(t, err)
|
||||
|
||||
// 检查状态
|
||||
assert.True(t, mgr.IsRunning())
|
||||
assert.Equal(t, "running", mgr.Status())
|
||||
|
||||
// 停止 Manager
|
||||
err = mgr.Stop()
|
||||
require.NoError(t, err)
|
||||
|
||||
@@ -57,17 +54,14 @@ func TestManager_Restart(t *testing.T) {
|
||||
mgr := manager.NewManager(args)
|
||||
defer mgr.Cleanup()
|
||||
|
||||
// 启动
|
||||
err = mgr.Start()
|
||||
require.NoError(t, err)
|
||||
assert.True(t, mgr.IsRunning())
|
||||
|
||||
// 重启
|
||||
err = mgr.Restart()
|
||||
require.NoError(t, err)
|
||||
assert.True(t, mgr.IsRunning())
|
||||
|
||||
// 停止
|
||||
err = mgr.Stop()
|
||||
require.NoError(t, err)
|
||||
}
|
||||
@@ -122,7 +116,7 @@ func TestManager_Ready(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
ready, err := mgr.Ready(ctx)
|
||||
require.NoError(t, err)
|
||||
assert.False(t, ready) // 未加载模型,应该返回 false
|
||||
assert.False(t, ready)
|
||||
}
|
||||
|
||||
func TestManager_PoweronPoweroff(t *testing.T) {
|
||||
@@ -136,7 +130,7 @@ func TestManager_PoweronPoweroff(t *testing.T) {
|
||||
args.Port = port
|
||||
args.Host = "127.0.0.1"
|
||||
args.EnableWebSocket = true
|
||||
args.WorkDir = "../../testdata" // 假设有测试数据目录
|
||||
args.WorkDir = "../../testdata"
|
||||
|
||||
mgr := manager.NewManager(args)
|
||||
defer mgr.Cleanup()
|
||||
@@ -149,18 +143,16 @@ func TestManager_PoweronPoweroff(t *testing.T) {
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// 测试 Poweron(这里会失败,因为没有真实的模型文件)
|
||||
_, err = mgr.Poweron(ctx, manager.PoweronRequest{
|
||||
Path: "nonexistent",
|
||||
})
|
||||
assert.Error(t, err) // 预期失败
|
||||
assert.Error(t, err)
|
||||
|
||||
// 测试 Poweroff
|
||||
_, err = mgr.Poweroff(ctx, manager.PoweroffRequest{
|
||||
Time: 0,
|
||||
Force: true,
|
||||
})
|
||||
// Poweroff 可能成功也可能失败,取决于引擎状态
|
||||
|
||||
t.Logf("Poweroff result: %v", err)
|
||||
}
|
||||
|
||||
@@ -187,12 +179,11 @@ func TestManager_Reboot(t *testing.T) {
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// 测试 Reboot(未加载引擎时会失败)
|
||||
_, err = mgr.Reboot(ctx, manager.RebootRequest{
|
||||
Time: 0,
|
||||
Force: false,
|
||||
})
|
||||
assert.Error(t, err) // 预期失败
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestManager_Compute(t *testing.T) {
|
||||
@@ -218,12 +209,11 @@ func TestManager_Compute(t *testing.T) {
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// 测试 Compute(未加载引擎时会失败)
|
||||
_, err = mgr.Compute(ctx, manager.ComputeRequest{
|
||||
Text: "Hello",
|
||||
HTML: false,
|
||||
})
|
||||
assert.Error(t, err) // 预期失败,因为引擎未加载
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestManager_Translate(t *testing.T) {
|
||||
@@ -249,9 +239,8 @@ func TestManager_Translate(t *testing.T) {
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// 测试 Translate(未加载引擎时会失败)
|
||||
_, err = mgr.Translate(ctx, "Hello")
|
||||
assert.Error(t, err) // 预期失败
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestManager_TranslateHTML(t *testing.T) {
|
||||
@@ -277,9 +266,8 @@ func TestManager_TranslateHTML(t *testing.T) {
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// 测试 TranslateHTML(未加载引擎时会失败)
|
||||
_, err = mgr.TranslateHTML(ctx, "<p>Hello</p>")
|
||||
assert.Error(t, err) // 预期失败
|
||||
assert.Error(t, err)
|
||||
}
|
||||
|
||||
func TestManager_MultipleManagers(t *testing.T) {
|
||||
@@ -308,12 +296,10 @@ func TestManager_MultipleManagers(t *testing.T) {
|
||||
|
||||
time.Sleep(2 * time.Second)
|
||||
|
||||
// 验证所有 Manager 都在运行
|
||||
for i, mgr := range managers {
|
||||
assert.True(t, mgr.IsRunning(), "Manager %d should be running", i)
|
||||
}
|
||||
|
||||
// 停止所有 Manager
|
||||
for i, mgr := range managers {
|
||||
err := mgr.Stop()
|
||||
assert.NoError(t, err)
|
||||
@@ -327,7 +313,6 @@ func TestManager_NotStarted(t *testing.T) {
|
||||
mgr := manager.NewManager(args)
|
||||
defer mgr.Cleanup()
|
||||
|
||||
// 未启动时调用方法应该返回错误
|
||||
ctx := context.Background()
|
||||
|
||||
_, err := mgr.Ready(ctx)
|
||||
@@ -344,8 +329,6 @@ func TestManager_FullWorkflow(t *testing.T) {
|
||||
t.Skip("Skipping integration test in short mode")
|
||||
}
|
||||
|
||||
// 这是一个完整的工作流测试
|
||||
// 需要真实的模型文件才能完全通过
|
||||
t.Skip("Requires real model files")
|
||||
|
||||
args := manager.NewWorkerArgs()
|
||||
@@ -359,7 +342,6 @@ func TestManager_FullWorkflow(t *testing.T) {
|
||||
mgr := manager.NewManager(args)
|
||||
defer mgr.Cleanup()
|
||||
|
||||
// 1. 启动
|
||||
err = mgr.Start()
|
||||
require.NoError(t, err)
|
||||
defer mgr.Stop()
|
||||
@@ -368,38 +350,32 @@ func TestManager_FullWorkflow(t *testing.T) {
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// 2. 检查就绪状态
|
||||
ready, err := mgr.Ready(ctx)
|
||||
require.NoError(t, err)
|
||||
assert.False(t, ready)
|
||||
|
||||
// 3. 加载模型
|
||||
resp, err := mgr.Poweron(ctx, manager.PoweronRequest{
|
||||
Path: "path/to/model",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.NotNil(t, resp)
|
||||
|
||||
// 4. 等待引擎就绪
|
||||
time.Sleep(2 * time.Second)
|
||||
|
||||
ready, err = mgr.Ready(ctx)
|
||||
require.NoError(t, err)
|
||||
assert.True(t, ready)
|
||||
|
||||
// 5. 翻译文本
|
||||
result, err := mgr.Translate(ctx, "Hello, world!")
|
||||
require.NoError(t, err)
|
||||
assert.NotEmpty(t, result)
|
||||
t.Logf("Translation result: %s", result)
|
||||
|
||||
// 6. 翻译 HTML
|
||||
htmlResult, err := mgr.TranslateHTML(ctx, "<p>Hello, world!</p>")
|
||||
require.NoError(t, err)
|
||||
assert.NotEmpty(t, htmlResult)
|
||||
t.Logf("HTML translation result: %s", htmlResult)
|
||||
|
||||
// 7. 重启引擎
|
||||
rebootResp, err := mgr.Reboot(ctx, manager.RebootRequest{
|
||||
Time: 0,
|
||||
Force: false,
|
||||
@@ -407,7 +383,6 @@ func TestManager_FullWorkflow(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
assert.NotNil(t, rebootResp)
|
||||
|
||||
// 8. 关闭引擎
|
||||
poweroffResp, err := mgr.Poweroff(ctx, manager.PoweroffRequest{
|
||||
Time: 0,
|
||||
Force: true,
|
||||
|
||||
@@ -8,16 +8,14 @@ import (
|
||||
"github.com/xxnuo/MTranServer/internal/logger"
|
||||
)
|
||||
|
||||
// Auth 认证中间件
|
||||
func Auth(apiToken string) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
// 如果未设置 token,则不进行认证
|
||||
|
||||
if apiToken == "" {
|
||||
c.Next()
|
||||
return
|
||||
}
|
||||
|
||||
// 从 Authorization header 获取 token
|
||||
token := c.GetHeader("Authorization")
|
||||
if token != "" {
|
||||
token = strings.TrimPrefix(token, "Bearer ")
|
||||
@@ -25,7 +23,6 @@ func Auth(apiToken string) gin.HandlerFunc {
|
||||
token = c.Query("token")
|
||||
}
|
||||
|
||||
// 验证 token
|
||||
if token != apiToken {
|
||||
logger.Warn("Unauthorized access attempt from %s to %s", c.ClientIP(), c.Request.URL.Path)
|
||||
c.JSON(http.StatusUnauthorized, gin.H{
|
||||
|
||||
@@ -6,7 +6,6 @@ import (
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
// CORS 中间件
|
||||
func CORS() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
c.Writer.Header().Set("Access-Control-Allow-Origin", "*")
|
||||
|
||||
@@ -8,18 +8,15 @@ import (
|
||||
"github.com/xxnuo/MTranServer/internal/logger"
|
||||
)
|
||||
|
||||
// Logger 自定义日志中间件,将 Gin 的日志输出到我们的日志系统
|
||||
func Logger() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
// 开始时间
|
||||
|
||||
start := time.Now()
|
||||
path := c.Request.URL.Path
|
||||
raw := c.Request.URL.RawQuery
|
||||
|
||||
// 处理请求
|
||||
c.Next()
|
||||
|
||||
// 结束时间
|
||||
end := time.Now()
|
||||
latency := end.Sub(start)
|
||||
|
||||
@@ -32,7 +29,6 @@ func Logger() gin.HandlerFunc {
|
||||
path = path + "?" + raw
|
||||
}
|
||||
|
||||
// 根据状态码选择日志级别
|
||||
logFunc := logger.Info
|
||||
if statusCode >= 500 {
|
||||
logFunc = logger.Error
|
||||
@@ -40,7 +36,6 @@ func Logger() gin.HandlerFunc {
|
||||
logFunc = logger.Warn
|
||||
}
|
||||
|
||||
// 构建日志消息
|
||||
msg := fmt.Sprintf("%s %s %d %v %s",
|
||||
method,
|
||||
path,
|
||||
@@ -57,7 +52,6 @@ func Logger() gin.HandlerFunc {
|
||||
}
|
||||
}
|
||||
|
||||
// Recovery 自定义恢复中间件,将 panic 信息输出到我们的日志系统
|
||||
func Recovery() gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
defer func() {
|
||||
|
||||
@@ -19,12 +19,10 @@ const (
|
||||
AttachmentsBaseUrl = "https://firefox-settings-attachments.cdn.mozilla.net"
|
||||
)
|
||||
|
||||
// RecordsData records.json 的结构
|
||||
type RecordsData struct {
|
||||
Data []RecordItem `json:"data"`
|
||||
}
|
||||
|
||||
// RecordItem 单个记录项
|
||||
type RecordItem struct {
|
||||
Hash string `json:"hash,omitempty"`
|
||||
Name string `json:"name"`
|
||||
@@ -37,7 +35,6 @@ type RecordItem struct {
|
||||
ID string `json:"id"`
|
||||
}
|
||||
|
||||
// Attachment 附件信息
|
||||
type Attachment struct {
|
||||
Hash string `json:"hash"`
|
||||
Size int64 `json:"size"`
|
||||
@@ -50,7 +47,6 @@ var (
|
||||
GlobalRecords *RecordsData
|
||||
)
|
||||
|
||||
// GetLanguagePairs 获取所有可用的语言对
|
||||
func (r *RecordsData) GetLanguagePairs() []string {
|
||||
pairMap := make(map[string]bool)
|
||||
for _, record := range r.Data {
|
||||
@@ -65,7 +61,6 @@ func (r *RecordsData) GetLanguagePairs() []string {
|
||||
return pairs
|
||||
}
|
||||
|
||||
// HasLanguagePair 检查是否支持指定的语言对
|
||||
func (r *RecordsData) HasLanguagePair(fromLang, toLang string) bool {
|
||||
for _, record := range r.Data {
|
||||
if record.FromLang == fromLang && record.ToLang == toLang {
|
||||
@@ -75,7 +70,6 @@ func (r *RecordsData) HasLanguagePair(fromLang, toLang string) bool {
|
||||
return false
|
||||
}
|
||||
|
||||
// GetVersions 获取指定语言对的所有可用版本
|
||||
func (r *RecordsData) GetVersions(fromLang, toLang string) []string {
|
||||
versionMap := make(map[string]bool)
|
||||
for _, record := range r.Data {
|
||||
@@ -91,16 +85,12 @@ func (r *RecordsData) GetVersions(fromLang, toLang string) []string {
|
||||
return versions
|
||||
}
|
||||
|
||||
// InitRecords 检测默认配置目录下是否存在 records.json
|
||||
// 存在则解析本地 records.json
|
||||
// 不存在则写出默认内嵌的 records.json 到配置目录然后解析
|
||||
func InitRecords() error {
|
||||
cfg := config.GetConfig()
|
||||
recordsPath := filepath.Join(cfg.ConfigDir, "records.json")
|
||||
|
||||
// 检查文件是否存在
|
||||
if _, err := os.Stat(recordsPath); os.IsNotExist(err) {
|
||||
// 不存在,写出内嵌的 records.json
|
||||
|
||||
logger.Info("Initializing records.json from embedded data")
|
||||
if err := os.MkdirAll(cfg.ConfigDir, 0755); err != nil {
|
||||
return fmt.Errorf("Failed to create config directory: %w", err)
|
||||
@@ -110,7 +100,6 @@ func InitRecords() error {
|
||||
}
|
||||
}
|
||||
|
||||
// 解析本地 records.json
|
||||
logger.Debug("Loading records.json from %s", recordsPath)
|
||||
fileData, err := os.ReadFile(recordsPath)
|
||||
if err != nil {
|
||||
@@ -127,13 +116,11 @@ func InitRecords() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// DownloadRecords 更新 records.json,从远程下载 records.json 到配置目录
|
||||
// 然后解析本地 records.json
|
||||
func DownloadRecords() error {
|
||||
cfg := config.GetConfig()
|
||||
|
||||
logger.Info("Updating records.json from remote")
|
||||
// 下载 records.json
|
||||
|
||||
d := downloader.New(cfg.ConfigDir)
|
||||
if err := d.Download(RecordsUrl, RecordsFileName, &downloader.DownloadOptions{
|
||||
Overwrite: true,
|
||||
@@ -141,25 +128,17 @@ func DownloadRecords() error {
|
||||
return fmt.Errorf("Failed to download records.json: %w", err)
|
||||
}
|
||||
|
||||
// 解析本地 records.json
|
||||
return InitRecords()
|
||||
}
|
||||
|
||||
// DownloadModel 解析 records.json,找到对应的模型属性
|
||||
// 检查配置目录下是否存在对应的模型文件,可通过 sha256 校验
|
||||
// 不存在则下载到语言对子目录
|
||||
// 参数:toLang 目标语言,fromLang 源语言,version 模型版本
|
||||
// records 里同一个 fromLang 和 toLang 会存在多个 version 的模型
|
||||
// 需要根据 version 下载对应的模型,未指定 version 则下载最新版本
|
||||
func DownloadModel(toLang string, fromLang string, version string) error {
|
||||
// 确保 records 已加载
|
||||
|
||||
if GlobalRecords == nil {
|
||||
if err := InitRecords(); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
// 找到匹配的模型记录
|
||||
var matchedRecords []RecordItem
|
||||
for _, record := range GlobalRecords.Data {
|
||||
if record.ToLang == toLang && record.FromLang == fromLang {
|
||||
@@ -173,10 +152,9 @@ func DownloadModel(toLang string, fromLang string, version string) error {
|
||||
return fmt.Errorf("No model found for %s -> %s (version: %s)", fromLang, toLang, version)
|
||||
}
|
||||
|
||||
// 如果未指定版本,找最新版本
|
||||
targetRecords := matchedRecords
|
||||
if version == "" {
|
||||
// 按 fileType 分组,每组找最新版本
|
||||
|
||||
fileTypeMap := make(map[string][]RecordItem)
|
||||
for _, record := range matchedRecords {
|
||||
fileTypeMap[record.FileType] = append(fileTypeMap[record.FileType], record)
|
||||
@@ -195,17 +173,15 @@ func DownloadModel(toLang string, fromLang string, version string) error {
|
||||
}
|
||||
}
|
||||
|
||||
// 构建语言对子目录
|
||||
cfg := config.GetConfig()
|
||||
langPairDir := filepath.Join(cfg.ModelDir, fmt.Sprintf("%s_%s", fromLang, toLang))
|
||||
|
||||
// 创建语言对子目录
|
||||
if err := os.MkdirAll(langPairDir, 0755); err != nil {
|
||||
return fmt.Errorf("Failed to create language pair directory: %w", err)
|
||||
}
|
||||
|
||||
logger.Info("Downloading model files for %s -> %s", fromLang, toLang)
|
||||
// 下载所有需要的文件到语言对子目录
|
||||
|
||||
d := downloader.New(langPairDir)
|
||||
|
||||
for _, record := range targetRecords {
|
||||
@@ -226,37 +202,30 @@ func DownloadModel(toLang string, fromLang string, version string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetModelFiles 根据语言对查找模型文件路径
|
||||
// 返回 map[string]string,key 为文件类型:model, lex, vocab_src, vocab_trg
|
||||
// 支持单个词表文件同时用于源语言和目标语言
|
||||
func GetModelFiles(modelDir, fromLang, toLang string) (map[string]string, error) {
|
||||
// 确保 records 已加载
|
||||
|
||||
if GlobalRecords == nil {
|
||||
if err := InitRecords(); err != nil {
|
||||
return nil, fmt.Errorf("failed to init records: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
// 构建语言对子目录
|
||||
langPairDir := filepath.Join(modelDir, fmt.Sprintf("%s_%s", fromLang, toLang))
|
||||
|
||||
files := make(map[string]string)
|
||||
fileTypeMap := make(map[string]string) // fileType -> fullPath
|
||||
fileTypeMap := make(map[string]string)
|
||||
|
||||
// 从 records 中查找匹配的文件
|
||||
for _, record := range GlobalRecords.Data {
|
||||
if record.FromLang == fromLang && record.ToLang == toLang {
|
||||
filename := record.Attachment.Filename
|
||||
fullPath := filepath.Join(langPairDir, filename)
|
||||
|
||||
// 检查文件是否存在
|
||||
if _, err := os.Stat(fullPath); err == nil {
|
||||
fileTypeMap[record.FileType] = fullPath
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 映射 fileType 到所需的 key
|
||||
if modelPath, ok := fileTypeMap["model"]; ok {
|
||||
files["model"] = modelPath
|
||||
} else {
|
||||
@@ -269,14 +238,12 @@ func GetModelFiles(modelDir, fromLang, toLang string) (map[string]string, error)
|
||||
return nil, fmt.Errorf("lex file not found for %s -> %s", fromLang, toLang)
|
||||
}
|
||||
|
||||
// 处理词表文件:可能是单个 vocab 文件或分开的 srcvocab/trgvocab
|
||||
// 优先查找 vocab(单个词表文件)
|
||||
if vocabPath, ok := fileTypeMap["vocab"]; ok {
|
||||
// 单个词表文件同时用于源语言和目标语言
|
||||
|
||||
files["vocab_src"] = vocabPath
|
||||
files["vocab_trg"] = vocabPath
|
||||
} else {
|
||||
// 查找分开的词表文件
|
||||
|
||||
if srcvocabPath, ok := fileTypeMap["srcvocab"]; ok {
|
||||
files["vocab_src"] = srcvocabPath
|
||||
} else {
|
||||
@@ -293,13 +260,11 @@ func GetModelFiles(modelDir, fromLang, toLang string) (map[string]string, error)
|
||||
return files, nil
|
||||
}
|
||||
|
||||
// IsModelDownloaded 检查指定语言对的模型是否已下载
|
||||
func IsModelDownloaded(modelDir, fromLang, toLang string) bool {
|
||||
_, err := GetModelFiles(modelDir, fromLang, toLang)
|
||||
return err == nil
|
||||
}
|
||||
|
||||
// GetSupportedLanguages 获取所有支持的语言列表
|
||||
func GetSupportedLanguages() ([]string, error) {
|
||||
if GlobalRecords == nil {
|
||||
if err := InitRecords(); err != nil {
|
||||
@@ -320,7 +285,6 @@ func GetSupportedLanguages() ([]string, error) {
|
||||
return langs, nil
|
||||
}
|
||||
|
||||
// ValidateLanguagePair 验证语言对是否有效
|
||||
func ValidateLanguagePair(fromLang, toLang string) error {
|
||||
if GlobalRecords == nil {
|
||||
if err := InitRecords(); err != nil {
|
||||
|
||||
@@ -12,45 +12,37 @@ import (
|
||||
)
|
||||
|
||||
func TestInitRecords(t *testing.T) {
|
||||
// 保存原始配置
|
||||
|
||||
oldConfig := config.GlobalConfig
|
||||
defer func() { config.GlobalConfig = oldConfig }()
|
||||
|
||||
// 创建临时测试目录
|
||||
tmpDir := t.TempDir()
|
||||
|
||||
// 设置测试配置
|
||||
config.GlobalConfig = &config.Config{
|
||||
ConfigDir: tmpDir,
|
||||
ModelDir: filepath.Join(tmpDir, "models"),
|
||||
}
|
||||
|
||||
// 重置缓存
|
||||
models.GlobalRecords = nil
|
||||
|
||||
// 测试初始化
|
||||
err := models.InitRecords()
|
||||
if err != nil {
|
||||
t.Fatalf("initRecords() error = %v", err)
|
||||
}
|
||||
|
||||
// 检查 records.json 是否被写出
|
||||
recordsPath := filepath.Join(tmpDir, "records.json")
|
||||
if _, err := os.Stat(recordsPath); os.IsNotExist(err) {
|
||||
t.Fatal("records.json was not created")
|
||||
}
|
||||
|
||||
// 检查缓存是否被设置
|
||||
if models.GlobalRecords == nil {
|
||||
t.Fatal("GlobalRecords was not set")
|
||||
}
|
||||
|
||||
// 检查数据是否正确解析
|
||||
if len(models.GlobalRecords.Data) == 0 {
|
||||
t.Fatal("GlobalRecords.Data is empty")
|
||||
}
|
||||
|
||||
// 再次调用 initRecords,应该使用已存在的文件
|
||||
models.GlobalRecords = nil
|
||||
err = models.InitRecords()
|
||||
if err != nil {
|
||||
@@ -63,7 +55,7 @@ func TestInitRecords(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestRecordsDataStructure(t *testing.T) {
|
||||
// 测试 JSON 解析
|
||||
|
||||
var records models.RecordsData
|
||||
err := json.Unmarshal(data.RecordsJson, &records)
|
||||
if err != nil {
|
||||
@@ -74,7 +66,6 @@ func TestRecordsDataStructure(t *testing.T) {
|
||||
t.Fatal("No records found in embedded data")
|
||||
}
|
||||
|
||||
// 验证第一条记录的结构
|
||||
firstRecord := records.Data[0]
|
||||
if firstRecord.Name == "" {
|
||||
t.Error("Record name is empty")
|
||||
@@ -103,7 +94,7 @@ func TestRecordsDataStructure(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestFindModelRecords(t *testing.T) {
|
||||
// 解析内嵌数据
|
||||
|
||||
var records models.RecordsData
|
||||
err := json.Unmarshal(data.RecordsJson, &records)
|
||||
if err != nil {
|
||||
@@ -176,14 +167,13 @@ func TestFindModelRecords(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestVersionGrouping(t *testing.T) {
|
||||
// 解析内嵌数据
|
||||
|
||||
var records models.RecordsData
|
||||
err := json.Unmarshal(data.RecordsJson, &records)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to unmarshal records: %v", err)
|
||||
}
|
||||
|
||||
// 找到 en->pl 的所有记录
|
||||
var matchedRecords []models.RecordItem
|
||||
for _, record := range records.Data {
|
||||
if record.ToLang == "pl" && record.FromLang == "en" {
|
||||
@@ -195,13 +185,11 @@ func TestVersionGrouping(t *testing.T) {
|
||||
t.Skip("No en->pl records found for testing")
|
||||
}
|
||||
|
||||
// 按 fileType 分组
|
||||
fileTypeMap := make(map[string][]models.RecordItem)
|
||||
for _, record := range matchedRecords {
|
||||
fileTypeMap[record.FileType] = append(fileTypeMap[record.FileType], record)
|
||||
}
|
||||
|
||||
// 验证每种文件类型都有记录
|
||||
expectedTypes := []string{"model", "vocab", "lex"}
|
||||
for _, fileType := range expectedTypes {
|
||||
if records, exists := fileTypeMap[fileType]; !exists || len(records) == 0 {
|
||||
@@ -209,7 +197,6 @@ func TestVersionGrouping(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// 验证版本分组
|
||||
for fileType, fileRecords := range fileTypeMap {
|
||||
if len(fileRecords) > 1 {
|
||||
t.Logf("FileType %s has %d versions", fileType, len(fileRecords))
|
||||
@@ -222,40 +209,32 @@ func TestDownloadRecords(t *testing.T) {
|
||||
t.Skip("Skipping download test in short mode")
|
||||
}
|
||||
|
||||
// 保存原始配置
|
||||
oldConfig := config.GlobalConfig
|
||||
defer func() { config.GlobalConfig = oldConfig }()
|
||||
|
||||
// 创建临时测试目录
|
||||
tmpDir := t.TempDir()
|
||||
|
||||
// 设置测试配置
|
||||
config.GlobalConfig = &config.Config{
|
||||
ConfigDir: tmpDir,
|
||||
ModelDir: filepath.Join(tmpDir, "models"),
|
||||
}
|
||||
|
||||
// 重置缓存
|
||||
models.GlobalRecords = nil
|
||||
|
||||
// 测试下载
|
||||
err := models.DownloadRecords()
|
||||
if err != nil {
|
||||
t.Fatalf("downloadRecords() error = %v", err)
|
||||
}
|
||||
|
||||
// 检查 records.json 是否被下载
|
||||
recordsPath := filepath.Join(tmpDir, "records.json")
|
||||
if _, err := os.Stat(recordsPath); os.IsNotExist(err) {
|
||||
t.Fatal("records.json was not downloaded")
|
||||
}
|
||||
|
||||
// 检查缓存是否被设置
|
||||
if models.GlobalRecords == nil {
|
||||
t.Fatal("GlobalRecords was not set after download")
|
||||
}
|
||||
|
||||
// 检查数据是否正确解析
|
||||
if len(models.GlobalRecords.Data) == 0 {
|
||||
t.Fatal("GlobalRecords.Data is empty after download")
|
||||
}
|
||||
@@ -268,19 +247,16 @@ func TestRealDownloadModel(t *testing.T) {
|
||||
|
||||
t.Log("This test is real world test.")
|
||||
|
||||
// 初始化 records
|
||||
err := models.InitRecords()
|
||||
if err != nil {
|
||||
t.Fatalf("initRecords() error = %v", err)
|
||||
}
|
||||
|
||||
// 测试下载模型(选择一个较小的模型进行测试)
|
||||
err = models.DownloadModel("ja", "en", "")
|
||||
if err != nil {
|
||||
t.Fatalf("downloadModel() error = %v", err)
|
||||
}
|
||||
|
||||
// 验证模型文件已下载
|
||||
modelDir := config.GetConfig().ModelDir
|
||||
files, err := os.ReadDir(modelDir)
|
||||
if err != nil {
|
||||
@@ -291,7 +267,6 @@ func TestRealDownloadModel(t *testing.T) {
|
||||
t.Fatal("No model files were downloaded")
|
||||
}
|
||||
|
||||
// 验证下载的文件与 records 中的记录匹配
|
||||
downloadedFiles := make(map[string]bool)
|
||||
for _, file := range files {
|
||||
downloadedFiles[file.Name()] = true
|
||||
@@ -303,7 +278,7 @@ func TestRealDownloadModel(t *testing.T) {
|
||||
if !downloadedFiles[record.Attachment.Filename] {
|
||||
t.Errorf("Expected file %s was not downloaded", record.Attachment.Filename)
|
||||
}
|
||||
// 验证文件类型
|
||||
|
||||
found := false
|
||||
for _, ft := range expectedFileTypes {
|
||||
if record.FileType == ft {
|
||||
@@ -323,35 +298,28 @@ func TestDownloadModelLatestVersion(t *testing.T) {
|
||||
t.Skip("Skipping model download test in short mode")
|
||||
}
|
||||
|
||||
// 保存原始配置
|
||||
oldConfig := config.GlobalConfig
|
||||
defer func() { config.GlobalConfig = oldConfig }()
|
||||
|
||||
// 创建临时测试目录
|
||||
tmpDir := t.TempDir()
|
||||
|
||||
// 设置测试配置
|
||||
config.GlobalConfig = &config.Config{
|
||||
ConfigDir: tmpDir,
|
||||
ModelDir: filepath.Join(tmpDir, "models"),
|
||||
}
|
||||
|
||||
// 重置缓存
|
||||
models.GlobalRecords = nil
|
||||
|
||||
// 初始化 records
|
||||
err := models.InitRecords()
|
||||
if err != nil {
|
||||
t.Fatalf("initRecords() error = %v", err)
|
||||
}
|
||||
|
||||
// 测试下载最新版本的模型(不指定版本号)
|
||||
err = models.DownloadModel("de", "en", "")
|
||||
if err != nil {
|
||||
t.Fatalf("downloadModel() error = %v", err)
|
||||
}
|
||||
|
||||
// 验证模型文件已下载
|
||||
modelDir := filepath.Join(tmpDir, "models")
|
||||
files, err := os.ReadDir(modelDir)
|
||||
if err != nil {
|
||||
@@ -368,29 +336,23 @@ func TestDownloadModelNonExistent(t *testing.T) {
|
||||
t.Skip("Skipping test in short mode")
|
||||
}
|
||||
|
||||
// 保存原始配置
|
||||
oldConfig := config.GlobalConfig
|
||||
defer func() { config.GlobalConfig = oldConfig }()
|
||||
|
||||
// 创建临时测试目录
|
||||
tmpDir := t.TempDir()
|
||||
|
||||
// 设置测试配置
|
||||
config.GlobalConfig = &config.Config{
|
||||
ConfigDir: tmpDir,
|
||||
ModelDir: filepath.Join(tmpDir, "models"),
|
||||
}
|
||||
|
||||
// 重置缓存
|
||||
models.GlobalRecords = nil
|
||||
|
||||
// 初始化 records
|
||||
err := models.InitRecords()
|
||||
if err != nil {
|
||||
t.Fatalf("initRecords() error = %v", err)
|
||||
}
|
||||
|
||||
// 测试下载不存在的语言对
|
||||
err = models.DownloadModel("zz", "yy", "")
|
||||
if err == nil {
|
||||
t.Fatal("Expected error for non-existent language pair, got nil")
|
||||
|
||||
@@ -14,33 +14,27 @@ import (
|
||||
"github.com/xxnuo/MTranServer/ui"
|
||||
)
|
||||
|
||||
// Setup 设置所有路由
|
||||
func Setup(r *gin.Engine, apiToken string) {
|
||||
// 添加 CORS 中间件
|
||||
|
||||
r.Use(middleware.CORS())
|
||||
|
||||
// 配置 Swagger
|
||||
docs.SwaggerInfo.BasePath = "/"
|
||||
r.GET("/docs/*any", ginSwagger.WrapHandler(swaggerFiles.Handler))
|
||||
|
||||
// 无需认证的路由
|
||||
r.GET("/version", handlers.HandleVersion)
|
||||
r.GET("/health", handlers.HandleHealth)
|
||||
r.GET("/__heartbeat__", handlers.HandleHeartbeat)
|
||||
r.GET("/__lbheartbeat__", handlers.HandleLBHeartbeat)
|
||||
|
||||
// 需要认证的路由
|
||||
auth := r.Group("/")
|
||||
if apiToken != "" {
|
||||
auth.Use(middleware.Auth(apiToken))
|
||||
}
|
||||
|
||||
// 内置接口
|
||||
auth.GET("/languages", handlers.HandleLanguages)
|
||||
auth.POST("/translate", handlers.HandleTranslate)
|
||||
auth.POST("/translate/batch", handlers.HandleTranslateBatch)
|
||||
|
||||
// 插件兼容接口
|
||||
r.POST("/imme", handlers.HandleImmeTranslate(apiToken))
|
||||
r.POST("/kiss", handlers.HandleKissTranslate(apiToken))
|
||||
r.POST("/deepl", handlers.HandleDeeplTranslate(apiToken))
|
||||
@@ -48,13 +42,12 @@ func Setup(r *gin.Engine, apiToken string) {
|
||||
r.GET("/google/translate_a/single", handlers.HandleGoogleTranslateSingle(apiToken))
|
||||
r.POST("/hcfy", handlers.HandleHcfyTranslate(apiToken))
|
||||
|
||||
// 前端静态文件服务
|
||||
cfg := config.GetConfig()
|
||||
if cfg.EnableWebUI {
|
||||
distFS, err := ui.GetDistFS()
|
||||
if err == nil {
|
||||
r.StaticFS("/ui", http.FS(distFS))
|
||||
// 根路径重定向到 /ui
|
||||
|
||||
r.GET("/", func(c *gin.Context) {
|
||||
c.Redirect(http.StatusMovedPermanently, "/ui/")
|
||||
})
|
||||
|
||||
@@ -13,11 +13,9 @@ import (
|
||||
"github.com/xxnuo/MTranServer/internal/routes"
|
||||
)
|
||||
|
||||
// TestIntegrationServerSetup 测试服务器完整设置
|
||||
func TestIntegrationServerSetup(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
// 初始化测试数据
|
||||
models.GlobalRecords = &models.RecordsData{
|
||||
Data: []models.RecordItem{
|
||||
{FromLang: "en", ToLang: "zh-Hans"},
|
||||
@@ -28,7 +26,6 @@ func TestIntegrationServerSetup(t *testing.T) {
|
||||
r := gin.New()
|
||||
routes.Setup(r, "test-token")
|
||||
|
||||
// 测试无需认证的端点
|
||||
t.Run("PublicEndpoints", func(t *testing.T) {
|
||||
endpoints := []string{"/version", "/health", "/__heartbeat__", "/__lbheartbeat__"}
|
||||
for _, endpoint := range endpoints {
|
||||
@@ -39,7 +36,6 @@ func TestIntegrationServerSetup(t *testing.T) {
|
||||
}
|
||||
})
|
||||
|
||||
// 测试需要认证的端点
|
||||
t.Run("AuthenticatedEndpoints", func(t *testing.T) {
|
||||
w := httptest.NewRecorder()
|
||||
req, _ := http.NewRequest("GET", "/languages", nil)
|
||||
@@ -53,7 +49,6 @@ func TestIntegrationServerSetup(t *testing.T) {
|
||||
assert.Contains(t, response, "languages")
|
||||
})
|
||||
|
||||
// 测试认证失败
|
||||
t.Run("AuthenticationFailure", func(t *testing.T) {
|
||||
w := httptest.NewRecorder()
|
||||
req, _ := http.NewRequest("GET", "/languages", nil)
|
||||
@@ -62,14 +57,12 @@ func TestIntegrationServerSetup(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
// TestIntegrationCORS 测试 CORS 完整流程
|
||||
func TestIntegrationCORS(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
r := gin.New()
|
||||
routes.Setup(r, "")
|
||||
|
||||
// 测试 OPTIONS 请求
|
||||
w := httptest.NewRecorder()
|
||||
req, _ := http.NewRequest("OPTIONS", "/version", nil)
|
||||
req.Header.Set("Origin", "http://example.com")
|
||||
@@ -79,7 +72,6 @@ func TestIntegrationCORS(t *testing.T) {
|
||||
assert.Equal(t, "*", w.Header().Get("Access-Control-Allow-Origin"))
|
||||
}
|
||||
|
||||
// TestIntegrationPluginEndpoints 测试插件端点
|
||||
func TestIntegrationPluginEndpoints(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
@@ -92,7 +84,6 @@ func TestIntegrationPluginEndpoints(t *testing.T) {
|
||||
r := gin.New()
|
||||
routes.Setup(r, "test-token")
|
||||
|
||||
// 测试沉浸式翻译端点(需要 token 在 query)
|
||||
t.Run("ImmeEndpoint", func(t *testing.T) {
|
||||
reqBody := map[string]interface{}{
|
||||
"from": "en",
|
||||
@@ -106,13 +97,10 @@ func TestIntegrationPluginEndpoints(t *testing.T) {
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
// 由于没有实际的翻译引擎,期望返回错误
|
||||
// 但至少应该通过认证和请求解析
|
||||
assert.NotEqual(t, http.StatusUnauthorized, w.Code)
|
||||
assert.NotEqual(t, http.StatusBadRequest, w.Code)
|
||||
})
|
||||
|
||||
// 测试简约翻译端点(需要 token 在 header KEY)
|
||||
t.Run("KissEndpoint", func(t *testing.T) {
|
||||
reqBody := map[string]interface{}{
|
||||
"from": "en",
|
||||
@@ -127,14 +115,11 @@ func TestIntegrationPluginEndpoints(t *testing.T) {
|
||||
req.Header.Set("KEY", "test-token")
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
// 由于没有实际的翻译引擎,期望返回错误
|
||||
// 但至少应该通过认证和请求解析
|
||||
assert.NotEqual(t, http.StatusUnauthorized, w.Code)
|
||||
assert.NotEqual(t, http.StatusBadRequest, w.Code)
|
||||
})
|
||||
}
|
||||
|
||||
// TestIntegrationAPIEndpoints 测试 API 端点
|
||||
func TestIntegrationAPIEndpoints(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
@@ -147,7 +132,6 @@ func TestIntegrationAPIEndpoints(t *testing.T) {
|
||||
r := gin.New()
|
||||
routes.Setup(r, "test-token")
|
||||
|
||||
// 测试翻译端点
|
||||
t.Run("TranslateEndpoint", func(t *testing.T) {
|
||||
reqBody := map[string]interface{}{
|
||||
"from": "en",
|
||||
@@ -163,13 +147,10 @@ func TestIntegrationAPIEndpoints(t *testing.T) {
|
||||
req.Header.Set("Authorization", "test-token")
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
// 由于没有实际的翻译引擎,期望返回错误
|
||||
// 但至少应该通过认证和请求解析
|
||||
assert.NotEqual(t, http.StatusUnauthorized, w.Code)
|
||||
assert.NotEqual(t, http.StatusBadRequest, w.Code)
|
||||
})
|
||||
|
||||
// 测试批量翻译端点
|
||||
t.Run("TranslateBatchEndpoint", func(t *testing.T) {
|
||||
reqBody := map[string]interface{}{
|
||||
"from": "en",
|
||||
@@ -185,13 +166,10 @@ func TestIntegrationAPIEndpoints(t *testing.T) {
|
||||
req.Header.Set("Authorization", "test-token")
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
// 由于没有实际的翻译引擎,期望返回错误
|
||||
// 但至少应该通过认证和请求解析
|
||||
assert.NotEqual(t, http.StatusUnauthorized, w.Code)
|
||||
assert.NotEqual(t, http.StatusBadRequest, w.Code)
|
||||
})
|
||||
|
||||
// 测试 Google 兼容端点
|
||||
t.Run("GoogleCompatEndpoint", func(t *testing.T) {
|
||||
reqBody := map[string]interface{}{
|
||||
"q": "Hello",
|
||||
@@ -207,21 +185,17 @@ func TestIntegrationAPIEndpoints(t *testing.T) {
|
||||
req.Header.Set("Authorization", "test-token")
|
||||
r.ServeHTTP(w, req)
|
||||
|
||||
// 由于没有实际的翻译引擎,期望返回错误
|
||||
// 但至少应该通过认证和请求解析
|
||||
assert.NotEqual(t, http.StatusUnauthorized, w.Code)
|
||||
assert.NotEqual(t, http.StatusBadRequest, w.Code)
|
||||
})
|
||||
}
|
||||
|
||||
// TestIntegrationInvalidRequests 测试无效请求
|
||||
func TestIntegrationInvalidRequests(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
r := gin.New()
|
||||
routes.Setup(r, "test-token")
|
||||
|
||||
// 测试无效的 JSON
|
||||
t.Run("InvalidJSON", func(t *testing.T) {
|
||||
w := httptest.NewRecorder()
|
||||
req, _ := http.NewRequest("POST", "/translate", bytes.NewBufferString("invalid json"))
|
||||
@@ -232,11 +206,9 @@ func TestIntegrationInvalidRequests(t *testing.T) {
|
||||
assert.Equal(t, http.StatusBadRequest, w.Code)
|
||||
})
|
||||
|
||||
// 测试缺少必需字段
|
||||
t.Run("MissingRequiredFields", func(t *testing.T) {
|
||||
reqBody := map[string]interface{}{
|
||||
"from": "en",
|
||||
// 缺少 to 和 text
|
||||
}
|
||||
body, _ := json.Marshal(reqBody)
|
||||
|
||||
|
||||
@@ -20,51 +20,39 @@ import (
|
||||
"github.com/xxnuo/MTranServer/internal/services"
|
||||
)
|
||||
|
||||
// Run 启动服务器
|
||||
func Run() error {
|
||||
// 加载配置
|
||||
|
||||
cfg := config.GetConfig()
|
||||
|
||||
// 初始化 records
|
||||
if err := models.InitRecords(); err != nil {
|
||||
return fmt.Errorf("failed to initialize records: %w", err)
|
||||
}
|
||||
|
||||
// 创建必要的目录
|
||||
if err := os.MkdirAll(cfg.ModelDir, 0755); err != nil {
|
||||
return fmt.Errorf("failed to create model directory: %w", err)
|
||||
}
|
||||
|
||||
// 初始化 worker 二进制文件
|
||||
if err := manager.EnsureWorkerBinary(cfg); err != nil {
|
||||
return fmt.Errorf("failed to initialize worker binary: %w", err)
|
||||
}
|
||||
|
||||
// 设置 Gin 模式
|
||||
// 始终使用 ReleaseMode,我们使用自定义的日志中间件
|
||||
gin.SetMode(gin.ReleaseMode)
|
||||
|
||||
// 创建 Gin 引擎(不使用默认中间件)
|
||||
r := gin.New()
|
||||
|
||||
// 添加自定义中间件
|
||||
r.Use(middleware.Recovery())
|
||||
r.Use(middleware.Logger())
|
||||
|
||||
// 注册路由
|
||||
routes.Setup(r, cfg.APIToken)
|
||||
|
||||
// 启动服务器
|
||||
addr := fmt.Sprintf("%s:%s", cfg.Host, cfg.Port)
|
||||
srv := &http.Server{
|
||||
Addr: addr,
|
||||
Handler: r,
|
||||
}
|
||||
|
||||
// 用于等待优雅关闭完成的通道
|
||||
shutdownDone := make(chan struct{})
|
||||
|
||||
// 优雅关闭
|
||||
go func() {
|
||||
sigChan := make(chan os.Signal, 1)
|
||||
signal.Notify(sigChan, syscall.SIGINT, syscall.SIGTERM)
|
||||
@@ -75,7 +63,6 @@ func Run() error {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer cancel()
|
||||
|
||||
// 关闭所有翻译引擎
|
||||
services.CleanupAllEngines()
|
||||
|
||||
if err := srv.Shutdown(ctx); err != nil {
|
||||
@@ -85,23 +72,20 @@ func Run() error {
|
||||
close(shutdownDone)
|
||||
}()
|
||||
|
||||
// 总是输出服务启动信息(即使在 warn/error 模式下)
|
||||
fmt.Fprintf(os.Stderr, "[INFO] %s HTTP Service URL: http://%s\n",
|
||||
time.Now().Format("2006/01/02 15:04:05"), addr)
|
||||
fmt.Fprintf(os.Stderr, "[INFO] %s Swagger UI: http://%s/docs/index.html\n",
|
||||
time.Now().Format("2006/01/02 15:04:05"), addr)
|
||||
|
||||
// 总是输出日志级别信息(即使在 warn/error 模式下)
|
||||
fmt.Fprintf(os.Stderr, "[INFO] %s Log level set to: %s\n",
|
||||
time.Now().Format("2006/01/02 15:04:05"), cfg.LogLevel)
|
||||
|
||||
if err := srv.ListenAndServe(); err != nil && err != http.ErrServerClosed {
|
||||
// 服务器启动失败,确保清理资源
|
||||
|
||||
services.CleanupAllEngines()
|
||||
return fmt.Errorf("failed to start server: %w", err)
|
||||
}
|
||||
|
||||
// 等待优雅关闭完成
|
||||
<-shutdownDone
|
||||
logger.Info("Server shutdown complete")
|
||||
|
||||
|
||||
@@ -16,7 +16,6 @@ import (
|
||||
func TestMain(m *testing.M) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
|
||||
// 初始化测试数据
|
||||
models.GlobalRecords = &models.RecordsData{
|
||||
Data: []models.RecordItem{
|
||||
{FromLang: "en", ToLang: "zh-Hans"},
|
||||
|
||||
@@ -13,7 +13,6 @@ var (
|
||||
detectorOnce sync.Once
|
||||
)
|
||||
|
||||
// initDetector 初始化语言检测器(懒加载)
|
||||
func initDetector() {
|
||||
detectorOnce.Do(func() {
|
||||
logger.Debug("Initializing language detector")
|
||||
@@ -25,31 +24,26 @@ func initDetector() {
|
||||
})
|
||||
}
|
||||
|
||||
// linguaToBCP47 将 lingua 语言代码转换为 BCP 47 格式
|
||||
func linguaToBCP47(lang lingua.Language) string {
|
||||
// 特殊处理中文
|
||||
|
||||
switch lang {
|
||||
case lingua.Chinese:
|
||||
// 默认返回简体中文
|
||||
|
||||
return "zh-Hans"
|
||||
default:
|
||||
// 其他语言直接返回 ISO 639-1 代码(小写)
|
||||
|
||||
code := lang.IsoCode639_1()
|
||||
return strings.ToLower(code.String())
|
||||
}
|
||||
}
|
||||
|
||||
// DetectLanguage 检测文本语言并返回 BCP 47 格式的语言代码
|
||||
// 如果检测失败或无法确定,返回空字符串
|
||||
func DetectLanguage(text string) string {
|
||||
if text == "" {
|
||||
return ""
|
||||
}
|
||||
|
||||
// 初始化检测器
|
||||
initDetector()
|
||||
|
||||
// 检测语言
|
||||
lang, exists := detector.DetectLanguageOf(text)
|
||||
if !exists {
|
||||
return ""
|
||||
@@ -58,27 +52,21 @@ func DetectLanguage(text string) string {
|
||||
return linguaToBCP47(lang)
|
||||
}
|
||||
|
||||
// DetectLanguageWithConfidence 检测文本语言并返回置信度
|
||||
// 返回 BCP 47 格式的语言代码和置信度(0.0-1.0)
|
||||
func DetectLanguageWithConfidence(text string, minConfidence float64) (string, float64) {
|
||||
if text == "" {
|
||||
return "", 0.0
|
||||
}
|
||||
|
||||
// 初始化检测器
|
||||
initDetector()
|
||||
|
||||
// 获取所有语言的置信度
|
||||
confidenceValues := detector.ComputeLanguageConfidenceValues(text)
|
||||
if len(confidenceValues) == 0 {
|
||||
return "", 0.0
|
||||
}
|
||||
|
||||
// 获取置信度最高的语言
|
||||
topResult := confidenceValues[0]
|
||||
confidence := topResult.Value()
|
||||
|
||||
// 如果置信度低于阈值,返回空
|
||||
if confidence < minConfidence {
|
||||
return "", confidence
|
||||
}
|
||||
@@ -86,17 +74,13 @@ func DetectLanguageWithConfidence(text string, minConfidence float64) (string, f
|
||||
return linguaToBCP47(topResult.Language()), confidence
|
||||
}
|
||||
|
||||
// NormalizeLanguageCode 标准化语言代码为 BCP 47 格式
|
||||
// 支持多种输入格式:zh, zh-CN, zh_CN, Chinese 等
|
||||
func NormalizeLanguageCode(code string) string {
|
||||
if code == "" {
|
||||
return ""
|
||||
}
|
||||
|
||||
// 转为小写并替换下划线为连字符
|
||||
code = strings.ToLower(strings.ReplaceAll(code, "_", "-"))
|
||||
|
||||
// 特殊处理中文
|
||||
switch code {
|
||||
case "zh", "zh-cn", "zh-hans", "chinese", "cmn":
|
||||
return "zh-Hans"
|
||||
@@ -104,10 +88,8 @@ func NormalizeLanguageCode(code string) string {
|
||||
return "zh-Hant"
|
||||
}
|
||||
|
||||
// 如果是标准的 BCP 47 格式(如 en-US),提取主语言代码
|
||||
parts := strings.Split(code, "-")
|
||||
mainCode := parts[0]
|
||||
|
||||
// 返回主语言代码
|
||||
return mainCode
|
||||
}
|
||||
|
||||
@@ -14,7 +14,6 @@ import (
|
||||
"github.com/xxnuo/MTranServer/internal/utils"
|
||||
)
|
||||
|
||||
// EngineInfo 引擎信息
|
||||
type EngineInfo struct {
|
||||
Manager *manager.Manager
|
||||
LastUsed time.Time
|
||||
@@ -25,28 +24,23 @@ type EngineInfo struct {
|
||||
}
|
||||
|
||||
var (
|
||||
// 存储已加载的翻译引擎 key: "fromLang-toLang"
|
||||
engines = make(map[string]*EngineInfo)
|
||||
engMu sync.RWMutex
|
||||
)
|
||||
|
||||
// resetIdleTimer 重置空闲计时器
|
||||
func (ei *EngineInfo) resetIdleTimer() {
|
||||
ei.mu.Lock()
|
||||
defer ei.mu.Unlock()
|
||||
|
||||
ei.LastUsed = time.Now()
|
||||
|
||||
// 停止旧的计时器
|
||||
if ei.stopTimer != nil {
|
||||
ei.stopTimer.Stop()
|
||||
}
|
||||
|
||||
// 从配置获取超时时间
|
||||
cfg := config.GetConfig()
|
||||
timeout := time.Duration(cfg.WorkerIdleTimeout) * time.Second
|
||||
|
||||
// 创建新的计时器
|
||||
ei.stopTimer = time.AfterFunc(timeout, func() {
|
||||
key := fmt.Sprintf("%s-%s", ei.FromLang, ei.ToLang)
|
||||
logger.Info("Engine %s idle timeout, stopping...", key)
|
||||
@@ -64,28 +58,23 @@ func (ei *EngineInfo) resetIdleTimer() {
|
||||
})
|
||||
}
|
||||
|
||||
// getOrCreateSingleEngine 获取或创建单个翻译引擎(内部函数)
|
||||
// 一个 worker 只对应一个语言方向的翻译(fromLang -> toLang)
|
||||
func getOrCreateSingleEngine(fromLang, toLang string) (*manager.Manager, error) {
|
||||
key := fmt.Sprintf("%s-%s", fromLang, toLang)
|
||||
|
||||
// 检查是否已存在
|
||||
engMu.RLock()
|
||||
if info, ok := engines[key]; ok {
|
||||
if info.Manager.IsRunning() {
|
||||
engMu.RUnlock()
|
||||
// 更新最后使用时间并重置空闲计时器
|
||||
|
||||
info.resetIdleTimer()
|
||||
return info.Manager, nil
|
||||
}
|
||||
}
|
||||
engMu.RUnlock()
|
||||
|
||||
// 创建新引擎
|
||||
engMu.Lock()
|
||||
defer engMu.Unlock()
|
||||
|
||||
// 再次检查(双重检查锁定)
|
||||
if info, ok := engines[key]; ok {
|
||||
if info.Manager.IsRunning() {
|
||||
info.resetIdleTimer()
|
||||
@@ -95,7 +84,6 @@ func getOrCreateSingleEngine(fromLang, toLang string) (*manager.Manager, error)
|
||||
|
||||
logger.Info("Creating new engine for %s -> %s", fromLang, toLang)
|
||||
|
||||
// 下载模型(如果需要)
|
||||
cfg := config.GetConfig()
|
||||
if cfg.EnableOfflineMode {
|
||||
logger.Info("Offline mode enabled, skipping model download")
|
||||
@@ -106,13 +94,11 @@ func getOrCreateSingleEngine(fromLang, toLang string) (*manager.Manager, error)
|
||||
}
|
||||
}
|
||||
|
||||
// 查找模型文件
|
||||
modelFiles, err := models.GetModelFiles(cfg.ModelDir, fromLang, toLang)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to find model files: %w", err)
|
||||
}
|
||||
|
||||
// 创建 Worker,分配独立端口
|
||||
port, err := utils.GetFreePort()
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to allocate port: %w", err)
|
||||
@@ -120,22 +106,18 @@ func getOrCreateSingleEngine(fromLang, toLang string) (*manager.Manager, error)
|
||||
args := manager.NewWorkerArgs()
|
||||
args.Port = port
|
||||
|
||||
// WorkDir 设置为语言对子目录
|
||||
langPairDir := filepath.Join(cfg.ModelDir, fmt.Sprintf("%s_%s", fromLang, toLang))
|
||||
args.WorkDir = langPairDir
|
||||
|
||||
m := manager.NewManager(args)
|
||||
|
||||
// 启动 Manager
|
||||
if err := m.Start(); err != nil {
|
||||
return nil, fmt.Errorf("failed to start manager: %w", err)
|
||||
}
|
||||
|
||||
// 加载模型
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
|
||||
defer cancel()
|
||||
|
||||
// 提取文件名(相对于 WorkDir)
|
||||
poweronReq := manager.PoweronRequest{
|
||||
ModelPath: filepath.Base(modelFiles["model"]),
|
||||
LexicalShortlistPath: filepath.Base(modelFiles["lex"]),
|
||||
@@ -147,7 +129,6 @@ func getOrCreateSingleEngine(fromLang, toLang string) (*manager.Manager, error)
|
||||
return nil, fmt.Errorf("failed to load model: %w", err)
|
||||
}
|
||||
|
||||
// 等待引擎就绪
|
||||
for i := 0; i < 30; i++ {
|
||||
ready, err := m.Ready(ctx)
|
||||
if err == nil && ready {
|
||||
@@ -156,7 +137,6 @@ func getOrCreateSingleEngine(fromLang, toLang string) (*manager.Manager, error)
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
}
|
||||
|
||||
// 创建引擎信息并设置空闲计时器
|
||||
info := &EngineInfo{
|
||||
Manager: m,
|
||||
LastUsed: time.Now(),
|
||||
@@ -171,41 +151,31 @@ func getOrCreateSingleEngine(fromLang, toLang string) (*manager.Manager, error)
|
||||
return m, nil
|
||||
}
|
||||
|
||||
// needsPivotTranslation 检查是否需要通过英语中转
|
||||
func needsPivotTranslation(fromLang, toLang string) bool {
|
||||
// 如果源语言或目标语言是英语,不需要中转
|
||||
|
||||
if fromLang == "en" || toLang == "en" {
|
||||
return false
|
||||
}
|
||||
|
||||
// 检查是否存在直接的语言对
|
||||
if models.GlobalRecords != nil && models.GlobalRecords.HasLanguagePair(fromLang, toLang) {
|
||||
return false
|
||||
}
|
||||
|
||||
// 需要通过英语中转
|
||||
return true
|
||||
}
|
||||
|
||||
// GetOrCreateEngine 获取或创建翻译引擎
|
||||
// 如果需要跨英语翻译(如 zh-Hans -> ja),返回第一步的引擎(zh-Hans -> en)
|
||||
// 调用者需要使用 TranslateWithPivot 来完成完整的翻译
|
||||
func GetOrCreateEngine(fromLang, toLang string) (*manager.Manager, error) {
|
||||
// 如果不需要中转,直接创建单个引擎
|
||||
|
||||
if !needsPivotTranslation(fromLang, toLang) {
|
||||
return getOrCreateSingleEngine(fromLang, toLang)
|
||||
}
|
||||
|
||||
// 需要中转,返回第一步的引擎
|
||||
logger.Debug("Translation %s -> %s requires pivot through English", fromLang, toLang)
|
||||
return getOrCreateSingleEngine(fromLang, "en")
|
||||
}
|
||||
|
||||
// TranslateWithPivot 处理可能需要中转的翻译
|
||||
// 如果需要中转(如 zh-Hans -> ja),会自动创建两个引擎并执行两步翻译
|
||||
// 支持 auto 模式:自动检测源语言
|
||||
func TranslateWithPivot(ctx context.Context, fromLang, toLang, text string, isHTML bool) (string, error) {
|
||||
// 处理 auto 模式:自动检测源语言
|
||||
|
||||
if fromLang == "auto" {
|
||||
detected := DetectLanguage(text)
|
||||
if detected == "" {
|
||||
@@ -215,13 +185,11 @@ func TranslateWithPivot(ctx context.Context, fromLang, toLang, text string, isHT
|
||||
fromLang = detected
|
||||
}
|
||||
|
||||
// 如果源语言和目标语言相同,直接返回原文
|
||||
if fromLang == toLang {
|
||||
logger.Debug("Source and target languages are the same (%s), returning original text", fromLang)
|
||||
return text, nil
|
||||
}
|
||||
|
||||
// 如果不需要中转,直接翻译
|
||||
if !needsPivotTranslation(fromLang, toLang) {
|
||||
m, err := getOrCreateSingleEngine(fromLang, toLang)
|
||||
if err != nil {
|
||||
@@ -233,7 +201,6 @@ func TranslateWithPivot(ctx context.Context, fromLang, toLang, text string, isHT
|
||||
return m.Translate(ctx, text)
|
||||
}
|
||||
|
||||
// 需要中转:第一步 fromLang -> en
|
||||
logger.Debug("Step 1: Translating %s -> en", fromLang)
|
||||
m1, err := getOrCreateSingleEngine(fromLang, "en")
|
||||
if err != nil {
|
||||
@@ -250,7 +217,6 @@ func TranslateWithPivot(ctx context.Context, fromLang, toLang, text string, isHT
|
||||
return "", fmt.Errorf("failed in first step (%s -> en): %w", fromLang, err)
|
||||
}
|
||||
|
||||
// 第二步 en -> toLang
|
||||
logger.Debug("Step 2: Translating en -> %s", toLang)
|
||||
m2, err := getOrCreateSingleEngine("en", toLang)
|
||||
if err != nil {
|
||||
@@ -270,7 +236,6 @@ func TranslateWithPivot(ctx context.Context, fromLang, toLang, text string, isHT
|
||||
return finalText, nil
|
||||
}
|
||||
|
||||
// CleanupAllEngines 清理所有翻译引擎
|
||||
func CleanupAllEngines() {
|
||||
engMu.Lock()
|
||||
defer engMu.Unlock()
|
||||
@@ -282,7 +247,6 @@ func CleanupAllEngines() {
|
||||
|
||||
logger.Info("Cleaning up %d engine(s)...", len(engines))
|
||||
|
||||
// 使用 WaitGroup 并发清理所有引擎以加快关闭速度
|
||||
var wg sync.WaitGroup
|
||||
for key, info := range engines {
|
||||
wg.Add(1)
|
||||
@@ -296,14 +260,12 @@ func CleanupAllEngines() {
|
||||
|
||||
logger.Debug("Stopping engine: %s", k)
|
||||
|
||||
// 停止空闲计时器
|
||||
ei.mu.Lock()
|
||||
if ei.stopTimer != nil {
|
||||
ei.stopTimer.Stop()
|
||||
}
|
||||
ei.mu.Unlock()
|
||||
|
||||
// 清理 Manager
|
||||
if err := ei.Manager.Cleanup(); err != nil {
|
||||
logger.Error("Failed to cleanup engine %s: %v", k, err)
|
||||
} else {
|
||||
@@ -312,7 +274,6 @@ func CleanupAllEngines() {
|
||||
}(key, info)
|
||||
}
|
||||
|
||||
// 等待所有清理完成,最多等待 15 秒
|
||||
done := make(chan struct{})
|
||||
go func() {
|
||||
wg.Wait()
|
||||
@@ -326,6 +287,5 @@ func CleanupAllEngines() {
|
||||
logger.Warn("Engine cleanup timeout after 15 seconds")
|
||||
}
|
||||
|
||||
// 清空 engines map
|
||||
engines = make(map[string]*EngineInfo)
|
||||
}
|
||||
|
||||
@@ -5,7 +5,6 @@ import (
|
||||
"strconv"
|
||||
)
|
||||
|
||||
// GetEnv gets the environment variable value or returns the default value
|
||||
func GetEnv(key string, defaultValue string) string {
|
||||
if value := os.Getenv(key); value != "" {
|
||||
return value
|
||||
@@ -13,7 +12,6 @@ func GetEnv(key string, defaultValue string) string {
|
||||
return defaultValue
|
||||
}
|
||||
|
||||
// ParseBoolEnv parses boolean values from environment variables
|
||||
func GetBoolEnv(key string, defaultValue bool) bool {
|
||||
if value := os.Getenv(key); value != "" {
|
||||
result, err := strconv.ParseBool(value)
|
||||
@@ -24,7 +22,6 @@ func GetBoolEnv(key string, defaultValue bool) bool {
|
||||
return defaultValue
|
||||
}
|
||||
|
||||
// GetIntEnv gets the environment variable value as int or returns the default value
|
||||
func GetIntEnv(key string, defaultValue int) int {
|
||||
if value := os.Getenv(key); value != "" {
|
||||
result, err := strconv.Atoi(value)
|
||||
|
||||
@@ -8,7 +8,6 @@ import (
|
||||
"os"
|
||||
)
|
||||
|
||||
// VerifySHA256 校验文件的 SHA256
|
||||
func VerifySHA256(filepath, expectedHash string) error {
|
||||
file, err := os.Open(filepath)
|
||||
if err != nil {
|
||||
@@ -29,7 +28,6 @@ func VerifySHA256(filepath, expectedHash string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// ComputeSHA256 计算文件的 SHA256
|
||||
func ComputeSHA256(filepath string) (string, error) {
|
||||
file, err := os.Open(filepath)
|
||||
if err != nil {
|
||||
|
||||
@@ -7,7 +7,7 @@ import (
|
||||
)
|
||||
|
||||
func TestCalculateSHA256(t *testing.T) {
|
||||
// 创建临时文件
|
||||
|
||||
tempDir, err := os.MkdirTemp("", "file-test-*")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
@@ -22,7 +22,6 @@ func TestCalculateSHA256(t *testing.T) {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
// 计算 SHA256
|
||||
hash, err := ComputeSHA256(filePath)
|
||||
if err != nil {
|
||||
t.Fatalf("计算 SHA256 失败: %v", err)
|
||||
|
||||
@@ -2,7 +2,6 @@ package utils
|
||||
|
||||
import "net"
|
||||
|
||||
// GetFreePort 获取一个未被占用的随机端口
|
||||
func GetFreePort() (int, error) {
|
||||
addr, err := net.ResolveTCPAddr("tcp", "localhost:0")
|
||||
if err != nil {
|
||||
|
||||
@@ -5,11 +5,6 @@ import (
|
||||
"strings"
|
||||
)
|
||||
|
||||
// 比较版本号,返回最大的版本号
|
||||
// 支持 semver 格式,但不完全符合 semver 规范
|
||||
// 样例:
|
||||
// - 1.0.0, 1.0.1, 1.1.0 返回 1.1.0
|
||||
// - 1.0.0-alpha.1, 1.0.0-alpha.2, 1.0.0-alpha.3 返回 1.0.0-alpha.3
|
||||
func GetLargestVersion(versions []string) string {
|
||||
if len(versions) == 0 {
|
||||
return ""
|
||||
@@ -24,24 +19,19 @@ func GetLargestVersion(versions []string) string {
|
||||
return largest
|
||||
}
|
||||
|
||||
// compareVersions 比较两个版本号
|
||||
// 返回值: v1 > v2 返回 1, v1 < v2 返回 -1, v1 == v2 返回 0
|
||||
func compareVersions(v1, v2 string) int {
|
||||
// 分割版本号和预发布标识
|
||||
|
||||
parts1 := strings.Split(v1, "-")
|
||||
parts2 := strings.Split(v2, "-")
|
||||
|
||||
version1 := parts1[0]
|
||||
version2 := parts2[0]
|
||||
|
||||
// 比较主版本号
|
||||
cmp := compareNumericVersions(version1, version2)
|
||||
if cmp != 0 {
|
||||
return cmp
|
||||
}
|
||||
|
||||
// 版本号相同,比较预发布标识
|
||||
// 无预发布标识 > 有预发布标识
|
||||
if len(parts1) == 1 && len(parts2) > 1 {
|
||||
return 1
|
||||
}
|
||||
@@ -52,11 +42,9 @@ func compareVersions(v1, v2 string) int {
|
||||
return 0
|
||||
}
|
||||
|
||||
// 比较预发布标识
|
||||
return comparePrerelease(parts1[1], parts2[1])
|
||||
}
|
||||
|
||||
// compareNumericVersions 比较数字版本号 (如 1.0.0)
|
||||
func compareNumericVersions(v1, v2 string) int {
|
||||
segments1 := strings.Split(v1, ".")
|
||||
segments2 := strings.Split(v2, ".")
|
||||
@@ -87,9 +75,8 @@ func compareNumericVersions(v1, v2 string) int {
|
||||
return 0
|
||||
}
|
||||
|
||||
// comparePrerelease 比较预发布标识
|
||||
func comparePrerelease(p1, p2 string) int {
|
||||
// 分割预发布标识的各个部分
|
||||
|
||||
parts1 := strings.Split(p1, ".")
|
||||
parts2 := strings.Split(p2, ".")
|
||||
|
||||
@@ -109,7 +96,6 @@ func comparePrerelease(p1, p2 string) int {
|
||||
part1 := parts1[i]
|
||||
part2 := parts2[i]
|
||||
|
||||
// 尝试作为数字比较
|
||||
num1, err1 := strconv.Atoi(part1)
|
||||
num2, err2 := strconv.Atoi(part2)
|
||||
|
||||
@@ -121,7 +107,7 @@ func comparePrerelease(p1, p2 string) int {
|
||||
return -1
|
||||
}
|
||||
} else {
|
||||
// 字符串比较
|
||||
|
||||
if part1 > part2 {
|
||||
return 1
|
||||
}
|
||||
|
||||
@@ -1,9 +1,7 @@
|
||||
package version
|
||||
|
||||
// Version 会在编译时通过 -ldflags 注入
|
||||
var Version = "v0.0.0-dev"
|
||||
|
||||
// GetVersion 获取版本号
|
||||
func GetVersion() string {
|
||||
return Version
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user