Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 12 additions & 6 deletions pslab/pico/transport.py
Original file line number Diff line number Diff line change
Expand Up @@ -156,8 +156,12 @@ def connect(self) -> None:
if self._socket is not None:
return
sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
sock.settimeout(self.timeout)
sock.connect((self.host, self.port))
try:
sock.settimeout(self.timeout)
sock.connect((self.host, self.port))
except BaseException:
sock.close()
raise
self._socket = sock

def close(self) -> None:
Expand Down Expand Up @@ -233,7 +237,8 @@ def query(self, command: str) -> str:

self._write_line(command)
line = self.transport.readline()
if not line:
# pyserial returns a partial line without the terminator on timeout.
if not line.endswith(b"\n"):
raise ScpiTimeoutError(f"Timed out waiting for response to {command!r}.")
return line.decode("ascii", errors="replace").strip()

Expand All @@ -248,9 +253,10 @@ def read_block(self) -> bytes:

marker = self._read_exact(1)
if marker != b"#":
rest = self.transport.readline()
message = (marker + rest).decode("ascii", errors="replace").strip()
raise ScpiError(message)
line = marker + self.transport.readline()
if not line.endswith(b"\n"):
raise ScpiTimeoutError("Timed out waiting for SCPI error response.")
raise ScpiError(line.decode("ascii", errors="replace").strip())

digit_count_text = self._read_exact(1)
try:
Expand Down
38 changes: 38 additions & 0 deletions tests/test_pico_scpi_transport.py
Original file line number Diff line number Diff line change
Expand Up @@ -234,3 +234,41 @@ def test_wifi_readline_wraps_socket_timeout():

with pytest.raises(ScpiTimeoutError):
transport.readline()


class TruncatingTransport(FakeTransport):
"""Return a response without its newline, as pyserial does on timeout."""

def write(self, data):
self.output.extend(b"12")
return len(data)


def test_query_rejects_response_without_newline():
client = ScpiClient(TruncatingTransport())

with pytest.raises(ScpiTimeoutError):
client.error_count()


def test_query_block_rejects_error_line_without_newline():
client = ScpiClient(TruncatingTransport())

with pytest.raises(ScpiTimeoutError):
client.query_block("LA:READ?")


def test_wifi_connect_closes_socket_when_connect_fails(monkeypatch):
class RefusingSocket(FakeSocket):
def connect(self, address):
raise ConnectionRefusedError

sock = RefusingSocket([])
monkeypatch.setattr(socket, "socket", lambda *args: sock)
transport = PicoWifiTransport("example.test")

with pytest.raises(ConnectionRefusedError):
transport.connect()

assert sock.closed
assert not transport.is_open