396 lines
13 KiB
Python
396 lines
13 KiB
Python
import asyncio
|
||
import logging
|
||
import shlex
|
||
import signal
|
||
import sys
|
||
from uuid import uuid4
|
||
|
||
from yuxi.channel.context import ChannelContext
|
||
from yuxi.channel.message.models import MessageType, PeerInfo, PeerKind, UnifiedMessage
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
RICH_AVAILABLE = False
|
||
_rich_console = None
|
||
_RICH_THEMES: dict = {}
|
||
|
||
try:
|
||
from rich.console import Console
|
||
from rich.markdown import Markdown
|
||
from rich.theme import Theme
|
||
|
||
_RICH_THEMES = {
|
||
"dark": None,
|
||
"light": Theme(
|
||
{
|
||
"markdown.code": "bright_black on white",
|
||
"markdown.heading": "bold blue",
|
||
}
|
||
),
|
||
"no-color": Theme(inherit=False),
|
||
}
|
||
|
||
def _init_rich_console(theme: str = "dark") -> None:
|
||
global _rich_console, RICH_AVAILABLE
|
||
try:
|
||
_rich_console = Console(theme=_RICH_THEMES.get(theme))
|
||
RICH_AVAILABLE = True
|
||
except Exception:
|
||
_rich_console = None
|
||
RICH_AVAILABLE = False
|
||
|
||
_init_rich_console("dark")
|
||
except ImportError:
|
||
pass
|
||
|
||
PROMPT_TOOLKIT_AVAILABLE = False
|
||
_prompt_session = None
|
||
_kb = None
|
||
try:
|
||
from prompt_toolkit import PromptSession
|
||
from prompt_toolkit.completion import Completer, Completion
|
||
from prompt_toolkit.history import InMemoryHistory
|
||
from prompt_toolkit.key_binding import KeyBindings
|
||
|
||
class _TerminalCompleter(Completer):
|
||
_COMMANDS = ["help", "exit", "quit", "clear", "history", "model"]
|
||
|
||
def get_completions(self, document, complete_event):
|
||
text = document.text_before_cursor.strip()
|
||
if text.startswith("/"):
|
||
word = text[1:].split()[0] if len(text) > 1 else ""
|
||
for cmd in self._COMMANDS:
|
||
if cmd.startswith(word):
|
||
yield Completion(
|
||
f"/{cmd}",
|
||
start_position=-len(text),
|
||
display_meta=f"/{cmd}",
|
||
)
|
||
|
||
_kb = KeyBindings()
|
||
|
||
@_kb.add("escape", "enter")
|
||
def _multiline_send(event):
|
||
event.current_buffer.validate_and_handle()
|
||
|
||
try:
|
||
_prompt_session = PromptSession(
|
||
history=InMemoryHistory(),
|
||
completer=_TerminalCompleter(),
|
||
)
|
||
PROMPT_TOOLKIT_AVAILABLE = True
|
||
except Exception:
|
||
logger.debug("prompt_toolkit PromptSession unavailable (no console), falling back to readline")
|
||
except ImportError:
|
||
pass
|
||
|
||
|
||
class TerminalAdapter:
|
||
def __init__(self):
|
||
self._msg_counter = 0
|
||
self._ctx: ChannelContext | None = None
|
||
self._original_sigint = None
|
||
self._typing_task: asyncio.Task | None = None
|
||
self._is_typing = False
|
||
self._history: list[str] = []
|
||
self._commands = {
|
||
"help": self._cmd_help,
|
||
"exit": self._cmd_exit,
|
||
"quit": self._cmd_exit,
|
||
"clear": self._cmd_clear,
|
||
"history": self._cmd_history,
|
||
"model": self._cmd_model,
|
||
}
|
||
|
||
def list_account_ids(self, config: dict) -> list[str]:
|
||
return ["default"]
|
||
|
||
async def resolve_account(self, account_id: str) -> dict:
|
||
return {
|
||
"account_id": "default",
|
||
"enabled": True,
|
||
"configured": True,
|
||
"name": "Terminal",
|
||
}
|
||
|
||
def is_configured(self, account: dict) -> bool:
|
||
return True
|
||
|
||
def is_enabled(self, account: dict) -> bool:
|
||
return True
|
||
|
||
async def start(self, ctx: ChannelContext) -> None:
|
||
logger.info("Terminal channel starting (stdin listener)")
|
||
self._ctx = ctx
|
||
|
||
if RICH_AVAILABLE:
|
||
account_config = ctx.config.get("accounts", {}).get("default", {})
|
||
theme = account_config.get("color_theme", "dark")
|
||
_init_rich_console(theme)
|
||
|
||
self._print_greeting()
|
||
|
||
def _sigint_handler(_signum, _frame):
|
||
sys.stdout.write("\n\n[中断请求,正在取消当前回复...]\n")
|
||
sys.stdout.flush()
|
||
ctx.cancel_event.set()
|
||
|
||
self._original_sigint = signal.signal(signal.SIGINT, _sigint_handler)
|
||
|
||
try:
|
||
while not ctx.cancel_event.is_set():
|
||
try:
|
||
line = await self._read_stdin_line()
|
||
except EOFError:
|
||
logger.info("Terminal stdin closed (EOF)")
|
||
sys.stdout.write("\nGoodbye!\n")
|
||
sys.stdout.flush()
|
||
break
|
||
|
||
if line is None:
|
||
break
|
||
|
||
text = line.strip()
|
||
if not text:
|
||
if PROMPT_TOOLKIT_AVAILABLE:
|
||
continue
|
||
self._print_prompt()
|
||
continue
|
||
|
||
if text.startswith("/") and self._handle_command(text):
|
||
if not PROMPT_TOOLKIT_AVAILABLE:
|
||
self._print_prompt()
|
||
continue
|
||
|
||
if text.startswith('"""'):
|
||
text = await self._read_multiline_input(text)
|
||
if text is None or not text.strip():
|
||
if not PROMPT_TOOLKIT_AVAILABLE:
|
||
self._print_prompt()
|
||
continue
|
||
|
||
self._enqueue_message(text)
|
||
|
||
finally:
|
||
if self._original_sigint is not None:
|
||
signal.signal(signal.SIGINT, self._original_sigint)
|
||
self._original_sigint = None
|
||
|
||
async def stop(self, ctx: ChannelContext) -> None:
|
||
logger.info("Terminal channel stopping")
|
||
ctx.cancel_event.set()
|
||
if self._original_sigint is not None:
|
||
signal.signal(signal.SIGINT, self._original_sigint)
|
||
self._original_sigint = None
|
||
await self.clear_typing("terminal-default")
|
||
sys.stdout.write("\n")
|
||
sys.stdout.flush()
|
||
|
||
async def send_text(self, target_id: str, content: str) -> None:
|
||
sys.stdout.write("\n")
|
||
if RICH_AVAILABLE:
|
||
_rich_console.print(Markdown(content))
|
||
else:
|
||
sys.stdout.write(content)
|
||
sys.stdout.write("\n")
|
||
sys.stdout.flush()
|
||
if not PROMPT_TOOLKIT_AVAILABLE:
|
||
self._print_prompt()
|
||
|
||
async def send_typing(self, target_id: str, thread_id: str | None = None) -> None:
|
||
if not sys.stdin.isatty():
|
||
return
|
||
self._is_typing = True
|
||
self._typing_task = asyncio.create_task(self._typing_animation())
|
||
|
||
async def clear_typing(self, target_id: str, thread_id: str | None = None) -> None:
|
||
self._is_typing = False
|
||
if self._typing_task and not self._typing_task.done():
|
||
self._typing_task.cancel()
|
||
try:
|
||
await self._typing_task
|
||
except asyncio.CancelledError:
|
||
pass
|
||
self._typing_task = None
|
||
sys.stdout.write("\r\033[K")
|
||
sys.stdout.flush()
|
||
|
||
def _enqueue_message(self, text: str) -> None:
|
||
self._history.append(text)
|
||
if len(self._history) > 100:
|
||
self._history.pop(0)
|
||
self._msg_counter += 1
|
||
msg_id = f"term:{uuid4().hex[:12]}"
|
||
|
||
msg = UnifiedMessage(
|
||
msg_id=msg_id,
|
||
channel_type="terminal",
|
||
account_id="default",
|
||
content=text,
|
||
message_type=MessageType.TEXT,
|
||
sender=PeerInfo(
|
||
kind=PeerKind.DIRECT,
|
||
id="terminal-user",
|
||
display_name="Dev",
|
||
),
|
||
timestamp=None,
|
||
)
|
||
|
||
asyncio.create_task(self._put_message(msg))
|
||
|
||
async def _put_message(self, msg: UnifiedMessage) -> None:
|
||
if self._ctx is None:
|
||
return
|
||
if self._ctx.queue.qsize() >= self._ctx.queue.maxsize * 0.8:
|
||
self._print_warning("消息队列繁忙,请稍候...")
|
||
await self._ctx.queue.put(msg)
|
||
|
||
def _handle_command(self, text: str) -> bool:
|
||
try:
|
||
parts = shlex.split(text)
|
||
except ValueError:
|
||
return False
|
||
if not parts or not parts[0].startswith("/"):
|
||
return False
|
||
|
||
cmd_name = parts[0][1:].lower()
|
||
handler = self._commands.get(cmd_name)
|
||
if handler and self._ctx:
|
||
handler(self._ctx)
|
||
return True
|
||
return False
|
||
|
||
@staticmethod
|
||
def _cmd_help(ctx: ChannelContext) -> None:
|
||
help_text = (
|
||
"\n可用命令:\n"
|
||
" /help 显示此帮助信息\n"
|
||
" /exit, /quit 退出 Terminal 渠道\n"
|
||
" /clear 清屏\n"
|
||
" /history 显示输入历史\n"
|
||
" /model <name> 切换模型\n"
|
||
"\n快捷键: Ctrl+C 中断回复 Ctrl+D 退出\n"
|
||
)
|
||
sys.stdout.write(f"{help_text}\n")
|
||
sys.stdout.flush()
|
||
|
||
@staticmethod
|
||
def _cmd_exit(ctx: ChannelContext) -> None:
|
||
sys.stdout.write("\nGoodbye!\n")
|
||
sys.stdout.flush()
|
||
ctx.cancel_event.set()
|
||
|
||
@staticmethod
|
||
def _cmd_clear(ctx: ChannelContext) -> None:
|
||
sys.stdout.write("\033[2J\033[H")
|
||
sys.stdout.flush()
|
||
|
||
@staticmethod
|
||
def _cmd_model(ctx: ChannelContext) -> None:
|
||
sys.stdout.write("\n模型切换请在系统中配置\n")
|
||
sys.stdout.flush()
|
||
|
||
def _cmd_history(self, ctx: ChannelContext) -> None:
|
||
if not self._history:
|
||
sys.stdout.write("\n暂无历史记录\n")
|
||
sys.stdout.flush()
|
||
return
|
||
lines = ["\n历史记录:"]
|
||
for i, entry in enumerate(self._history, 1):
|
||
preview = entry[:60] + "..." if len(entry) > 60 else entry
|
||
lines.append(f" {i}. {preview}")
|
||
lines.append("")
|
||
sys.stdout.write("\n".join(lines) + "\n")
|
||
sys.stdout.flush()
|
||
|
||
async def _read_stdin_line(self, multiline: bool = False) -> str | None:
|
||
loop = asyncio.get_event_loop()
|
||
|
||
if PROMPT_TOOLKIT_AVAILABLE:
|
||
try:
|
||
line = await loop.run_in_executor(
|
||
None,
|
||
lambda: _prompt_session.prompt(
|
||
"... " if multiline else "> ",
|
||
multiline=multiline,
|
||
key_bindings=_kb if multiline else None,
|
||
),
|
||
)
|
||
return line
|
||
except EOFError:
|
||
raise
|
||
except KeyboardInterrupt:
|
||
return ""
|
||
|
||
line = await loop.run_in_executor(None, sys.stdin.readline)
|
||
if not line:
|
||
return None
|
||
return line
|
||
|
||
async def _read_multiline_input(self, first_line: str) -> str | None:
|
||
remaining = first_line[3:]
|
||
|
||
if PROMPT_TOOLKIT_AVAILABLE:
|
||
try:
|
||
loop = asyncio.get_event_loop()
|
||
return await loop.run_in_executor(
|
||
None,
|
||
lambda: _prompt_session.prompt(
|
||
"... ",
|
||
multiline=True,
|
||
key_bindings=_kb,
|
||
default=remaining,
|
||
),
|
||
)
|
||
except (KeyboardInterrupt, EOFError):
|
||
return None
|
||
|
||
lines = []
|
||
if remaining.strip():
|
||
lines.append(remaining.strip())
|
||
while self._ctx and not self._ctx.cancel_event.is_set():
|
||
sys.stdout.write("... ")
|
||
sys.stdout.flush()
|
||
try:
|
||
next_line = await self._read_stdin_line()
|
||
except EOFError:
|
||
break
|
||
if next_line is None:
|
||
break
|
||
stripped = next_line.strip()
|
||
if stripped == '"""':
|
||
break
|
||
lines.append(next_line.rstrip("\n"))
|
||
return "\n".join(lines) if lines else None
|
||
|
||
async def _typing_animation(self) -> None:
|
||
frames = ["⠋", "⠙", "⠹", "⠸", "⠼", "⠴", "⠦", "⠧", "⠇", "⠏"]
|
||
idx = 0
|
||
while self._is_typing:
|
||
frame = frames[idx % len(frames)]
|
||
if RICH_AVAILABLE:
|
||
sys.stdout.write(f"\r[bold yellow]{frame}[/bold yellow] 思考中...")
|
||
else:
|
||
sys.stdout.write(f"\r{frame} 思考中...")
|
||
sys.stdout.flush()
|
||
idx += 1
|
||
await asyncio.sleep(0.1)
|
||
|
||
@staticmethod
|
||
def _print_greeting() -> None:
|
||
sys.stdout.write("Terminal 渠道已启动,输入问题与 Agent 对话")
|
||
sys.stdout.write("\n按 Ctrl+C 中断回复,Ctrl+D 退出\n")
|
||
if not PROMPT_TOOLKIT_AVAILABLE:
|
||
sys.stdout.write("\n> ")
|
||
sys.stdout.flush()
|
||
|
||
@staticmethod
|
||
def _print_prompt() -> None:
|
||
sys.stdout.write("> ")
|
||
sys.stdout.flush()
|
||
|
||
@staticmethod
|
||
def _print_warning(text: str) -> None:
|
||
sys.stdout.write(f"\n[{text}]\n")
|
||
sys.stdout.flush()
|