package message import ( "context" "database/sql" "strconv" "git.asio.asia/nixevol/NixMsg/internal/app/port" "git.asio.asia/nixevol/NixMsg/internal/protocol" ) // Ack 处理确认(DEVELOPMENT 7.6)。 func (a *App) Ack(ctx context.Context, endpointID string, req *protocol.Ack) (AckResult, error) { if req == nil || req.ID == "" || req.From == "" { return AckResult{}, errCode(protocol.CodeBadRequest, "invalid ack") } nowMs := a.now().UnixMilli() var out AckResult var seq int64 var ackLatencySec float64 var observeAck bool var changed bool var wroteReceipt bool err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error { err := tx.QueryRow(`SELECT seq FROM messages WHERE sender_id = ? AND id = ?`, req.From, req.ID).Scan(&seq) if err == sql.ErrNoRows { return errCode(protocol.CodeNotFound, "message not found") } if err != nil { return err } var pushedAt sql.NullInt64 _ = tx.QueryRow(` SELECT pushed_at FROM deliveries WHERE seq = ? AND endpoint_id = ? AND state = 'pending'`, seq, endpointID).Scan(&pushedAt) res, err := tx.Exec(` UPDATE deliveries SET state = ?, reason = '', pushed_conn = NULL, updated_at = ? WHERE seq = ? AND endpoint_id = ? AND state = 'pending'`, DeliveryAccepted, nowMs, seq, endpointID) if err != nil { return err } aff, _ := res.RowsAffected() if aff > 0 { changed = true out.Result = DeliveryAccepted if pushedAt.Valid && pushedAt.Int64 > 0 && nowMs >= pushedAt.Int64 { ackLatencySec = float64(nowMs-pushedAt.Int64) / 1000.0 observeAck = true } wrote, e := insertReceiptTx(tx, req.From, seq, endpointID, DeliveryAccepted, "", nowMs) if e != nil { return e } wroteReceipt = wrote return TryFinalizeTx(tx, seq, nowMs, a.lim.RecordRetentionDays) } var state string err = tx.QueryRow(` SELECT state FROM deliveries WHERE seq = ? AND endpoint_id = ?`, seq, endpointID).Scan(&state) if err == sql.ErrNoRows { return errCode(protocol.CodeNotFound, "delivery not found") } if err != nil { return err } out.Result = state return nil }) if err != nil { return out, err } if observeAck && a.met != nil { a.met.AckSeconds.Observe(ackLatencySec) } if changed { a.WakePush(endpointID) } if wroteReceipt { a.WakePush(req.From) } return out, nil } // Recall 处理撤回。 func (a *App) Recall(ctx context.Context, senderID string, req *protocol.Recall) (protocol.RecallData, error) { if req == nil || req.ID == "" { return protocol.RecallData{}, errCode(protocol.CodeBadRequest, "invalid recall") } nowMs := a.now().UnixMilli() var data protocol.RecallData var revokes []revokeJob err := a.db.Queue.Do(ctx, func(tx *sql.Tx) error { var seq int64 var state string err := tx.QueryRow(`SELECT seq, state FROM messages WHERE sender_id = ? AND id = ?`, senderID, req.ID).Scan(&seq, &state) if err == sql.ErrNoRows { var one int e2 := tx.QueryRow(`SELECT 1 FROM send_keys WHERE sender_id = ? AND msg_id = ?`, senderID, req.ID).Scan(&one) if e2 == nil { data.Result = "failed" return nil } return errCode(protocol.CodeNotFound, "message not found") } if err != nil { return err } if state == StateScheduled { if _, e := tx.Exec(`UPDATE messages SET state = ?, reason = ?, completed_at = ? WHERE seq = ? AND state = 'scheduled'`, StateCompleted, ReasonRecalled, nowMs, seq); e != nil { return e } if _, e := tx.Exec(`DELETE FROM message_bodies WHERE seq = ?`, seq); e != nil { return e } if a.lim.RecordRetentionDays == 0 { if _, e := tx.Exec(`DELETE FROM messages WHERE seq = ?`, seq); e != nil { return e } } data.Result = "recalled" return nil } rows, err := tx.Query(` SELECT endpoint_id, pushed_conn FROM deliveries WHERE seq = ? AND state = 'pending'`, seq) if err != nil { return err } type pend struct { ep string pushed sql.NullString } var pending []pend for rows.Next() { var p pend if err := rows.Scan(&p.ep, &p.pushed); err != nil { _ = rows.Close() return err } pending = append(pending, p) } _ = rows.Close() recalled := 0 for _, p := range pending { res, err := tx.Exec(` UPDATE deliveries SET state = ?, reason = ?, pushed_conn = NULL, updated_at = ? WHERE seq = ? AND endpoint_id = ? AND state = 'pending'`, DeliveryRecalled, ReasonRecalled, nowMs, seq, p.ep) if err != nil { return err } aff, _ := res.RowsAffected() if aff == 0 { continue } recalled++ if p.pushed.Valid && p.pushed.String != "" { revokes = append(revokes, revokeJob{ endpointID: p.ep, connID: port.ConnID(p.pushed.String), msgID: req.ID, from: senderID, reason: ReasonRecalled, }) } } var accepted, other int _ = tx.QueryRow(`SELECT COUNT(*) FROM deliveries WHERE seq = ? AND state = 'accepted'`, seq).Scan(&accepted) _ = tx.QueryRow(` SELECT COUNT(*) FROM deliveries WHERE seq = ? AND state IN ('expired','dropped','rejected')`, seq).Scan(&other) data.Recalled = recalled data.Accepted = accepted data.Other = other switch { case recalled > 0 && accepted == 0: data.Result = "recalled" case recalled > 0 && accepted > 0: data.Result = "partial" default: data.Result = "failed" } return TryFinalizeTx(tx, seq, nowMs, a.lim.RecordRetentionDays) }) if err != nil { return data, err } a.mu.Lock() a.pendingRevoke = append(a.pendingRevoke, revokes...) a.mu.Unlock() a.flushRevokes(ctx) return data, nil } // Status 查询自己发出的消息状态。 func (a *App) Status(ctx context.Context, senderID string, req *protocol.Status) (any, error) { if req == nil || req.ID == "" { return nil, errCode(protocol.CodeBadRequest, "invalid status") } limit := req.Limit if limit <= 0 { limit = 100 } if limit > 200 { limit = 200 } var seq int64 var state, reason string err := a.db.Read.QueryRowContext(ctx, ` SELECT seq, state, reason FROM messages WHERE sender_id = ? AND id = ?`, senderID, req.ID).Scan(&seq, &state, &reason) if err == sql.ErrNoRows { return nil, errCode(protocol.CodeNotFound, "message not found") } if err != nil { return nil, err } type counts struct { Pending int `json:"pending"` Accepted int `json:"accepted"` Recalled int `json:"recalled"` Expired int `json:"expired"` Dropped int `json:"dropped"` Rejected int `json:"rejected"` } var c counts rows, err := a.db.Read.QueryContext(ctx, ` SELECT state, COUNT(*) FROM deliveries WHERE seq = ? GROUP BY state`, seq) if err != nil { return nil, err } for rows.Next() { var st string var n int if scanErr := rows.Scan(&st, &n); scanErr != nil { _ = rows.Close() return nil, scanErr } switch st { case DeliveryPending: c.Pending = n case DeliveryAccepted: c.Accepted = n case DeliveryRecalled: c.Recalled = n case DeliveryExpired: c.Expired = n case DeliveryDropped: c.Dropped = n case DeliveryRejected: c.Rejected = n } } _ = rows.Close() q := `SELECT endpoint_id, state, reason FROM deliveries WHERE seq = ?` args := []any{seq} if req.Cursor != "" { q += ` AND endpoint_id > ?` args = append(args, req.Cursor) } q += ` ORDER BY endpoint_id ASC LIMIT ?` args = append(args, limit) drows, err := a.db.Read.QueryContext(ctx, q, args...) if err != nil { return nil, err } defer func() { _ = drows.Close() }() type item struct { EndpointID string `json:"endpoint_id"` State string `json:"state"` Reason string `json:"reason"` } items := make([]item, 0) var nextCursor string for drows.Next() { var it item if err := drows.Scan(&it.EndpointID, &it.State, &it.Reason); err != nil { return nil, err } items = append(items, it) nextCursor = it.EndpointID } return map[string]any{ "id": req.ID, "state": state, "reason": reason, "counts": c, "deliveries": items, "next_cursor": nextCursor, }, nil } // ReceiptAck 确认回执已收下。 func (a *App) ReceiptAck(ctx context.Context, endpointID string, req *protocol.ReceiptAck) error { if req == nil || req.ReceiptID == "" { return errCode(protocol.CodeBadRequest, "invalid receipt_ack") } rid, err := strconv.ParseInt(req.ReceiptID, 10, 64) if err != nil { return errCode(protocol.CodeBadRequest, "invalid receipt_id") } var acked bool err = a.db.Queue.Do(ctx, func(tx *sql.Tx) error { res, err := tx.Exec(`UPDATE receipts SET acked = 1 WHERE receipt_id = ? AND sender_id = ? AND acked = 0`, rid, endpointID) if err != nil { return err } aff, _ := res.RowsAffected() acked = aff > 0 return nil }) if err != nil { return err } if acked { live, ok := a.lookupConn(endpointID) if ok { a.unmarkReceipt(string(live.ConnID), rid) } a.WakePush(endpointID) } return nil }