"""Assistant CRUD backed by capability-to-model-resource bindings.""" import uuid from models import RuntimeTool from db.models import ( Assistant, AssistantModelBinding, AssistantToolBinding, KnowledgeBase, ModelResource, Tool, ) from db.session import get_session from fastapi import APIRouter, Depends, HTTPException from schemas import AssistantOut, AssistantUpsert from services.auth import require_admin from services.masking import mask, resolve_incoming_key from services.node_specs import graph_references, normalize_graph, validate_graph from services.tool_policy import policy_for_tool from services.workflow_engine import WorkflowEngine from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession router = APIRouter( prefix="/api/assistants", tags=["assistants"], dependencies=[Depends(require_admin)], ) CAPABILITIES = ("LLM", "ASR", "TTS", "Realtime", "Embedding", "Agent") def _validate_workflow(body: AssistantUpsert) -> None: """workflow 类型:保存前校验图结构,不通过则 400。其他类型跳过。""" if body.type != "workflow": return body.graph = normalize_graph(body.graph or {}) errors = validate_graph(body.graph) if errors: raise HTTPException(400, "工作流校验失败:" + ";".join(errors)) # Graph settings are the source of truth. The flat flag is only a session # capability summary used by clients to decide whether to open the camera. body.vision_enabled = WorkflowEngine(body.graph).uses_vision() body.vision_model_resource_id = None refs = graph_references(body.graph) body.tool_ids = list(dict.fromkeys([*body.tool_ids, *sorted(refs["tools"])])) async def _validate_workflow_references( session: AsyncSession, body: AssistantUpsert ) -> None: if body.type != "workflow" or not body.graph.get("nodes"): return graph = body.graph engine = WorkflowEngine(graph) settings = graph.get("settings") or {} resource_expectations: dict[str, str] = {} vision_resource_ids: set[str] = set() for key, capability in ( ("defaultLlmResourceId", "LLM"), ("defaultAsrResourceId", "ASR"), ("defaultTtsResourceId", "TTS"), ): if settings.get(key): resource_expectations[str(settings[key])] = capability knowledge_ids: set[str] = ( {str(settings["knowledgeBaseId"])} if settings.get("knowledgeBaseId") else set() ) for node in graph.get("nodes") or []: data = node.get("data") or {} if node.get("type") == "agent" and data.get("inheritGlobalConfig", True): continue if data.get("llmResourceId"): resource_expectations[str(data["llmResourceId"])] = "LLM" if data.get("asrResourceId"): resource_expectations[str(data["asrResourceId"])] = "ASR" if data.get("ttsResourceId"): resource_expectations[str(data["ttsResourceId"])] = "TTS" if data.get("knowledgeBaseId"): knowledge_ids.add(str(data["knowledgeBaseId"])) for node_id, node in engine.nodes.items(): if node.get("type") != "agent": continue stage = engine.agent_stage_config(node_id) if not stage.vision_enabled: continue resource_id = ( stage.vision_model_resource_id or stage.llm_resource_id ) if not resource_id: raise HTTPException( 400, f"Agent 节点 {node_id} 开启视觉理解时必须选择" "支持图片输入的大语言模型或视觉模型", ) resource_expectations[resource_id] = "LLM" vision_resource_ids.add(resource_id) for resource_id, capability in resource_expectations.items(): resource = await session.get(ModelResource, resource_id) if not resource or not resource.enabled or resource.capability != capability: raise HTTPException(400, f"Workflow 引用了无效的 {capability} 资源:{resource_id}") if resource_id in vision_resource_ids and not resource.support_image_input: raise HTTPException( 400, f"Workflow 视觉模型必须支持图片输入:{resource_id}", ) for knowledge_id in knowledge_ids: knowledge = await session.get(KnowledgeBase, knowledge_id) if not knowledge or knowledge.status != "active": raise HTTPException(400, f"Workflow 引用了无效知识库:{knowledge_id}") async def _validate_startup_actions( session: AsyncSession, body: AssistantUpsert, ) -> None: """Keep startup deterministic and reject unsupported lifecycle/tool pairs.""" actions = body.startup.actions if not actions: return for action in actions: tool = await session.get(Tool, action.tool_id) if not tool or tool.status != "active": raise HTTPException(400, f"启动 Action 引用了无效工具:{action.tool_id}") runtime_tool = RuntimeTool( id=tool.id, name=tool.name, function_name=tool.function_name, type=tool.type, definition=tool.definition or {}, ) if action.phase == "preflight" and runtime_tool.type not in {"http", "mcp"}: raise HTTPException(400, "preflight Action 仅支持 HTTP 或 MCP 工具") if action.phase == "opening" and runtime_tool.type not in { "http", "mcp", "client", }: raise HTTPException(400, "opening Action 不支持该工具类型") if runtime_tool.type != "client": continue policy = policy_for_tool(runtime_tool) if action.required and not policy.wait_for_response: raise HTTPException(400, "必需的 Client 启动 Action 必须等待客户端响应") if runtime_tool.function_name != "show_message": continue if policy.response_wait_mode != "session": raise HTTPException(400, "show_message 启动 Action 必须使用会话内等待") message = str(action.arguments.get("message") or "").strip() buttons = action.arguments.get("actions") if not message or len(message) > 2000: raise HTTPException(400, "show_message 重要信息必须为 1-2000 个字符") if not isinstance(buttons, list) or not buttons: raise HTTPException(400, "show_message 至少需要一个确认按钮") first_button = buttons[0] if isinstance(buttons[0], dict) else {} if not str(first_button.get("id") or "").strip() or not str( first_button.get("label") or "" ).strip(): raise HTTPException(400, "show_message 确认按钮必须配置 id 和文字") if action.required and action.arguments.get("dismissible") is not False: raise HTTPException(400, "必需的 show_message 启动 Action 不允许跳过确认") async def _validate_vision_model( session: AsyncSession, body: AssistantUpsert ) -> None: # Workflow 的视觉模型属于图中的全局/Agent 配置,由上面的图引用校验处理。 if body.type == "workflow": return if body.vision_enabled: if body.vision_model_resource_id: resource = await session.get(ModelResource, body.vision_model_resource_id) else: resource_id = body.model_resource_ids.get("LLM") resource = ( await session.get(ModelResource, resource_id) if resource_id else None ) if not resource or resource.capability != "LLM": raise HTTPException(400, "视觉模型必须引用 LLM 模型资源") if not resource.support_image_input: raise HTTPException(400, "视觉模型必须支持图片输入") async def _validate_knowledge_base(session: AsyncSession, body: AssistantUpsert) -> None: if body.runtime_mode != "pipeline" or body.type not in {"prompt", "workflow"}: body.knowledge_base_id = None return if body.knowledge_base_id and not await session.get(KnowledgeBase, body.knowledge_base_id): raise HTTPException(400, "知识库不存在") async def _sync_bindings( session: AsyncSession, assistant_id: str, resource_ids: dict[str, str] ) -> None: for capability in CAPABILITIES: resource_id = resource_ids.get(capability) binding = await session.get(AssistantModelBinding, (assistant_id, capability)) if not resource_id: if binding: await session.delete(binding) continue resource = await session.get(ModelResource, resource_id) if not resource or resource.capability != capability: raise HTTPException(400, f"{capability} 绑定必须引用同能力的模型资源") if binding: binding.model_resource_id = resource_id else: session.add( AssistantModelBinding( assistant_id=assistant_id, capability=capability, model_resource_id=resource_id, config={}, ) ) async def _resource_ids(session: AsyncSession, assistant_id: str) -> dict[str, str]: bindings = ( await session.execute( select(AssistantModelBinding).where( AssistantModelBinding.assistant_id == assistant_id ) ) ).scalars().all() return {binding.capability: binding.model_resource_id for binding in bindings} async def _sync_tool_bindings( session: AsyncSession, assistant_id: str, assistant_type: str, tool_ids: list[str] ) -> None: requested = ( list(dict.fromkeys(tool_ids)) if assistant_type in {"prompt", "workflow"} else [] ) if requested: tools = ( await session.execute(select(Tool).where(Tool.id.in_(requested))) ).scalars().all() found = {tool.id for tool in tools if tool.status == "active"} missing = [tool_id for tool_id in requested if tool_id not in found] if missing: raise HTTPException(400, f"工具不存在或未启用: {', '.join(missing)}") existing = ( await session.execute( select(AssistantToolBinding).where( AssistantToolBinding.assistant_id == assistant_id ) ) ).scalars().all() existing_by_id = {binding.tool_id: binding for binding in existing} for tool_id, binding in existing_by_id.items(): if tool_id not in requested: await session.delete(binding) for tool_id in requested: if tool_id not in existing_by_id: session.add(AssistantToolBinding(assistant_id=assistant_id, tool_id=tool_id)) async def _tool_ids(session: AsyncSession, assistant_id: str) -> list[str]: rows = ( await session.execute( select(AssistantToolBinding.tool_id) .where(AssistantToolBinding.assistant_id == assistant_id) .order_by(AssistantToolBinding.created_at) ) ).scalars().all() return list(rows) async def _to_out(session: AsyncSession, assistant: Assistant) -> AssistantOut: return AssistantOut( id=assistant.id, name=assistant.name, type=assistant.type, # type: ignore[arg-type] runtime_mode=assistant.runtime_mode, # type: ignore[arg-type] greeting=assistant.greeting, enable_interrupt=assistant.enable_interrupt, turn_config=assistant.turn_config or {}, startup=assistant.startup or {}, vision_enabled=assistant.vision_enabled, vision_model_resource_id=assistant.vision_model_resource_id, model_resource_ids=await _resource_ids(session, assistant.id), knowledge_base_id=assistant.knowledge_base_id, knowledge_retrieval_config=assistant.knowledge_retrieval_config or {}, tool_ids=await _tool_ids(session, assistant.id), prompt=assistant.prompt, dynamic_variable_definitions=assistant.dynamic_variable_definitions or {}, api_url=assistant.api_url, api_key=mask(assistant.api_key), app_id=assistant.app_id, graph=( normalize_graph(assistant.graph or {}) if assistant.type == "workflow" else {} ), updated_at=assistant.updated_at.isoformat() if assistant.updated_at else None, ) @router.get("", response_model=list[AssistantOut]) async def list_assistants(session: AsyncSession = Depends(get_session)): rows = ( await session.execute(select(Assistant).order_by(Assistant.updated_at.desc())) ).scalars().all() return [await _to_out(session, assistant) for assistant in rows] @router.post("", response_model=AssistantOut) async def create_assistant( body: AssistantUpsert, session: AsyncSession = Depends(get_session) ): _validate_workflow(body) await _validate_workflow_references(session, body) await _validate_startup_actions(session, body) await _validate_vision_model(session, body) await _validate_knowledge_base(session, body) data = body.model_dump() resource_ids = data.pop("model_resource_ids") tool_ids = data.pop("tool_ids") assistant = Assistant(id=f"asst_{uuid.uuid4().hex[:12]}", **data) session.add(assistant) await session.flush() await _sync_bindings(session, assistant.id, resource_ids) await _sync_tool_bindings(session, assistant.id, assistant.type, tool_ids) await session.commit() await session.refresh(assistant) return await _to_out(session, assistant) @router.get("/{assistant_id}", response_model=AssistantOut) async def get_assistant( assistant_id: str, session: AsyncSession = Depends(get_session) ): assistant = await session.get(Assistant, assistant_id) if not assistant: raise HTTPException(404, "助手不存在") return await _to_out(session, assistant) @router.post("/{assistant_id}/duplicate", response_model=AssistantOut) async def duplicate_assistant( assistant_id: str, session: AsyncSession = Depends(get_session) ): source = await session.get(Assistant, assistant_id) if not source: raise HTTPException(404, "助手不存在") assistant = Assistant( id=f"asst_{uuid.uuid4().hex[:12]}", name=f"{source.name} 副本", type=source.type, runtime_mode=source.runtime_mode, greeting=source.greeting, enable_interrupt=source.enable_interrupt, turn_config=dict(source.turn_config or {}), startup=dict(source.startup or {}), vision_enabled=source.vision_enabled, vision_model_resource_id=source.vision_model_resource_id, knowledge_base_id=source.knowledge_base_id, knowledge_retrieval_config=dict(source.knowledge_retrieval_config or {}), prompt=source.prompt, dynamic_variable_definitions=dict(source.dynamic_variable_definitions or {}), api_url=source.api_url, api_key=source.api_key, app_id=source.app_id, graph=dict(source.graph or {}), ) session.add(assistant) await session.flush() await _sync_bindings(session, assistant.id, await _resource_ids(session, source.id)) await _sync_tool_bindings( session, assistant.id, assistant.type, await _tool_ids(session, source.id) ) await session.commit() await session.refresh(assistant) return await _to_out(session, assistant) @router.put("/{assistant_id}", response_model=AssistantOut) async def update_assistant( assistant_id: str, body: AssistantUpsert, session: AsyncSession = Depends(get_session), ): assistant = await session.get(Assistant, assistant_id) if not assistant: raise HTTPException(404, "助手不存在") _validate_workflow(body) await _validate_workflow_references(session, body) await _validate_startup_actions(session, body) await _validate_vision_model(session, body) await _validate_knowledge_base(session, body) data = body.model_dump() resource_ids = data.pop("model_resource_ids") tool_ids = data.pop("tool_ids") data["api_key"] = resolve_incoming_key(data["api_key"], assistant.api_key) for key, value in data.items(): setattr(assistant, key, value) await _sync_bindings(session, assistant.id, resource_ids) await _sync_tool_bindings(session, assistant.id, assistant.type, tool_ids) await session.commit() await session.refresh(assistant) return await _to_out(session, assistant) @router.delete("/{assistant_id}") async def delete_assistant( assistant_id: str, session: AsyncSession = Depends(get_session) ): assistant = await session.get(Assistant, assistant_id) if not assistant: raise HTTPException(404, "助手不存在") await session.delete(assistant) await session.commit() return {"ok": True}