Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions .agents/scripts/check_runtime_imports.py
Original file line number Diff line number Diff line change
Expand Up @@ -45,6 +45,7 @@
OPTIONAL_DEPS: set[str] = {
"bilibili_api",
"litellm",
"mcp",
"schedule",
}

Expand Down
2 changes: 2 additions & 0 deletions .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -61,6 +61,8 @@ config.yaml

# === Development ===
dev/
.claude/
CLAUDE.md

# === WebUI frontend ===
ncatbot/webui/frontend/node_modules/
Expand Down
94 changes: 93 additions & 1 deletion ncatbot/adapter/ai/adapter.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,25 @@
LOG = get_log("AIAdapter")


def _parse_kv(text: str, sep: str = "=") -> Dict[str, str]:
"""解析 ``KEY=VALUE`` 逗号分隔文本为字典。"""
result: Dict[str, str] = {}
for part in text.split(","):
part = part.strip()
if not part:
continue
key, _, value = part.partition(sep)
key = key.strip()
if key:
result[key] = value.strip()
return result


def _parse_headers(text: str) -> Dict[str, str]:
"""解析 ``Key:Value`` 逗号分隔文本为请求头字典。"""
return _parse_kv(text, sep=":")


class AIAdapter(BaseAdapter):
"""AI 适配器 — 通过 litellm 统一调用 100+ LLM 提供商

Expand All @@ -45,7 +64,10 @@ class AIAdapter(BaseAdapter):
description = "AI 适配器(基于 litellm 的多模型统一接口)"
supported_protocols: List[str] = ["litellm"]
platform = "ai"
pip_dependencies: Dict[str, str] = {"litellm": ">=1.40.0"}
pip_dependencies: Dict[str, str] = {
"litellm": ">=1.40.0",
"mcp": ">=1.25.0,<2.0.0",
}

@classmethod
def cli_configure(cls) -> Dict[str, Any]:
Expand All @@ -71,8 +93,78 @@ def cli_configure(cls) -> Dict[str, Any]:
cfg["base_url"] = base_url
if completion_model:
cfg["completion_model"] = completion_model

mcp_servers = cls._cli_configure_mcp()
if mcp_servers:
cfg["mcp_servers"] = mcp_servers
return cfg

@classmethod
def _cli_configure_mcp(cls) -> Dict[str, Any]:
"""交互式收集 MCP 服务器配置(LiteLLM 兼容格式)。"""
import click

servers: Dict[str, Any] = {}
if not click.confirm(
"是否添加 MCP 服务器(让模型调用外部工具)?", default=False
):
return servers

while True:
click.echo(
click.style(
"\n ── MCP 服务器 ──\n"
" 传输类型: http (Streamable HTTP) / sse / stdio\n"
" http/sse 需 URL,stdio 需 command + args",
dim=True,
)
)
name = click.prompt("服务器名称(如 deepwiki)").strip()
if not name:
click.echo(click.style("服务器名称不能为空,已跳过", fg="yellow"))
continue

transport = click.prompt(
"传输类型", type=click.Choice(["http", "sse", "stdio"]), default="http"
)

entry: Dict[str, Any] = {"transport": transport}
if transport == "stdio":
command = click.prompt("启动命令(如 npx / uvx)")
entry["command"] = command
args = click.prompt(
"参数(空格分隔,可直接回车跳过)",
default="",
show_default=False,
)
if args:
entry["args"] = args.split()
env = click.prompt(
"环境变量(逗号分隔 KEY=VALUE,可回车跳过)",
default="",
show_default=False,
)
parsed_env = _parse_kv(env)
if parsed_env:
entry["env"] = parsed_env
else:
url = click.prompt("MCP 服务器 URL")
entry["url"] = url
headers = click.prompt(
"请求头(逗号分隔 Key:Value,可回车跳过)",
default="",
show_default=False,
)
parsed_headers = _parse_headers(headers)
if parsed_headers:
entry["headers"] = parsed_headers

servers[name] = entry
if not click.confirm("继续添加 MCP 服务器?", default=False):
break

return servers

def __init__(self, **kwargs: Any) -> None:
super().__init__(**kwargs)
self._ai_config = AIConfig(**self._raw_config)
Expand Down
3 changes: 2 additions & 1 deletion ncatbot/adapter/ai/api/__init__.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
"""AI 平台 API 子模块"""

from .bot_api import AIBotAPI
from .mcp import MCPSessionManager

__all__ = ["AIBotAPI"]
__all__ = ["AIBotAPI", "MCPSessionManager"]
109 changes: 106 additions & 3 deletions ncatbot/adapter/ai/api/bot_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
from ncatbot.utils import get_log

from ..config import AIConfig
from .mcp import MCPSessionManager

# 接受的输入类型:str / list[dict] / MessageArray / 单个 MessageSegment
ChatInput = Union[str, List[dict], "MessageArray", "MessageSegment"]
Expand Down Expand Up @@ -63,6 +64,8 @@ async def chat(
temperature: Optional[float] = None,
max_tokens: Optional[int] = None,
nickname_map: Optional[Dict[str, str]] = None,
mcp_servers: Optional[Dict[str, dict]] = None,
max_tool_calls: int = 10,
**kwargs: Any,
) -> Any:
"""Chat Completion
Expand All @@ -83,6 +86,12 @@ async def chat(
nickname_map:
``{user_id: 昵称}`` 映射,用于将 ``At`` 段转为可读文本。
缺省时 At 段渲染为 ``@{user_id}``。
mcp_servers:
MCP 服务器配置字典(格式见 ``AIConfig.mcp_servers``)。
缺省使用配置中的 ``mcp_servers``;为空则不启用 MCP 工具。
启用后模型可调用 MCP 工具,工具名按 ``{server}_{tool}`` 命名空间区分。
max_tool_calls:
单轮对话中允许的最多工具调用轮数(默认 10),防止死循环。

Returns
-------
Expand All @@ -106,14 +115,104 @@ async def chat(
if resolved_max_tokens is not None:
call_kwargs["max_tokens"] = resolved_max_tokens

return await self._call_with_fallback(
merged_mcp_servers = (
mcp_servers if mcp_servers is not None else self._config.mcp_servers
)
if not merged_mcp_servers:
return await self._call_with_fallback(
acompletion,
resolved_model,
self._config.completion_model,
messages=messages,
**call_kwargs,
)

return await self._chat_with_tools(
acompletion,
resolved_model,
self._config.completion_model,
messages=messages,
**call_kwargs,
messages,
merged_mcp_servers,
max_tool_calls,
call_kwargs,
)

async def _chat_with_tools(
self,
acompletion: Any,
model: str,
default_model: str,
messages: List[dict],
mcp_servers: Dict[str, dict],
max_tool_calls: int,
call_kwargs: Dict[str, Any],
) -> Any:
"""带 MCP 工具调用循环的 Chat Completion。

模型请求工具 → 执行 MCP 工具 → 回传结果 → 继续对话,
直到模型不再请求工具或达到 ``max_tool_calls`` 上限。
"""
async with MCPSessionManager(mcp_servers) as mcp:
tools = await mcp.load_tools()
if not tools:
LOG.warning("MCP 服务器未加载到任何工具,按普通对话处理(无 tools)")
return await self._call_with_fallback(
acompletion,
model,
default_model,
messages=messages,
**call_kwargs,
)

msgs = list(messages)
resp: Any = None
for _ in range(max_tool_calls):
resp = await self._call_with_fallback(
acompletion,
model,
default_model,
messages=msgs,
tools=tools,
**call_kwargs,
)
message = resp.choices[0].message
tool_calls = getattr(message, "tool_calls", None)
if not tool_calls:
return resp

msgs.append(self._assistant_tool_message(message, tool_calls))
for tool_call in tool_calls:
try:
result_text = await mcp.call_openai_tool(tool_call)
except Exception as exc: # noqa: BLE001
LOG.error("MCP 工具调用失败: %s", exc)
result_text = f"[MCP 工具调用失败: {exc}]"
msgs.append(
{
"role": "tool",
"tool_call_id": tool_call.id,
"content": result_text,
}
)

LOG.warning("MCP 工具调用达到上限 %d 轮,返回最后一次响应", max_tool_calls)
return resp

@staticmethod
def _assistant_tool_message(message: Any, tool_calls: List[Any]) -> dict:
"""构造带 tool_calls 的 assistant 消息,供下一轮请求使用。"""
dumped = []
for tc in tool_calls:
if hasattr(tc, "model_dump"):
dumped.append(tc.model_dump())
else:
dumped.append(dict(tc))
return {
"role": "assistant",
"content": getattr(message, "content", None),
"tool_calls": dumped,
}

async def embeddings(
self,
input_text: Union[str, List[str]],
Expand Down Expand Up @@ -232,6 +331,8 @@ async def chat_text(
temperature: Optional[float] = None,
max_tokens: Optional[int] = None,
nickname_map: Optional[Dict[str, str]] = None,
mcp_servers: Optional[Dict[str, dict]] = None,
max_tool_calls: int = 10,
**kwargs: Any,
) -> str:
"""Chat Completion — 直接返回文本
Expand All @@ -244,6 +345,8 @@ async def chat_text(
temperature=temperature,
max_tokens=max_tokens,
nickname_map=nickname_map,
mcp_servers=mcp_servers,
max_tool_calls=max_tool_calls,
**kwargs,
)
return resp.choices[0].message.content or ""
Expand Down
Loading