zrcoder 0.3.0__tar.gz → 0.4.0__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- {zrcoder-0.3.0 → zrcoder-0.4.0}/PKG-INFO +1 -1
- {zrcoder-0.3.0 → zrcoder-0.4.0}/core/agent/core.py +2 -2
- {zrcoder-0.3.0 → zrcoder-0.4.0}/core/agent/hooks.py +33 -1
- {zrcoder-0.3.0 → zrcoder-0.4.0}/core/agent/tools.py +30 -0
- {zrcoder-0.3.0 → zrcoder-0.4.0}/core/cli.py +37 -45
- zrcoder-0.4.0/core/permissions.py +234 -0
- {zrcoder-0.3.0 → zrcoder-0.4.0}/core/ui/commands.py +20 -11
- zrcoder-0.4.0/core/ui/input.py +202 -0
- {zrcoder-0.3.0 → zrcoder-0.4.0}/pyproject.toml +1 -1
- {zrcoder-0.3.0 → zrcoder-0.4.0}/pyproject.toml.orig +1 -1
- zrcoder-0.3.0/core/agent/openai_compat.py +0 -44
- {zrcoder-0.3.0 → zrcoder-0.4.0}/LICENSE +0 -0
- {zrcoder-0.3.0 → zrcoder-0.4.0}/README.md +0 -0
- {zrcoder-0.3.0 → zrcoder-0.4.0}/core/__init__.py +0 -0
- {zrcoder-0.3.0 → zrcoder-0.4.0}/core/agent/__init__.py +0 -0
- {zrcoder-0.3.0 → zrcoder-0.4.0}/core/session.py +0 -0
- {zrcoder-0.3.0 → zrcoder-0.4.0}/core/ui/__init__.py +0 -0
- {zrcoder-0.3.0 → zrcoder-0.4.0}/core/ui/render.py +0 -0
|
@@ -8,10 +8,10 @@ from urllib.parse import urlparse, urlunparse
|
|
|
8
8
|
|
|
9
9
|
from dotenv import load_dotenv
|
|
10
10
|
from pydantic_ai import Agent
|
|
11
|
+
from pydantic_ai.models.openai import OpenAIChatModel
|
|
11
12
|
from pydantic_ai.providers.openai import OpenAIProvider
|
|
12
13
|
|
|
13
14
|
from .hooks import hooks
|
|
14
|
-
from .openai_compat import CompatibleOpenAIChatModel
|
|
15
15
|
from .tools import TOOLS
|
|
16
16
|
|
|
17
17
|
load_dotenv()
|
|
@@ -41,7 +41,7 @@ def _normalize_openai_base_url(url: str) -> str:
|
|
|
41
41
|
|
|
42
42
|
MODEL_NAME = os.environ.get("MODEL_NAME", "deepseek-v4-flash")
|
|
43
43
|
|
|
44
|
-
model =
|
|
44
|
+
model = OpenAIChatModel(
|
|
45
45
|
MODEL_NAME,
|
|
46
46
|
provider=OpenAIProvider(base_url=_normalize_openai_base_url(BASE_URL), api_key=API_KEY),
|
|
47
47
|
)
|
|
@@ -2,7 +2,8 @@
|
|
|
2
2
|
挂在 Agent 上的 hooks:
|
|
3
3
|
1. API 调用元数据记录(/api-detail 命令用)
|
|
4
4
|
2. API 请求失败时的自动重试(wrap_model_request)
|
|
5
|
-
3.
|
|
5
|
+
3. 工具调用权限检查(on_tool_execute)
|
|
6
|
+
4. 工具执行异常的兜底处理(on_tool_execute_error)
|
|
6
7
|
"""
|
|
7
8
|
|
|
8
9
|
from __future__ import annotations
|
|
@@ -14,6 +15,7 @@ from typing import Any
|
|
|
14
15
|
from pydantic_ai.capabilities import Hooks
|
|
15
16
|
from pydantic_ai.exceptions import ModelAPIError, ModelHTTPError, UnexpectedModelBehavior
|
|
16
17
|
|
|
18
|
+
from core import permissions
|
|
17
19
|
from core.ui.render import console
|
|
18
20
|
|
|
19
21
|
MAX_RETRIES = 3
|
|
@@ -132,6 +134,36 @@ async def _retry_on_error(ctx: Any, *, request_context: Any, handler: Any) -> An
|
|
|
132
134
|
raise RuntimeError("unreachable") # pragma: no cover
|
|
133
135
|
|
|
134
136
|
|
|
137
|
+
# ---------- 工具调用权限检查 ----------
|
|
138
|
+
|
|
139
|
+
|
|
140
|
+
@hooks.on.tool_execute
|
|
141
|
+
async def _check_permission(ctx: Any, *, call: Any, tool_def: Any, args: Any, handler: Any) -> Any:
|
|
142
|
+
"""
|
|
143
|
+
工具执行前的权限关卡。allow 就调用 handler 真正执行;
|
|
144
|
+
ask 就弹审批列表;deny 则把拒绝原因当作工具结果回填,让模型自行纠正。
|
|
145
|
+
"""
|
|
146
|
+
decision = permissions.compute_decision(call.tool_name, args)
|
|
147
|
+
if decision == "allow":
|
|
148
|
+
# 放行,handler(args) 才是真正执行工具的那一步
|
|
149
|
+
return await handler(args)
|
|
150
|
+
|
|
151
|
+
# decision == "ask",弹审批让用户决定
|
|
152
|
+
choice = await permissions.prompt_approval(call.tool_name, args)
|
|
153
|
+
if choice == "once":
|
|
154
|
+
return await handler(args)
|
|
155
|
+
if choice == "always":
|
|
156
|
+
# 记进会话白名单,本会话内这个工具不再询问
|
|
157
|
+
permissions.state.session_allowed.add(call.tool_name)
|
|
158
|
+
return await handler(args)
|
|
159
|
+
|
|
160
|
+
# 拒绝:不执行工具,把拒绝原因回填给模型,让它停下来等用户发话,而不是自作主张绕过去
|
|
161
|
+
return (
|
|
162
|
+
f"用户拒绝了对 {call.tool_name} 的调用,这次调用没有执行。"
|
|
163
|
+
"请停下手上的事,等用户告诉你接下来该怎么做。"
|
|
164
|
+
)
|
|
165
|
+
|
|
166
|
+
|
|
135
167
|
# ---------- 工具执行异常兜底 ----------
|
|
136
168
|
|
|
137
169
|
|
|
@@ -6,10 +6,16 @@ Coding Agent 用到的三个工具:读文件、写文件、跑 shell 命令。
|
|
|
6
6
|
意料之外的异常由 hooks 里的 on_tool_execute_error 统一兜底。
|
|
7
7
|
"""
|
|
8
8
|
|
|
9
|
+
from __future__ import annotations
|
|
10
|
+
|
|
11
|
+
import re
|
|
9
12
|
import subprocess
|
|
13
|
+
from typing import Any
|
|
10
14
|
|
|
11
15
|
from pydantic_ai.exceptions import ModelRetry
|
|
12
16
|
|
|
17
|
+
from core import permissions
|
|
18
|
+
|
|
13
19
|
|
|
14
20
|
def read_file(path: str) -> str:
|
|
15
21
|
"""
|
|
@@ -62,5 +68,29 @@ def run_command(command: str) -> str:
|
|
|
62
68
|
return f"[错误] 无法执行命令 ({e})"
|
|
63
69
|
|
|
64
70
|
|
|
71
|
+
# 高危命令的特征:删除文件、提权、直写磁盘
|
|
72
|
+
DANGEROUS_PATTERNS = [
|
|
73
|
+
r"\brm\b",
|
|
74
|
+
r"\bsudo\b",
|
|
75
|
+
r"\bdd\b",
|
|
76
|
+
r"\bmkfs\w*\b",
|
|
77
|
+
]
|
|
78
|
+
|
|
79
|
+
|
|
80
|
+
def run_command_self_check(args: dict[str, Any]) -> str | None:
|
|
81
|
+
"""
|
|
82
|
+
run_command 的权限自检:扫一遍命令字符串,命中高危特征就要求审批。
|
|
83
|
+
误伤的代价不过是多弹一次窗,绝不能把真正的高危命令漏过去。
|
|
84
|
+
"""
|
|
85
|
+
command = str(args.get("command", ""))
|
|
86
|
+
if any(re.search(pattern, command) for pattern in DANGEROUS_PATTERNS):
|
|
87
|
+
return "ask"
|
|
88
|
+
# 没命中高危特征,交给通用规则决定
|
|
89
|
+
return None
|
|
90
|
+
|
|
91
|
+
|
|
92
|
+
# 把自检挂到权限模块的注册表上
|
|
93
|
+
permissions.register_self_check("run_command", run_command_self_check)
|
|
94
|
+
|
|
65
95
|
# Pydantic AI 支持 tools=[plain_function],从函数签名 + docstring 自动生成 JSON Schema
|
|
66
96
|
TOOLS = [read_file, write_file, run_command]
|
|
@@ -5,44 +5,27 @@
|
|
|
5
5
|
from __future__ import annotations
|
|
6
6
|
|
|
7
7
|
import asyncio
|
|
8
|
+
import inspect
|
|
8
9
|
from typing import Any, Literal
|
|
9
10
|
|
|
10
|
-
from prompt_toolkit import PromptSession
|
|
11
11
|
from pydantic_ai import Agent
|
|
12
12
|
from pydantic_graph import End
|
|
13
13
|
|
|
14
14
|
from core.agent import MODEL_NAME, agent, api_call_log
|
|
15
15
|
from core.session import append_messages, new_session_id
|
|
16
|
-
from core.ui.commands import COMMANDS, SessionState,
|
|
16
|
+
from core.ui.commands import COMMANDS, SessionState, print_part
|
|
17
|
+
from core.ui.input import Repl
|
|
17
18
|
from core.ui.render import console, print_welcome_banner
|
|
18
19
|
|
|
19
|
-
# PromptSession 比内置 input() 好用:支持左右移光标编辑,并记住本次输入历史
|
|
20
|
-
prompt_session: PromptSession[str] = PromptSession()
|
|
21
|
-
|
|
22
20
|
CommandAction = Literal["pass", "continue", "break"]
|
|
23
21
|
|
|
24
22
|
|
|
25
|
-
def
|
|
26
|
-
"""
|
|
27
|
-
打印上横线并读一行用户输入;回车后再补一条下横线。
|
|
28
|
-
返回 None 表示用户希望退出(Ctrl-C / Ctrl-D)。
|
|
29
|
-
"""
|
|
30
|
-
print_divider()
|
|
31
|
-
try:
|
|
32
|
-
user_input = prompt_session.prompt("❯ ").strip()
|
|
33
|
-
except (EOFError, KeyboardInterrupt):
|
|
34
|
-
print()
|
|
35
|
-
return None
|
|
36
|
-
print_divider()
|
|
37
|
-
return user_input
|
|
38
|
-
|
|
39
|
-
|
|
40
|
-
def handle_command(user_input: str, state: SessionState) -> CommandAction:
|
|
23
|
+
async def handle_command(user_input: str, state: SessionState) -> CommandAction:
|
|
41
24
|
"""
|
|
42
25
|
处理以 / 开头的命令。
|
|
43
|
-
返回 'pass'
|
|
44
|
-
返回 'continue'
|
|
45
|
-
返回 'break'
|
|
26
|
+
返回 'pass':不是命令,交给 Agent;
|
|
27
|
+
返回 'continue':命令已处理,进入下一轮;
|
|
28
|
+
返回 'break':命令要求退出程序。
|
|
46
29
|
"""
|
|
47
30
|
if not user_input.startswith("/"):
|
|
48
31
|
return "pass"
|
|
@@ -51,7 +34,11 @@ def handle_command(user_input: str, state: SessionState) -> CommandAction:
|
|
|
51
34
|
if command is None:
|
|
52
35
|
console.print(f"未知命令:/{cmd_name},输入 /help 查看可用命令\n")
|
|
53
36
|
return "continue"
|
|
54
|
-
|
|
37
|
+
result = command.handler(state)
|
|
38
|
+
# 个别命令(如 /resume)要弹交互式列表,是异步的,需要 await
|
|
39
|
+
if inspect.isawaitable(result):
|
|
40
|
+
result = await result
|
|
41
|
+
return "continue" if result else "break"
|
|
55
42
|
|
|
56
43
|
|
|
57
44
|
def apply_result(state: SessionState, result: Any) -> None:
|
|
@@ -92,35 +79,40 @@ async def run_agent_loop(user_input: str, state: SessionState) -> None:
|
|
|
92
79
|
console.print()
|
|
93
80
|
|
|
94
81
|
|
|
95
|
-
def
|
|
82
|
+
async def async_main() -> None:
|
|
96
83
|
state = SessionState(
|
|
97
84
|
model_name=MODEL_NAME,
|
|
98
85
|
session_id=new_session_id(),
|
|
99
86
|
)
|
|
100
87
|
print_welcome_banner("Zrcoder")
|
|
101
88
|
|
|
102
|
-
|
|
103
|
-
|
|
104
|
-
user_input = read_user_input()
|
|
105
|
-
if user_input is None:
|
|
106
|
-
break
|
|
107
|
-
if not user_input:
|
|
108
|
-
continue
|
|
89
|
+
# 常驻输入区:输入框整个会话期间不消失
|
|
90
|
+
repl = Repl(state)
|
|
109
91
|
|
|
110
|
-
|
|
111
|
-
|
|
92
|
+
async def on_submit(user_input: str) -> None:
|
|
93
|
+
# 每次回车提交一行输入,都走这里
|
|
94
|
+
# 先处理 / 开头的命令
|
|
95
|
+
action = await handle_command(user_input, state)
|
|
112
96
|
if action == "break":
|
|
113
|
-
|
|
97
|
+
# 命令要求退出,结束常驻输入区
|
|
98
|
+
repl.exit()
|
|
99
|
+
return
|
|
114
100
|
if action == "continue":
|
|
115
|
-
|
|
116
|
-
|
|
117
|
-
# 核心 Agent
|
|
118
|
-
|
|
119
|
-
|
|
120
|
-
|
|
121
|
-
|
|
122
|
-
|
|
123
|
-
|
|
101
|
+
return
|
|
102
|
+
|
|
103
|
+
# 核心 Agent 循环:开请求时显示 working...,结束 / 被打断后由 Repl 统一隐藏;
|
|
104
|
+
# 中途按 ESC / Ctrl+C 会打断
|
|
105
|
+
repl.start_working()
|
|
106
|
+
await run_agent_loop(user_input, state)
|
|
107
|
+
|
|
108
|
+
await repl.run(on_submit)
|
|
109
|
+
|
|
110
|
+
|
|
111
|
+
def main() -> None:
|
|
112
|
+
try:
|
|
113
|
+
asyncio.run(async_main())
|
|
114
|
+
except (KeyboardInterrupt, EOFError):
|
|
115
|
+
pass
|
|
124
116
|
|
|
125
117
|
|
|
126
118
|
if __name__ == "__main__":
|
|
@@ -0,0 +1,234 @@
|
|
|
1
|
+
"""
|
|
2
|
+
权限管控:工具执行前,根据当前权限模式决定放行、询问还是拒绝。
|
|
3
|
+
"""
|
|
4
|
+
|
|
5
|
+
from __future__ import annotations
|
|
6
|
+
|
|
7
|
+
from collections.abc import Callable
|
|
8
|
+
from dataclasses import dataclass, field
|
|
9
|
+
from typing import Any
|
|
10
|
+
|
|
11
|
+
from prompt_toolkit.application import Application, in_terminal
|
|
12
|
+
from prompt_toolkit.formatted_text import FormattedText
|
|
13
|
+
from prompt_toolkit.key_binding import KeyBindings
|
|
14
|
+
from prompt_toolkit.key_binding.key_processor import KeyPressEvent
|
|
15
|
+
from prompt_toolkit.layout import HSplit, Layout, Window
|
|
16
|
+
from prompt_toolkit.layout.controls import FormattedTextControl
|
|
17
|
+
from prompt_toolkit.styles import Style
|
|
18
|
+
|
|
19
|
+
# 三种权限模式
|
|
20
|
+
# default:写文件、跑命令要授权,读文件等只读操作自动放行
|
|
21
|
+
DEFAULT = "default"
|
|
22
|
+
# acceptEdits:读写文件自动放行,其他工具仍需授权
|
|
23
|
+
ACCEPT_EDITS = "acceptEdits"
|
|
24
|
+
# bypass:一切放行,不再询问
|
|
25
|
+
BYPASS = "bypass"
|
|
26
|
+
|
|
27
|
+
# 权限模式按 Shift+Tab 循环切换的顺序
|
|
28
|
+
MODES = [DEFAULT, ACCEPT_EDITS, BYPASS]
|
|
29
|
+
|
|
30
|
+
# 只读工具,任何模式都自动放行(读取不会改动系统,放行没风险)
|
|
31
|
+
READONLY_TOOLS = {"read_file"}
|
|
32
|
+
# 编辑文件类工具,acceptEdits 模式下自动放行
|
|
33
|
+
EDIT_TOOLS = {"write_file"}
|
|
34
|
+
|
|
35
|
+
# 工具自检:接收 args,返回 "ask" 表示要求审批,返回 None 表示交给通用规则
|
|
36
|
+
SelfCheck = Callable[[dict[str, Any]], str | None]
|
|
37
|
+
|
|
38
|
+
# 工具自检注册表:通用权限规则只认工具名,但危不危险往往取决于参数,只有工具自己最懂参数的语义
|
|
39
|
+
TOOL_SELF_CHECKS: dict[str, SelfCheck] = {}
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
def register_self_check(tool_name: str, check: SelfCheck) -> None:
|
|
43
|
+
"""
|
|
44
|
+
工具模块调用它挂上自己的自检函数;
|
|
45
|
+
自检接收 args,返回 "ask" 表示要求审批,返回 None 表示交给通用规则。
|
|
46
|
+
"""
|
|
47
|
+
TOOL_SELF_CHECKS[tool_name] = check
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
@dataclass
|
|
51
|
+
class PermissionState:
|
|
52
|
+
"""
|
|
53
|
+
进程级权限状态:当前模式,加上本会话的放行白名单。
|
|
54
|
+
权限模式是进程级的:跨 /new、/resume 保持,不写进 jsonl,重启程序才回到 default。
|
|
55
|
+
白名单是会话级的:用 /new、/resume 切换会话时会清空。
|
|
56
|
+
"""
|
|
57
|
+
|
|
58
|
+
# 当前权限模式
|
|
59
|
+
mode: str = DEFAULT
|
|
60
|
+
# 本会话内点过「不再询问」的工具名,后续直接放行
|
|
61
|
+
session_allowed: set[str] = field(default_factory=set)
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
state = PermissionState()
|
|
65
|
+
|
|
66
|
+
|
|
67
|
+
def compute_decision(tool_name: str, args: dict[str, Any]) -> str:
|
|
68
|
+
"""
|
|
69
|
+
纯规则判定:在当前模式下,对给定的工具调用返回 "allow" 或 "ask"。
|
|
70
|
+
"""
|
|
71
|
+
# bypass 模式:全部放行,连工具自检都不再过问(用户主动选择了这个模式,后果自负)
|
|
72
|
+
if state.mode == BYPASS:
|
|
73
|
+
return "allow"
|
|
74
|
+
# 工具自检:自检要求审批的调用,会话白名单也盖不过
|
|
75
|
+
check = TOOL_SELF_CHECKS.get(tool_name)
|
|
76
|
+
if check and check(args) == "ask":
|
|
77
|
+
return "ask"
|
|
78
|
+
# 本会话点过「不再询问」的工具:放行
|
|
79
|
+
if tool_name in state.session_allowed:
|
|
80
|
+
return "allow"
|
|
81
|
+
# 只读工具:任何模式都自动放行
|
|
82
|
+
if tool_name in READONLY_TOOLS:
|
|
83
|
+
return "allow"
|
|
84
|
+
# acceptEdits 模式:编辑文件放行,命令等其他工具仍要审批
|
|
85
|
+
if state.mode == ACCEPT_EDITS and tool_name in EDIT_TOOLS:
|
|
86
|
+
return "allow"
|
|
87
|
+
# default 模式,或没命中任何放行规则:询问用户
|
|
88
|
+
return "ask"
|
|
89
|
+
|
|
90
|
+
|
|
91
|
+
def cycle_mode() -> str:
|
|
92
|
+
"""
|
|
93
|
+
按 default -> acceptEdits -> bypass -> default 循环切换。
|
|
94
|
+
"""
|
|
95
|
+
index = MODES.index(state.mode)
|
|
96
|
+
state.mode = MODES[(index + 1) % len(MODES)]
|
|
97
|
+
return state.mode
|
|
98
|
+
|
|
99
|
+
|
|
100
|
+
# 可能装着整份文件内容的参数,预览里只留开头;命令、路径等参数完整展示
|
|
101
|
+
BULKY_ARGS = {"content", "old_string", "new_string"}
|
|
102
|
+
|
|
103
|
+
|
|
104
|
+
def _format_call(tool_name: str, args: dict[str, Any]) -> str:
|
|
105
|
+
"""
|
|
106
|
+
把一次工具调用渲染成审批预览,例如 run_command(command=npm test)。
|
|
107
|
+
"""
|
|
108
|
+
|
|
109
|
+
def show(key: str, value: Any) -> str:
|
|
110
|
+
text = " ".join(str(value).split())
|
|
111
|
+
if key in BULKY_ARGS and len(text) > 60:
|
|
112
|
+
return text[:60] + "..."
|
|
113
|
+
return text
|
|
114
|
+
|
|
115
|
+
inner = ", ".join(f"{key}={show(key, value)}" for key, value in args.items())
|
|
116
|
+
return f"{tool_name}({inner})"
|
|
117
|
+
|
|
118
|
+
|
|
119
|
+
_STYLE = Style.from_dict(
|
|
120
|
+
{
|
|
121
|
+
"question": "bold",
|
|
122
|
+
"label-current": "#3b82f6 bold",
|
|
123
|
+
"label": "",
|
|
124
|
+
"footer": "#6b7280",
|
|
125
|
+
}
|
|
126
|
+
)
|
|
127
|
+
|
|
128
|
+
|
|
129
|
+
class _ApprovalPicker:
|
|
130
|
+
"""
|
|
131
|
+
手绘单选审批 picker,视觉对齐 ask_user_question 的 picker。
|
|
132
|
+
"""
|
|
133
|
+
|
|
134
|
+
def __init__(self, question: str, options: list[tuple[str, str]]) -> None:
|
|
135
|
+
# options 是 (value, label) 列表
|
|
136
|
+
self.question = question
|
|
137
|
+
self.options = options
|
|
138
|
+
self.cursor = 0
|
|
139
|
+
# 选中项的 value;取消时保持 None,run() 据此返回 deny
|
|
140
|
+
self.result: str | None = None
|
|
141
|
+
self.app = self._build_app()
|
|
142
|
+
|
|
143
|
+
def _render_question(self) -> FormattedText:
|
|
144
|
+
return FormattedText([("class:question", self.question)])
|
|
145
|
+
|
|
146
|
+
def _render_options(self) -> FormattedText:
|
|
147
|
+
lines: list[tuple[str, str]] = []
|
|
148
|
+
for i, (_value, label) in enumerate(self.options):
|
|
149
|
+
is_cursor = i == self.cursor
|
|
150
|
+
pointer = "❯" if is_cursor else " "
|
|
151
|
+
cls_label = "class:label-current" if is_cursor else "class:label"
|
|
152
|
+
lines.append((cls_label, f" {pointer} {i + 1}. {label}"))
|
|
153
|
+
lines.append(("", "\n"))
|
|
154
|
+
return FormattedText(lines)
|
|
155
|
+
|
|
156
|
+
def _render_footer(self) -> FormattedText:
|
|
157
|
+
return FormattedText([("class:footer", " ↑↓ 选择 · Enter 确认 · Esc 取消")])
|
|
158
|
+
|
|
159
|
+
def _move(self, delta: int) -> None:
|
|
160
|
+
self.cursor = (self.cursor + delta) % len(self.options)
|
|
161
|
+
|
|
162
|
+
def _build_app(self) -> Application[None]:
|
|
163
|
+
kb = KeyBindings()
|
|
164
|
+
|
|
165
|
+
# 堆叠装饰器让一个回调绑定多个键(kb.add 多参数是「按键序列」而非「任选其一」)
|
|
166
|
+
@kb.add("up")
|
|
167
|
+
@kb.add("k")
|
|
168
|
+
def _(event: KeyPressEvent) -> None:
|
|
169
|
+
self._move(-1)
|
|
170
|
+
|
|
171
|
+
@kb.add("down")
|
|
172
|
+
@kb.add("j")
|
|
173
|
+
def _(event: KeyPressEvent) -> None:
|
|
174
|
+
self._move(1)
|
|
175
|
+
|
|
176
|
+
@kb.add("enter")
|
|
177
|
+
def _(event: KeyPressEvent) -> None:
|
|
178
|
+
self.result = self.options[self.cursor][0]
|
|
179
|
+
self.app.exit()
|
|
180
|
+
|
|
181
|
+
@kb.add("escape")
|
|
182
|
+
@kb.add("c-c")
|
|
183
|
+
def _(event: KeyPressEvent) -> None:
|
|
184
|
+
self.app.exit()
|
|
185
|
+
|
|
186
|
+
# always_hide_cursor 藏掉终端光标,否则会落在左上角压住问句首字
|
|
187
|
+
layout = Layout(
|
|
188
|
+
HSplit(
|
|
189
|
+
[
|
|
190
|
+
Window(
|
|
191
|
+
FormattedTextControl(self._render_question),
|
|
192
|
+
wrap_lines=True,
|
|
193
|
+
dont_extend_height=True,
|
|
194
|
+
always_hide_cursor=True,
|
|
195
|
+
),
|
|
196
|
+
Window(
|
|
197
|
+
FormattedTextControl(self._render_options),
|
|
198
|
+
dont_extend_height=True,
|
|
199
|
+
always_hide_cursor=True,
|
|
200
|
+
),
|
|
201
|
+
Window(
|
|
202
|
+
FormattedTextControl(self._render_footer), height=1, always_hide_cursor=True
|
|
203
|
+
),
|
|
204
|
+
]
|
|
205
|
+
)
|
|
206
|
+
)
|
|
207
|
+
return Application(
|
|
208
|
+
layout=layout,
|
|
209
|
+
key_bindings=kb,
|
|
210
|
+
style=_STYLE,
|
|
211
|
+
full_screen=False,
|
|
212
|
+
mouse_support=False,
|
|
213
|
+
# 选完擦掉整个 picker,滚动区里不留下选项
|
|
214
|
+
erase_when_done=True,
|
|
215
|
+
)
|
|
216
|
+
|
|
217
|
+
async def run(self) -> str:
|
|
218
|
+
await self.app.run_async()
|
|
219
|
+
return self.result if self.result is not None else "deny"
|
|
220
|
+
|
|
221
|
+
|
|
222
|
+
async def prompt_approval(tool_name: str, args: dict[str, Any]) -> str:
|
|
223
|
+
"""
|
|
224
|
+
工具执行前弹出审批 picker,返回 "once" / "always" / "deny"。
|
|
225
|
+
"""
|
|
226
|
+
options = [
|
|
227
|
+
("once", "允许"),
|
|
228
|
+
("always", f"允许,且本会话不再询问 {tool_name}"),
|
|
229
|
+
("deny", "拒绝"),
|
|
230
|
+
]
|
|
231
|
+
async with in_terminal():
|
|
232
|
+
question = f"是否允许执行 {_format_call(tool_name, args)}?"
|
|
233
|
+
choice = await _ApprovalPicker(question, options).run()
|
|
234
|
+
return choice
|
|
@@ -1,16 +1,19 @@
|
|
|
1
|
-
from
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from collections.abc import Awaitable, Callable, Iterable
|
|
2
4
|
from dataclasses import dataclass, field
|
|
3
5
|
from datetime import datetime
|
|
4
6
|
from typing import Any
|
|
5
7
|
|
|
6
8
|
import questionary
|
|
9
|
+
from prompt_toolkit.application import in_terminal
|
|
7
10
|
from rich.console import Console, ConsoleOptions, RenderResult
|
|
8
11
|
from rich.markdown import Heading, Markdown
|
|
9
12
|
from rich.markup import escape
|
|
10
13
|
from rich.padding import Padding
|
|
11
14
|
from rich.rule import Rule
|
|
12
15
|
|
|
13
|
-
from core import session
|
|
16
|
+
from core import permissions, session
|
|
14
17
|
|
|
15
18
|
from .render import console, print_step
|
|
16
19
|
|
|
@@ -50,8 +53,8 @@ class SessionState:
|
|
|
50
53
|
class Command:
|
|
51
54
|
name: str
|
|
52
55
|
description: str
|
|
53
|
-
# handler 返回 False
|
|
54
|
-
handler: Callable[[SessionState], bool]
|
|
56
|
+
# handler 返回 False 表示主循环应当退出;可同步也可异步(如 /resume)
|
|
57
|
+
handler: Callable[[SessionState], bool | Awaitable[bool]]
|
|
55
58
|
|
|
56
59
|
|
|
57
60
|
def print_divider() -> None:
|
|
@@ -163,13 +166,15 @@ def cmd_help(state: SessionState) -> bool:
|
|
|
163
166
|
|
|
164
167
|
def cmd_new(state: SessionState) -> bool:
|
|
165
168
|
"""
|
|
166
|
-
开启新会话:清空历史、token 计数、API
|
|
169
|
+
开启新会话:清空历史、token 计数、API 调用记录,换一个新的会话 ID。
|
|
167
170
|
"""
|
|
168
171
|
state.history.clear()
|
|
169
172
|
state.input_tokens = 0
|
|
170
173
|
state.output_tokens = 0
|
|
171
174
|
state.last_api_calls.clear()
|
|
172
175
|
state.session_id = session.new_session_id()
|
|
176
|
+
# 权限白名单是会话级的,「本会话不再询问」不该带进新会话
|
|
177
|
+
permissions.state.session_allowed.clear()
|
|
173
178
|
console.print("已开启新会话\n")
|
|
174
179
|
return True
|
|
175
180
|
|
|
@@ -184,9 +189,10 @@ def _summary_line(mtime: datetime, prompt: str) -> str:
|
|
|
184
189
|
return f"{mtime:%m-%d %H:%M} {prompt}"
|
|
185
190
|
|
|
186
191
|
|
|
187
|
-
def cmd_resume(state: SessionState) -> bool:
|
|
192
|
+
async def cmd_resume(state: SessionState) -> bool:
|
|
188
193
|
"""
|
|
189
194
|
列出当前项目的历史会话,选中后恢复对话历史。
|
|
195
|
+
它跑在 REPL 的事件循环里,所以是异步的:in_terminal 把终端让给 questionary,结束后再恢复输入框。
|
|
190
196
|
"""
|
|
191
197
|
sessions = session.list_sessions()
|
|
192
198
|
if not sessions:
|
|
@@ -197,10 +203,10 @@ def cmd_resume(state: SessionState) -> bool:
|
|
|
197
203
|
questionary.Choice(title=_summary_line(mtime, prompt), value=sid)
|
|
198
204
|
for sid, mtime, prompt in sessions
|
|
199
205
|
]
|
|
200
|
-
|
|
201
|
-
|
|
202
|
-
|
|
203
|
-
|
|
206
|
+
async with in_terminal():
|
|
207
|
+
selected = await questionary.select(
|
|
208
|
+
"选择要恢复的会话(上下键移动,回车确认):", choices=choices
|
|
209
|
+
).ask_async()
|
|
204
210
|
# 用户按 Ctrl+C 取消选择
|
|
205
211
|
if selected is None:
|
|
206
212
|
return True
|
|
@@ -208,6 +214,8 @@ def cmd_resume(state: SessionState) -> bool:
|
|
|
208
214
|
# 还原对话历史,并把会话 ID 切换成选中的旧会话,后续消息继续追加到同一个文件
|
|
209
215
|
state.history = session.load_history(selected)
|
|
210
216
|
state.session_id = selected
|
|
217
|
+
# 权限白名单是会话级的,切换会话后清空
|
|
218
|
+
permissions.state.session_allowed.clear()
|
|
211
219
|
|
|
212
220
|
# jsonl 里每条模型回复都带 usage,把会话的 token 用量累加回来
|
|
213
221
|
state.input_tokens = sum(
|
|
@@ -223,6 +231,7 @@ def cmd_resume(state: SessionState) -> bool:
|
|
|
223
231
|
console.print(f"\n已恢复会话 {selected[:8]},共 {len(state.history)} 条消息:\n")
|
|
224
232
|
for msg in state.history:
|
|
225
233
|
for part in msg.parts:
|
|
234
|
+
# 回放和实时输出共用同一套 part 渲染逻辑
|
|
226
235
|
print_part(part)
|
|
227
236
|
console.print()
|
|
228
237
|
return True
|
|
@@ -230,7 +239,7 @@ def cmd_resume(state: SessionState) -> bool:
|
|
|
230
239
|
|
|
231
240
|
def cmd_status(state: SessionState) -> bool:
|
|
232
241
|
console.print(f"模型: {state.model_name}")
|
|
233
|
-
console.print(f"
|
|
242
|
+
console.print(f"权限模式: {permissions.state.mode}")
|
|
234
243
|
console.print(f"历史消息条数: {len(state.history)}")
|
|
235
244
|
console.print(f"累计输入 tokens:{state.input_tokens}")
|
|
236
245
|
console.print(f"累计输出 tokens:{state.output_tokens}\n")
|
|
@@ -0,0 +1,202 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import asyncio
|
|
4
|
+
import time
|
|
5
|
+
from collections.abc import Awaitable, Callable
|
|
6
|
+
|
|
7
|
+
from prompt_toolkit.application import Application
|
|
8
|
+
from prompt_toolkit.buffer import Buffer
|
|
9
|
+
from prompt_toolkit.filters import Condition
|
|
10
|
+
from prompt_toolkit.formatted_text import HTML
|
|
11
|
+
from prompt_toolkit.history import InMemoryHistory
|
|
12
|
+
from prompt_toolkit.key_binding import KeyBindings
|
|
13
|
+
from prompt_toolkit.key_binding.key_processor import KeyPressEvent
|
|
14
|
+
from prompt_toolkit.layout.containers import ConditionalContainer, HSplit, Window
|
|
15
|
+
from prompt_toolkit.layout.controls import BufferControl, FormattedTextControl
|
|
16
|
+
from prompt_toolkit.layout.dimension import Dimension
|
|
17
|
+
from prompt_toolkit.layout.layout import Layout
|
|
18
|
+
from prompt_toolkit.patch_stdout import patch_stdout
|
|
19
|
+
from rich.markup import escape
|
|
20
|
+
|
|
21
|
+
from core import permissions
|
|
22
|
+
|
|
23
|
+
from .commands import SessionState
|
|
24
|
+
from .render import console
|
|
25
|
+
|
|
26
|
+
# 常驻输入区:输入框整个会话期间挂在屏幕底部不消失,
|
|
27
|
+
# Agent 输出通过 patch_stdout 打印在它上方,请求期间按 ESC / Ctrl+C 能立刻打断。
|
|
28
|
+
|
|
29
|
+
# working 指示器的转圈动画帧
|
|
30
|
+
_SPINNER = "⠋⠙⠹⠸⠼⠴⠦⠧⠇⠏"
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
class Repl:
|
|
34
|
+
# 常驻输入区,由 main 注入 on_submit 回调来处理每一行输入
|
|
35
|
+
|
|
36
|
+
def __init__(self, state: SessionState) -> None:
|
|
37
|
+
self.state = state
|
|
38
|
+
# on_submit 由 main 注入,处理一行输入(命令或交给 Agent)
|
|
39
|
+
self._on_submit: Callable[[str], Awaitable[None]] | None = None
|
|
40
|
+
# 当前处理输入的后台任务,ESC / Ctrl+C 据此打断;None 表示空闲
|
|
41
|
+
self._task: asyncio.Task[None] | None = None
|
|
42
|
+
# 是否正在请求模型,决定上方 working... 指示器的显隐
|
|
43
|
+
self.working = False
|
|
44
|
+
self._work_start = 0.0
|
|
45
|
+
self._frame = 0
|
|
46
|
+
self._buffer = Buffer(multiline=False, history=InMemoryHistory())
|
|
47
|
+
self.app = self._build_app()
|
|
48
|
+
|
|
49
|
+
def _prompt_prefix(self, line_number: int, wrap_count: int) -> HTML:
|
|
50
|
+
return HTML("<ansicyan>❯ </ansicyan>")
|
|
51
|
+
|
|
52
|
+
def _working_line(self) -> HTML:
|
|
53
|
+
# 输入框上方那行:转圈帧 + 已耗时 + 打断提示,只在请求期间显示
|
|
54
|
+
frame = _SPINNER[self._frame % len(_SPINNER)]
|
|
55
|
+
elapsed = int(time.monotonic() - self._work_start)
|
|
56
|
+
return HTML(
|
|
57
|
+
f"<ansigreen>{frame}</ansigreen> <b>Working…</b>"
|
|
58
|
+
f"<ansibrightblack>(已耗时 {elapsed}s · 按 esc 打断)</ansibrightblack>"
|
|
59
|
+
)
|
|
60
|
+
|
|
61
|
+
def _mode_line(self) -> HTML:
|
|
62
|
+
# 输入框下方那行:当前权限模式 + 切换提示
|
|
63
|
+
mode = permissions.state.mode
|
|
64
|
+
return HTML(
|
|
65
|
+
f" <ansimagenta><b>▶▶ {mode}</b></ansimagenta>"
|
|
66
|
+
f"<ansibrightblack>(Shift+Tab 切换)</ansibrightblack>"
|
|
67
|
+
)
|
|
68
|
+
|
|
69
|
+
def _divider(self) -> Window:
|
|
70
|
+
# 一条横向分割线
|
|
71
|
+
return Window(height=1, char="─", style="fg:ansibrightblack")
|
|
72
|
+
|
|
73
|
+
def _build_app(self) -> Application[None]:
|
|
74
|
+
# 从上到下:working 指示器、分割线、输入行、分割线、模式行
|
|
75
|
+
layout = Layout(
|
|
76
|
+
HSplit(
|
|
77
|
+
[
|
|
78
|
+
ConditionalContainer(
|
|
79
|
+
Window(FormattedTextControl(self._working_line), height=1),
|
|
80
|
+
filter=Condition(lambda: self.working),
|
|
81
|
+
),
|
|
82
|
+
self._divider(),
|
|
83
|
+
# 输入行只占内容高度,多行自动换行撑开
|
|
84
|
+
Window(
|
|
85
|
+
BufferControl(buffer=self._buffer),
|
|
86
|
+
get_line_prefix=self._prompt_prefix,
|
|
87
|
+
height=Dimension(min=1),
|
|
88
|
+
wrap_lines=True,
|
|
89
|
+
dont_extend_height=True,
|
|
90
|
+
),
|
|
91
|
+
self._divider(),
|
|
92
|
+
Window(FormattedTextControl(self._mode_line), height=1),
|
|
93
|
+
]
|
|
94
|
+
)
|
|
95
|
+
)
|
|
96
|
+
return Application(layout=layout, key_bindings=self._build_key_bindings())
|
|
97
|
+
|
|
98
|
+
def _build_key_bindings(self) -> KeyBindings:
|
|
99
|
+
kb = KeyBindings()
|
|
100
|
+
|
|
101
|
+
@kb.add("enter")
|
|
102
|
+
def _(event: KeyPressEvent) -> None:
|
|
103
|
+
self._on_enter()
|
|
104
|
+
|
|
105
|
+
@kb.add("c-c")
|
|
106
|
+
def _(event: KeyPressEvent) -> None:
|
|
107
|
+
# 请求中打断当前 Agent loop;空闲时退出程序
|
|
108
|
+
if self._task is not None:
|
|
109
|
+
self._task.cancel()
|
|
110
|
+
else:
|
|
111
|
+
event.app.exit()
|
|
112
|
+
|
|
113
|
+
@kb.add("escape")
|
|
114
|
+
def _(event: KeyPressEvent) -> None:
|
|
115
|
+
# 请求中打断;空闲时清空输入行
|
|
116
|
+
if self._task is not None:
|
|
117
|
+
self._task.cancel()
|
|
118
|
+
else:
|
|
119
|
+
self._buffer.reset()
|
|
120
|
+
|
|
121
|
+
@kb.add("c-d")
|
|
122
|
+
def _(event: KeyPressEvent) -> None:
|
|
123
|
+
# 空闲且输入为空时退出
|
|
124
|
+
if self._task is None and not self._buffer.text:
|
|
125
|
+
event.app.exit()
|
|
126
|
+
|
|
127
|
+
@kb.add("s-tab")
|
|
128
|
+
def _(event: KeyPressEvent) -> None:
|
|
129
|
+
# 循环切换权限模式
|
|
130
|
+
permissions.cycle_mode()
|
|
131
|
+
|
|
132
|
+
return kb
|
|
133
|
+
|
|
134
|
+
def _on_enter(self) -> None:
|
|
135
|
+
# 请求中不接受新提交(输入框仍在,只是回车不触发新一轮)
|
|
136
|
+
if self._task is not None:
|
|
137
|
+
return
|
|
138
|
+
text = self._buffer.text.strip()
|
|
139
|
+
if not text:
|
|
140
|
+
self._buffer.reset()
|
|
141
|
+
return
|
|
142
|
+
# 存进输入历史,清空输入行
|
|
143
|
+
self._buffer.history.append_string(text)
|
|
144
|
+
self._buffer.reset()
|
|
145
|
+
# 把这行回显到上方滚动区,留下记录(输入框常驻,不回显的话提交后这行就没了)
|
|
146
|
+
self._echo_input(text)
|
|
147
|
+
# 把处理丢进后台任务,回车处理立刻返回,输入框继续渲染、随时能打断
|
|
148
|
+
self._task = self.app.create_background_task(self._process(text))
|
|
149
|
+
|
|
150
|
+
def _echo_input(self, text: str) -> None:
|
|
151
|
+
# 回显刚提交的一行:上下分割线夹住 ❯ 文本,和输入框观感一致
|
|
152
|
+
rule = "─" * console.width
|
|
153
|
+
console.print(f"[bright_black]{rule}[/]")
|
|
154
|
+
console.print(f"[cyan]❯[/] {escape(text)}")
|
|
155
|
+
console.print(f"[bright_black]{rule}[/]")
|
|
156
|
+
console.print()
|
|
157
|
+
|
|
158
|
+
async def _process(self, text: str) -> None:
|
|
159
|
+
# 后台任务:交给 on_submit,统一兜住打断和异常
|
|
160
|
+
assert self._on_submit is not None
|
|
161
|
+
try:
|
|
162
|
+
await self._on_submit(text)
|
|
163
|
+
except asyncio.CancelledError:
|
|
164
|
+
console.print("\n[bold yellow]已中断[/]\n")
|
|
165
|
+
except Exception as e:
|
|
166
|
+
console.print(f"\n[bold red]✗ {type(e).__name__}: {e}[/]\n")
|
|
167
|
+
finally:
|
|
168
|
+
self._task = None
|
|
169
|
+
self.working = False
|
|
170
|
+
self.app.invalidate()
|
|
171
|
+
|
|
172
|
+
def start_working(self) -> None:
|
|
173
|
+
# 进入请求中状态,上方显示 working...;请求结束 / 被打断后由 _process 的 finally 统一清掉
|
|
174
|
+
self.working = True
|
|
175
|
+
self._work_start = time.monotonic()
|
|
176
|
+
self.app.invalidate()
|
|
177
|
+
|
|
178
|
+
def exit(self) -> None:
|
|
179
|
+
# 结束常驻输入区(/exit 命令用)
|
|
180
|
+
self.app.exit()
|
|
181
|
+
|
|
182
|
+
async def run(self, on_submit: Callable[[str], Awaitable[None]]) -> None:
|
|
183
|
+
# 运行 REPL 直到用户退出;on_submit 是处理一行输入的异步回调
|
|
184
|
+
self._on_submit = on_submit
|
|
185
|
+
# 转圈动画的心跳:请求期间定时重绘
|
|
186
|
+
ticker = asyncio.ensure_future(self._tick())
|
|
187
|
+
try:
|
|
188
|
+
# patch_stdout 让 Agent 的 rich 输出打印在输入框上方而不是冲掉它
|
|
189
|
+
with patch_stdout(raw=True):
|
|
190
|
+
await self.app.run_async()
|
|
191
|
+
finally:
|
|
192
|
+
ticker.cancel()
|
|
193
|
+
|
|
194
|
+
async def _tick(self) -> None:
|
|
195
|
+
try:
|
|
196
|
+
while True:
|
|
197
|
+
await asyncio.sleep(0.1)
|
|
198
|
+
if self.working:
|
|
199
|
+
self._frame += 1
|
|
200
|
+
self.app.invalidate()
|
|
201
|
+
except asyncio.CancelledError:
|
|
202
|
+
pass
|
|
@@ -1,44 +0,0 @@
|
|
|
1
|
-
"""
|
|
2
|
-
对 OpenAI 兼容网关做响应清洗:部分服务商偶发返回不规范字段,
|
|
3
|
-
导致 pydantic-ai 校验 ChatCompletion 失败(表现为「好一次、坏一次」)。
|
|
4
|
-
"""
|
|
5
|
-
|
|
6
|
-
from __future__ import annotations
|
|
7
|
-
|
|
8
|
-
from typing import Any
|
|
9
|
-
|
|
10
|
-
from openai.types import chat
|
|
11
|
-
from pydantic_ai.models import openai as openai_model
|
|
12
|
-
from pydantic_ai.models.openai import OpenAIChatModel
|
|
13
|
-
|
|
14
|
-
|
|
15
|
-
def _coerce_chat_completion_payload(data: dict[str, Any]) -> dict[str, Any]:
|
|
16
|
-
"""
|
|
17
|
-
尽量把兼容网关的松散 JSON 修成 OpenAI ChatCompletion 能过校验的形状。
|
|
18
|
-
"""
|
|
19
|
-
if data.get("object") != "chat.completion":
|
|
20
|
-
data["object"] = "chat.completion"
|
|
21
|
-
|
|
22
|
-
choices = data.get("choices")
|
|
23
|
-
if isinstance(choices, list):
|
|
24
|
-
for i, choice in enumerate(choices):
|
|
25
|
-
if not isinstance(choice, dict):
|
|
26
|
-
continue
|
|
27
|
-
index = choice.get("index", i)
|
|
28
|
-
if isinstance(index, str) and index.isdigit():
|
|
29
|
-
choice["index"] = int(index)
|
|
30
|
-
elif not isinstance(index, int):
|
|
31
|
-
choice["index"] = i
|
|
32
|
-
if choice.get("finish_reason") is None:
|
|
33
|
-
choice["finish_reason"] = "stop"
|
|
34
|
-
return data
|
|
35
|
-
|
|
36
|
-
|
|
37
|
-
class CompatibleOpenAIChatModel(OpenAIChatModel):
|
|
38
|
-
"""
|
|
39
|
-
覆盖校验钩子:先 coerce,再走与父类相同的 _ChatCompletion 校验。
|
|
40
|
-
"""
|
|
41
|
-
|
|
42
|
-
def _validate_completion(self, response: chat.ChatCompletion) -> Any:
|
|
43
|
-
data = _coerce_chat_completion_payload(response.model_dump())
|
|
44
|
-
return openai_model._ChatCompletion.model_validate(data)
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|