Files
NotesAgentic/backend/app/agent/runtime.py

479 lines
19 KiB
Python

"""Agent 运行时:负责模型轮次、工具调用、权限确认与事件发布。"""
from __future__ import annotations
import asyncio
import json
from collections.abc import AsyncIterator
from dataclasses import dataclass, field
from datetime import datetime, timezone
from typing import TYPE_CHECKING
from uuid import uuid4
from app.agent.permissions import PermissionManager, PermissionMode
from app.agent.tools import ToolExecutionContext, ToolNotFoundError, ToolRegistry
from app.contracts import (
AgentEvent,
AgentEventType,
AgentRun,
AgentRunCreateRequest,
AgentRunStatus,
Citation,
Message,
MessageRole,
ModelRequest,
ToolCall,
ToolResult,
)
from app.providers.registry import ProviderRegistry
from app.providers.base import ProviderError
if TYPE_CHECKING:
from app.extensions import AgentConfiguration, SkillRuntime
class AgentRunNotFoundError(LookupError):
pass
class AgentCapacityError(RuntimeError):
pass
TERMINAL_STATUSES = {
AgentRunStatus.completed,
AgentRunStatus.failed,
AgentRunStatus.cancelled,
}
MAX_RUN_RECORDS = 200
MAX_EVENTS_PER_RUN = 2_000
MAX_TOOL_CALLS_PER_TURN = 50
@dataclass(slots=True)
class RunRecord:
"""单次运行的可变上下文,仅由 AgentRuntime 持有。"""
run: AgentRun
request: AgentRunCreateRequest
skill_config: AgentConfiguration | None = None
allowed_tools: list[str] = field(default_factory=list)
events: list[AgentEvent] = field(default_factory=list)
subscribers: set[asyncio.Queue[AgentEvent]] = field(default_factory=set)
task: asyncio.Task[None] | None = None
class AgentRuntime:
"""进程内 Agent 编排器;对外返回深拷贝,避免调用方修改运行状态。"""
def __init__(
self,
providers: ProviderRegistry,
tools: ToolRegistry,
permissions: PermissionManager,
skills: SkillRuntime | None = None,
) -> None:
self.providers = providers
self.tools = tools
self.permissions = permissions
self.skills = skills
self._records: dict[str, RunRecord] = {}
async def create_run(self, request: AgentRunCreateRequest) -> AgentRun:
self._prune_records()
provider = self.providers.get(request.provider_id)
skill_config = None
if request.skill_id:
if self.skills is None:
raise RuntimeError("Skill Runtime is not configured.")
skill_config = self.skills.build_agent_configuration(
request.skill_id, provider.config.capabilities
)
now = datetime.now(timezone.utc)
run = AgentRun(
run_id=f"run_{uuid4().hex}",
status=AgentRunStatus.queued,
input=request.input,
provider_id=request.provider_id,
model=request.model,
skill_id=request.skill_id,
max_steps=request.max_steps,
token_budget=request.token_budget,
created_at=now,
updated_at=now,
)
allowed_tools = list(request.allowed_tools)
if skill_config is not None:
# 同时指定 Skill 与工具白名单时取交集,避免 Skill 扩大调用权限。
allowed_tools = (
[name for name in skill_config.allowed_tools if name in allowed_tools]
if allowed_tools
else list(skill_config.allowed_tools)
)
record = RunRecord(
run=run,
request=request,
skill_config=skill_config,
allowed_tools=allowed_tools,
)
self._records[run.run_id] = record
record.task = asyncio.create_task(self._execute(record), name=run.run_id)
return run.model_copy(deep=True)
def get_run(self, run_id: str) -> AgentRun:
return self._get_record(run_id).run.model_copy(deep=True)
def list_runs(self, limit: int, offset: int) -> tuple[list[AgentRun], int]:
records = sorted(
self._records.values(), key=lambda item: item.run.created_at, reverse=True
)
items = [item.run.model_copy(deep=True) for item in records[offset : offset + limit]]
return items, len(records)
async def cancel(self, run_id: str) -> AgentRun:
record = self._get_record(run_id)
if record.run.status in TERMINAL_STATUSES:
return record.run.model_copy(deep=True)
record.run.cancelled = True
record.run.status = AgentRunStatus.cancelled
record.run.updated_at = datetime.now(timezone.utc)
self.permissions.cancel_run(run_id)
self._publish(record, AgentEventType.run_cancelled, {})
if record.task and not record.task.done():
record.task.cancel()
return record.run.model_copy(deep=True)
def resolve_permission(self, run_id: str, request_id: str, decision: str) -> bool:
self._get_record(run_id)
return self.permissions.resolve(run_id, request_id, decision)
async def events(self, run_id: str) -> AsyncIterator[AgentEvent]:
record = self._get_record(run_id)
# 先回放快照再订阅实时事件,使晚加入的 SSE 客户端也能恢复界面状态。
# TODO(agent): 持久化事件并支持 Last-Event-ID,进程重启后仍可续传。
queue: asyncio.Queue[AgentEvent] = asyncio.Queue()
record.subscribers.add(queue)
history = [event.model_copy(deep=True) for event in record.events]
try:
for event in history:
yield event
if record.run.status in TERMINAL_STATUSES:
return
while True:
event = await queue.get()
yield event.model_copy(deep=True)
if event.event in {
AgentEventType.run_completed,
AgentEventType.run_failed,
AgentEventType.run_cancelled,
}:
return
finally:
record.subscribers.discard(queue)
async def wait(self, run_id: str) -> AgentRun:
record = self._get_record(run_id)
if record.task:
try:
await asyncio.shield(record.task)
except asyncio.CancelledError:
pass
return record.run.model_copy(deep=True)
async def _execute(self, record: RunRecord) -> None:
try:
async with asyncio.timeout(record.request.run_timeout_seconds):
await self._run_loop(record)
except asyncio.CancelledError:
if record.run.status != AgentRunStatus.cancelled:
self._finish_cancelled(record)
except TimeoutError:
self._fail(record, "AGENT_TIMEOUT", "Agent run exceeded its timeout.")
except ProviderError as exc:
self._fail(record, exc.code, exc.message)
except Exception as exc:
self._fail(record, "AGENT_FAILED", str(exc))
async def _run_loop(self, record: RunRecord) -> None:
record.run.status = AgentRunStatus.running
record.run.updated_at = datetime.now(timezone.utc)
self._publish(
record,
AgentEventType.run_started,
{"provider_id": record.request.provider_id, "model": record.request.model},
)
messages = [Message(role=MessageRole.user, content=record.request.input)]
allowed_tools = self.tools.definitions(record.allowed_tools)
provider = self.providers.get(record.request.provider_id).adapter
for step in range(1, record.request.max_steps + 1):
record.run.current_step = step
record.run.updated_at = datetime.now(timezone.utc)
turn = await provider.complete(
ModelRequest(
provider_id=record.request.provider_id,
model=record.request.model,
system=(record.skill_config.system_prompt if record.skill_config else None),
messages=messages,
tools=allowed_tools,
metadata=self._request_metadata(record),
)
)
record.run.token_usage += turn.input_tokens + turn.output_tokens
self._publish(
record,
AgentEventType.usage,
{"token_usage": record.run.token_usage},
)
if (
record.request.token_budget is not None
and record.run.token_usage > record.request.token_budget
):
self._fail(record, "TOKEN_BUDGET_EXCEEDED", "Agent token budget exceeded.")
return
if turn.tool_calls:
if len(turn.tool_calls) > MAX_TOOL_CALLS_PER_TURN:
self._fail(
record,
"TOO_MANY_TOOL_CALLS",
f"Provider requested more than {MAX_TOOL_CALLS_PER_TURN} tools in one turn.",
)
return
calls = [
ToolCall(
tool_call_id=item.tool_call_id,
name=item.name,
arguments=item.arguments,
)
for item in turn.tool_calls
]
messages.append(
Message(role=MessageRole.assistant, content=turn.text or "", tool_calls=calls)
)
# 工具可以并发执行,但结果按模型原始调用顺序写回上下文,保证轮次可复现。
semaphore = asyncio.Semaphore(record.request.max_concurrent_tools)
async def execute(call: ToolCall) -> ToolResult:
async with semaphore:
return await self._execute_tool(record, call)
results = await asyncio.gather(*(execute(call) for call in calls))
for call, result in zip(calls, results):
record.run.tool_results.append(result)
self._collect_citations(record, result)
messages.append(
Message(
role=MessageRole.tool,
name=call.name,
tool_call_id=call.tool_call_id,
content=json.dumps(result.model_dump(mode="json"), ensure_ascii=False),
)
)
continue
if turn.text is not None:
record.run.output = turn.text
self._publish(record, AgentEventType.text_delta, {"text": turn.text})
record.run.status = AgentRunStatus.completed
record.run.updated_at = datetime.now(timezone.utc)
self._publish(
record,
AgentEventType.run_completed,
{"output": turn.text, "token_usage": record.run.token_usage},
)
return
self._fail(record, "EMPTY_MODEL_RESPONSE", "Provider returned no text or tool call.")
return
self._fail(record, "MAX_STEPS_EXCEEDED", "Agent reached its maximum step count.")
async def _execute_tool(self, record: RunRecord, call: ToolCall) -> ToolResult:
self._publish(record, AgentEventType.tool_call, call.model_dump(mode="json"))
try:
registered = self.tools.get(call.name)
except ToolNotFoundError:
registered = None
if registered is not None and call.name not in record.allowed_tools:
result = ToolResult(
tool_call_id=call.tool_call_id,
name=call.name,
success=False,
error_code="TOOL_NOT_ALLOWED",
error_message="Tool is not included in allowed_tools.",
)
self._publish(record, AgentEventType.tool_result, result.model_dump(mode="json"))
return result
permission = registered.definition.permission if registered else None
if permission == "network.request" and not record.request.allow_network:
result = ToolResult(
tool_call_id=call.tool_call_id,
name=call.name,
success=False,
error_code="NETWORK_NOT_ALLOWED",
error_message="Agent run does not allow network tools.",
)
self._publish(record, AgentEventType.tool_result, result.model_dump(mode="json"))
return result
mode = self.permissions.mode_for(permission)
if mode == PermissionMode.deny:
result = self._permission_denied(call)
elif mode == PermissionMode.confirm and permission:
# 运行状态必须在等待期间可见,前端才能展示并处理权限确认卡片。
ticket = self.permissions.create_ticket(record.run.run_id, permission)
record.run.status = AgentRunStatus.waiting_permission
self._publish(
record,
AgentEventType.permission_required,
{
"request_id": ticket.request_id,
"permission": permission,
"tool_call": call.model_dump(mode="json"),
},
)
try:
decision = await self.permissions.wait(
ticket, timeout=record.request.tool_timeout_seconds
)
except TimeoutError:
record.run.status = AgentRunStatus.running
result = ToolResult(
tool_call_id=call.tool_call_id,
name=call.name,
success=False,
error_code="PERMISSION_TIMEOUT",
error_message="Tool permission confirmation timed out.",
)
self._publish(
record, AgentEventType.tool_result, result.model_dump(mode="json")
)
return result
record.run.status = AgentRunStatus.running
result = (
await self._invoke_tool(record, call)
if decision in {"allow_once", "allow_session"}
else self._permission_denied(call)
)
else:
result = await self._invoke_tool(record, call)
self._publish(record, AgentEventType.tool_result, result.model_dump(mode="json"))
return result
async def _invoke_tool(self, record: RunRecord, call: ToolCall) -> ToolResult:
try:
return await asyncio.wait_for(
self.tools.execute(call, ToolExecutionContext(run_id=record.run.run_id)),
timeout=record.request.tool_timeout_seconds,
)
except TimeoutError:
return ToolResult(
tool_call_id=call.tool_call_id,
name=call.name,
success=False,
error_code="TOOL_TIMEOUT",
error_message="Tool execution timed out.",
)
@staticmethod
def _permission_denied(call: ToolCall) -> ToolResult:
return ToolResult(
tool_call_id=call.tool_call_id,
name=call.name,
success=False,
error_code="PERMISSION_DENIED",
error_message="Tool permission was denied.",
)
def _finish_cancelled(self, record: RunRecord) -> None:
record.run.cancelled = True
record.run.status = AgentRunStatus.cancelled
record.run.updated_at = datetime.now(timezone.utc)
self._publish(record, AgentEventType.run_cancelled, {})
def _fail(self, record: RunRecord, code: str, message: str) -> None:
if record.run.status in TERMINAL_STATUSES:
return
record.run.status = AgentRunStatus.failed
record.run.error_code = code
record.run.error_message = message
record.run.updated_at = datetime.now(timezone.utc)
self._publish(
record,
AgentEventType.run_failed,
{"code": code, "message": message},
)
def _publish(
self, record: RunRecord, event_type: AgentEventType, data: dict[str, object]
) -> None:
event = AgentEvent(
event=event_type,
run_id=record.run.run_id,
sequence=len(record.events),
data=data,
timestamp=datetime.now(timezone.utc),
)
record.events.append(event)
# 内存事件只保留最近窗口;完整审计轨迹应由后续持久化层承担。
if len(record.events) > MAX_EVENTS_PER_RUN:
del record.events[: len(record.events) - MAX_EVENTS_PER_RUN]
for queue in record.subscribers:
queue.put_nowait(event)
@staticmethod
def _request_metadata(record: RunRecord) -> dict[str, object]:
metadata = dict(record.request.metadata)
if record.skill_config is not None:
metadata["skill_id"] = record.skill_config.skill_id
metadata["retrieval"] = record.skill_config.retrieval.model_dump(mode="json")
return metadata
def _collect_citations(self, record: RunRecord, result: ToolResult) -> None:
if not result.success or not isinstance(result.output, dict):
return
items = result.output.get("items")
if not isinstance(items, list):
return
known = {citation.citation_id for citation in record.run.citations}
for item in items:
if not isinstance(item, dict) or not isinstance(item.get("citation"), dict):
continue
try:
citation = Citation.model_validate(item["citation"])
except ValueError:
continue
if citation.citation_id in known:
continue
known.add(citation.citation_id)
record.run.citations.append(citation)
self._publish(record, AgentEventType.citation, citation.model_dump(mode="json"))
def _get_record(self, run_id: str) -> RunRecord:
try:
return self._records[run_id]
except KeyError as exc:
raise AgentRunNotFoundError(run_id) from exc
def _prune_records(self) -> None:
# 只清理终态记录,绝不为了容量取消仍在执行或等待授权的任务。
overflow = len(self._records) - MAX_RUN_RECORDS + 1
if overflow <= 0:
return
terminal = sorted(
(
record
for record in self._records.values()
if record.run.status in TERMINAL_STATUSES
),
key=lambda record: record.run.updated_at,
)
for record in terminal[:overflow]:
self._records.pop(record.run.run_id, None)
if len(self._records) >= MAX_RUN_RECORDS:
raise AgentCapacityError("Too many active Agent runs.")