diff --git a/bin/gen_update.go b/bin/gen_update.go index 13e8ddc..11aeeda 100644 --- a/bin/gen_update.go +++ b/bin/gen_update.go @@ -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) } diff --git a/bin/utils.go b/bin/utils.go index 1449619..c72f0aa 100644 --- a/bin/utils.go +++ b/bin/utils.go @@ -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) } diff --git a/cmd/mtranserver/main.go b/cmd/mtranserver/main.go index 21d651a..40b572d 100644 --- a/cmd/mtranserver/main.go +++ b/cmd/mtranserver/main.go @@ -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) } diff --git a/data/gen_records.go b/data/gen_records.go index 7d0a580..5355be4 100644 --- a/data/gen_records.go +++ b/data/gen_records.go @@ -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(), diff --git a/data/utils.go b/data/utils.go index 8425335..38ff096 100644 --- a/data/utils.go +++ b/data/utils.go @@ -1,3 +1,3 @@ package data -//go:generate env GOOS= GOARCH= go run gen_records.go +//go:generate go run gen_records.go diff --git a/internal/config/config.go b/internal/config/config.go index 6922452..b3168e7 100644 --- a/internal/config/config.go +++ b/internal/config/config.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 ( diff --git a/internal/downloader/downloader.go b/internal/downloader/downloader.go index 0d51080..a98e827 100644 --- a/internal/downloader/downloader.go +++ b/internal/downloader/downloader.go @@ -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) diff --git a/internal/downloader/downloader_test.go b/internal/downloader/downloader_test.go index d2089ae..11aa8a0 100644 --- a/internal/downloader/downloader_test.go +++ b/internal/downloader/downloader_test.go @@ -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) diff --git a/internal/handlers/deepl.go b/internal/handlers/deepl.go index 33a8140..858568c 100644 --- a/internal/handlers/deepl.go +++ b/internal/handlers/deepl.go @@ -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) diff --git a/internal/handlers/google.go b/internal/handlers/google.go index 3dab804..0583b4c 100644 --- a/internal/handlers/google.go +++ b/internal/handlers/google.go @@ -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{}{ diff --git a/internal/handlers/hcfy.go b/internal/handlers/hcfy.go index 5d77491..cf3fd28 100644 --- a/internal/handlers/hcfy.go +++ b/internal/handlers/hcfy.go @@ -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 { diff --git a/internal/handlers/imme.go b/internal/handlers/imme.go index e419b77..06cc682 100644 --- a/internal/handlers/imme.go +++ b/internal/handlers/imme.go @@ -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, diff --git a/internal/handlers/kiss.go b/internal/handlers/kiss.go index 3a5b28e..6829d9a 100644 --- a/internal/handlers/kiss.go +++ b/internal/handlers/kiss.go @@ -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() diff --git a/internal/handlers/language.go b/internal/handlers/language.go index af9767c..77bc4ad 100644 --- a/internal/handlers/language.go +++ b/internal/handlers/language.go @@ -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 diff --git a/internal/handlers/language_test.go b/internal/handlers/language_test.go index 7007ee6..287e35a 100644 --- a/internal/handlers/language_test.go +++ b/internal/handlers/language_test.go @@ -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() { diff --git a/internal/handlers/translate.go b/internal/handlers/translate.go index d58743b..45dce7c 100644 --- a/internal/handlers/translate.go +++ b/internal/handlers/translate.go @@ -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) diff --git a/internal/logger/logger.go b/internal/logger/logger.go index c1042d3..f887749 100644 --- a/internal/logger/logger.go +++ b/internal/logger/logger.go @@ -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...) } diff --git a/internal/manager/client.go b/internal/manager/client.go index f8a6405..31ce334 100644 --- a/internal/manager/client.go +++ b/internal/manager/client.go @@ -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 { diff --git a/internal/manager/client_test.go b/internal/manager/client_test.go index fbf0f48..429824e 100644 --- a/internal/manager/client_test.go +++ b/internal/manager/client_test.go @@ -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)), diff --git a/internal/manager/daemon.go b/internal/manager/daemon.go index e0ff18d..bafc1e8 100644 --- a/internal/manager/daemon.go +++ b/internal/manager/daemon.go @@ -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) } diff --git a/internal/manager/daemon_test.go b/internal/manager/daemon_test.go index b5c7091..8a691af 100644 --- a/internal/manager/daemon_test.go +++ b/internal/manager/daemon_test.go @@ -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) diff --git a/internal/manager/manager.go b/internal/manager/manager.go index cece911..ff500fb 100644 --- a/internal/manager/manager.go +++ b/internal/manager/manager.go @@ -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, diff --git a/internal/manager/manager_test.go b/internal/manager/manager_test.go index f641c67..95b2ad7 100644 --- a/internal/manager/manager_test.go +++ b/internal/manager/manager_test.go @@ -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, "

Hello

") - 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, "

Hello, world!

") 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, diff --git a/internal/middleware/auth.go b/internal/middleware/auth.go index 697d4b4..e770513 100644 --- a/internal/middleware/auth.go +++ b/internal/middleware/auth.go @@ -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{ diff --git a/internal/middleware/cors.go b/internal/middleware/cors.go index 2153755..b3f366f 100644 --- a/internal/middleware/cors.go +++ b/internal/middleware/cors.go @@ -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", "*") diff --git a/internal/middleware/logger.go b/internal/middleware/logger.go index c7c84b0..0070861 100644 --- a/internal/middleware/logger.go +++ b/internal/middleware/logger.go @@ -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() { diff --git a/internal/models/records.go b/internal/models/records.go index 9b34aa5..9220509 100644 --- a/internal/models/records.go +++ b/internal/models/records.go @@ -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 { diff --git a/internal/models/records_test.go b/internal/models/records_test.go index d33dfed..083eb4e 100644 --- a/internal/models/records_test.go +++ b/internal/models/records_test.go @@ -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") diff --git a/internal/routes/routes.go b/internal/routes/routes.go index 50b5b64..bb47c63 100644 --- a/internal/routes/routes.go +++ b/internal/routes/routes.go @@ -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/") }) diff --git a/internal/server/integration_test.go b/internal/server/integration_test.go index 2d87151..14f8f23 100644 --- a/internal/server/integration_test.go +++ b/internal/server/integration_test.go @@ -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) diff --git a/internal/server/server.go b/internal/server/server.go index ea9d1ce..3b2a710 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -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") diff --git a/internal/server/server_test.go b/internal/server/server_test.go index fb5401f..8e76f13 100644 --- a/internal/server/server_test.go +++ b/internal/server/server_test.go @@ -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"}, diff --git a/internal/services/detector.go b/internal/services/detector.go index 617f598..9cefb97 100644 --- a/internal/services/detector.go +++ b/internal/services/detector.go @@ -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 } diff --git a/internal/services/engine.go b/internal/services/engine.go index 316245b..b19c834 100644 --- a/internal/services/engine.go +++ b/internal/services/engine.go @@ -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) } diff --git a/internal/utils/env.go b/internal/utils/env.go index a95b664..798fa8d 100644 --- a/internal/utils/env.go +++ b/internal/utils/env.go @@ -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) diff --git a/internal/utils/file.go b/internal/utils/file.go index 8e93251..6c1ea0a 100644 --- a/internal/utils/file.go +++ b/internal/utils/file.go @@ -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 { diff --git a/internal/utils/file_test.go b/internal/utils/file_test.go index 5e192ec..e426478 100644 --- a/internal/utils/file_test.go +++ b/internal/utils/file_test.go @@ -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) diff --git a/internal/utils/port.go b/internal/utils/port.go index 4acabf7..52027cc 100644 --- a/internal/utils/port.go +++ b/internal/utils/port.go @@ -2,7 +2,6 @@ package utils import "net" -// GetFreePort 获取一个未被占用的随机端口 func GetFreePort() (int, error) { addr, err := net.ResolveTCPAddr("tcp", "localhost:0") if err != nil { diff --git a/internal/utils/version.go b/internal/utils/version.go index 2fdd1a6..53f6560 100644 --- a/internal/utils/version.go +++ b/internal/utils/version.go @@ -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 } diff --git a/internal/version/version.go b/internal/version/version.go index 6b2286d..c702b9a 100644 --- a/internal/version/version.go +++ b/internal/version/version.go @@ -1,9 +1,7 @@ package version -// Version 会在编译时通过 -ldflags 注入 var Version = "v0.0.0-dev" -// GetVersion 获取版本号 func GetVersion() string { return Version } diff --git a/ui/ui.go b/ui/ui.go index eb246ee..33e9449 100644 --- a/ui/ui.go +++ b/ui/ui.go @@ -8,7 +8,6 @@ import ( //go:embed all:dist var distFS embed.FS -// GetDistFS returns the embedded dist filesystem func GetDistFS() (fs.FS, error) { return fs.Sub(distFS, "dist") }