diff --git a/admin/assets/js/01-core.js b/admin/assets/js/01-core.js index a0f172b..9383bf2 100644 --- a/admin/assets/js/01-core.js +++ b/admin/assets/js/01-core.js @@ -683,7 +683,12 @@ function clientTrafficHTML(c) { const up = Number(c.uplink_bytes || 0); const down = Number(c.downlink_bytes || 0); const total = Number(c.total_bytes || (up + down) || 0); - return `${escapeHTML(formatBytes(total))}
↑ ${escapeHTML(formatBytes(up))} · ↓ ${escapeHTML(formatBytes(down))}
`; + const quota = Number(c.data_quota_bytes || 0); + const quotaLabel = quota > 0 ? formatBytes(quota) : "∞"; + const state = c.quota_exceeded + ? (c.quota_action === "throttle" ? ` · ${t("throttled")}` : ` · ${t("blocked")}`) + : ""; + return `${escapeHTML(formatBytes(total))} / ${escapeHTML(quotaLabel)}${escapeHTML(state)}
↑ ${escapeHTML(formatBytes(up))} · ↓ ${escapeHTML(formatBytes(down))}
`; } function updateCell(row, name, html) { diff --git a/admin/assets/js/03-ssh-users.js b/admin/assets/js/03-ssh-users.js index 24fdc50..22c34db 100644 --- a/admin/assets/js/03-ssh-users.js +++ b/admin/assets/js/03-ssh-users.js @@ -9,6 +9,10 @@ cancelUserBtn.addEventListener("click", () => { function prepareNewSSHUser() { userForm.reset(); fTotpPeriod.value = 60; fTotpWindow.value = 1; fTotpDigits.value = 6; + fQuotaAction.value = "block"; + fQuotaThrottle.value = 1; + fUsageDisplay.value = "0 B"; + fResetUsage.checked = false; const heading = document.getElementById("userFormHeading"); const title = document.getElementById("userFormTitle"); if (heading) heading.textContent = t("Create user"); @@ -68,11 +72,12 @@ const USER_SORT_EXTRACT = { max: u => u.max_connections || 0, up: u => u.limit_mbps_up || 0, down: u => u.limit_mbps_down || 0, + usage: u => Number(u.total_bytes || ((u.total_uplink_bytes || 0) + (u.total_downlink_bytes || 0)) || 0), expires: u => u.expires_at ? new Date(u.expires_at).getTime() : Infinity, owner: u => String(u.owner_username || "").toLowerCase(), }; // Columns that default to descending on first click (most/online first). -const USER_SORT_DEFAULT_DESC = new Set(["status", "conn", "max", "up", "down"]); +const USER_SORT_DEFAULT_DESC = new Set(["status", "conn", "max", "up", "down", "usage"]); let userSort = { key: "username", dir: "asc" }; let lastUsersData = []; @@ -139,6 +144,12 @@ function renderUsers(users) { const on = (u.active_conns || 0) > 0; if (on) online++; if (isExpiredDate(u.expires_at)) expiredCount++; + const totalBytes = Number(u.total_bytes || ((u.total_uplink_bytes || 0) + (u.total_downlink_bytes || 0)) || 0); + const quotaBytes = Number(u.data_quota_bytes || 0); + const quotaLabel = quotaBytes > 0 ? formatBytes(quotaBytes) : "∞"; + const quotaState = u.quota_exceeded + ? (u.quota_action === "throttle" ? ` · ${t("throttled")}` : ` · ${t("blocked")}`) + : ""; const tr = document.createElement("tr"); const cells = [ u.username, @@ -148,6 +159,7 @@ function renderUsers(users) { u.max_connections || 0, u.limit_mbps_up || 0, u.limit_mbps_down || 0, + `${formatBytes(totalBytes)} / ${quotaLabel}${quotaState}`, u.expires_at ? fmtDate(u.expires_at) : "—", ]; if (isSA) cells.push(u.owner_username || "—"); @@ -194,6 +206,12 @@ function fillUserForm(u) { fMaxConn.value = u.max_connections || ""; fUp.value = u.limit_mbps_up || ""; fDown.value = u.limit_mbps_down || ""; + fQuotaGB.value = u.data_quota_bytes ? (Number(u.data_quota_bytes) / (1024 ** 3)).toFixed(2).replace(/\.00$/, "") : "0"; + fQuotaAction.value = u.quota_action === "throttle" ? "throttle" : "block"; + fQuotaThrottle.value = u.quota_throttle_mbps || 1; + const totalBytes = Number(u.total_bytes || ((u.total_uplink_bytes || 0) + (u.total_downlink_bytes || 0)) || 0); + fUsageDisplay.value = `${formatBytes(totalBytes)} (↑ ${formatBytes(u.total_uplink_bytes || 0)} · ↓ ${formatBytes(u.total_downlink_bytes || 0)})`; + fResetUsage.checked = false; fExpires.value = u.expires_at ? localFromISO(u.expires_at) : ""; const heading = document.getElementById("userFormHeading"); const title = document.getElementById("userFormTitle"); @@ -218,6 +236,10 @@ userForm.addEventListener("submit", async e => { expires_at: isoFromLocal(fExpires.value), limit_mbps_up: parseInt(fUp.value||"0",10), limit_mbps_down: parseInt(fDown.value||"0",10), + data_quota_bytes: Math.round((parseFloat(fQuotaGB.value || "0") || 0) * (1024 ** 3)), + quota_action: fQuotaAction.value === "throttle" ? "throttle" : "block", + quota_throttle_mbps: parseInt(fQuotaThrottle.value || "1", 10) || 1, + reset_usage: !!fResetUsage.checked, server_id: selectedSSHServer(), }; try { @@ -225,6 +247,7 @@ userForm.addEventListener("submit", async e => { if (!res.ok) throw new Error(await res.text()); userStatus.textContent = t("Saved."); fPassword.value = ""; + fResetUsage.checked = false; loadUsers(); if (currentRole === "reseller") loadMe(); showPanelToast(t("SSH user saved successfully."), "success", t("SSH / SlowDNS")); diff --git a/admin/assets/js/04-xray.js b/admin/assets/js/04-xray.js index 80c4262..60ef477 100644 --- a/admin/assets/js/04-xray.js +++ b/admin/assets/js/04-xray.js @@ -325,6 +325,12 @@ function prepareXrayClientCreator(preferredTag = "") { if (uuid) uuid.value = genUUID(); const maxConns = document.getElementById("xCreateMaxConns"); if (maxConns) maxConns.value = "0"; + const quotaGB = document.getElementById("xCreateQuotaGB"); + if (quotaGB) quotaGB.value = "0"; + const quotaAction = document.getElementById("xCreateQuotaAction"); + if (quotaAction) quotaAction.value = "block"; + const quotaThrottle = document.getElementById("xCreateQuotaThrottle"); + if (quotaThrottle) quotaThrottle.value = "1"; const status = document.getElementById("xCreateClientStatus"); if (status) status.textContent = xrayCreatorInbounds.length ? t("Ready to create a new Xray client.") : t("Waiting for a compatible inbound."); updateXrayCreatorInboundLabel(); @@ -351,6 +357,9 @@ async function submitXrayClientCreator(event) { name: (document.getElementById("xCreateName")?.value || "").trim(), expires_at: isoFromLocal(document.getElementById("xCreateExpiry")?.value || ""), max_connections: parseInt(document.getElementById("xCreateMaxConns")?.value || "0", 10) || 0, + data_quota_bytes: Math.round((parseFloat(document.getElementById("xCreateQuotaGB")?.value || "0") || 0) * (1024 ** 3)), + quota_action: document.getElementById("xCreateQuotaAction")?.value === "throttle" ? "throttle" : "block", + quota_throttle_mbps: parseInt(document.getElementById("xCreateQuotaThrottle")?.value || "1", 10) || 1, server_id: selectedXrayServer(), }; if (button) button.disabled = true; diff --git a/admin/assets/js/09-xray-wizard.js b/admin/assets/js/09-xray-wizard.js index 08fbd42..af04428 100644 --- a/admin/assets/js/09-xray-wizard.js +++ b/admin/assets/js/09-xray-wizard.js @@ -6,6 +6,11 @@ function openEditXrayClient(tag, client) { document.getElementById("editXrayEmail").value = client.email || ""; document.getElementById("editXrayExpiry").value = client.expires_at ? localFromISO(client.expires_at) : ""; document.getElementById("editXrayMaxConns").value = client.max_conns || 0; + document.getElementById("editXrayQuotaGB").value = client.data_quota_bytes ? (Number(client.data_quota_bytes) / (1024 ** 3)).toFixed(2).replace(/\.00$/, "") : "0"; + document.getElementById("editXrayQuotaAction").value = client.quota_action === "throttle" ? "throttle" : "block"; + document.getElementById("editXrayQuotaThrottle").value = client.quota_throttle_mbps || 1; + document.getElementById("editXrayUsage").value = `${formatBytes(client.total_bytes || 0)} (↑ ${formatBytes(client.uplink_bytes || 0)} · ↓ ${formatBytes(client.downlink_bytes || 0)})`; + document.getElementById("editXrayResetUsage").checked = false; document.getElementById("editXrayClientStatus").textContent = ""; document.getElementById("editXrayClientPanel").classList.remove("hidden"); document.getElementById("editXrayClientPanel").scrollIntoView({ behavior:"smooth", block:"nearest" }); @@ -26,6 +31,10 @@ async function saveEditXrayClient() { email: document.getElementById("editXrayEmail").value.trim(), expires_at: isoFromLocal(document.getElementById("editXrayExpiry").value), max_connections: parseInt(document.getElementById("editXrayMaxConns").value || "0", 10), + data_quota_bytes: Math.round((parseFloat(document.getElementById("editXrayQuotaGB").value || "0") || 0) * (1024 ** 3)), + quota_action: document.getElementById("editXrayQuotaAction").value === "throttle" ? "throttle" : "block", + quota_throttle_mbps: parseInt(document.getElementById("editXrayQuotaThrottle").value || "1", 10) || 1, + reset_usage: !!document.getElementById("editXrayResetUsage").checked, server_id: selectedXrayServer(), }; try { diff --git a/admin/index.html b/admin/index.html index 2908cbc..86e94e0 100644 --- a/admin/index.html +++ b/admin/index.html @@ -273,7 +273,7 @@ - + @@ -310,6 +310,11 @@
+
+
+
+
+
@@ -372,6 +377,11 @@
+
+
+
+
+
@@ -428,6 +438,9 @@
+
+
+
Preencha os dados do novo cliente.
@@ -1542,15 +1555,15 @@ - + - - + + - + diff --git a/main.go b/main.go index 0780176..35b9b28 100644 --- a/main.go +++ b/main.go @@ -371,6 +371,13 @@ type UserConfig struct { LimitMbpsUp int `json:"limit_mbps_up"` // Mbps upstream LimitMbpsDown int `json:"limit_mbps_down"` // Mbps downstream + // Persistent data quota. Zero means unlimited. When the total uploaded + + // downloaded bytes reaches the quota, QuotaAction either blocks traffic or + // throttles the account to QuotaThrottleMbps. + DataQuotaBytes int64 `json:"data_quota_bytes"` + QuotaAction string `json:"quota_action"` + QuotaThrottleMbps int `json:"quota_throttle_mbps"` + // OwnerUsername is the reseller who created this SSH user. Empty = superadmin-owned. OwnerUsername string `json:"owner_username,omitempty"` } @@ -383,6 +390,17 @@ type UserState struct { mu sync.Mutex ActiveConns int conns map[*ssh.ServerConn]struct{} // active SSH connections for this user + + // Persistent per-user tunnel traffic. totalBytes includes reservations made + // by concurrent copy loops, while directional totals only include bytes that + // were actually written. The pending counters are flushed to PostgreSQL. + TotalUplinkBytes int64 + TotalDownlinkBytes int64 + totalBytes int64 + pendingUplinkBytes int64 + pendingDownlinkBytes int64 + quotaLimiter *rate.Limiter + quotaLimiterMbps int } type UserManager struct { @@ -1373,6 +1391,11 @@ func (s *Store) EnsureUsersSchema(ctx context.Context) error { expires_at TEXT, limit_mbps_up INT NOT NULL DEFAULT 0, limit_mbps_down INT NOT NULL DEFAULT 0, + data_quota_bytes BIGINT NOT NULL DEFAULT 0, + quota_action TEXT NOT NULL DEFAULT 'block', + quota_throttle_mbps INT NOT NULL DEFAULT 1, + total_uplink_bytes BIGINT NOT NULL DEFAULT 0, + total_downlink_bytes BIGINT NOT NULL DEFAULT 0, totp_secret TEXT NOT NULL DEFAULT '', totp_period INT NOT NULL DEFAULT 60, totp_window INT NOT NULL DEFAULT 1, @@ -1386,6 +1409,11 @@ func (s *Store) EnsureUsersSchema(ctx context.Context) error { `ALTER TABLE ssh_users ADD COLUMN IF NOT EXISTS totp_digits INT NOT NULL DEFAULT 6`, `ALTER TABLE ssh_users ADD COLUMN IF NOT EXISTS allow_static_password BOOLEAN NOT NULL DEFAULT FALSE`, `ALTER TABLE ssh_users ADD COLUMN IF NOT EXISTS use_pam BOOLEAN NOT NULL DEFAULT FALSE`, + `ALTER TABLE ssh_users ADD COLUMN IF NOT EXISTS data_quota_bytes BIGINT NOT NULL DEFAULT 0`, + `ALTER TABLE ssh_users ADD COLUMN IF NOT EXISTS quota_action TEXT NOT NULL DEFAULT 'block'`, + `ALTER TABLE ssh_users ADD COLUMN IF NOT EXISTS quota_throttle_mbps INT NOT NULL DEFAULT 1`, + `ALTER TABLE ssh_users ADD COLUMN IF NOT EXISTS total_uplink_bytes BIGINT NOT NULL DEFAULT 0`, + `ALTER TABLE ssh_users ADD COLUMN IF NOT EXISTS total_downlink_bytes BIGINT NOT NULL DEFAULT 0`, `ALTER TABLE ssh_users ALTER COLUMN password SET DEFAULT ''`, } for _, stmt := range stmts { @@ -1433,6 +1461,8 @@ func (s *Store) migrateSSHPasswords(ctx context.Context) error { func (s *Store) LoadUsers(ctx context.Context) (map[string]*UserState, error) { rows, err := s.db.QueryContext(ctx, ` SELECT username, password, max_connections, expires_at, limit_mbps_up, limit_mbps_down, + COALESCE(data_quota_bytes, 0), COALESCE(quota_action, 'block'), COALESCE(quota_throttle_mbps, 1), + COALESCE(total_uplink_bytes, 0), COALESCE(total_downlink_bytes, 0), COALESCE(totp_secret, ''), COALESCE(totp_period, 60), COALESCE(totp_window, 1), COALESCE(totp_digits, 6), COALESCE(allow_static_password, FALSE), COALESCE(use_pam, FALSE), COALESCE(owner_username, '') @@ -1451,6 +1481,11 @@ func (s *Store) LoadUsers(ctx context.Context) (map[string]*UserState, error) { expiresAt sql.NullString limitUp int limitDown int + dataQuotaBytes int64 + quotaAction string + quotaThrottleMbps int + totalUplinkBytes int64 + totalDownlinkBytes int64 totpSecret string totpPeriod int totpWindow int @@ -1460,6 +1495,7 @@ func (s *Store) LoadUsers(ctx context.Context) (map[string]*UserState, error) { ownerUsername string ) if err := rows.Scan(&username, &password, &maxConnections, &expiresAt, &limitUp, &limitDown, + &dataQuotaBytes, "aAction, "aThrottleMbps, &totalUplinkBytes, &totalDownlinkBytes, &totpSecret, &totpPeriod, &totpWindow, &totpDigits, &allowStaticPassword, &usePAM, &ownerUsername); err != nil { return nil, err } @@ -1474,6 +1510,9 @@ func (s *Store) LoadUsers(ctx context.Context) (map[string]*UserState, error) { MaxConnections: maxConnections, LimitMbpsUp: limitUp, LimitMbpsDown: limitDown, + DataQuotaBytes: dataQuotaBytes, + QuotaAction: normalizeQuotaAction(quotaAction), + QuotaThrottleMbps: quotaThrottleMbps, TOTPSecret: totpSecret, TOTPPeriod: totpPeriod, TOTPWindow: totpWindow, @@ -1484,6 +1523,7 @@ func (s *Store) LoadUsers(ctx context.Context) (map[string]*UserState, error) { } st := &UserState{Cfg: cfg} + initSSHRuntimeUsage(st, totalUplinkBytes, totalDownlinkBytes) if expiresAt.Valid && expiresAt.String != "" { t, err := time.Parse(time.RFC3339, expiresAt.String) if err != nil { @@ -1510,15 +1550,19 @@ func (s *Store) UpsertUser(ctx context.Context, u UserConfig) error { _, err = s.db.ExecContext(ctx, ` INSERT INTO ssh_users ( username, password, max_connections, expires_at, limit_mbps_up, limit_mbps_down, + data_quota_bytes, quota_action, quota_throttle_mbps, totp_secret, totp_period, totp_window, totp_digits, allow_static_password, use_pam, owner_username ) - VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13) + VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13, $14, $15, $16) ON CONFLICT (username) DO UPDATE SET password = EXCLUDED.password, max_connections = EXCLUDED.max_connections, expires_at = EXCLUDED.expires_at, limit_mbps_up = EXCLUDED.limit_mbps_up, limit_mbps_down = EXCLUDED.limit_mbps_down, + data_quota_bytes = EXCLUDED.data_quota_bytes, + quota_action = EXCLUDED.quota_action, + quota_throttle_mbps = EXCLUDED.quota_throttle_mbps, totp_secret = EXCLUDED.totp_secret, totp_period = EXCLUDED.totp_period, totp_window = EXCLUDED.totp_window, @@ -1527,6 +1571,7 @@ func (s *Store) UpsertUser(ctx context.Context, u UserConfig) error { use_pam = EXCLUDED.use_pam`, // owner_username is intentionally excluded from UPDATE — ownership is set at creation only. u.Username, storedPassword, u.MaxConnections, u.ExpiresAt, u.LimitMbpsUp, u.LimitMbpsDown, + u.DataQuotaBytes, normalizeQuotaAction(u.QuotaAction), quotaThrottleMbpsOrDefault(u.QuotaThrottleMbps), u.TOTPSecret, u.TOTPPeriod, u.TOTPWindow, u.TOTPDigits, u.AllowStaticPassword, u.UsePAM, u.OwnerUsername) return err } @@ -1739,6 +1784,13 @@ type UserDTO struct { ExpiresAt *time.Time `json:"expires_at,omitempty"` LimitUpMbps int `json:"limit_mbps_up"` LimitDownMbps int `json:"limit_mbps_down"` + DataQuotaBytes int64 `json:"data_quota_bytes"` + QuotaAction string `json:"quota_action"` + QuotaThrottleMbps int `json:"quota_throttle_mbps"` + TotalUplinkBytes int64 `json:"total_uplink_bytes"` + TotalDownlinkBytes int64 `json:"total_downlink_bytes"` + TotalBytes int64 `json:"total_bytes"` + QuotaExceeded bool `json:"quota_exceeded"` TOTPSecret string `json:"totp_secret,omitempty"` TOTPPeriod int `json:"totp_period"` TOTPWindow int `json:"totp_window"` @@ -1773,6 +1825,9 @@ func handleListUsers(w http.ResponseWriter, r *http.Request) { cfg := u.Cfg expires := u.ExpiresAt u.mu.Unlock() + totalUp := atomic.LoadInt64(&u.TotalUplinkBytes) + totalDown := atomic.LoadInt64(&u.TotalDownlinkBytes) + totalBytes := atomic.LoadInt64(&u.totalBytes) // Resellers only see their own users if sess != nil && sess.Role == RoleReseller && cfg.OwnerUsername != sess.Username { @@ -1786,6 +1841,13 @@ func handleListUsers(w http.ResponseWriter, r *http.Request) { ExpiresAt: expires, LimitUpMbps: cfg.LimitMbpsUp, LimitDownMbps: cfg.LimitMbpsDown, + DataQuotaBytes: cfg.DataQuotaBytes, + QuotaAction: normalizeQuotaAction(cfg.QuotaAction), + QuotaThrottleMbps: quotaThrottleMbpsOrDefault(cfg.QuotaThrottleMbps), + TotalUplinkBytes: totalUp, + TotalDownlinkBytes: totalDown, + TotalBytes: totalBytes, + QuotaExceeded: cfg.DataQuotaBytes > 0 && totalBytes >= cfg.DataQuotaBytes, TOTPSecret: cfg.TOTPSecret, TOTPPeriod: cfg.TOTPPeriod, TOTPWindow: cfg.TOTPWindow, @@ -1809,6 +1871,10 @@ type UserPayload struct { ExpiresAt string `json:"expires_at"` LimitUpMbps int `json:"limit_mbps_up"` LimitDownMbps int `json:"limit_mbps_down"` + DataQuotaBytes int64 `json:"data_quota_bytes"` + QuotaAction string `json:"quota_action"` + QuotaThrottleMbps int `json:"quota_throttle_mbps"` + ResetUsage bool `json:"reset_usage,omitempty"` TOTPSecret string `json:"totp_secret"` TOTPPeriod int `json:"totp_period"` TOTPWindow int `json:"totp_window"` @@ -1839,6 +1905,10 @@ func handleCreateUser(store *Store) http.HandlerFunc { http.Error(w, "username required", http.StatusBadRequest) return } + if err := validateQuotaConfig(p.DataQuotaBytes, p.QuotaAction, p.QuotaThrottleMbps); err != nil { + http.Error(w, err.Error(), http.StatusBadRequest) + return + } ctx := r.Context() if ms, remote, err := managedServerFromID(ctx, store, p.ServerID); err != nil { @@ -1964,6 +2034,9 @@ func handleCreateUser(store *Store) http.HandlerFunc { ExpiresAt: p.ExpiresAt, LimitMbpsUp: p.LimitUpMbps, LimitMbpsDown: p.LimitDownMbps, + DataQuotaBytes: p.DataQuotaBytes, + QuotaAction: normalizeQuotaAction(p.QuotaAction), + QuotaThrottleMbps: quotaThrottleMbpsOrDefault(p.QuotaThrottleMbps), TOTPSecret: strings.TrimSpace(p.TOTPSecret), TOTPPeriod: p.TOTPPeriod, TOTPWindow: p.TOTPWindow, @@ -1978,9 +2051,14 @@ func handleCreateUser(store *Store) http.HandlerFunc { http.Error(w, "db error", http.StatusInternalServerError) return } - // Force-disconnect all active sessions for this user so new config applies. userMgr.DisconnectUser(p.Username) + if p.ResetUsage { + if err := resetSSHUserTrafficAccounting(ctx, store, p.Username); err != nil { + http.Error(w, "could not reset usage", http.StatusInternalServerError) + return + } + } reloadUsersFromDB(ctx, store) w.WriteHeader(http.StatusCreated) @@ -2191,6 +2269,10 @@ func passwordCallback(meta ssh.ConnMetadata, pass []byte) (*ssh.Permissions, err log.Printf("user %s tried to connect but account is expired", meta.User()) return nil, fmt.Errorf("account expired") } + if sshUserQuotaBlocked(u) { + log.Printf("user %s tried to connect after reaching the data quota", meta.User()) + return nil, errDataQuotaExceeded + } if err := ownerIsActive(u.Cfg.OwnerUsername); err != nil { return nil, fmt.Errorf("authentication failed: %w", err) } @@ -2246,6 +2328,9 @@ func publicKeyCallback(meta ssh.ConnMetadata, key ssh.PublicKey) (*ssh.Permissio log.Printf("user %s tried to connect but account is expired", meta.User()) return nil, fmt.Errorf("account expired") } + if sshUserQuotaBlocked(u) { + return nil, errDataQuotaExceeded + } if err := ownerIsActive(u.Cfg.OwnerUsername); err != nil { return nil, fmt.Errorf("authentication failed: %w", err) } @@ -2377,6 +2462,10 @@ type directTCPIPReq struct { } func handleDirectTCPIP(newChan ssh.NewChannel, u *UserState, upLimiter, downLimiter *rate.Limiter) { + if sshUserQuotaBlocked(u) { + newChan.Reject(ssh.Prohibited, "data quota exceeded") + return + } var req directTCPIPReq if err := ssh.Unmarshal(newChan.ExtraData(), &req); err != nil { newChan.Reject(ssh.Prohibited, "bad direct-tcpip request") @@ -2421,7 +2510,7 @@ func handleDirectTCPIP(newChan ssh.NewChannel, u *UserState, upLimiter, downLimi // upstream: SSH channel -> backend, in its own goroutine. go func() { - _, _ = copyWithRateLimit(backend, ch, upLimiter) + _, _ = copyWithRateLimit(sshQuotaWriter{w: backend, user: u, uplink: true}, ch, upLimiter) // Signal to the backend that we are done writing. if cw, ok := backend.(interface{ CloseWrite() error }); ok { _ = cw.CloseWrite() @@ -2432,7 +2521,7 @@ func handleDirectTCPIP(newChan ssh.NewChannel, u *UserState, upLimiter, downLimi // downstream: backend -> SSH channel, run in this goroutine. // handleDirectTCPIP already runs as its own goroutine (see handleConn), // so reusing it here avoids spawning a third goroutine per channel. - _, _ = copyWithRateLimit(ch, backend, downLimiter) + _, _ = copyWithRateLimit(sshQuotaWriter{w: ch, user: u, uplink: false}, backend, downLimiter) closeAll() } @@ -3074,6 +3163,7 @@ func main() { // Optional: initialize interface totals persistence (best-effort). if store != nil { statsStore = store + startSSHUserTrafficFlusher(store) ctx := context.Background() if err := store.EnsureXrayClientsSchema(ctx); err != nil { log.Printf("xray clients table: %v", err) diff --git a/quota.go b/quota.go new file mode 100644 index 0000000..c537ea4 --- /dev/null +++ b/quota.go @@ -0,0 +1,304 @@ +package main + +import ( + "context" + "errors" + "fmt" + "io" + "log" + "strings" + "sync" + "sync/atomic" + "time" + + "golang.org/x/time/rate" +) + +const ( + quotaActionBlock = "block" + quotaActionThrottle = "throttle" +) + +var errDataQuotaExceeded = errors.New("data quota exceeded") + +func normalizeQuotaAction(v string) string { + if strings.EqualFold(strings.TrimSpace(v), quotaActionThrottle) { + return quotaActionThrottle + } + return quotaActionBlock +} + +func quotaThrottleMbpsOrDefault(v int) int { + if v <= 0 { + return 1 + } + return v +} + +type sshTrafficDelta struct { + Uplink int64 + Downlink int64 +} + +var sshTrafficPersistenceMu sync.Mutex + +func (s *Store) AddSSHUserTrafficBatch(ctx context.Context, deltas map[string]sshTrafficDelta) error { + if s == nil || len(deltas) == 0 { + return nil + } + tx, err := s.db.BeginTx(ctx, nil) + if err != nil { + return err + } + stmt, err := tx.PrepareContext(ctx, ` + UPDATE ssh_users SET + total_uplink_bytes = GREATEST(total_uplink_bytes + GREATEST($2::BIGINT, 0), 0), + total_downlink_bytes = GREATEST(total_downlink_bytes + GREATEST($3::BIGINT, 0), 0) + WHERE username = $1`) + if err != nil { + _ = tx.Rollback() + return err + } + defer stmt.Close() + for username, d := range deltas { + if strings.TrimSpace(username) == "" || (d.Uplink == 0 && d.Downlink == 0) { + continue + } + if _, err := stmt.ExecContext(ctx, username, d.Uplink, d.Downlink); err != nil { + _ = tx.Rollback() + return err + } + } + return tx.Commit() +} + +func (s *Store) ResetSSHUserTraffic(ctx context.Context, username string) error { + if s == nil || strings.TrimSpace(username) == "" { + return nil + } + _, err := s.db.ExecContext(ctx, ` + UPDATE ssh_users + SET total_uplink_bytes = 0, total_downlink_bytes = 0 + WHERE username = $1`, username) + return err +} + +func initSSHRuntimeUsage(u *UserState, uplink, downlink int64) { + if u == nil { + return + } + if uplink < 0 { + uplink = 0 + } + if downlink < 0 { + downlink = 0 + } + atomic.StoreInt64(&u.TotalUplinkBytes, uplink) + atomic.StoreInt64(&u.TotalDownlinkBytes, downlink) + atomic.StoreInt64(&u.totalBytes, uplink+downlink) + atomic.StoreInt64(&u.pendingUplinkBytes, 0) + atomic.StoreInt64(&u.pendingDownlinkBytes, 0) +} + +func resetSSHRuntimeUsage(username string) { + u, ok := userMgr.Get(username) + if !ok || u == nil { + return + } + initSSHRuntimeUsage(u, 0, 0) + u.mu.Lock() + u.quotaLimiter = nil + u.quotaLimiterMbps = 0 + u.mu.Unlock() +} + +func resetSSHUserTrafficAccounting(ctx context.Context, store *Store, username string) error { + sshTrafficPersistenceMu.Lock() + defer sshTrafficPersistenceMu.Unlock() + if err := store.ResetSSHUserTraffic(ctx, username); err != nil { + return err + } + resetSSHRuntimeUsage(username) + return nil +} + +func sshUserQuotaBlocked(u *UserState) bool { + if u == nil { + return false + } + u.mu.Lock() + quota := u.Cfg.DataQuotaBytes + action := normalizeQuotaAction(u.Cfg.QuotaAction) + u.mu.Unlock() + return quota > 0 && action == quotaActionBlock && atomic.LoadInt64(&u.totalBytes) >= quota +} + +func sshQuotaLimiter(u *UserState, mbps int) *rate.Limiter { + mbps = quotaThrottleMbpsOrDefault(mbps) + u.mu.Lock() + defer u.mu.Unlock() + if u.quotaLimiter == nil || u.quotaLimiterMbps != mbps { + bps := mbpsToBytesPerSec(mbps) + burst := int(bps) + if burst < copyBufSize { + burst = copyBufSize + } + u.quotaLimiter = rate.NewLimiter(rate.Limit(bps), burst) + u.quotaLimiterMbps = mbps + } + return u.quotaLimiter +} + +func reserveSSHUserBytes(u *UserState, requested int) (allowed int, throttle *rate.Limiter, stopAfter bool) { + if u == nil || requested <= 0 { + return 0, nil, false + } + u.mu.Lock() + quota := u.Cfg.DataQuotaBytes + action := normalizeQuotaAction(u.Cfg.QuotaAction) + throttleMbps := u.Cfg.QuotaThrottleMbps + u.mu.Unlock() + + n := int64(requested) + if quota <= 0 { + atomic.AddInt64(&u.totalBytes, n) + return requested, nil, false + } + + if action == quotaActionThrottle { + previous := atomic.AddInt64(&u.totalBytes, n) - n + if previous+n > quota { + return requested, sshQuotaLimiter(u, throttleMbps), false + } + return requested, nil, false + } + + for { + used := atomic.LoadInt64(&u.totalBytes) + remaining := quota - used + if remaining <= 0 { + return 0, nil, true + } + take := n + if take > remaining { + take = remaining + } + if atomic.CompareAndSwapInt64(&u.totalBytes, used, used+take) { + return int(take), nil, take < n || used+take >= quota + } + } +} + +func finishSSHUserReservation(u *UserState, uplink bool, reserved, written int) { + if u == nil || reserved <= 0 { + return + } + if written < 0 { + written = 0 + } + if written > reserved { + written = reserved + } + if written < reserved { + atomic.AddInt64(&u.totalBytes, -int64(reserved-written)) + } + if written == 0 { + return + } + if uplink { + atomic.AddInt64(&u.TotalUplinkBytes, int64(written)) + atomic.AddInt64(&u.pendingUplinkBytes, int64(written)) + } else { + atomic.AddInt64(&u.TotalDownlinkBytes, int64(written)) + atomic.AddInt64(&u.pendingDownlinkBytes, int64(written)) + } +} + +type sshQuotaWriter struct { + w io.Writer + user *UserState + uplink bool +} + +func (qw sshQuotaWriter) Write(p []byte) (int, error) { + allowed, quotaLimiter, stopAfter := reserveSSHUserBytes(qw.user, len(p)) + if allowed <= 0 { + return 0, errDataQuotaExceeded + } + if quotaLimiter != nil { + if err := quotaLimiter.WaitN(context.Background(), allowed); err != nil { + finishSSHUserReservation(qw.user, qw.uplink, allowed, 0) + return 0, err + } + } + n, err := qw.w.Write(p[:allowed]) + finishSSHUserReservation(qw.user, qw.uplink, allowed, n) + if err != nil { + return n, err + } + if stopAfter || allowed < len(p) { + return n, errDataQuotaExceeded + } + return n, nil +} + +func startSSHUserTrafficFlusher(store *Store) { + if store == nil { + return + } + go func() { + ticker := time.NewTicker(5 * time.Second) + defer ticker.Stop() + for range ticker.C { + flushSSHUserTraffic(store) + } + }() +} + +func flushSSHUserTraffic(store *Store) { + if store == nil { + return + } + sshTrafficPersistenceMu.Lock() + defer sshTrafficPersistenceMu.Unlock() + deltas := make(map[string]sshTrafficDelta) + states := make(map[string]*UserState) + for _, u := range userMgr.List() { + if u == nil || strings.TrimSpace(u.Cfg.Username) == "" { + continue + } + up := atomic.SwapInt64(&u.pendingUplinkBytes, 0) + down := atomic.SwapInt64(&u.pendingDownlinkBytes, 0) + if up == 0 && down == 0 { + continue + } + username := u.Cfg.Username + deltas[username] = sshTrafficDelta{Uplink: up, Downlink: down} + states[username] = u + } + if len(deltas) == 0 { + return + } + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + if err := store.AddSSHUserTrafficBatch(ctx, deltas); err != nil { + log.Printf("ssh traffic flush failed: %v", err) + for username, d := range deltas { + if u := states[username]; u != nil { + atomic.AddInt64(&u.pendingUplinkBytes, d.Uplink) + atomic.AddInt64(&u.pendingDownlinkBytes, d.Downlink) + } + } + } +} + +func validateQuotaConfig(quotaBytes int64, action string, throttleMbps int) error { + if quotaBytes < 0 { + return fmt.Errorf("data_quota_bytes must be non-negative") + } + action = normalizeQuotaAction(action) + if quotaBytes > 0 && action == quotaActionThrottle && throttleMbps < 0 { + return fmt.Errorf("quota_throttle_mbps must be non-negative") + } + return nil +} diff --git a/xray_clients.go b/xray_clients.go index d5a436b..98024f0 100644 --- a/xray_clients.go +++ b/xray_clients.go @@ -18,6 +18,9 @@ type XrayClientMeta struct { OwnerUsername string ExpiresAt *time.Time MaxConns int + DataQuotaBytes int64 + QuotaAction string + QuotaThrottleMbps int CreatedAt time.Time TotalUplinkBytes int64 TotalDownlinkBytes int64 @@ -35,6 +38,9 @@ func (s *Store) EnsureXrayClientsSchema(ctx context.Context) error { owner_username TEXT NOT NULL DEFAULT '', expires_at TIMESTAMPTZ, max_conns INT NOT NULL DEFAULT 0, + data_quota_bytes BIGINT NOT NULL DEFAULT 0, + quota_action TEXT NOT NULL DEFAULT 'block', + quota_throttle_mbps INT NOT NULL DEFAULT 1, created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), total_uplink_bytes BIGINT NOT NULL DEFAULT 0, total_downlink_bytes BIGINT NOT NULL DEFAULT 0, @@ -42,6 +48,9 @@ func (s *Store) EnsureXrayClientsSchema(ctx context.Context) error { active_connections INT NOT NULL DEFAULT 0 )`, `ALTER TABLE xray_clients ADD COLUMN IF NOT EXISTS owner_username TEXT NOT NULL DEFAULT ''`, + `ALTER TABLE xray_clients ADD COLUMN IF NOT EXISTS data_quota_bytes BIGINT NOT NULL DEFAULT 0`, + `ALTER TABLE xray_clients ADD COLUMN IF NOT EXISTS quota_action TEXT NOT NULL DEFAULT 'block'`, + `ALTER TABLE xray_clients ADD COLUMN IF NOT EXISTS quota_throttle_mbps INT NOT NULL DEFAULT 1`, `ALTER TABLE xray_clients ADD COLUMN IF NOT EXISTS total_uplink_bytes BIGINT NOT NULL DEFAULT 0`, `ALTER TABLE xray_clients ADD COLUMN IF NOT EXISTS total_downlink_bytes BIGINT NOT NULL DEFAULT 0`, `ALTER TABLE xray_clients ADD COLUMN IF NOT EXISTS last_active TIMESTAMPTZ`, @@ -61,16 +70,20 @@ func (s *Store) UpsertXrayClientMeta(ctx context.Context, m XrayClientMeta) erro expiresAt = *m.ExpiresAt } _, err := s.db.ExecContext(ctx, ` - INSERT INTO xray_clients (uuid, name, email, inbound_tag, owner_username, expires_at, max_conns) - VALUES ($1, $2, $3, $4, $5, $6, $7) + INSERT INTO xray_clients (uuid, name, email, inbound_tag, owner_username, expires_at, max_conns, data_quota_bytes, quota_action, quota_throttle_mbps) + VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10) ON CONFLICT (uuid) DO UPDATE SET name = EXCLUDED.name, email = EXCLUDED.email, inbound_tag = CASE WHEN EXCLUDED.inbound_tag <> '' THEN EXCLUDED.inbound_tag ELSE xray_clients.inbound_tag END, owner_username = CASE WHEN EXCLUDED.owner_username <> '' THEN EXCLUDED.owner_username ELSE xray_clients.owner_username END, - expires_at = EXCLUDED.expires_at, - max_conns = EXCLUDED.max_conns`, - m.UUID, m.Name, m.Email, m.InboundTag, m.OwnerUsername, expiresAt, m.MaxConns) + expires_at = EXCLUDED.expires_at, + max_conns = EXCLUDED.max_conns, + data_quota_bytes = EXCLUDED.data_quota_bytes, + quota_action = EXCLUDED.quota_action, + quota_throttle_mbps = EXCLUDED.quota_throttle_mbps`, + m.UUID, m.Name, m.Email, m.InboundTag, m.OwnerUsername, expiresAt, m.MaxConns, + m.DataQuotaBytes, normalizeQuotaAction(m.QuotaAction), quotaThrottleMbpsOrDefault(m.QuotaThrottleMbps)) return err } @@ -79,10 +92,13 @@ func (s *Store) GetXrayClientMeta(ctx context.Context, uuid string) (*XrayClient var expiresAt sql.NullTime var lastActive sql.NullTime err := s.db.QueryRowContext(ctx, ` - SELECT uuid, name, email, inbound_tag, COALESCE(owner_username, ''), expires_at, max_conns, created_at, + SELECT uuid, name, email, inbound_tag, COALESCE(owner_username, ''), expires_at, max_conns, + COALESCE(data_quota_bytes, 0), COALESCE(quota_action, 'block'), COALESCE(quota_throttle_mbps, 1), created_at, COALESCE(total_uplink_bytes, 0), COALESCE(total_downlink_bytes, 0), last_active, COALESCE(active_connections, 0) FROM xray_clients WHERE uuid = $1`, uuid). - Scan(&m.UUID, &m.Name, &m.Email, &m.InboundTag, &m.OwnerUsername, &expiresAt, &m.MaxConns, &m.CreatedAt, &m.TotalUplinkBytes, &m.TotalDownlinkBytes, &lastActive, &m.ActiveConnections) + Scan(&m.UUID, &m.Name, &m.Email, &m.InboundTag, &m.OwnerUsername, &expiresAt, &m.MaxConns, + &m.DataQuotaBytes, &m.QuotaAction, &m.QuotaThrottleMbps, &m.CreatedAt, + &m.TotalUplinkBytes, &m.TotalDownlinkBytes, &lastActive, &m.ActiveConnections) if err != nil { return nil, err } @@ -97,12 +113,16 @@ func (s *Store) GetXrayClientMeta(ctx context.Context, uuid string) (*XrayClient func (s *Store) DeleteXrayClientMeta(ctx context.Context, uuid string) error { _, err := s.db.ExecContext(ctx, `DELETE FROM xray_clients WHERE uuid = $1`, uuid) + if err == nil { + xrayMgr.removeNativeQuotaPolicy(uuid) + } return err } func (s *Store) ListAllXrayClients(ctx context.Context) ([]*XrayClientMeta, error) { rows, err := s.db.QueryContext(ctx, ` - SELECT uuid, name, email, inbound_tag, COALESCE(owner_username, ''), expires_at, max_conns, created_at, + SELECT uuid, name, email, inbound_tag, COALESCE(owner_username, ''), expires_at, max_conns, + COALESCE(data_quota_bytes, 0), COALESCE(quota_action, 'block'), COALESCE(quota_throttle_mbps, 1), created_at, COALESCE(total_uplink_bytes, 0), COALESCE(total_downlink_bytes, 0), last_active, COALESCE(active_connections, 0) FROM xray_clients ORDER BY created_at DESC`) if err != nil { @@ -114,7 +134,8 @@ func (s *Store) ListAllXrayClients(ctx context.Context) ([]*XrayClientMeta, erro func (s *Store) ListXrayClientsByOwner(ctx context.Context, ownerUsername string) ([]*XrayClientMeta, error) { rows, err := s.db.QueryContext(ctx, ` - SELECT uuid, name, email, inbound_tag, COALESCE(owner_username, ''), expires_at, max_conns, created_at, + SELECT uuid, name, email, inbound_tag, COALESCE(owner_username, ''), expires_at, max_conns, + COALESCE(data_quota_bytes, 0), COALESCE(quota_action, 'block'), COALESCE(quota_throttle_mbps, 1), created_at, COALESCE(total_uplink_bytes, 0), COALESCE(total_downlink_bytes, 0), last_active, COALESCE(active_connections, 0) FROM xray_clients WHERE owner_username = $1 ORDER BY created_at DESC`, ownerUsername) if err != nil { @@ -132,7 +153,8 @@ func (s *Store) CountXrayClientsByOwner(ctx context.Context, ownerUsername strin func (s *Store) ListExpiredXrayClients(ctx context.Context) ([]*XrayClientMeta, error) { rows, err := s.db.QueryContext(ctx, ` - SELECT uuid, name, email, inbound_tag, COALESCE(owner_username, ''), expires_at, max_conns, created_at, + SELECT uuid, name, email, inbound_tag, COALESCE(owner_username, ''), expires_at, max_conns, + COALESCE(data_quota_bytes, 0), COALESCE(quota_action, 'block'), COALESCE(quota_throttle_mbps, 1), created_at, COALESCE(total_uplink_bytes, 0), COALESCE(total_downlink_bytes, 0), last_active, COALESCE(active_connections, 0) FROM xray_clients WHERE expires_at IS NOT NULL AND expires_at <= NOW()`) if err != nil { @@ -148,7 +170,9 @@ func scanXrayClientMetaRows(rows *sql.Rows) ([]*XrayClientMeta, error) { m := &XrayClientMeta{} var expiresAt sql.NullTime var lastActive sql.NullTime - if err := rows.Scan(&m.UUID, &m.Name, &m.Email, &m.InboundTag, &m.OwnerUsername, &expiresAt, &m.MaxConns, &m.CreatedAt, &m.TotalUplinkBytes, &m.TotalDownlinkBytes, &lastActive, &m.ActiveConnections); err != nil { + if err := rows.Scan(&m.UUID, &m.Name, &m.Email, &m.InboundTag, &m.OwnerUsername, &expiresAt, &m.MaxConns, + &m.DataQuotaBytes, &m.QuotaAction, &m.QuotaThrottleMbps, &m.CreatedAt, + &m.TotalUplinkBytes, &m.TotalDownlinkBytes, &lastActive, &m.ActiveConnections); err != nil { return nil, err } if expiresAt.Valid { @@ -304,3 +328,18 @@ func startXrayClientExpiryChecker(store *Store) { } }() } + +// ResetXrayClientTraffic clears a client's persistent usage without removing +// the account or changing its expiry/quota policy. +func (s *Store) ResetXrayClientTraffic(ctx context.Context, uuid string) error { + if s == nil || uuid == "" { + return nil + } + _, err := s.db.ExecContext(ctx, ` + UPDATE xray_clients SET + total_uplink_bytes = 0, + total_downlink_bytes = 0, + last_active = NULL + WHERE uuid = $1`, uuid) + return err +} diff --git a/xray_config_store.go b/xray_config_store.go index e7e0043..257f0d8 100644 --- a/xray_config_store.go +++ b/xray_config_store.go @@ -64,7 +64,8 @@ func (s *Store) UpsertXrayConfig(ctx context.Context, configKey string, data []b func (s *Store) ListXrayClientsByInbound(ctx context.Context, inboundTag string) ([]*XrayClientMeta, error) { rows, err := s.db.QueryContext(ctx, ` - SELECT uuid, name, email, inbound_tag, COALESCE(owner_username, ''), expires_at, max_conns, created_at, + SELECT uuid, name, email, inbound_tag, COALESCE(owner_username, ''), expires_at, max_conns, + COALESCE(data_quota_bytes, 0), COALESCE(quota_action, 'block'), COALESCE(quota_throttle_mbps, 1), created_at, COALESCE(total_uplink_bytes, 0), COALESCE(total_downlink_bytes, 0), last_active, COALESCE(active_connections, 0) FROM xray_clients WHERE inbound_tag = $1 ORDER BY created_at DESC`, inboundTag) if err != nil { diff --git a/xray_integration.go b/xray_integration.go index a40989c..0697a0d 100644 --- a/xray_integration.go +++ b/xray_integration.go @@ -259,8 +259,12 @@ type XrayManager struct { pollStarted bool nativeDBMu sync.Mutex + nativeTrafficPersistMu sync.Mutex nativeTrafficPending map[string]xrayPendingTraffic nativeStatsFlushStarted bool + + nativeQuotaMu sync.RWMutex + nativeQuotaByUUID map[string]*xrayNativeQuotaState } type xrayTrafficCounters struct { @@ -296,6 +300,8 @@ func initXrayManager(cfg *XrayConfig) { } xrayMgr.mu.Unlock() + xrayMgr.reloadNativeQuotaPolicies() + // In native mode the in-process emulator records traffic directly, so the // external `xray api statsquery` poller is not started (it would overwrite // the native counters with errors from a non-existent CLI endpoint). @@ -514,7 +520,7 @@ func (m *XrayManager) recordNativeDisconnect(uuid, email string) { // recordNativeTraffic accumulates in-process byte counters for a client and // queues DB persistence. Used by the native emulator instead of external // `xray api statsquery` polling. -func (m *XrayManager) recordNativeTraffic(uuid, email string, up, down int64) { +func (m *XrayManager) recordNativeTraffic(uuid, email string, up, down int64, generation uint64) { uuid = strings.TrimSpace(uuid) email = strings.TrimSpace(email) if email == "" { @@ -523,18 +529,14 @@ func (m *XrayManager) recordNativeTraffic(uuid, email string, up, down int64) { if email == "" || (up == 0 && down == 0) { return } - now := time.Now() - m.statsMu.Lock() - if m.statsByEmail == nil { - m.statsByEmail = make(map[string]xrayRuntimeStat) + state := m.nativeQuotaState(uuid) + if state != nil { + state.mu.Lock() + defer state.mu.Unlock() + if generation != state.generation { + return + } } - st := m.statsByEmail[email] - st.Email = email - st.Uplink += up - st.Downlink += down - st.LastActive = now - m.statsByEmail[email] = st - m.statsMu.Unlock() if statsStore != nil && uuid != "" { m.nativeDBMu.Lock() @@ -548,6 +550,20 @@ func (m *XrayManager) recordNativeTraffic(uuid, email string, up, down int64) { m.nativeTrafficPending[uuid] = p m.nativeDBMu.Unlock() } + + now := time.Now() + m.statsMu.Lock() + if m.statsByEmail == nil { + m.statsByEmail = make(map[string]xrayRuntimeStat) + } + st := m.statsByEmail[email] + st.Email = email + st.Uplink += up + st.Downlink += down + st.LastActive = now + m.statsByEmail[email] = st + m.statsMu.Unlock() + } func (m *XrayManager) startNativeStatsFlusher() { @@ -594,6 +610,8 @@ func (m *XrayManager) flushNativeStatsToDB() { if statsStore == nil { return } + m.nativeTrafficPersistMu.Lock() + defer m.nativeTrafficPersistMu.Unlock() m.nativeDBMu.Lock() pending := m.nativeTrafficPending m.nativeTrafficPending = nil @@ -2027,12 +2045,16 @@ type XrayClientInfo struct { TotalBytes int64 `json:"total_bytes,omitempty"` ActiveConnections int `json:"active_connections,omitempty"` // Metadata from PostgreSQL (enriched by handleXrayInbounds) - Name string `json:"name,omitempty"` - ExpiresAt *time.Time `json:"expires_at,omitempty"` - ExpirationDays int `json:"expiration_days"` - MaxConns int `json:"max_conns"` - OwnerUsername string `json:"owner_username,omitempty"` - Expired bool `json:"expired,omitempty"` + Name string `json:"name,omitempty"` + ExpiresAt *time.Time `json:"expires_at,omitempty"` + ExpirationDays int `json:"expiration_days"` + MaxConns int `json:"max_conns"` + DataQuotaBytes int64 `json:"data_quota_bytes"` + QuotaAction string `json:"quota_action"` + QuotaThrottleMbps int `json:"quota_throttle_mbps"` + QuotaExceeded bool `json:"quota_exceeded,omitempty"` + OwnerUsername string `json:"owner_username,omitempty"` + Expired bool `json:"expired,omitempty"` } // XrayInboundInfo is returned by /api/xray/inbounds. @@ -2338,6 +2360,10 @@ func handleXrayInbounds(w http.ResponseWriter, r *http.Request) { Name: m.Name, ExpiresAt: m.ExpiresAt, MaxConns: m.MaxConns, + DataQuotaBytes: m.DataQuotaBytes, + QuotaAction: normalizeQuotaAction(m.QuotaAction), + QuotaThrottleMbps: quotaThrottleMbpsOrDefault(m.QuotaThrottleMbps), + QuotaExceeded: m.DataQuotaBytes > 0 && m.TotalUplinkBytes+m.TotalDownlinkBytes >= m.DataQuotaBytes, OwnerUsername: m.OwnerUsername, UplinkBytes: m.TotalUplinkBytes, DownlinkBytes: m.TotalDownlinkBytes, @@ -2421,6 +2447,7 @@ func applyXrayRuntimeStats(c *XrayClientInfo) { c.DownlinkBytes = st.Downlink } c.TotalBytes = c.UplinkBytes + c.DownlinkBytes + c.QuotaExceeded = c.DataQuotaBytes > 0 && c.TotalBytes >= c.DataQuotaBytes if st.ActiveConnections > c.ActiveConnections { c.ActiveConnections = st.ActiveConnections } @@ -2437,14 +2464,17 @@ func handleXrayClientAdd(w http.ResponseWriter, r *http.Request) { return } var req struct { - InboundTag string `json:"inbound_tag"` - UUID string `json:"uuid"` - Email string `json:"email"` - Name string `json:"name"` - ExpiresAt string `json:"expires_at"` // RFC3339 or YYYY-MM-DD or empty - MaxConnections int `json:"max_connections"` - OwnerUsername string `json:"owner_username,omitempty"` - ServerID string `json:"server_id,omitempty"` + InboundTag string `json:"inbound_tag"` + UUID string `json:"uuid"` + Email string `json:"email"` + Name string `json:"name"` + ExpiresAt string `json:"expires_at"` // RFC3339 or YYYY-MM-DD or empty + MaxConnections int `json:"max_connections"` + DataQuotaBytes int64 `json:"data_quota_bytes"` + QuotaAction string `json:"quota_action"` + QuotaThrottleMbps int `json:"quota_throttle_mbps"` + OwnerUsername string `json:"owner_username,omitempty"` + ServerID string `json:"server_id,omitempty"` } if err := json.NewDecoder(r.Body).Decode(&req); err != nil { http.Error(w, "invalid json", http.StatusBadRequest) @@ -2454,6 +2484,10 @@ func handleXrayClientAdd(w http.ResponseWriter, r *http.Request) { http.Error(w, "inbound_tag and uuid required", http.StatusBadRequest) return } + if err := validateQuotaConfig(req.DataQuotaBytes, req.QuotaAction, req.QuotaThrottleMbps); err != nil { + http.Error(w, err.Error(), http.StatusBadRequest) + return + } if ms, remote, err := managedServerFromID(r.Context(), statsStore, req.ServerID); err != nil { http.Error(w, err.Error(), http.StatusBadRequest) return @@ -2543,12 +2577,15 @@ func handleXrayClientAdd(w http.ResponseWriter, r *http.Request) { } if statsStore != nil { meta := XrayClientMeta{ - UUID: req.UUID, - Name: req.Name, - Email: req.Email, - InboundTag: req.InboundTag, - OwnerUsername: ownerUsername, - MaxConns: req.MaxConnections, + UUID: req.UUID, + Name: req.Name, + Email: req.Email, + InboundTag: req.InboundTag, + OwnerUsername: ownerUsername, + MaxConns: req.MaxConnections, + DataQuotaBytes: req.DataQuotaBytes, + QuotaAction: normalizeQuotaAction(req.QuotaAction), + QuotaThrottleMbps: quotaThrottleMbpsOrDefault(req.QuotaThrottleMbps), } if req.ExpiresAt != "" { var t time.Time @@ -2565,6 +2602,8 @@ func handleXrayClientAdd(w http.ResponseWriter, r *http.Request) { } if err := statsStore.UpsertXrayClientMeta(r.Context(), meta); err != nil { xrayLogf("xray: save meta for %s: %v", req.UUID, err) + } else { + xrayMgr.setNativeQuotaPolicy(&meta) } } xrayMgr.restartIfExternalRunning() @@ -2579,12 +2618,16 @@ func handleXrayClientUpdate(w http.ResponseWriter, r *http.Request) { return } var req struct { - UUID string `json:"uuid"` - Name string `json:"name"` - Email string `json:"email"` - ExpiresAt string `json:"expires_at"` - MaxConnections int `json:"max_connections"` - ServerID string `json:"server_id,omitempty"` + UUID string `json:"uuid"` + Name string `json:"name"` + Email string `json:"email"` + ExpiresAt string `json:"expires_at"` + MaxConnections int `json:"max_connections"` + DataQuotaBytes int64 `json:"data_quota_bytes"` + QuotaAction string `json:"quota_action"` + QuotaThrottleMbps int `json:"quota_throttle_mbps"` + ResetUsage bool `json:"reset_usage,omitempty"` + ServerID string `json:"server_id,omitempty"` } if err := json.NewDecoder(r.Body).Decode(&req); err != nil { http.Error(w, "invalid json", http.StatusBadRequest) @@ -2594,6 +2637,10 @@ func handleXrayClientUpdate(w http.ResponseWriter, r *http.Request) { http.Error(w, "uuid required", http.StatusBadRequest) return } + if err := validateQuotaConfig(req.DataQuotaBytes, req.QuotaAction, req.QuotaThrottleMbps); err != nil { + http.Error(w, err.Error(), http.StatusBadRequest) + return + } if ms, remote, err := managedServerFromID(r.Context(), statsStore, req.ServerID); err != nil { http.Error(w, err.Error(), http.StatusBadRequest) return @@ -2629,12 +2676,17 @@ func handleXrayClientUpdate(w http.ResponseWriter, r *http.Request) { } meta := XrayClientMeta{ - UUID: req.UUID, - Name: req.Name, - Email: req.Email, - InboundTag: existing.InboundTag, - OwnerUsername: existing.OwnerUsername, - MaxConns: req.MaxConnections, + UUID: req.UUID, + Name: req.Name, + Email: req.Email, + InboundTag: existing.InboundTag, + OwnerUsername: existing.OwnerUsername, + MaxConns: req.MaxConnections, + DataQuotaBytes: req.DataQuotaBytes, + QuotaAction: normalizeQuotaAction(req.QuotaAction), + QuotaThrottleMbps: quotaThrottleMbpsOrDefault(req.QuotaThrottleMbps), + TotalUplinkBytes: existing.TotalUplinkBytes, + TotalDownlinkBytes: existing.TotalDownlinkBytes, } if req.ExpiresAt != "" { for _, layout := range []string{time.RFC3339, "2006-01-02T15:04", "2006-01-02"} { @@ -2648,6 +2700,15 @@ func handleXrayClientUpdate(w http.ResponseWriter, r *http.Request) { http.Error(w, "update failed: "+err.Error(), http.StatusInternalServerError) return } + if req.ResetUsage { + if err := xrayMgr.resetNativeTrafficAccounting(r.Context(), statsStore, req.UUID, existing.Email); err != nil { + http.Error(w, "usage reset failed: "+err.Error(), http.StatusInternalServerError) + return + } + meta.TotalUplinkBytes = 0 + meta.TotalDownlinkBytes = 0 + } + xrayMgr.setNativeQuotaPolicy(&meta) if req.Email != "" { if err := xrayMgr.UpdateXrayClientEmail(req.UUID, req.Email); err != nil { xrayLogf("xray: update config email for %s: %v", req.UUID, err) diff --git a/xray_native.go b/xray_native.go index 0a8fe25..0dd5cdf 100644 --- a/xray_native.go +++ b/xray_native.go @@ -352,6 +352,10 @@ func (ib *nativeInbound) handleVLESS(stream net.Conn, remote net.Addr) { xrayLogf("native xray: inbound %q rejected unknown VLESS uuid from %s", ib.tag, remote) return } + if xrayMgr.nativeQuotaBlocked(client.uuid) { + xrayLogf("native xray: inbound %q rejected VLESS user %s after data quota", ib.tag, client.email) + return + } if addonLen := int(head[17]); addonLen > 0 { if _, err := io.CopyN(io.Discard, stream, int64(addonLen)); err != nil { @@ -678,7 +682,7 @@ func nativeTunnel(client io.ReadWriteCloser, backend net.Conn, uuid, email strin xrayGo("native xray TCP uplink", func() { // client -> backend defer wg.Done() defer closeAll() - _, _ = copyWithRateLimit(meteredWriter{w: backend, meter: upMeter}, client, up) + _, _ = copyWithRateLimit(xrayQuotaMeteredWriter{w: backend, meter: upMeter}, client, up) if cw, ok := backend.(interface{ CloseWrite() error }); ok { _ = cw.CloseWrite() } @@ -688,7 +692,7 @@ func nativeTunnel(client io.ReadWriteCloser, backend net.Conn, uuid, email strin xrayGo("native xray TCP downlink", func() { // backend -> client defer wg.Done() defer closeAll() - _, _ = copyWithRateLimit(meteredWriter{w: client, meter: downMeter}, backend, down) + _, _ = copyWithRateLimit(xrayQuotaMeteredWriter{w: client, meter: downMeter}, backend, down) }) wg.Wait() @@ -700,15 +704,17 @@ func nativeTunnel(client io.ReadWriteCloser, backend net.Conn, uuid, email strin // trafficMeter accumulates bytes for one direction and flushes them to the // stats manager in batches to avoid locking on every write. type trafficMeter struct { - uuid string - email string - uplink bool - n int64 + uuid string + email string + uplink bool + n int64 + quotaGeneration uint64 } const trafficFlushThreshold = 1024 * 1024 func (t *trafficMeter) add(n int) { + t.syncQuotaGeneration() t.n += int64(n) if t.n >= trafficFlushThreshold { t.flush() @@ -716,29 +722,28 @@ func (t *trafficMeter) add(n int) { } func (t *trafficMeter) flush() { + t.syncQuotaGeneration() if t.n == 0 || t.email == "" { return } if t.uplink { - xrayMgr.recordNativeTraffic(t.uuid, t.email, t.n, 0) + xrayMgr.recordNativeTraffic(t.uuid, t.email, t.n, 0, t.quotaGeneration) } else { - xrayMgr.recordNativeTraffic(t.uuid, t.email, 0, t.n) + xrayMgr.recordNativeTraffic(t.uuid, t.email, 0, t.n, t.quotaGeneration) } t.n = 0 } -// meteredWriter counts bytes as they are written through to the wrapped writer. -type meteredWriter struct { - w io.Writer - meter *trafficMeter -} - -func (mw meteredWriter) Write(p []byte) (int, error) { - n, err := mw.w.Write(p) - if n > 0 { - mw.meter.add(n) +func (t *trafficMeter) syncQuotaGeneration() { + generation := xrayMgr.nativeQuotaGeneration(t.uuid) + if t.quotaGeneration == 0 { + t.quotaGeneration = generation + return + } + if generation != t.quotaGeneration { + t.n = 0 + t.quotaGeneration = generation } - return n, err } func (ib *nativeInbound) upLimiter() *rate.Limiter { return newByteLimiter(ib.upBytesPerSec) } diff --git a/xray_native_mux.go b/xray_native_mux.go index 6ce619f..2057a99 100644 --- a/xray_native_mux.go +++ b/xray_native_mux.go @@ -481,6 +481,16 @@ func (s *nativeMuxSession) writeBackendItem(item nativeMuxUplinkItem) bool { return false } } + quotaLimiter, quotaErr := reserveNativePacketQuota(s.upMeter, len(payload)) + if quotaErr != nil { + return false + } + if quotaLimiter != nil { + if err := quotaLimiter.WaitN(s.ctx, len(payload)); err != nil { + finishNativePacketQuota(s.upMeter, len(payload), 0) + return false + } + } var n int var err error @@ -492,6 +502,7 @@ func (s *nativeMuxSession) writeBackendItem(item nativeMuxUplinkItem) bool { if isNativeDNSSinkTarget(item.host) || invalidNativeDestination(item.host, item.port) { // AdGuard/blocked endpoints must be ignored at the cheapest possible // point. Do not resolve, dial, log loudly, or keep the mux child busy. + finishNativePacketQuota(s.upMeter, len(payload), 0) xrayTracef("native xray: VLESS mux UDP fast-ignored override sink session=%d target=%s:%d", s.id, item.host, item.port) return true } @@ -503,6 +514,7 @@ func (s *nativeMuxSession) writeBackendItem(item nativeMuxUplinkItem) bool { s.lastUDPPort = item.port s.lastUDPAddr = addr } else { + finishNativePacketQuota(s.upMeter, len(payload), 0) xrayTracef("native xray: VLESS mux UDP override resolve failed session=%d target=%s:%d: %v", s.id, item.host, item.port, rerr) return true } @@ -512,9 +524,7 @@ func (s *nativeMuxSession) writeBackendItem(item nativeMuxUplinkItem) bool { if s.network == nativeMuxNetworkUDP && err == nil { _ = s.udp.SetReadDeadline(time.Now().Add(nativeMuxUDPIdleTimeout())) } - if n > 0 { - s.upMeter.add(n) - } + finishNativePacketQuota(s.upMeter, len(payload), n) if err != nil { xrayLogf("native xray: VLESS mux backend write failed session=%d: %v", s.id, err) return false @@ -571,14 +581,25 @@ func (s *nativeMuxSession) readTCPBackendLoop() { if err := s.waitDownRate(n); err != nil { return } - s.downMeter.add(n) + quotaLimiter, quotaErr := reserveNativePacketQuota(s.downMeter, n) + if quotaErr != nil { + return + } + if quotaLimiter != nil { + if err := quotaLimiter.WaitN(s.ctx, n); err != nil { + finishNativePacketQuota(s.downMeter, n, 0) + return + } + } s.writeMu.Lock() werr := writeNativeMuxData(s.client, s.id, nativeMuxStatusKeep, buf[:n]) s.writeMu.Unlock() if werr != nil { + finishNativePacketQuota(s.downMeter, n, 0) xrayLogf("native xray: VLESS mux TCP client write failed session=%d: %v", s.id, werr) return } + finishNativePacketQuota(s.downMeter, n, n) } } @@ -602,7 +623,16 @@ func (s *nativeMuxSession) readUDPBackendLoop() bool { if err := s.waitDownRate(n); err != nil { return true } - s.downMeter.add(n) + quotaLimiter, quotaErr := reserveNativePacketQuota(s.downMeter, n) + if quotaErr != nil { + return true + } + if quotaLimiter != nil { + if err := quotaLimiter.WaitN(s.ctx, n); err != nil { + finishNativePacketQuota(s.downMeter, n, 0) + return true + } + } s.writeMu.Lock() // Include the UDP source endpoint on XUDP responses so clients that rely on // full-cone packet addressing can associate the datagram with the correct @@ -610,9 +640,11 @@ func (s *nativeMuxSession) readUDPBackendLoop() bool { werr := writeNativeMuxPacketData(s.client, s.id, nativeMuxStatusKeep, buf[:n], addr, s.xudp) s.writeMu.Unlock() if werr != nil { + finishNativePacketQuota(s.downMeter, n, 0) xrayLogf("native xray: VLESS mux UDP client write failed session=%d: %v", s.id, werr) return true } + finishNativePacketQuota(s.downMeter, n, n) } } diff --git a/xray_native_udp.go b/xray_native_udp.go index e90a854..23c519e 100644 --- a/xray_native_udp.go +++ b/xray_native_udp.go @@ -60,10 +60,16 @@ func nativeVLESSUDPTunnel(client io.ReadWriteCloser, backend net.Conn, uuid, ema if err := waitNativeRate(up, len(payload)); err != nil { return } - n, err := backend.Write(payload) - if n > 0 { - upMeter.add(n) + quotaLimiter, err := reserveNativePacketQuota(upMeter, len(payload)) + if err != nil { + return } + if err := waitNativeRate(quotaLimiter, len(payload)); err != nil { + finishNativePacketQuota(upMeter, len(payload), 0) + return + } + n, err := backend.Write(payload) + finishNativePacketQuota(upMeter, len(payload), n) if err != nil { xrayLogf("native xray: VLESS UDP backend write failed: %v", err) return @@ -94,11 +100,20 @@ func nativeVLESSUDPTunnel(client io.ReadWriteCloser, backend net.Conn, uuid, ema if err := waitNativeRate(down, n); err != nil { return } + quotaLimiter, err := reserveNativePacketQuota(downMeter, n) + if err != nil { + return + } + if err := waitNativeRate(quotaLimiter, n); err != nil { + finishNativePacketQuota(downMeter, n, 0) + return + } if err := writeVLESSLengthPacket(client, buf[:n]); err != nil { + finishNativePacketQuota(downMeter, n, 0) xrayLogf("native xray: VLESS UDP client write failed: %v", err) return } - downMeter.add(n) + finishNativePacketQuota(downMeter, n, n) } }) @@ -336,10 +351,16 @@ func nativeVMessUDPTunnel(client nativeVMessStream, backend net.Conn, uuid, emai if err := waitNativeRate(up, len(pkt)); err != nil { return } - n, err := backend.Write(pkt) - if n > 0 { - upMeter.add(n) + quotaLimiter, err := reserveNativePacketQuota(upMeter, len(pkt)) + if err != nil { + return } + if err := waitNativeRate(quotaLimiter, len(pkt)); err != nil { + finishNativePacketQuota(upMeter, len(pkt), 0) + return + } + n, err := backend.Write(pkt) + finishNativePacketQuota(upMeter, len(pkt), n) if err != nil { xrayLogf("native xray: VMess UDP backend write failed: %v", err) return @@ -370,11 +391,20 @@ func nativeVMessUDPTunnel(client nativeVMessStream, backend net.Conn, uuid, emai if err := waitNativeRate(down, n); err != nil { return } + quotaLimiter, err := reserveNativePacketQuota(downMeter, n) + if err != nil { + return + } + if err := waitNativeRate(quotaLimiter, n); err != nil { + finishNativePacketQuota(downMeter, n, 0) + return + } if err := client.WritePacket(buf[:n]); err != nil { + finishNativePacketQuota(downMeter, n, 0) xrayLogf("native xray: VMess UDP client write failed: %v", err) return } - downMeter.add(n) + finishNativePacketQuota(downMeter, n, n) } }) diff --git a/xray_quota.go b/xray_quota.go new file mode 100644 index 0000000..01baeea --- /dev/null +++ b/xray_quota.go @@ -0,0 +1,309 @@ +package main + +import ( + "context" + "io" + "strings" + "sync" + + "golang.org/x/time/rate" +) + +type xrayNativeQuotaState struct { + mu sync.Mutex + usedBytes int64 + quotaBytes int64 + action string + throttleMbps int + limiter *rate.Limiter + generation uint64 +} + +func (m *XrayManager) reloadNativeQuotaPolicies() { + if statsStore == nil { + return + } + metas, err := statsStore.ListAllXrayClients(context.Background()) + if err != nil { + xrayLogf("xray native quota: load policies failed: %v", err) + return + } + next := make(map[string]*xrayNativeQuotaState, len(metas)) + for _, meta := range metas { + if meta == nil || strings.TrimSpace(meta.UUID) == "" { + continue + } + next[meta.UUID] = newXrayNativeQuotaState(meta) + } + m.nativeQuotaMu.Lock() + m.nativeQuotaByUUID = next + m.nativeQuotaMu.Unlock() +} + +func newXrayNativeQuotaState(meta *XrayClientMeta) *xrayNativeQuotaState { + used := meta.TotalUplinkBytes + meta.TotalDownlinkBytes + if used < 0 { + used = 0 + } + return &xrayNativeQuotaState{ + usedBytes: used, + quotaBytes: meta.DataQuotaBytes, + action: normalizeQuotaAction(meta.QuotaAction), + throttleMbps: quotaThrottleMbpsOrDefault(meta.QuotaThrottleMbps), + generation: 1, + } +} + +func (m *XrayManager) setNativeQuotaPolicy(meta *XrayClientMeta) { + if meta == nil || strings.TrimSpace(meta.UUID) == "" { + return + } + uuid := strings.TrimSpace(meta.UUID) + m.nativeQuotaMu.Lock() + if m.nativeQuotaByUUID == nil { + m.nativeQuotaByUUID = make(map[string]*xrayNativeQuotaState) + } + existing := m.nativeQuotaByUUID[uuid] + if existing == nil { + m.nativeQuotaByUUID[uuid] = newXrayNativeQuotaState(meta) + m.nativeQuotaMu.Unlock() + return + } + m.nativeQuotaMu.Unlock() + + existing.mu.Lock() + existing.quotaBytes = meta.DataQuotaBytes + existing.action = normalizeQuotaAction(meta.QuotaAction) + existing.throttleMbps = quotaThrottleMbpsOrDefault(meta.QuotaThrottleMbps) + existing.limiter = nil + existing.mu.Unlock() +} + +func (m *XrayManager) removeNativeQuotaPolicy(uuid string) { + uuid = strings.TrimSpace(uuid) + if uuid == "" { + return + } + m.nativeQuotaMu.Lock() + delete(m.nativeQuotaByUUID, uuid) + m.nativeQuotaMu.Unlock() +} + +func (m *XrayManager) resetNativeQuotaUsage(uuid string) { + uuid = strings.TrimSpace(uuid) + m.nativeQuotaMu.RLock() + state := m.nativeQuotaByUUID[uuid] + m.nativeQuotaMu.RUnlock() + if state == nil { + return + } + state.mu.Lock() + state.usedBytes = 0 + state.limiter = nil + state.generation++ + if state.generation == 0 { + state.generation = 1 + } + state.mu.Unlock() +} + +func (m *XrayManager) nativeQuotaGeneration(uuid string) uint64 { + state := m.nativeQuotaState(uuid) + if state == nil { + return 0 + } + state.mu.Lock() + generation := state.generation + state.mu.Unlock() + return generation +} + +func (m *XrayManager) resetNativeTrafficAccounting(ctx context.Context, store *Store, uuid, email string) error { + state := m.nativeQuotaState(uuid) + if state != nil { + state.mu.Lock() + defer state.mu.Unlock() + } + + m.nativeTrafficPersistMu.Lock() + defer m.nativeTrafficPersistMu.Unlock() + m.nativeDBMu.Lock() + key := strings.TrimSpace(uuid) + var pending xrayPendingTraffic + hadPending := false + if m.nativeTrafficPending != nil { + pending, hadPending = m.nativeTrafficPending[key] + delete(m.nativeTrafficPending, key) + } + err := store.ResetXrayClientTraffic(ctx, uuid) + if err != nil && hadPending { + if m.nativeTrafficPending == nil { + m.nativeTrafficPending = make(map[string]xrayPendingTraffic) + } + m.nativeTrafficPending[key] = pending + } + m.nativeDBMu.Unlock() + if err != nil { + return err + } + if state != nil { + state.usedBytes = 0 + state.limiter = nil + state.generation++ + if state.generation == 0 { + state.generation = 1 + } + } + + m.statsMu.Lock() + for _, key := range []string{strings.TrimSpace(email), strings.TrimSpace(uuid)} { + if key == "" { + continue + } + if runtime, ok := m.statsByEmail[key]; ok { + runtime.Uplink = 0 + runtime.Downlink = 0 + m.statsByEmail[key] = runtime + } + } + m.statsMu.Unlock() + return nil +} + +func (m *XrayManager) nativeQuotaState(uuid string) *xrayNativeQuotaState { + m.nativeQuotaMu.RLock() + state := m.nativeQuotaByUUID[strings.TrimSpace(uuid)] + m.nativeQuotaMu.RUnlock() + return state +} + +func (m *XrayManager) nativeQuotaBlocked(uuid string) bool { + state := m.nativeQuotaState(uuid) + if state == nil { + return false + } + state.mu.Lock() + defer state.mu.Unlock() + return state.quotaBytes > 0 && normalizeQuotaAction(state.action) == quotaActionBlock && state.usedBytes >= state.quotaBytes +} + +func (m *XrayManager) reserveNativeQuota(uuid string, requested int) (allowed int, limiter *rate.Limiter, stopAfter bool) { + if requested <= 0 { + return 0, nil, false + } + state := m.nativeQuotaState(uuid) + if state == nil { + return requested, nil, false + } + state.mu.Lock() + defer state.mu.Unlock() + + n := int64(requested) + if state.quotaBytes <= 0 { + state.usedBytes += n + return requested, nil, false + } + if normalizeQuotaAction(state.action) == quotaActionThrottle { + previous := state.usedBytes + state.usedBytes += n + if previous+n > state.quotaBytes { + if state.limiter == nil { + bps := mbpsToBytesPerSec(quotaThrottleMbpsOrDefault(state.throttleMbps)) + burst := int(bps) + if burst < copyBufSize { + burst = copyBufSize + } + state.limiter = rate.NewLimiter(rate.Limit(bps), burst) + } + return requested, state.limiter, false + } + return requested, nil, false + } + + remaining := state.quotaBytes - state.usedBytes + if remaining <= 0 { + return 0, nil, true + } + take := n + if take > remaining { + take = remaining + } + state.usedBytes += take + return int(take), nil, take < n +} + +func (m *XrayManager) finishNativeQuotaReservation(uuid string, reserved, written int) { + if reserved <= 0 || written >= reserved { + return + } + if written < 0 { + written = 0 + } + state := m.nativeQuotaState(uuid) + if state == nil { + return + } + state.mu.Lock() + state.usedBytes -= int64(reserved - written) + if state.usedBytes < 0 { + state.usedBytes = 0 + } + state.mu.Unlock() +} + +type xrayQuotaMeteredWriter struct { + w io.Writer + meter *trafficMeter +} + +func (mw xrayQuotaMeteredWriter) Write(p []byte) (int, error) { + if mw.meter == nil { + return mw.w.Write(p) + } + allowed, limiter, stopAfter := xrayMgr.reserveNativeQuota(mw.meter.uuid, len(p)) + if allowed <= 0 { + return 0, errDataQuotaExceeded + } + if limiter != nil { + if err := limiter.WaitN(context.Background(), allowed); err != nil { + xrayMgr.finishNativeQuotaReservation(mw.meter.uuid, allowed, 0) + return 0, err + } + } + n, err := mw.w.Write(p[:allowed]) + xrayMgr.finishNativeQuotaReservation(mw.meter.uuid, allowed, n) + if n > 0 { + mw.meter.add(n) + } + if err != nil { + return n, err + } + if stopAfter || allowed < len(p) || xrayMgr.nativeQuotaBlocked(mw.meter.uuid) { + return n, errDataQuotaExceeded + } + return n, nil +} + +func reserveNativePacketQuota(meter *trafficMeter, n int) (*rate.Limiter, error) { + if meter == nil || n <= 0 { + return nil, nil + } + allowed, limiter, stopAfter := xrayMgr.reserveNativeQuota(meter.uuid, n) + if allowed != n || stopAfter { + if allowed > 0 { + xrayMgr.finishNativeQuotaReservation(meter.uuid, allowed, 0) + } + return nil, errDataQuotaExceeded + } + return limiter, nil +} + +func finishNativePacketQuota(meter *trafficMeter, reserved, written int) { + if meter == nil || reserved <= 0 { + return + } + xrayMgr.finishNativeQuotaReservation(meter.uuid, reserved, written) + if written > 0 { + meter.add(written) + } +} diff --git a/xray_vmess.go b/xray_vmess.go index 537d12d..bce182a 100644 --- a/xray_vmess.go +++ b/xray_vmess.go @@ -673,6 +673,10 @@ func (ib *nativeInbound) handleVMess(stream net.Conn, remote net.Addr) { log.Printf("native xray: inbound %q rejected unknown/expired VMess auth id from %s", ib.tag, remote) return } + if xrayMgr.nativeQuotaBlocked(client.uuid) { + log.Printf("native xray: inbound %q rejected VMess user %s after data quota", ib.tag, client.email) + return + } header, err := openVMessHeader(client.cmdKey, authid, stream) if err != nil {
UserStatusAuthConnMaxUpDnExpiresConnMaxUpDnUsage / QuotaExpires Actions