ソースを参照

feat: relaunch process in user session when in Session 0

caesar 2 ヶ月 前
コミット
d9e18d5f12
2 ファイル変更103 行追加32 行削除
  1. 16 18
      main.go
  2. 87 14
      msgbox_windows.go

+ 16 - 18
main.go

@@ -20,9 +20,8 @@ var appIcon []byte
 
 const (
 	httpTimeout     = 15 * time.Second
-	tickInterval    = 10 * time.Minute
-	refreshMinutes  = 10
-	lowBalanceAlert = 1.0 // 单位:元 (CNY)
+	refreshMinutes  = 10                         // 自动刷新间隔(分钟)
+	lowBalanceAlert = 1.0                        // 低余额告警阈值(元 CNY)
 	defaultAPIURL   = "https://api.deepseek.com/user/balance"
 )
 
@@ -57,6 +56,7 @@ type balanceResult struct {
 }
 
 func main() {
+	relaunchInUserSession()
 	exeDir = executableDir()
 	apiKey, balanceURL = loadConfig()
 	systray.Run(onReady, func() {})
@@ -85,12 +85,14 @@ func onReady() {
 	autoRefresh := true
 	alertOn := false
 
-	balanceCh := make(chan balanceResult, 1) // timer 和首查均通过此 channel 传回结果
-	var ticker *time.Ticker                   // 仅 startTicker 写入
+	tickDuration := refreshMinutes * time.Minute
+
+	balanceCh := make(chan balanceResult, 1)
+	var ticker *time.Ticker    // 仅 startTicker 写入
 	var tickerStop chan struct{}
 
 	startTicker := func() {
-		ticker = time.NewTicker(tickInterval)
+		ticker = time.NewTicker(tickDuration)
 		tickerStop = make(chan struct{})
 		stop := tickerStop
 
@@ -268,17 +270,16 @@ func alertLowBalance(cny float64) {
 }
 
 func showAbout(autoRefresh, alertOn bool) {
-	auto := "否"
-	if autoRefresh {
-		auto = "是"
-	}
-	alert := "否"
-	if alertOn {
-		alert = "是"
-	}
 	showMessageBox("DeepSeek Balance Tray",
 		fmt.Sprintf("自动刷新: %s\n刷新间隔: %d 分钟\n低余额告警: %s (阈值 ¥%.0f)",
-			auto, refreshMinutes, alert, lowBalanceAlert))
+			yesNo(autoRefresh), refreshMinutes, yesNo(alertOn), lowBalanceAlert))
+}
+
+func yesNo(b bool) string {
+	if b {
+		return "是"
+	}
+	return "否"
 }
 
 // ---------- 配置加载 ----------
@@ -325,9 +326,6 @@ func readDotEnv(path string) map[string]string {
 			continue
 		}
 		k := strings.TrimSpace(line[:eq])
-		if k == "" {
-			continue
-		}
 		v := strings.TrimSpace(line[eq+1:])
 		v = strings.Trim(v, `"'`)
 		m[k] = v

+ 87 - 14
msgbox_windows.go

@@ -1,14 +1,22 @@
 package main
 
 import (
+	"os"
 	"os/exec"
 	"syscall"
 	"unsafe"
 )
 
 var (
-	user32  = syscall.NewLazyDLL("user32.dll")
-	msgBoxW = user32.NewProc("MessageBoxW")
+	user32   = syscall.NewLazyDLL("user32.dll")
+	kernel32 = syscall.NewLazyDLL("kernel32.dll")
+	wtsapi32 = syscall.NewLazyDLL("wtsapi32.dll")
+	msgBoxW          = user32.NewProc("MessageBoxW")
+	consoleSessionID = kernel32.NewProc("WTSGetActiveConsoleSessionId")
+	queryUserToken   = wtsapi32.NewProc("WTSQueryUserToken")
+	createEnvBlock   = user32.NewProc("CreateEnvironmentBlock")
+	destroyEnvBlock  = user32.NewProc("DestroyEnvironmentBlock")
+	createAsUser     = kernel32.NewProc("CreateProcessAsUserW")
 )
 
 const (
@@ -18,19 +26,15 @@ const (
 	mbTopmost       = 0x00040000
 )
 
-func showMessageBox(title, msg string) {
-	ptitle, err := syscall.UTF16PtrFromString(title)
-	if err != nil {
-		return
-	}
-	ptext, err := syscall.UTF16PtrFromString(msg)
-	if err != nil {
-		return
-	}
+func mustUTF16Ptr(s string) *uint16 {
+	p, _ := syscall.UTF16PtrFromString(s) // NUL-free strings; ignore impossible error
+	return p
+}
 
+func showMessageBox(title, msg string) {
 	msgBoxW.Call(0,
-		uintptr(unsafe.Pointer(ptext)),
-		uintptr(unsafe.Pointer(ptitle)),
+		uintptr(unsafe.Pointer(mustUTF16Ptr(msg))),
+		uintptr(unsafe.Pointer(mustUTF16Ptr(title))),
 		uintptr(mbOK|mbIconWarning|mbSetForeground|mbTopmost),
 	)
 }
@@ -40,4 +44,73 @@ func runNotepad(path string) {
 	if err != nil {
 		showMessageBox("错误", "无法打开记事本: "+err.Error())
 	}
-}
+}
+
+// relaunchInUserSession 若当前进程在 Session 0(如服务)中运行,
+// 则在活跃用户会话中重新启动自身,然后退出当前进程。
+// 调用方应在 main() 最开始调用此函数。
+func relaunchInUserSession() {
+	sessionID, _, _ := consoleSessionID.Call()
+	if sessionID == 0 {
+		// 当前已在用户会话或 sessionID 获取失败,不做处理
+		return
+	}
+
+	// 仅在 Session 0 时才需要重新启动
+	var ourSession uint32
+	kernel32.NewProc("ProcessIdToSessionId").Call(
+		uintptr(os.Getpid()), uintptr(unsafe.Pointer(&ourSession)),
+	)
+	if ourSession != 0 {
+		return
+	}
+
+	// 获取用户令牌
+	var token syscall.Handle
+	ret, _, _ := queryUserToken.Call(sessionID, uintptr(unsafe.Pointer(&token)))
+	if ret == 0 || token == 0 {
+		return
+	}
+	defer syscall.CloseHandle(token)
+
+	// 创建用户环境块
+	var envBlock unsafe.Pointer
+	ret, _, _ = createEnvBlock.Call(uintptr(unsafe.Pointer(&envBlock)), uintptr(token), 0)
+	if ret == 0 {
+		return
+	}
+	defer destroyEnvBlock.Call(uintptr(envBlock))
+
+	// 获取当前 EXE 路径
+	exePath, err := os.Executable()
+	if err != nil {
+		return
+	}
+	exePtr := mustUTF16Ptr(exePath)
+
+	// 准备 STARTUPINFO 和 PROCESS_INFORMATION
+	var si syscall.StartupInfo
+	si.Cb = uint32(unsafe.Sizeof(si))
+	var pi syscall.ProcessInformation
+
+	ret, _, _ = createAsUser.Call(
+		uintptr(token),
+		0,
+		uintptr(unsafe.Pointer(exePtr)),
+		0,
+		0,
+		0,
+		0, // no extra flags; -H windowsgui already linked
+		uintptr(unsafe.Pointer(envBlock)),
+		0,
+		uintptr(unsafe.Pointer(&si)),
+		uintptr(unsafe.Pointer(&pi)),
+	)
+	if ret != 0 {
+		syscall.CloseHandle(pi.Thread)
+		syscall.CloseHandle(pi.Process)
+	}
+
+	os.Exit(0)
+}
+