1
0

client_link.go 14 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461
  1. package service
  2. import (
  3. "fmt"
  4. "strings"
  5. "github.com/google/uuid"
  6. "github.com/mhsanaei/3x-ui/v3/internal/database"
  7. "github.com/mhsanaei/3x-ui/v3/internal/database/model"
  8. "gorm.io/gorm"
  9. "gorm.io/gorm/clause"
  10. )
  11. // applyClientRecordMerge merges incoming client-record fields onto row using the
  12. // same rules everywhere a client record is persisted: scalar quota / lifecycle /
  13. // subscription fields are applied unconditionally (so clearing them takes
  14. // effect), while credentials and identifiers are only overwritten when the
  15. // incoming value is non-empty (so a partial update preserves the stored UUID /
  16. // password / keys). CreatedAt keeps the earliest known value. Email, UpdatedAt,
  17. // and the Id primary key are intentionally not touched here — callers handle
  18. // those separately. Shared by SyncInbound (per-inbound persistence) and Update
  19. // (the no-attached-inbound fallback) so the two paths cannot diverge.
  20. func applyClientRecordMerge(row *model.ClientRecord, incoming *model.ClientRecord) {
  21. if incoming.UUID != "" {
  22. row.UUID = incoming.UUID
  23. }
  24. if incoming.Password != "" {
  25. row.Password = incoming.Password
  26. }
  27. if incoming.Auth != "" {
  28. row.Auth = incoming.Auth
  29. }
  30. if incoming.Secret != "" {
  31. row.Secret = incoming.Secret
  32. }
  33. if incoming.AdTag != "" {
  34. row.AdTag = incoming.AdTag
  35. }
  36. row.Flow = incoming.Flow
  37. if incoming.Security != "" {
  38. row.Security = incoming.Security
  39. }
  40. if incoming.Reverse != "" {
  41. row.Reverse = incoming.Reverse
  42. }
  43. if incoming.PrivateKey != "" {
  44. row.PrivateKey = incoming.PrivateKey
  45. }
  46. if incoming.PublicKey != "" {
  47. row.PublicKey = incoming.PublicKey
  48. }
  49. if incoming.AllowedIPs != "" {
  50. row.AllowedIPs = incoming.AllowedIPs
  51. }
  52. row.PreSharedKey = incoming.PreSharedKey
  53. row.KeepAlive = incoming.KeepAlive
  54. row.SubID = incoming.SubID
  55. row.LimitIP = incoming.LimitIP
  56. row.TotalGB = incoming.TotalGB
  57. row.ExpiryTime = incoming.ExpiryTime
  58. row.Enable = incoming.Enable
  59. row.TgID = incoming.TgID
  60. if incoming.Group != "" {
  61. row.Group = incoming.Group
  62. }
  63. row.Comment = incoming.Comment
  64. row.Reset = incoming.Reset
  65. row.ResetDay = incoming.ResetDay
  66. row.ResetWeekday = incoming.ResetWeekday
  67. row.ResetMax = incoming.ResetMax
  68. // Guarded like Group and AdTag: a node snapshot rebuilt from settings that
  69. // predate the cycle would otherwise silently erase it.
  70. if incoming.TrafficReset != "" {
  71. row.TrafficReset = incoming.TrafficReset
  72. }
  73. if incoming.TrafficResetDay > 0 {
  74. row.TrafficResetDay = incoming.TrafficResetDay
  75. }
  76. if incoming.CreatedAt > 0 && (row.CreatedAt == 0 || incoming.CreatedAt < row.CreatedAt) {
  77. row.CreatedAt = incoming.CreatedAt
  78. }
  79. }
  80. // SyncInbound makes the inbound's client records and links match clients
  81. // exactly: links for clients no longer in the set are removed.
  82. func (s *ClientService) SyncInbound(tx *gorm.DB, inboundId int, clients []model.Client) error {
  83. return s.syncInboundClients(tx, inboundId, clients, nil, true)
  84. }
  85. // ApplyInboundClientDelta persists only the clients an edit actually changed
  86. // plus the emails it detached, leaving every other link on the inbound alone —
  87. // the whole point being that a one-client edit must not rewrite the inbound's
  88. // entire membership set (#6252).
  89. func (s *ClientService) ApplyInboundClientDelta(tx *gorm.DB, inboundId int, changed []model.Client, detachEmails []string) error {
  90. return s.syncInboundClients(tx, inboundId, changed, detachEmails, false)
  91. }
  92. func (s *ClientService) syncInboundClients(tx *gorm.DB, inboundId int, clients []model.Client, detachEmails []string, prune bool) error {
  93. if err := validateClientsRenewal(clients); err != nil {
  94. return err
  95. }
  96. if tx == nil {
  97. tx = database.GetDB()
  98. }
  99. if err := s.validateTuicIdentities(tx, inboundId, clients, detachEmails, prune); err != nil {
  100. return err
  101. }
  102. emails := make([]string, 0, len(clients))
  103. seen := make(map[string]struct{}, len(clients))
  104. for i := range clients {
  105. email := strings.TrimSpace(clients[i].Email)
  106. if email == "" {
  107. continue
  108. }
  109. if _, ok := seen[email]; ok {
  110. continue
  111. }
  112. seen[email] = struct{}{}
  113. emails = append(emails, email)
  114. }
  115. existing := make(map[string]*model.ClientRecord, len(emails))
  116. const selectChunk = 400
  117. for start := 0; start < len(emails); start += selectChunk {
  118. end := min(start+selectChunk, len(emails))
  119. var rows []model.ClientRecord
  120. if err := tx.Where("email IN ?", emails[start:end]).Find(&rows).Error; err != nil {
  121. return err
  122. }
  123. for i := range rows {
  124. r := rows[i]
  125. existing[r.Email] = &r
  126. }
  127. }
  128. idByEmail := make(map[string]int, len(emails))
  129. pending := make(map[string]*model.ClientRecord, len(emails))
  130. toCreate := make([]*model.ClientRecord, 0, len(emails))
  131. for i := range clients {
  132. email := strings.TrimSpace(clients[i].Email)
  133. if email == "" {
  134. continue
  135. }
  136. incoming := clients[i].ToRecord()
  137. // ToRecord copies the raw email; store the trimmed key this function
  138. // looks up by, or a padded email is inserted and never found again.
  139. incoming.Email = email
  140. row, ok := existing[email]
  141. if !ok {
  142. if _, dup := pending[email]; !dup {
  143. pending[email] = incoming
  144. toCreate = append(toCreate, incoming)
  145. }
  146. continue
  147. }
  148. before := *row
  149. applyClientRecordMerge(row, incoming)
  150. preservedUpdatedAt := max(incoming.UpdatedAt, row.UpdatedAt)
  151. row.UpdatedAt = preservedUpdatedAt
  152. idByEmail[email] = row.Id
  153. if *row == before {
  154. continue
  155. }
  156. if err := tx.Save(row).Error; err != nil {
  157. return err
  158. }
  159. if err := tx.Model(&model.ClientRecord{}).
  160. Where("id = ?", row.Id).
  161. UpdateColumn("updated_at", preservedUpdatedAt).Error; err != nil {
  162. return err
  163. }
  164. }
  165. if len(toCreate) > 0 {
  166. // Capture enable before Create: gorm default:true drops explicit false (#6478).
  167. // Restate disabled rows after CreateInBatches.
  168. wantEnable := make([]bool, len(toCreate))
  169. for i, rec := range toCreate {
  170. wantEnable[i] = rec.Enable
  171. }
  172. if err := tx.CreateInBatches(toCreate, 200).Error; err != nil {
  173. return err
  174. }
  175. disabledIDs := make([]int, 0)
  176. for i, rec := range toCreate {
  177. idByEmail[rec.Email] = rec.Id
  178. if !wantEnable[i] {
  179. disabledIDs = append(disabledIDs, rec.Id)
  180. }
  181. }
  182. for _, batch := range chunkInts(disabledIDs, sqlInChunk) {
  183. if err := tx.Model(&model.ClientRecord{}).Where("id IN ?", batch).
  184. UpdateColumn("enable", false).Error; err != nil {
  185. return err
  186. }
  187. }
  188. }
  189. wantedFlow := make(map[int]string, len(clients))
  190. wantedIds := make([]int, 0, len(clients))
  191. for i := range clients {
  192. email := strings.TrimSpace(clients[i].Email)
  193. if email == "" {
  194. continue
  195. }
  196. id, ok := idByEmail[email]
  197. if !ok {
  198. continue
  199. }
  200. if _, dup := wantedFlow[id]; dup {
  201. continue
  202. }
  203. wantedFlow[id] = clients[i].Flow
  204. wantedIds = append(wantedIds, id)
  205. }
  206. if err := s.reconcileInboundLinks(tx, inboundId, wantedFlow, wantedIds, detachEmails, prune); err != nil {
  207. return err
  208. }
  209. affected := map[int]struct{}{inboundId: {}}
  210. for _, ids := range chunkInts(wantedIds, sqlInChunk) {
  211. var inboundIDs []int
  212. if err := tx.Model(&model.ClientInbound{}).Distinct("inbound_id").Where("client_id IN ?", ids).Pluck("inbound_id", &inboundIDs).Error; err != nil {
  213. return err
  214. }
  215. for _, id := range inboundIDs {
  216. affected[id] = struct{}{}
  217. }
  218. }
  219. affectedIDs := make([]int, 0, len(affected))
  220. for id := range affected {
  221. affectedIDs = append(affectedIDs, id)
  222. }
  223. for _, ids := range chunkInts(affectedIDs, sqlInChunk) {
  224. var duplicates int64
  225. err := tx.Raw(`SELECT COUNT(*) FROM (
  226. SELECT ci.inbound_id, LOWER(c.uuid) FROM clients c
  227. JOIN client_inbounds ci ON ci.client_id = c.id
  228. JOIN inbounds i ON i.id = ci.inbound_id
  229. WHERE i.protocol = ? AND c.uuid <> '' AND ci.inbound_id IN ?
  230. GROUP BY ci.inbound_id, LOWER(c.uuid) HAVING COUNT(*) > 1
  231. ) AS duplicate_uuids`, model.TUIC, ids).Scan(&duplicates).Error
  232. if err != nil {
  233. return err
  234. }
  235. if duplicates > 0 {
  236. return fmt.Errorf("TUIC: duplicate client UUID within inbound")
  237. }
  238. }
  239. return nil
  240. }
  241. // reconcileInboundLinks writes only the client_inbounds rows that differ. prune
  242. // also removes links absent from wantedFlow, which only a full sync may do.
  243. func (s *ClientService) reconcileInboundLinks(tx *gorm.DB, inboundId int, wantedFlow map[int]string, wantedIds []int, detachEmails []string, prune bool) error {
  244. var current []model.ClientInbound
  245. if prune {
  246. if err := tx.Where("inbound_id = ?", inboundId).Find(&current).Error; err != nil {
  247. return err
  248. }
  249. } else {
  250. for _, batch := range chunkInts(wantedIds, sqlInChunk) {
  251. var rows []model.ClientInbound
  252. if err := tx.Where("inbound_id = ? AND client_id IN ?", inboundId, batch).Find(&rows).Error; err != nil {
  253. return err
  254. }
  255. current = append(current, rows...)
  256. }
  257. }
  258. var toDelete []int
  259. toUpdate := make(map[string][]int)
  260. have := make(map[int]struct{}, len(current))
  261. for _, link := range current {
  262. have[link.ClientId] = struct{}{}
  263. flow, keep := wantedFlow[link.ClientId]
  264. if !keep {
  265. if prune {
  266. toDelete = append(toDelete, link.ClientId)
  267. }
  268. continue
  269. }
  270. // Plain compare, not non-empty-wins: clearing a flow must persist "".
  271. if flow != link.FlowOverride {
  272. toUpdate[flow] = append(toUpdate[flow], link.ClientId)
  273. }
  274. }
  275. if len(detachEmails) > 0 {
  276. for _, batch := range chunkStrings(detachEmails, sqlInChunk) {
  277. var ids []int
  278. if err := tx.Model(&model.ClientRecord{}).Where("email IN ?", batch).Pluck("id", &ids).Error; err != nil {
  279. return err
  280. }
  281. for _, id := range ids {
  282. if _, keep := wantedFlow[id]; !keep {
  283. toDelete = append(toDelete, id)
  284. }
  285. }
  286. }
  287. }
  288. toInsert := make([]model.ClientInbound, 0, len(wantedIds))
  289. for _, id := range wantedIds {
  290. if _, exists := have[id]; exists {
  291. continue
  292. }
  293. toInsert = append(toInsert, model.ClientInbound{
  294. ClientId: id,
  295. InboundId: inboundId,
  296. FlowOverride: wantedFlow[id],
  297. })
  298. }
  299. for _, batch := range chunkInts(toDelete, sqlInChunk) {
  300. if err := tx.Where("inbound_id = ? AND client_id IN ?", inboundId, batch).
  301. Delete(&model.ClientInbound{}).Error; err != nil {
  302. return err
  303. }
  304. }
  305. for flow, ids := range toUpdate {
  306. for _, batch := range chunkInts(ids, sqlInChunk) {
  307. if err := tx.Model(&model.ClientInbound{}).
  308. Where("inbound_id = ? AND client_id IN ?", inboundId, batch).
  309. Update("flow_override", flow).Error; err != nil {
  310. return err
  311. }
  312. }
  313. }
  314. if len(toInsert) > 0 {
  315. // The delete this replaced also serialized concurrent syncs of one
  316. // inbound; without the clause a racing node poll aborts its whole tx.
  317. if err := tx.Clauses(clause.OnConflict{
  318. Columns: []clause.Column{{Name: "client_id"}, {Name: "inbound_id"}},
  319. DoUpdates: clause.AssignmentColumns([]string{"flow_override"}),
  320. }).CreateInBatches(toInsert, 200).Error; err != nil {
  321. return err
  322. }
  323. }
  324. return nil
  325. }
  326. func (s *ClientService) DetachInbound(tx *gorm.DB, inboundId int) error {
  327. if tx == nil {
  328. tx = database.GetDB()
  329. }
  330. return tx.Where("inbound_id = ?", inboundId).Delete(&model.ClientInbound{}).Error
  331. }
  332. func (s *ClientService) ListForInbound(tx *gorm.DB, inboundId int) ([]model.Client, error) {
  333. if tx == nil {
  334. tx = database.GetDB()
  335. }
  336. type joinedRow struct {
  337. model.ClientRecord
  338. FlowOverride string
  339. }
  340. var rows []joinedRow
  341. err := tx.Table("clients").
  342. Select("clients.*, client_inbounds.flow_override AS flow_override").
  343. Joins("JOIN client_inbounds ON client_inbounds.client_id = clients.id").
  344. Where("client_inbounds.inbound_id = ?", inboundId).
  345. Order("clients.id ASC").
  346. Find(&rows).Error
  347. if err != nil {
  348. return nil, err
  349. }
  350. out := make([]model.Client, 0, len(rows))
  351. for i := range rows {
  352. c := rows[i].ToClient()
  353. c.Flow = rows[i].FlowOverride
  354. out = append(out, *c)
  355. }
  356. return out, nil
  357. }
  358. // ListForInboundBySubId is ListForInbound narrowed to one subscription id —
  359. // both filter columns are indexed, so the subscription server resolves a
  360. // subscriber's clients without touching the inbound's settings JSON.
  361. func (s *ClientService) ListForInboundBySubId(tx *gorm.DB, inboundId int, subId string) ([]model.Client, error) {
  362. if tx == nil {
  363. tx = database.GetDB()
  364. }
  365. type joinedRow struct {
  366. model.ClientRecord
  367. FlowOverride string
  368. }
  369. var rows []joinedRow
  370. err := tx.Table("clients").
  371. Select("clients.*, client_inbounds.flow_override AS flow_override").
  372. Joins("JOIN client_inbounds ON client_inbounds.client_id = clients.id").
  373. Where("client_inbounds.inbound_id = ? AND clients.sub_id = ?", inboundId, subId).
  374. Order("clients.id ASC").
  375. Find(&rows).Error
  376. if err != nil {
  377. return nil, err
  378. }
  379. out := make([]model.Client, 0, len(rows))
  380. for i := range rows {
  381. c := rows[i].ToClient()
  382. c.Flow = rows[i].FlowOverride
  383. out = append(out, *c)
  384. }
  385. return out, nil
  386. }
  387. func (s *ClientService) validateTuicIdentities(tx *gorm.DB, inboundID int, changed []model.Client, detached []string, prune bool) error {
  388. var inbound model.Inbound
  389. if err := tx.Select("protocol").Where("id = ?", inboundID).Take(&inbound).Error; err != nil {
  390. return err
  391. }
  392. if inbound.Protocol != model.TUIC {
  393. return nil
  394. }
  395. candidates := append([]model.Client(nil), changed...)
  396. if !prune {
  397. current, err := s.ListForInbound(tx, inboundID)
  398. if err != nil {
  399. return err
  400. }
  401. excluded := make(map[string]bool)
  402. for _, client := range changed {
  403. excluded[client.Email] = true
  404. }
  405. for _, email := range detached {
  406. excluded[email] = true
  407. }
  408. for _, client := range current {
  409. if !excluded[client.Email] {
  410. candidates = append(candidates, client)
  411. }
  412. }
  413. }
  414. seen := make(map[uuid.UUID]string)
  415. for _, client := range candidates {
  416. if client.ID == "" {
  417. continue
  418. }
  419. id, err := uuid.Parse(client.ID)
  420. if err != nil {
  421. return fmt.Errorf("TUIC: invalid client UUID")
  422. }
  423. if email, exists := seen[id]; exists && email != client.Email {
  424. return fmt.Errorf("TUIC: duplicate client UUID within inbound")
  425. }
  426. seen[id] = client.Email
  427. }
  428. return nil
  429. }