summaryrefslogtreecommitdiff
path: root/tools/test_srcipts/test_cdc_speed.py
blob: 9a54bfdcdcf2312fc7c4449316e7ff7102becf02 (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
#
# Copyright (c) 2022, sakumisu
#
# SPDX-License-Identifier: Apache-2.0
#
import argparse
import time

import serial


DEFAULT_PORT = "COM17"
DEFAULT_BAUDRATE = 2000000
DEFAULT_LENGTH = 409600
MAX_CHUNK_LENGTH = 4 * 1024 * 1024
DEFAULT_REPORT_INTERVAL = 1.0
TEST_PATTERN = bytes([0xAA]) * MAX_CHUNK_LENGTH


def parse_size(value: str) -> int:
    text = value.strip().upper()
    units = {
        "B": 1,
        "K": 1024,
        "KB": 1024,
        "M": 1024 * 1024,
        "MB": 1024 * 1024,
        "G": 1024 * 1024 * 1024,
        "GB": 1024 * 1024 * 1024,
    }

    for unit, factor in sorted(units.items(), key=lambda item: len(item[0]), reverse=True):
        if text.endswith(unit):
            number = text[:-len(unit)].strip()
            return int(float(number) * factor)

    return int(float(text))


def format_speed(byte_count: int, duration: float) -> str:
    if duration <= 0:
        return "0.00 MB/s (0.00 Mbps)"

    mb_per_sec = byte_count / duration / 1024 / 1024
    mbps = byte_count * 8 / duration / 1000 / 1000
    return f"{mb_per_sec:.2f} MB/s ({mbps:.2f} Mbps)"


def print_progress(prefix: str, done_bytes: int, bytes_in_window: int, start_time: float, window_start: float) -> None:
    now = time.time()
    avg_speed = format_speed(done_bytes, now - start_time)
    instant_speed = format_speed(bytes_in_window, now - window_start)
    print(f"{prefix}: total {done_bytes} bytes, instant {instant_speed}, average {avg_speed}")


def run_tx(ser: serial.Serial, packet_size: int, report_interval: float) -> None:
    sent_bytes = 0
    start_time = time.time()
    window_start = start_time
    window_bytes = 0

    if packet_size <= len(TEST_PATTERN):
        payload = TEST_PATTERN[:packet_size]
    else:
        payload = bytes([0xAA]) * packet_size

    try:
        while True:
            sent = ser.write(payload)
            sent_bytes += sent
            window_bytes += sent

            now = time.time()
            if now - window_start >= report_interval:
                print_progress("TX", sent_bytes, window_bytes, start_time, window_start)
                window_start = now
                window_bytes = 0
    except KeyboardInterrupt:
        pass

    if window_bytes > 0:
        print_progress("TX", sent_bytes, window_bytes, start_time, window_start)

    total_time = time.time() - start_time
    print(f"TX stopped: {format_speed(sent_bytes, total_time)}")


def run_rx(ser: serial.Serial, packet_size: int, report_interval: float) -> None:
    recv_bytes = 0
    start_time = time.time()
    window_start = start_time
    window_bytes = 0

    try:
        while True:
            data = ser.read(packet_size)
            if not data:
                continue

            received = len(data)
            recv_bytes += received
            window_bytes += received

            now = time.time()
            if now - window_start >= report_interval:
                print_progress("RX", recv_bytes, window_bytes, start_time, window_start)
                window_start = now
                window_bytes = 0
    except KeyboardInterrupt:
        pass

    if window_bytes > 0:
        print_progress("RX", recv_bytes, window_bytes, start_time, window_start)

    total_time = time.time() - start_time
    print(f"RX stopped: {format_speed(recv_bytes, total_time)}")


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description="CDC speed test tool")
    parser.add_argument("--port", default=DEFAULT_PORT, help=f"serial port, default: {DEFAULT_PORT}")
    parser.add_argument("--mode", choices=["tx", "rx"], required=True, help="tx: send data, rx: receive data")
    parser.add_argument("--chunk-length", type=parse_size, default=DEFAULT_LENGTH,
                        help=f"packet size per send/read, max 4M, default: {DEFAULT_LENGTH}, supports suffixes like 4K, 400K, 1M")
    parser.add_argument("--report-interval", type=float, default=DEFAULT_REPORT_INTERVAL,
                        help=f"speed report interval in seconds, default: {DEFAULT_REPORT_INTERVAL}")
    return parser.parse_args()


def main() -> None:
    args = parse_args()

    print("Please remove usb log and change read & write buffer size then test speed")
    if args.chunk_length > MAX_CHUNK_LENGTH:
        raise ValueError("chunk length must be <= 4M")

    with serial.Serial(args.port, DEFAULT_BAUDRATE, timeout=None) as ser:
        if args.mode == "tx":
            print(f"Start TX test on {args.port}, packet_size={args.chunk_length} bytes")
            print("Press Ctrl+C to stop")
            ser.setDTR(0)
            run_tx(ser, args.chunk_length, args.report_interval)
        else:
            print(f"Start RX test on {args.port}, packet_size={args.chunk_length} bytes")
            print("Press Ctrl+C to stop")
            ser.setDTR(1)
            run_rx(ser, args.chunk_length, args.report_interval)


if __name__ == "__main__":
    main()