Copilot commented on code in PR #1125:
URL:
https://github.com/apache/incubator-seata-go/pull/1125#discussion_r3409094245
##########
pkg/rm/tcc/tcc_service.go:
##########
@@ -84,9 +85,47 @@ func (t *TCCServiceProxy) Prepare(ctx context.Context,
params interface{}) (inte
}
}
- // to set up the fence phase
tm.SetFencePhase(ctx, enum.FencePhasePrepare)
- return t.TCCResource.Prepare(ctx, params)
+ result, err := t.TCCResource.Prepare(ctx, params)
+ if err != nil {
+ return nil, err
+ }
+
+ bac := tm.GetBusinessActionContext(ctx)
+ if bac != nil && bac.IsDelayReport {
+ if err := t.reportActionContext(ctx, bac); err != nil {
+ return nil, fmt.Errorf("report action context failed
after prepare: %w", err)
+ }
+ }
+
+ return result, nil
Review Comment:
`Prepare` currently fails the whole TCC prepare if the post-prepare
`BranchReport` (used to update ActionContext) fails. Since phase-two
`Commit/Rollback` already treat END_TRANSACTION as best-effort and can fall
back to RocketMQ check-back, this makes the integration unnecessarily brittle
(a transient TC/RM network issue will cause Prepare to return error even though
the half-message was prepared successfully). Consider treating
`reportActionContext` failure as non-fatal: log a warning and continue, so the
system can still rely on broker check-back.
##########
pkg/integration/rocketmq/tcc_rocketmq_action.go:
##########
@@ -78,27 +82,80 @@ func (a *TCCRocketMQAction) Prepare(ctx context.Context,
params interface{}) (bo
bac.ActionContext[ActionContextKeyQueueId] =
result.MessageQueue.QueueId
bac.ActionContext[ActionContextKeyBrokerName] =
result.MessageQueue.BrokerName
}
+ bac.ActionContext[ActionContextKeyTopic] = msg.Topic
log.Infof("[TCCRocketMQ] Prepare success, xid=%s, branchId=%d,
msgId=%s", xid, bac.BranchId, result.MsgID)
return true, nil
}
func (a *TCCRocketMQAction) Commit(ctx context.Context, bac
*tm.BusinessActionContext) (bool, error) {
- // Commit is a no-op because RocketMQ transactional messages use a
check-back mechanism.
- // When the global transaction commits, RocketMQ will invoke
CheckLocalTransaction
- // via SeataTransactionListener to determine the final message
disposition.
- // The message has already been sent to the broker during Prepare phase
with an
- // initial state of UnknowState, pending the check-back resolution.
- log.Infof("[TCCRocketMQ] Commit (no-op, rely on check-back), xid=%s,
branchId=%d", bac.Xid, bac.BranchId)
+ topic := getStringFromMap(bac.ActionContext, ActionContextKeyTopic)
+ header := a.buildEndTransactionHeader(bac, topic,
commitOrRollbackCommit)
+ brokerName := getStringFromMap(bac.ActionContext,
ActionContextKeyBrokerName)
+ err := sendEndTransaction(
Review Comment:
`Commit` will attempt to resolve broker address / send END_TRANSACTION even
when required metadata (e.g. `topic` or `brokerName`) is missing from
`ActionContext`. This can happen if the delayed ActionContext report didn’t
reach the TC; in that case we should skip the active notification immediately
and fall back to check-back, avoiding unnecessary network work and noisy errors.
##########
pkg/integration/rocketmq/end_transaction_sender.go:
##########
@@ -0,0 +1,477 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package rocketmq
+
+import (
+ "bytes"
+ "encoding/binary"
+ "encoding/json"
+ "fmt"
+ "io"
+ "net"
+ "sort"
+ "sync/atomic"
+ "time"
+)
+
+const (
+ reqEndTransaction int16 = 37
+
+ reqGetRouteInfoByTopic int16 = 105
+
+ rmqProtocolVersion int16 = 317
+
+ rmqLanguageGo byte = 9
+
+ rmqCodecType byte = 1
+
+ rmqHeaderFixedLength = 21
+
+ commitOrRollbackCommit = 8
+
+ commitOrRollbackRollback = 12
+)
+
+var opaqueCounter int32
+
+type endTransactionRequestHeader struct {
+ Topic string
+ ProducerGroup string
+ TranStateTableOffset int64
+ CommitLogOffset int64
+ CommitOrRollback int
+ FromTransactionCheck bool
+ MsgID string
+ TransactionId string
+}
+
+func (h *endTransactionRequestHeader) Encode() map[string]string {
+ return map[string]string{
+ "topic": h.Topic,
+ "producerGroup": h.ProducerGroup,
+ "tranStateTableOffset": fmt.Sprintf("%d",
h.TranStateTableOffset),
+ "commitLogOffset": fmt.Sprintf("%d", h.CommitLogOffset),
+ "commitOrRollback": fmt.Sprintf("%d", h.CommitOrRollback),
+ "fromTransactionCheck": fmt.Sprintf("%v",
h.FromTransactionCheck),
+ "msgId": h.MsgID,
+ "transactionId": h.TransactionId,
+ }
+}
+
+type remotingCommand struct {
+ Code int16
+ Language byte
+ Version int16
+ Opaque int32
+ Flag int32
+ Remark string
+ ExtFields map[string]string
+ Body []byte
+}
+
+func newEndTransactionCommand(header *endTransactionRequestHeader)
*remotingCommand {
+ return &remotingCommand{
+ Code: reqEndTransaction,
+ Language: rmqLanguageGo,
+ Version: rmqProtocolVersion,
+ Opaque: atomic.AddInt32(&opaqueCounter, 1),
+ ExtFields: header.Encode(),
+ }
+}
+
+func (cmd *remotingCommand) encode() ([]byte, error) {
+ headerBytes, err := cmd.encodeHeader()
+ if err != nil {
+ return nil, fmt.Errorf("encode header failed: %w", err)
+ }
+
+ frameSize := 4 + len(headerBytes) + len(cmd.Body)
+ buf := bytes.NewBuffer(make([]byte, 0, 4+frameSize))
+
+ if err := binary.Write(buf, binary.BigEndian, int32(frameSize)); err !=
nil {
+ return nil, err
+ }
+ if err := binary.Write(buf, binary.BigEndian,
markProtocolType(int32(len(headerBytes)))); err != nil {
+ return nil, err
+ }
+ if _, err := buf.Write(headerBytes); err != nil {
+ return nil, err
+ }
+ if len(cmd.Body) > 0 {
+ if _, err := buf.Write(cmd.Body); err != nil {
+ return nil, err
+ }
+ }
+
+ return buf.Bytes(), nil
+}
+
+func (cmd *remotingCommand) encodeHeader() ([]byte, error) {
+ extBytes, err := encodeExtFields(cmd.ExtFields)
+ if err != nil {
+ return nil, err
+ }
+
+ buf := bytes.NewBuffer(make([]byte, 0,
rmqHeaderFixedLength+len(cmd.Remark)+len(extBytes)))
+
+ if err := binary.Write(buf, binary.BigEndian, cmd.Code); err != nil {
+ return nil, err
+ }
+ if err := buf.WriteByte(cmd.Language); err != nil {
+ return nil, err
+ }
+ if err := binary.Write(buf, binary.BigEndian, cmd.Version); err != nil {
+ return nil, err
+ }
+ if err := binary.Write(buf, binary.BigEndian, cmd.Opaque); err != nil {
+ return nil, err
+ }
+ if err := binary.Write(buf, binary.BigEndian, cmd.Flag); err != nil {
+ return nil, err
+ }
+ if err := binary.Write(buf, binary.BigEndian, int32(len(cmd.Remark)));
err != nil {
+ return nil, err
+ }
+ if len(cmd.Remark) > 0 {
+ if _, err := buf.Write([]byte(cmd.Remark)); err != nil {
+ return nil, err
+ }
+ }
+ if err := binary.Write(buf, binary.BigEndian, int32(len(extBytes)));
err != nil {
+ return nil, err
+ }
+ if len(extBytes) > 0 {
+ if _, err := buf.Write(extBytes); err != nil {
+ return nil, err
+ }
+ }
+
+ return buf.Bytes(), nil
+}
+
+func encodeExtFields(fields map[string]string) ([]byte, error) {
+ if len(fields) == 0 {
+ return []byte{}, nil
+ }
+ keys := make([]string, 0, len(fields))
+ for k := range fields {
+ keys = append(keys, k)
+ }
+ sort.Strings(keys)
+
+ buf := bytes.NewBuffer(nil)
+ for _, key := range keys {
+ value := fields[key]
+ if err := binary.Write(buf, binary.BigEndian, int16(len(key)));
err != nil {
+ return nil, err
+ }
+ if _, err := buf.Write([]byte(key)); err != nil {
+ return nil, err
+ }
+ if err := binary.Write(buf, binary.BigEndian,
int32(len(value))); err != nil {
+ return nil, err
+ }
+ if _, err := buf.Write([]byte(value)); err != nil {
+ return nil, err
+ }
+ }
+ return buf.Bytes(), nil
+}
+
+func markProtocolType(source int32) []byte {
+ result := make([]byte, 4)
+ result[0] = rmqCodecType
+ result[1] = byte((source >> 16) & 0xFF)
+ result[2] = byte((source >> 8) & 0xFF)
+ result[3] = byte(source & 0xFF)
+ return result
+}
+
+type brokerAddrResolver interface {
+ ResolveBrokerAddr(nameServerAddrs []string, topic string, brokerName
string) (string, error)
+}
+
+type tcpSender interface {
+ Send(addr string, data []byte, timeout time.Duration) error
+}
+
+type defaultBrokerAddrResolver struct{}
+
+func (r *defaultBrokerAddrResolver) ResolveBrokerAddr(nameServerAddrs
[]string, topic string, brokerName string) (string, error) {
+ for _, nsAddr := range nameServerAddrs {
+ addr, err := queryBrokerAddrFromNameServer(nsAddr, topic,
brokerName)
+ if err == nil && addr != "" {
+ return addr, nil
+ }
+ }
+ return "", fmt.Errorf("broker %s addr not found from name servers",
brokerName)
+}
+
+type defaultTCPSender struct{}
+
+func (s *defaultTCPSender) Send(addr string, data []byte, timeout
time.Duration) error {
+ conn, err := net.DialTimeout("tcp", addr, timeout)
+ if err != nil {
+ return fmt.Errorf("dial broker %s failed: %w", addr, err)
+ }
+ defer conn.Close()
+
+ if err := conn.SetWriteDeadline(time.Now().Add(timeout)); err != nil {
+ return fmt.Errorf("set write deadline failed: %w", err)
+ }
+ if _, err := conn.Write(data); err != nil {
+ return fmt.Errorf("write to broker %s failed: %w", addr, err)
+ }
+
+ if err := conn.SetReadDeadline(time.Now().Add(timeout)); err != nil {
+ return fmt.Errorf("set read deadline failed: %w", err)
+ }
+ frameLenBuf := make([]byte, 4)
+ if _, err := readFull(conn, frameLenBuf); err != nil {
+ return fmt.Errorf("read response frame length from broker %s
failed: %w", addr, err)
+ }
+ frameLen := int(binary.BigEndian.Uint32(frameLenBuf))
+ if frameLen < 8 {
+ return fmt.Errorf("broker %s response frame too short: %d",
addr, frameLen)
+ }
+ frameBuf := make([]byte, frameLen)
+ if _, err := readFull(conn, frameBuf); err != nil {
+ return fmt.Errorf("read response frame from broker %s failed:
%w", addr, err)
+ }
+
+ code := int16(binary.BigEndian.Uint16(frameBuf[4:6]))
+ if code != 0 {
+ return fmt.Errorf("broker %s rejected END_TRANSACTION,
responseCode=%d", addr, code)
+ }
+ return nil
+}
+
+func sendEndTransaction(
+ nameServerAddrs []string,
+ topic string,
+ brokerName string,
+ header *endTransactionRequestHeader,
+ timeout time.Duration,
+ resolver brokerAddrResolver,
+ sender tcpSender,
+) error {
+ brokerAddr, err := resolver.ResolveBrokerAddr(nameServerAddrs, topic,
brokerName)
+ if err != nil {
+ return fmt.Errorf("resolve broker addr failed: %w", err)
+ }
+
+ cmd := newEndTransactionCommand(header)
+
+ data, err := cmd.encode()
+ if err != nil {
+ return fmt.Errorf("encode command failed: %w", err)
+ }
+
+ if err := sender.Send(brokerAddr, data, timeout); err != nil {
+ return fmt.Errorf("send to broker %s failed: %w", brokerAddr,
err)
+ }
+
+ return nil
+}
+
+func queryBrokerAddrFromNameServer(nameServerAddr string, topic string,
brokerName string) (string, error) {
+ cmd := &remotingCommand{
+ Code: reqGetRouteInfoByTopic,
+ Language: rmqLanguageGo,
+ Version: rmqProtocolVersion,
+ Opaque: atomic.AddInt32(&opaqueCounter, 1),
+ ExtFields: map[string]string{
+ "topic": topic,
+ },
+ }
+
+ data, err := cmd.encode()
+ if err != nil {
+ return "", fmt.Errorf("encode route request failed: %w", err)
+ }
+
+ conn, err := net.DialTimeout("tcp", nameServerAddr, 3*time.Second)
+ if err != nil {
+ return "", fmt.Errorf("dial name server %s failed: %w",
nameServerAddr, err)
+ }
+ defer conn.Close()
+
+ if err := conn.SetWriteDeadline(time.Now().Add(3 * time.Second)); err
!= nil {
+ return "", err
+ }
+ if _, err := conn.Write(data); err != nil {
+ return "", fmt.Errorf("send route request failed: %w", err)
+ }
+
+ if err := conn.SetReadDeadline(time.Now().Add(3 * time.Second)); err !=
nil {
+ return "", err
+ }
+ frameLenBuf := make([]byte, 4)
+ if _, err := readFull(conn, frameLenBuf); err != nil {
+ return "", fmt.Errorf("read frame length failed: %w", err)
+ }
+ frameLen := int(binary.BigEndian.Uint32(frameLenBuf))
+ if frameLen < 8 {
+ return "", fmt.Errorf("name server %s response frame too short:
%d", nameServerAddr, frameLen)
+ }
+
+ frameBuf := make([]byte, frameLen)
+ if _, err := readFull(conn, frameBuf); err != nil {
+ return "", fmt.Errorf("read frame data failed: %w", err)
+ }
Review Comment:
`queryBrokerAddrFromNameServer` also allocates `frameBuf` directly from a
peer-provided length. Bound the maximum frame size to avoid OOM on malformed
responses.
##########
pkg/integration/rocketmq/tcc_rocketmq_action.go:
##########
@@ -78,27 +82,80 @@ func (a *TCCRocketMQAction) Prepare(ctx context.Context,
params interface{}) (bo
bac.ActionContext[ActionContextKeyQueueId] =
result.MessageQueue.QueueId
bac.ActionContext[ActionContextKeyBrokerName] =
result.MessageQueue.BrokerName
}
+ bac.ActionContext[ActionContextKeyTopic] = msg.Topic
log.Infof("[TCCRocketMQ] Prepare success, xid=%s, branchId=%d,
msgId=%s", xid, bac.BranchId, result.MsgID)
return true, nil
}
func (a *TCCRocketMQAction) Commit(ctx context.Context, bac
*tm.BusinessActionContext) (bool, error) {
- // Commit is a no-op because RocketMQ transactional messages use a
check-back mechanism.
- // When the global transaction commits, RocketMQ will invoke
CheckLocalTransaction
- // via SeataTransactionListener to determine the final message
disposition.
- // The message has already been sent to the broker during Prepare phase
with an
- // initial state of UnknowState, pending the check-back resolution.
- log.Infof("[TCCRocketMQ] Commit (no-op, rely on check-back), xid=%s,
branchId=%d", bac.Xid, bac.BranchId)
+ topic := getStringFromMap(bac.ActionContext, ActionContextKeyTopic)
+ header := a.buildEndTransactionHeader(bac, topic,
commitOrRollbackCommit)
+ brokerName := getStringFromMap(bac.ActionContext,
ActionContextKeyBrokerName)
+ err := sendEndTransaction(
+ a.producer.config.NameServerAddrs,
+ topic,
+ brokerName,
+ header,
+ a.producer.config.SendMsgTimeout,
+ a.resolver,
+ a.sender,
+ )
+ if err != nil {
+ log.Warnf("[TCCRocketMQ] Commit send END_TRANSACTION failed,
fallback to check-back, xid=%s, branchId=%d, err=%v",
+ bac.Xid, bac.BranchId, err)
+ return true, nil
+ }
+ log.Infof("[TCCRocketMQ] Commit send END_TRANSACTION success, xid=%s,
branchId=%d", bac.Xid, bac.BranchId)
return true, nil
}
func (a *TCCRocketMQAction) Rollback(ctx context.Context, bac
*tm.BusinessActionContext) (bool, error) {
- // Rollback is a no-op because RocketMQ transactional messages use a
check-back mechanism.
- // When the global transaction rolls back, RocketMQ will invoke
CheckLocalTransaction
- // via SeataTransactionListener, which queries the TC for the global
status and returns
- // RollbackMessageState, causing the broker to discard the message.
- log.Infof("[TCCRocketMQ] Rollback (no-op, rely on check-back), xid=%s,
branchId=%d", bac.Xid, bac.BranchId)
+ topic := getStringFromMap(bac.ActionContext, ActionContextKeyTopic)
+ header := a.buildEndTransactionHeader(bac, topic,
commitOrRollbackRollback)
+ brokerName := getStringFromMap(bac.ActionContext,
ActionContextKeyBrokerName)
+ err := sendEndTransaction(
Review Comment:
Same as `Commit`: `Rollback` should avoid attempting END_TRANSACTION when
`topic` / `brokerName` are missing from `ActionContext` (e.g. delayed report
failed). Skipping early reduces unnecessary NameServer/Broker traffic and makes
fallback behavior clearer.
##########
pkg/integration/rocketmq/end_transaction_sender.go:
##########
@@ -0,0 +1,477 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package rocketmq
+
+import (
+ "bytes"
+ "encoding/binary"
+ "encoding/json"
+ "fmt"
+ "io"
+ "net"
+ "sort"
+ "sync/atomic"
+ "time"
+)
+
+const (
+ reqEndTransaction int16 = 37
+
+ reqGetRouteInfoByTopic int16 = 105
+
+ rmqProtocolVersion int16 = 317
+
+ rmqLanguageGo byte = 9
+
+ rmqCodecType byte = 1
+
+ rmqHeaderFixedLength = 21
+
+ commitOrRollbackCommit = 8
+
+ commitOrRollbackRollback = 12
+)
+
+var opaqueCounter int32
+
+type endTransactionRequestHeader struct {
+ Topic string
+ ProducerGroup string
+ TranStateTableOffset int64
+ CommitLogOffset int64
+ CommitOrRollback int
+ FromTransactionCheck bool
+ MsgID string
+ TransactionId string
+}
+
+func (h *endTransactionRequestHeader) Encode() map[string]string {
+ return map[string]string{
+ "topic": h.Topic,
+ "producerGroup": h.ProducerGroup,
+ "tranStateTableOffset": fmt.Sprintf("%d",
h.TranStateTableOffset),
+ "commitLogOffset": fmt.Sprintf("%d", h.CommitLogOffset),
+ "commitOrRollback": fmt.Sprintf("%d", h.CommitOrRollback),
+ "fromTransactionCheck": fmt.Sprintf("%v",
h.FromTransactionCheck),
+ "msgId": h.MsgID,
+ "transactionId": h.TransactionId,
+ }
+}
+
+type remotingCommand struct {
+ Code int16
+ Language byte
+ Version int16
+ Opaque int32
+ Flag int32
+ Remark string
+ ExtFields map[string]string
+ Body []byte
+}
+
+func newEndTransactionCommand(header *endTransactionRequestHeader)
*remotingCommand {
+ return &remotingCommand{
+ Code: reqEndTransaction,
+ Language: rmqLanguageGo,
+ Version: rmqProtocolVersion,
+ Opaque: atomic.AddInt32(&opaqueCounter, 1),
+ ExtFields: header.Encode(),
+ }
+}
+
+func (cmd *remotingCommand) encode() ([]byte, error) {
+ headerBytes, err := cmd.encodeHeader()
+ if err != nil {
+ return nil, fmt.Errorf("encode header failed: %w", err)
+ }
+
+ frameSize := 4 + len(headerBytes) + len(cmd.Body)
+ buf := bytes.NewBuffer(make([]byte, 0, 4+frameSize))
+
+ if err := binary.Write(buf, binary.BigEndian, int32(frameSize)); err !=
nil {
+ return nil, err
+ }
+ if err := binary.Write(buf, binary.BigEndian,
markProtocolType(int32(len(headerBytes)))); err != nil {
+ return nil, err
+ }
+ if _, err := buf.Write(headerBytes); err != nil {
+ return nil, err
+ }
+ if len(cmd.Body) > 0 {
+ if _, err := buf.Write(cmd.Body); err != nil {
+ return nil, err
+ }
+ }
+
+ return buf.Bytes(), nil
+}
+
+func (cmd *remotingCommand) encodeHeader() ([]byte, error) {
+ extBytes, err := encodeExtFields(cmd.ExtFields)
+ if err != nil {
+ return nil, err
+ }
+
+ buf := bytes.NewBuffer(make([]byte, 0,
rmqHeaderFixedLength+len(cmd.Remark)+len(extBytes)))
+
+ if err := binary.Write(buf, binary.BigEndian, cmd.Code); err != nil {
+ return nil, err
+ }
+ if err := buf.WriteByte(cmd.Language); err != nil {
+ return nil, err
+ }
+ if err := binary.Write(buf, binary.BigEndian, cmd.Version); err != nil {
+ return nil, err
+ }
+ if err := binary.Write(buf, binary.BigEndian, cmd.Opaque); err != nil {
+ return nil, err
+ }
+ if err := binary.Write(buf, binary.BigEndian, cmd.Flag); err != nil {
+ return nil, err
+ }
+ if err := binary.Write(buf, binary.BigEndian, int32(len(cmd.Remark)));
err != nil {
+ return nil, err
+ }
+ if len(cmd.Remark) > 0 {
+ if _, err := buf.Write([]byte(cmd.Remark)); err != nil {
+ return nil, err
+ }
+ }
+ if err := binary.Write(buf, binary.BigEndian, int32(len(extBytes)));
err != nil {
+ return nil, err
+ }
+ if len(extBytes) > 0 {
+ if _, err := buf.Write(extBytes); err != nil {
+ return nil, err
+ }
+ }
+
+ return buf.Bytes(), nil
+}
+
+func encodeExtFields(fields map[string]string) ([]byte, error) {
+ if len(fields) == 0 {
+ return []byte{}, nil
+ }
+ keys := make([]string, 0, len(fields))
+ for k := range fields {
+ keys = append(keys, k)
+ }
+ sort.Strings(keys)
+
+ buf := bytes.NewBuffer(nil)
+ for _, key := range keys {
+ value := fields[key]
+ if err := binary.Write(buf, binary.BigEndian, int16(len(key)));
err != nil {
+ return nil, err
+ }
+ if _, err := buf.Write([]byte(key)); err != nil {
+ return nil, err
+ }
+ if err := binary.Write(buf, binary.BigEndian,
int32(len(value))); err != nil {
+ return nil, err
+ }
+ if _, err := buf.Write([]byte(value)); err != nil {
+ return nil, err
+ }
+ }
+ return buf.Bytes(), nil
+}
+
+func markProtocolType(source int32) []byte {
+ result := make([]byte, 4)
+ result[0] = rmqCodecType
+ result[1] = byte((source >> 16) & 0xFF)
+ result[2] = byte((source >> 8) & 0xFF)
+ result[3] = byte(source & 0xFF)
+ return result
+}
+
+type brokerAddrResolver interface {
+ ResolveBrokerAddr(nameServerAddrs []string, topic string, brokerName
string) (string, error)
+}
+
+type tcpSender interface {
+ Send(addr string, data []byte, timeout time.Duration) error
+}
+
+type defaultBrokerAddrResolver struct{}
+
+func (r *defaultBrokerAddrResolver) ResolveBrokerAddr(nameServerAddrs
[]string, topic string, brokerName string) (string, error) {
+ for _, nsAddr := range nameServerAddrs {
+ addr, err := queryBrokerAddrFromNameServer(nsAddr, topic,
brokerName)
+ if err == nil && addr != "" {
+ return addr, nil
+ }
+ }
+ return "", fmt.Errorf("broker %s addr not found from name servers",
brokerName)
+}
+
+type defaultTCPSender struct{}
+
+func (s *defaultTCPSender) Send(addr string, data []byte, timeout
time.Duration) error {
+ conn, err := net.DialTimeout("tcp", addr, timeout)
+ if err != nil {
+ return fmt.Errorf("dial broker %s failed: %w", addr, err)
+ }
+ defer conn.Close()
+
+ if err := conn.SetWriteDeadline(time.Now().Add(timeout)); err != nil {
+ return fmt.Errorf("set write deadline failed: %w", err)
+ }
+ if _, err := conn.Write(data); err != nil {
+ return fmt.Errorf("write to broker %s failed: %w", addr, err)
+ }
+
+ if err := conn.SetReadDeadline(time.Now().Add(timeout)); err != nil {
+ return fmt.Errorf("set read deadline failed: %w", err)
+ }
+ frameLenBuf := make([]byte, 4)
+ if _, err := readFull(conn, frameLenBuf); err != nil {
+ return fmt.Errorf("read response frame length from broker %s
failed: %w", addr, err)
+ }
+ frameLen := int(binary.BigEndian.Uint32(frameLenBuf))
+ if frameLen < 8 {
+ return fmt.Errorf("broker %s response frame too short: %d",
addr, frameLen)
+ }
+ frameBuf := make([]byte, frameLen)
+ if _, err := readFull(conn, frameBuf); err != nil {
+ return fmt.Errorf("read response frame from broker %s failed:
%w", addr, err)
+ }
+
+ code := int16(binary.BigEndian.Uint16(frameBuf[4:6]))
+ if code != 0 {
+ return fmt.Errorf("broker %s rejected END_TRANSACTION,
responseCode=%d", addr, code)
+ }
+ return nil
+}
+
+func sendEndTransaction(
+ nameServerAddrs []string,
+ topic string,
+ brokerName string,
+ header *endTransactionRequestHeader,
+ timeout time.Duration,
+ resolver brokerAddrResolver,
+ sender tcpSender,
+) error {
+ brokerAddr, err := resolver.ResolveBrokerAddr(nameServerAddrs, topic,
brokerName)
+ if err != nil {
+ return fmt.Errorf("resolve broker addr failed: %w", err)
+ }
+
+ cmd := newEndTransactionCommand(header)
+
+ data, err := cmd.encode()
+ if err != nil {
+ return fmt.Errorf("encode command failed: %w", err)
+ }
+
+ if err := sender.Send(brokerAddr, data, timeout); err != nil {
+ return fmt.Errorf("send to broker %s failed: %w", brokerAddr,
err)
+ }
+
+ return nil
+}
+
+func queryBrokerAddrFromNameServer(nameServerAddr string, topic string,
brokerName string) (string, error) {
+ cmd := &remotingCommand{
+ Code: reqGetRouteInfoByTopic,
+ Language: rmqLanguageGo,
+ Version: rmqProtocolVersion,
+ Opaque: atomic.AddInt32(&opaqueCounter, 1),
+ ExtFields: map[string]string{
+ "topic": topic,
+ },
+ }
+
+ data, err := cmd.encode()
+ if err != nil {
+ return "", fmt.Errorf("encode route request failed: %w", err)
+ }
+
+ conn, err := net.DialTimeout("tcp", nameServerAddr, 3*time.Second)
+ if err != nil {
+ return "", fmt.Errorf("dial name server %s failed: %w",
nameServerAddr, err)
+ }
+ defer conn.Close()
+
+ if err := conn.SetWriteDeadline(time.Now().Add(3 * time.Second)); err
!= nil {
+ return "", err
+ }
+ if _, err := conn.Write(data); err != nil {
+ return "", fmt.Errorf("send route request failed: %w", err)
+ }
+
+ if err := conn.SetReadDeadline(time.Now().Add(3 * time.Second)); err !=
nil {
+ return "", err
+ }
+ frameLenBuf := make([]byte, 4)
+ if _, err := readFull(conn, frameLenBuf); err != nil {
+ return "", fmt.Errorf("read frame length failed: %w", err)
+ }
+ frameLen := int(binary.BigEndian.Uint32(frameLenBuf))
+ if frameLen < 8 {
+ return "", fmt.Errorf("name server %s response frame too short:
%d", nameServerAddr, frameLen)
+ }
+
+ frameBuf := make([]byte, frameLen)
+ if _, err := readFull(conn, frameBuf); err != nil {
+ return "", fmt.Errorf("read frame data failed: %w", err)
+ }
+
+ if len(frameBuf) < 4 {
+ return "", fmt.Errorf("response frame too short")
+ }
+ oriHeaderLen := binary.BigEndian.Uint32(frameBuf[0:4])
+ headerLen := int(oriHeaderLen & 0xFFFFFF)
+ if len(frameBuf) < 4+headerLen {
+ return "", fmt.Errorf("response header truncated")
+ }
+
+ respCode := int16(binary.BigEndian.Uint16(frameBuf[4:6]))
+ if respCode != 0 {
+ return "", fmt.Errorf("name server %s returned error code %d
for route query", nameServerAddr, respCode)
+ }
+
+ _, err = decodeResponseHeader(frameBuf[4 : 4+headerLen])
+ if err != nil {
+ return "", fmt.Errorf("decode response header failed: %w", err)
+ }
+
+ bodyStart := 4 + headerLen
+ if bodyStart >= len(frameBuf) {
+ return "", fmt.Errorf("response has no body")
+ }
+ body := frameBuf[bodyStart:]
+
+ return parseBrokerAddrFromRouteBody(body, brokerName)
+}
+
+func readFull(conn net.Conn, buf []byte) (int, error) {
+ return io.ReadFull(conn, buf)
+}
+
+func decodeResponseHeader(data []byte) (map[string]string, error) {
+ buf := bytes.NewReader(data)
+
+ if buf.Len() < 13 {
+ return nil, fmt.Errorf("header too short for fixed fields")
+ }
+ discard := make([]byte, 13)
+ if _, err := buf.Read(discard); err != nil {
+ return nil, err
+ }
+
+ var remarkLen int32
+ if err := binary.Read(buf, binary.BigEndian, &remarkLen); err != nil {
+ return nil, err
+ }
+ if remarkLen > 0 {
+ discardRemark := make([]byte, remarkLen)
+ if _, err := buf.Read(discardRemark); err != nil {
+ return nil, err
+ }
+ }
+
+ var extLen int32
+ if err := binary.Read(buf, binary.BigEndian, &extLen); err != nil {
+ return nil, err
+ }
+ extFields := make(map[string]string)
+ if extLen > 0 {
+ extData := make([]byte, extLen)
+ if _, err := buf.Read(extData); err != nil {
+ return nil, err
+ }
+ extBuf := bytes.NewReader(extData)
Review Comment:
`decodeResponseHeader` uses `remarkLen`/`extLen` read from the network
directly in `make(...)`. Negative or oversized lengths can panic the process or
cause large allocations. Validate lengths are non-negative and within remaining
buffer size, and use `io.ReadFull` to ensure complete reads.
##########
pkg/integration/rocketmq/end_transaction_sender_test.go:
##########
@@ -0,0 +1,475 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package rocketmq
+
+import (
+ "bytes"
+ "encoding/binary"
+ "encoding/json"
+ "errors"
+ "testing"
+ "time"
+
+ "github.com/stretchr/testify/assert"
+ "github.com/stretchr/testify/require"
+)
+
+type stubBrokerAddrResolver struct {
+ addr string
+ err error
+
+ calls int
+ lastNsAddrs []string
+ lastTopic string
+ lastBroker string
+}
+
+func (s *stubBrokerAddrResolver) ResolveBrokerAddr(nameServerAddrs []string,
topic string, brokerName string) (string, error) {
+ s.calls++
+ s.lastNsAddrs = nameServerAddrs
+ s.lastTopic = topic
+ s.lastBroker = brokerName
+ if s.err != nil {
+ return "", s.err
+ }
+ return s.addr, nil
+}
+
+type stubTCPSender struct {
+ err error
+
+ calls int
+ lastAddr string
+ lastData []byte
+}
+
+func (s *stubTCPSender) Send(addr string, data []byte, timeout time.Duration)
error {
+ s.calls++
+ s.lastAddr = addr
+ s.lastData = data
+ if s.err != nil {
+ return s.err
+ }
+ return nil
+}
+
+func TestEndTransactionRequestHeader_Encode(t *testing.T) {
+ header := &endTransactionRequestHeader{
+ Topic: "test-topic",
+ ProducerGroup: "test-group",
+ TranStateTableOffset: 42,
+ CommitLogOffset: 1024,
+ CommitOrRollback: commitOrRollbackCommit,
+ FromTransactionCheck: false,
+ MsgID: "msg-001",
+ TransactionId: "tx-001",
+ }
+
+ result := header.Encode()
+
+ assert.Equal(t, "test-topic", result["topic"])
+ assert.Equal(t, "test-group", result["producerGroup"])
+ assert.Equal(t, "42", result["tranStateTableOffset"])
+ assert.Equal(t, "1024", result["commitLogOffset"])
+ assert.Equal(t, "8", result["commitOrRollback"])
+ assert.Equal(t, "false", result["fromTransactionCheck"])
+ assert.Equal(t, "msg-001", result["msgId"])
+ assert.Equal(t, "tx-001", result["transactionId"])
+}
+
+func TestEndTransactionRequestHeader_EncodeRollback(t *testing.T) {
+ header := &endTransactionRequestHeader{
+ Topic: "test-topic",
+ ProducerGroup: "test-group",
+ TranStateTableOffset: 10,
+ CommitLogOffset: 2048,
+ CommitOrRollback: commitOrRollbackRollback,
+ FromTransactionCheck: false,
+ MsgID: "msg-002",
+ TransactionId: "tx-002",
+ }
+
+ result := header.Encode()
+
+ assert.Equal(t, "12", result["commitOrRollback"])
+}
+
+func TestNewEndTransactionCommand(t *testing.T) {
+ header := &endTransactionRequestHeader{
+ Topic: "test-topic",
+ ProducerGroup: "test-group",
+ CommitOrRollback: commitOrRollbackCommit,
+ MsgID: "msg-001",
+ TransactionId: "tx-001",
+ }
+
+ cmd := newEndTransactionCommand(header)
+
+ assert.Equal(t, reqEndTransaction, cmd.Code)
+ assert.Equal(t, rmqLanguageGo, cmd.Language)
+ assert.Equal(t, rmqProtocolVersion, cmd.Version)
+ assert.NotZero(t, cmd.Opaque)
+ assert.Equal(t, "test-topic", cmd.ExtFields["topic"])
+ assert.Equal(t, "test-group", cmd.ExtFields["producerGroup"])
+ assert.Equal(t, "8", cmd.ExtFields["commitOrRollback"])
+}
+
+func TestRemotingCommand_Encode_FrameFormat(t *testing.T) {
+ cmd := &remotingCommand{
+ Code: reqEndTransaction,
+ Language: rmqLanguageGo,
+ Version: rmqProtocolVersion,
+ Opaque: 1,
+ Flag: 0,
+ Remark: "",
+ ExtFields: map[string]string{},
+ }
+
+ data, err := cmd.encode()
+ require.NoError(t, err)
+ require.NotNil(t, data)
+
+ reader := bytes.NewReader(data)
+
+ var frameSize int32
+ err = binary.Read(reader, binary.BigEndian, &frameSize)
+ require.NoError(t, err)
+ assert.Equal(t, int32(len(data)-4), frameSize)
+
+ var headerLenRaw int32
+ err = binary.Read(reader, binary.BigEndian, &headerLenRaw)
+ require.NoError(t, err)
+ codecType := byte((headerLenRaw >> 24) & 0xFF)
+ headerLen := headerLenRaw & 0xFFFFFF
+ assert.Equal(t, rmqCodecType, codecType)
+ assert.Equal(t, int32(rmqHeaderFixedLength), headerLen)
+
+ var code int16
+ err = binary.Read(reader, binary.BigEndian, &code)
+ require.NoError(t, err)
+ assert.Equal(t, reqEndTransaction, code)
+
+ language, err := reader.ReadByte()
+ require.NoError(t, err)
+ assert.Equal(t, rmqLanguageGo, language)
+
+ var version int16
+ err = binary.Read(reader, binary.BigEndian, &version)
+ require.NoError(t, err)
+ assert.Equal(t, rmqProtocolVersion, version)
+
+ var opaque int32
+ err = binary.Read(reader, binary.BigEndian, &opaque)
+ require.NoError(t, err)
+ assert.Equal(t, int32(1), opaque)
+
+ var flag int32
+ err = binary.Read(reader, binary.BigEndian, &flag)
+ require.NoError(t, err)
+ assert.Equal(t, int32(0), flag)
+
+ var remarkLen int32
+ err = binary.Read(reader, binary.BigEndian, &remarkLen)
+ require.NoError(t, err)
+ assert.Equal(t, int32(0), remarkLen)
+
+ var extLen int32
+ err = binary.Read(reader, binary.BigEndian, &extLen)
+ require.NoError(t, err)
+ assert.Equal(t, int32(0), extLen)
+}
+
+func TestRemotingCommand_Encode_WithExtFields(t *testing.T) {
+ cmd := &remotingCommand{
+ Code: reqEndTransaction,
+ Language: rmqLanguageGo,
+ Version: rmqProtocolVersion,
+ Opaque: 42,
+ ExtFields: map[string]string{
+ "producerGroup": "test-group",
+ },
+ }
+
+ data, err := cmd.encode()
+ require.NoError(t, err)
+ require.NotNil(t, data)
+
+ frameSize := int32(binary.BigEndian.Uint32(data[0:4]))
+ assert.Equal(t, int32(len(data)-4), frameSize)
+
+ headerLenRaw := binary.BigEndian.Uint32(data[4:8])
+ headerLen := int(headerLenRaw & 0xFFFFFF)
+ assert.Equal(t, rmqHeaderFixedLength+29, headerLen)
+}
+
+func TestEncodeExtFields_Empty(t *testing.T) {
+ result, err := encodeExtFields(map[string]string{})
+ require.NoError(t, err)
+ assert.Empty(t, result)
+}
+
+func TestEncodeExtFields_SingleField(t *testing.T) {
+ fields := map[string]string{
+ "key": "value",
+ }
+
+ result, err := encodeExtFields(fields)
+ require.NoError(t, err)
+
+ reader := bytes.NewReader(result)
+
+ var keyLen int16
+ err = binary.Read(reader, binary.BigEndian, &keyLen)
+ require.NoError(t, err)
+ assert.Equal(t, int16(3), keyLen)
+
+ keyBuf := make([]byte, keyLen)
+ _, err = reader.Read(keyBuf)
+ require.NoError(t, err)
+ assert.Equal(t, "key", string(keyBuf))
+
+ var valueLen int32
+ err = binary.Read(reader, binary.BigEndian, &valueLen)
+ require.NoError(t, err)
+ assert.Equal(t, int32(5), valueLen)
+
+ valueBuf := make([]byte, valueLen)
+ _, err = reader.Read(valueBuf)
+ require.NoError(t, err)
+ assert.Equal(t, "value", string(valueBuf))
+}
+
+func TestMarkProtocolType(t *testing.T) {
+ result := markProtocolType(100)
+
+ assert.Equal(t, rmqCodecType, result[0])
+ assert.Equal(t, byte(0x00), result[1])
+ assert.Equal(t, byte(0x00), result[2])
+ assert.Equal(t, byte(0x64), result[3])
+}
+
+func TestParseBrokerAddrFromRouteBody_FoundMaster(t *testing.T) {
+ route := topicRouteData{
+ BrokerDataList: []brokerData{
+ {
+ BrokerName: "broker-a",
+ BrokerAddrs: map[string]string{
+ "0": "192.168.1.100:10911",
+ "1": "192.168.1.101:10911",
+ },
+ },
+ },
+ }
+ body, _ := json.Marshal(route)
+
+ addr, err := parseBrokerAddrFromRouteBody(body, "broker-a")
+
+ require.NoError(t, err)
+ assert.Equal(t, "192.168.1.100:10911", addr)
+}
+
+func TestParseBrokerAddrFromRouteBody_MasterNotFound(t *testing.T) {
+ route := topicRouteData{
+ BrokerDataList: []brokerData{
+ {
+ BrokerName: "broker-a",
+ BrokerAddrs: map[string]string{
+ "1": "192.168.1.101:10911",
+ },
+ },
+ },
+ }
+ body, _ := json.Marshal(route)
+
+ _, err := parseBrokerAddrFromRouteBody(body, "broker-a")
+
+ require.Error(t, err)
+ assert.Contains(t, err.Error(), "master (addr key '0') not found")
+}
+
+func TestParseBrokerAddrFromRouteBody_BrokerNotFound(t *testing.T) {
+ route := topicRouteData{
+ BrokerDataList: []brokerData{
+ {
+ BrokerName: "broker-a",
+ BrokerAddrs: map[string]string{"0":
"192.168.1.100:10911"},
+ },
+ },
+ }
+ body, _ := json.Marshal(route)
+
+ _, err := parseBrokerAddrFromRouteBody(body, "broker-b")
+
+ require.Error(t, err)
+ assert.Contains(t, err.Error(), "broker-b")
+}
+
+func TestParseBrokerAddrFromRouteBody_InvalidJSON(t *testing.T) {
+ _, err := parseBrokerAddrFromRouteBody([]byte("invalid json"),
"broker-a")
+
+ require.Error(t, err)
+ assert.Contains(t, err.Error(), "unmarshal")
+}
+
+func TestGetQueueOffsetFromActionContext_Float64(t *testing.T) {
+ ctx := map[string]interface{}{
+ ActionContextKeyQueueOffset: float64(42),
+ }
+ assert.Equal(t, int64(42), getQueueOffsetFromActionContext(ctx))
+}
+
+func TestGetQueueOffsetFromActionContext_Int64(t *testing.T) {
+ ctx := map[string]interface{}{
+ ActionContextKeyQueueOffset: int64(42),
+ }
+ assert.Equal(t, int64(42), getQueueOffsetFromActionContext(ctx))
+}
+
+func TestGetQueueOffsetFromActionContext_Int(t *testing.T) {
+ ctx := map[string]interface{}{
+ ActionContextKeyQueueOffset: int(42),
+ }
+ assert.Equal(t, int64(42), getQueueOffsetFromActionContext(ctx))
+}
+
+func TestGetQueueOffsetFromActionContext_Missing(t *testing.T) {
+ ctx := map[string]interface{}{}
+ assert.Equal(t, int64(0), getQueueOffsetFromActionContext(ctx))
+}
+
+func TestGetQueueOffsetFromActionContext_InvalidType(t *testing.T) {
+ ctx := map[string]interface{}{
+ ActionContextKeyQueueOffset: "not a number",
+ }
+ assert.Equal(t, int64(0), getQueueOffsetFromActionContext(ctx))
+}
+
+func TestGetStringFromMap_Found(t *testing.T) {
+ ctx := map[string]interface{}{
+ ActionContextKeyMsgId: "msg-001",
+ }
+ assert.Equal(t, "msg-001", getStringFromMap(ctx, ActionContextKeyMsgId))
+}
+
+func TestGetStringFromMap_Missing(t *testing.T) {
+ ctx := map[string]interface{}{}
+ assert.Equal(t, "", getStringFromMap(ctx, ActionContextKeyMsgId))
+}
+
+func TestGetStringFromMap_WrongType(t *testing.T) {
+ ctx := map[string]interface{}{
+ ActionContextKeyMsgId: 12345,
+ }
+ assert.Equal(t, "", getStringFromMap(ctx, ActionContextKeyMsgId))
+}
+
+func TestSendEndTransaction_Success(t *testing.T) {
+ resolver := &stubBrokerAddrResolver{addr: "192.168.1.100:10911"}
+ sender := &stubTCPSender{}
+ header := &endTransactionRequestHeader{
+ Topic: "test-topic",
+ ProducerGroup: "test-group",
+ TranStateTableOffset: 42,
+ CommitLogOffset: 1024,
+ CommitOrRollback: commitOrRollbackCommit,
+ MsgID: "msg-001",
+ TransactionId: "tx-001",
+ }
+
+ err := sendEndTransaction(
+ []string{"nameserver:9876"},
+ "test-topic",
+ "broker-a",
+ header,
+ 3*time.Second,
+ resolver,
+ sender,
+ )
+
+ require.NoError(t, err)
+ assert.Equal(t, 1, resolver.calls)
+ assert.Equal(t, []string{"nameserver:9876"}, resolver.lastNsAddrs)
+ assert.Equal(t, "test-topic", resolver.lastTopic)
+ assert.Equal(t, "broker-a", resolver.lastBroker)
+ assert.Equal(t, 1, sender.calls)
+ assert.Equal(t, "192.168.1.100:10911", sender.lastAddr)
+ assert.NotEmpty(t, sender.lastData)
+
+ frameSize := int32(binary.BigEndian.Uint32(sender.lastData[0:4]))
+ assert.Equal(t, int32(len(sender.lastData)-4), frameSize)
+}
+
+func TestSendEndTransaction_ResolverError(t *testing.T) {
+ resolver := &stubBrokerAddrResolver{err: errors.New("name server
unreachable")}
+ sender := &stubTCPSender{}
+ header := &endTransactionRequestHeader{
+ Topic: "test-topic",
+ ProducerGroup: "test-group",
+ }
+
+ err := sendEndTransaction(
+ []string{"nameserver:9876"},
+ "test-topic",
+ "broker-a",
+ header,
+ 3*time.Second,
+ resolver,
+ sender,
+ )
+
+ require.Error(t, err)
+ assert.Contains(t, err.Error(), "resolve broker addr failed")
+ assert.Equal(t, 0, sender.calls)
+}
+
+func TestSendEndTransaction_SendError(t *testing.T) {
+ resolver := &stubBrokerAddrResolver{addr: "192.168.1.100:10911"}
+ sender := &stubTCPSender{err: errors.New("connection refused")}
+ header := &endTransactionRequestHeader{
+ Topic: "test-topic",
+ ProducerGroup: "test-group",
+ }
+
+ err := sendEndTransaction(
+ []string{"nameserver:9876"},
+ "test-topic",
+ "broker-a",
+ header,
+ 3*time.Second,
+ resolver,
+ sender,
+ )
+
+ require.Error(t, err)
+ assert.Contains(t, err.Error(), "send to broker")
+ assert.Equal(t, 1, resolver.calls)
+ assert.Equal(t, 1, sender.calls)
+}
+
+func TestDefaultBrokerAddrResolver_TriesAllNameServers(t *testing.T) {
+ resolver := &defaultBrokerAddrResolver{}
+
+ _, err := resolver.ResolveBrokerAddr(
+ []string{"invalid-addr-1:9876", "invalid-addr-2:9876"},
+ "test-topic",
+ "broker-a",
+ )
Review Comment:
This unit test performs real network dials (via `defaultBrokerAddrResolver`
-> `net.DialTimeout`) against non-existent hostnames. Depending on DNS /
environment, this can be slow or flaky. Prefer localhost with closed ports to
fail fast, or refactor to stub the NameServer query.
##########
pkg/integration/rocketmq/end_transaction_sender.go:
##########
@@ -0,0 +1,477 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package rocketmq
+
+import (
+ "bytes"
+ "encoding/binary"
+ "encoding/json"
+ "fmt"
+ "io"
+ "net"
+ "sort"
+ "sync/atomic"
+ "time"
+)
+
+const (
+ reqEndTransaction int16 = 37
+
+ reqGetRouteInfoByTopic int16 = 105
+
+ rmqProtocolVersion int16 = 317
+
+ rmqLanguageGo byte = 9
+
+ rmqCodecType byte = 1
+
+ rmqHeaderFixedLength = 21
+
+ commitOrRollbackCommit = 8
+
+ commitOrRollbackRollback = 12
+)
+
+var opaqueCounter int32
+
+type endTransactionRequestHeader struct {
+ Topic string
+ ProducerGroup string
+ TranStateTableOffset int64
+ CommitLogOffset int64
+ CommitOrRollback int
+ FromTransactionCheck bool
+ MsgID string
+ TransactionId string
+}
+
+func (h *endTransactionRequestHeader) Encode() map[string]string {
+ return map[string]string{
+ "topic": h.Topic,
+ "producerGroup": h.ProducerGroup,
+ "tranStateTableOffset": fmt.Sprintf("%d",
h.TranStateTableOffset),
+ "commitLogOffset": fmt.Sprintf("%d", h.CommitLogOffset),
+ "commitOrRollback": fmt.Sprintf("%d", h.CommitOrRollback),
+ "fromTransactionCheck": fmt.Sprintf("%v",
h.FromTransactionCheck),
+ "msgId": h.MsgID,
+ "transactionId": h.TransactionId,
+ }
+}
+
+type remotingCommand struct {
+ Code int16
+ Language byte
+ Version int16
+ Opaque int32
+ Flag int32
+ Remark string
+ ExtFields map[string]string
+ Body []byte
+}
+
+func newEndTransactionCommand(header *endTransactionRequestHeader)
*remotingCommand {
+ return &remotingCommand{
+ Code: reqEndTransaction,
+ Language: rmqLanguageGo,
+ Version: rmqProtocolVersion,
+ Opaque: atomic.AddInt32(&opaqueCounter, 1),
+ ExtFields: header.Encode(),
+ }
+}
+
+func (cmd *remotingCommand) encode() ([]byte, error) {
+ headerBytes, err := cmd.encodeHeader()
+ if err != nil {
+ return nil, fmt.Errorf("encode header failed: %w", err)
+ }
+
+ frameSize := 4 + len(headerBytes) + len(cmd.Body)
+ buf := bytes.NewBuffer(make([]byte, 0, 4+frameSize))
+
+ if err := binary.Write(buf, binary.BigEndian, int32(frameSize)); err !=
nil {
+ return nil, err
+ }
+ if err := binary.Write(buf, binary.BigEndian,
markProtocolType(int32(len(headerBytes)))); err != nil {
+ return nil, err
+ }
+ if _, err := buf.Write(headerBytes); err != nil {
+ return nil, err
+ }
+ if len(cmd.Body) > 0 {
+ if _, err := buf.Write(cmd.Body); err != nil {
+ return nil, err
+ }
+ }
+
+ return buf.Bytes(), nil
+}
+
+func (cmd *remotingCommand) encodeHeader() ([]byte, error) {
+ extBytes, err := encodeExtFields(cmd.ExtFields)
+ if err != nil {
+ return nil, err
+ }
+
+ buf := bytes.NewBuffer(make([]byte, 0,
rmqHeaderFixedLength+len(cmd.Remark)+len(extBytes)))
+
+ if err := binary.Write(buf, binary.BigEndian, cmd.Code); err != nil {
+ return nil, err
+ }
+ if err := buf.WriteByte(cmd.Language); err != nil {
+ return nil, err
+ }
+ if err := binary.Write(buf, binary.BigEndian, cmd.Version); err != nil {
+ return nil, err
+ }
+ if err := binary.Write(buf, binary.BigEndian, cmd.Opaque); err != nil {
+ return nil, err
+ }
+ if err := binary.Write(buf, binary.BigEndian, cmd.Flag); err != nil {
+ return nil, err
+ }
+ if err := binary.Write(buf, binary.BigEndian, int32(len(cmd.Remark)));
err != nil {
+ return nil, err
+ }
+ if len(cmd.Remark) > 0 {
+ if _, err := buf.Write([]byte(cmd.Remark)); err != nil {
+ return nil, err
+ }
+ }
+ if err := binary.Write(buf, binary.BigEndian, int32(len(extBytes)));
err != nil {
+ return nil, err
+ }
+ if len(extBytes) > 0 {
+ if _, err := buf.Write(extBytes); err != nil {
+ return nil, err
+ }
+ }
+
+ return buf.Bytes(), nil
+}
+
+func encodeExtFields(fields map[string]string) ([]byte, error) {
+ if len(fields) == 0 {
+ return []byte{}, nil
+ }
+ keys := make([]string, 0, len(fields))
+ for k := range fields {
+ keys = append(keys, k)
+ }
+ sort.Strings(keys)
+
+ buf := bytes.NewBuffer(nil)
+ for _, key := range keys {
+ value := fields[key]
+ if err := binary.Write(buf, binary.BigEndian, int16(len(key)));
err != nil {
+ return nil, err
+ }
+ if _, err := buf.Write([]byte(key)); err != nil {
+ return nil, err
+ }
+ if err := binary.Write(buf, binary.BigEndian,
int32(len(value))); err != nil {
+ return nil, err
+ }
+ if _, err := buf.Write([]byte(value)); err != nil {
+ return nil, err
+ }
+ }
+ return buf.Bytes(), nil
+}
+
+func markProtocolType(source int32) []byte {
+ result := make([]byte, 4)
+ result[0] = rmqCodecType
+ result[1] = byte((source >> 16) & 0xFF)
+ result[2] = byte((source >> 8) & 0xFF)
+ result[3] = byte(source & 0xFF)
+ return result
+}
+
+type brokerAddrResolver interface {
+ ResolveBrokerAddr(nameServerAddrs []string, topic string, brokerName
string) (string, error)
+}
+
+type tcpSender interface {
+ Send(addr string, data []byte, timeout time.Duration) error
+}
+
+type defaultBrokerAddrResolver struct{}
+
+func (r *defaultBrokerAddrResolver) ResolveBrokerAddr(nameServerAddrs
[]string, topic string, brokerName string) (string, error) {
+ for _, nsAddr := range nameServerAddrs {
+ addr, err := queryBrokerAddrFromNameServer(nsAddr, topic,
brokerName)
+ if err == nil && addr != "" {
+ return addr, nil
+ }
+ }
+ return "", fmt.Errorf("broker %s addr not found from name servers",
brokerName)
+}
+
+type defaultTCPSender struct{}
+
+func (s *defaultTCPSender) Send(addr string, data []byte, timeout
time.Duration) error {
+ conn, err := net.DialTimeout("tcp", addr, timeout)
+ if err != nil {
+ return fmt.Errorf("dial broker %s failed: %w", addr, err)
+ }
+ defer conn.Close()
+
+ if err := conn.SetWriteDeadline(time.Now().Add(timeout)); err != nil {
+ return fmt.Errorf("set write deadline failed: %w", err)
+ }
+ if _, err := conn.Write(data); err != nil {
+ return fmt.Errorf("write to broker %s failed: %w", addr, err)
+ }
+
+ if err := conn.SetReadDeadline(time.Now().Add(timeout)); err != nil {
+ return fmt.Errorf("set read deadline failed: %w", err)
+ }
+ frameLenBuf := make([]byte, 4)
+ if _, err := readFull(conn, frameLenBuf); err != nil {
+ return fmt.Errorf("read response frame length from broker %s
failed: %w", addr, err)
+ }
+ frameLen := int(binary.BigEndian.Uint32(frameLenBuf))
+ if frameLen < 8 {
+ return fmt.Errorf("broker %s response frame too short: %d",
addr, frameLen)
+ }
+ frameBuf := make([]byte, frameLen)
+ if _, err := readFull(conn, frameBuf); err != nil {
+ return fmt.Errorf("read response frame from broker %s failed:
%w", addr, err)
+ }
Review Comment:
`defaultTCPSender.Send` allocates a buffer of size `frameLen` directly from
the broker-provided length. A malformed response could cause very large
allocations (OOM) or other resource issues. Add a sane upper bound before
allocating/reading the frame.
##########
pkg/integration/rocketmq/end_transaction_sender.go:
##########
@@ -0,0 +1,477 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements. See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License. You may obtain a copy of the License at
+ *
+ * http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package rocketmq
+
+import (
+ "bytes"
+ "encoding/binary"
+ "encoding/json"
+ "fmt"
+ "io"
+ "net"
+ "sort"
+ "sync/atomic"
+ "time"
+)
+
+const (
+ reqEndTransaction int16 = 37
+
+ reqGetRouteInfoByTopic int16 = 105
+
+ rmqProtocolVersion int16 = 317
+
+ rmqLanguageGo byte = 9
+
+ rmqCodecType byte = 1
+
+ rmqHeaderFixedLength = 21
+
+ commitOrRollbackCommit = 8
+
+ commitOrRollbackRollback = 12
+)
+
+var opaqueCounter int32
+
+type endTransactionRequestHeader struct {
+ Topic string
+ ProducerGroup string
+ TranStateTableOffset int64
+ CommitLogOffset int64
+ CommitOrRollback int
+ FromTransactionCheck bool
+ MsgID string
+ TransactionId string
+}
+
+func (h *endTransactionRequestHeader) Encode() map[string]string {
+ return map[string]string{
+ "topic": h.Topic,
+ "producerGroup": h.ProducerGroup,
+ "tranStateTableOffset": fmt.Sprintf("%d",
h.TranStateTableOffset),
+ "commitLogOffset": fmt.Sprintf("%d", h.CommitLogOffset),
+ "commitOrRollback": fmt.Sprintf("%d", h.CommitOrRollback),
+ "fromTransactionCheck": fmt.Sprintf("%v",
h.FromTransactionCheck),
+ "msgId": h.MsgID,
+ "transactionId": h.TransactionId,
+ }
+}
+
+type remotingCommand struct {
+ Code int16
+ Language byte
+ Version int16
+ Opaque int32
+ Flag int32
+ Remark string
+ ExtFields map[string]string
+ Body []byte
+}
+
+func newEndTransactionCommand(header *endTransactionRequestHeader)
*remotingCommand {
+ return &remotingCommand{
+ Code: reqEndTransaction,
+ Language: rmqLanguageGo,
+ Version: rmqProtocolVersion,
+ Opaque: atomic.AddInt32(&opaqueCounter, 1),
+ ExtFields: header.Encode(),
+ }
+}
+
+func (cmd *remotingCommand) encode() ([]byte, error) {
+ headerBytes, err := cmd.encodeHeader()
+ if err != nil {
+ return nil, fmt.Errorf("encode header failed: %w", err)
+ }
+
+ frameSize := 4 + len(headerBytes) + len(cmd.Body)
+ buf := bytes.NewBuffer(make([]byte, 0, 4+frameSize))
+
+ if err := binary.Write(buf, binary.BigEndian, int32(frameSize)); err !=
nil {
+ return nil, err
+ }
+ if err := binary.Write(buf, binary.BigEndian,
markProtocolType(int32(len(headerBytes)))); err != nil {
+ return nil, err
+ }
+ if _, err := buf.Write(headerBytes); err != nil {
+ return nil, err
+ }
+ if len(cmd.Body) > 0 {
+ if _, err := buf.Write(cmd.Body); err != nil {
+ return nil, err
+ }
+ }
+
+ return buf.Bytes(), nil
+}
+
+func (cmd *remotingCommand) encodeHeader() ([]byte, error) {
+ extBytes, err := encodeExtFields(cmd.ExtFields)
+ if err != nil {
+ return nil, err
+ }
+
+ buf := bytes.NewBuffer(make([]byte, 0,
rmqHeaderFixedLength+len(cmd.Remark)+len(extBytes)))
+
+ if err := binary.Write(buf, binary.BigEndian, cmd.Code); err != nil {
+ return nil, err
+ }
+ if err := buf.WriteByte(cmd.Language); err != nil {
+ return nil, err
+ }
+ if err := binary.Write(buf, binary.BigEndian, cmd.Version); err != nil {
+ return nil, err
+ }
+ if err := binary.Write(buf, binary.BigEndian, cmd.Opaque); err != nil {
+ return nil, err
+ }
+ if err := binary.Write(buf, binary.BigEndian, cmd.Flag); err != nil {
+ return nil, err
+ }
+ if err := binary.Write(buf, binary.BigEndian, int32(len(cmd.Remark)));
err != nil {
+ return nil, err
+ }
+ if len(cmd.Remark) > 0 {
+ if _, err := buf.Write([]byte(cmd.Remark)); err != nil {
+ return nil, err
+ }
+ }
+ if err := binary.Write(buf, binary.BigEndian, int32(len(extBytes)));
err != nil {
+ return nil, err
+ }
+ if len(extBytes) > 0 {
+ if _, err := buf.Write(extBytes); err != nil {
+ return nil, err
+ }
+ }
+
+ return buf.Bytes(), nil
+}
+
+func encodeExtFields(fields map[string]string) ([]byte, error) {
+ if len(fields) == 0 {
+ return []byte{}, nil
+ }
+ keys := make([]string, 0, len(fields))
+ for k := range fields {
+ keys = append(keys, k)
+ }
+ sort.Strings(keys)
+
+ buf := bytes.NewBuffer(nil)
+ for _, key := range keys {
+ value := fields[key]
+ if err := binary.Write(buf, binary.BigEndian, int16(len(key)));
err != nil {
+ return nil, err
+ }
+ if _, err := buf.Write([]byte(key)); err != nil {
+ return nil, err
+ }
+ if err := binary.Write(buf, binary.BigEndian,
int32(len(value))); err != nil {
+ return nil, err
+ }
+ if _, err := buf.Write([]byte(value)); err != nil {
+ return nil, err
+ }
+ }
+ return buf.Bytes(), nil
+}
+
+func markProtocolType(source int32) []byte {
+ result := make([]byte, 4)
+ result[0] = rmqCodecType
+ result[1] = byte((source >> 16) & 0xFF)
+ result[2] = byte((source >> 8) & 0xFF)
+ result[3] = byte(source & 0xFF)
+ return result
+}
+
+type brokerAddrResolver interface {
+ ResolveBrokerAddr(nameServerAddrs []string, topic string, brokerName
string) (string, error)
+}
+
+type tcpSender interface {
+ Send(addr string, data []byte, timeout time.Duration) error
+}
+
+type defaultBrokerAddrResolver struct{}
+
+func (r *defaultBrokerAddrResolver) ResolveBrokerAddr(nameServerAddrs
[]string, topic string, brokerName string) (string, error) {
+ for _, nsAddr := range nameServerAddrs {
+ addr, err := queryBrokerAddrFromNameServer(nsAddr, topic,
brokerName)
+ if err == nil && addr != "" {
+ return addr, nil
+ }
+ }
+ return "", fmt.Errorf("broker %s addr not found from name servers",
brokerName)
+}
+
+type defaultTCPSender struct{}
+
+func (s *defaultTCPSender) Send(addr string, data []byte, timeout
time.Duration) error {
+ conn, err := net.DialTimeout("tcp", addr, timeout)
+ if err != nil {
+ return fmt.Errorf("dial broker %s failed: %w", addr, err)
+ }
+ defer conn.Close()
+
+ if err := conn.SetWriteDeadline(time.Now().Add(timeout)); err != nil {
+ return fmt.Errorf("set write deadline failed: %w", err)
+ }
+ if _, err := conn.Write(data); err != nil {
+ return fmt.Errorf("write to broker %s failed: %w", addr, err)
+ }
+
+ if err := conn.SetReadDeadline(time.Now().Add(timeout)); err != nil {
+ return fmt.Errorf("set read deadline failed: %w", err)
+ }
+ frameLenBuf := make([]byte, 4)
+ if _, err := readFull(conn, frameLenBuf); err != nil {
+ return fmt.Errorf("read response frame length from broker %s
failed: %w", addr, err)
+ }
+ frameLen := int(binary.BigEndian.Uint32(frameLenBuf))
+ if frameLen < 8 {
+ return fmt.Errorf("broker %s response frame too short: %d",
addr, frameLen)
+ }
+ frameBuf := make([]byte, frameLen)
+ if _, err := readFull(conn, frameBuf); err != nil {
+ return fmt.Errorf("read response frame from broker %s failed:
%w", addr, err)
+ }
+
+ code := int16(binary.BigEndian.Uint16(frameBuf[4:6]))
+ if code != 0 {
+ return fmt.Errorf("broker %s rejected END_TRANSACTION,
responseCode=%d", addr, code)
+ }
+ return nil
+}
+
+func sendEndTransaction(
+ nameServerAddrs []string,
+ topic string,
+ brokerName string,
+ header *endTransactionRequestHeader,
+ timeout time.Duration,
+ resolver brokerAddrResolver,
+ sender tcpSender,
+) error {
+ brokerAddr, err := resolver.ResolveBrokerAddr(nameServerAddrs, topic,
brokerName)
+ if err != nil {
+ return fmt.Errorf("resolve broker addr failed: %w", err)
+ }
+
+ cmd := newEndTransactionCommand(header)
+
+ data, err := cmd.encode()
+ if err != nil {
+ return fmt.Errorf("encode command failed: %w", err)
+ }
+
+ if err := sender.Send(brokerAddr, data, timeout); err != nil {
+ return fmt.Errorf("send to broker %s failed: %w", brokerAddr,
err)
+ }
+
+ return nil
+}
+
+func queryBrokerAddrFromNameServer(nameServerAddr string, topic string,
brokerName string) (string, error) {
+ cmd := &remotingCommand{
+ Code: reqGetRouteInfoByTopic,
+ Language: rmqLanguageGo,
+ Version: rmqProtocolVersion,
+ Opaque: atomic.AddInt32(&opaqueCounter, 1),
+ ExtFields: map[string]string{
+ "topic": topic,
+ },
+ }
+
+ data, err := cmd.encode()
+ if err != nil {
+ return "", fmt.Errorf("encode route request failed: %w", err)
+ }
+
+ conn, err := net.DialTimeout("tcp", nameServerAddr, 3*time.Second)
+ if err != nil {
+ return "", fmt.Errorf("dial name server %s failed: %w",
nameServerAddr, err)
+ }
+ defer conn.Close()
+
+ if err := conn.SetWriteDeadline(time.Now().Add(3 * time.Second)); err
!= nil {
+ return "", err
+ }
+ if _, err := conn.Write(data); err != nil {
+ return "", fmt.Errorf("send route request failed: %w", err)
+ }
+
+ if err := conn.SetReadDeadline(time.Now().Add(3 * time.Second)); err !=
nil {
+ return "", err
+ }
+ frameLenBuf := make([]byte, 4)
+ if _, err := readFull(conn, frameLenBuf); err != nil {
+ return "", fmt.Errorf("read frame length failed: %w", err)
+ }
+ frameLen := int(binary.BigEndian.Uint32(frameLenBuf))
+ if frameLen < 8 {
+ return "", fmt.Errorf("name server %s response frame too short:
%d", nameServerAddr, frameLen)
+ }
+
+ frameBuf := make([]byte, frameLen)
+ if _, err := readFull(conn, frameBuf); err != nil {
+ return "", fmt.Errorf("read frame data failed: %w", err)
+ }
+
+ if len(frameBuf) < 4 {
+ return "", fmt.Errorf("response frame too short")
+ }
+ oriHeaderLen := binary.BigEndian.Uint32(frameBuf[0:4])
+ headerLen := int(oriHeaderLen & 0xFFFFFF)
+ if len(frameBuf) < 4+headerLen {
+ return "", fmt.Errorf("response header truncated")
+ }
+
+ respCode := int16(binary.BigEndian.Uint16(frameBuf[4:6]))
+ if respCode != 0 {
+ return "", fmt.Errorf("name server %s returned error code %d
for route query", nameServerAddr, respCode)
+ }
+
+ _, err = decodeResponseHeader(frameBuf[4 : 4+headerLen])
+ if err != nil {
+ return "", fmt.Errorf("decode response header failed: %w", err)
+ }
+
+ bodyStart := 4 + headerLen
+ if bodyStart >= len(frameBuf) {
+ return "", fmt.Errorf("response has no body")
+ }
+ body := frameBuf[bodyStart:]
+
+ return parseBrokerAddrFromRouteBody(body, brokerName)
+}
+
+func readFull(conn net.Conn, buf []byte) (int, error) {
+ return io.ReadFull(conn, buf)
+}
+
+func decodeResponseHeader(data []byte) (map[string]string, error) {
+ buf := bytes.NewReader(data)
+
+ if buf.Len() < 13 {
+ return nil, fmt.Errorf("header too short for fixed fields")
+ }
+ discard := make([]byte, 13)
+ if _, err := buf.Read(discard); err != nil {
+ return nil, err
+ }
+
+ var remarkLen int32
+ if err := binary.Read(buf, binary.BigEndian, &remarkLen); err != nil {
+ return nil, err
+ }
+ if remarkLen > 0 {
+ discardRemark := make([]byte, remarkLen)
+ if _, err := buf.Read(discardRemark); err != nil {
+ return nil, err
+ }
+ }
+
+ var extLen int32
+ if err := binary.Read(buf, binary.BigEndian, &extLen); err != nil {
+ return nil, err
+ }
+ extFields := make(map[string]string)
+ if extLen > 0 {
+ extData := make([]byte, extLen)
+ if _, err := buf.Read(extData); err != nil {
+ return nil, err
+ }
+ extBuf := bytes.NewReader(extData)
+ for extBuf.Len() > 0 {
+ var kLen int16
+ if err := binary.Read(extBuf, binary.BigEndian, &kLen);
err != nil {
+ break
+ }
+ key := make([]byte, kLen)
+ if _, err := extBuf.Read(key); err != nil {
+ break
+ }
+ var vLen int32
+ if err := binary.Read(extBuf, binary.BigEndian, &vLen);
err != nil {
+ break
+ }
+ value := make([]byte, vLen)
+ if _, err := extBuf.Read(value); err != nil {
+ break
+ }
+ extFields[string(key)] = string(value)
+ }
Review Comment:
The ext-fields parsing loop in `decodeResponseHeader` can panic on negative
lengths (`kLen`/`vLen`) or silently ignore truncated data by `break`ing on
errors. It’s safer to validate lengths against remaining buffer and return an
error on malformed headers.
--
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.
To unsubscribe, e-mail: [email protected]
For queries about this service, please contact Infrastructure at:
[email protected]
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]