feat: harden sessions and track downstream acknowledgements
This commit is contained in:
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user