sms.go 17 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476477478479480481482483484485486487488489490491492493494495496497498499500501502503504505506507508509510511512513514515516517518519520521522523524525526527528529530531532533534535536537538539540541542543544545546547548549550551552553554555556557558559560561562563564565566567568569570571572573574575576577578579580581582583584585586587588589590591
  1. package store
  2. import (
  3. "context"
  4. "database/sql"
  5. "encoding/json"
  6. "errors"
  7. "fmt"
  8. "strconv"
  9. "strings"
  10. "time"
  11. )
  12. type contextQueryExecer interface {
  13. contextExecer
  14. QueryRowContext(context.Context, string, ...any) *sql.Row
  15. }
  16. // SaveSMSMessage inserts a new message or updates an existing record. A
  17. // non-empty (device_id, message_id) pair is idempotent for modem retries.
  18. func (s *Store) SaveSMSMessage(ctx context.Context, value SMSMessage) (SMSMessage, error) {
  19. tx, err := s.db.BeginTx(ctx, nil)
  20. if err != nil {
  21. return SMSMessage{}, fmt.Errorf("begin SMS update: %w", err)
  22. }
  23. defer tx.Rollback()
  24. saved, err := saveSMSMessage(ctx, tx, value)
  25. if err != nil {
  26. return SMSMessage{}, err
  27. }
  28. if err := tx.Commit(); err != nil {
  29. return SMSMessage{}, fmt.Errorf("commit SMS update: %w", err)
  30. }
  31. return saved, nil
  32. }
  33. func saveSMSMessage(
  34. ctx context.Context,
  35. executor contextQueryExecer,
  36. value SMSMessage,
  37. ) (SMSMessage, error) {
  38. value.DeviceID = strings.TrimSpace(value.DeviceID)
  39. value.Peer = strings.TrimSpace(value.Peer)
  40. value.Direction = strings.ToLower(strings.TrimSpace(value.Direction))
  41. if value.DeviceID == "" {
  42. return SMSMessage{}, errors.New("SMS device id is required")
  43. }
  44. if value.Peer == "" {
  45. return SMSMessage{}, errors.New("SMS peer is required")
  46. }
  47. switch value.Direction {
  48. case "inbound", "outbound", "received", "sent":
  49. default:
  50. return SMSMessage{}, fmt.Errorf("unsupported SMS direction %q", value.Direction)
  51. }
  52. if value.PartsTotal == 0 {
  53. value.PartsTotal = 1
  54. }
  55. if value.PartsTotal < 1 {
  56. return SMSMessage{}, errors.New("SMS parts total must be positive")
  57. }
  58. extra, err := normalizeJSONObject(value.Extra)
  59. if err != nil {
  60. return SMSMessage{}, fmt.Errorf("normalize SMS extra data: %w", err)
  61. }
  62. now := time.Now().UTC()
  63. if value.Timestamp.IsZero() {
  64. value.Timestamp = now
  65. }
  66. if value.CreatedAt.IsZero() {
  67. value.CreatedAt = now
  68. }
  69. if value.UpdatedAt.IsZero() {
  70. value.UpdatedAt = now
  71. }
  72. if value.ID > 0 {
  73. result, err := executor.ExecContext(ctx, `
  74. UPDATE sms_messages SET
  75. message_id = ?, device_id = ?, imsi = ?, peer = ?,
  76. direction = ?, body = ?, message_time = ?, status = ?,
  77. source = ?, parts_total = ?, delivery_state = ?, is_read = ?,
  78. extra_json = ?, updated_at = ?
  79. WHERE id = ?
  80. `,
  81. value.MessageID, value.DeviceID, value.IMSI, value.Peer,
  82. value.Direction, value.Body, value.Timestamp.Unix(), value.Status,
  83. value.Source, value.PartsTotal, value.DeliveryState,
  84. boolInt(value.Read), string(extra), value.UpdatedAt.Unix(), value.ID,
  85. )
  86. if err != nil {
  87. return SMSMessage{}, fmt.Errorf("update SMS %d: %w", value.ID, err)
  88. }
  89. if err := requireAffected(result); err != nil {
  90. return SMSMessage{}, err
  91. }
  92. return scanSMSMessage(executor.QueryRowContext(ctx, smsMessageSelect+` WHERE id = ?`, value.ID))
  93. }
  94. result, err := executor.ExecContext(ctx, `
  95. INSERT INTO sms_messages (
  96. message_id, device_id, imsi, peer, direction, body, message_time,
  97. status, source, parts_total, delivery_state, is_read, extra_json,
  98. created_at, updated_at
  99. ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
  100. ON CONFLICT(device_id, message_id) WHERE message_id <> '' DO UPDATE SET
  101. imsi = excluded.imsi,
  102. peer = excluded.peer,
  103. direction = excluded.direction,
  104. body = excluded.body,
  105. message_time = MIN(sms_messages.message_time, excluded.message_time),
  106. status = excluded.status,
  107. source = excluded.source,
  108. parts_total = excluded.parts_total,
  109. delivery_state = excluded.delivery_state,
  110. is_read = excluded.is_read,
  111. extra_json = excluded.extra_json,
  112. updated_at = excluded.updated_at
  113. `,
  114. value.MessageID, value.DeviceID, value.IMSI, value.Peer,
  115. value.Direction, value.Body, value.Timestamp.Unix(), value.Status,
  116. value.Source, value.PartsTotal, value.DeliveryState,
  117. boolInt(value.Read), string(extra), value.CreatedAt.Unix(),
  118. value.UpdatedAt.Unix(),
  119. )
  120. if err != nil {
  121. return SMSMessage{}, fmt.Errorf("save SMS: %w", err)
  122. }
  123. if value.MessageID != "" {
  124. return scanSMSMessage(executor.QueryRowContext(
  125. ctx,
  126. smsMessageSelect+` WHERE device_id = ? AND message_id = ?`,
  127. value.DeviceID,
  128. value.MessageID,
  129. ))
  130. }
  131. id, err := result.LastInsertId()
  132. if err != nil {
  133. return SMSMessage{}, fmt.Errorf("read inserted SMS id: %w", err)
  134. }
  135. return scanSMSMessage(executor.QueryRowContext(ctx, smsMessageSelect+` WHERE id = ?`, id))
  136. }
  137. func (s *Store) SaveSMSMessages(ctx context.Context, values []SMSMessage) error {
  138. tx, err := s.db.BeginTx(ctx, nil)
  139. if err != nil {
  140. return fmt.Errorf("begin SMS batch: %w", err)
  141. }
  142. defer tx.Rollback()
  143. for index, value := range values {
  144. if _, err := saveSMSMessage(ctx, tx, value); err != nil {
  145. return fmt.Errorf("save SMS batch item %d: %w", index, err)
  146. }
  147. }
  148. if err := tx.Commit(); err != nil {
  149. return fmt.Errorf("commit SMS batch: %w", err)
  150. }
  151. return nil
  152. }
  153. func (s *Store) SMSMessage(ctx context.Context, id int64) (SMSMessage, error) {
  154. return scanSMSMessage(s.db.QueryRowContext(ctx, smsMessageSelect+` WHERE id = ?`, id))
  155. }
  156. // LatestSMSMessageID returns the current durable cursor used by notification
  157. // consumers. Starting at this value avoids replaying the entire SMS archive
  158. // whenever the service or a notification provider is restarted.
  159. func (s *Store) LatestSMSMessageID(ctx context.Context) (int64, error) {
  160. var id int64
  161. if err := s.db.QueryRowContext(ctx, `SELECT COALESCE(MAX(id), 0) FROM sms_messages`).Scan(&id); err != nil {
  162. return 0, fmt.Errorf("read latest SMS id: %w", err)
  163. }
  164. return id, nil
  165. }
  166. // ListInboundSMSAfterID returns newly inserted inbound messages in durable ID
  167. // order. Telegram advances this cursor only after considering each item, so
  168. // timestamp corrections and duplicate modem synchronisations cannot reorder or
  169. // duplicate notifications.
  170. func (s *Store) ListInboundSMSAfterID(ctx context.Context, afterID int64, limit int) ([]SMSMessage, error) {
  171. if afterID < 0 {
  172. afterID = 0
  173. }
  174. rows, err := s.db.QueryContext(ctx, smsMessageSelect+`
  175. WHERE id > ? AND direction IN ('inbound', 'received')
  176. ORDER BY id ASC
  177. LIMIT ?`, afterID, normalizedLimit(limit))
  178. if err != nil {
  179. return nil, fmt.Errorf("list new inbound SMS messages: %w", err)
  180. }
  181. defer rows.Close()
  182. values := make([]SMSMessage, 0)
  183. for rows.Next() {
  184. value, scanErr := scanSMSMessage(rows)
  185. if scanErr != nil {
  186. return nil, fmt.Errorf("scan new inbound SMS message: %w", scanErr)
  187. }
  188. values = append(values, value)
  189. }
  190. if err := rows.Err(); err != nil {
  191. return nil, fmt.Errorf("iterate new inbound SMS messages: %w", err)
  192. }
  193. return values, nil
  194. }
  195. // ApplySMSDeliveryReport attaches a TP-STATUS report to the newest matching
  196. // outbound submission and advances its aggregate delivery state. Multipart
  197. // messages become delivered only after every submitted part is reported.
  198. func (s *Store) ApplySMSDeliveryReport(ctx context.Context, report SMSDeliveryReport) (SMSMessage, error) {
  199. if report.DeviceID == "" || report.MessageReference < 0 || report.MessageReference > 255 {
  200. return SMSMessage{}, errors.New("invalid SMS delivery report identity")
  201. }
  202. if report.ReceivedAt.IsZero() {
  203. report.ReceivedAt = time.Now().UTC()
  204. }
  205. tx, err := s.db.BeginTx(ctx, nil)
  206. if err != nil {
  207. return SMSMessage{}, fmt.Errorf("begin SMS delivery report: %w", err)
  208. }
  209. defer tx.Rollback()
  210. query := smsMessageSelect + `
  211. WHERE device_id = ?
  212. AND direction IN ('outbound', 'sent')
  213. AND (? = '' OR imsi = ?)
  214. AND (? = '' OR peer = ?)
  215. AND (? = '' OR source = ?)
  216. ORDER BY created_at DESC, id DESC
  217. LIMIT 256`
  218. rows, err := tx.QueryContext(
  219. ctx,
  220. query,
  221. report.DeviceID,
  222. report.IMSI, report.IMSI,
  223. report.Peer, report.Peer,
  224. report.Source, report.Source,
  225. )
  226. if err != nil {
  227. return SMSMessage{}, fmt.Errorf("find SMS delivery target: %w", err)
  228. }
  229. var target SMSMessage
  230. var targetExtra map[string]any
  231. for rows.Next() {
  232. candidate, scanErr := scanSMSMessage(rows)
  233. if scanErr != nil {
  234. _ = rows.Close()
  235. return SMSMessage{}, scanErr
  236. }
  237. extra := make(map[string]any)
  238. if json.Unmarshal(candidate.Extra, &extra) != nil || !smsExtraHasReference(extra, report.MessageReference) {
  239. continue
  240. }
  241. target, targetExtra = candidate, extra
  242. break
  243. }
  244. if err := rows.Close(); err != nil {
  245. return SMSMessage{}, err
  246. }
  247. if target.ID == 0 {
  248. return SMSMessage{}, ErrNotFound
  249. }
  250. reports, _ := targetExtra["delivery_reports"].(map[string]any)
  251. if reports == nil {
  252. reports = make(map[string]any)
  253. }
  254. reportValue := map[string]any{
  255. "status_code": report.StatusCode,
  256. "delivery_state": report.DeliveryState,
  257. "received_at": report.ReceivedAt.UTC(),
  258. }
  259. if report.ServiceCenterTime != nil {
  260. reportValue["service_center_timestamp"] = report.ServiceCenterTime.UTC()
  261. }
  262. if report.DischargeTime != nil {
  263. reportValue["discharge_timestamp"] = report.DischargeTime.UTC()
  264. }
  265. reports[strconv.Itoa(report.MessageReference)] = reportValue
  266. targetExtra["delivery_reports"] = reports
  267. target.DeliveryState = aggregateSMSDeliveryState(targetExtra, reports)
  268. target.Extra, err = json.Marshal(targetExtra)
  269. if err != nil {
  270. return SMSMessage{}, fmt.Errorf("encode SMS delivery reports: %w", err)
  271. }
  272. target.UpdatedAt = time.Now().UTC()
  273. saved, err := saveSMSMessage(ctx, tx, target)
  274. if err != nil {
  275. return SMSMessage{}, err
  276. }
  277. if err := tx.Commit(); err != nil {
  278. return SMSMessage{}, fmt.Errorf("commit SMS delivery report: %w", err)
  279. }
  280. return saved, nil
  281. }
  282. func smsExtraHasReference(extra map[string]any, reference int) bool {
  283. if numberAsInt(extra["message_reference"]) == reference {
  284. return true
  285. }
  286. parts, _ := extra["part_results"].([]any)
  287. for _, value := range parts {
  288. part, _ := value.(map[string]any)
  289. if numberAsInt(part["reference"]) == reference ||
  290. numberAsInt(part["messageReference"]) == reference ||
  291. numberAsInt(part["message_reference"]) == reference {
  292. return true
  293. }
  294. }
  295. return false
  296. }
  297. func aggregateSMSDeliveryState(extra map[string]any, reports map[string]any) string {
  298. parts, _ := extra["part_results"].([]any)
  299. references := make([]int, 0, len(parts))
  300. for _, value := range parts {
  301. part, _ := value.(map[string]any)
  302. reference := numberAsInt(part["reference"])
  303. if reference < 0 {
  304. reference = numberAsInt(part["messageReference"])
  305. }
  306. if reference < 0 {
  307. reference = numberAsInt(part["message_reference"])
  308. }
  309. if reference >= 0 {
  310. references = append(references, reference)
  311. }
  312. }
  313. if len(references) == 0 {
  314. if reference := numberAsInt(extra["message_reference"]); reference >= 0 {
  315. references = append(references, reference)
  316. }
  317. }
  318. if len(references) == 0 {
  319. return "unknown"
  320. }
  321. delivered := 0
  322. for _, reference := range references {
  323. value, found := reports[strconv.Itoa(reference)]
  324. if !found {
  325. continue
  326. }
  327. report, _ := value.(map[string]any)
  328. state, _ := report["delivery_state"].(string)
  329. switch state {
  330. case "delivered":
  331. delivered++
  332. case "permanent_error", "failed", "rejected":
  333. return "failed"
  334. }
  335. }
  336. if delivered == len(references) {
  337. return "delivered"
  338. }
  339. return "pending_delivery_report"
  340. }
  341. func numberAsInt(value any) int {
  342. switch number := value.(type) {
  343. case float64:
  344. return int(number)
  345. case int:
  346. return number
  347. case json.Number:
  348. parsed, err := strconv.Atoi(string(number))
  349. if err == nil {
  350. return parsed
  351. }
  352. }
  353. return -1
  354. }
  355. func (s *Store) ListSMSMessages(ctx context.Context, filter SMSFilter) ([]SMSMessage, error) {
  356. where, args := smsWhere(filter, "")
  357. query := smsMessageSelect + where + ` ORDER BY message_time DESC, id DESC LIMIT ?`
  358. args = append(args, normalizedLimit(filter.Limit))
  359. rows, err := s.db.QueryContext(ctx, query, args...)
  360. if err != nil {
  361. return nil, fmt.Errorf("list SMS messages: %w", err)
  362. }
  363. defer rows.Close()
  364. values := make([]SMSMessage, 0)
  365. for rows.Next() {
  366. value, err := scanSMSMessage(rows)
  367. if err != nil {
  368. return nil, fmt.Errorf("scan SMS message: %w", err)
  369. }
  370. values = append(values, value)
  371. }
  372. if err := rows.Err(); err != nil {
  373. return nil, fmt.Errorf("iterate SMS messages: %w", err)
  374. }
  375. return values, nil
  376. }
  377. func (s *Store) DeleteSMSMessage(ctx context.Context, id int64) error {
  378. result, err := s.db.ExecContext(ctx, `DELETE FROM sms_messages WHERE id = ?`, id)
  379. if err != nil {
  380. return fmt.Errorf("delete SMS %d: %w", id, err)
  381. }
  382. return requireAffected(result)
  383. }
  384. func (s *Store) DeleteSMSThread(
  385. ctx context.Context,
  386. deviceID string,
  387. imsi string,
  388. peer string,
  389. ) (int64, error) {
  390. result, err := s.db.ExecContext(ctx, `
  391. DELETE FROM sms_messages
  392. WHERE device_id = ? AND imsi = ? AND peer = ?
  393. `, deviceID, imsi, peer)
  394. if err != nil {
  395. return 0, fmt.Errorf("delete SMS thread: %w", err)
  396. }
  397. affected, err := result.RowsAffected()
  398. if err != nil {
  399. return 0, fmt.Errorf("read deleted SMS count: %w", err)
  400. }
  401. if affected == 0 {
  402. return 0, ErrNotFound
  403. }
  404. return affected, nil
  405. }
  406. func (s *Store) MarkSMSThreadRead(
  407. ctx context.Context,
  408. deviceID string,
  409. imsi string,
  410. peer string,
  411. ) (int64, error) {
  412. result, err := s.db.ExecContext(ctx, `
  413. UPDATE sms_messages
  414. SET is_read = 1, updated_at = ?
  415. WHERE device_id = ? AND imsi = ? AND peer = ?
  416. AND direction IN ('inbound', 'received') AND is_read = 0
  417. `, time.Now().UTC().Unix(), deviceID, imsi, peer)
  418. if err != nil {
  419. return 0, fmt.Errorf("mark SMS thread read: %w", err)
  420. }
  421. affected, err := result.RowsAffected()
  422. if err != nil {
  423. return 0, fmt.Errorf("read marked SMS count: %w", err)
  424. }
  425. return affected, nil
  426. }
  427. // ListSMSContacts derives contacts and thread counters from messages. No
  428. // duplicated contact/thread table can drift out of sync with message history.
  429. func (s *Store) ListSMSContacts(ctx context.Context, filter SMSFilter) ([]SMSContact, error) {
  430. where, args := smsWhere(filter, "m.")
  431. query := `
  432. WITH ranked AS (
  433. SELECT
  434. m.id, m.device_id, m.imsi, m.peer, m.body, m.message_time,
  435. m.direction,
  436. ROW_NUMBER() OVER (
  437. PARTITION BY m.device_id, m.imsi, m.peer
  438. ORDER BY m.message_time DESC, m.id DESC
  439. ) AS row_number,
  440. SUM(CASE
  441. WHEN m.direction IN ('inbound', 'received') AND m.is_read = 0
  442. THEN 1 ELSE 0
  443. END) OVER (
  444. PARTITION BY m.device_id, m.imsi, m.peer
  445. ) AS unread_count,
  446. COUNT(*) OVER (
  447. PARTITION BY m.device_id, m.imsi, m.peer
  448. ) AS message_count
  449. FROM sms_messages m` + where + `
  450. )
  451. SELECT
  452. r.device_id,
  453. COALESCE(d.name, ''),
  454. r.imsi,
  455. COALESCE(NULLIF(dr.phone_number, ''), NULLIF(vr.local_phone, ''), ''),
  456. r.peer,
  457. r.peer,
  458. r.body,
  459. r.message_time,
  460. r.direction,
  461. r.id,
  462. r.unread_count,
  463. r.message_count
  464. FROM ranked r
  465. LEFT JOIN devices d ON d.id = r.device_id
  466. LEFT JOIN device_runtime dr ON dr.device_id = r.device_id
  467. LEFT JOIN vowifi_runtime vr ON vr.device_id = r.device_id
  468. WHERE r.row_number = 1
  469. ORDER BY r.message_time DESC, r.id DESC
  470. LIMIT ?`
  471. args = append(args, normalizedLimit(filter.Limit))
  472. rows, err := s.db.QueryContext(ctx, query, args...)
  473. if err != nil {
  474. return nil, fmt.Errorf("list SMS contacts: %w", err)
  475. }
  476. defer rows.Close()
  477. values := make([]SMSContact, 0)
  478. for rows.Next() {
  479. var value SMSContact
  480. var timestamp int64
  481. if err := rows.Scan(
  482. &value.DeviceID, &value.DeviceName, &value.IMSI,
  483. &value.LocalPhone, &value.Peer, &value.DisplayName,
  484. &value.LastMessage, &timestamp, &value.Direction,
  485. &value.LastSMSID, &value.UnreadCount, &value.MessageCount,
  486. ); err != nil {
  487. return nil, fmt.Errorf("scan SMS contact: %w", err)
  488. }
  489. value.LastTimestamp = time.Unix(timestamp, 0).UTC()
  490. values = append(values, value)
  491. }
  492. if err := rows.Err(); err != nil {
  493. return nil, fmt.Errorf("iterate SMS contacts: %w", err)
  494. }
  495. return values, nil
  496. }
  497. const smsMessageSelect = `
  498. SELECT id, message_id, device_id, imsi, peer, direction, body,
  499. message_time, status, source, parts_total, delivery_state, is_read,
  500. extra_json, created_at, updated_at
  501. FROM sms_messages`
  502. func scanSMSMessage(row rowScanner) (SMSMessage, error) {
  503. var value SMSMessage
  504. var messageTime, createdAt, updatedAt int64
  505. var read int
  506. var extra string
  507. err := row.Scan(
  508. &value.ID, &value.MessageID, &value.DeviceID, &value.IMSI,
  509. &value.Peer, &value.Direction, &value.Body, &messageTime,
  510. &value.Status, &value.Source, &value.PartsTotal,
  511. &value.DeliveryState, &read, &extra, &createdAt, &updatedAt,
  512. )
  513. if errors.Is(err, sql.ErrNoRows) {
  514. return SMSMessage{}, ErrNotFound
  515. }
  516. if err != nil {
  517. return SMSMessage{}, err
  518. }
  519. value.Read = read != 0
  520. value.Extra = []byte(extra)
  521. value.Timestamp = time.Unix(messageTime, 0).UTC()
  522. value.CreatedAt = time.Unix(createdAt, 0).UTC()
  523. value.UpdatedAt = time.Unix(updatedAt, 0).UTC()
  524. return value, nil
  525. }
  526. func smsWhere(filter SMSFilter, prefix string) (string, []any) {
  527. clauses := make([]string, 0, 6)
  528. args := make([]any, 0, 6)
  529. if filter.DeviceID != "" {
  530. clauses = append(clauses, prefix+`device_id = ?`)
  531. args = append(args, filter.DeviceID)
  532. }
  533. if filter.IMSI != "" {
  534. clauses = append(clauses, prefix+`imsi = ?`)
  535. args = append(args, filter.IMSI)
  536. }
  537. if filter.Peer != "" {
  538. clauses = append(clauses, prefix+`peer = ?`)
  539. args = append(args, filter.Peer)
  540. }
  541. if !filter.Since.IsZero() {
  542. clauses = append(clauses, prefix+`message_time >= ?`)
  543. args = append(args, filter.Since.UTC().Unix())
  544. }
  545. if !filter.Until.IsZero() {
  546. clauses = append(clauses, prefix+`message_time < ?`)
  547. args = append(args, filter.Until.UTC().Unix())
  548. }
  549. if filter.BeforeID > 0 {
  550. clauses = append(clauses, prefix+`id < ?`)
  551. args = append(args, filter.BeforeID)
  552. }
  553. if len(clauses) == 0 {
  554. return "", args
  555. }
  556. return " WHERE " + strings.Join(clauses, " AND "), args
  557. }
  558. func normalizedLimit(value int) int {
  559. if value <= 0 {
  560. return 100
  561. }
  562. if value > 1000 {
  563. return 1000
  564. }
  565. return value
  566. }