You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Proto消息递归截断复制问题:嵌套消息列表处理失败求助

问题:Go中复制Proto消息并截断嵌套列表时递归报错

我正在编写一个Go方法,接收大型Proto消息,复制到同类型消息中并截断指定属性的长度。比如原消息中PhoneNumbers数组有100个值,希望新消息只保留前50个。方法接收一个map,指定各属性保留的元素数量。

Proto定义如下:

syntax = "proto2";

package main;

option go_package = "temp/src";

message Person {
  optional string name = 1;
  optional int32 age = 2;
  repeated int32 departmentIds = 3;
  repeated Address addresses = 4;
}

message Address{
  optional string address = 1;
  repeated PhoneNumber phoneNumbers = 3;
}

message PhoneNumber {
  optional int32 countryCode = 1;
  optional int32 stateCode = 2;
  optional string number = 3;
}

当前代码对顶层属性处理正常,但处理嵌套消息列表时失败,递归中value := orgMsg.ProtoReflect().Get(fd)报错,原因是orgMsg没有对应的子FieldDescriptor值。代码如下:

func main() {
        //Create Person instance
    attributeCount := map[string]int{
        "departmentIds": 5,
        "addresses/phoneNumbers":     3,
    }
    newPerson := &Person{}
    stripMessage(person, newPerson, attributeCount)
    fmt.Println(newPerson)
}

func stripMessage(orgMsg proto.Message, newMsg proto.Message, attributeToCount map[string]int) error {
    if len(attributeToCount) == 0 {
        proto.Merge(newMsg, orgMsg)
        return nil
    }

    mds := newMsg.ProtoReflect().Descriptor()
    dm, err := NewDynamicProtoRand(orgMsg, mds, attributeToCount)
    if err != nil {
        return nil
    }
    proto.Merge(newMsg, dm)
    return nil
}

func getValue(orgMsg proto.Message, fd protoreflect.FieldDescriptor, attributeCount map[string]int) (protoreflect.Value, error) {
    // process recursively
    rm, err := NewDynamicProtoRand(orgMsg, fd.Message(), attributeCount)
    if err != nil {
        return protoreflect.Value{}, err
    }
    return protoreflect.ValueOfMessage(rm), nil
}

// NewDynamicProtoRand created dynamicpb with assiging random value to proto
func NewDynamicProtoRand(orgMsg proto.Message, mds protoreflect.MessageDescriptor, attributeCount map[string]int) (*dynamicpb.Message, error) {
    dm := dynamicpb.NewMessage(mds)
    fds := mds.Fields()
    for k := 0; k < fds.Len(); k++ {
        fd := fds.Get(k)

        if fd.IsList() {
            list := dm.Mutable(fd).List()
            if fd.Kind() == protoreflect.MessageKind {
                getValue(orgMsg, fd, attributeCount)
            } else {
                values := orgMsg.ProtoReflect().Get(fd).List()
                var lenForAttribute int
                var attributePresent bool
                valuesLen := values.Len()
                for mapKey, mapValue := range attributeCount {
                    if mapKey == fd.TextName() {
                        lenForAttribute = mapValue
                        attributePresent = true
                        break
                    }
                }
                if lenForAttribute > valuesLen {
                    log.Fatalf("Cannot truncate list provided len:%d is greater then lenth:%d", lenForAttribute, valuesLen)
                } else {
                    var requiredValues int
                    if attributePresent {
                        requiredValues = lenForAttribute
                    } else {
                        requiredValues = valuesLen
                    }
                    for l := 0; l < requiredValues; l++ {
                        list.Append(values.Get(l))
                    }
                }
                dm.Set(fd, protoreflect.ValueOfList(list))
                continue
            }
        }

        value := orgMsg.ProtoReflect().Get(fd)

        dm.Set(fd, value)
    }

    return dm, nil
}

解决方案

问题根源

  1. 递归时传递错误的消息实例:处理嵌套消息列表时,直接传递顶层orgMsg给递归函数,导致在子层级试图从顶层消息获取子字段,必然失败。
  2. 未处理嵌套路径匹配:原代码只匹配顶层字段名,无法识别addresses/phoneNumbers这种嵌套路径的截断规则。
  3. 嵌套列表处理逻辑缺失:处理消息类型的列表时,既没有遍历原始列表元素,也没有将递归处理后的子消息添加到新列表中。

修正后的代码

import (
    "fmt"
    "log"
    "strings"

    "google.golang.org/protobuf/proto"
    "google.golang.org/protobuf/reflect/protoreflect"
    "google.golang.org/protobuf/types/dynamicpb"
)

func main() {
    // 示例:创建原始Person实例
    person := &Person{
        Name:          proto.String("John Doe"),
        Age:           proto.Int32(30),
        DepartmentIds: []int32{1, 2, 3, 4, 5, 6, 7},
        Addresses: []*Address{
            {
                Address: proto.String("123 Main St"),
                PhoneNumbers: []*PhoneNumber{
                    {CountryCode: proto.Int32(1), StateCode: proto.Int32(555), Number: proto.String("1234567")},
                    {CountryCode: proto.Int32(1), StateCode: proto.Int32(555), Number: proto.String("7654321")},
                    {CountryCode: proto.Int32(1), StateCode: proto.Int32(555), Number: proto.String("9876543")},
                    {CountryCode: proto.Int32(1), StateCode: proto.Int32(555), Number: proto.String("3456789")},
                },
            },
        },
    }

    attributeCount := map[string]int{
        "departmentIds":          5,
        "addresses/phoneNumbers": 3,
    }
    newPerson := &Person{}
    if err := stripMessage(person, newPerson, attributeCount); err != nil {
        log.Fatal(err)
    }
    fmt.Println(newPerson)
}

func stripMessage(orgMsg proto.Message, newMsg proto.Message, attributeToCount map[string]int) error {
    if len(attributeToCount) == 0 {
        proto.Merge(newMsg, orgMsg)
        return nil
    }

    dm, err := copyAndTruncateMessage(orgMsg.ProtoReflect(), attributeToCount)
    if err != nil {
        return err
    }
    proto.Merge(newMsg, dm)
    return nil
}

// copyAndTruncateMessage 递归复制消息并根据路径规则截断列表
func copyAndTruncateMessage(orgMsg protoreflect.Message, attributeCount map[string]int) (*dynamicpb.Message, error) {
    mds := orgMsg.Descriptor()
    dm := dynamicpb.NewMessage(mds)

    fds := mds.Fields()
    for k := 0; k < fds.Len(); k++ {
        fd := fds.Get(k)
        if !orgMsg.Has(fd) {
            continue // 跳过原始消息中不存在的字段
        }

        if fd.IsList() {
            orgList := orgMsg.Get(fd).List()
            newList := dm.Mutable(fd).List()

            // 获取当前字段对应的截断长度
            truncateLen := getTruncateLength(fd.TextName(), attributeCount, mds.FullName())
            maxLen := orgList.Len()
            if truncateLen > 0 && truncateLen < maxLen {
                maxLen = truncateLen
            }

            if fd.Kind() == protoreflect.MessageKind {
                // 处理消息类型的列表:递归复制每个子消息
                for i := 0; i < maxLen; i++ {
                    orgSubMsg := orgList.Get(i).Message()
                    subMsg, err := copyAndTruncateMessage(orgSubMsg, attributeCount)
                    if err != nil {
                        return nil, err
                    }
                    newList.Append(protoreflect.ValueOfMessage(subMsg))
                }
            } else {
                // 处理基本类型的列表:直接复制前N个元素
                for i := 0; i < maxLen; i++ {
                    newList.Append(orgList.Get(i))
                }
            }
            dm.Set(fd, protoreflect.ValueOfList(newList))
            continue
        }

        // 处理普通字段(非列表)
        if fd.Kind() == protoreflect.MessageKind {
            // 嵌套消息:递归复制
            orgSubMsg := orgMsg.Get(fd).Message()
            subMsg, err := copyAndTruncateMessage(orgSubMsg, attributeCount)
            if err != nil {
                return nil, err
            }
            dm.Set(fd, protoreflect.ValueOfMessage(subMsg))
        } else {
            // 基本类型:直接复制值
            dm.Set(fd, orgMsg.Get(fd))
        }
    }

    return dm, nil
}

// getTruncateLength 根据字段路径获取截断长度
func getTruncateLength(fieldName string, attributeCount map[string]int, parentFullName protoreflect.FullName) int {
    // 优先匹配当前层级的直接字段规则
    if length, ok := attributeCount[fieldName]; ok {
        return length
    }

    // 匹配嵌套路径规则(如addresses/phoneNumbers)
    parentFieldPrefix := fieldName + "/"
    for path, length := range attributeCount {
        if strings.HasPrefix(path, parentFieldPrefix) {
            // 嵌套路径的当前层级不截断,子层级处理剩余路径
            return 0
        }
    }

    // 检查是否有针对当前消息下字段的完整路径规则
    fullFieldPath := string(parentFullName) + "." + fieldName
    for path, length := range attributeCount {
        if path == fullFieldPath {
            return length
        }
    }

    return 0 // 0表示不截断,保留全部
}

关键改进点

  • 递归传递正确的消息实例:处理嵌套消息时,传递当前层级的原始子消息给递归函数,而非顶层消息。
  • 支持嵌套路径解析:getTruncateLength函数处理类似addresses/phoneNumbers的路径,在递归到子消息时自动匹配剩余路径部分。
  • 完善列表处理逻辑:无论是基本类型列表还是消息类型列表,都先获取原始列表,再根据截断规则复制前N个元素;消息类型列表会递归处理每个子消息。
  • 跳过不存在的字段:添加if !orgMsg.Has(fd)判断,避免尝试获取原始消息中不存在的字段导致错误。

内容的提问来源于stack exchange,提问作者prblm_solver

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.09 09:17:09