mirror of
https://github.com/MHSanaei/3x-ui.git
synced 2026-08-30 06:57:14 +00:00
refactor: focused service files, leaf subpackages, and an internal/ layout (#5167)
* refactor(service): split client.go into focused files
client.go had grown to 4455 lines mixing ~10 responsibilities. Split it
verbatim into cohesive same-package files (no behavior change):
client.go foundation: ClientService, ClientWithAttachments,
ClientCreatePayload, ErrClientNotInInbound, sqlInChunk
client_locks.go inbound mutation locks, delete tombstones, compactOrphans
client_lookup.go read-only lookups (GetByID, List, EffectiveFlow, ...)
client_link.go inbound association sync (SyncInbound, DetachInbound, ...)
client_crud.go single-client CRUD + validation + protocol defaults
client_inbound_apply.go low-level inbound-settings mutators + by-email setters
client_bulk.go bulk attach/detach/adjust/delete/create + DelDepleted
client_traffic.go traffic-reset paths
client_groups.go client group management
client_paging.go paged listing, filtering, sorting, summary
Every declaration moved unchanged (verified: identical func/type/const/var
signature set before vs after). Imports redistributed per file via goimports.
go build ./..., go vet, and go test ./web/service/... all pass.
* refactor(service): split inbound.go into focused files
inbound.go was 4100 lines. Split it verbatim into cohesive same-package
files (no behavior change):
inbound.go core inbound CRUD + InboundService (keeps pkg doc)
inbound_protocol.go protocol / stream capability helpers
inbound_node.go node/runtime/remote coordination + online tracking
inbound_traffic.go traffic accounting, reset, client stats
inbound_client_ips.go per-client IP tracking
inbound_clients.go client lookups within inbounds + copy-clients
inbound_disable.go auto-disable invalid inbounds/clients
inbound_migration.go DB migrations
inbound_sublink.go subscription link providers
inbound_util.go generic slice/string helpers
Identical func/type/const/var signature set before vs after; package doc
comment preserved on inbound.go. Imports redistributed via goimports.
Build, vet, and go test ./web/service/... all pass.
* refactor(service): split tgbot.go into focused files
tgbot.go was 3738 lines dominated by a 1246-line answerCallback. Split it
verbatim into cohesive same-package files (no behavior change):
tgbot.go lifecycle, bot setup, caches, small utils
tgbot_router.go incoming update / command / callback dispatch
tgbot_send.go outbound messaging primitives
tgbot_client.go client views, actions, subscription links
tgbot_inbound.go inbound listing / pickers
tgbot_report.go server usage, exhausted, online, backups, notifications
Identical func/type/const/var signature set before vs after. Imports
redistributed via goimports. Build, vet, and go test ./web/service/... pass.
* refactor(client): dedupe single-field by-email setters
ResetClientIpLimitByEmail, ResetClientExpiryTimeByEmail, and
ResetClientTrafficLimitByEmail shared an identical ~50-line body that
resolves the inbound by email, confirms the client exists, rewrites a
single-client settings payload, and delegates to UpdateInboundClient.
Extract that into applyClientFieldByEmail(inboundSvc, email, mutate) and
reduce each setter to a 3-line wrapper. Behavior is unchanged: same checks
and error strings, same single-client payload contract, same totalGB guard.
SetClientTelegramUserID (resolves by traffic id, different error text) and
ToggleClientEnableByEmail/SetClientEnableByEmail (different return shape and
a pre-read of the old state) intentionally keep their own bodies.
* refactor(service): extract panel/ subpackage
Move the panel-administration leaf services out of the flat service
package into web/service/panel/ (package panel):
user.go UserService (auth / 2FA / LDAP)
panel.go PanelService (restart / self-update) + version helpers
panel_other.go non-unix RestartPanel
panel_unix.go unix RestartPanel
api_token.go ApiTokenService
websocket.go WebSocketService
panel_test.go version/shellQuote unit tests
These are leaves: they depend on core (SettingService, Release) but no
core file references them, so the extraction creates no import cycle.
Core references are now qualified (service.SettingService, service.Release);
callers in main.go, web/web.go, and web/controller/* updated to panel.*.
Build, vet, and go test ./web/... pass.
* refactor(service): extract integration/ subpackage
Move the external-provider integration leaves into web/service/integration/
(package integration):
warp.go WarpService (Cloudflare WARP)
nord.go NordService (NordVPN)
custom_geo.go CustomGeoService (custom geo asset management)
*_test.go custom_geo / panel-proxy tests
These depend on core (SettingService, ServerService, XraySettingService) but
no core file references them. xray_setting.go stays in core because it calls
the unexported SettingService.saveSetting. The shared isBlockedIP SSRF helper
(used by core url_safety.go and by custom_geo) now has a small copy in each
package rather than being exported. Core references qualified; callers in
web/web.go, web/job/*, and web/controller/* updated to integration.*.
Build, vet, and go test ./web/... pass.
* refactor(service): extract tgbot/ subpackage
Move the Telegram bot (6 files + test) into web/service/tgbot/ (package
tgbot). It is a leaf: it embeds five core services (Inbound/Client/Setting/
Server/Xray) and the core never references it, so no import cycle.
To support the package boundary without changing behavior:
- core exposes XrayProcess() *xray.Process so tgbot keeps calling the
exact same running-process methods it used via the package-level `p`;
- three core methods tgbot calls are exported: ClientService.checkIs-
EnabledByEmail -> CheckIsEnabledByEmail, InboundService.getAllEmails ->
GetAllEmails (callers updated in-package);
- tgbot's embedded-field types and the few core type refs (Status,
ClientCreatePayload, SanitizePublicHTTPURL) are now service-qualified.
Callers in main.go, web/web.go, web/job/*, and web/controller/* updated to
tgbot.*. Build, vet, and go test ./web/... pass.
* refactor(service): extract outbound/ subpackage
OutboundService (outbound.go) imports only neutral packages (config,
database, model, xray) and its production code is referenced by no core or
sibling service file — only by web/controller/xray_setting.go and
web/job/xray_traffic_job.go. Move it to web/service/outbound/ (package
outbound); no core qualification needed inside. Callers updated to outbound.*.
The one coupling was a tiny pure test helper, outboundsContainTag, used by
both outbound.go and the core outbound_subscription_test.go; it now has a
small copy in that test file rather than being shared across the boundary.
Build, vet, and go test ./web/... pass.
* refactor(util): move wireguard into its own subpackage
util/wireguard.go was the lone file of the root `util` package (24 lines,
one exported func GenerateWireguardKeypair), while every other util concern
lives in a focused subpackage (util/common, util/crypto, util/netsafe, ...).
Move it to util/wireguard/ (package wireguard) for consistency; its only
importer, web/service/integration/warp.go, is updated. The root `util`
package no longer exists.
* refactor(sub): drop redundant sub prefix from filenames
Inside package sub the subXxx.go prefix just repeats the package name
(like client_*.go did inside service). Rename for consistency; content and
type names are unchanged:
subController.go -> controller.go
subService.go -> service.go
subClashService.go -> clash_service.go
subJsonService.go -> json_service.go
(+ matching _test.go files)
* refactor(controller): rename xui.go -> spa.go
XUIController serves the panel's single-page-app shell; spa.go names that
role plainly (the other controller files are domain-named). File rename only
— the type stays XUIController. api_docs_test.go keys route base paths by
filename, so its "xui.go" case is updated to "spa.go".
* refactor: move backend packages under internal/
Adopt the idiomatic Go application layout: the backend packages now live
under internal/ (a boundary the toolchain enforces), signalling private
implementation instead of a library-style flat root. No runtime behavior
changes — only import paths and a few build/config paths move.
Moved: config, database, logger, mtproto, sub, util, web, xray -> internal/.
main.go stays at the repo root and tools/openapigen stays under tools/ (both
still import internal/* because the internal rule keys off the module root).
The module path github.com/mhsanaei/3x-ui/v3 is unchanged; 149 .go files had
their import prefix rewritten to .../internal/<pkg>.
Couplings the Go compiler can't see, updated to the new layout:
- frontend i18n imports of web/translation (react.ts, setup.components.ts)
- vite outDir + eslint/tsconfig ignore globs -> internal/web/dist
- Dockerfile COPY paths for web/dist and web/translation
- locale.go os.DirFS("web") disk fallback -> "internal/web"
- .gitignore and ci.yml go:embed stub for internal/web/dist
- api_docs_test.go repo-root relative walk (one level deeper)
- tools/openapigen filesystem package paths; ApiTokenView repointed to the
web/service/panel subpackage and codegen regenerated (clears a stale
type the ci.yml codegen check was failing on)
Verified: go build/vet/test (all packages), and frontend typecheck, lint,
vitest (478 tests), and production build into internal/web/dist.
* fix(config): keep test runs from writing logs into the source tree
GetLogFolder() returns a CWD-relative "./log" on Windows. Under `go test`
the working directory is each package's own folder, so InitLogger (called by
tests in web/job, web/service, xray, web/websocket) created stray log/
directories scattered through the source tree (e.g. internal/web/job/log/).
Redirect to a shared temp folder when testing.Testing() reports a test run.
Production behavior is unchanged: Windows still uses ./log next to the binary
and Linux /var/log/x-ui. The log files were always gitignored (*.log) and
never committed; this just stops the noise at the source.
* docs: move subscription-template guide out of root into docs/
sub_templates/ was a top-level folder holding only a README and no actual
templates (3x-ui ships none by design), referenced nowhere and unlinked from
any doc — it read like an empty placeholder cluttering the repo root.
Move the guide to docs/custom-subscription-templates.md (a proper docs home),
reword its intro to read as documentation rather than a folder note, link it
from the Features list in README.md, and drop the empty sub_templates/ folder.
* fix: update stale web/ path references after the internal/ move
The internal/ migration rewrote Go import paths but left some references to
the old top-level layout in docs, comments, and a few runtime disk paths.
Functional (dev-mode only): the disk-serving fallbacks that read the Vite
build from disk when running from source still pointed at web/dist/, which
moved to internal/web/dist/ — so `os.DirFS`/`os.Stat`/`os.ReadFile` in
internal/web/web.go and internal/sub/{sub,controller}.go are corrected.
Production was unaffected (it serves the embedded FS; verified by the Docker
build), but `go run` with a live frontend build silently fell back to embed.
Docs/comments: frontend/README.md, CONTRIBUTING.md, the claude-issue-bot and
release workflows, the openapigen -root help text, and assorted Go comments
now reference internal/web, internal/database, internal/sub, internal/xray,
etc. Package-name mentions (the "web" package), root paths (main.go,
frontend/, install scripts, /etc/x-ui), routes (/panel/api/xray), and the
historical "web/assets no longer exists" note were intentionally left as-is.
* refactor(web): remove the legacy /xui -> /panel redirect middleware
RedirectMiddleware existed only for backward compatibility with the old
`/xui` URL scheme (301-redirecting /xui and /xui/API to /panel and
/panel/api). That cutover was long ago, so drop the middleware, its
registration in initRouter, and the now-inaccurate "URL redirection"
mention in the middleware package doc. Old /xui URLs now 404 like any other
unknown path. HTTPS auto-redirect and auth redirects are unrelated and stay.
* build: fix .dockerignore for internal/ layout and exclude runtime dir
- web/dist -> internal/web/dist: the embedded frontend moved under internal/,
so the stale exclude no longer matched and the locally-built dist could be
sent to the build context (the frontend stage rebuilds it fresh anyway).
- exclude x-ui/: the local runtime directory (SQLite db, geo .dat files, xray
binaries, certs — ~150MB) was being shipped into the build context for no
reason. Verified the pattern excludes only the directory and still keeps
x-ui.sh, which the Dockerfile copies to /usr/bin/x-ui.
This commit is contained in:
@@ -0,0 +1,784 @@
|
||||
package integration
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/mhsanaei/3x-ui/v3/internal/config"
|
||||
"github.com/mhsanaei/3x-ui/v3/internal/database"
|
||||
"github.com/mhsanaei/3x-ui/v3/internal/database/model"
|
||||
"github.com/mhsanaei/3x-ui/v3/internal/logger"
|
||||
"github.com/mhsanaei/3x-ui/v3/internal/util/netproxy"
|
||||
"github.com/mhsanaei/3x-ui/v3/internal/util/netsafe"
|
||||
"github.com/mhsanaei/3x-ui/v3/internal/web/service"
|
||||
)
|
||||
|
||||
const (
|
||||
customGeoTypeGeosite = "geosite"
|
||||
customGeoTypeGeoip = "geoip"
|
||||
minDatBytes = 64
|
||||
customGeoProbeTimeout = 12 * time.Second
|
||||
)
|
||||
|
||||
var (
|
||||
customGeoAliasPattern = regexp.MustCompile(`^[a-z0-9_-]+$`)
|
||||
reservedCustomAliases = map[string]struct{}{
|
||||
"geoip": {}, "geosite": {},
|
||||
"geoip_ir": {}, "geosite_ir": {},
|
||||
"geoip_ru": {}, "geosite_ru": {},
|
||||
}
|
||||
ErrCustomGeoInvalidType = errors.New("custom_geo_invalid_type")
|
||||
ErrCustomGeoAliasRequired = errors.New("custom_geo_alias_required")
|
||||
ErrCustomGeoAliasPattern = errors.New("custom_geo_alias_pattern")
|
||||
ErrCustomGeoAliasReserved = errors.New("custom_geo_alias_reserved")
|
||||
ErrCustomGeoURLRequired = errors.New("custom_geo_url_required")
|
||||
ErrCustomGeoInvalidURL = errors.New("custom_geo_invalid_url")
|
||||
ErrCustomGeoURLScheme = errors.New("custom_geo_url_scheme")
|
||||
ErrCustomGeoURLHost = errors.New("custom_geo_url_host")
|
||||
ErrCustomGeoDuplicateAlias = errors.New("custom_geo_duplicate_alias")
|
||||
ErrCustomGeoNotFound = errors.New("custom_geo_not_found")
|
||||
ErrCustomGeoDownload = errors.New("custom_geo_download")
|
||||
ErrCustomGeoSSRFBlocked = errors.New("custom_geo_ssrf_blocked")
|
||||
ErrCustomGeoPathTraversal = errors.New("custom_geo_path_traversal")
|
||||
)
|
||||
|
||||
type CustomGeoUpdateAllItem struct {
|
||||
Id int `json:"id"`
|
||||
Alias string `json:"alias"`
|
||||
FileName string `json:"fileName"`
|
||||
}
|
||||
|
||||
type CustomGeoUpdateAllFailure struct {
|
||||
Id int `json:"id"`
|
||||
Alias string `json:"alias"`
|
||||
FileName string `json:"fileName"`
|
||||
Err string `json:"error"`
|
||||
}
|
||||
|
||||
type CustomGeoUpdateAllResult struct {
|
||||
Succeeded []CustomGeoUpdateAllItem `json:"succeeded"`
|
||||
Failed []CustomGeoUpdateAllFailure `json:"failed"`
|
||||
}
|
||||
|
||||
type CustomGeoService struct {
|
||||
serverService *service.ServerService
|
||||
updateAllGetAll func() ([]model.CustomGeoResource, error)
|
||||
updateAllApply func(id int, onStartup bool) (string, error)
|
||||
updateAllRestart func() error
|
||||
getPanelProxy func() (string, error)
|
||||
}
|
||||
|
||||
func NewCustomGeoService() *CustomGeoService {
|
||||
s := &CustomGeoService{
|
||||
serverService: &service.ServerService{},
|
||||
}
|
||||
s.updateAllGetAll = s.GetAll
|
||||
s.updateAllApply = s.applyDownloadAndPersist
|
||||
s.updateAllRestart = func() error { return s.serverService.RestartXrayService() }
|
||||
s.getPanelProxy = (&service.SettingService{}).GetPanelProxy
|
||||
return s
|
||||
}
|
||||
|
||||
func NormalizeAliasKey(alias string) string {
|
||||
return strings.ToLower(strings.ReplaceAll(alias, "-", "_"))
|
||||
}
|
||||
|
||||
func (s *CustomGeoService) fileNameFor(typ, alias string) string {
|
||||
if typ == customGeoTypeGeoip {
|
||||
return fmt.Sprintf("geoip_%s.dat", alias)
|
||||
}
|
||||
return fmt.Sprintf("geosite_%s.dat", alias)
|
||||
}
|
||||
|
||||
func (s *CustomGeoService) validateType(typ string) error {
|
||||
if typ != customGeoTypeGeosite && typ != customGeoTypeGeoip {
|
||||
return ErrCustomGeoInvalidType
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *CustomGeoService) validateAlias(alias string) error {
|
||||
if alias == "" {
|
||||
return ErrCustomGeoAliasRequired
|
||||
}
|
||||
if !customGeoAliasPattern.MatchString(alias) {
|
||||
return ErrCustomGeoAliasPattern
|
||||
}
|
||||
if _, ok := reservedCustomAliases[NormalizeAliasKey(alias)]; ok {
|
||||
return ErrCustomGeoAliasReserved
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *CustomGeoService) sanitizeURL(raw string) (string, error) {
|
||||
if raw == "" {
|
||||
return "", ErrCustomGeoURLRequired
|
||||
}
|
||||
u, err := url.Parse(raw)
|
||||
if err != nil {
|
||||
return "", ErrCustomGeoInvalidURL
|
||||
}
|
||||
if u.Scheme != "http" && u.Scheme != "https" {
|
||||
return "", ErrCustomGeoURLScheme
|
||||
}
|
||||
if u.Host == "" {
|
||||
return "", ErrCustomGeoURLHost
|
||||
}
|
||||
if err := checkSSRF(context.Background(), u.Hostname()); err != nil {
|
||||
return "", err
|
||||
}
|
||||
// Reconstruct URL from parsed components to break taint propagation.
|
||||
clean := &url.URL{
|
||||
Scheme: u.Scheme,
|
||||
Host: u.Host,
|
||||
Path: u.Path,
|
||||
RawPath: u.RawPath,
|
||||
RawQuery: u.RawQuery,
|
||||
Fragment: u.Fragment,
|
||||
}
|
||||
return clean.String(), nil
|
||||
}
|
||||
|
||||
func localDatFileNeedsRepair(path string) bool {
|
||||
safePath, err := sanitizeDestPath(path)
|
||||
if err != nil {
|
||||
return true
|
||||
}
|
||||
fi, err := os.Stat(safePath)
|
||||
if err != nil {
|
||||
return true
|
||||
}
|
||||
if fi.IsDir() {
|
||||
return true
|
||||
}
|
||||
return fi.Size() < int64(minDatBytes)
|
||||
}
|
||||
|
||||
func CustomGeoLocalFileNeedsRepair(path string) bool {
|
||||
return localDatFileNeedsRepair(path)
|
||||
}
|
||||
|
||||
func isBlockedIP(ip net.IP) bool {
|
||||
return netsafe.IsBlockedIP(ip)
|
||||
}
|
||||
|
||||
// checkSSRFDefault validates that the given host does not resolve to a private/internal IP.
|
||||
// It is context-aware so that dial context cancellation/deadlines are respected during DNS resolution.
|
||||
func checkSSRFDefault(ctx context.Context, hostname string) error {
|
||||
ips, err := net.DefaultResolver.LookupIPAddr(ctx, hostname)
|
||||
if err != nil {
|
||||
return fmt.Errorf("%w: cannot resolve host %s", ErrCustomGeoSSRFBlocked, hostname)
|
||||
}
|
||||
for _, ipAddr := range ips {
|
||||
if isBlockedIP(ipAddr.IP) {
|
||||
return fmt.Errorf("%w: %s resolves to blocked address %s", ErrCustomGeoSSRFBlocked, hostname, ipAddr.IP)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// checkSSRF is the active SSRF guard. Override in tests to allow localhost test servers.
|
||||
var checkSSRF = checkSSRFDefault
|
||||
|
||||
func ssrfSafeTransport() http.RoundTripper {
|
||||
base, ok := http.DefaultTransport.(*http.Transport)
|
||||
if !ok {
|
||||
base = &http.Transport{}
|
||||
}
|
||||
cloned := base.Clone()
|
||||
cloned.DialContext = func(ctx context.Context, network, addr string) (net.Conn, error) {
|
||||
host, _, err := net.SplitHostPort(addr)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: %v", ErrCustomGeoSSRFBlocked, err)
|
||||
}
|
||||
if err := checkSSRF(ctx, host); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var dialer net.Dialer
|
||||
return dialer.DialContext(ctx, network, addr)
|
||||
}
|
||||
return cloned
|
||||
}
|
||||
|
||||
func (s *CustomGeoService) httpClient(timeout time.Duration) *http.Client {
|
||||
proxyURL := ""
|
||||
if s.getPanelProxy != nil {
|
||||
if p, err := s.getPanelProxy(); err != nil {
|
||||
logger.Warning("custom geo: read panel proxy:", err)
|
||||
} else {
|
||||
proxyURL = strings.TrimSpace(p)
|
||||
}
|
||||
}
|
||||
if proxyURL != "" {
|
||||
client, err := netproxy.NewHTTPClient(proxyURL, timeout)
|
||||
if err != nil {
|
||||
logger.Warningf("custom geo: invalid panel proxy %q, using direct connection: %v", proxyURL, err)
|
||||
} else {
|
||||
return client
|
||||
}
|
||||
}
|
||||
return &http.Client{Timeout: timeout, Transport: ssrfSafeTransport()}
|
||||
}
|
||||
|
||||
func (s *CustomGeoService) probeCustomGeoURLWithGET(rawURL string) error {
|
||||
sanitizedURL, err := s.sanitizeURL(rawURL)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
client := s.httpClient(customGeoProbeTimeout)
|
||||
req, err := http.NewRequest(http.MethodGet, sanitizedURL, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
req.Header.Set("Range", "bytes=0-0")
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
_, _ = io.Copy(io.Discard, io.LimitReader(resp.Body, 256))
|
||||
switch resp.StatusCode {
|
||||
case http.StatusOK, http.StatusPartialContent:
|
||||
return nil
|
||||
default:
|
||||
return fmt.Errorf("get range status %d", resp.StatusCode)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *CustomGeoService) probeCustomGeoURL(rawURL string) error {
|
||||
sanitizedURL, err := s.sanitizeURL(rawURL)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
client := s.httpClient(customGeoProbeTimeout)
|
||||
req, err := http.NewRequest(http.MethodHead, sanitizedURL, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
_ = resp.Body.Close()
|
||||
sc := resp.StatusCode
|
||||
if sc >= 200 && sc < 300 {
|
||||
return nil
|
||||
}
|
||||
if sc == http.StatusMethodNotAllowed || sc == http.StatusNotImplemented {
|
||||
return s.probeCustomGeoURLWithGET(rawURL)
|
||||
}
|
||||
return fmt.Errorf("head status %d", sc)
|
||||
}
|
||||
|
||||
func (s *CustomGeoService) EnsureOnStartup() {
|
||||
list, err := s.GetAll()
|
||||
if err != nil {
|
||||
logger.Warning("custom geo startup: load list:", err)
|
||||
return
|
||||
}
|
||||
n := len(list)
|
||||
if n == 0 {
|
||||
logger.Info("custom geo startup: no custom geofiles configured")
|
||||
return
|
||||
}
|
||||
logger.Infof("custom geo startup: checking %d custom geofile(s)", n)
|
||||
for i := range list {
|
||||
r := &list[i]
|
||||
sanitizedURL, err := s.sanitizeURL(r.Url)
|
||||
if err != nil {
|
||||
logger.Warningf("custom geo startup id=%d: invalid url: %v", r.Id, err)
|
||||
continue
|
||||
}
|
||||
r.Url = sanitizedURL
|
||||
s.syncLocalPath(r)
|
||||
localPath := r.LocalPath
|
||||
if !localDatFileNeedsRepair(localPath) {
|
||||
logger.Infof("custom geo startup id=%d alias=%s path=%s: present", r.Id, r.Alias, localPath)
|
||||
continue
|
||||
}
|
||||
logger.Infof("custom geo startup id=%d alias=%s path=%s: missing or needs repair, probing source", r.Id, r.Alias, localPath)
|
||||
if err := s.probeCustomGeoURL(r.Url); err != nil {
|
||||
logger.Warningf("custom geo startup id=%d alias=%s url=%s: probe: %v (attempting download anyway)", r.Id, r.Alias, r.Url, err)
|
||||
}
|
||||
_, _ = s.applyDownloadAndPersist(r.Id, true)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *CustomGeoService) downloadToPath(resourceURL, destPath string, lastModifiedHeader string) (skipped bool, newLastModified string, err error) {
|
||||
safeDestPath, err := sanitizeDestPath(destPath)
|
||||
if err != nil {
|
||||
return false, "", fmt.Errorf("%w: %v", ErrCustomGeoDownload, err)
|
||||
}
|
||||
|
||||
skipped, lm, err := s.downloadToPathOnce(resourceURL, safeDestPath, lastModifiedHeader, false)
|
||||
if err != nil {
|
||||
return false, "", err
|
||||
}
|
||||
if skipped {
|
||||
if _, statErr := os.Stat(safeDestPath); statErr == nil && !localDatFileNeedsRepair(safeDestPath) {
|
||||
return true, lm, nil
|
||||
}
|
||||
return s.downloadToPathOnce(resourceURL, safeDestPath, lastModifiedHeader, true)
|
||||
}
|
||||
return false, lm, nil
|
||||
}
|
||||
|
||||
// sanitizeDestPath ensures destPath is inside the bin folder, preventing path traversal.
|
||||
// It resolves symlinks to prevent symlink-based escapes.
|
||||
// Returns the cleaned absolute path that is safe to use in file operations.
|
||||
func sanitizeDestPath(destPath string) (string, error) {
|
||||
baseDirAbs, err := filepath.Abs(config.GetBinFolderPath())
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("%w: %v", ErrCustomGeoPathTraversal, err)
|
||||
}
|
||||
// Resolve symlinks in base directory to get the real path.
|
||||
if resolved, evalErr := filepath.EvalSymlinks(baseDirAbs); evalErr == nil {
|
||||
baseDirAbs = resolved
|
||||
}
|
||||
destPathAbs, err := filepath.Abs(destPath)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("%w: %v", ErrCustomGeoPathTraversal, err)
|
||||
}
|
||||
// Resolve symlinks for the parent directory of the destination path.
|
||||
destDir := filepath.Dir(destPathAbs)
|
||||
if resolved, evalErr := filepath.EvalSymlinks(destDir); evalErr == nil {
|
||||
destPathAbs = filepath.Join(resolved, filepath.Base(destPathAbs))
|
||||
}
|
||||
// Verify the resolved path is within the safe base directory using prefix check.
|
||||
safeDirPrefix := baseDirAbs + string(filepath.Separator)
|
||||
if !strings.HasPrefix(destPathAbs, safeDirPrefix) {
|
||||
return "", ErrCustomGeoPathTraversal
|
||||
}
|
||||
return destPathAbs, nil
|
||||
}
|
||||
|
||||
func (s *CustomGeoService) downloadToPathOnce(resourceURL, destPath string, lastModifiedHeader string, forceFull bool) (skipped bool, newLastModified string, err error) {
|
||||
safeDestPath, err := sanitizeDestPath(destPath)
|
||||
if err != nil {
|
||||
return false, "", fmt.Errorf("%w: %v", ErrCustomGeoDownload, err)
|
||||
}
|
||||
sanitizedURL, err := s.sanitizeURL(resourceURL)
|
||||
if err != nil {
|
||||
return false, "", fmt.Errorf("%w: %v", ErrCustomGeoDownload, err)
|
||||
}
|
||||
|
||||
var req *http.Request
|
||||
req, err = http.NewRequest(http.MethodGet, sanitizedURL, nil)
|
||||
if err != nil {
|
||||
return false, "", fmt.Errorf("%w: %v", ErrCustomGeoDownload, err)
|
||||
}
|
||||
|
||||
if !forceFull {
|
||||
if fi, statErr := os.Stat(safeDestPath); statErr == nil && !localDatFileNeedsRepair(safeDestPath) {
|
||||
if !fi.ModTime().IsZero() {
|
||||
req.Header.Set("If-Modified-Since", fi.ModTime().UTC().Format(http.TimeFormat))
|
||||
} else if lastModifiedHeader != "" {
|
||||
if t, perr := time.Parse(http.TimeFormat, lastModifiedHeader); perr == nil {
|
||||
req.Header.Set("If-Modified-Since", t.UTC().Format(http.TimeFormat))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
client := s.httpClient(10 * time.Minute)
|
||||
// lgtm[go/request-forgery]
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return false, "", fmt.Errorf("%w: %v", ErrCustomGeoDownload, err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
var serverModTime time.Time
|
||||
if lm := resp.Header.Get("Last-Modified"); lm != "" {
|
||||
if parsed, perr := time.Parse(http.TimeFormat, lm); perr == nil {
|
||||
serverModTime = parsed
|
||||
newLastModified = lm
|
||||
}
|
||||
}
|
||||
|
||||
updateModTime := func() {
|
||||
if !serverModTime.IsZero() {
|
||||
_ = os.Chtimes(safeDestPath, serverModTime, serverModTime)
|
||||
}
|
||||
}
|
||||
|
||||
if resp.StatusCode == http.StatusNotModified {
|
||||
if forceFull {
|
||||
return false, "", fmt.Errorf("%w: unexpected 304 on unconditional get", ErrCustomGeoDownload)
|
||||
}
|
||||
updateModTime()
|
||||
return true, newLastModified, nil
|
||||
}
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return false, "", fmt.Errorf("%w: unexpected status %d", ErrCustomGeoDownload, resp.StatusCode)
|
||||
}
|
||||
|
||||
binDir := filepath.Dir(safeDestPath)
|
||||
if err = os.MkdirAll(binDir, 0o755); err != nil {
|
||||
return false, "", fmt.Errorf("%w: %v", ErrCustomGeoDownload, err)
|
||||
}
|
||||
|
||||
safeTmpPath, err := sanitizeDestPath(safeDestPath + ".tmp")
|
||||
if err != nil {
|
||||
return false, "", fmt.Errorf("%w: %v", ErrCustomGeoDownload, err)
|
||||
}
|
||||
out, err := os.Create(safeTmpPath)
|
||||
if err != nil {
|
||||
return false, "", fmt.Errorf("%w: %v", ErrCustomGeoDownload, err)
|
||||
}
|
||||
n, err := io.Copy(out, resp.Body)
|
||||
closeErr := out.Close()
|
||||
if err != nil {
|
||||
_ = os.Remove(safeTmpPath)
|
||||
return false, "", fmt.Errorf("%w: %v", ErrCustomGeoDownload, err)
|
||||
}
|
||||
if closeErr != nil {
|
||||
_ = os.Remove(safeTmpPath)
|
||||
return false, "", fmt.Errorf("%w: %v", ErrCustomGeoDownload, closeErr)
|
||||
}
|
||||
if n < minDatBytes {
|
||||
_ = os.Remove(safeTmpPath)
|
||||
return false, "", fmt.Errorf("%w: file too small", ErrCustomGeoDownload)
|
||||
}
|
||||
|
||||
if err = os.Rename(safeTmpPath, safeDestPath); err != nil {
|
||||
_ = os.Remove(safeTmpPath)
|
||||
return false, "", fmt.Errorf("%w: %v", ErrCustomGeoDownload, err)
|
||||
}
|
||||
|
||||
updateModTime()
|
||||
if newLastModified == "" && resp.Header.Get("Last-Modified") != "" {
|
||||
newLastModified = resp.Header.Get("Last-Modified")
|
||||
}
|
||||
return false, newLastModified, nil
|
||||
}
|
||||
|
||||
func (s *CustomGeoService) resolveDestPath(r *model.CustomGeoResource) string {
|
||||
if r.LocalPath != "" {
|
||||
return r.LocalPath
|
||||
}
|
||||
return filepath.Join(config.GetBinFolderPath(), s.fileNameFor(r.Type, r.Alias))
|
||||
}
|
||||
|
||||
func (s *CustomGeoService) syncLocalPath(r *model.CustomGeoResource) {
|
||||
p := filepath.Join(config.GetBinFolderPath(), s.fileNameFor(r.Type, r.Alias))
|
||||
r.LocalPath = p
|
||||
}
|
||||
|
||||
func (s *CustomGeoService) syncAndSanitizeLocalPath(r *model.CustomGeoResource) error {
|
||||
s.syncLocalPath(r)
|
||||
safePath, err := sanitizeDestPath(r.LocalPath)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
r.LocalPath = safePath
|
||||
return nil
|
||||
}
|
||||
|
||||
func removeSafePathIfExists(path string) error {
|
||||
safePath, err := sanitizeDestPath(path)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err := os.Stat(safePath); err == nil {
|
||||
if err := os.Remove(safePath); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *CustomGeoService) Create(r *model.CustomGeoResource) error {
|
||||
if err := s.validateType(r.Type); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := s.validateAlias(r.Alias); err != nil {
|
||||
return err
|
||||
}
|
||||
sanitizedURL, err := s.sanitizeURL(r.Url)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
r.Url = sanitizedURL
|
||||
var existing int64
|
||||
database.GetDB().Model(&model.CustomGeoResource{}).
|
||||
Where("geo_type = ? AND alias = ?", r.Type, r.Alias).Count(&existing)
|
||||
if existing > 0 {
|
||||
return ErrCustomGeoDuplicateAlias
|
||||
}
|
||||
if err := s.syncAndSanitizeLocalPath(r); err != nil {
|
||||
return err
|
||||
}
|
||||
skipped, lm, err := s.downloadToPath(r.Url, r.LocalPath, r.LastModified)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
now := time.Now().Unix()
|
||||
r.LastUpdatedAt = now
|
||||
r.LastModified = lm
|
||||
if err = database.GetDB().Create(r).Error; err != nil {
|
||||
_ = removeSafePathIfExists(r.LocalPath)
|
||||
return err
|
||||
}
|
||||
logger.Infof("custom geo created id=%d type=%s alias=%s skipped=%v", r.Id, r.Type, r.Alias, skipped)
|
||||
if err = s.serverService.RestartXrayService(); err != nil {
|
||||
logger.Warning("custom geo create: restart xray:", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *CustomGeoService) Update(id int, r *model.CustomGeoResource) error {
|
||||
var cur model.CustomGeoResource
|
||||
if err := database.GetDB().First(&cur, id).Error; err != nil {
|
||||
if database.IsNotFound(err) {
|
||||
return ErrCustomGeoNotFound
|
||||
}
|
||||
return err
|
||||
}
|
||||
if err := s.validateType(r.Type); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := s.validateAlias(r.Alias); err != nil {
|
||||
return err
|
||||
}
|
||||
sanitizedURL, err := s.sanitizeURL(r.Url)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
r.Url = sanitizedURL
|
||||
if cur.Type != r.Type || cur.Alias != r.Alias {
|
||||
var cnt int64
|
||||
database.GetDB().Model(&model.CustomGeoResource{}).
|
||||
Where("geo_type = ? AND alias = ? AND id <> ?", r.Type, r.Alias, id).
|
||||
Count(&cnt)
|
||||
if cnt > 0 {
|
||||
return ErrCustomGeoDuplicateAlias
|
||||
}
|
||||
}
|
||||
oldPath := s.resolveDestPath(&cur)
|
||||
r.Id = id
|
||||
if err := s.syncAndSanitizeLocalPath(r); err != nil {
|
||||
return err
|
||||
}
|
||||
if oldPath != r.LocalPath && oldPath != "" {
|
||||
if err := removeSafePathIfExists(oldPath); err != nil && !errors.Is(err, ErrCustomGeoPathTraversal) {
|
||||
logger.Warningf("custom geo remove old path %s: %v", oldPath, err)
|
||||
}
|
||||
}
|
||||
_, lm, err := s.downloadToPath(r.Url, r.LocalPath, cur.LastModified)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
r.LastUpdatedAt = time.Now().Unix()
|
||||
r.LastModified = lm
|
||||
err = database.GetDB().Model(&model.CustomGeoResource{}).Where("id = ?", id).Updates(map[string]any{
|
||||
"geo_type": r.Type,
|
||||
"alias": r.Alias,
|
||||
"url": r.Url,
|
||||
"local_path": r.LocalPath,
|
||||
"last_updated_at": r.LastUpdatedAt,
|
||||
"last_modified": r.LastModified,
|
||||
}).Error
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
logger.Infof("custom geo updated id=%d", id)
|
||||
if err = s.serverService.RestartXrayService(); err != nil {
|
||||
logger.Warning("custom geo update: restart xray:", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *CustomGeoService) Delete(id int) (displayName string, err error) {
|
||||
var r model.CustomGeoResource
|
||||
if err := database.GetDB().First(&r, id).Error; err != nil {
|
||||
if database.IsNotFound(err) {
|
||||
return "", ErrCustomGeoNotFound
|
||||
}
|
||||
return "", err
|
||||
}
|
||||
displayName = s.fileNameFor(r.Type, r.Alias)
|
||||
p := s.resolveDestPath(&r)
|
||||
if _, err := sanitizeDestPath(p); err != nil {
|
||||
return displayName, err
|
||||
}
|
||||
if err := database.GetDB().Delete(&model.CustomGeoResource{}, id).Error; err != nil {
|
||||
return displayName, err
|
||||
}
|
||||
if p != "" {
|
||||
if err := removeSafePathIfExists(p); err != nil {
|
||||
logger.Warningf("custom geo delete file %s: %v", p, err)
|
||||
}
|
||||
}
|
||||
logger.Infof("custom geo deleted id=%d", id)
|
||||
if err := s.serverService.RestartXrayService(); err != nil {
|
||||
logger.Warning("custom geo delete: restart xray:", err)
|
||||
}
|
||||
return displayName, nil
|
||||
}
|
||||
|
||||
func (s *CustomGeoService) GetAll() ([]model.CustomGeoResource, error) {
|
||||
var list []model.CustomGeoResource
|
||||
err := database.GetDB().Order("id asc").Find(&list).Error
|
||||
return list, err
|
||||
}
|
||||
|
||||
func (s *CustomGeoService) applyDownloadAndPersist(id int, onStartup bool) (displayName string, err error) {
|
||||
var r model.CustomGeoResource
|
||||
if err := database.GetDB().First(&r, id).Error; err != nil {
|
||||
if database.IsNotFound(err) {
|
||||
return "", ErrCustomGeoNotFound
|
||||
}
|
||||
return "", err
|
||||
}
|
||||
displayName = s.fileNameFor(r.Type, r.Alias)
|
||||
if err := s.syncAndSanitizeLocalPath(&r); err != nil {
|
||||
return displayName, err
|
||||
}
|
||||
sanitizedURL, sanitizeErr := s.sanitizeURL(r.Url)
|
||||
if sanitizeErr != nil {
|
||||
return displayName, sanitizeErr
|
||||
}
|
||||
skipped, lm, err := s.downloadToPath(sanitizedURL, r.LocalPath, r.LastModified)
|
||||
if err != nil {
|
||||
if onStartup {
|
||||
logger.Warningf("custom geo startup download id=%d: %v", id, err)
|
||||
} else {
|
||||
logger.Warningf("custom geo manual update id=%d: %v", id, err)
|
||||
}
|
||||
return displayName, err
|
||||
}
|
||||
now := time.Now().Unix()
|
||||
updates := map[string]any{
|
||||
"last_modified": lm,
|
||||
"local_path": r.LocalPath,
|
||||
"last_updated_at": now,
|
||||
}
|
||||
if err = database.GetDB().Model(&model.CustomGeoResource{}).Where("id = ?", id).Updates(updates).Error; err != nil {
|
||||
if onStartup {
|
||||
logger.Warningf("custom geo startup id=%d: persist metadata: %v", id, err)
|
||||
} else {
|
||||
logger.Warningf("custom geo manual update id=%d: persist metadata: %v", id, err)
|
||||
}
|
||||
return displayName, err
|
||||
}
|
||||
if skipped {
|
||||
if onStartup {
|
||||
logger.Infof("custom geo startup download skipped (not modified) id=%d", id)
|
||||
} else {
|
||||
logger.Infof("custom geo manual update skipped (not modified) id=%d", id)
|
||||
}
|
||||
} else {
|
||||
if onStartup {
|
||||
logger.Infof("custom geo startup download ok id=%d", id)
|
||||
} else {
|
||||
logger.Infof("custom geo manual update ok id=%d", id)
|
||||
}
|
||||
}
|
||||
return displayName, nil
|
||||
}
|
||||
|
||||
func (s *CustomGeoService) TriggerUpdate(id int) (string, error) {
|
||||
displayName, err := s.applyDownloadAndPersist(id, false)
|
||||
if err != nil {
|
||||
return displayName, err
|
||||
}
|
||||
if err = s.serverService.RestartXrayService(); err != nil {
|
||||
logger.Warning("custom geo manual update: restart xray:", err)
|
||||
}
|
||||
return displayName, nil
|
||||
}
|
||||
|
||||
func (s *CustomGeoService) TriggerUpdateAll() (*CustomGeoUpdateAllResult, error) {
|
||||
var list []model.CustomGeoResource
|
||||
var err error
|
||||
if s.updateAllGetAll != nil {
|
||||
list, err = s.updateAllGetAll()
|
||||
} else {
|
||||
list, err = s.GetAll()
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
res := &CustomGeoUpdateAllResult{}
|
||||
if len(list) == 0 {
|
||||
return res, nil
|
||||
}
|
||||
for _, r := range list {
|
||||
var name string
|
||||
var applyErr error
|
||||
if s.updateAllApply != nil {
|
||||
name, applyErr = s.updateAllApply(r.Id, false)
|
||||
} else {
|
||||
name, applyErr = s.applyDownloadAndPersist(r.Id, false)
|
||||
}
|
||||
if applyErr != nil {
|
||||
res.Failed = append(res.Failed, CustomGeoUpdateAllFailure{
|
||||
Id: r.Id, Alias: r.Alias, FileName: name, Err: applyErr.Error(),
|
||||
})
|
||||
continue
|
||||
}
|
||||
res.Succeeded = append(res.Succeeded, CustomGeoUpdateAllItem{
|
||||
Id: r.Id, Alias: r.Alias, FileName: name,
|
||||
})
|
||||
}
|
||||
if len(res.Succeeded) > 0 {
|
||||
var restartErr error
|
||||
if s.updateAllRestart != nil {
|
||||
restartErr = s.updateAllRestart()
|
||||
} else {
|
||||
restartErr = s.serverService.RestartXrayService()
|
||||
}
|
||||
if restartErr != nil {
|
||||
logger.Warning("custom geo update all: restart xray:", restartErr)
|
||||
}
|
||||
}
|
||||
return res, nil
|
||||
}
|
||||
|
||||
type CustomGeoAliasItem struct {
|
||||
Alias string `json:"alias"`
|
||||
Type string `json:"type"`
|
||||
FileName string `json:"fileName"`
|
||||
ExtExample string `json:"extExample"`
|
||||
}
|
||||
|
||||
type CustomGeoAliasesResponse struct {
|
||||
Geosite []CustomGeoAliasItem `json:"geosite"`
|
||||
Geoip []CustomGeoAliasItem `json:"geoip"`
|
||||
}
|
||||
|
||||
func (s *CustomGeoService) GetAliasesForUI() (CustomGeoAliasesResponse, error) {
|
||||
list, err := s.GetAll()
|
||||
if err != nil {
|
||||
logger.Warning("custom geo GetAliasesForUI:", err)
|
||||
return CustomGeoAliasesResponse{}, err
|
||||
}
|
||||
var out CustomGeoAliasesResponse
|
||||
for _, r := range list {
|
||||
fn := s.fileNameFor(r.Type, r.Alias)
|
||||
ex := fmt.Sprintf("ext:%s:tag", fn)
|
||||
item := CustomGeoAliasItem{
|
||||
Alias: r.Alias,
|
||||
Type: r.Type,
|
||||
FileName: fn,
|
||||
ExtExample: ex,
|
||||
}
|
||||
if r.Type == customGeoTypeGeoip {
|
||||
out.Geoip = append(out.Geoip, item)
|
||||
} else {
|
||||
out.Geosite = append(out.Geosite, item)
|
||||
}
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
@@ -0,0 +1,348 @@
|
||||
package integration
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"github.com/mhsanaei/3x-ui/v3/internal/database/model"
|
||||
)
|
||||
|
||||
// disableSSRFCheck disables the SSRF guard for the duration of a test,
|
||||
// allowing httptest servers on localhost. It restores the original on cleanup.
|
||||
func disableSSRFCheck(t *testing.T) {
|
||||
t.Helper()
|
||||
orig := checkSSRF
|
||||
checkSSRF = func(_ context.Context, _ string) error { return nil }
|
||||
t.Cleanup(func() { checkSSRF = orig })
|
||||
}
|
||||
|
||||
func TestNormalizeAliasKey(t *testing.T) {
|
||||
if got := NormalizeAliasKey("GeoIP-IR"); got != "geoip_ir" {
|
||||
t.Fatalf("got %q", got)
|
||||
}
|
||||
if got := NormalizeAliasKey("a-b_c"); got != "a_b_c" {
|
||||
t.Fatalf("got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNewCustomGeoService(t *testing.T) {
|
||||
s := NewCustomGeoService()
|
||||
if err := s.validateAlias("ok_alias-1"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTriggerUpdateAllAllSuccess(t *testing.T) {
|
||||
s := CustomGeoService{}
|
||||
s.updateAllGetAll = func() ([]model.CustomGeoResource, error) {
|
||||
return []model.CustomGeoResource{
|
||||
{Id: 1, Alias: "a"},
|
||||
{Id: 2, Alias: "b"},
|
||||
}, nil
|
||||
}
|
||||
s.updateAllApply = func(id int, onStartup bool) (string, error) {
|
||||
return fmt.Sprintf("geo_%d.dat", id), nil
|
||||
}
|
||||
restartCalls := 0
|
||||
s.updateAllRestart = func() error {
|
||||
restartCalls++
|
||||
return nil
|
||||
}
|
||||
|
||||
res, err := s.TriggerUpdateAll()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(res.Succeeded) != 2 || len(res.Failed) != 0 {
|
||||
t.Fatalf("unexpected result: %+v", res)
|
||||
}
|
||||
if restartCalls != 1 {
|
||||
t.Fatalf("expected 1 restart, got %d", restartCalls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTriggerUpdateAllPartialSuccess(t *testing.T) {
|
||||
s := CustomGeoService{}
|
||||
s.updateAllGetAll = func() ([]model.CustomGeoResource, error) {
|
||||
return []model.CustomGeoResource{
|
||||
{Id: 1, Alias: "ok"},
|
||||
{Id: 2, Alias: "bad"},
|
||||
}, nil
|
||||
}
|
||||
s.updateAllApply = func(id int, onStartup bool) (string, error) {
|
||||
if id == 2 {
|
||||
return "geo_2.dat", ErrCustomGeoDownload
|
||||
}
|
||||
return "geo_1.dat", nil
|
||||
}
|
||||
restartCalls := 0
|
||||
s.updateAllRestart = func() error {
|
||||
restartCalls++
|
||||
return nil
|
||||
}
|
||||
|
||||
res, err := s.TriggerUpdateAll()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(res.Succeeded) != 1 || len(res.Failed) != 1 {
|
||||
t.Fatalf("unexpected result: %+v", res)
|
||||
}
|
||||
if restartCalls != 1 {
|
||||
t.Fatalf("expected 1 restart, got %d", restartCalls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTriggerUpdateAllAllFailure(t *testing.T) {
|
||||
s := CustomGeoService{}
|
||||
s.updateAllGetAll = func() ([]model.CustomGeoResource, error) {
|
||||
return []model.CustomGeoResource{
|
||||
{Id: 1, Alias: "a"},
|
||||
{Id: 2, Alias: "b"},
|
||||
}, nil
|
||||
}
|
||||
s.updateAllApply = func(id int, onStartup bool) (string, error) {
|
||||
return fmt.Sprintf("geo_%d.dat", id), ErrCustomGeoDownload
|
||||
}
|
||||
restartCalls := 0
|
||||
s.updateAllRestart = func() error {
|
||||
restartCalls++
|
||||
return nil
|
||||
}
|
||||
|
||||
res, err := s.TriggerUpdateAll()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if len(res.Succeeded) != 0 || len(res.Failed) != 2 {
|
||||
t.Fatalf("unexpected result: %+v", res)
|
||||
}
|
||||
if restartCalls != 0 {
|
||||
t.Fatalf("expected 0 restart, got %d", restartCalls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCustomGeoValidateAlias(t *testing.T) {
|
||||
s := CustomGeoService{}
|
||||
if err := s.validateAlias(""); !errors.Is(err, ErrCustomGeoAliasRequired) {
|
||||
t.Fatal("empty alias")
|
||||
}
|
||||
if err := s.validateAlias("Bad"); !errors.Is(err, ErrCustomGeoAliasPattern) {
|
||||
t.Fatal("uppercase")
|
||||
}
|
||||
if err := s.validateAlias("a b"); !errors.Is(err, ErrCustomGeoAliasPattern) {
|
||||
t.Fatal("space")
|
||||
}
|
||||
if err := s.validateAlias("ok_alias-1"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := s.validateAlias("geoip"); !errors.Is(err, ErrCustomGeoAliasReserved) {
|
||||
t.Fatal("reserved")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCustomGeoValidateURL(t *testing.T) {
|
||||
s := CustomGeoService{}
|
||||
if _, err := s.sanitizeURL(""); !errors.Is(err, ErrCustomGeoURLRequired) {
|
||||
t.Fatal("empty")
|
||||
}
|
||||
if _, err := s.sanitizeURL("ftp://x"); !errors.Is(err, ErrCustomGeoURLScheme) {
|
||||
t.Fatal("ftp")
|
||||
}
|
||||
if sanitized, err := s.sanitizeURL("https://example.com/a.dat"); err != nil {
|
||||
t.Fatal(err)
|
||||
} else if sanitized != "https://example.com/a.dat" {
|
||||
t.Fatalf("unexpected sanitized URL: %s", sanitized)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCustomGeoValidateType(t *testing.T) {
|
||||
s := CustomGeoService{}
|
||||
if err := s.validateType("geosite"); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := s.validateType("x"); !errors.Is(err, ErrCustomGeoInvalidType) {
|
||||
t.Fatal("bad type")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCustomGeoDownloadToPath(t *testing.T) {
|
||||
disableSSRFCheck(t)
|
||||
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("X-Test", "1")
|
||||
if r.Header.Get("If-Modified-Since") != "" {
|
||||
w.WriteHeader(http.StatusNotModified)
|
||||
return
|
||||
}
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_, _ = w.Write(make([]byte, minDatBytes+1))
|
||||
}))
|
||||
defer ts.Close()
|
||||
dir := t.TempDir()
|
||||
t.Setenv("XUI_BIN_FOLDER", dir)
|
||||
dest := filepath.Join(dir, "geoip_t.dat")
|
||||
s := CustomGeoService{}
|
||||
skipped, _, err := s.downloadToPath(ts.URL, dest, "")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if skipped {
|
||||
t.Fatal("expected download")
|
||||
}
|
||||
st, err := os.Stat(dest)
|
||||
if err != nil || st.Size() < minDatBytes {
|
||||
t.Fatalf("file %v", err)
|
||||
}
|
||||
skipped2, _, err2 := s.downloadToPath(ts.URL, dest, "")
|
||||
if err2 != nil || !skipped2 {
|
||||
t.Fatalf("304 expected skipped=%v err=%v", skipped2, err2)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCustomGeoDownloadToPath_missingLocalSendsNoIMSFromDB(t *testing.T) {
|
||||
disableSSRFCheck(t)
|
||||
lm := "Wed, 21 Oct 2015 07:28:00 GMT"
|
||||
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Header.Get("If-Modified-Since") != "" {
|
||||
w.WriteHeader(http.StatusNotModified)
|
||||
return
|
||||
}
|
||||
w.Header().Set("Last-Modified", lm)
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_, _ = w.Write(make([]byte, minDatBytes+1))
|
||||
}))
|
||||
defer ts.Close()
|
||||
dir := t.TempDir()
|
||||
t.Setenv("XUI_BIN_FOLDER", dir)
|
||||
dest := filepath.Join(dir, "geoip_rebuild.dat")
|
||||
s := CustomGeoService{}
|
||||
skipped, _, err := s.downloadToPath(ts.URL, dest, lm)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if skipped {
|
||||
t.Fatal("must not treat as not-modified when local file is missing")
|
||||
}
|
||||
if _, err := os.Stat(dest); err != nil {
|
||||
t.Fatal("file should exist after container-style rebuild")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCustomGeoDownloadToPath_repairSkipsConditional(t *testing.T) {
|
||||
disableSSRFCheck(t)
|
||||
lm := "Wed, 21 Oct 2015 07:28:00 GMT"
|
||||
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Header.Get("If-Modified-Since") != "" {
|
||||
w.WriteHeader(http.StatusNotModified)
|
||||
return
|
||||
}
|
||||
w.Header().Set("Last-Modified", lm)
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_, _ = w.Write(make([]byte, minDatBytes+1))
|
||||
}))
|
||||
defer ts.Close()
|
||||
dir := t.TempDir()
|
||||
t.Setenv("XUI_BIN_FOLDER", dir)
|
||||
dest := filepath.Join(dir, "geoip_bad.dat")
|
||||
if err := os.WriteFile(dest, make([]byte, minDatBytes-1), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
s := CustomGeoService{}
|
||||
skipped, _, err := s.downloadToPath(ts.URL, dest, lm)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if skipped {
|
||||
t.Fatal("corrupt local file must be re-downloaded, not 304")
|
||||
}
|
||||
st, err := os.Stat(dest)
|
||||
if err != nil || st.Size() < minDatBytes {
|
||||
t.Fatalf("file repaired: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCustomGeoFileNameFor(t *testing.T) {
|
||||
s := CustomGeoService{}
|
||||
if s.fileNameFor("geoip", "a") != "geoip_a.dat" {
|
||||
t.Fatal("geoip name")
|
||||
}
|
||||
if s.fileNameFor("geosite", "b") != "geosite_b.dat" {
|
||||
t.Fatal("geosite name")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLocalDatFileNeedsRepair(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
t.Setenv("XUI_BIN_FOLDER", dir)
|
||||
if !localDatFileNeedsRepair(filepath.Join(dir, "missing.dat")) {
|
||||
t.Fatal("missing")
|
||||
}
|
||||
smallPath := filepath.Join(dir, "small.dat")
|
||||
if err := os.WriteFile(smallPath, make([]byte, minDatBytes-1), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !localDatFileNeedsRepair(smallPath) {
|
||||
t.Fatal("small")
|
||||
}
|
||||
okPath := filepath.Join(dir, "ok.dat")
|
||||
if err := os.WriteFile(okPath, make([]byte, minDatBytes), 0o644); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if localDatFileNeedsRepair(okPath) {
|
||||
t.Fatal("ok size")
|
||||
}
|
||||
dirPath := filepath.Join(dir, "isdir.dat")
|
||||
if err := os.Mkdir(dirPath, 0o755); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !localDatFileNeedsRepair(dirPath) {
|
||||
t.Fatal("dir should need repair")
|
||||
}
|
||||
if !CustomGeoLocalFileNeedsRepair(dirPath) {
|
||||
t.Fatal("exported wrapper dir")
|
||||
}
|
||||
if CustomGeoLocalFileNeedsRepair(okPath) {
|
||||
t.Fatal("exported wrapper ok file")
|
||||
}
|
||||
}
|
||||
|
||||
func TestProbeCustomGeoURL_HEADOK(t *testing.T) {
|
||||
disableSSRFCheck(t)
|
||||
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method == http.MethodHead {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
return
|
||||
}
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))
|
||||
defer ts.Close()
|
||||
if err := (&CustomGeoService{}).probeCustomGeoURL(ts.URL); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProbeCustomGeoURL_HEAD405GETRange(t *testing.T) {
|
||||
disableSSRFCheck(t)
|
||||
ts := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method == http.MethodHead {
|
||||
w.WriteHeader(http.StatusMethodNotAllowed)
|
||||
return
|
||||
}
|
||||
if r.Method == http.MethodGet && r.Header.Get("Range") != "" {
|
||||
w.WriteHeader(http.StatusPartialContent)
|
||||
_, _ = w.Write([]byte{0})
|
||||
return
|
||||
}
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
}))
|
||||
defer ts.Close()
|
||||
if err := (&CustomGeoService{}).probeCustomGeoURL(ts.URL); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,151 @@
|
||||
package integration
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"github.com/mhsanaei/3x-ui/v3/internal/util/common"
|
||||
"github.com/mhsanaei/3x-ui/v3/internal/web/service"
|
||||
)
|
||||
|
||||
type NordService struct {
|
||||
service.SettingService
|
||||
}
|
||||
|
||||
var nordHTTPClient = &http.Client{Timeout: 15 * time.Second}
|
||||
|
||||
// maxResponseSize limits the maximum size of NordVPN API responses (10 MB).
|
||||
const maxResponseSize = 10 << 20
|
||||
|
||||
func (s *NordService) GetCountries() (string, error) {
|
||||
resp, err := nordHTTPClient.Get("https://api.nordvpn.com/v1/countries")
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return "", common.NewErrorf("NordVPN API error: %s", resp.Status)
|
||||
}
|
||||
body, err := io.ReadAll(io.LimitReader(resp.Body, maxResponseSize))
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return string(body), nil
|
||||
}
|
||||
|
||||
func (s *NordService) GetServers(countryId string) (string, error) {
|
||||
// Validate countryId is numeric to prevent URL injection
|
||||
for _, c := range countryId {
|
||||
if c < '0' || c > '9' {
|
||||
return "", common.NewError("invalid country ID")
|
||||
}
|
||||
}
|
||||
url := fmt.Sprintf("https://api.nordvpn.com/v2/servers?limit=0&filters[servers_technologies][id]=35&filters[country_id]=%s", countryId)
|
||||
resp, err := nordHTTPClient.Get(url)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return "", common.NewErrorf("NordVPN API error: %s", resp.Status)
|
||||
}
|
||||
body, err := io.ReadAll(io.LimitReader(resp.Body, maxResponseSize))
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
var data map[string]any
|
||||
if err := json.Unmarshal(body, &data); err != nil {
|
||||
return string(body), nil
|
||||
}
|
||||
|
||||
servers, ok := data["servers"].([]any)
|
||||
if !ok {
|
||||
return string(body), nil
|
||||
}
|
||||
|
||||
var filtered []any
|
||||
for _, s := range servers {
|
||||
if server, ok := s.(map[string]any); ok {
|
||||
if load, ok := server["load"].(float64); ok && load > 7 {
|
||||
filtered = append(filtered, s)
|
||||
}
|
||||
}
|
||||
}
|
||||
data["servers"] = filtered
|
||||
|
||||
result, _ := json.Marshal(data)
|
||||
return string(result), nil
|
||||
}
|
||||
|
||||
func (s *NordService) SetKey(privateKey string) (string, error) {
|
||||
if privateKey == "" {
|
||||
return "", common.NewError("private key cannot be empty")
|
||||
}
|
||||
nordData := map[string]string{
|
||||
"private_key": privateKey,
|
||||
"token": "",
|
||||
}
|
||||
data, _ := json.Marshal(nordData)
|
||||
err := s.SettingService.SetNord(string(data))
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return string(data), nil
|
||||
}
|
||||
|
||||
func (s *NordService) GetCredentials(token string) (string, error) {
|
||||
url := "https://api.nordvpn.com/v1/users/services/credentials"
|
||||
req, err := http.NewRequest("GET", url, nil)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
req.SetBasicAuth("token", token)
|
||||
|
||||
resp, err := nordHTTPClient.Do(req)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return "", common.NewErrorf("NordVPN API error: %s", resp.Status)
|
||||
}
|
||||
|
||||
body, err := io.ReadAll(io.LimitReader(resp.Body, maxResponseSize))
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
var creds map[string]any
|
||||
if err := json.Unmarshal(body, &creds); err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
privateKey, ok := creds["nordlynx_private_key"].(string)
|
||||
if !ok || privateKey == "" {
|
||||
return "", common.NewError("failed to retrieve NordLynx private key")
|
||||
}
|
||||
|
||||
nordData := map[string]string{
|
||||
"private_key": privateKey,
|
||||
"token": token,
|
||||
}
|
||||
data, _ := json.Marshal(nordData)
|
||||
err = s.SettingService.SetNord(string(data))
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
return string(data), nil
|
||||
}
|
||||
|
||||
func (s *NordService) GetNordData() (string, error) {
|
||||
return s.SettingService.GetNord()
|
||||
}
|
||||
|
||||
func (s *NordService) DelNordData() error {
|
||||
return s.SettingService.SetNord("")
|
||||
}
|
||||
@@ -0,0 +1,103 @@
|
||||
package integration
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/mhsanaei/3x-ui/v3/internal/util/netproxy"
|
||||
)
|
||||
|
||||
func recordingProxy(t *testing.T, hits *int64) *httptest.Server {
|
||||
t.Helper()
|
||||
return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
atomic.AddInt64(hits, 1)
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_, _ = w.Write(make([]byte, minDatBytes+1))
|
||||
}))
|
||||
}
|
||||
|
||||
func originServer(t *testing.T, hits *int64) *httptest.Server {
|
||||
t.Helper()
|
||||
return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
atomic.AddInt64(hits, 1)
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_, _ = w.Write(make([]byte, minDatBytes+1))
|
||||
}))
|
||||
}
|
||||
|
||||
func TestPanelProxy_NetproxyHelperRoutesThroughProxy(t *testing.T) {
|
||||
var proxyHits, originHits int64
|
||||
proxy := recordingProxy(t, &proxyHits)
|
||||
defer proxy.Close()
|
||||
origin := originServer(t, &originHits)
|
||||
defer origin.Close()
|
||||
|
||||
client, err := netproxy.NewHTTPClient(proxy.URL, 5*time.Second)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
resp, err := client.Get(origin.URL)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_ = resp.Body.Close()
|
||||
|
||||
if atomic.LoadInt64(&proxyHits) != 1 {
|
||||
t.Fatalf("expected panel proxy to be hit once, got %d (origin hits=%d)", proxyHits, originHits)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPanelProxy_CustomGeoDownloadUsesProxy(t *testing.T) {
|
||||
disableSSRFCheck(t)
|
||||
|
||||
var proxyHits, originHits int64
|
||||
proxy := recordingProxy(t, &proxyHits)
|
||||
defer proxy.Close()
|
||||
origin := originServer(t, &originHits)
|
||||
defer origin.Close()
|
||||
|
||||
dir := t.TempDir()
|
||||
t.Setenv("XUI_BIN_FOLDER", dir)
|
||||
dest := filepath.Join(dir, "geosite_repro.dat")
|
||||
|
||||
s := CustomGeoService{getPanelProxy: func() (string, error) { return proxy.URL, nil }}
|
||||
if _, _, err := s.downloadToPath(origin.URL, dest, ""); err != nil {
|
||||
t.Fatalf("download failed: %v", err)
|
||||
}
|
||||
if _, err := os.Stat(dest); err != nil {
|
||||
t.Fatalf("expected file to be written: %v", err)
|
||||
}
|
||||
|
||||
if got := atomic.LoadInt64(&proxyHits); got != 1 {
|
||||
t.Fatalf("custom geo download did not route through the Panel Network Proxy "+
|
||||
"(proxy hits=%d, origin hits=%d)", got, atomic.LoadInt64(&originHits))
|
||||
}
|
||||
}
|
||||
|
||||
func TestPanelProxy_CustomGeoDownloadDirectWhenUnset(t *testing.T) {
|
||||
disableSSRFCheck(t)
|
||||
|
||||
var proxyHits, originHits int64
|
||||
proxy := recordingProxy(t, &proxyHits)
|
||||
defer proxy.Close()
|
||||
origin := originServer(t, &originHits)
|
||||
defer origin.Close()
|
||||
|
||||
dir := t.TempDir()
|
||||
t.Setenv("XUI_BIN_FOLDER", dir)
|
||||
dest := filepath.Join(dir, "geosite_direct.dat")
|
||||
|
||||
s := CustomGeoService{}
|
||||
if _, _, err := s.downloadToPath(origin.URL, dest, ""); err != nil {
|
||||
t.Fatalf("download failed: %v", err)
|
||||
}
|
||||
if atomic.LoadInt64(&proxyHits) != 0 || atomic.LoadInt64(&originHits) != 1 {
|
||||
t.Fatalf("expected direct connection (proxy=0, origin=1), got proxy=%d origin=%d",
|
||||
atomic.LoadInt64(&proxyHits), atomic.LoadInt64(&originHits))
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,267 @@
|
||||
package integration
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"os"
|
||||
"time"
|
||||
|
||||
"github.com/mhsanaei/3x-ui/v3/internal/logger"
|
||||
"github.com/mhsanaei/3x-ui/v3/internal/util/common"
|
||||
"github.com/mhsanaei/3x-ui/v3/internal/util/wireguard"
|
||||
"github.com/mhsanaei/3x-ui/v3/internal/web/service"
|
||||
)
|
||||
|
||||
// WarpService provides business logic for Cloudflare WARP integration.
|
||||
// It manages WARP configuration and connectivity settings.
|
||||
type WarpService struct {
|
||||
service.SettingService
|
||||
}
|
||||
|
||||
const (
|
||||
warpAPIBase = "https://api.cloudflareclient.com/v0a4005"
|
||||
warpClientVer = "a-6.30-3596"
|
||||
)
|
||||
|
||||
func (s *WarpService) GetWarpData() (string, error) {
|
||||
return s.SettingService.GetWarp()
|
||||
}
|
||||
|
||||
func (s *WarpService) DelWarpData() error {
|
||||
return s.SettingService.SetWarp("")
|
||||
}
|
||||
|
||||
func (s *WarpService) GetWarpConfig() (string, error) {
|
||||
warpData, err := s.loadWarpCreds()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
url := fmt.Sprintf("%s/reg/%s", warpAPIBase, warpData["device_id"])
|
||||
req, err := http.NewRequest(http.MethodGet, url, nil)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
req.Header.Set("Authorization", "Bearer "+warpData["access_token"])
|
||||
|
||||
body, err := s.doWarpRequest(req)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return string(body), nil
|
||||
}
|
||||
|
||||
func (s *WarpService) RegWarp(secretKey string, publicKey string) (string, error) {
|
||||
hostName, _ := os.Hostname()
|
||||
reqBody, err := json.Marshal(map[string]any{
|
||||
"key": publicKey,
|
||||
"tos": time.Now().UTC().Format("2006-01-02T15:04:05.000Z"),
|
||||
"type": "PC",
|
||||
"model": "x-ui",
|
||||
"name": hostName,
|
||||
})
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
req, err := http.NewRequest(http.MethodPost, warpAPIBase+"/reg", bytes.NewReader(reqBody))
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
req.Header.Set("CF-Client-Version", warpClientVer)
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
|
||||
body, err := s.doWarpRequest(req)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
var rsp map[string]any
|
||||
if err := json.Unmarshal(body, &rsp); err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
deviceID, ok := rsp["id"].(string)
|
||||
if !ok {
|
||||
return "", common.NewError("warp register: missing 'id' in response")
|
||||
}
|
||||
token, ok := rsp["token"].(string)
|
||||
if !ok {
|
||||
return "", common.NewError("warp register: missing 'token' in response")
|
||||
}
|
||||
account, ok := rsp["account"].(map[string]any)
|
||||
if !ok {
|
||||
return "", common.NewError("warp register: missing 'account' in response")
|
||||
}
|
||||
license, ok := account["license"].(string)
|
||||
if !ok {
|
||||
return "", common.NewError("warp register: missing 'account.license' in response")
|
||||
}
|
||||
|
||||
warpData := map[string]string{
|
||||
"access_token": token,
|
||||
"device_id": deviceID,
|
||||
"license_key": license,
|
||||
"private_key": secretKey,
|
||||
}
|
||||
if config, ok := rsp["config"].(map[string]any); ok {
|
||||
if clientID, ok := config["client_id"].(string); ok {
|
||||
warpData["client_id"] = clientID
|
||||
}
|
||||
}
|
||||
warpJSON, err := json.MarshalIndent(warpData, "", " ")
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if err := s.SettingService.SetWarp(string(warpJSON)); err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
result, err := json.MarshalIndent(map[string]any{
|
||||
"data": warpData,
|
||||
"config": json.RawMessage(body),
|
||||
}, "", " ")
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return string(result), nil
|
||||
}
|
||||
|
||||
func (s *WarpService) SetWarpLicense(license string) (string, error) {
|
||||
warpData, err := s.loadWarpCreds()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
url := fmt.Sprintf("%s/reg/%s/account", warpAPIBase, warpData["device_id"])
|
||||
reqBody, err := json.Marshal(map[string]string{"license": license})
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
req, err := http.NewRequest(http.MethodPut, url, bytes.NewReader(reqBody))
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
req.Header.Set("Authorization", "Bearer "+warpData["access_token"])
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
|
||||
body, err := s.doWarpRequest(req)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
var response map[string]any
|
||||
if err := json.Unmarshal(body, &response); err != nil {
|
||||
return "", err
|
||||
}
|
||||
if _, ok := response["id"].(string); !ok {
|
||||
return "", common.NewErrorf("warp set license failed: unexpected response: %s", string(body))
|
||||
}
|
||||
|
||||
warpData["license_key"] = license
|
||||
newWarpData, err := json.MarshalIndent(warpData, "", " ")
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if err := s.SettingService.SetWarp(string(newWarpData)); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return string(newWarpData), nil
|
||||
}
|
||||
|
||||
func (s *WarpService) ChangeWarpIP() (string, error) {
|
||||
warpDataMap, err := s.loadWarpCreds()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
privKey, pubKey, err := wireguard.GenerateWireguardKeypair()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
result, err := s.RegWarp(privKey, pubKey)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
var parsed struct {
|
||||
Data map[string]string `json:"data"`
|
||||
Config map[string]interface{} `json:"config"`
|
||||
}
|
||||
if err := json.Unmarshal([]byte(result), &parsed); err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
xraySvc := service.XraySettingService{}
|
||||
if err := xraySvc.UpdateWarpXraySetting(parsed.Data, parsed.Config); err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
if license, ok := warpDataMap["license_key"]; ok && len(license) >= 26 {
|
||||
if _, licErr := s.SetWarpLicense(license); licErr != nil {
|
||||
logger.Warning("ChangeWarpIP: failed to re-apply WARP license: ", licErr)
|
||||
}
|
||||
}
|
||||
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// loadWarpCreds reads the stored warp JSON and ensures access_token + device_id are set.
|
||||
func (s *WarpService) loadWarpCreds() (map[string]string, error) {
|
||||
warp, err := s.SettingService.GetWarp()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var data map[string]string
|
||||
if err := json.Unmarshal([]byte(warp), &data); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if data["access_token"] == "" || data["device_id"] == "" {
|
||||
return nil, common.NewError("warp not registered: missing access_token or device_id")
|
||||
}
|
||||
return data, nil
|
||||
}
|
||||
|
||||
// doWarpRequest sends the request and returns the response body on 2xx.
|
||||
// Non-2xx responses are returned as errors including the status code and body.
|
||||
func (s *WarpService) doWarpRequest(req *http.Request) ([]byte, error) {
|
||||
client := s.NewProxiedHTTPClient(15 * time.Second)
|
||||
resp, err := client.Do(req)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if resp.StatusCode < 200 || resp.StatusCode >= 300 {
|
||||
if msg := parseWarpError(body); msg != "" {
|
||||
return nil, common.NewError(msg)
|
||||
}
|
||||
return nil, common.NewErrorf("warp api %s %s returned status %d: %s",
|
||||
req.Method, req.URL.Path, resp.StatusCode, string(body))
|
||||
}
|
||||
return body, nil
|
||||
}
|
||||
|
||||
func parseWarpError(body []byte) string {
|
||||
var env struct {
|
||||
Errors []struct {
|
||||
Message string `json:"message"`
|
||||
} `json:"errors"`
|
||||
}
|
||||
if err := json.Unmarshal(body, &env); err != nil {
|
||||
return ""
|
||||
}
|
||||
if len(env.Errors) == 0 || env.Errors[0].Message == "" {
|
||||
return ""
|
||||
}
|
||||
return env.Errors[0].Message
|
||||
}
|
||||
Reference in New Issue
Block a user