init pro_sdk
This commit is contained in:
48
tests/test_blowfish.py
Normal file
48
tests/test_blowfish.py
Normal file
@@ -0,0 +1,48 @@
|
||||
"""Blowfish 加解密测试 (用已知明密文对)."""
|
||||
import os
|
||||
import pytest
|
||||
from go1_pro_sdk import Blowfish, verify_state
|
||||
from go1_pro_sdk.codec.blowfish import KNOWN_PAIRS
|
||||
|
||||
|
||||
STATE_FILE = os.path.join(
|
||||
os.path.dirname(__file__), '..', 'go1_pro_sdk', '_data', 'blowfish_state.bin'
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def bf():
|
||||
return Blowfish.from_state_file(STATE_FILE)
|
||||
|
||||
|
||||
def test_state_loads(bf):
|
||||
assert len(bf.P) == 18
|
||||
assert len(bf.S) == 4
|
||||
assert all(len(s) == 256 for s in bf.S)
|
||||
|
||||
|
||||
def test_state_verify(bf):
|
||||
"""3/3 已知明密文对应该全部匹配."""
|
||||
assert verify_state(bf) is True
|
||||
|
||||
|
||||
def test_encrypt_known_pairs(bf):
|
||||
for pt, ct in KNOWN_PAIRS:
|
||||
assert bf.encrypt_block(pt) == ct
|
||||
|
||||
|
||||
def test_decrypt_known_pairs(bf):
|
||||
for pt, ct in KNOWN_PAIRS:
|
||||
assert bf.decrypt_block(ct) == pt
|
||||
|
||||
|
||||
def test_encrypt_decrypt_roundtrip(bf):
|
||||
data = b'\x12\x34\x56\x78' * 100 # 400B
|
||||
encrypted = bf.encrypt_ecb(data)
|
||||
decrypted = bf.decrypt_ecb(encrypted)
|
||||
assert data == decrypted
|
||||
|
||||
|
||||
def test_encrypt_size_must_be_8x(bf):
|
||||
with pytest.raises(AssertionError):
|
||||
bf.encrypt_ecb(b'\x00' * 7)
|
||||
94
tests/test_lowcmd_builder.py
Normal file
94
tests/test_lowcmd_builder.py
Normal file
@@ -0,0 +1,94 @@
|
||||
"""LowCmd 序列化测试 (PRO 格式)."""
|
||||
import os
|
||||
import pytest
|
||||
from go1_pro_sdk import (
|
||||
Blowfish, LowCmd, MotorCmd, MotorMode,
|
||||
build_low_cmd_plain, build_low_cmd_encrypted,
|
||||
)
|
||||
|
||||
|
||||
STATE_FILE = os.path.join(
|
||||
os.path.dirname(__file__), '..', 'go1_pro_sdk', '_data', 'blowfish_state.bin'
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def bf():
|
||||
return Blowfish.from_state_file(STATE_FILE)
|
||||
|
||||
|
||||
def test_plain_length_616():
|
||||
cmd = LowCmd()
|
||||
plain = build_low_cmd_plain(cmd)
|
||||
assert len(plain) == 616
|
||||
|
||||
|
||||
def test_plain_head():
|
||||
cmd = LowCmd()
|
||||
plain = build_low_cmd_plain(cmd)
|
||||
assert plain[:2] == b'\xfe\xef'
|
||||
assert plain[2] == 0xff
|
||||
|
||||
|
||||
def test_plain_sn_version_all_zero():
|
||||
"""PRO 真实包 SN/version 全 0."""
|
||||
cmd = LowCmd()
|
||||
plain = build_low_cmd_plain(cmd)
|
||||
assert plain[4:12] == b'\x00' * 8 # SN
|
||||
assert plain[12:20] == b'\x00' * 8 # version
|
||||
|
||||
|
||||
def test_plain_bandwidth_be():
|
||||
"""bandWidth 是 BE 字节序 (3a c0, 不是 c0 3a)."""
|
||||
cmd = LowCmd()
|
||||
plain = build_low_cmd_plain(cmd)
|
||||
assert plain[20:22] == b'\x3a\xc0'
|
||||
|
||||
|
||||
def test_plain_padding_zeros():
|
||||
"""610..612 是固定 0000."""
|
||||
cmd = LowCmd()
|
||||
plain = build_low_cmd_plain(cmd)
|
||||
assert plain[610:612] == b'\x00\x00'
|
||||
|
||||
|
||||
def test_plain_crc_at_612():
|
||||
"""CRC 在 612..616, 不是 EDU 的 610..614."""
|
||||
cmd = LowCmd()
|
||||
plain = build_low_cmd_plain(cmd)
|
||||
from go1_pro_sdk.utils.common import gen_crc
|
||||
expected_crc = gen_crc(plain[:612])
|
||||
assert plain[612:616] == expected_crc
|
||||
|
||||
|
||||
def test_encrypted_damping_matches_real(bf):
|
||||
"""全 damping 加密后, 前 16B 应跟真实 Legged_sport 抓包一致."""
|
||||
cmd = LowCmd()
|
||||
enc = build_low_cmd_encrypted(cmd, bf)
|
||||
assert len(enc) == 616
|
||||
# 真实抓包前 16B (strace 抓的 Legged_sport sendto 内容)
|
||||
expected = bytes.fromhex('e79e1ccc4bae9cf812aa0354236b66e3')
|
||||
assert enc[:16] == expected
|
||||
|
||||
|
||||
def test_set_motor_by_name():
|
||||
cmd = LowCmd()
|
||||
cmd.set_motor('FR_1', MotorCmd(mode=MotorMode.Servo, q=1.2, Kp=5, Kd=1))
|
||||
assert cmd.motorCmd[1].mode == MotorMode.Servo
|
||||
assert cmd.motorCmd[1].q == 1.2
|
||||
assert cmd.motorCmd[1].Kp == 5
|
||||
|
||||
|
||||
def test_set_motor_by_index():
|
||||
cmd = LowCmd()
|
||||
cmd.set_motor(5, MotorCmd(q=2.0))
|
||||
assert cmd.motorCmd[5].q == 2.0
|
||||
|
||||
|
||||
def test_all_damping():
|
||||
cmd = LowCmd()
|
||||
cmd.motorCmd[0].mode = MotorMode.Servo
|
||||
cmd.motorCmd[0].q = 1.5
|
||||
cmd.all_damping()
|
||||
assert cmd.motorCmd[0].mode == MotorMode.Damping
|
||||
assert cmd.motorCmd[0].q == 0.0
|
||||
82
tests/test_safety.py
Normal file
82
tests/test_safety.py
Normal file
@@ -0,0 +1,82 @@
|
||||
"""Safety 保护层测试."""
|
||||
import pytest
|
||||
from go1_pro_sdk import (
|
||||
LowCmd, MotorCmd, MotorMode,
|
||||
apply_safety, position_limit, power_protect,
|
||||
PowerProtectViolation,
|
||||
TAU_MAX,
|
||||
)
|
||||
|
||||
|
||||
class _FakeMotor:
|
||||
tauEst = 0.0
|
||||
q = 0.0
|
||||
|
||||
|
||||
class _FakeState:
|
||||
def __init__(self):
|
||||
self.motorState = [_FakeMotor() for _ in range(20)]
|
||||
|
||||
|
||||
def test_position_limit_clamps_hip_over():
|
||||
cmd = LowCmd()
|
||||
cmd.motorCmd[0].q = 5.0 # FR_0 hip 超限
|
||||
position_limit(cmd)
|
||||
assert cmd.motorCmd[0].q == 0.78
|
||||
|
||||
|
||||
def test_position_limit_clamps_thigh_under():
|
||||
cmd = LowCmd()
|
||||
cmd.motorCmd[1].q = -3.0
|
||||
position_limit(cmd)
|
||||
assert cmd.motorCmd[1].q == -0.60
|
||||
|
||||
|
||||
def test_power_protect_clamps():
|
||||
cmd = LowCmd()
|
||||
for i in range(12):
|
||||
cmd.motorCmd[i].tau = 10.0
|
||||
state = _FakeState()
|
||||
n = power_protect(cmd, state, factor=1)
|
||||
# factor=1 → 限制到 TAU_MAX * 0.1
|
||||
assert cmd.motorCmd[0].tau == TAU_MAX['hip'] * 0.1
|
||||
assert cmd.motorCmd[2].tau == TAU_MAX['knee'] * 0.1
|
||||
|
||||
|
||||
def test_power_protect_critical_command_raises():
|
||||
cmd = LowCmd()
|
||||
cmd.motorCmd[2].tau = 200.0 # > 5 × 35.55
|
||||
with pytest.raises(PowerProtectViolation):
|
||||
power_protect(cmd, _FakeState(), factor=5)
|
||||
|
||||
|
||||
def test_power_protect_overload_actual_raises():
|
||||
"""实测 tauEst 已超 max, 也应 raise."""
|
||||
cmd = LowCmd()
|
||||
cmd.motorCmd[2].tau = 1.0
|
||||
state = _FakeState()
|
||||
state.motorState[2].tauEst = 40.0 # > 35.55
|
||||
with pytest.raises(PowerProtectViolation):
|
||||
power_protect(cmd, state, factor=1)
|
||||
|
||||
|
||||
def test_apply_safety_degraded_mode():
|
||||
"""raise_on_critical=False 应当降级到全 damping."""
|
||||
cmd = LowCmd()
|
||||
cmd.motorCmd[2].tau = 200.0
|
||||
cmd.motorCmd[2].Kp = 10
|
||||
apply_safety(cmd, _FakeState(), power_factor=5,
|
||||
position_limit_on=False, raise_on_critical=False)
|
||||
assert cmd.motorCmd[0].tau == 0
|
||||
assert cmd.motorCmd[0].Kp == 0
|
||||
assert cmd.motorCmd[0].mode == 0
|
||||
|
||||
|
||||
def test_apply_safety_normal_flow():
|
||||
"""正常命令应能通过. 返回 LowCmd."""
|
||||
cmd = LowCmd()
|
||||
cmd.motorCmd[0].q = 0.5
|
||||
cmd.motorCmd[0].tau = 1.0
|
||||
result = apply_safety(cmd, _FakeState(), power_factor=5)
|
||||
assert result is cmd
|
||||
assert cmd.motorCmd[0].q == 0.5 # 在限位内, 不变
|
||||
Reference in New Issue
Block a user