55import asyncio
66import argparse
77import traceback
8+ from contextlib import asynccontextmanager
89from pathlib import Path
910from typing import (
1011 TYPE_CHECKING ,
1415 ContextManager ,
1516 Type ,
1617 TypeAlias ,
18+ AsyncIterator ,
1719)
1820from concurrent .futures import ThreadPoolExecutor
1921
2022if TYPE_CHECKING :
23+ import telnetlib3
2124 from telnetlib3 .stream_reader import TelnetReader
2225 from telnetlib3 .stream_writer import TelnetWriter
2326
@@ -308,7 +311,8 @@ def target(term: RemoteTerminal) -> int:
308311 pass
309312
310313
311- async def run_server (
314+ @asynccontextmanager
315+ async def run_telnet_server (
312316 bind : str ,
313317 port : int ,
314318 robot_check : bool ,
@@ -317,7 +321,7 @@ async def run_server(
317321 console_cls : type [Console ],
318322 namespace : argparse .Namespace ,
319323 executor : ThreadPoolExecutor ,
320- ) -> None :
324+ ) -> AsyncIterator [ telnetlib3 . Server ] :
321325 import telnetlib3
322326
323327 shell = make_telnet_shell (namespace , console_cls , idle_timeout , executor )
@@ -366,7 +370,15 @@ async def guarded_shell(reader: TelnetReader, writer: TelnetWriter) -> None:
366370 assert sockets is not None
367371 actual_bind , actual_port = sockets [0 ].getsockname ()[:2 ]
368372 print (f"Running telnet server on { actual_bind } :{ actual_port } ..." , flush = True )
369- await asyncio .Future ()
373+ try :
374+ yield server
375+ finally :
376+ assert server ._server is not None
377+ server ._server .close ()
378+ for client in server .clients :
379+ if client .reader is not None :
380+ client .reader .feed_eof ()
381+ await server .wait_closed ()
370382
371383
372384def main (
@@ -425,8 +437,9 @@ def main(
425437
426438 try :
427439 with ThreadPoolExecutor (max_workers = 32 ) as executor :
428- asyncio .run (
429- run_server (
440+
441+ async def async_main () -> None :
442+ async with run_telnet_server (
430443 bind ,
431444 port ,
432445 robot_check ,
@@ -435,8 +448,10 @@ def main(
435448 console_cls ,
436449 namespace ,
437450 executor ,
438- )
439- )
451+ ):
452+ await asyncio .Future ()
453+
454+ asyncio .run (async_main ())
440455 except KeyboardInterrupt :
441456 pass
442457
0 commit comments