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
@@ -311,7 +314,8 @@ def target(term: RemoteTerminal) -> int:
311314 pass
312315
313316
314- async def run_server (
317+ @asynccontextmanager
318+ async def run_telnet_server (
315319 bind : str ,
316320 port : int ,
317321 robot_check : bool ,
@@ -320,7 +324,7 @@ async def run_server(
320324 console_cls : type [Console ],
321325 namespace : argparse .Namespace ,
322326 executor : ThreadPoolExecutor ,
323- ) -> None :
327+ ) -> AsyncIterator [ telnetlib3 . Server ] :
324328 import telnetlib3
325329
326330 shell = make_telnet_shell (namespace , console_cls , idle_timeout , executor )
@@ -369,7 +373,15 @@ async def guarded_shell(reader: TelnetReader, writer: TelnetWriter) -> None:
369373 assert sockets is not None
370374 actual_bind , actual_port = sockets [0 ].getsockname ()[:2 ]
371375 print (f"Running telnet server on { actual_bind } :{ actual_port } ..." , flush = True )
372- await asyncio .Future ()
376+ try :
377+ yield server
378+ finally :
379+ assert server ._server is not None
380+ server ._server .close ()
381+ for client in server .clients :
382+ if client .reader is not None :
383+ client .reader .feed_eof ()
384+ await server .wait_closed ()
373385
374386
375387def main (
@@ -428,8 +440,9 @@ def main(
428440
429441 try :
430442 with ThreadPoolExecutor (max_workers = 32 ) as executor :
431- asyncio .run (
432- run_server (
443+
444+ async def async_main () -> None :
445+ async with run_telnet_server (
433446 bind ,
434447 port ,
435448 robot_check ,
@@ -438,8 +451,10 @@ def main(
438451 console_cls ,
439452 namespace ,
440453 executor ,
441- )
442- )
454+ ):
455+ await asyncio .Future ()
456+
457+ asyncio .run (async_main ())
443458 except KeyboardInterrupt :
444459 pass
445460
0 commit comments