fix: 优化超时输入倒计时提示

Co-authored-by: factory-droid[bot] <138933559+factory-droid[bot]@users.noreply.github.com>
This commit is contained in:
sladro 2025-10-15 10:07:21 +08:00
parent bf7eb9ca9d
commit d887e7c313
2 changed files with 452 additions and 40 deletions

View File

@ -3,16 +3,97 @@
提供带超时机制的输入函数,超时后自动使用默认值
"""
from __future__ import annotations
import math
import os
import sys
import threading
from queue import Queue, Empty
from typing import Optional, List
import time
from contextlib import contextmanager
from queue import Empty, Queue
from typing import Any, Iterable, List, Optional, Tuple
from rich.console import Console
console = Console()
class _InteractiveInputUnavailable(Exception):
"""指示当前环境不支持逐字符交互读取"""
_COUNTDOWN_TEMPLATE = "输入将在 {seconds} 秒后自动使用默认值。开始输入后倒计时会暂停。"
_INACTIVITY_GRACE_SECONDS = 2.0
class _CountdownRenderer:
"""负责渲染倒计时提示和输入行"""
_ESC = "\x1b"
def __init__(self, prompt: str, timeout: Optional[int]):
self.prompt = prompt
self.timeout = timeout
self.active = bool(timeout and timeout > 0)
self._has_drawn = False
self._last_seconds: Optional[int] = None
def render(self, buffer: List[str], remaining: Optional[float]) -> None:
if not self.active:
if not self._has_drawn:
sys.stdout.write(self.prompt + "".join(buffer))
sys.stdout.flush()
self._has_drawn = True
else:
self._render_prompt(buffer)
return
seconds = self._format_seconds(remaining)
if not self._has_drawn:
self._draw_initial(buffer, seconds)
return
if seconds != self._last_seconds:
self._rewrite_message(seconds)
self._render_prompt(buffer)
def finish(self) -> None:
if not self._has_drawn:
return
sys.stdout.write("\n")
sys.stdout.flush()
self._has_drawn = False
def _format_seconds(self, remaining: Optional[float]) -> int:
if remaining is None:
return int(self.timeout or 0)
return max(int(math.ceil(remaining)), 0)
def _draw_initial(self, buffer: List[str], seconds: int) -> None:
message = _COUNTDOWN_TEMPLATE.format(seconds=seconds)
sys.stdout.write(message + "\n")
sys.stdout.write(self.prompt + "".join(buffer))
sys.stdout.flush()
self._has_drawn = True
self._last_seconds = seconds
def _rewrite_message(self, seconds: int) -> None:
message = _COUNTDOWN_TEMPLATE.format(seconds=seconds)
sys.stdout.write("\r")
sys.stdout.write(self._ESC + "[2K")
sys.stdout.write(self._ESC + "[F")
sys.stdout.write(self._ESC + "[2K")
sys.stdout.write(message + "\n")
self._last_seconds = seconds
def _render_prompt(self, buffer: List[str]) -> None:
sys.stdout.write("\r")
sys.stdout.write(self._ESC + "[2K")
sys.stdout.write(self.prompt + "".join(buffer))
sys.stdout.flush()
def timeout_prompt(
prompt: str,
default: Optional[str] = None,
@ -30,44 +111,41 @@ def timeout_prompt(
Returns:
用户输入或默认值
"""
result_queue = Queue()
# 无超时时直接使用标准输入
if timeout is not None and timeout <= 0:
return _blocking_prompt(prompt, default, choices)
def input_thread():
"""输入线程"""
if choices is not None:
# 预处理,便于后续验证
valid_choices = set(choices)
else:
valid_choices = None
# 首选逐字符读取,支持用户输入时暂停倒计时
if _can_use_interactive_input():
try:
user_input = input(prompt)
result_queue.put(user_input)
except EOFError:
# 处理EOF情况
result_queue.put(None)
except Exception as e:
console.print(f"[red]输入错误: {e}[/red]")
result_queue.put(None)
result, timed_out = _interactive_readline(prompt, timeout)
except _InteractiveInputUnavailable:
pass
else:
if timed_out:
console.print(
f"[yellow]⏰ {timeout}秒内未完成输入,自动使用默认值: {default}[/yellow]"
)
return default or ""
# 启动输入线程
thread = threading.Thread(target=input_thread, daemon=True)
thread.start()
result = (result or "").strip()
if not result:
return default or ""
try:
# 等待输入,带超时
result = result_queue.get(timeout=timeout)
if valid_choices is not None and result not in valid_choices:
console.print(f"[yellow]输入无效,使用默认值: {default}[/yellow]")
return default or ""
# 验证选项
if choices and result and result not in choices:
console.print(f"[yellow]输入无效,使用默认值: {default}[/yellow]")
return default or ""
return result
# 空输入使用默认值
if not result or result.strip() == "":
return default or ""
return result.strip()
except Empty:
# 超时 - 需要打印换行以清除输入行
print() # 打印空行以清理输入提示
console.print(f"[yellow]⏰ {timeout}秒内未完成输入,自动使用默认值: {default}[/yellow]")
return default or ""
# 回退到线程+input方案(无法检测逐字符输入,但兼容非TTY环境)
return _threaded_prompt(prompt, default, timeout, valid_choices)
def timeout_choice_prompt(
@ -87,20 +165,274 @@ def timeout_choice_prompt(
Returns:
选择的选项索引(字符串形式,"1","2")
"""
# 显示选项
console.print(f"\n[blue]{prompt}:[/blue]")
for i, option in enumerate(options, 1):
console.print(f"{i}. {option}")
# 构建有效选项列表
valid_choices = [str(i) for i in range(1, len(options) + 1)]
# 获取输入
choice = timeout_prompt(
return timeout_prompt(
"请输入选择: ",
default=default,
timeout=timeout,
choices=valid_choices
choices=valid_choices,
)
return choice
def _blocking_prompt(
prompt: str,
default: Optional[str],
choices: Optional[Iterable[str]],
) -> str:
try:
user_input = input(prompt)
except EOFError:
user_input = None
user_input = (user_input or "").strip()
if not user_input:
return default or ""
if choices is not None and user_input not in choices:
console.print(f"[yellow]输入无效,使用默认值: {default}[/yellow]")
return default or ""
return user_input
def _threaded_prompt(
prompt: str,
default: Optional[str],
timeout: Optional[int],
choices: Optional[Iterable[str]],
) -> str:
result_queue: "Queue[Optional[str]]" = Queue()
def input_thread():
try:
user_input = input(prompt)
result_queue.put(user_input)
except EOFError:
result_queue.put(None)
except Exception as exc: # pragma: no cover - 极端情况下记录错误
console.print(f"[red]输入错误: {exc}[/red]")
result_queue.put(None)
if timeout and timeout > 0:
console.print(_COUNTDOWN_TEMPLATE.format(seconds=timeout))
thread = threading.Thread(target=input_thread, daemon=True)
thread.start()
try:
result = result_queue.get(timeout=timeout or None)
except Empty:
print()
console.print(
f"[yellow]⏰ {timeout}秒内未完成输入,自动使用默认值: {default}[/yellow]"
)
return default or ""
result = (result or "").strip()
if not result:
return default or ""
if choices is not None and result not in choices:
console.print(f"[yellow]输入无效,使用默认值: {default}[/yellow]")
return default or ""
return result
def _can_use_interactive_input() -> bool:
stdin = sys.stdin
stdout = sys.stdout
return stdin.isatty() and stdout.isatty()
def _interactive_readline(prompt: str, timeout: int) -> Tuple[str, bool]:
if os.name == "nt": # Windows
return _interactive_readline_windows(prompt, timeout)
return _interactive_readline_posix(prompt, timeout)
def _interactive_readline_windows(prompt: str, timeout: int) -> Tuple[str, bool]:
try:
import msvcrt
except ImportError: # pragma: no cover - 非Windows环境
raise _InteractiveInputUnavailable from None
buffer: List[str] = []
timer_active = timeout > 0
deadline = time.monotonic() + timeout if timer_active else None
pause_remaining = float(timeout) if timer_active else 0.0
last_activity: Optional[float] = None
renderer = _CountdownRenderer(prompt, timeout)
renderer.render(buffer, timeout if timer_active else None)
while True:
now = time.monotonic()
if timer_active and deadline is not None:
remaining = max(deadline - now, 0.0)
if remaining <= 0:
renderer.render(buffer, 0.0)
renderer.finish()
return "", True
else:
remaining = pause_remaining
if timeout and pause_remaining <= 0:
renderer.render(buffer, 0.0)
renderer.finish()
return "", True
if (
timeout
and last_activity is not None
and (now - last_activity) >= _INACTIVITY_GRACE_SECONDS
):
timer_active = True
deadline = now + pause_remaining
last_activity = None
continue
renderer.render(buffer, remaining if timeout else None)
if msvcrt.kbhit():
ch = msvcrt.getwch()
if ch in ("\r", "\n"):
renderer.finish()
return "".join(buffer), False
if ch == "\003": # Ctrl+C
renderer.finish()
raise KeyboardInterrupt
if ch in ("\x00", "\xe0"): # 特殊功能键,忽略额外字节
msvcrt.getwch()
continue
if ch in ("\b", "\x7f"):
if buffer:
buffer.pop()
if timer_active:
pause_remaining = max(remaining, 0.0)
timer_active = False
deadline = None
last_activity = time.monotonic()
renderer.render(buffer, pause_remaining if not timer_active else remaining)
continue
buffer.append(ch)
if timer_active:
pause_remaining = max(remaining, 0.0)
timer_active = False
deadline = None
last_activity = time.monotonic()
renderer.render(buffer, pause_remaining if not timer_active else remaining)
continue
time.sleep(0.05)
def _interactive_readline_posix(prompt: str, timeout: int) -> Tuple[str, bool]:
import select
import termios
import tty
fd = sys.stdin.fileno()
if fd < 0:
raise _InteractiveInputUnavailable
buffer: List[str] = []
timer_active = timeout > 0
deadline = time.monotonic() + timeout if timer_active else None
pause_remaining = float(timeout) if timer_active else 0.0
last_activity: Optional[float] = None
renderer = _CountdownRenderer(prompt, timeout)
renderer.render(buffer, timeout if timer_active else None)
with _raw_mode(sys.stdin):
while True:
now = time.monotonic()
if timer_active and deadline is not None:
remaining = max(deadline - now, 0.0)
if remaining <= 0:
renderer.render(buffer, 0.0)
renderer.finish()
return "", True
wait_timeout = min(0.05, remaining)
else:
remaining = pause_remaining
wait_timeout = 0.05
if timeout and pause_remaining <= 0:
renderer.render(buffer, 0.0)
renderer.finish()
return "", True
if (
timeout
and last_activity is not None
and (now - last_activity) >= _INACTIVITY_GRACE_SECONDS
):
timer_active = True
deadline = now + pause_remaining
last_activity = None
continue
renderer.render(buffer, remaining if timeout else None)
ready, _, _ = select.select([sys.stdin], [], [], wait_timeout)
if not ready:
continue
ch = sys.stdin.read(1)
if ch == "":
raise _InteractiveInputUnavailable
if ch in ("\r", "\n"):
renderer.finish()
return "".join(buffer), False
if ch == "\x03": # Ctrl+C
renderer.finish()
raise KeyboardInterrupt
if ch in ("\x7f", "\b"):
if buffer:
buffer.pop()
if timer_active:
pause_remaining = max(remaining, 0.0)
timer_active = False
deadline = None
last_activity = time.monotonic()
renderer.render(buffer, pause_remaining if not timer_active else remaining)
continue
buffer.append(ch)
if timer_active:
pause_remaining = max(remaining, 0.0)
timer_active = False
deadline = None
last_activity = time.monotonic()
renderer.render(buffer, pause_remaining if not timer_active else remaining)
@contextmanager
def _raw_mode(file_obj: Any):
import termios
import tty
if not file_obj.isatty():
raise _InteractiveInputUnavailable
fd = file_obj.fileno()
old_attrs = termios.tcgetattr(fd)
try:
tty.setraw(fd)
yield
finally: # pragma: no branch - 确保恢复终端
termios.tcsetattr(fd, termios.TCSADRAIN, old_attrs)

View File

@ -0,0 +1,80 @@
"""timeout_input 实用函数的单元测试"""
import importlib.util
import sys
from pathlib import Path
import pytest
PROJECT_SRC = Path(__file__).resolve().parents[2] / "src"
MODULE_PATH = PROJECT_SRC / "bidmaster" / "utils" / "timeout_input.py"
spec = importlib.util.spec_from_file_location("timeout_input", MODULE_PATH)
timeout_input = importlib.util.module_from_spec(spec)
assert spec and spec.loader # 保持类型检查
spec.loader.exec_module(timeout_input)
@pytest.fixture()
def interactive_enabled(monkeypatch):
monkeypatch.setattr(timeout_input, "_can_use_interactive_input", lambda: True)
def test_timeout_prompt_returns_interactive_result(monkeypatch, interactive_enabled):
monkeypatch.setattr(
timeout_input,
"_interactive_readline",
lambda prompt, timeout: (" 用户提示 ", False),
)
result = timeout_input.timeout_prompt("输入: ", default="默认", timeout=60)
assert result == "用户提示"
def test_timeout_prompt_uses_default_on_timeout(monkeypatch, interactive_enabled):
monkeypatch.setattr(
timeout_input,
"_interactive_readline",
lambda prompt, timeout: ("", True),
)
result = timeout_input.timeout_prompt("输入: ", default="默认", timeout=60)
assert result == "默认"
def test_timeout_prompt_falls_back_to_threaded(monkeypatch, interactive_enabled):
def fake_threaded(prompt, default, timeout, choices):
assert prompt == "输入: "
assert default == "兜底"
assert timeout == 60
assert choices is None
return "线程输入"
monkeypatch.setattr(
timeout_input,
"_interactive_readline",
lambda prompt, timeout: (_ for _ in ()).throw(
timeout_input._InteractiveInputUnavailable()
),
)
monkeypatch.setattr(timeout_input, "_threaded_prompt", fake_threaded)
result = timeout_input.timeout_prompt("输入: ", default="兜底", timeout=60)
assert result == "线程输入"
def test_timeout_prompt_invalid_choice_returns_default(monkeypatch, interactive_enabled):
monkeypatch.setattr(
timeout_input,
"_interactive_readline",
lambda prompt, timeout: ("3", False),
)
result = timeout_input.timeout_prompt(
"选择:", default="1", timeout=60, choices=["1", "2"]
)
assert result == "1"