mirror of
https://github.com/odriverobotics/ODrive.git
synced 2026-09-28 16:13:37 +08:00
cancel receiver thread correctly on app shutdown
This commit is contained in:
+1
-1
@@ -251,7 +251,7 @@ def launch_dfu(args, app_shutdown_token):
|
||||
|
||||
# Scan for ODrives not in DFU mode and put them into DFU mode once they appear
|
||||
# We only scan on USB because DFU is only possible over USB
|
||||
odrive.discovery.find_all(args.path, serial_number, put_odrive_into_dfu_mode, find_odrive_cancellation_token)
|
||||
odrive.discovery.find_all(args.path, serial_number, put_odrive_into_dfu_mode, find_odrive_cancellation_token, app_shutdown_token)
|
||||
|
||||
# Poll libUSB until a device in DFU mode is found
|
||||
while not app_shutdown_token.is_set():
|
||||
|
||||
@@ -24,7 +24,8 @@ def noprint(text):
|
||||
|
||||
def find_all(path, serial_number,
|
||||
did_discover_object_callback,
|
||||
cancellation_token, printer=noprint):
|
||||
search_cancellation_token,
|
||||
channel_termination_token, printer=noprint):
|
||||
"""
|
||||
Starts scanning for ODrives that match the specified path spec and calls
|
||||
the callback for each ODrive that is found.
|
||||
@@ -75,21 +76,23 @@ def find_all(path, serial_number,
|
||||
the_rest = ':'.join(search_spec.split(':')[1:])
|
||||
if prefix in channel_types:
|
||||
threading.Thread(target=channel_types[prefix],
|
||||
args=(the_rest, serial_number, did_discover_channel, cancellation_token, printer)).start()
|
||||
args=(the_rest, serial_number, did_discover_channel, search_cancellation_token, channel_termination_token, printer)).start()
|
||||
else:
|
||||
raise Exception("Invalid path spec \"{}\"".format(search_spec))
|
||||
|
||||
|
||||
def find_any(path="usb", serial_number=None, cancellation_token=None, timeout=None, printer=noprint):
|
||||
def find_any(path="usb", serial_number=None,
|
||||
search_cancellation_token=None, channel_termination_token=None,
|
||||
timeout=None, printer=noprint):
|
||||
"""
|
||||
Blocks until the first matching ODrive is connected and then returns that device
|
||||
"""
|
||||
result = [ None ]
|
||||
done_signal = Event(cancellation_token)
|
||||
done_signal = Event(search_cancellation_token)
|
||||
def did_discover_object(obj):
|
||||
result[0] = obj
|
||||
done_signal.set()
|
||||
find_all(path, serial_number, did_discover_object, done_signal, printer)
|
||||
find_all(path, serial_number, did_discover_object, done_signal, channel_termination_token, printer)
|
||||
try:
|
||||
done_signal.wait(timeout=timeout)
|
||||
finally:
|
||||
|
||||
@@ -206,7 +206,7 @@ class Channel(PacketSink):
|
||||
_resend_timeout = 0.1 # [s]
|
||||
_send_attempts = 5
|
||||
|
||||
def __init__(self, name, input, output, printer):
|
||||
def __init__(self, name, input, output, cancellation_token, printer):
|
||||
"""
|
||||
Params:
|
||||
input: A PacketSource where this channel will source packets from on
|
||||
@@ -223,8 +223,8 @@ class Channel(PacketSink):
|
||||
self._expected_acks = {}
|
||||
self._responses = {}
|
||||
self._my_lock = threading.Lock()
|
||||
self._channel_broken = Event()
|
||||
self.start_receiver_thread(Event()) # TODO: use app_shutdown_token
|
||||
self._channel_broken = Event(cancellation_token)
|
||||
self.start_receiver_thread(Event(self._channel_broken)) # TODO: use app_shutdown_token
|
||||
|
||||
def start_receiver_thread(self, cancellation_token):
|
||||
"""
|
||||
@@ -246,7 +246,7 @@ class Channel(PacketSink):
|
||||
# Process response
|
||||
# This should not throw an exception, otherwise the channel breaks
|
||||
self.process_packet(response)
|
||||
print("receiver thread is exiting")
|
||||
#print("receiver thread is exiting")
|
||||
except Exception:
|
||||
self._printer("receiver thread is exiting: " + traceback.format_exc())
|
||||
finally:
|
||||
@@ -296,7 +296,7 @@ class Channel(PacketSink):
|
||||
self._my_lock.release()
|
||||
# Wait for ACK until the resend timeout is exceeded
|
||||
try:
|
||||
if wait_any(ack_event, self._channel_broken, timeout=self._resend_timeout) != 0:
|
||||
if wait_any(self._resend_timeout, ack_event, self._channel_broken) != 0:
|
||||
raise ChannelBrokenException()
|
||||
except odrive.utils.TimeoutException:
|
||||
attempt += 1
|
||||
|
||||
@@ -6,6 +6,7 @@ PacketSource/PacketSink interfaces for serial ports.
|
||||
import os
|
||||
import re
|
||||
import time
|
||||
import traceback
|
||||
import serial
|
||||
import serial.tools.list_ports
|
||||
import odrive.protocol
|
||||
@@ -53,10 +54,11 @@ def find_pyserial_ports():
|
||||
return [x.device for x in serial.tools.list_ports.comports()]
|
||||
|
||||
|
||||
def discover_channels(path, serial_number, callback, cancellation_token, printer):
|
||||
def discover_channels(path, serial_number, callback, cancellation_token, channel_termination_token, printer):
|
||||
"""
|
||||
Scans for serial ports that match the path spec.
|
||||
This function blocks until cancellation_token is set.
|
||||
Channels spawned by this function run until channel_termination_token is set.
|
||||
"""
|
||||
if path == None:
|
||||
# This regex should match all desired port names on macOS,
|
||||
@@ -86,7 +88,7 @@ def discover_channels(path, serial_number, callback, cancellation_token, printer
|
||||
output_stream = odrive.protocol.PacketToStreamConverter(serial_device)
|
||||
channel = odrive.protocol.Channel(
|
||||
"serial port {}@{}".format(port_name, ODRIVE_BAUDRATE),
|
||||
input_stream, output_stream, printer)
|
||||
input_stream, output_stream, channel_termination_token, printer)
|
||||
channel.serial_device = serial_device
|
||||
except serial.serialutil.SerialException:
|
||||
printer("Serial device init failed. Ignoring this port. More info: " + traceback.format_exc())
|
||||
|
||||
@@ -78,6 +78,7 @@ def launch_shell(args, logger, printer, app_shutdown_token):
|
||||
odrive.discovery.find_all(args.path, args.serial_number,
|
||||
lambda dev: did_discover_device(dev, logger, app_shutdown_token),
|
||||
app_shutdown_token,
|
||||
app_shutdown_token,
|
||||
printer=printer)
|
||||
|
||||
# Check if IPython is installed
|
||||
|
||||
@@ -139,10 +139,11 @@ class USBBulkTransport(odrive.protocol.PacketSource, odrive.protocol.PacketSink)
|
||||
return 64
|
||||
|
||||
|
||||
def discover_channels(path, serial_number, callback, cancellation_token, printer):
|
||||
def discover_channels(path, serial_number, callback, cancellation_token, channel_termination_token, printer):
|
||||
"""
|
||||
Scans for USB devices that match the path spec.
|
||||
This function blocks until cancellation_token is set.
|
||||
Channels spawned by this function run until channel_termination_token is set.
|
||||
"""
|
||||
if path == None or path == "":
|
||||
bus = None
|
||||
@@ -181,7 +182,7 @@ def discover_channels(path, serial_number, callback, cancellation_token, printer
|
||||
bulk_device.init()
|
||||
channel = odrive.protocol.Channel(
|
||||
"USB device bus {} device {}".format(usb_device.bus, usb_device.address),
|
||||
bulk_device, bulk_device, printer)
|
||||
bulk_device, bulk_device, channel_termination_token, printer)
|
||||
channel.usb_device = usb_device # for debugging only
|
||||
except usb.core.USBError as ex:
|
||||
if ex.errno == 13:
|
||||
|
||||
@@ -144,7 +144,7 @@ class Event():
|
||||
self._subscribers = []
|
||||
self._mutex = threading.Lock()
|
||||
if not trigger is None:
|
||||
trigger.subscribe(self.set())
|
||||
trigger.subscribe(lambda: self.set())
|
||||
|
||||
def is_set(self):
|
||||
return self._evt.is_set()
|
||||
@@ -170,6 +170,8 @@ class Event():
|
||||
handler is invoked immediately.
|
||||
Returns a function that can be invoked to unsubscribe.
|
||||
"""
|
||||
if handler is None:
|
||||
raise TypeError
|
||||
self._mutex.acquire()
|
||||
try:
|
||||
self._subscribers.append(handler)
|
||||
@@ -200,11 +202,12 @@ class Event():
|
||||
self.set()
|
||||
threading.Thread(target=delayed_trigger, daemon=True).start()
|
||||
|
||||
def wait_any(*events, timeout=None):
|
||||
def wait_any(timeout=None, *events):
|
||||
"""
|
||||
Blocks until any of the specified events are triggered.
|
||||
Returns the index of the event that was triggerd or raises
|
||||
a TimeoutException
|
||||
Param timeout: A timeout in seconds
|
||||
"""
|
||||
or_event = threading.Event()
|
||||
subscriptions = []
|
||||
|
||||
+12
-6
@@ -3,13 +3,22 @@
|
||||
ODrive command line utility
|
||||
"""
|
||||
|
||||
from __future__ import print_function
|
||||
import sys
|
||||
import argparse
|
||||
import odrive.discovery
|
||||
from odrive.utils import Logger, Event
|
||||
|
||||
# Flush stdout by default
|
||||
import functools
|
||||
print = functools.partial(print, flush=True)
|
||||
# Source:
|
||||
# https://stackoverflow.com/questions/230751/how-to-flush-output-of-python-print
|
||||
old_print = print
|
||||
def print(*args, **kwargs):
|
||||
kwargs.pop('flush', False)
|
||||
old_print(*args, **kwargs)
|
||||
file = kwargs.get('file', sys.stdout)
|
||||
# Why might file=None? IDK, but it works for print(i, file=None)
|
||||
file.flush() if file is not None else sys.stdout.flush()
|
||||
|
||||
|
||||
## Parse arguments ##
|
||||
@@ -69,10 +78,6 @@ if args.command is None:
|
||||
args.command = 'shell'
|
||||
args.no_ipython = False
|
||||
|
||||
# We are interactively printing status messages, so flush by default
|
||||
import functools
|
||||
print = functools.partial(print, flush=True)
|
||||
|
||||
# TODO: deprecate printer - use logger instead
|
||||
if (args.verbose):
|
||||
printer = print
|
||||
@@ -103,6 +108,7 @@ try:
|
||||
odrive.shell.launch_shell(args, logger, printer, app_shutdown_token)
|
||||
|
||||
elif args.command == 'dfu':
|
||||
print_version()
|
||||
import odrive.dfu
|
||||
odrive.dfu.launch_dfu(args, app_shutdown_token)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user