Python gRPC

常用protobuf数据类型

Protocol Buffers(protobuf)的核心是使用定义清晰的数据类型来构建高效的数据结构。理解这些类型和.proto文件的写法,是使用gRPC的基础。下面我将为你系统地介绍常用数据类型和.proto文件的结构与编写规范。

📊 常用 Protobuf 数据类型

Protobuf 支持多种数据类型,可以满足各种业务场景的需求。它们主要分为以下几类 :

1. 标量类型

这是最基本的数据类型,对应各种编程语言中的基础类型 。

Protobuf 类型 说明 对应 Python 类型
double / float 64位/32位浮点数 float
int32 / int64 使用变长编码的有符号整数。对于负数,效率不如 sintN int
uint32 / uint64 使用变长编码的无符号整数 。 int
sint32 / sint64 使用变长编码的有符号整数。编码负数时比 int32/int64 效率更高 int
fixed32 / fixed64 固定4字节/8字节的无符号整数。如果值经常大于2^28,比 uint32/uint64 更高效 。 int
sfixed32 / sfixed64 固定4字节/8字节的有符号整数 。 int
bool 布尔值 bool
string 字符串,必须包含 UTF-8 编码或 7-bit ASCII 文本 。 str (Unicode)
bytes 字节序列,可以包含任意字节数据 。 bytes

2. 复合类型

这些类型允许你构建更复杂的数据结构 。

  • 枚举(enum:定义一组预定义的常量值 。

    enum DeviceStatus {
      DEVICE_STATUS_UNKNOWN = 0; // 枚举值必须从0开始
      DEVICE_STATUS_ONLINE = 1;
      DEVICE_STATUS_OFFLINE = 2;
    }
    
  • 消息类型(message:用户自定义的复杂数据类型,可以包含任何其他类型(包括其他message),类似于编程语言中的类 。

    message FirmwareInfo {
      string product_id = 1;
      string version = 2;
      int64 size_bytes = 3;
    }
    
  • 嵌套类型:在一个 message 内部定义另一个 message,用于组织紧密相关的数据结构 。

3. 特殊类型

  • 重复字段(repeated:表示该字段可以包含零个或多个值,类似于数组或列表 。

    message UploadResult {
      repeated string error_messages = 1; // 可以包含多个错误消息
    }
    
  • 映射类型(map:定义键值对集合,类似于字典 。

    message UpdateReport {
      map<string, string> metadata = 1; // 键和值都是字符串
    }
    
  • 可选字段与 oneof

    • 在 proto3 语法中,所有字段都是"可选的"。如果一个字段未设置,它会返回一个默认值(如空字符串、0、false)。
    • oneof:表示一组字段中,最多只有一个字段可以被同时设置,常用于处理互斥的选项 。
      message UpdateRequest {
        oneof payload {
          string manifest_url = 1;
          bytes firmware_image = 2;
        }
      }
      

📄 .proto 文件编写规范

了解了数据类型,我们来看看如何在一个.proto文件中定义消息和服务。一个规范的.proto文件通常包含以下几个部分:

1. 版本声明

文件的第一行(非注释)必须指定使用的语法版本。目前主流是 proto3,而更新的特性则在 editions 中演进 。

syntax = "proto3";
// 或
edition = "2023";

2. 包声明(可选)

package 用于避免不同文件中的消息类型命名冲突,尤其在生成特定语言的代码时会作为命名空间 。

package firmware_update;

3. 导入(可选)

可以使用 import 语句引用其他.proto文件中定义的类型,实现代码复用。

import "google/protobuf/timestamp.proto";

4. 定义消息 (message)

这是核心部分,使用上面介绍的各种数据类型定义你的数据结构和字段 。

5. 字段编号

在定义消息时,必须为每个字段赋予一个唯一的编号

  • 作用:这些编号用于在二进制编码中标识字段。一旦使用,就不应更改
  • 范围1536,870,911,其中 1900019999 是预留的 。
  • 优化最频繁使用的字段,建议分配 115 的编号,因为它们编码后只占1个字节,而 162047 则占2个字节 。

6. 定义服务 (service)

如果用于gRPC,需要在此处定义RPC服务接口,指定方法名、请求参数和返回结果 。

service FirmwareService {
  rpc Upload(stream FirmwareChunk) returns (UploadStatus);
  rpc GetVersion(VersionRequest) returns (VersionInfo);
}

7. 注释

使用 ///* ... */ 为你的定义添加清晰的注释,提升代码可读性 。

8. 保留字段 (reserved)

当你需要删除一个字段时,务必将它的编号(和名称)添加到 reserved 列表中。这可以防止未来其他开发者重复使用这些编号,从而避免因数据错乱导致的严重bug 。

message DeprecatedMessage {
  reserved 2, 15, 9 to 11;
  reserved "old_field_name";
}

💡 与固件升级的联系

结合固件升级场景,可以尝试编写一个包含以下核心内容的.proto文件:

  1. 定义一个 FirmwareChunk 消息(可能包含 bytes dataint64 sequence_number)。
  2. 定义一个 UploadStatus 消息(可能包含 int32 codestring message)。
  3. 定义一个 FirmwareService 服务,包含一个客户端流式RPC方法 Upload(stream FirmwareChunk) returns (UploadStatus)
  4. 如果需要查询版本,可以定义 VersionRequest(可能为空)和 VersionInfo 消息,并添加一个一元RPC方法。

📘 Python gRPC详细教程

本教程整合了CSDN博客《gRPC Python 详细入门教程(一)》的核心内容,并针对你关注的流式传输(特别是客户端流式RPC)进行深入解析,辅以固件上传的实战代码示例。

一、环境准备与快速入门

在开始编写代码前,需要安装必要的Python包。

1.1 安装gRPC核心库

pip install grpcio

1.2 安装gRPC工具(包含protoc编译器及插件)

pip install grpcio-tools

1.3 快速体验:运行示例程序

  1. 获取示例代码
    git clone -b v1.74.0 --depth 1 --shallow-submodules https://github.com/grpc/grpc
    cd grpc/examples/python/helloworld
    
  2. 运行服务器
    python greeter_server.py
    
  3. 运行客户端(新终端):
    python greeter_client.py
    
    如果看到"Greeter client received: Hello, you!",则环境搭建成功。

二、核心概念:Protocol Buffers与服务定义

gRPC的核心是使用Protocol Buffers(protobuf)作为接口定义语言(IDL)。

2.1 .proto文件结构

一个典型的.proto文件包含:

  • 消息类型(Message):定义数据结构。
  • 服务(Service):定义RPC方法接口。

示例helloworld.proto):

syntax = "proto3";

// 定义服务
service Greeter {
  rpc SayHello (HelloRequest) returns (HelloReply) {}
}

// 定义请求消息
message HelloRequest {
  string name = 1;
}

// 定义响应消息
message HelloReply {
  string message = 1;
}

2.2 四种服务方法类型

gRPC支持四种RPC类型,这在处理不同业务场景时非常关键:

类型 语法 描述 应用场景
一元 RPC (Unary) rpc Method(Request) returns (Response) 客户端发送一个请求,服务器返回一个响应。 简单的查询、状态获取(如GetSWInfos)。
服务器流式 RPC rpc Method(Request) returns (stream Response) 客户端发送一个请求,服务器返回一个响应流。 服务器推送大量数据,如地图路线点列表。
客户端流式 RPC rpc Method(stream Request) returns (Response) 客户端发送一个请求流,服务器返回一个响应。 大文件上传(如固件升级)、数据采集。
双向流式 RPC rpc Method(stream Request) returns (stream Response) 双方使用流同时发送和接收消息。 实时聊天、游戏状态同步。

你的固件升级功能正属于客户端流式RPC


三、从.proto生成Python代码

使用grpcio-tools提供的protoc编译器生成Python代码。

3.1 基本生成命令

python -m grpc_tools.protoc \
    -I../../protos \                # 指定.proto文件搜索路径
    --python_out=. \                 # 生成消息类的输出目录
    --grpc_python_out=. \            # 生成gRPC服务类的输出目录
    ../../protos/route_guide.proto   # 要编译的.proto文件

执行后会生成两个核心文件:

  • route_guide_pb2.py:包含protobuf消息类。
  • route_guide_pb2_grpc.py:包含gRPC服务相关的Stub(客户端存根)Servicer(服务端抽象基类)

3.2 代码生成与导入逻辑

生成的*_pb2_grpc.py文件会自动导入*_pb2文件中的消息类,因此在客户端/服务端代码中,你只需导入*_pb2_grpc即可。


四、实现Python gRPC服务器

以CSDN博客中的RouteGuide服务为例,展示如何实现各类RPC方法。

4.1 创建Servicer子类

你需要继承生成的RouteGuideServicer基类,并实现所有RPC方法。

# route_guide_server.py (简化版)
import route_guide_pb2
import route_guide_pb2_grpc

class RouteGuideServicer(route_guide_pb2_grpc.RouteGuideServicer):
    # 1. 实现一元RPC: GetFeature
    def GetFeature(self, request, context):
        # request 是 Point 对象
        feature = get_feature_from_db(request.latitude, request.longitude)
        if feature is None:
            return route_guide_pb2.Feature(name="", location=request)
        return feature

    # 2. 实现客户端流式RPC: RecordRoute
    def RecordRoute(self, request_iterator, context):
        # request_iterator 是一个迭代器,用于逐个获取客户端发来的Point
        point_count = 0
        for point in request_iterator:
            point_count += 1
            # 处理每个point,例如记录距离、时间等
            pass
        # 返回一个RouteSummary响应
        return route_guide_pb2.RouteSummary(point_count=point_count)

    # 3. 实现双向流式RPC: RouteChat
    def RouteChat(self, request_iterator, context):
        # 对于客户端发来的每个RouteNote,回复一个或多个RouteNote
        for note in request_iterator:
            # 处理收到的note,并可能发送回复
            yield route_guide_pb2.RouteNote(message=f"Echo: {note.message}", location=note.location)

4.2 启动gRPC服务器

def serve():
    server = grpc.server(futures.ThreadPoolExecutor(max_workers=10))
    route_guide_pb2_grpc.add_RouteGuideServicer_to_server(
        RouteGuideServicer(), server)
    server.add_insecure_port('[::]:50051')
    server.start()
    server.wait_for_termination()

五、实现Python gRPC客户端

客户端通过Stub调用远程方法。

5.1 创建客户端Stub

import grpc
import route_guide_pb2
import route_guide_pb2_grpc

channel = grpc.insecure_channel('localhost:50051')
stub = route_guide_pb2_grpc.RouteGuideStub(channel)

5.2 调用不同类型RPC

一元RPC调用

point = route_guide_pb2.Point(latitude=409146138, longitude=-746188906)
feature = stub.GetFeature(point)
print(feature.name)

客户端流式RPC调用(核心:生成器)
这是你固件升级功能的核心模式。通过一个生成器函数来产生请求流。

def generate_points():
    # 模拟从文件或数据源逐个产生Point
    points = [
        route_guide_pb2.Point(latitude=407838351, longitude=-746143763),
        route_guide_pb2.Point(latitude=408122808, longitude=-743999179),
    ]
    for point in points:
        print(f"Sending point: {point}")
        yield point
    # 生成器结束,自动通知服务器流结束

# 调用客户端流式RPC,将生成器作为参数传入
summary = stub.RecordRoute(generate_points())
print(f"Route summary: {summary}")

双向流式RPC调用
同时使用生成器发送,并通过迭代响应流接收。

def generate_notes():
    notes = [
        route_guide_pb2.RouteNote(message="First", location=point1),
        route_guide_pb2.RouteNote(message="Second", location=point2),
    ]
    for note in notes:
        yield note

# 调用双向流式RPC,返回的也是一个迭代器
responses = stub.RouteChat(generate_notes())
for response in responses:
    print(f"Received echo: {response.message}")

六、gRPC-Python四种服务方法

gRPC 支持四种服务方法类型,它们基于客户端-服务器交互模式的不同,能够灵活应对各种分布式系统通信需求。本文将通过一个完整的固件升级系统案例,逐一剖析这四种类型在 Python 中的实现,并附上可运行的代码片段,帮助你快速掌握 gRPC 的核心用法。

1. 定义统一的 .proto 文件

为了体现四种类型在同一个系统中的应用,我们设计一个固件升级服务的 .proto 文件:

syntax = "proto3";

package firmware;

// 固件升级服务
service FirmwareService {
  // 一元 RPC:获取当前版本信息
  rpc GetVersion(VersionRequest) returns (VersionInfo);

  // 服务器流式 RPC:下载升级日志
  rpc DownloadLogs(LogRequest) returns (stream LogEntry);

  // 客户端流式 RPC:上传固件文件
  rpc UploadFirmware(stream FirmwareChunk) returns (UploadStatus);

  // 双向流式 RPC:实时监控升级状态
  rpc MonitorUpgrade(stream ControlCommand) returns (stream ProgressReport);
}

// 一元 RPC 的消息定义
message VersionRequest {}
message VersionInfo {
  string version = 1;
  string build_time = 2;
}

// 服务器流式 RPC 的消息定义
message LogRequest {
  int32 max_lines = 1;  // 最多返回的行数
}
message LogEntry {
  string line = 1;
  int64 timestamp = 2;
}

// 客户端流式 RPC 的消息定义(固件上传)
message FirmwareChunk {
  bytes data = 1;
  string product_id = 2;  // 第一个块携带产品 ID
}
message UploadStatus {
  int32 code = 1;
  string message = 2;
}

// 双向流式 RPC 的消息定义
message ControlCommand {
  enum Command {
    START = 0;
    PAUSE = 1;
    RESUME = 2;
    CANCEL = 3;
  }
  Command cmd = 1;
}
message ProgressReport {
  int32 percent = 1;           // 升级进度百分比
  string stage = 2;            // 当前阶段(校验、烧录、重启等)
  string message = 3;          // 附加信息
}

使用以下命令生成 Python 代码:

python -m grpc_tools.protoc -I. --python_out=. --grpc_python_out=. firmware.proto

生成的文件:firmware_pb2.py(消息类)和 firmware_pb2_grpc.py(服务类与存根)。


2. 一元 RPC:获取版本信息

2.1 服务端实现

import grpc
import firmware_pb2
import firmware_pb2_grpc
from concurrent import futures

class FirmwareServicer(firmware_pb2_grpc.FirmwareServiceServicer):
    def GetVersion(self, request, context):
        # 模拟从设备获取版本信息
        return firmware_pb2.VersionInfo(
            version="v2.1.0",
            build_time="2026-02-28 10:00:00"
        )

def serve():
    server = grpc.server(futures.ThreadPoolExecutor(max_workers=10))
    firmware_pb2_grpc.add_FirmwareServiceServicer_to_server(FirmwareServicer(), server)
    server.add_insecure_port('[::]:50051')
    print("Server started on port 50051")
    server.start()
    server.wait_for_termination()

if __name__ == '__main__':
    serve()

2.2 客户端调用

import grpc
import firmware_pb2
import firmware_pb2_grpc

channel = grpc.insecure_channel('localhost:50051')
stub = firmware_pb2_grpc.FirmwareServiceStub(channel)

response = stub.GetVersion(firmware_pb2.VersionRequest())
print(f"Device version: {response.version}, built at {response.build_time}")

要点:一元 RPC 是最简单的模式,直接调用方法并等待响应,超时可通过 timeout 参数设置。


3. 服务器流式 RPC:下载升级日志

3.1 服务端实现

class FirmwareServicer(firmware_pb2_grpc.FirmwareServiceServicer):
    # ... 其他方法

    def DownloadLogs(self, request, context):
        # 模拟日志数据
        logs = [
            ("2026-02-28 10:05:01", "Starting firmware update"),
            ("2026-02-28 10:05:03", "Validating firmware image"),
            ("2026-02-28 10:05:05", "Burning flash..."),
            ("2026-02-28 10:05:10", "Rebooting device"),
            ("2026-02-28 10:06:00", "Device online, new firmware v2.1.0"),
        ]
        # 根据请求限制返回行数
        for i, (timestamp, line) in enumerate(logs):
            if i >= request.max_lines:
                break
            # 将时间字符串转换为时间戳(简单示例)
            # 实际可使用 datetime 转换
            yield firmware_pb2.LogEntry(
                line=line,
                timestamp=0  # 简化处理
            )

3.2 客户端调用

def download_logs(stub):
    request = firmware_pb2.LogRequest(max_lines=10)
    try:
        for log_entry in stub.DownloadLogs(request):
            print(f"[LOG] {log_entry.line}")
    except grpc.RpcError as e:
        print(f"Error downloading logs: {e.code()}")

要点:服务端方法变为生成器,每次 yield 一条日志;客户端使用 for 循环迭代接收。网络错误通过异常捕获。


4. 客户端流式 RPC:固件上传

4.1 服务端实现

class FirmwareServicer(firmware_pb2_grpc.FirmwareServiceServicer):
    def UploadFirmware(self, request_iterator, context):
        # request_iterator 是一个迭代器,逐个接收客户端发送的块
        total_bytes = 0
        product_id = ""
        for chunk in request_iterator:
            if not product_id and chunk.product_id:
                product_id = chunk.product_id  # 第一个块携带产品ID
            total_bytes += len(chunk.data)
            # 可以在这里进行校验、写入临时文件等操作
        print(f"Received firmware for {product_id}, total {total_bytes} bytes")
        # 返回上传结果
        return firmware_pb2.UploadStatus(code=0, message="Upload successful")

4.2 客户端调用

import os

def upload_firmware(stub, file_path, product_id):
    def chunk_generator():
        # 发送产品ID(第一个块)
        first_chunk = firmware_pb2.FirmwareChunk(product_id=product_id, data=b'')
        yield first_chunk

        with open(file_path, 'rb') as f:
            while True:
                data = f.read(64 * 1024)  # 64KB
                if not data:
                    break
                yield firmware_pb2.FirmwareChunk(data=data)

    try:
        response = stub.UploadFirmware(chunk_generator(), timeout=600)
        print(f"Upload result: code={response.code}, message={response.message}")
        return response.code == 0
    except grpc.RpcError as e:
        print(f"Upload failed: {e.code()}")
        return False

要点:客户端通过生成器产生流式请求;服务端通过 request_iterator 迭代接收。超时设置为 10 分钟以应对大文件传输。


5. 双向流式 RPC:实时监控升级状态

5.1 服务端实现

import time

class FirmwareServicer(firmware_pb2_grpc.FirmwareServiceServicer):
    def MonitorUpgrade(self, request_iterator, context):
        # 该生成器接收客户端的控制命令,同时发送进度报告
        for cmd in request_iterator:
            # 处理控制命令
            if cmd.cmd == firmware_pb2.ControlCommand.START:
                # 开始模拟升级过程
                for percent in range(0, 101, 10):
                    # 检查客户端是否取消或暂停
                    # 实际应用中可能需要更复杂的同步机制
                    yield firmware_pb2.ProgressReport(
                        percent=percent,
                        stage="burning" if percent < 80 else "rebooting",
                        message=f"Progress: {percent}%"
                    )
                    time.sleep(1)  # 模拟耗时
            elif cmd.cmd == firmware_pb2.ControlCommand.CANCEL:
                print("Upgrade cancelled by client")
                break  # 停止发送
            # 可以处理 PAUSE/RESUME 等

5.2 客户端调用

def monitor_upgrade(stub):
    def command_generator():
        # 发送开始命令
        yield firmware_pb2.ControlCommand(cmd=firmware_pb2.ControlCommand.START)
        # 假设 5 秒后取消升级
        time.sleep(5)
        yield firmware_pb2.ControlCommand(cmd=firmware_pb2.ControlCommand.CANCEL)

    try:
        responses = stub.MonitorUpgrade(command_generator())
        for report in responses:
            print(f"Progress: {report.percent}% - {report.stage} - {report.message}")
            if report.percent >= 100:
                print("Upgrade completed.")
    except grpc.RpcError as e:
        print(f"Monitor error: {e.code()}")

要点:双向流式 RPC 允许双方独立发送消息。客户端使用生成器发送命令,同时通过迭代接收服务器的进度报告;服务器在迭代客户端命令的同时可以随时 yield 响应。需要注意并发控制和取消机制。


6. 四种类型对比总结

特性 一元 RPC 服务器流式 客户端流式 双向流式
请求 1 1 多个(流) 多个(流)
响应 1 多个(流) 1 多个(流)
客户端实现 普通函数调用 迭代接收 生成器发送 生成器发送 + 迭代接收
服务端实现 普通函数 生成器发送 迭代器接收 迭代器接收 + 生成器发送
典型场景 查询、配置 日志下载、数据推送 文件上传、数据采集 实时交互、控制台

通过这四种类型的灵活组合,gRPC 能够轻松应对从简单请求到复杂流式交互的各种分布式通信需求。在你的固件升级系统中,可以同时使用这四种模式,构建一个功能完备、交互丰富的升级服务。


七、实战解析:固件升级的客户端流式RPC

结合你的具体需求,我们将固件上传流程映射到Python gRPC的客户端流式RPC模式。

7.1 假设的.proto定义

service FirmwareUpdate {
  // 客户端流式RPC:上传固件
  rpc UploadFirmware(stream FirmwareChunk) returns (UploadStatus);
}

message FirmwareChunk {
  bytes data = 1;          // 固件数据块
  string product_id = 2;   // 可选:第一个块携带产品ID
}

message UploadStatus {
  int32 code = 1;
  string message = 2;
}

7.2 Python客户端实现(带超时和进度)

import grpc
import firmware_pb2
import firmware_pb2_grpc
import os

def upload_firmware(file_path, product_id, server_addr='192.168.2.61:50050'):
    # 1. 创建channel和stub
    channel = grpc.insecure_channel(server_addr)
    stub = firmware_pb2_grpc.FirmwareUpdateStub(channel)

    # 2. 定义生成器函数:分块读取文件并yield
    def chunk_generator():
        file_size = os.path.getsize(file_path)
        sent = 0
        # 第一个块可以携带元信息,例如产品ID
        first_chunk = firmware_pb2.FirmwareChunk(product_id=product_id, data=b'')
        yield first_chunk

        with open(file_path, 'rb') as f:
            while True:
                chunk_data = f.read(1024 * 1024)  # 1MB per chunk
                if not chunk_data:
                    break
                sent += len(chunk_data)
                progress = (sent / file_size) * 100
                print(f"Progress: {progress:.1f}%")
                yield firmware_pb2.FirmwareChunk(data=chunk_data)

    # 3. 调用流式RPC,设置超时
    try:
        response = stub.UploadFirmware(
            chunk_generator(),
            timeout=600  # 10分钟超时,适用于大文件
        )
        if response.code == 0:
            print("Upload successful:", response.message)
        else:
            print("Upload failed:", response.message)
    except grpc.RpcError as e:
        print(f"gRPC error: {e.code()} - {e.details()}")

关键点阐述:

  1. 生成器即请求流chunk_generator函数每次yield一个FirmwareChunk消息,gRPC运行时负责将其打包成HTTP/2 DATA帧发送。
  2. 流结束信号:当生成器函数执行完毕(文件读完),gRPC自动发送半关闭信号,通知服务器客户端已无更多数据。
  3. 超时控制timeout参数为整个RPC调用设置了截止时间,包含所有数据块的发送时间和服务器处理时间。
  4. 错误处理:通过捕获grpc.RpcError可以处理网络问题、超时(DEADLINE_EXCEEDED)或服务器返回的错误。

八、总结与对比(Python vs C++)

特性 Python gRPC C++ gRPC (基于你的博客)
请求流实现 生成器 (Generator)。通过yield产生消息,隐式控制流。 ClientWriter对象。显式调用Write()WritesDone()Finish()
代码风格 简洁、声明式,符合Python惯用法。 控制精细、显式,符合C++ RAII理念。
错误处理 主要使用异常grpc.RpcError)。 检查每个Write()的返回值,通过Finish()返回的Status判断。
超时设置 在Stub方法调用时直接传入timeout参数。 通过ClientContextset_deadline()方法设置。
资源管理 生成器内的with语句自动管理文件;gRPC资源通过引用计数管理。 智能指针管理ClientWriter,RAII管理资源。

核心结论:尽管实现风格迥异,但两者底层都基于相同的gRPC C Core,因此网络传输、流控、消息序列化等机制完全一致。Python的生成器模型让流式编程更加直观,非常适合快速开发和测试工具,与你固件升级客户端的定位完美契合。

希望这份详尽的阐述能帮助你深入理解Python gRPC,特别是客户端流式RPC的实现原理。如果后续在编码过程中遇到具体问题,欢迎随时继续探讨。

firmware_client.py
#!/usr/bin/env python3
"""
Firmware Update Client 类定义
"""
import os
import sys
import time
import hashlib
import traceback
import logging
from typing import Generator, List, Optional, Dict, Any
from datetime import datetime

try:
    import grpc
except ImportError:
    print("[ERROR] 需要安装 grpc: pip3 install grpcio grpcio-tools")
    sys.exit(1)

# 添加 proto_python 目录到 Python 路径
current_dir = os.path.dirname(os.path.abspath(__file__))
sys.path.insert(0, os.path.join(current_dir, 'proto_python'))

try:
    import test_interface_pb2 as pb2
    import service_ms_test_interface_pb2_grpc as pb2_grpc
    import common_enums_pb2
    import type_pb2
    from google.protobuf.timestamp_pb2 import Timestamp
    from google.protobuf import empty_pb2
    print("[SUCCESS] 导入 protobuf 模块成功")
except ImportError as e:
    print(f"[ERROR] 导入失败: {e}")
    sys.exit(1)


class CompleteFirmwareClient:
    """完整固件升级客户端 - 修复版"""
    
    # 常量定义
    FIRMWARE_CHUNK_SIZE = 1024 * 1024  # 1MB
    PROGRESS_UPDATE_INTERVAL = 10    # 每10个块更新一次进度
    INITIAL_WAIT_SECONDS = 120       # 升级后等待设备重启的初始时间
    MAX_POLL_ATTEMPTS = 120          # 最大轮询次数(10分钟,每次5秒)
    POLLING_INTERVAL_SECONDS = 5     # 轮询间隔(秒)
    PROGRESS_LOG_INTERVAL = 12       # 每12次轮询(即1分钟)打印一次日志
    
    # 状态符号
    SUCCESS_SYMBOL = "[OK]"
    FAILURE_SYMBOL = "[FAIL]"
    DEBUG_SYMBOL = "[DEBUG]"
    
    def __init__(self, server_addr: str, verbose: bool = False):
        self.server_addr = server_addr
        self.verbose = verbose
        self.channel = None
        self.firmware_stub = None
        self.config_stub = None
        self.start_time = None
        self.connected = False
        self.logger = logging.getLogger(self.__class__.__name__)
        if self.verbose:
            self.logger.setLevel(logging.DEBUG)
    
    def connect(self) -> bool:
        """连接服务器并初始化两个存根"""
        self.logger.info(f"连接到 {self.server_addr}")
        
        # 重试连接
        max_retries = 3
        for retry in range(max_retries):
            try:
                self.channel = grpc.insecure_channel(
                    self.server_addr,
                    options=[
                        ('grpc.max_send_message_length', 100 * 1024 * 1024),
                        ('grpc.max_receive_message_length', 100 * 1024 * 1024),
                        ('grpc.keepalive_time_ms', 10000),
                        ('grpc.keepalive_timeout_ms', 5000),
                        ('grpc.keepalive_permit_without_calls', 1),
                        ('grpc.http2.max_pings_without_data', 0),
                    ]
                )
                
                # 初始化两个存根
                self.firmware_stub = pb2_grpc.TestInterfaceFirmwareUpdateServiceStub(self.channel)
                self.config_stub = pb2_grpc.TestInterfaceConfigurationServiceStub(self.channel)
                
                # 测试通道
                grpc.channel_ready_future(self.channel).result(timeout=10)
                self.logger.info(f"Firmware client connected to: {self.server_addr}")
                self.connected = True
                
                # 测试固件服务连接
                if self.test_connection():
                    return True
                else:
                    self.logger.warning(f"连接测试失败,重试 {retry+1}/{max_retries}")
                    
            except Exception as e:
                self.logger.warning(f"连接尝试 {retry+1}/{max_retries} 失败: {e}")
                time.sleep(2)
        
        self.logger.error(f"连接失败,尝试 {max_retries} 次后仍无法连接")
        return False
    
    def test_connection(self) -> bool:
        """测试固件服务连接"""
        print(f"{self.DEBUG_SYMBOL} 测试固件连接...")
        try:
            # 创建测试请求生成器
            def generate_test_request():
                request = pb2.FirmwareUpdateRequest()
                request.dut_position = 0
                yield request
            
            # 调用流式RPC测试
            response = self.firmware_stub.FirmwareUpdate(
                generate_test_request(),
                timeout=10
            )
            
            print(f"{self.SUCCESS_SYMBOL} Firmware connection test succeeded")
            return True
            
        except grpc.RpcError as e:
            print(f"{self.FAILURE_SYMBOL} Firmware connection test failed: {e.details()}")
            return False
        except Exception as e:
            print(f"{self.FAILURE_SYMBOL} Firmware connection test error: {e}")
            return False
    
    def create_session(self) -> type_pb2.SessionId:
        """创建会话"""
        session = type_pb2.SessionId()
        session.device_id = "python_complete_client"
        
        now = Timestamp()
        now.GetCurrentTime()
        session.start_time.CopyFrom(now)
        
        return session
    
    def calculate_md5_hash(self, filepath: str) -> str:
        """计算文件的MD5哈希值"""
        md5 = hashlib.md5()
        with open(filepath, 'rb') as f:
            while chunk := f.read(8192):
                md5.update(chunk)
        return md5.hexdigest().lower()
    
    def upload_with_manual_info(self, filepath: str, product_num: str, 
                               rstate: str, dut_position: int = 0) -> bool:
        """上传固件"""
        print(f"Testing upload: {filepath}")
        
        # 检查文件
        if not os.path.exists(filepath):
            print(f"{self.FAILURE_SYMBOL} Cannot open file: {filepath}")
            return False
        
        filesize = os.path.getsize(filepath)
        print(f"File size: {filesize} bytes")
        
        try:
            filename = os.path.basename(filepath)
            filehash = self.calculate_md5_hash(filepath)
            
            shared_session = self.create_session()
            
            # 创建请求生成器
            def generate_upload_requests():
                # 1. 发送软件项信息
                request1 = pb2.FirmwareUpdateRequest()
                request1.session.CopyFrom(shared_session)
                request1.dut_position = dut_position
                
                # 创建SoftwareItem
                item = pb2.SoftwareItem()
                item.description = f"Firmware update"
                item.product_number = product_num
                item.rstate = rstate
                item.sw_type = common_enums_pb2.SOFTWARE_TYPE_INITIAL_FLASH_IMAGE
                item.filename = filename
                item.hash = filehash
                item.total_size = filesize
                
                request1.item.CopyFrom(item)
                yield request1
                
                print(f"{self.SUCCESS_SYMBOL} Software item sent")
                
                # 2. 发送文件数据
                chunk_count = 0
                total_sent = 0
                
                with open(filepath, 'rb') as f:
                    while True:
                        chunk = f.read(self.FIRMWARE_CHUNK_SIZE)
                        if not chunk:
                            break
                        
                        bytes_read = len(chunk)
                        request = pb2.FirmwareUpdateRequest()
                        request.session.CopyFrom(shared_session)
                        request.dut_position = dut_position
                        
                        # 创建SoftwareItemContent
                        content = pb2.SoftwareItemContent()
                        content.swType = item.sw_type 
                        content.data = chunk
                        
                        request.content.CopyFrom(content)
                        
                        total_sent += bytes_read
                        chunk_count += 1
                        
                        # 显示进度
                        if (chunk_count % self.PROGRESS_UPDATE_INTERVAL == 0 or 
                            total_sent == filesize):
                            if filesize > 0:
                                progress = (total_sent * 100) // filesize
                            else:
                                progress = 0
                            print(f"Progress: {progress}% ({total_sent}/{filesize} bytes)")
                        
                        yield request
                
                print(f"{self.SUCCESS_SYMBOL} All data sent ({chunk_count} chunks)")
            
            # 调用流式RPC,设置较长超时时间
            response = self.firmware_stub.FirmwareUpdate(
                generate_upload_requests(),
                timeout=600  # 10分钟超时
            )
            
            # 检查响应
            if hasattr(response, 'code'):
                print(f"Server response: {response.message} (code: {response.code})")
                return response.code == 0
            else:
                print(f"{self.DEBUG_SYMBOL} Unexpected server response: {response}")
                return True  # 即使响应格式不同,也认为成功
                
        except grpc.RpcError as e:
            self.logger.error(f"Upload failed: {e.details()}")
            if self.verbose:
                self.logger.error(f"Error code: {e.code()}")
            return False
        except Exception as e:
            self.logger.exception(f"Upload error: {e}")
            return False
    
    def parse_sw_info_response(self, response) -> List[Dict[str, Any]]:
        """
        解析 GetSWInfos 响应。
        此函数根据 'test_interface.proto' 中定义的 SWInfos 结构进行解析。
        """
        sw_infos = []
        try:
            # 响应包含一个名为 sw_info 的重复字段
            for sw_info_item in response.sw_info:
                info = {
                    'name': sw_info_item.sw_name,
                    'product_number': sw_info_item.product_number,
                    'rstate': sw_info_item.product_rstate,
                    'comment': sw_info_item.comment
                }
                sw_infos.append(info)
        except AttributeError as e:
            self.logger.error(f"解析软件信息响应时发生属性错误: {e}")
            self.logger.debug(f"收到的响应类型: {type(response)}")
            self.logger.debug(f"收到的响应内容: {response}")
            if self.verbose:
                traceback.print_exc()
        except Exception as e:
            self.logger.error(f"解析软件信息时发生未知错误: {e}")
            if self.verbose:
                traceback.print_exc()
        
        return sw_infos
    
    def wait_for_device_and_get_sw_info(self, max_retries: Optional[int] = None) -> bool:
        """等待设备重启并获取软件信息 - 与C++版本一致"""
        if max_retries is None:
            max_retries = self.MAX_POLL_ATTEMPTS
        
        print(f"Waiting for device to start upgrade ({self.INITIAL_WAIT_SECONDS} seconds)...")
        
        # 等待设备开始升级(2分钟)
        for i in range(self.INITIAL_WAIT_SECONDS):
            time.sleep(1)
            if (i + 1) % 30 == 0:
                print(f"Still waiting for upgrade to start... ({i+1}/{self.INITIAL_WAIT_SECONDS} seconds)")
        
        print(f"Starting to poll GetSWInfos every {self.POLLING_INTERVAL_SECONDS} seconds (max {max_retries} attempts)...")
        
        # 轮询设备状态
        for attempt in range(1, max_retries + 1):
            response = None
            try:
                # --- 核心修复:在每次循环中创建全新的临时连接 ---
                # 这可以避免因服务器重启导致的旧连接失效问题
                with grpc.insecure_channel(self.server_addr) as channel:
                    # 等待通道就绪,超时时间要短于轮询间隔
                    grpc.channel_ready_future(channel).result(timeout=self.POLLING_INTERVAL_SECONDS - 1)
                    
                    temp_stub = pb2_grpc.TestInterfaceConfigurationServiceStub(channel)
                    request = empty_pb2.Empty()
                    # 使用一个较短的RPC超时
                    response = temp_stub.GetSWInfos(request, timeout=2)

            except grpc.RpcError as e:
                # 在轮询期间,UNAVAILABLE 和 DEADLINE_EXCEEDED 是正常现象
                if e.code() not in (grpc.StatusCode.UNAVAILABLE, grpc.StatusCode.DEADLINE_EXCEEDED):
                    if self.verbose:
                        print(f"{self.DEBUG_SYMBOL} 轮询时发生非预期的RPC错误 (attempt {attempt}): {e.details()}")
            except KeyboardInterrupt:
                print("\n轮询被用户中断")
                return False
            except Exception as e:
                # 捕获其他异常,例如 channel_ready_future 超时
                if self.verbose:
                    print(f"{self.DEBUG_SYMBOL} 轮询异常 (attempt {attempt}): {e}")

            # 检查是否成功获取响应
            if response is not None:
                print(f"{self.SUCCESS_SYMBOL} Device is back online! New firmware is running.")
                
                # 解析软件信息
                sw_infos = self.parse_sw_info_response(response)
                
                if sw_infos:
                    print("Current software information:")
                    for info in sw_infos:
                        info_str = f"  - Name: {info.get('name', 'Unknown')}"
                        if 'product_number' in info:
                            info_str += f"\n    Product Number: {info['product_number']}"
                        if 'rstate' in info:
                            info_str += f"\n    R-State: {info['rstate']}"
                        if 'comment' in info and info['comment']:
                            info_str += f"\n    Comment: {info['comment']}"
                        print(info_str)
                else:
                    print(f"{self.DEBUG_SYMBOL} 解析到响应但无软件信息,原始响应: {response}")
                
                return True
            
            # 设备尚未准备好,继续轮询
            if attempt % self.PROGRESS_LOG_INTERVAL == 0:
                print(f"Still waiting for device response... (attempt {attempt}/{max_retries})")
            
            time.sleep(self.POLLING_INTERVAL_SECONDS)

        print(f"{self.FAILURE_SYMBOL} Device did not respond after {max_retries * self.POLLING_INTERVAL_SECONDS} seconds of polling")
        return False
    
    def firmware_upgrade(self, filepath: str, product_num: str, 
                        rstate: str, dut_position: int = 0, 
                        max_poll_attempts: int = None) -> bool:
        """完整固件升级流程"""
        self.start_time = time.time()
        self.logger.info("Starting complete firmware upgrade process...")
        
        # 使用指定的轮询次数
        if max_poll_attempts is not None:
            self.MAX_POLL_ATTEMPTS = max_poll_attempts
        
        # 1. 执行固件上传
        self.logger.info("\n" + "="*60)
        self.logger.info("Step 1: Uploading firmware...")
        self.logger.info("="*60)
        
        if not self.upload_with_manual_info(filepath, product_num, rstate, dut_position):
            self.logger.error("Firmware upload failed")
            return False
        
        upload_time = time.time() - self.start_time
        self.logger.info(f"Firmware upload completed in {upload_time:.2f} seconds")
        
        # 2. 等待设备重启并重新上线
        self.logger.info("\n" + "="*60)
        self.logger.info("Step 2: Waiting for device to restart and come back online...")
        self.logger.info("="*60)
        
        if not self.wait_for_device_and_get_sw_info(max_poll_attempts):
            self.logger.error("Device did not come back online or failed to get software info")
            return False
        
        total_time = time.time() - self.start_time
        self.logger.info("\n" + "="*60)
        self.logger.info("Firmware upgrade completed successfully!")
        self.logger.info(f"Total time: {total_time:.2f} seconds")
        self.logger.info("="*60)
        
        return True
    
    def verify_file(self, filepath: str) -> bool:
        """验证文件并计算哈希值"""
        if not os.path.exists(filepath):
            print(f"{self.FAILURE_SYMBOL} 文件不存在: {filepath}")
            return False
        
        filename = os.path.basename(filepath)
        filesize = os.path.getsize(filepath)
        
        # 计算所有哈希值
        md5_hash = self.calculate_md5_hash(filepath)
        
        print(f"\n{'='*60}")
        print("文件验证报告")
        print('='*60)
        print(f"文件名: {filename}")
        print(f"文件大小: {filesize:,} 字节 ({filesize/1024/1024:.2f} MB)")
        print(f"修改时间: {datetime.fromtimestamp(os.path.getmtime(filepath))}")
        print('='*60)
        print("哈希值:")
        print(f"  MD5 (小写):    {md5_hash}")
        print('='*60)
        
        print(f"\n{self.DEBUG_SYMBOL} 此服务器期望使用: MD5 (小写)")
        
        return True
    
    def close(self):
        """关闭连接"""
        if self.channel:
            self.channel.close()
            if self.start_time:
                elapsed_time = time.time() - self.start_time
                self.logger.info(f"连接已关闭,总运行时间: {elapsed_time:.2f} 秒")
            else:
                self.logger.info("连接已关闭")
            self.connected = False
client_main.py
#!/usr/bin/env python3
"""
Firmware Update Client - 命令行入口
"""
import argparse
import os
import sys
import traceback

from firmware_client import CompleteFirmwareClient


def main():
    parser = argparse.ArgumentParser(
        description='完整固件升级客户端 - 修复版',
        formatter_class=argparse.RawDescriptionHelpFormatter,
        epilog='''
示例:
  %(prog)s test localhost:50050
  %(prog)s upload localhost:50050 firmware.bin "CXP9024418/1" R1A
  %(prog)s upgrade localhost:50050 firmware.bin "CXP9024418/1" R1A
  %(prog)s verify firmware.bin

注意: 完整升级流程需要较长时间,建议增加轮询次数
        '''
    )
    
    subparsers = parser.add_subparsers(dest='command', help='命令', required=True)
    
    # test 命令
    test_parser = subparsers.add_parser('test', help='测试服务器连接')
    test_parser.add_argument('server', help='服务器地址 (主机:端口)')
    test_parser.add_argument('--verbose', '-v', action='store_true', help='详细输出')
    
    # upload 命令
    upload_parser = subparsers.add_parser('upload', help='上传固件(不等待重启)')
    upload_parser.add_argument('server', help='服务器地址')
    upload_parser.add_argument('file', help='固件文件路径')
    upload_parser.add_argument('product', help='产品编号')
    upload_parser.add_argument('rstate', help='状态代码')
    upload_parser.add_argument('--position', '-p', type=int, default=0, help='DUT位置')
    upload_parser.add_argument('--verbose', '-v', action='store_true', help='详细输出')
    
    # upgrade 命令
    upgrade_parser = subparsers.add_parser('upgrade', help='完整固件升级(上传+等待重启)')
    upgrade_parser.add_argument('server', help='服务器地址')
    upgrade_parser.add_argument('file', help='固件文件路径')
    upgrade_parser.add_argument('product', help='产品编号')
    upgrade_parser.add_argument('rstate', help='状态代码')
    upgrade_parser.add_argument('--position', '-p', type=int, default=0, help='DUT位置')
    upgrade_parser.add_argument('--verbose', '-v', action='store_true', help='详细输出')
    upgrade_parser.add_argument('--max-poll', type=int, default=120, help='最大轮询次数(默认120,即10分钟)')
    
    # verify 命令
    verify_parser = subparsers.add_parser('verify', help='验证文件哈希值')
    verify_parser.add_argument('file', help='文件路径')
    verify_parser.add_argument('--verbose', '-v', action='store_true', help='详细输出')
    
    args = parser.parse_args()
    
    client = None
    
    try:
        if args.command == 'test':
            client = CompleteFirmwareClient(args.server, args.verbose)
            if client.connect():
                return 0
            return 1
        
        elif args.command == 'upload':
            if not os.path.exists(args.file):
                print(f"{CompleteFirmwareClient.FAILURE_SYMBOL} Error: File not found: {args.file}")
                return 1
            
            client = CompleteFirmwareClient(args.server, args.verbose)
            if not client.connect():
                return 1
            
            success = client.upload_with_manual_info(
                args.file, args.product, args.rstate, args.position
            )
            return 0 if success else 1
        
        elif args.command == 'upgrade':
            if not os.path.exists(args.file):
                print(f"{CompleteFirmwareClient.FAILURE_SYMBOL} Error: File not found: {args.file}")
                return 1
            
            print(f"\n准备执行完整固件升级:")
            print(f"  文件: {os.path.basename(args.file)}")
            print(f"  大小: {os.path.getsize(args.file):,} 字节")
            print(f"  产品: {args.product}")
            print(f"  状态: {args.rstate}")
            print(f"  位置: {args.position}")
            print(f"  最大轮询: {args.max_poll}次 (约{args.max_poll * 5 / 60:.1f}分钟)")
            
            confirm = input("\n确认执行完整升级流程? (y/N): ").strip().lower()
            if confirm not in ['y', 'yes', '是']:
                print(f"{CompleteFirmwareClient.DEBUG_SYMBOL} 操作已取消")
                return 0
            
            client = CompleteFirmwareClient(args.server, args.verbose)
            if not client.connect():
                return 1
            
            success = client.firmware_upgrade(
                args.file, args.product, args.rstate, args.position, args.max_poll
            )
            return 0 if success else 1
        
        elif args.command == 'verify':
            if not os.path.exists(args.file):
                print(f"{CompleteFirmwareClient.FAILURE_SYMBOL} Error: File not found: {args.file}")
                return 1
            
            client = CompleteFirmwareClient("localhost:50050", args.verbose)
            success = client.verify_file(args.file)
            return 0 if success else 1
    
    except KeyboardInterrupt:
        print(f"\n{CompleteFirmwareClient.DEBUG_SYMBOL} 操作被用户中断")
        return 130
    except Exception as e:
        print(f"\n{CompleteFirmwareClient.FAILURE_SYMBOL} 发生未预期错误: {e}")
        traceback.print_exc()
        return 1
    finally:
        if client:
            client.close()


if __name__ == '__main__':
    sys.exit(main())
posted @ 2026-02-27 15:16  mo686  阅读(51)  评论(0)    收藏  举报