Files
countbot-ai__countbot/backend/app.py
T
w 08e66ccc02 发布(0.7.0): 修复问题并增强技能与工具体验
1. 解决多个 issue 反馈问题,修复已知 Bug
2. 优化前端界面与交互体验
3. 优化 tool 调用链路,较此前版本节省约 70% token 用量
4. 增加模型思考控制开关,提升整体响应体感
5. 新增 find-skills,全面接入腾讯云 SkillsHub,支持通过对话进行 skills 管理
6. 新增 ima-knowledge-base、ima-notes,全面接入 IMA 知识库和笔记,支持知识库与笔记内容管理和写入
7. 优化 README,并补充 0.7.0 发布说明
8. 发布说明:https://654321.ai/docs/releases/v0.7.0
2026-03-29 22:37:11 +08:00

579 lines
21 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""FastAPI 应用入口"""
import asyncio
from contextlib import asynccontextmanager
from pathlib import Path
from fastapi import FastAPI, HTTPException
from fastapi.middleware.cors import CORSMiddleware
from fastapi.staticfiles import StaticFiles
from loguru import logger
from backend.utils.logger import setup_logger
from backend.database import get_db_session_factory
from backend.version import APP_VERSION
setup_logger()
def _create_shared_components(config, config_loader=None):
"""创建共享组件WebSocket 和渠道处理器共用)"""
from loguru import logger
from backend.modules.providers import create_provider
from backend.modules.providers.runtime import (
find_first_selectable_provider,
get_provider_runtime_state,
)
from backend.modules.agent.context import ContextBuilder
from backend.modules.agent.memory import MemoryStore
from backend.modules.agent.skills import SkillsLoader
from backend.modules.agent.subagent import SubagentManager
from backend.modules.tools.setup import register_all_tools
from backend.modules.workspace import (
seed_bundled_workspace_resources,
workspace_manager,
)
logger.info("Getting provider metadata...")
provider_id = config.model.provider
runtime_state = get_provider_runtime_state(config, provider_id)
if not runtime_state.selectable:
fallback_state = find_first_selectable_provider(config)
if fallback_state and fallback_state.provider_id != provider_id:
logger.warning(
f"共享组件默认 provider '{provider_id}' 不可用({runtime_state.reason}"
f"已回退到 '{fallback_state.provider_id}'"
)
runtime_state = fallback_state
provider_id = fallback_state.provider_id
else:
logger.warning(
f"共享组件默认 provider '{provider_id}' 当前不可用({runtime_state.reason}"
"将继续使用现有配置完成启动,实际请求阶段会再校验"
)
logger.info("Setting up workspace...")
workspace, used_fallback = workspace_manager.resolve_workspace_path_or_default(
config.workspace.path
)
workspace_manager.activate_workspace_path(workspace)
if used_fallback:
config.workspace.path = str(workspace)
logger.warning(f"共享组件启动时已回退到默认工作空间: {workspace}")
logger.info("Creating LLM provider...")
provider = create_provider(
api_key=runtime_state.api_key or None,
api_base=runtime_state.api_base,
default_model=config.model.model,
api_mode=config.model.api_mode,
timeout=120.0,
max_retries=3,
provider_id=provider_id,
)
logger.info("Creating memory and skills directories...")
memory_dir = workspace / "memory"
memory_dir.mkdir(parents=True, exist_ok=True)
seed_bundled_workspace_resources(workspace)
# Skills 目录始终从 workspace/skills 加载
# 确保用户修改 workspace 路径后skills 也在新路径下
skills_dir = workspace / "skills"
skills_dir.mkdir(parents=True, exist_ok=True)
logger.info(f"Workspace: {workspace}")
logger.info(f"Skills directory: {skills_dir}")
logger.info("Initializing memory store...")
memory = MemoryStore(memory_dir)
logger.info("Loading skills...")
skills = SkillsLoader(skills_dir)
logger.info("Building context builder...")
context_builder = ContextBuilder(
workspace=workspace,
memory=memory,
skills=skills,
persona_config=config.persona,
)
logger.info("Creating subagent manager...")
subagent_manager = SubagentManager(
provider=provider,
workspace=workspace,
model=config.model.model,
temperature=config.model.temperature,
max_tokens=config.model.max_tokens,
db_session_factory=get_db_session_factory(),
config_loader=config_loader,
skills=skills,
)
logger.info("Preparing tool parameters...")
tool_params = dict(
workspace=workspace,
command_timeout=config.security.command_timeout,
max_output_length=config.security.max_output_length,
allow_dangerous=not config.security.dangerous_commands_blocked,
restrict_to_workspace=config.security.restrict_to_workspace,
custom_deny_patterns=config.security.custom_deny_patterns,
custom_allow_patterns=(
config.security.custom_allow_patterns
if config.security.command_whitelist_enabled
else None
),
audit_log_enabled=config.security.audit_log_enabled,
subagent_manager=subagent_manager,
skills_loader=skills,
)
logger.info("Registering all tools...")
tool_registry = register_all_tools(**tool_params, memory_store=memory)
logger.info(f"Registered {len(tool_registry)} tools")
return dict(
provider=provider,
workspace=workspace,
context_builder=context_builder,
subagent_manager=subagent_manager,
tool_registry=tool_registry,
tool_params=tool_params,
memory=memory,
skills=skills,
)
@asynccontextmanager
async def lifespan(app: FastAPI):
"""应用生命周期管理"""
from backend.database import init_db, get_db_session_factory
from backend.modules.config.loader import config_loader
from backend.modules.channels.manager import ChannelManager
from backend.modules.messaging.enterprise_queue import EnterpriseMessageQueue
from backend.modules.messaging.rate_limiter import RateLimiter
from backend.modules.channels.handler import ChannelMessageHandler
from backend.modules.cron.executor import CronExecutor
from backend.modules.cron.scheduler import CronScheduler
from backend.modules.cron.service import CronService
from backend.modules.agent.loop import AgentLoop
from backend.modules.session.manager import SessionManager
from backend.modules.tools.setup import register_all_tools
from backend.api.channels import set_channel_manager
from backend.modules.auth.middleware import (
clear_remote_setup_secret,
ensure_remote_setup_secret,
get_remote_setup_secret_ttl_minutes,
)
from backend.modules.auth.router import get_password_hash as get_auth_password_hash
# 初始化数据库和配置
logger.info("Starting CountBot backend...")
await init_db()
logger.info("Database initialized")
await config_loader.load()
logger.info("Configuration loaded")
config = config_loader.config
if await get_auth_password_hash():
clear_remote_setup_secret(app)
else:
setup_secret = ensure_remote_setup_secret(app)
setup_secret_ttl_minutes = get_remote_setup_secret_ttl_minutes()
logger.info(
f"远程首次初始化入口:将此路径拼接到上方 Network 地址后访问(有效期 {setup_secret_ttl_minutes} 分钟,初始化成功后立即失效) -> /setup/{setup_secret}"
)
logger.info(
f"Remote first-time setup entry: append this path to the Network URL above (valid for {setup_secret_ttl_minutes} minutes and expires immediately after setup succeeds) -> /setup/{setup_secret}"
)
# 创建共享组件
logger.info("Creating shared components...")
shared = _create_shared_components(config, config_loader)
app.state.shared = shared
app.state.skills = shared["skills"]
app.state.memory = shared["memory"]
logger.info("Shared components created")
# 设置全局 SubagentManager
from backend.api.chat import set_global_subagent_manager
set_global_subagent_manager(shared["subagent_manager"])
logger.info("Creating message queue and rate limiter...")
message_queue = EnterpriseMessageQueue(
enable_dedup=True,
dedup_window=10
)
app.state.message_queue = message_queue
rate_limiter = RateLimiter(rate=10, per=60)
logger.info("Message queue and rate limiter created")
# 创建渠道消息处理器
logger.info("Creating message handler...")
message_handler = ChannelMessageHandler(
provider=shared["provider"],
workspace=shared["workspace"],
model=config.model.model,
bus=message_queue,
context_builder=shared["context_builder"],
tool_params=shared["tool_params"],
subagent_manager=shared["subagent_manager"],
max_iterations=config.model.max_iterations,
rate_limiter=rate_limiter,
temperature=config.model.temperature,
max_tokens=config.model.max_tokens,
thinking_enabled=config.model.thinking_enabled,
max_history_messages=config.persona.max_history_messages,
memory_store=shared["memory"],
)
app.state.message_handler = message_handler
logger.info("Message handler created")
# 创建渠道管理器
logger.info("Creating channel manager...")
channel_manager = ChannelManager(config, message_queue)
app.state.channel_manager = channel_manager
set_channel_manager(channel_manager)
message_handler.set_channel_manager(channel_manager)
logger.info("Channel manager created")
# 初始化 OSS 上传器(可选)
logger.info("Initializing OSS uploader (optional)...")
try:
from backend.modules.tools.image_uploader import init_oss_uploader
oss_config = None
if hasattr(config.channels, "qq") and hasattr(config.channels.qq, "oss"):
oss_config = config.channels.qq.oss.model_dump()
init_oss_uploader(oss_config)
logger.info("OSS uploader initialized")
except Exception as e:
logger.warning(f"OSS uploader init failed (optional): {e}")
# 启动后台任务(不等待完成)
app.state.background_tasks = []
app.state.channel_manager_task = None
if channel_manager.enabled_channels:
task = asyncio.create_task(channel_manager.start_all())
app.state.channel_manager_task = task
app.state.background_tasks.append(task)
logger.info(f"Started {len(channel_manager.enabled_channels)} channel(s) in background")
task = asyncio.create_task(message_handler.start_processing())
app.state.background_tasks.append(task)
logger.info("Started message handler in background")
# 初始化定时任务系统
logger.info("Initializing cron system...")
cron_tool_registry = register_all_tools(
**shared["tool_params"],
)
cron_agent = AgentLoop(
provider=shared["provider"],
workspace=shared["workspace"],
tools=cron_tool_registry,
context_builder=shared["context_builder"],
subagent_manager=shared["subagent_manager"],
model=config.model.model,
max_iterations=config.model.max_iterations,
temperature=config.model.temperature,
max_tokens=config.model.max_tokens,
thinking_enabled=config.model.thinking_enabled,
)
session_manager = SessionManager(shared["workspace"])
logger.info("Cron agent and session manager created")
# 初始化心跳服务
logger.info("Initializing heartbeat service...")
db_session_factory = get_db_session_factory()
from backend.modules.agent.heartbeat import HeartbeatService, ensure_heartbeat_job
heartbeat_config = config.persona.heartbeat
heartbeat_service = HeartbeatService(
provider=shared["provider"],
model=config.model.model,
workspace=shared["workspace"],
db_session_factory=db_session_factory,
ai_name=config.persona.ai_name or "小C",
user_name=config.persona.user_name or "主人",
user_address=config.persona.user_address or "",
personality=config.persona.personality or "professional",
custom_personality=config.persona.custom_personality or "",
idle_threshold_hours=heartbeat_config.idle_threshold_hours,
quiet_start=heartbeat_config.quiet_start,
quiet_end=heartbeat_config.quiet_end,
max_greets_per_day=heartbeat_config.max_greets_per_day,
)
logger.info("Heartbeat service created")
logger.info("Creating cron executor...")
cron_executor = CronExecutor(
agent=cron_agent,
bus=message_queue,
session_manager=session_manager,
channel_manager=channel_manager,
heartbeat_service=heartbeat_service,
)
logger.info("Cron executor created")
async def on_cron_execute(
job_id: str,
message: str,
channel: str,
account_id: str,
chat_id: str,
deliver_response: bool,
) -> str:
return await cron_executor.execute(
job_id, message, channel, account_id, chat_id, deliver_response
)
logger.info("Creating cron scheduler...")
scheduler = CronScheduler(
db_session_factory=db_session_factory,
on_execute=on_cron_execute,
)
await scheduler.start()
logger.info("Cron scheduler started")
# 注册内置心跳任务
logger.info("Ensuring heartbeat job...")
await ensure_heartbeat_job(db_session_factory, heartbeat_config=heartbeat_config)
await scheduler.trigger_reschedule()
logger.info("Heartbeat job ensured")
app.state.cron_scheduler = scheduler
app.state.cron_executor = cron_executor
# 注册进程退出清理处理器(备用机制)
import atexit
def cleanup_on_exit() -> None:
"""进程退出时的清理函数"""
logger.info("atexit cleanup triggered")
try:
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
try:
loop.run_until_complete(channel_manager.stop_all())
finally:
loop.close()
except RuntimeError as e:
logger.debug(f"Event loop already closed: {e}")
except Exception as e:
logger.error(f"Error in atexit cleanup: {e}")
atexit.register(cleanup_on_exit)
logger.info("Backend started successfully")
yield
# 正常关闭流程
logger.info("Initiating graceful shutdown...")
await channel_manager.stop_all()
await scheduler.stop()
logger.info("Backend shutdown complete")
app = FastAPI(
title="CountBot Desktop API",
description="CountBot backend API",
version=APP_VERSION,
lifespan=lifespan,
docs_url=None,
redoc_url=None,
openapi_url=None,
)
# 保存绑定地址用于认证判断
import os as _os
app.state.bind_host = _os.getenv("HOST", "127.0.0.1")
def get_tool_registry():
"""返回全局共享工具注册表,供 XiaozhiChannel 等频道内部调用。"""
try:
return app.state.shared.get("tool_registry")
except AttributeError:
return None
app.add_middleware(
CORSMiddleware,
allow_origins=[],
allow_credentials=False,
allow_methods=["GET", "POST", "PUT", "DELETE", "OPTIONS"],
allow_headers=["Content-Type", "Authorization"],
)
# 远程访问认证中间件
from backend.modules.auth.middleware import (
RemoteAuthMiddleware,
has_valid_remote_setup_secret,
)
from backend.modules.auth.router import get_password_hash
app.add_middleware(RemoteAuthMiddleware, get_password_hash_fn=get_password_hash)
# 注册 API 路由
from backend.api.chat import router as chat_router
from backend.api.settings import router as settings_router
from backend.api.tools import router as tools_router
from backend.api.memory import router as memory_router
from backend.api.skills import router as skills_router
from backend.api.cron import router as cron_router
from backend.api.tasks import router as tasks_router
from backend.api.system import router as system_router
from backend.api.channels import router as channels_router
from backend.api.queue import router as queue_router
from backend.api.auth import router as auth_router
from backend.api.personalities import router as personalities_router
from backend.api.agent_teams import router as agent_teams_router
app.include_router(auth_router)
app.include_router(chat_router)
app.include_router(settings_router)
app.include_router(tools_router)
app.include_router(memory_router)
app.include_router(skills_router)
app.include_router(cron_router)
app.include_router(tasks_router)
app.include_router(system_router)
app.include_router(channels_router)
app.include_router(queue_router)
app.include_router(personalities_router)
app.include_router(agent_teams_router)
# WebSocket 端点
from fastapi import WebSocket
from backend.ws.connection import handle_websocket
@app.websocket("/ws/chat")
async def websocket_endpoint(websocket: WebSocket):
"""WebSocket 聊天端点,复用共享组件"""
from backend.modules.agent.loop import AgentLoop
from backend.modules.providers import create_provider
from backend.modules.providers.registry import get_provider_metadata
from backend.modules.tools.setup import register_all_tools
from backend.modules.auth.middleware import AUTH_COOKIE_NAME, is_direct_local_client
from backend.modules.auth.utils import validate_session as validate_ws_session
from backend.modules.auth.router import get_password_hash as get_ws_password_hash
client_ip = websocket.client.host if websocket.client and websocket.client.host else None
is_local = is_direct_local_client(client_ip, websocket.headers.keys())
if not is_local:
auth_enabled = bool(await get_ws_password_hash())
if not auth_enabled:
await websocket.close(code=4003, reason="Authentication setup required")
return
token = websocket.cookies.get(AUTH_COOKIE_NAME)
if not token:
auth_header = websocket.headers.get("authorization", "")
if auth_header.startswith("Bearer "):
token = auth_header[7:]
if not token or not validate_ws_session(token):
await websocket.close(code=4001, reason="Authentication required")
return
shared = websocket.app.state.shared
# 每个 WebSocket 连接使用独立的工具注册表(会话隔离)
tool_registry = register_all_tools(
**shared["tool_params"],
memory_store=shared["memory"],
)
from backend.modules.config.loader import config_loader
config = config_loader.config
from backend.modules.providers.runtime import (
build_provider_unavailable_message,
find_first_selectable_provider,
get_provider_runtime_state,
)
# 根据当前配置创建 provider支持动态切换
provider_id = config.model.provider
runtime_state = get_provider_runtime_state(config, provider_id)
if not runtime_state.selectable:
fallback_state = find_first_selectable_provider(config)
if fallback_state:
logger.warning(
f"WebSocket 默认 provider '{provider_id}' 不可用({runtime_state.reason}"
f"已回退到 '{fallback_state.provider_id}'"
)
runtime_state = fallback_state
provider_id = fallback_state.provider_id
else:
await websocket.close(
code=1011,
reason=build_provider_unavailable_message(provider_id, runtime_state.reason),
)
return
provider = create_provider(
api_key=runtime_state.api_key or None,
api_base=runtime_state.api_base,
default_model=config.model.model,
api_mode=config.model.api_mode,
timeout=120.0,
max_retries=3,
provider_id=provider_id,
)
agent_loop = AgentLoop(
provider=provider,
workspace=shared["workspace"],
tools=tool_registry,
context_builder=shared["context_builder"],
subagent_manager=shared["subagent_manager"],
model=config.model.model,
max_iterations=config.model.max_iterations,
temperature=config.model.temperature,
max_tokens=config.model.max_tokens,
thinking_enabled=config.model.thinking_enabled,
)
await handle_websocket(websocket, agent_loop=agent_loop)
@app.get("/api/health")
async def health_check():
return {"status": "ok", "version": APP_VERSION}
# 挂载前端静态文件
from backend.utils.paths import APPLICATION_ROOT
frontend_dist = APPLICATION_ROOT / "frontend" / "dist"
if frontend_dist.exists():
from fastapi.responses import FileResponse
import mimetypes
# 确保 Windows 上正确识别 JavaScript 模块的 MIME 类型
mimetypes.add_type("application/javascript", ".js")
mimetypes.add_type("text/css", ".css")
mimetypes.add_type("image/svg+xml", ".svg")
# SPA 路由回退(必须在 StaticFiles 之前注册)
@app.get("/login")
async def spa_login():
return FileResponse(str(frontend_dist / "index.html"))
@app.get("/setup/{setup_secret}")
async def spa_setup(setup_secret: str):
if await get_password_hash():
raise HTTPException(status_code=404)
if not has_valid_remote_setup_secret(app, setup_secret):
raise HTTPException(status_code=404)
return FileResponse(str(frontend_dist / "index.html"))
app.mount("/", StaticFiles(directory=str(frontend_dist), html=True), name="static")