diff --git a/python/helpers/pydev/pydev_console/io.py b/python/helpers/pydev/pydev_console/io.py new file mode 100644 index 000000000000..482a3832d3fc --- /dev/null +++ b/python/helpers/pydev/pydev_console/io.py @@ -0,0 +1,76 @@ +import threading + + +class PipeIO(object): + """Thread-safe pipe with blocking reads and writes + """ + MAX_BUFFER_SIZE = 4096 + + def __init__(self): + self.lock = threading.RLock() + self.bytes_consumed = threading.Condition(self.lock) + self.bytes_produced = threading.Condition(self.lock) + + self.buffer = bytearray() + self.read_pos = 0 + + def _bytes_available(self): + return self.read_pos < len(self.buffer) + + def _reset_buffer(self): + self.buffer = bytearray() + self.read_pos = 0 + + def read(self, sz): + """Reads `sz` bytes at most + + Blocks until some data is available in buffer. + + :param sz: the maximum count of bytes to read + :return: bytes read + """ + self.lock.acquire() + try: + while not self._bytes_available(): + self.bytes_produced.wait() + + read_until_pos = min(self.read_pos + sz, len(self.buffer)) + result = bytes(self.buffer[self.read_pos:read_until_pos]) + + self.read_pos = read_until_pos + self.bytes_consumed.notifyAll() + + return result + finally: + self.lock.release() + + def write(self, buf): + """Writes `buf` content + + Blocks until all `buf` written. + + :param buf: bytes to write + :return: None + """ + self.lock.acquire() + try: + buf_pos = 0 + while True: + if buf_pos == len(buf): + break + + if len(self.buffer) == self.MAX_BUFFER_SIZE: + while self.read_pos < self.MAX_BUFFER_SIZE: + self.bytes_consumed.wait() + + self._reset_buffer() + + bytes_to_write = min(len(buf) - buf_pos, self.MAX_BUFFER_SIZE - len(self.buffer)) + new_buf_pos = buf_pos + bytes_to_write + + self.buffer.extend(buf[buf_pos:new_buf_pos]) + self.bytes_produced.notifyAll() + + buf_pos = new_buf_pos + finally: + self.lock.release() diff --git a/python/helpers/pydev/pydev_console/thrift_rpc.py b/python/helpers/pydev/pydev_console/thrift_rpc.py index 51d60cc48a37..bba4e8e915fe 100644 --- a/python/helpers/pydev/pydev_console/thrift_rpc.py +++ b/python/helpers/pydev/pydev_console/thrift_rpc.py @@ -1,20 +1,18 @@ +import socket import threading -from pydev_console.thrift_transport import TBidirectionalClientTransport, TSyncClient +from pydev_console.thrift_transport import TSyncClient, open_transports_as_client, _create_client_server_transports from thriftpy.protocol import TBinaryProtocolFactory from thriftpy.server import TThreadedServer from thriftpy.thrift import TProcessor def make_rpc_client(client_service, host, port, proto_factory=TBinaryProtocolFactory()): - # instantiate client - transport = TBidirectionalClientTransport(host, port) - protocol = proto_factory.get_protocol(transport) - transport.open() + client_transport, server_transport = open_transports_as_client((host, port)) - client = TSyncClient(client_service, protocol) + client_protocol = proto_factory.get_protocol(client_transport) - server_transport = transport.get_server_transport() + client = TSyncClient(client_service, client_protocol) return client, server_transport @@ -33,3 +31,42 @@ def start_rpc_server(server_transport, server_service, server_handler, proto_fac t.start() return server + + +def start_rpc_server_and_make_client(host, port, server_service, client_service, server_handler, + proto_factory=TBinaryProtocolFactory(strict_read=False, strict_write=False)): + server_socket = socket.socket(socket.AF_INET, socket.SOCK_STREAM) + server_socket.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) + + server_socket.bind((host, port)) + server_socket.listen(1) + + t = threading.Thread(target=_rpc_server, args=(server_socket, server_service, client_service, server_handler, proto_factory)) + # t.setDaemon(self.daemon) + t.start() + + return server_socket + + +def _rpc_server(server_socket, server_service, client_service, server_handler, proto_factory): + client_socket, address = server_socket.accept() + + client_transport, server_transport = _create_client_server_transports(client_socket) + + # setup server + processor = TProcessor(server_service, server_handler) + + # todo as `server.serve()` is excessive we may want to get rid of `server` as `TThreadedServer` + server = TThreadedServer(processor, server_transport, iprot_factory=proto_factory) + + client = server.trans.accept() + t = threading.Thread(target=server.handle, args=(client,)) + # t.setDaemon(self.daemon) + t.start() + + client_protocol = proto_factory.get_protocol(client_transport) + + client = TSyncClient(client_service, client_protocol) + + # todo fix broken encapsulation + server_handler.rpc_client = client diff --git a/python/helpers/pydev/pydev_console/thrift_transport.py b/python/helpers/pydev/pydev_console/thrift_transport.py index b0b26528169e..0a1468eb5ef0 100644 --- a/python/helpers/pydev/pydev_console/thrift_transport.py +++ b/python/helpers/pydev/pydev_console/thrift_transport.py @@ -2,8 +2,8 @@ import socket import struct import sys import threading -from io import BytesIO, SEEK_CUR +from pydev_console.io import PipeIO from thriftpy.thrift import TClient from thriftpy.transport import TTransportBase, readall @@ -11,28 +11,13 @@ REQUEST = 0 RESPONSE = 1 -def _read_anything(read_fn, sz): - buff = b'' - have = 0 - while have == 0: - chunk = read_fn(sz - have) - - have += len(chunk) - buff += chunk - - return buff - - class MultiplexedSocketReader(object): def __init__(self, s): self._socket = s - self._request_buffer = BytesIO() - self._response_buffer = BytesIO() - - self._request_buffer_lock = threading.RLock() - self._response_buffer_lock = threading.RLock() + self._request_pipe = PipeIO() + self._response_pipe = PipeIO() self._read_socket_lock = threading.RLock() @@ -40,19 +25,13 @@ class MultiplexedSocketReader(object): """ Invoked form server-side of the bidirectional transport. """ - return _read_anything(self._read_request_buffer, sz) + return self._request_pipe.read(sz) def read_response(self, sz): """ Invoked form client-side of the bidirectional transport. """ - return _read_anything(self._read_response_buffer, sz) - - def _read_request_buffer(self, sz): - return self._request_buffer.read(sz) - - def _read_response_buffer(self, sz): - return self._response_buffer.read(sz) + return self._response_pipe.read(sz) # noinspection PyUnusedLocal def _try_fill_buffer(self, sz): @@ -63,13 +42,9 @@ class MultiplexedSocketReader(object): direction, frame = self._read_frame() if direction == REQUEST: - with self._request_buffer_lock: - self._request_buffer.write(frame) - self._request_buffer.seek(-len(frame), SEEK_CUR) + self._request_pipe.write(frame) elif direction == RESPONSE: - with self._response_buffer_lock: - self._response_buffer.write(frame) - self._response_buffer.seek(-len(frame), SEEK_CUR) + self._response_pipe.write(frame) def _read_frame(self): # todo introduce sz argument @@ -110,7 +85,7 @@ class FramedWriter(object): MAX_BUFFER_SIZE = 4096 def __init__(self): - self._buffer = BytesIO() + self._buffer = bytearray() def _get_writer(self): raise NotImplementedError @@ -130,21 +105,21 @@ class FramedWriter(object): if buffer_size + bytes_to_write > self.MAX_BUFFER_SIZE: write_till_byte = bytes_written + (self.MAX_BUFFER_SIZE - buffer_size) - self._buffer.write(buf[bytes_written:write_till_byte]) + self._buffer.extend(buf[bytes_written:write_till_byte]) self.flush() bytes_written = write_till_byte else: # the whole buffer processed - self._buffer.write(buf[bytes_written:]) + self._buffer.extend(buf[bytes_written:]) bytes_written = buf_len def flush(self): # reset wbuf before write/flush to preserve state on underlying failure - out = self._buffer.getvalue() + out = bytes(self._buffer) # prepend the message with the direction byte out = struct.pack("b", self._get_write_direction()) + out - self._buffer = BytesIO() + self._buffer = bytearray() # N.B.: Doing this string concatenation is WAY cheaper than making # two separate calls to the underlying socket object. Socket writes in @@ -154,21 +129,18 @@ class FramedWriter(object): self._get_writer().write(struct.pack("!i", len(out)) + out) def close(self): - self._buffer.close() + self._buffer = bytearray() + pass class TBidirectionalClientTransport(TTransportBase, FramedWriter): - def __init__(self, host, port): + def __init__(self, client_socket, reader, writer): super(TBidirectionalClientTransport, self).__init__() - self.host = host - self.port = port - # the following properties will be initialized in `open()` - self._client_socket = None - self._reader = None - self._writer = None - self._server_transport = None + self._client_socket = client_socket + self._reader = reader + self._writer = writer def _get_writer(self): return self._writer @@ -182,29 +154,10 @@ class TBidirectionalClientTransport(TTransportBase, FramedWriter): """ return self._reader.read_response(sz) - def get_server_transport(self): - if not self._server_transport: - raise Exception - - return self._server_transport - def is_open(self): # todo we may try to monitor reads and writes and put a flag if they fail return self._client_socket - def open(self): - client_socket = socket.socket(socket.AF_INET, socket.SOCK_STREAM) - client_socket.connect((self.host, self.port)) - - self._client_socket = client_socket - - self._reader = MultiplexedSocketReader(self._client_socket) - self._reader.start_reading() - - self._writer = SocketWriter(self._client_socket) - - self._server_transport = TReversedServerTransport(self._reader.read_request, self._writer) - def close(self): # todo should we do something with buffer of the multiplexed reader self._client_socket.shutdown(socket.SHUT_RDWR) @@ -268,3 +221,34 @@ class TSyncClient(TClient): def _req(self, _api, *args, **kwargs): with self._lock: return super(TSyncClient, self)._req(_api, *args, **kwargs) + + +def open_transports_as_client(addr): + client_socket = socket.socket(socket.AF_INET, socket.SOCK_STREAM) + client_socket.connect(addr) + + return _create_client_server_transports(client_socket) + + +def open_transports_as_server(addr): + server_socket = socket.socket(socket.AF_INET, socket.SOCK_STREAM) + server_socket.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1) + + server_socket.bind(addr) + server_socket.listen(1) + + client_socket, address = server_socket.accept() + + raise _create_client_server_transports(client_socket) + + +def _create_client_server_transports(sock): + reader = MultiplexedSocketReader(sock) + reader.start_reading() + + writer = SocketWriter(sock) + + client_transport = TBidirectionalClientTransport(sock, reader, writer) + server_transport = TReversedServerTransport(client_transport._reader.read_request, client_transport._writer) + + return client_transport, server_transport diff --git a/python/helpers/pydev/pydevconsole.py b/python/helpers/pydev/pydevconsole.py index 1fa0ab1145d3..70d4af58eaf6 100644 --- a/python/helpers/pydev/pydevconsole.py +++ b/python/helpers/pydev/pydevconsole.py @@ -1,8 +1,9 @@ ''' Entry point module to start the interactive console. ''' +from _pydev_bundle._pydev_getopt import gnu_getopt from _pydev_imps._pydev_saved_modules import thread -from pydev_console.thrift_rpc import make_rpc_client, start_rpc_server +from pydev_console.thrift_rpc import make_rpc_client, start_rpc_server, start_rpc_server_and_make_client start_new_thread = thread.start_new_thread @@ -348,9 +349,8 @@ def enable_thrift_logging(): logger.addHandler(ch) -def start_server(host, port, client_port, client_host = None): - if not client_host: - client_host = host +def start_server(): + # 0. General stuff #replace exit (see comments on method) #note that this does not work in jython!!! (sys method can't be replaced). @@ -360,9 +360,42 @@ def start_server(host, port, client_port, client_host = None): enable_thrift_logging() + server_service = console_thrift.PythonConsole + client_service = console_thrift.IDE + + # 1. Start Python console server + + + interpreter = InterpreterInterface(threading.currentThread(), None, None) + + # `InterpreterInterface` implements all methods required for the handler + server_handler = interpreter + + server_socket = start_rpc_server_and_make_client('', 0, server_service, client_service, server_handler) + + + # 2. Print server port for the IDE + + _, server_port = server_socket.getsockname() + print(server_port) + + # 3. Wait for IDE to connect to the server + + process_exec_queue(interpreter) + + +def start_client(host, port): + #replace exit (see comments on method) + #note that this does not work in jython!!! (sys method can't be replaced). + sys.exit = do_exit + + from pydev_console.thrift_communication import console_thrift + + enable_thrift_logging() + client_service = console_thrift.IDE - client, server_transport = make_rpc_client(client_service, client_host, client_port) + client, server_transport = make_rpc_client(client_service, host, port) interpreter = InterpreterInterface(threading.currentThread(), None, client) @@ -545,18 +578,32 @@ if __name__ == '__main__': #'Variables' and 'Expressions' views stopped working when debugging interactive console import pydevconsole sys.stdin = pydevconsole.BaseStdIn(sys.stdin) - port, client_port = sys.argv[1:3] - from _pydev_bundle import pydev_localhost - if int(port) == 0 and int(client_port) == 0: - (h, p) = pydev_localhost.get_socket_name() + # parse command-line arguments + optlist, _ = gnu_getopt(sys.argv, 'm:h:p', ['mode=', 'host=', 'port=']) + mode = None + host = None + port = None + for opt, arg in optlist: + if opt in ('-m', '--mode'): + mode = arg + elif opt in ('-h', '--host'): + host = arg + elif opt in ('-p', '--port'): + port = int(arg) - client_port = p + if mode not in ('client', 'server'): + sys.exit(-1) - if len(sys.argv) > 4: - host = sys.argv[3] - client_host = sys.argv[4] - else: - host = client_host = pydev_localhost.get_localhost() + if mode == 'client': + if not port: + # port must be set for client + sys.exit(-1) - pydevconsole.start_server(host, int(port), int(client_port), client_host) + if not host: + from _pydev_bundle import pydev_localhost + host = client_host = pydev_localhost.get_localhost() + + pydevconsole.start_client(host, port) + elif mode == 'server': + pydevconsole.start_server() diff --git a/python/python-console/src/com/jetbrains/python/console/thrift/client/TNettyClientTransport.kt b/python/python-console/src/com/jetbrains/python/console/thrift/client/TNettyClientTransport.kt index 282442d43381..c91789903612 100644 --- a/python/python-console/src/com/jetbrains/python/console/thrift/client/TNettyClientTransport.kt +++ b/python/python-console/src/com/jetbrains/python/console/thrift/client/TNettyClientTransport.kt @@ -20,6 +20,7 @@ import org.apache.thrift.transport.TServerTransport import org.apache.thrift.transport.TTransport import java.io.PipedInputStream import java.io.PipedOutputStream +import java.util.concurrent.atomic.AtomicBoolean class TNettyClientTransport(private val host: String, private val port: Int) : TTransport() { @@ -37,11 +38,24 @@ class TNettyClientTransport(private val host: String, private val serverAcceptedTransport = TNettyTransport() val serverTransport: TServerTransport = object : TServerTransport() { + private val acceptedOnce = AtomicBoolean(false) + override fun listen() { // TODO ? } - override fun acceptImpl(): TTransport = serverAcceptedTransport + override fun acceptImpl(): TTransport { + if (acceptedOnce.compareAndSet(false, true)) { + return serverAcceptedTransport + } + + val lock = Object() + while (true) { + synchronized(lock) { + lock.wait() + } + } + } override fun close() { // TODO ? diff --git a/python/python-console/src/com/jetbrains/python/console/thrift/server/TNettyServerTransport.kt b/python/python-console/src/com/jetbrains/python/console/thrift/server/TNettyServerTransport.kt index 2ecf6eec403d..b6b7a7fa33c3 100644 --- a/python/python-console/src/com/jetbrains/python/console/thrift/server/TNettyServerTransport.kt +++ b/python/python-console/src/com/jetbrains/python/console/thrift/server/TNettyServerTransport.kt @@ -20,6 +20,7 @@ import org.apache.thrift.transport.TServerTransport import org.apache.thrift.transport.TTransport import org.apache.thrift.transport.TTransportException import java.util.concurrent.BlockingQueue +import java.util.concurrent.CountDownLatch import java.util.concurrent.LinkedBlockingQueue import java.util.concurrent.TimeUnit import java.util.concurrent.atomic.AtomicBoolean @@ -30,8 +31,19 @@ import java.util.concurrent.atomic.AtomicBoolean class TNettyServerTransport(port: Int) : TServerTransport() { private val nettyServer: NettyServer = NettyServer(port) + @Throws(TTransportException::class) override fun listen() { - nettyServer.listen() + try { + nettyServer.listen() + } + catch (e: InterruptedException) { + throw TTransportException(e) + } + } + + @Throws(InterruptedException::class) + fun waitForBind() { + nettyServer.waitForBind() } override fun acceptImpl(): TTransport = nettyServer.accept() @@ -60,6 +72,8 @@ class TNettyServerTransport(port: Int) : TServerTransport() { // the accepted connection to the worker. private val workerGroup = NioEventLoopGroup(0, ConcurrencyUtil.newNamedThreadFactory("Python Console NIO Event Loop Worker")) + private val serverBound = CountDownLatch(1) + fun listen() { // TODO check state! @@ -143,17 +157,16 @@ class TNettyServerTransport(port: Int) : TServerTransport() { // start the server. Here, we bind to the port 8080 of all NICs (network // interface cards) in the machine. You can now call the bind() method as // many times as you want (with different bind addresses.) - val f = b.bind(port).sync() // (7) + b.bind(port).sync() // (7) LOG.debug("Running Netty server on $port") - // TODO move to `close()` - /* - // Wait until the server socket is closed. - // In this example, this does not happen, but you can do that to gracefully - // shut down your server. - f.channel().closeFuture().sync() - */ + serverBound.countDown() + } + + @Throws(InterruptedException::class) + fun waitForBind() { + serverBound.await() } fun accept(): TTransport { @@ -183,6 +196,14 @@ class TNettyServerTransport(port: Int) : TServerTransport() { bossGroup.shutdownGracefully() // TODO close server channel! + + // TODO should we wait for channel close move to `close()` + /* + // Wait until the server socket is closed. + // In this example, this does not happen, but you can do that to gracefully + // shut down your server. + f.channel().closeFuture().sync() + */ } } } diff --git a/python/src/com/jetbrains/python/console/PydevConsoleCommunication.java b/python/src/com/jetbrains/python/console/PydevConsoleCommunication.java index 6f827601ebac..a736c5fc6878 100644 --- a/python/src/com/jetbrains/python/console/PydevConsoleCommunication.java +++ b/python/src/com/jetbrains/python/console/PydevConsoleCommunication.java @@ -25,6 +25,7 @@ import com.jetbrains.python.console.parsing.PythonConsoleData; import com.jetbrains.python.console.pydev.AbstractConsoleCommunication; import com.jetbrains.python.console.pydev.InterpreterResponse; import com.jetbrains.python.console.pydev.PydevCompletionVariant; +import com.jetbrains.python.console.thrift.client.TNettyClientTransport; import com.jetbrains.python.console.thrift.server.TNettyServerTransport; import com.jetbrains.python.debugger.*; import com.jetbrains.python.debugger.containerview.PyViewNumericContainerAction; @@ -32,14 +33,11 @@ import com.jetbrains.python.debugger.pydev.GetVariableCommand; import org.apache.thrift.TException; import org.apache.thrift.protocol.TBinaryProtocol; import org.apache.thrift.server.TThreadPoolServer; +import org.apache.thrift.transport.TServerTransport; import org.apache.thrift.transport.TTransport; import org.jetbrains.annotations.NotNull; import org.jetbrains.annotations.Nullable; -import java.lang.reflect.InvocationHandler; -import java.lang.reflect.Method; -import java.lang.reflect.Proxy; -import java.net.MalformedURLException; import java.util.ArrayList; import java.util.Collections; import java.util.List; @@ -47,7 +45,6 @@ import java.util.Map; import java.util.concurrent.CompletableFuture; import java.util.concurrent.ConcurrentHashMap; import java.util.concurrent.Future; -import java.util.concurrent.locks.ReentrantLock; import java.util.stream.Collectors; import static com.jetbrains.python.console.PydevConsoleCommunicationUtil.*; @@ -101,40 +98,23 @@ public class PydevConsoleCommunication extends AbstractConsoleCommunication impl @Nullable private XCompositeNode myCurrentRootNode; - private final int myClientPort; - /** - * Initializes the xml-rpc communication. - * - * @param port the port where the communication should happen. - * @param process this is the process that was spawned (server for the XML-RPC) - * @throws MalformedURLException + * Initializes the bidirectional RPC communication. */ - public PydevConsoleCommunication(Project project, int port, Process process, int clientPort) throws Exception { - this(project, null, port, process, clientPort); - } - - public PydevConsoleCommunication(Project project, String host, int port, Process process, int clientPort) throws Exception { + public PydevConsoleCommunication(Project project) { super(project); - myClientPort = clientPort; } - /** - * @return connection future - */ - @NotNull - public Future startServer() { + public void startServer(int port) throws InterruptedException { IDEHandler serverHandler = new IDEHandler(); IDE.Processor serverProcessor = new IDE.Processor<>(serverHandler); //noinspection IOResourceOpenedButNotSafelyClosed - TNettyServerTransport serverTransport = new TNettyServerTransport(myClientPort); + TNettyServerTransport serverTransport = new TNettyServerTransport(port); TThreadPoolServer server = new TThreadPoolServer( new TThreadPoolServer.Args(serverTransport).processor(serverProcessor).protocolFactory(new TBinaryProtocol.Factory()) .stopTimeoutVal(1)); - // @alexander todo do not `Thread.start()` here! - new Thread(() -> server.serve()).start(); - SettableFuture connectionFuture = SettableFuture.create(); + ApplicationManager.getApplication().executeOnPooledThread(() -> server.serve()); ApplicationManager.getApplication().executeOnPooledThread(() -> { TTransport clientTransport = serverTransport.getReverseTransport(); @@ -142,7 +122,7 @@ public class PydevConsoleCommunication extends AbstractConsoleCommunication impl PythonConsole.Client client = new PythonConsole.Client(clientProtocol); this.myServer = server; - this.myClient = syncPythonConsoleIface(client); + this.myClient = PythonConsoleClientUtil.synchronizedPythonConsoleClient(PydevConsoleCommunication.class.getClassLoader(), client); PyDebugValueExecutionService executionService = PyDebugValueExecutionService.getInstance(myProject); executionService.sessionStarted(this); @@ -152,19 +132,43 @@ public class PydevConsoleCommunication extends AbstractConsoleCommunication impl executionService.cancelSubmittedTasks(PydevConsoleCommunication.this); } }); - - // notify that we have connected - connectionFuture.set(null); }); - try { - Thread.sleep(1000L); - } - catch (InterruptedException e) { - Thread.currentThread().interrupt(); - } + serverTransport.waitForBind(); + } - return connectionFuture; + public void startClient(@NotNull String host, int port, @NotNull Process pythonConsoleProcess) { + ApplicationManager.getApplication().executeOnPooledThread(() -> { + TNettyClientTransport clientTransport = new TNettyClientTransport(host, port); + clientTransport.open(); + + TBinaryProtocol clientProtocol = new TBinaryProtocol(clientTransport); + PythonConsole.Client client = new PythonConsole.Client(clientProtocol); + + TServerTransport serverTransport = clientTransport.getServerTransport(); + + IDEHandler serverHandler = new IDEHandler(); + IDE.Processor serverProcessor = new IDE.Processor<>(serverHandler); + + TThreadPoolServer server = new TThreadPoolServer( + new TThreadPoolServer.Args(serverTransport).processor(serverProcessor).protocolFactory(new TBinaryProtocol.Factory()) + .stopTimeoutVal(1)); + + ApplicationManager.getApplication().executeOnPooledThread(() -> server.serve()); + + this.myServer = server; + this.myClient = PythonConsoleClientUtil.synchronizedPythonConsoleClient(PydevConsoleCommunication.class.getClassLoader(), client, + pythonConsoleProcess); + + PyDebugValueExecutionService executionService = PyDebugValueExecutionService.getInstance(myProject); + executionService.sessionStarted(this); + addFrameListener(new PyFrameListener() { + @Override + public void frameChanged() { + executionService.cancelSubmittedTasks(PydevConsoleCommunication.this); + } + }); + }); } public boolean handshake() { @@ -179,9 +183,6 @@ public class PydevConsoleCommunication extends AbstractConsoleCommunication impl return false; } - /** - * Sends {@link #CLOSE} message to the Python console script. - */ private void sendCloseMessageToScript() { if (this.myClient != null) { new Task.Backgroundable(myProject, "Close Console Communication", true) { @@ -787,23 +788,4 @@ public class PydevConsoleCommunication extends AbstractConsoleCommunication impl return execIPythonEditor(path); } } - - @NotNull - private static PythonConsole.Iface syncPythonConsoleIface(@NotNull PythonConsole.Iface iface) { - ReentrantLock lock = new ReentrantLock(); - return (PythonConsole.Iface)Proxy - .newProxyInstance(PydevConsoleCommunication.class.getClassLoader(), new Class[]{PythonConsole.Iface.class}, - new InvocationHandler() { - @Override - public Object invoke(Object proxy, Method method, Object[] args) throws Throwable { - lock.lock(); - try { - return method.invoke(iface, args); - } - finally { - lock.unlock(); - } - } - }); - } } diff --git a/python/src/com/jetbrains/python/console/PydevConsoleRunnerImpl.java b/python/src/com/jetbrains/python/console/PydevConsoleRunnerImpl.java index 77d5d45f9084..20fb8392aa71 100644 --- a/python/src/com/jetbrains/python/console/PydevConsoleRunnerImpl.java +++ b/python/src/com/jetbrains/python/console/PydevConsoleRunnerImpl.java @@ -44,7 +44,6 @@ import com.intellij.openapi.project.DumbAwareAction; import com.intellij.openapi.project.Project; import com.intellij.openapi.projectRoots.Sdk; import com.intellij.openapi.ui.Messages; -import com.intellij.openapi.util.Couple; import com.intellij.openapi.util.Disposer; import com.intellij.openapi.util.Key; import com.intellij.openapi.util.SystemInfo; @@ -97,7 +96,6 @@ import java.util.List; import java.util.Map; import java.util.Scanner; import java.util.concurrent.Future; -import java.util.concurrent.TimeUnit; import java.util.stream.Collectors; import static com.intellij.execution.runners.AbstractConsoleRunnerWithHistory.registerActionShortcuts; @@ -133,8 +131,6 @@ public class PydevConsoleRunnerImpl implements PydevConsoleRunner { private static final long HANDSHAKE_TIMEOUT = 60000; - private static final long CONNECTION_TIMEOUT = HANDSHAKE_TIMEOUT; - private RemoteConsoleProcessData myRemoteConsoleProcessData; private String myConsoleTitle = null; @@ -391,9 +387,7 @@ public class PydevConsoleRunnerImpl implements PydevConsoleRunner { } PythonHelper.CONSOLE.addToGroup(group, cmd); - for (int port : ports) { - group.addParameter(String.valueOf(port)); - } + group.addParameters("--mode=client", "--port=" + ports[1]); return cmd; } @@ -431,27 +425,19 @@ public class PydevConsoleRunnerImpl implements PydevConsoleRunner { EncodingEnvironmentUtil.setLocaleEnvironmentIfMac(envs, generalCommandLine.getCharset()); Future connectionFuture; + + myPydevConsoleCommunication = new PydevConsoleCommunication(myProject); // first of all - start server try { // todo use process in `PydevConsoleCommunication` - myPydevConsoleCommunication = new PydevConsoleCommunication(myProject, myPorts[0], null, myPorts[1]); // todo we might want to add a timeout here on start - connectionFuture = myPydevConsoleCommunication.startServer(); + myPydevConsoleCommunication.startServer(myPorts[1]); } catch (Exception e) { throw new ExecutionException(e.getMessage(), e); } - final Process server = generalCommandLine.createProcess(); - - try { - // todo we might want to add a timeout here - connectionFuture.get(CONNECTION_TIMEOUT, TimeUnit.SECONDS); - } - catch (Exception e) { - throw new ExecutionException(e.getMessage(), e); - } - return server; + return generalCommandLine.createProcess(); } } @@ -459,9 +445,9 @@ public class PydevConsoleRunnerImpl implements PydevConsoleRunner { return PYDEV_PYDEVCONSOLE_PY; } - public static Couple getRemotePortsFromProcess(Process process) throws ExecutionException { + public static int getRemotePortFromProcess(@NotNull Process process) throws ExecutionException { @SuppressWarnings("IOResourceOpenedButNotSafelyClosed") Scanner s = new Scanner(process.getInputStream()); - return Couple.of(readInt(s, process), readInt(s, process)); + return readInt(s, process); } private static int readInt(Scanner s, Process process) throws ExecutionException { diff --git a/python/src/com/jetbrains/python/console/PydevRemoteConsoleCommunication.java b/python/src/com/jetbrains/python/console/PydevRemoteConsoleCommunication.java index 995ce820e3d5..20e4b202af36 100644 --- a/python/src/com/jetbrains/python/console/PydevRemoteConsoleCommunication.java +++ b/python/src/com/jetbrains/python/console/PydevRemoteConsoleCommunication.java @@ -22,17 +22,9 @@ import com.intellij.openapi.project.Project; */ public class PydevRemoteConsoleCommunication extends PydevConsoleCommunication { /** - * Initializes the xml-rpc communication. - * - * @param port the port where the communication should happen. - * @param process this is the process that was spawned (server for the XML-RPC) + * Initializes the bidirectional RPC communication. */ - public PydevRemoteConsoleCommunication(Project project, int port, Process process, int clientPort) - throws Exception { - super(project, port, process, clientPort); - } - - public PydevRemoteConsoleCommunication(Project project, int port, Process process, int clientPort, String clientHost) throws Exception { - super(project, clientHost, port, process, clientPort); + public PydevRemoteConsoleCommunication(Project project) { + super(project); } } diff --git a/python/src/com/jetbrains/python/console/PythonConsoleClientUtil.kt b/python/src/com/jetbrains/python/console/PythonConsoleClientUtil.kt new file mode 100644 index 000000000000..0c5d809e0778 --- /dev/null +++ b/python/src/com/jetbrains/python/console/PythonConsoleClientUtil.kt @@ -0,0 +1,53 @@ +// Copyright 2000-2018 JetBrains s.r.o. Use of this source code is governed by the Apache 2.0 license that can be found in the LICENSE file. +@file:JvmName("PythonConsoleClientUtil") + +package com.jetbrains.python.console + +import com.intellij.openapi.application.ApplicationManager +import java.lang.reflect.InvocationHandler +import java.lang.reflect.Proxy +import java.util.concurrent.Callable +import java.util.concurrent.TimeUnit +import java.util.concurrent.TimeoutException +import java.util.concurrent.locks.ReentrantLock + +@JvmOverloads +fun synchronizedPythonConsoleClient(loader: ClassLoader, + delegate: PythonConsole.Iface, + pythonConsoleProcess: Process? = null): PythonConsole.Iface { + val lock = ReentrantLock() + return Proxy.newProxyInstance(loader, arrayOf>(PythonConsole.Iface::class.java), + InvocationHandler { _, method, args -> + val future = ApplicationManager.getApplication().executeOnPooledThread(Callable { + lock.lock() + try { + if (args == null) { + return@Callable method.invoke(delegate) + } + else { + return@Callable method.invoke(delegate, *args) + } + } + finally { + lock.unlock() + } + }) + + if (pythonConsoleProcess == null) { + return@InvocationHandler future.get() + } + + while (true) { + try { + return@InvocationHandler future.get(10L, TimeUnit.MILLISECONDS) + } + catch (e: TimeoutException) { + if (!pythonConsoleProcess.isAlive) { + val exitValue = pythonConsoleProcess.exitValue() + throw RuntimeException( + "Console already exited with value: $exitValue while waiting for an answer.\n") + } + } + } + }) as PythonConsole.Iface +}