Verified Commit 6b37d824 authored by Jakob Moser's avatar Jakob Moser
Browse files

Open new socket for every request

parent 79b4e185
Loading
Loading
Loading
Loading
Loading
+13 −16
Original line number Diff line number Diff line
import json
import socket
from collections.abc import Callable
from contextlib import AbstractContextManager
from dataclasses import dataclass, field
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Self
from socket import AF_UNIX, SOCK_STREAM
from socket import socket as Socket

from poolpay.wire.Packet import Packet


@dataclass
class Client(AbstractContextManager):
class Client:
    socket_path: Path
    _socket: socket.socket = field(init=False)

    def __enter__(self) -> Self:
        self._socket = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
        self._socket.connect(str(self.socket_path.resolve()))
        return self

    def __exit__(self, exc_type: Any, exc_value: Any, traceback: Any) -> None:
        self._socket.close()

    def send(self, packet: Packet) -> Packet:
        """
        Send the given packet to the server and returns the server's response.
        """
        self._socket.sendall((json.dumps(packet) + "\n").encode("utf-8"))
        socket = Socket(AF_UNIX, SOCK_STREAM)

        try:
            socket.connect(str(self.socket_path.resolve()))
            socket.sendall((json.dumps(packet) + "\n").encode("utf-8"))

            # TODO What if the response is longer than 1024 bytes?
        response = self._socket.recv(1024)
            response = socket.recv(1024)
            return json.loads(response.decode("utf-8"))
        finally:
            socket.close()

    def on_receive_push(self, handle_packet: Callable[[Packet], None]) -> None:
        """
+1 −1
Original line number Diff line number Diff line
@@ -66,7 +66,7 @@ def ingest(

    messages = tuple(json.loads(m) for m in messages_string.splitlines())

    with Client(socket_path) as client:
    client = Client(socket_path)
    for m in messages:
        client.send(m)

+3 −3
Original line number Diff line number Diff line
from collections.abc import Generator
from pathlib import Path
from threading import Thread

@@ -13,7 +14,7 @@ def socket_path() -> Path:


@pytest.fixture
def server(socket_path: Path) -> Server:
def server(socket_path: Path) -> Generator[Server]:
    server = Server(socket_path)
    try:
        Thread(target=server.start).start()
@@ -27,8 +28,7 @@ def test_can_send_multiple_messages(socket_path: Path, server: Server) -> None:
    Test that a client can send multiple messages in one connection
    (i.e., the server does not close the socket in between).
    """

    with Client(socket_path) as client:
    client = Client(socket_path)
    response_0 = client.send({"packet": 0})
    assert response_0 == {"status": "received"}