diff --git a/internal/daemon/daemon.go b/internal/daemon/daemon.go index 056badc..cc86c8b 100644 --- a/internal/daemon/daemon.go +++ b/internal/daemon/daemon.go @@ -129,6 +129,8 @@ func (d *daemon) startSession(conn net.Conn, req ipc.Request) { cfg := vpn.SessionConfig{ ServerURL: req.Config.ServerURL, + SNIHost: req.Config.SNIHost, + ServerIPs: req.Config.ServerIPs, Username: req.Config.Username, Password: req.Config.Password, Token: req.Config.Token, diff --git a/internal/db/db.go b/internal/db/db.go index 14a7463..4562c1e 100644 --- a/internal/db/db.go +++ b/internal/db/db.go @@ -8,6 +8,7 @@ package db import ( "database/sql" "fmt" + "strings" "lmvpn/internal/paths" @@ -40,15 +41,139 @@ func (s *Store) Close() error { } func (s *Store) migrate() error { - _, err := s.db.Exec(schema) - return err + _, err := s.db.Exec(schemaV2) + if err != nil { + return err + } + return s.migrateV2() } -const schema = ` +func (s *Store) migrateV2() error { + // Detect if migration is needed by checking if old server_url column + // still exists. If protocol column is present, v2 is already in place. + _, err := s.db.Exec(`SELECT protocol, host, server_ips, port, path FROM server_profiles LIMIT 0`) + if err == nil { + return nil // already v2 + } + + // Check if old table exists. + var count int + err = s.db.QueryRow(`SELECT COUNT(*) FROM sqlite_master WHERE type='table' AND name='server_profiles' AND sql LIKE '%server_url%'`).Scan(&count) + if err != nil || count == 0 { + return nil // nothing to migrate + } + + tx, err := s.db.Begin() + if err != nil { + return fmt.Errorf("migrate v2 begin: %w", err) + } + defer tx.Rollback() + + _, err = tx.Exec(`ALTER TABLE server_profiles RENAME TO server_profiles_old`) + if err != nil { + return fmt.Errorf("migrate v2 rename: %w", err) + } + + _, err = tx.Exec(schemaV2) + if err != nil { + return fmt.Errorf("migrate v2 create new: %w", err) + } + + rows, err := tx.Query(`SELECT id, name, server_url, username, auth_mode, routing_mode, custom_cidrs, mtu_override, auto_connect, created_at, last_connected_at FROM server_profiles_old`) + if err != nil { + return fmt.Errorf("migrate v2 read old: %w", err) + } + defer rows.Close() + + insert, err := tx.Prepare(`INSERT INTO server_profiles (id, name, protocol, host, server_ips, port, path, username, auth_mode, routing_mode, custom_cidrs, mtu_override, auto_connect, created_at, last_connected_at) VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?,?,?)`) + if err != nil { + return fmt.Errorf("migrate v2 prepare insert: %w", err) + } + defer insert.Close() + + for rows.Next() { + var id int64 + var name, serverURL, username, authMode, routingMode, customCIDRs string + var mtuOverride int + var autoConnect int + var createdAt, lastConnectedAt sql.NullString + if err := rows.Scan(&id, &name, &serverURL, &username, &authMode, &routingMode, &customCIDRs, &mtuOverride, &autoConnect, &createdAt, &lastConnectedAt); err != nil { + return fmt.Errorf("migrate v2 scan: %w", err) + } + protocol, host, ips, port, path := parseOldURL(serverURL) + _, err = insert.Exec(id, name, protocol, host, ips, port, path, username, authMode, routingMode, customCIDRs, mtuOverride, autoConnect, nullStr(createdAt), nullStr(lastConnectedAt)) + if err != nil { + return fmt.Errorf("migrate v2 insert: %w", err) + } + } + + if err := rows.Err(); err != nil { + return err + } + + _, err = tx.Exec(`DROP TABLE server_profiles_old`) + if err != nil { + return fmt.Errorf("migrate v2 drop old: %w", err) + } + + return tx.Commit() +} + +func parseOldURL(raw string) (protocol, host, ips, path string, port int) { + u := raw + switch { + case strings.HasPrefix(u, "wss://"): + protocol = "wss" + u = u[6:] + case strings.HasPrefix(u, "ws://"): + protocol = "ws" + u = u[5:] + default: + protocol = "wss" + } + + path = "/ws" + if i := strings.IndexByte(u, '/'); i >= 0 { + path = u[i:] + u = u[:i] + } + + port = 443 + if i := strings.LastIndexByte(u, ':'); i >= 0 { + if p, err := stringToInt(u[i+1:]); err == nil { + port = p + u = u[:i] + } + } + + host = u + return +} + +func stringToInt(s string) (int, error) { + var n int + for _, c := range s { + if c < '0' || c > '9' { + return 0, fmt.Errorf("not a number: %s", s) + } + n = n*10 + int(c-'0') + } + return n, nil +} + +func nullStr(s sql.NullString) sql.NullString { + return s +} + +const schemaV2 = ` CREATE TABLE IF NOT EXISTS server_profiles ( id INTEGER PRIMARY KEY AUTOINCREMENT, name TEXT NOT NULL, - server_url TEXT NOT NULL, + protocol TEXT NOT NULL DEFAULT 'wss', + host TEXT NOT NULL, + server_ips TEXT NOT NULL DEFAULT '', + port INTEGER NOT NULL DEFAULT 443, + path TEXT NOT NULL DEFAULT '/ws', username TEXT NOT NULL, auth_mode TEXT NOT NULL DEFAULT 'both', routing_mode TEXT NOT NULL DEFAULT 'full', diff --git a/internal/db/profile.go b/internal/db/profile.go index 64a014e..5bdad66 100644 --- a/internal/db/profile.go +++ b/internal/db/profile.go @@ -12,10 +12,12 @@ import ( func (s *Store) CreateProfile(p *model.ServerProfile) (int64, error) { res, err := s.db.Exec( `INSERT INTO server_profiles - (name, server_url, username, auth_mode, routing_mode, + (name, protocol, host, server_ips, port, path, + username, auth_mode, routing_mode, custom_cidrs, mtu_override, auto_connect) - VALUES (?,?,?,?,?,?,?,?)`, - p.Name, p.ServerURL, p.Username, p.AuthMode, p.RoutingMode, + VALUES (?,?,?,?,?,?,?,?,?,?,?,?)`, + p.Name, p.Protocol, p.Host, p.ServerIPs, p.Port, p.Path, + p.Username, p.AuthMode, p.RoutingMode, p.CustomCIDRs, p.MTUOverride, p.AutoConnect, ) if err != nil { @@ -32,12 +34,13 @@ func (s *Store) GetProfile(id int64) (*model.ServerProfile, error) { p := &model.ServerProfile{} var last sql.NullTime err := s.db.QueryRow( - `SELECT id, name, server_url, username, auth_mode, routing_mode, + `SELECT id, name, protocol, host, server_ips, port, path, + username, auth_mode, routing_mode, custom_cidrs, mtu_override, auto_connect, created_at, last_connected_at FROM server_profiles WHERE id = ?`, id, - ).Scan(&p.ID, &p.Name, &p.ServerURL, &p.Username, &p.AuthMode, - &p.RoutingMode, &p.CustomCIDRs, &p.MTUOverride, &p.AutoConnect, - &p.CreatedAt, &last) + ).Scan(&p.ID, &p.Name, &p.Protocol, &p.Host, &p.ServerIPs, &p.Port, &p.Path, + &p.Username, &p.AuthMode, &p.RoutingMode, &p.CustomCIDRs, &p.MTUOverride, + &p.AutoConnect, &p.CreatedAt, &last) if err != nil { return nil, fmt.Errorf("get profile %d: %w", id, err) } @@ -50,7 +53,8 @@ func (s *Store) GetProfile(id int64) (*model.ServerProfile, error) { // ListProfiles returns all saved profiles ordered by name. func (s *Store) ListProfiles() ([]model.ServerProfile, error) { rows, err := s.db.Query( - `SELECT id, name, server_url, username, auth_mode, routing_mode, + `SELECT id, name, protocol, host, server_ips, port, path, + username, auth_mode, routing_mode, custom_cidrs, mtu_override, auto_connect, created_at, last_connected_at FROM server_profiles ORDER BY name`) if err != nil { @@ -62,9 +66,9 @@ func (s *Store) ListProfiles() ([]model.ServerProfile, error) { for rows.Next() { var p model.ServerProfile var last sql.NullTime - if err := rows.Scan(&p.ID, &p.Name, &p.ServerURL, &p.Username, - &p.AuthMode, &p.RoutingMode, &p.CustomCIDRs, &p.MTUOverride, - &p.AutoConnect, &p.CreatedAt, &last); err != nil { + if err := rows.Scan(&p.ID, &p.Name, &p.Protocol, &p.Host, &p.ServerIPs, + &p.Port, &p.Path, &p.Username, &p.AuthMode, &p.RoutingMode, + &p.CustomCIDRs, &p.MTUOverride, &p.AutoConnect, &p.CreatedAt, &last); err != nil { return nil, err } if last.Valid { @@ -79,11 +83,13 @@ func (s *Store) ListProfiles() ([]model.ServerProfile, error) { func (s *Store) UpdateProfile(p *model.ServerProfile) error { _, err := s.db.Exec( `UPDATE server_profiles SET - name = ?, server_url = ?, username = ?, auth_mode = ?, - routing_mode = ?, custom_cidrs = ?, mtu_override = ?, auto_connect = ? + name = ?, protocol = ?, host = ?, server_ips = ?, port = ?, path = ?, + username = ?, auth_mode = ?, routing_mode = ?, + custom_cidrs = ?, mtu_override = ?, auto_connect = ? WHERE id = ?`, - p.Name, p.ServerURL, p.Username, p.AuthMode, - p.RoutingMode, p.CustomCIDRs, p.MTUOverride, p.AutoConnect, p.ID) + p.Name, p.Protocol, p.Host, p.ServerIPs, p.Port, p.Path, + p.Username, p.AuthMode, p.RoutingMode, + p.CustomCIDRs, p.MTUOverride, p.AutoConnect, p.ID) if err != nil { return fmt.Errorf("update profile %d: %w", p.ID, err) } diff --git a/internal/i18n/en.toml b/internal/i18n/en.toml index 1658b63..9ad2271 100644 --- a/internal/i18n/en.toml +++ b/internal/i18n/en.toml @@ -34,7 +34,7 @@ DlgDeleteProfileTitle = "Delete Profile" DlgDeleteProfileMsg = 'Delete profile "{{.name}}" and its stored credentials?' DlgProfileTitle = "Profile" DlgValidationTitle = "Validation" -DlgValidationMsg = "Name, Server URL, and Username are required." +DlgValidationMsg = "Name, Host, and Username are required." DlgDaemonError = "Daemon Error" DlgCredentialError = "Credential Error" DlgCredentialErrorMsg = "No password stored for this profile. Edit the profile to set it." @@ -44,6 +44,10 @@ DlgSaveError = "Save Error" DlgKeychainError = "Keychain Error" DlgError = "Error" +BtnResetDB = "Reset Database" +DlgResetDBTitle = "Reset Database" +DlgResetDBMsg = "Delete all profiles and connection logs? This cannot be undone." + TrayShowWindow = "Show Window" TrayConnect = "Connect" TrayDisconnect = "Disconnect" @@ -52,7 +56,11 @@ TrayLanguageAuto = "Auto" TrayQuit = "Quit" FieldName = "Name" -FieldServerURL = "Server URL" +FieldProtocol = "Protocol" +FieldHost = "Host" +FieldServerIPs = "Server IPs (CDN)" +FieldPort = "Port" +FieldPath = "Path" FieldUsername = "Username" FieldPassword = "Password" FieldAuthMode = "Auth Mode" @@ -63,6 +71,7 @@ FieldMTUOverride = "MTU Override" PlaceholderCIDRs = "10.0.0.0/8, 172.16.0.0/12" PlaceholderMTU = "0 = use server MTU" PlaceholderPasswordUnchanged = "(unchanged)" +PlaceholderServerIPs = "e.g. 1.2.3.4, 5.6.7.8" AuthModeBoth = "Both (JWT + Password)" AuthModeJWT = "JWT" diff --git a/internal/i18n/zh-Hans.toml b/internal/i18n/zh-Hans.toml index c577630..f97bcdf 100644 --- a/internal/i18n/zh-Hans.toml +++ b/internal/i18n/zh-Hans.toml @@ -34,7 +34,7 @@ DlgDeleteProfileTitle = "删除配置" DlgDeleteProfileMsg = '删除配置"{{.name}}"及其存储的凭据?' DlgProfileTitle = "配置" DlgValidationTitle = "验证" -DlgValidationMsg = "名称、服务器地址和用户名为必填项。" +DlgValidationMsg = "名称、主机名和用户名为必填项。" DlgDaemonError = "守护进程错误" DlgCredentialError = "凭据错误" DlgCredentialErrorMsg = "此配置未存储密码。请编辑配置以设置密码。" @@ -44,6 +44,10 @@ DlgSaveError = "保存错误" DlgKeychainError = "钥匙串错误" DlgError = "错误" +BtnResetDB = "重置数据库" +DlgResetDBTitle = "重置数据库" +DlgResetDBMsg = "删除所有配置和连接记录?此操作无法撤销。" + TrayShowWindow = "显示窗口" TrayConnect = "连接" TrayDisconnect = "断开连接" @@ -52,7 +56,11 @@ TrayLanguageAuto = "自动" TrayQuit = "退出" FieldName = "名称" -FieldServerURL = "服务器地址" +FieldProtocol = "协议" +FieldHost = "主机名" +FieldServerIPs = "服务器 IP(CDN 优选)" +FieldPort = "端口" +FieldPath = "路径" FieldUsername = "用户名" FieldPassword = "密码" FieldAuthMode = "认证方式" @@ -63,6 +71,7 @@ FieldMTUOverride = "MTU 覆盖" PlaceholderCIDRs = "10.0.0.0/8, 172.16.0.0/12" PlaceholderMTU = "0 = 使用服务器 MTU" PlaceholderPasswordUnchanged = "(未更改)" +PlaceholderServerIPs = "例: 1.2.3.4, 5.6.7.8" AuthModeBoth = "全部(JWT + 密码)" AuthModeJWT = "JWT" diff --git a/internal/ipc/ipc.go b/internal/ipc/ipc.go index 9dce92f..78b6cb3 100644 --- a/internal/ipc/ipc.go +++ b/internal/ipc/ipc.go @@ -5,8 +5,8 @@ // Protocol: newline-delimited JSON. Each message is one JSON object // followed by '\n'. // -// GUI → daemon: Request (start, stop, shutdown, stats) -// daemon → GUI: Event (state, stats, error) +// GUI → daemon: Request (start, stop, shutdown, stats) +// daemon → GUI: Event (state, stats, error) package ipc import ( @@ -49,6 +49,8 @@ type Request struct { // package (which needs root-only TUN) into the GUI. type ClientConfig struct { ServerURL string `json:"server_url"` + SNIHost string `json:"sni_host"` // TLS SNI hostname for CDN + ServerIPs []string `json:"server_ips"` // CDN edge IP list for failover Username string `json:"username"` Password string `json:"password"` Token string `json:"token"` diff --git a/internal/model/model.go b/internal/model/model.go index fd32b5f..575a7c1 100644 --- a/internal/model/model.go +++ b/internal/model/model.go @@ -2,7 +2,11 @@ // exchanged between application layers. package model -import "time" +import ( + "fmt" + "strings" + "time" +) // AuthMode selects how the client authenticates to a server. type AuthMode string @@ -26,7 +30,11 @@ const ( type ServerProfile struct { ID int64 `json:"id"` Name string `json:"name"` - ServerURL string `json:"server_url"` // e.g. wss://vpn.example.com/ws + Protocol string `json:"protocol"` // "wss" (default) or "ws" + Host string `json:"host"` // hostname for SNI, e.g. vpn.example.com + ServerIPs string `json:"server_ips"` // comma-separated CDN IPs, first used by default + Port int `json:"port"` // default 443 + Path string `json:"path"` // default "/ws" Username string `json:"username"` AuthMode AuthMode `json:"auth_mode"` RoutingMode RoutingMode `json:"routing_mode"` @@ -37,6 +45,58 @@ type ServerProfile struct { LastConnectedAt *time.Time `json:"last_connected_at"` } +// BuildServerURL constructs the WebSocket URL from the profile fields. +// If ip is provided, it is used as the host portion instead of Host +// (for CDN edge IP connections). +// Default ports are omitted from the URL (443 for wss, 80 for ws). +func (p *ServerProfile) BuildServerURL(ip ...string) string { + protocol := p.Protocol + if protocol == "" { + protocol = "wss" + } + + host := p.Host + if len(ip) > 0 && ip[0] != "" { + host = ip[0] + } + + port := p.Port + if port == 0 { + port = 443 + } + + path := p.Path + if path == "" { + path = "/ws" + } + if !strings.HasPrefix(path, "/") { + path = "/" + path + } + + isDefaultPort := (protocol == "wss" && port == 443) || (protocol == "ws" && port == 80) + + if isDefaultPort { + return fmt.Sprintf("%s://%s%s", protocol, host, path) + } + return fmt.Sprintf("%s://%s:%d%s", protocol, host, port, path) +} + +// GetServerIPList parses ServerIPs into a string slice. +func (p *ServerProfile) GetServerIPList() []string { + if p.ServerIPs == "" { + return nil + } + parts := strings.Split(p.ServerIPs, ",") + var out []string + for _, part := range parts { + s := strings.TrimSpace(part) + if s != "" { + out = append(out, s) + } + } + return out +} + // ConnectionStatus records the outcome of a connection attempt. type ConnectionStatus string diff --git a/internal/transport/transport.go b/internal/transport/transport.go index 6138920..e3d679b 100644 --- a/internal/transport/transport.go +++ b/internal/transport/transport.go @@ -15,6 +15,7 @@ package transport import ( "context" + "crypto/tls" "encoding/json" "fmt" "net/http" @@ -30,6 +31,7 @@ import ( // HandshakeConfig configures a single connection attempt. type HandshakeConfig struct { ServerURL string // e.g. wss://vpn.example.com/ws + SNIHost string // TLS SNI hostname for CDN edge connections Token string // JWT; if non-empty, used via ?token= (method A) Username string // for password auth (method B), or fallback Password string // for password auth (method B), or fallback @@ -55,23 +57,30 @@ func Connect(ctx context.Context, cfg HandshakeConfig) (*Conn, error) { WriteBufferSize: 4096, // match server (handler.go:18) } - // Build URL: append ?token= for JWT auth. - url := cfg.ServerURL - if cfg.Token != "" { - url = appendQuery(url, "token", cfg.Token) + if cfg.SNIHost != "" { + dialer.TLSClientConfig = &tls.Config{ + ServerName: cfg.SNIHost, + } + } + + // Build URL: append ?token= for JWT auth. + urlStr := cfg.ServerURL + if cfg.Token != "" { + urlStr = appendQuery(urlStr, "token", cfg.Token) } - // Omit Origin header (server allows empty Origin for non-browser - // clients — handler.go:19-29). header := http.Header{} header.Set("Origin", "") + if cfg.SNIHost != "" { + header.Set("Host", cfg.SNIHost) + } - ws, resp, err := dialer.DialContext(ctx, url, header) + ws, resp, err := dialer.DialContext(ctx, urlStr, header) if err != nil { if resp != nil { resp.Body.Close() } - return nil, fmt.Errorf("dial %s: %w", url, err) + return nil, fmt.Errorf("dial %s: %w", urlStr, err) } defer resp.Body.Close() diff --git a/internal/ui/app.go b/internal/ui/app.go index b4bee68..72e0c88 100644 --- a/internal/ui/app.go +++ b/internal/ui/app.go @@ -1,6 +1,7 @@ package ui import ( + "os" "sync" "lmvpn/internal/config" @@ -181,6 +182,60 @@ func (a *App) onDeleteProfile() { }, a.window).Show() } +// onResetDB deletes the SQLite database file after confirmation, +// then re-creates it. All profiles, credentials, and logs are lost. +func (a *App) onResetDB() { + dialog.NewCustomConfirm(i18n.T("DlgResetDBTitle"), + i18n.T("BtnDelete"), i18n.T("BtnCancel"), + widget.NewLabel(i18n.T("DlgResetDBMsg")), + func(ok bool) { + if !ok { + return + } + + // Disconnect if connected. + a.mu.Lock() + client := a.ipcClient + a.mu.Unlock() + if client != nil { + _ = ipc.SendStop(client) + } + + // Clear keychain entries for all profiles. + for _, p := range a.profiles { + _ = a.kc.DeleteAll(p.Name) + } + + // Close and delete database. + if a.db != nil { + a.db.Close() + } + if err := os.Remove(paths.DBPath()); err != nil { + showError(i18n.T("DlgError"), err.Error(), a.window) + return + } + + // Re-open (auto-creates new database). + store, err := db.Open() + if err != nil { + showError(i18n.T("DlgError"), err.Error(), a.window) + return + } + a.db = store + + // Reset UI state. + a.currentProfile = nil + a.loadProfiles() + a.stateLabel.SetText(i18n.T("StateDisconnected")) + a.ipLabel.SetText(i18n.T("IpNone")) + a.uptimeLabel.SetText(i18n.T("UptimeNone")) + a.rxLabel.SetText(i18n.T("RxZero")) + a.txLabel.SetText(i18n.T("TxZero")) + a.connectBtn.Enable() + a.disconnectBtn.Disable() + }, a.window).Show() +} + // changeLanguage switches the active language, persists the choice to // the config file, and rebuilds the UI so the new strings take effect // immediately. diff --git a/internal/ui/profile.go b/internal/ui/profile.go index 617433f..75dc0f9 100644 --- a/internal/ui/profile.go +++ b/internal/ui/profile.go @@ -17,6 +17,8 @@ import ( var ( authCodes = []string{string(model.AuthModeBoth), string(model.AuthModeJWT), string(model.AuthModePassword)} routeCodes = []string{string(model.RoutingFull), string(model.RoutingSplit), string(model.RoutingCustom)} + + protoCodes = []string{"wss", "ws"} ) func authModeLabels() []string { @@ -27,6 +29,10 @@ func routeModeLabels() []string { return []string{i18n.T("RoutingModeFull"), i18n.T("RoutingModeSplit"), i18n.T("RoutingModeCustom")} } +func protoLabels() []string { + return []string{"wss", "ws"} +} + // codeIndex returns the position of code in codes, or 0 if not found. func codeIndex(codes []string, code string) int { for i, c := range codes { @@ -57,7 +63,12 @@ func (a *App) showProfileDialog(editing *model.ServerProfile) { isNew := editing == nil nameEntry := widget.NewEntry() - serverEntry := widget.NewEntry() + protoSelect := widget.NewSelect(protoLabels(), nil) + hostEntry := widget.NewEntry() + ipsEntry := widget.NewEntry() + ipsEntry.SetPlaceHolder(i18n.T("PlaceholderServerIPs")) + portEntry := widget.NewEntry() + pathEntry := widget.NewEntry() userEntry := widget.NewEntry() passEntry := widget.NewPasswordEntry() authSelect := widget.NewSelect(authModeLabels(), nil) @@ -69,7 +80,13 @@ func (a *App) showProfileDialog(editing *model.ServerProfile) { if !isNew { nameEntry.SetText(editing.Name) - serverEntry.SetText(editing.ServerURL) + protoSelect.SetSelectedIndex(codeIndex(protoCodes, editing.Protocol)) + hostEntry.SetText(editing.Host) + ipsEntry.SetText(editing.ServerIPs) + if editing.Port > 0 { + portEntry.SetText(fmtInt(editing.Port)) + } + pathEntry.SetText(editing.Path) userEntry.SetText(editing.Username) authSelect.SetSelectedIndex(codeIndex(authCodes, string(editing.AuthMode))) routeSelect.SetSelectedIndex(codeIndex(routeCodes, string(editing.RoutingMode))) @@ -77,18 +94,39 @@ func (a *App) showProfileDialog(editing *model.ServerProfile) { mtuEntry.SetText(fmtInt(editing.MTUOverride)) passEntry.SetPlaceHolder(i18n.T("PlaceholderPasswordUnchanged")) } else { + protoSelect.SetSelectedIndex(0) // wss + portEntry.SetText("443") + pathEntry.SetText("/ws") authSelect.SetSelectedIndex(codeIndex(authCodes, string(model.AuthModeBoth))) routeSelect.SetSelectedIndex(codeIndex(routeCodes, string(model.RoutingFull))) mtuEntry.SetText("0") } form := container.NewVBox( - widget.NewLabel(i18n.T("FieldName")), nameEntry, - widget.NewLabel(i18n.T("FieldServerURL")), serverEntry, - widget.NewLabel(i18n.T("FieldUsername")), userEntry, - widget.NewLabel(i18n.T("FieldPassword")), passEntry, - widget.NewLabel(i18n.T("FieldAuthMode")), authSelect, - widget.NewLabel(i18n.T("FieldRoutingMode")), routeSelect, + widget.NewLabel(i18n.T("FieldName")), + nameEntry, + + container.NewBorder(nil, nil, + container.NewVBox(widget.NewLabel(i18n.T("FieldProtocol")), protoSelect), + container.NewVBox(widget.NewLabel(i18n.T("FieldPort")), portEntry), + container.NewVBox(widget.NewLabel(i18n.T("FieldHost")), hostEntry), + ), + + container.NewGridWithColumns(2, + container.NewVBox(widget.NewLabel(i18n.T("FieldPath")), pathEntry), + container.NewVBox(widget.NewLabel(i18n.T("FieldServerIPs")), ipsEntry), + ), + + container.NewGridWithColumns(2, + container.NewVBox(widget.NewLabel(i18n.T("FieldUsername")), userEntry), + container.NewVBox(widget.NewLabel(i18n.T("FieldPassword")), passEntry), + ), + + container.NewGridWithColumns(2, + container.NewVBox(widget.NewLabel(i18n.T("FieldAuthMode")), authSelect), + container.NewVBox(widget.NewLabel(i18n.T("FieldRoutingMode")), routeSelect), + ), + widget.NewLabel(i18n.T("FieldCustomCIDRs")), cidrEntry, widget.NewLabel(i18n.T("FieldMTUOverride")), mtuEntry, ) @@ -101,12 +139,17 @@ func (a *App) showProfileDialog(editing *model.ServerProfile) { }) saveBtn := widget.NewButton(i18n.T("BtnSave"), func() { - a.saveProfile(editing, nameEntry.Text, serverEntry.Text, + if a.saveProfile(editing, + nameEntry.Text, + selectedCode(protoCodes, protoSelect.SelectedIndex()), + hostEntry.Text, ipsEntry.Text, + portEntry.Text, pathEntry.Text, userEntry.Text, passEntry.Text, selectedCode(authCodes, authSelect.SelectedIndex()), selectedCode(routeCodes, routeSelect.SelectedIndex()), - cidrEntry.Text, mtuEntry.Text, isNew) - profileWin.Close() + cidrEntry.Text, mtuEntry.Text, isNew) { + profileWin.Close() + } }) saveBtn.Importance = widget.HighImportance @@ -114,25 +157,31 @@ func (a *App) showProfileDialog(editing *model.ServerProfile) { profileWin.Close() }) - profileWin.SetContent(container.NewBorder(nil, container.NewHBox(saveBtn, cancelBtn), nil, nil, form)) + profileWin.SetContent(container.NewBorder(nil, container.NewHBox(saveBtn, cancelBtn), nil, nil, container.NewVScroll(form))) profileWin.Resize(fyne.NewSize(460, 560)) profileWin.Show() } // saveProfile creates or updates a profile and stores credentials. +// Returns true on success, false if validation or DB operation failed. func (a *App) saveProfile(editing *model.ServerProfile, - name, server, user, password, authMode, routeMode, cidrs, mtuStr string, isNew bool) { - if name == "" || server == "" || user == "" { + name, protocol, host, ips, portStr, pathStr, user, password, authMode, routeMode, cidrs, mtuStr string, isNew bool) bool { + if name == "" || host == "" || user == "" { showError(i18n.T("DlgValidationTitle"), i18n.T("DlgValidationMsg"), a.window) - return + return false } + port := parseIntDefault(portStr, 443) mtu := parseIntDefault(mtuStr, 0) if isNew { p := &model.ServerProfile{ Name: name, - ServerURL: server, + Protocol: protocol, + Host: host, + ServerIPs: ips, + Port: port, + Path: pathStr, Username: user, AuthMode: model.AuthMode(authMode), RoutingMode: model.RoutingMode(routeMode), @@ -142,7 +191,7 @@ func (a *App) saveProfile(editing *model.ServerProfile, id, err := a.db.CreateProfile(p) if err != nil { showError(i18n.T("DlgSaveError"), err.Error(), a.window) - return + return false } _ = id if password != "" { @@ -153,7 +202,11 @@ func (a *App) saveProfile(editing *model.ServerProfile, } else { oldName := editing.Name editing.Name = name - editing.ServerURL = server + editing.Protocol = protocol + editing.Host = host + editing.ServerIPs = ips + editing.Port = port + editing.Path = pathStr editing.Username = user editing.AuthMode = model.AuthMode(authMode) editing.RoutingMode = model.RoutingMode(routeMode) @@ -161,7 +214,7 @@ func (a *App) saveProfile(editing *model.ServerProfile, editing.MTUOverride = mtu if err := a.db.UpdateProfile(editing); err != nil { showError(i18n.T("DlgSaveError"), err.Error(), a.window) - return + return false } if password != "" { _ = a.kc.DeleteAll(oldName) @@ -172,6 +225,7 @@ func (a *App) saveProfile(editing *model.ServerProfile, } a.loadProfiles() + return true } func fmtInt(n int) string { diff --git a/internal/ui/view.go b/internal/ui/view.go index 9b50ba3..04de61a 100644 --- a/internal/ui/view.go +++ b/internal/ui/view.go @@ -51,6 +51,8 @@ func (a *App) buildMainWindow() fyne.CanvasObject { editBtn := widget.NewButton(i18n.T("BtnEdit"), a.onEditProfile) deleteBtn := widget.NewButton(i18n.T("BtnDelete"), a.onDeleteProfile) + resetDBBtn := widget.NewButton(i18n.T("BtnResetDB"), a.onResetDB) + buttons := container.NewGridWithColumns(2, a.connectBtn, a.disconnectBtn, ) @@ -63,6 +65,7 @@ func (a *App) buildMainWindow() fyne.CanvasObject { a.profileSelect, buttons, profileButtons, + resetDBBtn, statusCard, ) } @@ -109,15 +112,25 @@ func (a *App) onConnect() { return } + p := a.currentProfile + serverURL := p.BuildServerURL() + sniHost := "" + serverIPs := p.GetServerIPList() + if len(serverIPs) > 0 { + sniHost = p.Host + } + // Build and send the start command. cfg := ipc.ClientConfig{ - ServerURL: a.currentProfile.ServerURL, - Username: a.currentProfile.Username, + ServerURL: serverURL, + SNIHost: sniHost, + ServerIPs: serverIPs, + Username: p.Username, Password: password, - AuthMode: string(a.currentProfile.AuthMode), - RoutingMode: string(a.currentProfile.RoutingMode), - CustomCIDRs: splitCIDRs(a.currentProfile.CustomCIDRs), - MTUOverride: a.currentProfile.MTUOverride, + AuthMode: string(p.AuthMode), + RoutingMode: string(p.RoutingMode), + CustomCIDRs: splitCIDRs(p.CustomCIDRs), + MTUOverride: p.MTUOverride, } if err := ipc.SendStart(client, cfg); err != nil { fyne.Do(func() { diff --git a/internal/vpn/session.go b/internal/vpn/session.go index 4654654..5bea185 100644 --- a/internal/vpn/session.go +++ b/internal/vpn/session.go @@ -28,6 +28,8 @@ import ( // SessionConfig describes how to connect to a VPN server. type SessionConfig struct { ServerURL string + SNIHost string // TLS SNI hostname for CDN + ServerIPs []string // CDN edge IPs for failover Username string Password string AuthMode model.AuthMode @@ -109,19 +111,29 @@ func (sm *SessionManager) Disconnect() { } } -// run is the main session loop with exponential-backoff reconnection. +// run is the main session loop with exponential-backoff reconnection +// and CDN IP failover. func (sm *SessionManager) run(ctx context.Context, cfg SessionConfig) { defer sm.setState(stats.StateDisconnected) backoff := time.Second maxBackoff := 60 * time.Second + // Build the full target list: original host first, then CDN IPs. + targets := append([]string{""}, cfg.ServerIPs...) // "" = use base URL + ipIndex := 0 + for { if ctx.Err() != nil { return } - err := sm.connectOnce(ctx, cfg) + targetIP := "" + if ipIndex > 0 && ipIndex < len(targets) { + targetIP = targets[ipIndex] + } + + err := sm.connectOnce(ctx, cfg, targetIP) if ctx.Err() != nil { sm.cleanup() return @@ -130,11 +142,20 @@ func (sm *SessionManager) run(ctx context.Context, cfg SessionConfig) { if err != nil { log.L().Error("VPN connection failed", "error", err) sm.setState(stats.StateReconnecting) + + // Try next CDN IP immediately. + ipIndex++ + if ipIndex < len(targets) { + log.L().Info("trying next CDN IP", "index", ipIndex, "ip", targets[ipIndex]) + continue + } + // All targets exhausted; reset and wait with backoff. + ipIndex = 0 } else { sm.setState(stats.StateReconnecting) + ipIndex = 0 } - // Wait before reconnecting, unless cancelled. select { case <-ctx.Done(): sm.cleanup() @@ -150,13 +171,20 @@ func (sm *SessionManager) run(ctx context.Context, cfg SessionConfig) { // connectOnce performs a single connection lifecycle: authenticate, // handshake, configure TUN, apply routes, pump packets until failure. -func (sm *SessionManager) connectOnce(ctx context.Context, cfg SessionConfig) error { +func (sm *SessionManager) connectOnce(ctx context.Context, cfg SessionConfig, targetIP string) error { sm.setState(stats.StateConnecting) + // Build URL for this attempt. If targetIP is set (CDN failover), + // build a URL with that IP. Otherwise use base ServerURL. + serverURL := cfg.ServerURL + if targetIP != "" { + serverURL = replaceHost(cfg.ServerURL, targetIP) + } + // Determine auth strategy and obtain JWT if needed. token := cfg.Token if token == "" && (cfg.AuthMode == model.AuthModeJWT || cfg.AuthMode == model.AuthModeBoth) { - httpBase, err := auth.WSURLToHTTP(cfg.ServerURL) + httpBase, err := wsURLToHTTP(serverURL) if err != nil { return fmt.Errorf("parse server URL: %w", err) } @@ -176,7 +204,8 @@ func (sm *SessionManager) connectOnce(ctx context.Context, cfg SessionConfig) er // Prepare the TUN + route setup callback (called during handshake, // between receiving init and sending ready). handshake := transport.HandshakeConfig{ - ServerURL: cfg.ServerURL, + ServerURL: serverURL, + SNIHost: cfg.SNIHost, Token: token, Username: cfg.Username, Password: cfg.Password, @@ -209,12 +238,6 @@ func (sm *SessionManager) connectOnce(ctx context.Context, cfg SessionConfig) er sm.conn = conn sm.mu.Unlock() - // If password auth was used (no token), the handshake already - // exchanged auth messages. For JWT, auth was implicit. - if token == "" { - // Password auth path already validated. - } - sm.stats.SetConnected(conn.Init().IP) sm.setState(stats.StateConnected) log.L().Info("VPN connected", @@ -431,6 +454,13 @@ func serverHostFromURL(wsURL string) string { break } } + // Strip port. + for i := 0; i < len(u); i++ { + if u[i] == ':' { + u = u[:i] + break + } + } // Strip path. for i := 0; i < len(u); i++ { if u[i] == '/' { @@ -440,3 +470,26 @@ func serverHostFromURL(wsURL string) string { } return u } + +// wsURLToHTTP converts a WebSocket URL to HTTP origin. +func wsURLToHTTP(wsURL string) (string, error) { + return auth.WSURLToHTTP(wsURL) +} + +// replaceHost substitutes the host portion of a URL string. +// e.g. wss://host:443/ws with 1.2.3.4 → wss://1.2.3.4:443/ws +func replaceHost(rawURL, newHost string) string { + u := rawURL + for _, prefix := range []string{"wss://", "ws://"} { + if len(u) > len(prefix) && u[:len(prefix)] == prefix { + rest := u[len(prefix):] + // Find end of host (either port or path). + end := 0 + for end < len(rest) && rest[end] != ':' && rest[end] != '/' { + end++ + } + return prefix + newHost + rest[end:] + } + } + return rawURL +}