feat: harden sessions and track downstream acknowledgements

This commit is contained in:
hectorzhao
2026-07-14 14:18:43 +08:00
parent 3d37adcc9f
commit 8c03663f24
43 changed files with 1733 additions and 150 deletions
+4 -4
View File
@@ -139,13 +139,13 @@ func (s Server) handleDownstreamReceipt(w http.ResponseWriter, r *http.Request)
http.Error(w, fmt.Sprintf("invalid downstream receipt: %v", err), http.StatusBadRequest)
return
}
delivered, err := inbound.PushReceipt(event)
result, err := inbound.PushReceiptWithResult(event)
if err != nil {
http.Error(w, err.Error(), http.StatusBadRequest)
return
}
w.Header().Set("Content-Type", "application/json")
_ = json.NewEncoder(w).Encode(map[string]any{"delivered": delivered})
_ = json.NewEncoder(w).Encode(result)
}
func (s Server) handleDownstreamUplink(w http.ResponseWriter, r *http.Request) {
@@ -158,13 +158,13 @@ func (s Server) handleDownstreamUplink(w http.ResponseWriter, r *http.Request) {
http.Error(w, fmt.Sprintf("invalid downstream uplink: %v", err), http.StatusBadRequest)
return
}
delivered, err := inbound.PushUplink(event)
result, err := inbound.PushUplinkWithResult(event)
if err != nil {
http.Error(w, err.Error(), http.StatusBadRequest)
return
}
w.Header().Set("Content-Type", "application/json")
_ = json.NewEncoder(w).Encode(map[string]any{"delivered": delivered})
_ = json.NewEncoder(w).Encode(result)
}
func (s Server) handleDownstreamRecoveryCandidates(w http.ResponseWriter, r *http.Request) {
+227 -31
View File
@@ -11,6 +11,8 @@ import (
"log"
"net"
"net/http"
"os"
"strconv"
"strings"
"sync"
"time"
@@ -20,6 +22,7 @@ import (
)
const defaultHTTPTimeout = 10 * time.Second
const defaultDownstreamAckTimeout = 30 * time.Second
type Server struct {
Addr string
@@ -86,15 +89,46 @@ type DownstreamUplink struct {
ReceivedAt string `json:"receivedAt,omitempty"`
}
type DownstreamSendResult struct {
Sent bool `json:"sent"`
ConnectionID string `json:"connectionId,omitempty"`
SequenceID string `json:"sequenceId,omitempty"`
MessageID string `json:"messageId,omitempty"`
SentAt string `json:"sentAt,omitempty"`
AckDeadlineAt string `json:"ackDeadlineAt,omitempty"`
}
type downstreamDeliveryLifecycleEvent struct {
Kind string
DeliveryID string
ConnectionID string
SequenceID uint32
MessageID uint64
Result uint32
ObservedAt time.Time
AckDeadlineAt time.Time
FailureType string
ErrorMessage string
}
type downstreamAckTracker struct {
deliveryID string
connectionID string
sequenceID uint32
messageID uint64
session *downstreamSession
timer *time.Timer
}
type downstreamConnectionEvent struct {
Account string `json:"account"`
ConnectionID string `json:"connectionId"`
Status string `json:"status"`
RemoteIP string `json:"remoteIp,omitempty"`
Protocol string `json:"protocol,omitempty"`
ConnectedAt string `json:"connectedAt,omitempty"`
ObservedAt string `json:"observedAt,omitempty"`
ErrorMessage string `json:"errorMessage,omitempty"`
Account string `json:"account"`
ConnectionID string `json:"connectionId"`
Status string `json:"status"`
RemoteIP string `json:"remoteIp,omitempty"`
Protocol string `json:"protocol,omitempty"`
ConnectedAt string `json:"connectedAt,omitempty"`
ObservedAt string `json:"observedAt,omitempty"`
ErrorMessage string `json:"errorMessage,omitempty"`
}
type downstreamSession struct {
@@ -113,6 +147,7 @@ type downstreamSession struct {
presence PresenceStore
instanceID string
report func(*downstreamSession, string, string)
deliveryReport func(downstreamDeliveryLifecycleEvent)
}
var downstreamRegistry = struct {
@@ -126,6 +161,11 @@ var downstreamRegistry = struct {
byConn: make(map[*cmpp.Conn]*downstreamSession),
}
var downstreamAckRegistry = struct {
sync.Mutex
items map[string]*downstreamAckTracker
}{items: make(map[string]*downstreamAckTracker)}
func (s Server) ListenAndServe() error {
addr := s.Addr
if addr == "" {
@@ -176,6 +216,7 @@ func (s Server) handleLogin(response *cmpp.Response, packet *cmpp.Packet, logger
presence: s.PresenceStore,
instanceID: s.gatewayInstanceID(),
report: s.reportConnection,
deliveryReport: s.reportDownstreamDelivery,
}
rememberAccount(session)
go s.reportConnection(&session, "connected", "")
@@ -273,6 +314,7 @@ func (s Server) handleSubmit(response *cmpp.Response, packet *cmpp.Packet, logge
presence: s.PresenceStore,
instanceID: s.gatewayInstanceID(),
report: session.report,
deliveryReport: session.deliveryReport,
})
if current := findSessionByConn(packet.Conn); current != nil && current.report != nil {
go current.report(current, "submit", "")
@@ -289,14 +331,20 @@ func (s Server) handleSubmit(response *cmpp.Response, packet *cmpp.Packet, logge
return false, nil
}
func (s Server) handleActivity(_ *cmpp.Response, packet *cmpp.Packet, _ *log.Logger) (bool, error) {
func (s Server) handleActivity(_ *cmpp.Response, packet *cmpp.Packet, logger *log.Logger) (bool, error) {
session := findSessionByConn(packet.Conn)
if session == nil || session.report == nil {
if session == nil {
return true, nil
}
switch packet.Packer.(type) {
switch response := packet.Packer.(type) {
case *cmpp.CmppActiveTestReqPkt, *cmpp.CmppActiveTestRspPkt:
go session.report(session, "heartbeat", "")
if session.report != nil {
go session.report(session, "heartbeat", "")
}
case *cmpp.Cmpp2DeliverRspPkt:
handleDownstreamAcknowledgement(packet.Conn, response.SeqId, response.MsgId, uint32(response.Result), logger)
case *cmpp.Cmpp3DeliverRspPkt:
handleDownstreamAcknowledgement(packet.Conn, response.SeqId, response.MsgId, response.Result, logger)
}
return true, nil
}
@@ -316,6 +364,31 @@ func (s Server) reportConnection(session *downstreamSession, status string, erro
}
}
func (s Server) reportDownstreamDelivery(event downstreamDeliveryLifecycleEvent) {
if strings.TrimSpace(event.DeliveryID) == "" {
return
}
payload := map[string]any{
"id": event.DeliveryID, "connectionId": event.ConnectionID,
"sequenceId": strconv.FormatUint(uint64(event.SequenceID), 10),
"messageId": strconv.FormatUint(event.MessageID, 10),
}
switch event.Kind {
case "sent":
payload["sentAt"] = formatRFC3339Nano(event.ObservedAt)
payload["ackDeadlineAt"] = formatRFC3339Nano(event.AckDeadlineAt)
_ = s.post(context.Background(), "/gateway/events/downstream/sent", payload, nil)
case "acknowledged":
payload["result"] = event.Result
payload["acknowledgedAt"] = formatRFC3339Nano(event.ObservedAt)
_ = s.post(context.Background(), "/gateway/events/downstream/acknowledged", payload, nil)
case "failed":
payload["failureType"] = event.FailureType
payload["errorMessage"] = event.ErrorMessage
_ = s.post(context.Background(), "/gateway/events/downstream/failed", payload, nil)
}
}
type inboundSubmitPacket struct {
protocol string
pkTotal uint8
@@ -444,19 +517,19 @@ func (s Server) flushPending(account string, logger *log.Logger) (pendingFlushRe
}
result.Deliveries = len(deliveries)
for _, delivery := range deliveries {
delivered, err := s.pushPendingDelivery(account, delivery)
sendResult, err := s.pushPendingDelivery(account, delivery)
if err != nil {
result.FailedCount++
result.LastError = err.Error()
_ = s.post(context.Background(), "/gateway/events/downstream/failed", map[string]string{
"id": delivery.ID,
"errorMessage": err.Error(),
"failureType": "send_failed",
}, nil)
continue
}
if delivered {
if sendResult.Sent {
result.DeliveredCount++
_ = s.post(context.Background(), "/gateway/events/downstream/delivered", map[string]string{"id": delivery.ID}, nil)
continue
}
result.WaitingCount++
@@ -464,26 +537,26 @@ func (s Server) flushPending(account string, logger *log.Logger) (pendingFlushRe
return result, nil
}
func (s Server) pushPendingDelivery(account string, delivery pendingDelivery) (bool, error) {
func (s Server) pushPendingDelivery(account string, delivery pendingDelivery) (DownstreamSendResult, error) {
switch delivery.DeliveryType {
case "receipt":
var event DownstreamReceipt
if err := json.Unmarshal(delivery.Payload, &event); err != nil {
return false, err
return DownstreamSendResult{}, err
}
event.DeliveryID = delivery.ID
event.Account = defaultString(event.Account, account)
return PushReceipt(event)
return PushReceiptWithResult(event)
case "uplink":
var event DownstreamUplink
if err := json.Unmarshal(delivery.Payload, &event); err != nil {
return false, err
return DownstreamSendResult{}, err
}
event.DeliveryID = delivery.ID
event.Account = defaultString(event.Account, account)
return PushUplink(event)
return PushUplinkWithResult(event)
default:
return false, fmt.Errorf("unsupported downstream delivery type %q", delivery.DeliveryType)
return DownstreamSendResult{}, fmt.Errorf("unsupported downstream delivery type %q", delivery.DeliveryType)
}
}
@@ -770,9 +843,14 @@ func onlineAccounts() []string {
}
func PushReceipt(event DownstreamReceipt) (bool, error) {
result, err := PushReceiptWithResult(event)
return result.Sent, err
}
func PushReceiptWithResult(event DownstreamReceipt) (DownstreamSendResult, error) {
session := findSession(event.MessageID, event.Account)
if session == nil {
return false, nil
return DownstreamSendResult{}, nil
}
stat := strings.TrimSpace(event.RawStatus)
if stat == "" {
@@ -794,20 +872,25 @@ func PushReceipt(event DownstreamReceipt) (bool, error) {
}
receiptBytes, err := receipt.Pack()
if err != nil {
return false, err
return DownstreamSendResult{}, err
}
deliver := downstreamDeliverPacket(session, session.gatewayMsgID, session.srcID, defaultString(event.PhoneNumber, session.phoneNumber), 0, 1, string(receiptBytes))
return sendDownstream(session, deliver)
return sendDownstream(session, deliver, event.DeliveryID)
}
func PushUplink(event DownstreamUplink) (bool, error) {
result, err := PushUplinkWithResult(event)
return result.Sent, err
}
func PushUplinkWithResult(event DownstreamUplink) (DownstreamSendResult, error) {
session := findSession(event.MessageID, event.Account)
if session == nil {
return false, nil
return DownstreamSendResult{}, nil
}
content, err := cmpputils.Utf8ToUcs2(event.Content)
if err != nil {
return false, err
return DownstreamSendResult{}, err
}
deliver := downstreamDeliverPacket(
session,
@@ -818,7 +901,7 @@ func PushUplink(event DownstreamUplink) (bool, error) {
0,
content,
)
return sendDownstream(session, deliver)
return sendDownstream(session, deliver, event.DeliveryID)
}
func downstreamDeliverPacket(session *downstreamSession, messageID uint64, destID string, sourceTerminalID string, msgFmt uint8, registerDelivery uint8, content string) cmpp.Packer {
@@ -850,21 +933,134 @@ func findSession(messageID string, account string) *downstreamSession {
return nil
}
func sendDownstream(session *downstreamSession, deliver cmpp.Packer) (bool, error) {
func sendDownstream(session *downstreamSession, deliver cmpp.Packer, deliveryID string) (DownstreamSendResult, error) {
session.mu.Lock()
defer session.mu.Unlock()
if err := session.conn.SendPkt(deliver, <-session.conn.SeqId); err != nil {
sequenceID := <-session.conn.SeqId
messageID := downstreamDeliverMessageID(deliver)
sentAt := time.Now().UTC()
ackDeadlineAt := sentAt.Add(downstreamAckTimeout())
tracker := registerDownstreamAck(session, deliveryID, sequenceID, messageID, ackDeadlineAt)
if err := session.conn.SendPkt(deliver, sequenceID); err != nil {
removeDownstreamAck(tracker)
if session.report != nil {
go session.report(session, "disconnected", err.Error())
}
forgetDownstream(session)
return false, err
return DownstreamSendResult{}, err
}
result := DownstreamSendResult{
Sent: true, ConnectionID: session.connectionID,
SequenceID: strconv.FormatUint(uint64(sequenceID), 10), MessageID: strconv.FormatUint(messageID, 10),
SentAt: formatRFC3339Nano(sentAt), AckDeadlineAt: formatRFC3339Nano(ackDeadlineAt),
}
if deliveryID != "" && session.deliveryReport != nil {
go session.deliveryReport(downstreamDeliveryLifecycleEvent{
Kind: "sent", DeliveryID: deliveryID, ConnectionID: session.connectionID,
SequenceID: sequenceID, MessageID: messageID, ObservedAt: sentAt, AckDeadlineAt: ackDeadlineAt,
})
}
session.touchPresence("connected", false, true)
if session.report != nil {
go session.report(session, "deliver", "")
}
return true, nil
return result, nil
}
func downstreamDeliverMessageID(deliver cmpp.Packer) uint64 {
switch packet := deliver.(type) {
case *cmpp.Cmpp2DeliverReqPkt:
return packet.MsgId
case *cmpp.Cmpp3DeliverReqPkt:
return packet.MsgId
default:
return 0
}
}
func downstreamAckKey(conn *cmpp.Conn, sequenceID uint32) string {
return fmt.Sprintf("%p:%d", conn, sequenceID)
}
func registerDownstreamAck(session *downstreamSession, deliveryID string, sequenceID uint32, messageID uint64, deadline time.Time) *downstreamAckTracker {
if session == nil || session.conn == nil || strings.TrimSpace(deliveryID) == "" {
return nil
}
tracker := &downstreamAckTracker{
deliveryID: deliveryID, connectionID: session.connectionID,
sequenceID: sequenceID, messageID: messageID, session: session,
}
key := downstreamAckKey(session.conn, sequenceID)
downstreamAckRegistry.Lock()
downstreamAckRegistry.items[key] = tracker
downstreamAckRegistry.Unlock()
tracker.timer = time.AfterFunc(time.Until(deadline), func() {
timedOut := takeDownstreamAck(session.conn, sequenceID)
if timedOut == nil || timedOut.session == nil || timedOut.session.deliveryReport == nil {
return
}
timedOut.session.deliveryReport(downstreamDeliveryLifecycleEvent{
Kind: "failed", DeliveryID: timedOut.deliveryID, ConnectionID: timedOut.connectionID,
SequenceID: timedOut.sequenceID, MessageID: timedOut.messageID, ObservedAt: time.Now().UTC(),
FailureType: "ack_timeout", ErrorMessage: "CMPP_DELIVER_RESP timeout",
})
})
return tracker
}
func takeDownstreamAck(conn *cmpp.Conn, sequenceID uint32) *downstreamAckTracker {
key := downstreamAckKey(conn, sequenceID)
downstreamAckRegistry.Lock()
tracker := downstreamAckRegistry.items[key]
delete(downstreamAckRegistry.items, key)
downstreamAckRegistry.Unlock()
if tracker != nil && tracker.timer != nil {
tracker.timer.Stop()
}
return tracker
}
func removeDownstreamAck(tracker *downstreamAckTracker) {
if tracker == nil || tracker.session == nil {
return
}
_ = takeDownstreamAck(tracker.session.conn, tracker.sequenceID)
}
func handleDownstreamAcknowledgement(conn *cmpp.Conn, sequenceID uint32, messageID uint64, result uint32, logger *log.Logger) {
tracker := takeDownstreamAck(conn, sequenceID)
if tracker == nil {
if logger != nil {
logger.Printf("cmpp inbound event=deliver_ack_unmatched seq=%d message_id=%d result=%d", sequenceID, messageID, result)
}
return
}
if tracker.messageID != messageID {
if logger != nil {
logger.Printf("cmpp inbound event=deliver_ack_message_mismatch delivery_id=%s seq=%d expected_message_id=%d actual_message_id=%d", tracker.deliveryID, sequenceID, tracker.messageID, messageID)
}
result = 1
}
if logger != nil {
logger.Printf("cmpp inbound event=deliver_acknowledged delivery_id=%s connection_id=%s seq=%d message_id=%d result=%d", tracker.deliveryID, tracker.connectionID, sequenceID, messageID, result)
}
if tracker.session != nil && tracker.session.deliveryReport != nil {
go tracker.session.deliveryReport(downstreamDeliveryLifecycleEvent{
Kind: "acknowledged", DeliveryID: tracker.deliveryID, ConnectionID: tracker.connectionID,
SequenceID: sequenceID, MessageID: messageID, Result: result, ObservedAt: time.Now().UTC(),
})
}
}
func downstreamAckTimeout() time.Duration {
configured, err := strconv.Atoi(strings.TrimSpace(os.Getenv("CMPP_DOWNSTREAM_ACK_TIMEOUT_SECONDS")))
if err != nil || configured <= 0 {
return defaultDownstreamAckTimeout
}
if configured < 5 {
configured = 5
}
return time.Duration(configured) * time.Second
}
func (s Server) pendingFlushInterval() time.Duration {
+77 -3
View File
@@ -100,6 +100,7 @@ func TestInboundServerAuthenticatesAndSubmits(t *testing.T) {
var gotAuth authRequest
var gotSubmit submitRequest
connectionEvents := make(chan downstreamConnectionEvent, 8)
acknowledgements := make(chan map[string]any, 1)
api := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/api/gateway/events/inbound/authenticate":
@@ -121,6 +122,15 @@ func TestInboundServerAuthenticatesAndSubmits(t *testing.T) {
}
connectionEvents <- event
w.WriteHeader(http.StatusOK)
case "/api/gateway/events/downstream/sent":
w.WriteHeader(http.StatusOK)
case "/api/gateway/events/downstream/acknowledged":
var event map[string]any
if err := json.NewDecoder(r.Body).Decode(&event); err != nil {
t.Fatalf("decode acknowledgement: %v", err)
}
acknowledgements <- event
w.WriteHeader(http.StatusOK)
default:
t.Fatalf("unexpected api path: %s", r.URL.Path)
}
@@ -176,14 +186,15 @@ func TestInboundServerAuthenticatesAndSubmits(t *testing.T) {
if rsp.Result != 0 || rsp.MsgId == 0 {
t.Fatalf("unexpected submit response: %+v", rsp)
}
delivered, err := PushReceipt(DownstreamReceipt{
sendResult, err := PushReceiptWithResult(DownstreamReceipt{
DeliveryID: "delivery-1",
MessageID: "MSG-1",
PhoneNumber: "13500002696",
ReceiptStatus: "delivered",
DeliveredAt: time.Now().UTC().Format(time.RFC3339Nano),
})
if err != nil || !delivered {
t.Fatalf("push receipt delivered=%v err=%v", delivered, err)
if err != nil || !sendResult.Sent {
t.Fatalf("push receipt sent=%v err=%v", sendResult.Sent, err)
}
deliver := recvDeliver(t, client)
if deliver.RegisterDelivery != 1 {
@@ -196,6 +207,17 @@ func TestInboundServerAuthenticatesAndSubmits(t *testing.T) {
if receipt.Stat != "DELIVRD" || receipt.DestTerminalId != "13500002696" {
t.Fatalf("unexpected pushed receipt: %+v", receipt)
}
if err := client.SendRspPkt(&cmpp.Cmpp3DeliverRspPkt{MsgId: deliver.MsgId, Result: 0}, deliver.SeqId); err != nil {
t.Fatalf("send deliver response: %v", err)
}
select {
case event := <-acknowledgements:
if event["id"] != "delivery-1" || event["result"] != float64(0) {
t.Fatalf("unexpected acknowledgement callback: %+v", event)
}
case <-time.After(2 * time.Second):
t.Fatal("expected downstream acknowledgement callback")
}
if gotAuth.Account != account || gotAuth.AuthSource == "" || gotAuth.RemoteIP == "" {
t.Fatalf("unexpected auth payload: %+v", gotAuth)
}
@@ -537,6 +559,50 @@ func TestRememberAndForgetAccountUpdatesPresenceStore(t *testing.T) {
}
}
func TestDownstreamDeliveryRequiresAcknowledgement(t *testing.T) {
resetDownstreamRegistry()
defer resetDownstreamRegistry()
events := make(chan downstreamDeliveryLifecycleEvent, 1)
conn := &cmpp.Conn{}
session := &downstreamSession{
conn: conn, connectionID: "conn-1",
deliveryReport: func(event downstreamDeliveryLifecycleEvent) { events <- event },
}
registerDownstreamAck(session, "delivery-1", 37, 9016479179509871733, time.Now().Add(time.Second))
handleDownstreamAcknowledgement(conn, 37, 9016479179509871733, 0, log.Default())
select {
case event := <-events:
if event.Kind != "acknowledged" || event.DeliveryID != "delivery-1" || event.Result != 0 || event.SequenceID != 37 {
t.Fatalf("unexpected acknowledgement event: %+v", event)
}
case <-time.After(time.Second):
t.Fatal("timed out waiting acknowledgement event")
}
}
func TestDownstreamDeliveryReportsAckTimeout(t *testing.T) {
resetDownstreamRegistry()
defer resetDownstreamRegistry()
events := make(chan downstreamDeliveryLifecycleEvent, 1)
session := &downstreamSession{
conn: &cmpp.Conn{}, connectionID: "conn-1",
deliveryReport: func(event downstreamDeliveryLifecycleEvent) { events <- event },
}
registerDownstreamAck(session, "delivery-timeout", 38, 9017467844344255865, time.Now().Add(20*time.Millisecond))
select {
case event := <-events:
if event.Kind != "failed" || event.FailureType != "ack_timeout" || event.DeliveryID != "delivery-timeout" {
t.Fatalf("unexpected timeout event: %+v", event)
}
case <-time.After(time.Second):
t.Fatal("timed out waiting acknowledgement timeout")
}
}
func recvDeliver(t *testing.T, client *cmpp.Client) *cmpp.Cmpp3DeliverReqPkt {
t.Helper()
deadline := time.Now().Add(2 * time.Second)
@@ -575,6 +641,14 @@ func resetDownstreamRegistry() {
downstreamRegistry.byAccount = make(map[string]*downstreamSession)
downstreamRegistry.byMessageID = make(map[string]*downstreamSession)
downstreamRegistry.byConn = make(map[*cmpp.Conn]*downstreamSession)
downstreamAckRegistry.Lock()
for _, tracker := range downstreamAckRegistry.items {
if tracker.timer != nil {
tracker.timer.Stop()
}
}
downstreamAckRegistry.items = make(map[string]*downstreamAckTracker)
downstreamAckRegistry.Unlock()
}
func reserveTCPAddr(t *testing.T) string {