chore: clean code

This commit is contained in:
xxnuo
2025-12-07 16:48:56 +08:00
parent 2b8e772412
commit fe0f663117
41 changed files with 109 additions and 642 deletions

View File

@@ -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)
}

View File

@@ -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)
}

View File

@@ -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)
}

View File

@@ -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(),

View File

@@ -1,3 +1,3 @@
package data
//go:generate env GOOS= GOARCH= go run gen_records.go
//go:generate go run gen_records.go

View File

@@ -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 (

View File

@@ -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)

View File

@@ -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)

View File

@@ -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)

View File

@@ -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{}{

View File

@@ -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 {

View File

@@ -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,

View File

@@ -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()

View File

@@ -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

View File

@@ -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() {

View File

@@ -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)

View File

@@ -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...)
}

View File

@@ -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 {

View File

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

View File

@@ -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)
}

View File

@@ -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)

View File

@@ -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,

View File

@@ -27,15 +27,12 @@ func TestManager_StartStop(t *testing.T) {
mgr := manager.NewManager(args)
defer mgr.Cleanup()
// 启动 Manager
err = mgr.Start()
require.NoError(t, err)
// 检查状态
assert.True(t, mgr.IsRunning())
assert.Equal(t, "running", mgr.Status())
// 停止 Manager
err = mgr.Stop()
require.NoError(t, err)
@@ -57,17 +54,14 @@ func TestManager_Restart(t *testing.T) {
mgr := manager.NewManager(args)
defer mgr.Cleanup()
// 启动
err = mgr.Start()
require.NoError(t, err)
assert.True(t, mgr.IsRunning())
// 重启
err = mgr.Restart()
require.NoError(t, err)
assert.True(t, mgr.IsRunning())
// 停止
err = mgr.Stop()
require.NoError(t, err)
}
@@ -122,7 +116,7 @@ func TestManager_Ready(t *testing.T) {
ctx := context.Background()
ready, err := mgr.Ready(ctx)
require.NoError(t, err)
assert.False(t, ready) // 未加载模型,应该返回 false
assert.False(t, ready)
}
func TestManager_PoweronPoweroff(t *testing.T) {
@@ -136,7 +130,7 @@ func TestManager_PoweronPoweroff(t *testing.T) {
args.Port = port
args.Host = "127.0.0.1"
args.EnableWebSocket = true
args.WorkDir = "../../testdata" // 假设有测试数据目录
args.WorkDir = "../../testdata"
mgr := manager.NewManager(args)
defer mgr.Cleanup()
@@ -149,18 +143,16 @@ func TestManager_PoweronPoweroff(t *testing.T) {
ctx := context.Background()
// 测试 Poweron这里会失败因为没有真实的模型文件
_, err = mgr.Poweron(ctx, manager.PoweronRequest{
Path: "nonexistent",
})
assert.Error(t, err) // 预期失败
assert.Error(t, err)
// 测试 Poweroff
_, err = mgr.Poweroff(ctx, manager.PoweroffRequest{
Time: 0,
Force: true,
})
// Poweroff 可能成功也可能失败,取决于引擎状态
t.Logf("Poweroff result: %v", err)
}
@@ -187,12 +179,11 @@ func TestManager_Reboot(t *testing.T) {
ctx := context.Background()
// 测试 Reboot未加载引擎时会失败
_, err = mgr.Reboot(ctx, manager.RebootRequest{
Time: 0,
Force: false,
})
assert.Error(t, err) // 预期失败
assert.Error(t, err)
}
func TestManager_Compute(t *testing.T) {
@@ -218,12 +209,11 @@ func TestManager_Compute(t *testing.T) {
ctx := context.Background()
// 测试 Compute未加载引擎时会失败
_, err = mgr.Compute(ctx, manager.ComputeRequest{
Text: "Hello",
HTML: false,
})
assert.Error(t, err) // 预期失败,因为引擎未加载
assert.Error(t, err)
}
func TestManager_Translate(t *testing.T) {
@@ -249,9 +239,8 @@ func TestManager_Translate(t *testing.T) {
ctx := context.Background()
// 测试 Translate未加载引擎时会失败
_, err = mgr.Translate(ctx, "Hello")
assert.Error(t, err) // 预期失败
assert.Error(t, err)
}
func TestManager_TranslateHTML(t *testing.T) {
@@ -277,9 +266,8 @@ func TestManager_TranslateHTML(t *testing.T) {
ctx := context.Background()
// 测试 TranslateHTML未加载引擎时会失败
_, err = mgr.TranslateHTML(ctx, "<p>Hello</p>")
assert.Error(t, err) // 预期失败
assert.Error(t, err)
}
func TestManager_MultipleManagers(t *testing.T) {
@@ -308,12 +296,10 @@ func TestManager_MultipleManagers(t *testing.T) {
time.Sleep(2 * time.Second)
// 验证所有 Manager 都在运行
for i, mgr := range managers {
assert.True(t, mgr.IsRunning(), "Manager %d should be running", i)
}
// 停止所有 Manager
for i, mgr := range managers {
err := mgr.Stop()
assert.NoError(t, err)
@@ -327,7 +313,6 @@ func TestManager_NotStarted(t *testing.T) {
mgr := manager.NewManager(args)
defer mgr.Cleanup()
// 未启动时调用方法应该返回错误
ctx := context.Background()
_, err := mgr.Ready(ctx)
@@ -344,8 +329,6 @@ func TestManager_FullWorkflow(t *testing.T) {
t.Skip("Skipping integration test in short mode")
}
// 这是一个完整的工作流测试
// 需要真实的模型文件才能完全通过
t.Skip("Requires real model files")
args := manager.NewWorkerArgs()
@@ -359,7 +342,6 @@ func TestManager_FullWorkflow(t *testing.T) {
mgr := manager.NewManager(args)
defer mgr.Cleanup()
// 1. 启动
err = mgr.Start()
require.NoError(t, err)
defer mgr.Stop()
@@ -368,38 +350,32 @@ func TestManager_FullWorkflow(t *testing.T) {
ctx := context.Background()
// 2. 检查就绪状态
ready, err := mgr.Ready(ctx)
require.NoError(t, err)
assert.False(t, ready)
// 3. 加载模型
resp, err := mgr.Poweron(ctx, manager.PoweronRequest{
Path: "path/to/model",
})
require.NoError(t, err)
assert.NotNil(t, resp)
// 4. 等待引擎就绪
time.Sleep(2 * time.Second)
ready, err = mgr.Ready(ctx)
require.NoError(t, err)
assert.True(t, ready)
// 5. 翻译文本
result, err := mgr.Translate(ctx, "Hello, world!")
require.NoError(t, err)
assert.NotEmpty(t, result)
t.Logf("Translation result: %s", result)
// 6. 翻译 HTML
htmlResult, err := mgr.TranslateHTML(ctx, "<p>Hello, world!</p>")
require.NoError(t, err)
assert.NotEmpty(t, htmlResult)
t.Logf("HTML translation result: %s", htmlResult)
// 7. 重启引擎
rebootResp, err := mgr.Reboot(ctx, manager.RebootRequest{
Time: 0,
Force: false,
@@ -407,7 +383,6 @@ func TestManager_FullWorkflow(t *testing.T) {
require.NoError(t, err)
assert.NotNil(t, rebootResp)
// 8. 关闭引擎
poweroffResp, err := mgr.Poweroff(ctx, manager.PoweroffRequest{
Time: 0,
Force: true,

View File

@@ -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{

View File

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

View File

@@ -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() {

View File

@@ -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]stringkey 为文件类型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 {

View File

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

View File

@@ -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/")
})

View File

@@ -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)

View File

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

View File

@@ -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"},

View File

@@ -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
}

View File

@@ -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)
}

View File

@@ -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)

View File

@@ -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 {

View File

@@ -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)

View File

@@ -2,7 +2,6 @@ package utils
import "net"
// GetFreePort 获取一个未被占用的随机端口
func GetFreePort() (int, error) {
addr, err := net.ResolveTCPAddr("tcp", "localhost:0")
if err != nil {

View File

@@ -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
}

View File

@@ -1,9 +1,7 @@
package version
// Version 会在编译时通过 -ldflags 注入
var Version = "v0.0.0-dev"
// GetVersion 获取版本号
func GetVersion() string {
return Version
}

View File

@@ -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")
}