实现 AI Core 与 Agent Core 基础功能
- 更新 README 描述从后端壳子到 AI Core/Agent Core - 添加 ToolCall 和 ToolResult 数据结构定义 - 扩展 AgentRun 模型增加输出、错误码、工具调用结果等字段 - 添加 mock 提供商类型支持 - 实现聊天、代理运行、工具调用和提供商管理的核心路由逻辑 - 集成容器化依赖注入和错误处理机制 - 更新 API 接口契约和文档说明
This commit is contained in:
+3
-1
@@ -1,6 +1,6 @@
|
||||
# Backend
|
||||
|
||||
FastAPI + Pydantic 的最小后端壳子。项目使用 uv 管理依赖和虚拟环境。
|
||||
FastAPI + Pydantic 的本地 AI Core / Agent Core。项目使用 uv 管理依赖和虚拟环境。
|
||||
|
||||
```powershell
|
||||
uv sync
|
||||
@@ -15,3 +15,5 @@ uv run uvicorn app.main:app --reload --host 127.0.0.1 --port 8000
|
||||
- API 文档:<http://127.0.0.1:8000/docs>
|
||||
|
||||
团队接口清单见 `../docs/后端接口契约-开发版.md`,机器可读契约以运行时的 `/openapi.json` 为准。
|
||||
|
||||
AI Core 与 Agent Core 的模块边界、Mock Provider 和 Tool Calling 调试方式见 `../docs/AI-Core与Agent-Core开发说明.md`。
|
||||
|
||||
@@ -0,0 +1,12 @@
|
||||
from app.agent.permissions import PermissionManager, PermissionMode, PermissionPolicy
|
||||
from app.agent.runtime import AgentRuntime, AgentRunNotFoundError
|
||||
from app.agent.tools import ToolRegistry
|
||||
|
||||
__all__ = [
|
||||
"AgentRunNotFoundError",
|
||||
"AgentRuntime",
|
||||
"PermissionManager",
|
||||
"PermissionMode",
|
||||
"PermissionPolicy",
|
||||
"ToolRegistry",
|
||||
]
|
||||
@@ -0,0 +1,46 @@
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
|
||||
from app.agent.tools import ToolExecutionContext, ToolRegistry
|
||||
from app.contracts import ToolDefinition
|
||||
|
||||
|
||||
class ToolArguments(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
|
||||
class EchoArguments(ToolArguments):
|
||||
text: str
|
||||
|
||||
|
||||
class AddArguments(ToolArguments):
|
||||
left: float
|
||||
right: float
|
||||
|
||||
|
||||
async def echo(arguments: EchoArguments, _: ToolExecutionContext) -> dict[str, str]:
|
||||
return {"text": arguments.text}
|
||||
|
||||
|
||||
async def add(arguments: AddArguments, _: ToolExecutionContext) -> dict[str, float]:
|
||||
return {"value": arguments.left + arguments.right}
|
||||
|
||||
|
||||
def register_builtin_tools(registry: ToolRegistry) -> None:
|
||||
registry.register(
|
||||
ToolDefinition(
|
||||
name="system.echo",
|
||||
description="Echo text for local Agent integration testing.",
|
||||
parameters=EchoArguments.model_json_schema(),
|
||||
),
|
||||
EchoArguments,
|
||||
echo,
|
||||
)
|
||||
registry.register(
|
||||
ToolDefinition(
|
||||
name="math.add",
|
||||
description="Add two numbers without external side effects.",
|
||||
parameters=AddArguments.model_json_schema(),
|
||||
),
|
||||
AddArguments,
|
||||
add,
|
||||
)
|
||||
@@ -0,0 +1,80 @@
|
||||
import asyncio
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum
|
||||
from uuid import uuid4
|
||||
|
||||
|
||||
class PermissionMode(str, Enum):
|
||||
allow = "allow"
|
||||
confirm = "confirm"
|
||||
deny = "deny"
|
||||
|
||||
|
||||
class PermissionPolicy:
|
||||
def __init__(self) -> None:
|
||||
self._rules: dict[str, PermissionMode] = {
|
||||
"notes.delete": PermissionMode.confirm,
|
||||
"notes.write": PermissionMode.confirm,
|
||||
"network.request": PermissionMode.confirm,
|
||||
"secrets.use": PermissionMode.confirm,
|
||||
}
|
||||
|
||||
def set_rule(self, permission: str, mode: PermissionMode) -> None:
|
||||
self._rules[permission] = mode
|
||||
|
||||
def mode_for(self, permission: str | None) -> PermissionMode:
|
||||
if permission is None:
|
||||
return PermissionMode.allow
|
||||
return self._rules.get(permission, PermissionMode.allow)
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class PermissionTicket:
|
||||
request_id: str
|
||||
run_id: str
|
||||
permission: str
|
||||
future: asyncio.Future[str]
|
||||
|
||||
|
||||
class PermissionManager:
|
||||
def __init__(self, policy: PermissionPolicy) -> None:
|
||||
self.policy = policy
|
||||
self._pending: dict[tuple[str, str], PermissionTicket] = {}
|
||||
self._session_grants: set[str] = set()
|
||||
|
||||
def mode_for(self, permission: str | None) -> PermissionMode:
|
||||
if permission in self._session_grants:
|
||||
return PermissionMode.allow
|
||||
return self.policy.mode_for(permission)
|
||||
|
||||
def create_ticket(self, run_id: str, permission: str) -> PermissionTicket:
|
||||
ticket = PermissionTicket(
|
||||
request_id=f"permission_{uuid4().hex}",
|
||||
run_id=run_id,
|
||||
permission=permission,
|
||||
future=asyncio.get_running_loop().create_future(),
|
||||
)
|
||||
self._pending[(run_id, ticket.request_id)] = ticket
|
||||
return ticket
|
||||
|
||||
async def wait(self, ticket: PermissionTicket, timeout: float) -> str:
|
||||
try:
|
||||
return await asyncio.wait_for(ticket.future, timeout=timeout)
|
||||
finally:
|
||||
self._pending.pop((ticket.run_id, ticket.request_id), None)
|
||||
|
||||
def resolve(self, run_id: str, request_id: str, decision: str) -> bool:
|
||||
ticket = self._pending.get((run_id, request_id))
|
||||
if ticket is None or ticket.future.done():
|
||||
return False
|
||||
if decision == "allow_session":
|
||||
self._session_grants.add(ticket.permission)
|
||||
ticket.future.set_result(decision)
|
||||
return True
|
||||
|
||||
def cancel_run(self, run_id: str) -> None:
|
||||
for key, ticket in list(self._pending.items()):
|
||||
if ticket.run_id == run_id:
|
||||
if not ticket.future.done():
|
||||
ticket.future.cancel()
|
||||
self._pending.pop(key, None)
|
||||
@@ -0,0 +1,346 @@
|
||||
import asyncio
|
||||
import json
|
||||
from collections.abc import AsyncIterator
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime, timezone
|
||||
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,
|
||||
Message,
|
||||
MessageRole,
|
||||
ModelRequest,
|
||||
ToolCall,
|
||||
ToolResult,
|
||||
)
|
||||
from app.providers.registry import ProviderRegistry
|
||||
|
||||
|
||||
class AgentRunNotFoundError(LookupError):
|
||||
pass
|
||||
|
||||
|
||||
TERMINAL_STATUSES = {
|
||||
AgentRunStatus.completed,
|
||||
AgentRunStatus.failed,
|
||||
AgentRunStatus.cancelled,
|
||||
}
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class RunRecord:
|
||||
run: AgentRun
|
||||
request: AgentRunCreateRequest
|
||||
events: list[AgentEvent] = field(default_factory=list)
|
||||
subscribers: set[asyncio.Queue[AgentEvent]] = field(default_factory=set)
|
||||
task: asyncio.Task[None] | None = None
|
||||
|
||||
|
||||
class AgentRuntime:
|
||||
def __init__(
|
||||
self,
|
||||
providers: ProviderRegistry,
|
||||
tools: ToolRegistry,
|
||||
permissions: PermissionManager,
|
||||
) -> None:
|
||||
self.providers = providers
|
||||
self.tools = tools
|
||||
self.permissions = permissions
|
||||
self._records: dict[str, RunRecord] = {}
|
||||
|
||||
async def create_run(self, request: AgentRunCreateRequest) -> AgentRun:
|
||||
self.providers.get(request.provider_id)
|
||||
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,
|
||||
)
|
||||
record = RunRecord(run=run, request=request)
|
||||
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)
|
||||
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 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.request.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,
|
||||
messages=messages,
|
||||
tools=allowed_tools,
|
||||
metadata=record.request.metadata,
|
||||
)
|
||||
)
|
||||
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:
|
||||
for provider_call in turn.tool_calls:
|
||||
call = ToolCall(
|
||||
tool_call_id=provider_call.tool_call_id,
|
||||
name=provider_call.name,
|
||||
arguments=provider_call.arguments,
|
||||
)
|
||||
result = await self._execute_tool(record, call)
|
||||
record.run.tool_results.append(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.request.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
|
||||
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)
|
||||
for queue in record.subscribers:
|
||||
queue.put_nowait(event)
|
||||
|
||||
def _get_record(self, run_id: str) -> RunRecord:
|
||||
try:
|
||||
return self._records[run_id]
|
||||
except KeyError as exc:
|
||||
raise AgentRunNotFoundError(run_id) from exc
|
||||
@@ -0,0 +1,108 @@
|
||||
import inspect
|
||||
from dataclasses import dataclass
|
||||
from time import perf_counter
|
||||
from typing import Any, Awaitable, Callable
|
||||
|
||||
from pydantic import BaseModel, ValidationError
|
||||
|
||||
from app.contracts import ToolCall, ToolDefinition, ToolResult
|
||||
|
||||
ToolExecutor = Callable[[BaseModel, "ToolExecutionContext"], Any | Awaitable[Any]]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ToolExecutionContext:
|
||||
run_id: str
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class RegisteredTool:
|
||||
definition: ToolDefinition
|
||||
arguments_model: type[BaseModel]
|
||||
executor: ToolExecutor
|
||||
|
||||
|
||||
class ToolNotFoundError(LookupError):
|
||||
pass
|
||||
|
||||
|
||||
class ToolRegistry:
|
||||
def __init__(self) -> None:
|
||||
self._tools: dict[str, RegisteredTool] = {}
|
||||
|
||||
def register(
|
||||
self,
|
||||
definition: ToolDefinition,
|
||||
arguments_model: type[BaseModel],
|
||||
executor: ToolExecutor,
|
||||
) -> None:
|
||||
if definition.name in self._tools:
|
||||
raise ValueError(f"Tool already registered: {definition.name}")
|
||||
self._tools[definition.name] = RegisteredTool(
|
||||
definition=definition,
|
||||
arguments_model=arguments_model,
|
||||
executor=executor,
|
||||
)
|
||||
|
||||
def unregister(self, name: str) -> None:
|
||||
self._tools.pop(name, None)
|
||||
|
||||
def get(self, name: str) -> RegisteredTool:
|
||||
try:
|
||||
return self._tools[name]
|
||||
except KeyError as exc:
|
||||
raise ToolNotFoundError(name) from exc
|
||||
|
||||
def definitions(self, allowed: list[str] | None = None) -> list[ToolDefinition]:
|
||||
names = set(allowed) if allowed is not None else None
|
||||
return [
|
||||
item.definition.model_copy(deep=True)
|
||||
for name, item in self._tools.items()
|
||||
if names is None or name in names
|
||||
]
|
||||
|
||||
async def execute(self, call: ToolCall, context: ToolExecutionContext) -> ToolResult:
|
||||
started = perf_counter()
|
||||
try:
|
||||
registered = self.get(call.name)
|
||||
except ToolNotFoundError:
|
||||
return ToolResult(
|
||||
tool_call_id=call.tool_call_id,
|
||||
name=call.name,
|
||||
success=False,
|
||||
error_code="TOOL_NOT_FOUND",
|
||||
error_message=f"Tool is not registered: {call.name}",
|
||||
)
|
||||
|
||||
try:
|
||||
arguments = registered.arguments_model.model_validate(call.arguments)
|
||||
except ValidationError as exc:
|
||||
return ToolResult(
|
||||
tool_call_id=call.tool_call_id,
|
||||
name=call.name,
|
||||
success=False,
|
||||
error_code="TOOL_ARGUMENT_INVALID",
|
||||
error_message=str(exc),
|
||||
duration_ms=round((perf_counter() - started) * 1000),
|
||||
)
|
||||
|
||||
try:
|
||||
output = registered.executor(arguments, context)
|
||||
if inspect.isawaitable(output):
|
||||
output = await output
|
||||
return ToolResult(
|
||||
tool_call_id=call.tool_call_id,
|
||||
name=call.name,
|
||||
success=True,
|
||||
output=output,
|
||||
duration_ms=round((perf_counter() - started) * 1000),
|
||||
)
|
||||
except Exception as exc: # Tool failures are isolated from the Agent loop.
|
||||
return ToolResult(
|
||||
tool_call_id=call.tool_call_id,
|
||||
name=call.name,
|
||||
success=False,
|
||||
error_code="TOOL_EXECUTION_FAILED",
|
||||
error_message=str(exc),
|
||||
duration_ms=round((perf_counter() - started) * 1000),
|
||||
)
|
||||
@@ -0,0 +1,49 @@
|
||||
from dataclasses import dataclass
|
||||
|
||||
from app.agent import AgentRuntime, PermissionManager, PermissionPolicy, ToolRegistry
|
||||
from app.agent.builtin_tools import register_builtin_tools
|
||||
from app.contracts import ModelCapability, ProviderConfig, ProviderType
|
||||
from app.providers import MockProvider, ProviderRegistry
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class ApplicationContainer:
|
||||
providers: ProviderRegistry
|
||||
tools: ToolRegistry
|
||||
permissions: PermissionManager
|
||||
agent: AgentRuntime
|
||||
|
||||
|
||||
def build_container() -> ApplicationContainer:
|
||||
providers = ProviderRegistry()
|
||||
providers.register(
|
||||
ProviderConfig(
|
||||
provider_id="mock",
|
||||
provider_type=ProviderType.mock,
|
||||
name="Mock Provider",
|
||||
default_model="mock-1",
|
||||
enabled=True,
|
||||
capabilities=[
|
||||
ModelCapability.chat,
|
||||
ModelCapability.tool_calling,
|
||||
ModelCapability.streaming,
|
||||
],
|
||||
),
|
||||
MockProvider(),
|
||||
)
|
||||
|
||||
tools = ToolRegistry()
|
||||
register_builtin_tools(tools)
|
||||
|
||||
policy = PermissionPolicy()
|
||||
permissions = PermissionManager(policy)
|
||||
agent = AgentRuntime(providers=providers, tools=tools, permissions=permissions)
|
||||
return ApplicationContainer(
|
||||
providers=providers,
|
||||
tools=tools,
|
||||
permissions=permissions,
|
||||
agent=agent,
|
||||
)
|
||||
|
||||
|
||||
container = build_container()
|
||||
@@ -155,6 +155,22 @@ class ToolDefinition(Contract):
|
||||
source: Literal["builtin", "plugin"] = "builtin"
|
||||
|
||||
|
||||
class ToolCall(Contract):
|
||||
tool_call_id: str
|
||||
name: str
|
||||
arguments: dict[str, Any] = Field(default_factory=dict)
|
||||
|
||||
|
||||
class ToolResult(Contract):
|
||||
tool_call_id: str
|
||||
name: str
|
||||
success: bool
|
||||
output: Any | None = None
|
||||
error_code: str | None = None
|
||||
error_message: str | None = None
|
||||
duration_ms: int | None = None
|
||||
|
||||
|
||||
class ToolListResponse(Contract):
|
||||
items: list[ToolDefinition] = Field(default_factory=list)
|
||||
|
||||
@@ -242,6 +258,11 @@ class AgentRun(Contract):
|
||||
max_steps: int
|
||||
token_budget: int | None = None
|
||||
cancelled: bool = False
|
||||
output: str | None = None
|
||||
error_code: str | None = None
|
||||
error_message: str | None = None
|
||||
token_usage: int = 0
|
||||
tool_results: list[ToolResult] = Field(default_factory=list)
|
||||
citations: list[Citation] = Field(default_factory=list)
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
@@ -371,6 +392,7 @@ class PluginListResponse(Contract):
|
||||
|
||||
# Providers
|
||||
class ProviderType(str, Enum):
|
||||
mock = "mock"
|
||||
openai_responses = "openai_responses"
|
||||
openai_chat = "openai_chat"
|
||||
openai_compatible = "openai_compatible"
|
||||
|
||||
+1
-1
@@ -13,7 +13,7 @@ settings = get_settings()
|
||||
app = FastAPI(
|
||||
title=settings.name,
|
||||
version=settings.version,
|
||||
description="AI 笔记软件的本地 FastAPI 服务壳子。",
|
||||
description="AI 笔记软件的本地 AI Core 与 Agent Core 服务。",
|
||||
)
|
||||
|
||||
app.add_middleware(
|
||||
|
||||
@@ -0,0 +1,11 @@
|
||||
from app.providers.base import ModelProvider, ProviderToolCall, ProviderTurn
|
||||
from app.providers.mock import MockProvider
|
||||
from app.providers.registry import ProviderRegistry
|
||||
|
||||
__all__ = [
|
||||
"MockProvider",
|
||||
"ModelProvider",
|
||||
"ProviderRegistry",
|
||||
"ProviderToolCall",
|
||||
"ProviderTurn",
|
||||
]
|
||||
@@ -0,0 +1,30 @@
|
||||
from collections.abc import AsyncIterator
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Protocol
|
||||
|
||||
from app.contracts import ModelEvent, ModelInfo, ModelRequest
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class ProviderToolCall:
|
||||
tool_call_id: str
|
||||
name: str
|
||||
arguments: dict[str, object] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class ProviderTurn:
|
||||
text: str | None = None
|
||||
tool_calls: list[ProviderToolCall] = field(default_factory=list)
|
||||
input_tokens: int = 0
|
||||
output_tokens: int = 0
|
||||
|
||||
|
||||
class ModelProvider(Protocol):
|
||||
async def complete(self, request: ModelRequest) -> ProviderTurn: ...
|
||||
|
||||
def stream(self, request: ModelRequest) -> AsyncIterator[ModelEvent]: ...
|
||||
|
||||
async def list_models(self) -> list[ModelInfo]: ...
|
||||
|
||||
async def test_connection(self, model: str | None = None) -> tuple[bool, str]: ...
|
||||
@@ -0,0 +1,127 @@
|
||||
import json
|
||||
import re
|
||||
from collections.abc import AsyncIterator
|
||||
from datetime import datetime, timezone
|
||||
from uuid import uuid4
|
||||
|
||||
from app.contracts import (
|
||||
MessageRole,
|
||||
ModelCapability,
|
||||
ModelEvent,
|
||||
ModelEventType,
|
||||
ModelInfo,
|
||||
ModelRequest,
|
||||
)
|
||||
from app.providers.base import ProviderToolCall, ProviderTurn
|
||||
|
||||
_TOOL_PATTERN = re.compile(r"^/tool\s+([\w.-]+)(?:\s+(\{.*\}))?\s*$", re.DOTALL)
|
||||
|
||||
|
||||
class MockProvider:
|
||||
"""离线开发 Provider,用于验证聊天、Tool Calling 和 Agent Loop。"""
|
||||
|
||||
async def complete(self, request: ModelRequest) -> ProviderTurn:
|
||||
if not request.messages:
|
||||
return ProviderTurn(text="Mock provider received an empty conversation.")
|
||||
|
||||
last_message = request.messages[-1]
|
||||
if last_message.role == MessageRole.tool:
|
||||
return ProviderTurn(
|
||||
text=f"Tool result received: {last_message.content}",
|
||||
input_tokens=len(last_message.content.split()),
|
||||
output_tokens=4,
|
||||
)
|
||||
|
||||
match = _TOOL_PATTERN.match(last_message.content.strip())
|
||||
if match:
|
||||
raw_arguments = match.group(2) or "{}"
|
||||
try:
|
||||
arguments = json.loads(raw_arguments)
|
||||
except json.JSONDecodeError:
|
||||
return ProviderTurn(text="Mock tool arguments must be valid JSON.")
|
||||
if not isinstance(arguments, dict):
|
||||
return ProviderTurn(text="Mock tool arguments must be a JSON object.")
|
||||
return ProviderTurn(
|
||||
tool_calls=[
|
||||
ProviderToolCall(
|
||||
tool_call_id=f"call_{uuid4().hex}",
|
||||
name=match.group(1),
|
||||
arguments=arguments,
|
||||
)
|
||||
],
|
||||
input_tokens=len(last_message.content.split()),
|
||||
)
|
||||
|
||||
text = f"Mock response: {last_message.content}"
|
||||
return ProviderTurn(
|
||||
text=text,
|
||||
input_tokens=len(last_message.content.split()),
|
||||
output_tokens=len(text.split()),
|
||||
)
|
||||
|
||||
async def stream(self, request: ModelRequest) -> AsyncIterator[ModelEvent]:
|
||||
turn = await self.complete(request)
|
||||
sequence = 0
|
||||
if turn.tool_calls:
|
||||
for call in turn.tool_calls:
|
||||
yield ModelEvent(
|
||||
event=ModelEventType.tool_call_start,
|
||||
sequence=sequence,
|
||||
data={
|
||||
"tool_call_id": call.tool_call_id,
|
||||
"name": call.name,
|
||||
"arguments": call.arguments,
|
||||
},
|
||||
timestamp=datetime.now(timezone.utc),
|
||||
)
|
||||
sequence += 1
|
||||
yield ModelEvent(
|
||||
event=ModelEventType.tool_call_end,
|
||||
sequence=sequence,
|
||||
data={"tool_call_id": call.tool_call_id},
|
||||
timestamp=datetime.now(timezone.utc),
|
||||
)
|
||||
sequence += 1
|
||||
elif turn.text:
|
||||
words = turn.text.split(" ")
|
||||
for index, word in enumerate(words):
|
||||
yield ModelEvent(
|
||||
event=ModelEventType.text_delta,
|
||||
sequence=sequence,
|
||||
data={"text": word + (" " if index < len(words) - 1 else "")},
|
||||
timestamp=datetime.now(timezone.utc),
|
||||
)
|
||||
sequence += 1
|
||||
|
||||
yield ModelEvent(
|
||||
event=ModelEventType.usage,
|
||||
sequence=sequence,
|
||||
data={
|
||||
"input_tokens": turn.input_tokens,
|
||||
"output_tokens": turn.output_tokens,
|
||||
},
|
||||
timestamp=datetime.now(timezone.utc),
|
||||
)
|
||||
yield ModelEvent(
|
||||
event=ModelEventType.done,
|
||||
sequence=sequence + 1,
|
||||
timestamp=datetime.now(timezone.utc),
|
||||
)
|
||||
|
||||
async def list_models(self) -> list[ModelInfo]:
|
||||
return [
|
||||
ModelInfo(
|
||||
model="mock-1",
|
||||
display_name="Mock Provider (Development)",
|
||||
capabilities=[
|
||||
ModelCapability.chat,
|
||||
ModelCapability.tool_calling,
|
||||
ModelCapability.streaming,
|
||||
],
|
||||
)
|
||||
]
|
||||
|
||||
async def test_connection(self, model: str | None = None) -> tuple[bool, str]:
|
||||
if model not in (None, "mock-1"):
|
||||
return False, f"Unknown mock model: {model}"
|
||||
return True, "Mock provider is ready."
|
||||
@@ -0,0 +1,55 @@
|
||||
from dataclasses import dataclass
|
||||
from time import perf_counter
|
||||
|
||||
from app.contracts import ModelInfo, ProviderConfig, ProviderTestResponse
|
||||
from app.providers.base import ModelProvider
|
||||
|
||||
|
||||
class ProviderNotFoundError(LookupError):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class RegisteredProvider:
|
||||
config: ProviderConfig
|
||||
adapter: ModelProvider
|
||||
|
||||
|
||||
class ProviderRegistry:
|
||||
def __init__(self) -> None:
|
||||
self._providers: dict[str, RegisteredProvider] = {}
|
||||
|
||||
def register(self, config: ProviderConfig, adapter: ModelProvider) -> None:
|
||||
if config.provider_id in self._providers:
|
||||
raise ValueError(f"Provider already registered: {config.provider_id}")
|
||||
self._providers[config.provider_id] = RegisteredProvider(config=config, adapter=adapter)
|
||||
|
||||
def unregister(self, provider_id: str) -> None:
|
||||
self._providers.pop(provider_id, None)
|
||||
|
||||
def get(self, provider_id: str) -> RegisteredProvider:
|
||||
try:
|
||||
provider = self._providers[provider_id]
|
||||
except KeyError as exc:
|
||||
raise ProviderNotFoundError(provider_id) from exc
|
||||
if not provider.config.enabled:
|
||||
raise ProviderNotFoundError(provider_id)
|
||||
return provider
|
||||
|
||||
def list_configs(self) -> list[ProviderConfig]:
|
||||
return [item.config.model_copy(deep=True) for item in self._providers.values()]
|
||||
|
||||
async def list_models(self, provider_id: str) -> list[ModelInfo]:
|
||||
return await self.get(provider_id).adapter.list_models()
|
||||
|
||||
async def test(self, provider_id: str, model: str | None = None) -> ProviderTestResponse:
|
||||
provider = self.get(provider_id)
|
||||
started = perf_counter()
|
||||
success, message = await provider.adapter.test_connection(model)
|
||||
latency_ms = round((perf_counter() - started) * 1000)
|
||||
return ProviderTestResponse(
|
||||
provider_id=provider_id,
|
||||
success=success,
|
||||
latency_ms=latency_ms,
|
||||
message=message,
|
||||
)
|
||||
+86
-33
@@ -5,8 +5,6 @@ from fastapi import APIRouter, Query
|
||||
from fastapi.responses import StreamingResponse
|
||||
|
||||
from app.contracts import (
|
||||
AgentEvent,
|
||||
AgentEventType,
|
||||
AgentRun,
|
||||
AgentRunCreateRequest,
|
||||
AgentRunListResponse,
|
||||
@@ -47,7 +45,10 @@ from app.contracts import (
|
||||
TranscriptionJob,
|
||||
TranscriptionRequest,
|
||||
)
|
||||
from app.errors import not_implemented
|
||||
from app.agent import AgentRunNotFoundError
|
||||
from app.container import container
|
||||
from app.errors import ApiError, not_implemented
|
||||
from app.providers.registry import ProviderNotFoundError
|
||||
|
||||
router = APIRouter(prefix="/api")
|
||||
not_implemented_response = {501: {"model": ErrorResponse, "description": "业务服务尚未实现"}}
|
||||
@@ -61,6 +62,30 @@ def as_sse(event: str, payload: str) -> str:
|
||||
return f"event: {event}\ndata: {payload}\n\n"
|
||||
|
||||
|
||||
def provider_or_404(provider_id: str):
|
||||
try:
|
||||
return container.providers.get(provider_id)
|
||||
except ProviderNotFoundError as exc:
|
||||
raise ApiError(
|
||||
404,
|
||||
"PROVIDER_NOT_FOUND",
|
||||
f"Provider is not registered or enabled: {provider_id}",
|
||||
{"provider_id": provider_id},
|
||||
) from exc
|
||||
|
||||
|
||||
def agent_run_or_404(run_id: str) -> AgentRun:
|
||||
try:
|
||||
return container.agent.get_run(run_id)
|
||||
except AgentRunNotFoundError as exc:
|
||||
raise ApiError(
|
||||
404,
|
||||
"AGENT_RUN_NOT_FOUND",
|
||||
f"Agent run does not exist: {run_id}",
|
||||
{"run_id": run_id},
|
||||
) from exc
|
||||
|
||||
|
||||
# Notes
|
||||
@router.get("/notes", response_model=NoteListResponse, tags=["Notes"])
|
||||
async def list_notes(
|
||||
@@ -131,16 +156,22 @@ async def search_notes(request: SearchRequest) -> SearchResponse:
|
||||
},
|
||||
tags=["Chat"],
|
||||
)
|
||||
async def chat(_: ChatRequest) -> StreamingResponse:
|
||||
async def chat(request: ChatRequest) -> StreamingResponse:
|
||||
provider = provider_or_404(request.provider_id)
|
||||
|
||||
async def stream() -> AsyncIterator[str]:
|
||||
error = ModelEvent(
|
||||
event=ModelEventType.error,
|
||||
data={"code": "NOT_IMPLEMENTED", "message": "Chat runtime is not implemented."},
|
||||
timestamp=utc_now(),
|
||||
)
|
||||
done = ModelEvent(event=ModelEventType.done, sequence=1, timestamp=utc_now())
|
||||
yield as_sse(error.event.value, error.model_dump_json())
|
||||
yield as_sse(done.event.value, done.model_dump_json())
|
||||
try:
|
||||
async for event in provider.adapter.stream(request):
|
||||
yield as_sse(event.event.value, event.model_dump_json())
|
||||
except Exception as exc:
|
||||
error = ModelEvent(
|
||||
event=ModelEventType.error,
|
||||
data={"code": "PROVIDER_ERROR", "message": str(exc)},
|
||||
timestamp=utc_now(),
|
||||
)
|
||||
done = ModelEvent(event=ModelEventType.done, sequence=1, timestamp=utc_now())
|
||||
yield as_sse(error.event.value, error.model_dump_json())
|
||||
yield as_sse(done.event.value, done.model_dump_json())
|
||||
|
||||
return StreamingResponse(stream(), media_type="text/event-stream")
|
||||
|
||||
@@ -150,7 +181,11 @@ async def chat(_: ChatRequest) -> StreamingResponse:
|
||||
async def list_agent_runs(
|
||||
limit: int = Query(default=50, ge=1, le=100), offset: int = Query(default=0, ge=0)
|
||||
) -> AgentRunListResponse:
|
||||
return AgentRunListResponse(page=PageMeta(limit=limit, offset=offset))
|
||||
items, total = container.agent.list_runs(limit=limit, offset=offset)
|
||||
return AgentRunListResponse(
|
||||
items=items,
|
||||
page=PageMeta(total=total, limit=limit, offset=offset),
|
||||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
@@ -160,8 +195,9 @@ async def list_agent_runs(
|
||||
responses=not_implemented_response,
|
||||
tags=["Agent"],
|
||||
)
|
||||
async def create_agent_run(_: AgentRunCreateRequest) -> AgentRun:
|
||||
not_implemented("agent.runs.create")
|
||||
async def create_agent_run(request: AgentRunCreateRequest) -> AgentRun:
|
||||
provider_or_404(request.provider_id)
|
||||
return await container.agent.create_run(request)
|
||||
|
||||
|
||||
@router.get(
|
||||
@@ -171,7 +207,7 @@ async def create_agent_run(_: AgentRunCreateRequest) -> AgentRun:
|
||||
tags=["Agent"],
|
||||
)
|
||||
async def get_agent_run(run_id: str) -> AgentRun:
|
||||
not_implemented(f"agent.runs.read:{run_id}")
|
||||
return agent_run_or_404(run_id)
|
||||
|
||||
|
||||
@router.post(
|
||||
@@ -181,7 +217,13 @@ async def get_agent_run(run_id: str) -> AgentRun:
|
||||
tags=["Agent"],
|
||||
)
|
||||
async def cancel_agent_run(run_id: str) -> OperationResponse:
|
||||
not_implemented(f"agent.runs.cancel:{run_id}")
|
||||
agent_run_or_404(run_id)
|
||||
run = await container.agent.cancel(run_id)
|
||||
return OperationResponse(
|
||||
status="completed",
|
||||
resource_id=run.run_id,
|
||||
message=f"Agent run status: {run.status.value}",
|
||||
)
|
||||
|
||||
|
||||
@router.get(
|
||||
@@ -196,15 +238,11 @@ async def cancel_agent_run(run_id: str) -> OperationResponse:
|
||||
tags=["Agent"],
|
||||
)
|
||||
async def agent_events(run_id: str) -> StreamingResponse:
|
||||
agent_run_or_404(run_id)
|
||||
|
||||
async def stream() -> AsyncIterator[str]:
|
||||
event = AgentEvent(
|
||||
event=AgentEventType.run_failed,
|
||||
run_id=run_id,
|
||||
sequence=0,
|
||||
data={"code": "NOT_IMPLEMENTED", "message": "Agent runtime is not implemented."},
|
||||
timestamp=utc_now(),
|
||||
)
|
||||
yield as_sse(event.event.value, event.model_dump_json())
|
||||
async for event in container.agent.events(run_id):
|
||||
yield as_sse(event.event.value, event.model_dump_json())
|
||||
|
||||
return StreamingResponse(stream(), media_type="text/event-stream")
|
||||
|
||||
@@ -216,14 +254,24 @@ async def agent_events(run_id: str) -> StreamingResponse:
|
||||
tags=["Agent"],
|
||||
)
|
||||
async def decide_agent_permission(
|
||||
run_id: str, request_id: str, _: PermissionDecisionRequest
|
||||
run_id: str, request_id: str, request: PermissionDecisionRequest
|
||||
) -> OperationResponse:
|
||||
not_implemented(f"agent.permissions:{run_id}:{request_id}")
|
||||
agent_run_or_404(run_id)
|
||||
if not container.agent.resolve_permission(run_id, request_id, request.decision):
|
||||
raise ApiError(
|
||||
404,
|
||||
"PERMISSION_REQUEST_NOT_FOUND",
|
||||
"Permission request does not exist or has already been resolved.",
|
||||
{"run_id": run_id, "request_id": request_id},
|
||||
)
|
||||
return OperationResponse(
|
||||
status="completed", resource_id=request_id, message=request.decision
|
||||
)
|
||||
|
||||
|
||||
@router.get("/tools", response_model=ToolListResponse, tags=["Agent"])
|
||||
async def list_tools() -> ToolListResponse:
|
||||
return ToolListResponse()
|
||||
return ToolListResponse(items=container.tools.definitions())
|
||||
|
||||
|
||||
# Skills
|
||||
@@ -340,7 +388,7 @@ async def uninstall_plugin(plugin_id: str) -> OperationResponse:
|
||||
# Providers
|
||||
@router.get("/providers", response_model=ProviderListResponse, tags=["Providers"])
|
||||
async def list_providers() -> ProviderListResponse:
|
||||
return ProviderListResponse()
|
||||
return ProviderListResponse(items=container.providers.list_configs())
|
||||
|
||||
|
||||
@router.get(
|
||||
@@ -350,7 +398,7 @@ async def list_providers() -> ProviderListResponse:
|
||||
tags=["Providers"],
|
||||
)
|
||||
async def get_provider(provider_id: str) -> ProviderConfig:
|
||||
not_implemented(f"providers.read:{provider_id}")
|
||||
return provider_or_404(provider_id).config.model_copy(deep=True)
|
||||
|
||||
|
||||
@router.post(
|
||||
@@ -390,7 +438,11 @@ async def delete_provider(provider_id: str) -> OperationResponse:
|
||||
tags=["Providers"],
|
||||
)
|
||||
async def list_provider_models(provider_id: str) -> ProviderModelsResponse:
|
||||
not_implemented(f"providers.models:{provider_id}")
|
||||
provider_or_404(provider_id)
|
||||
return ProviderModelsResponse(
|
||||
provider_id=provider_id,
|
||||
items=await container.providers.list_models(provider_id),
|
||||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
@@ -399,8 +451,9 @@ async def list_provider_models(provider_id: str) -> ProviderModelsResponse:
|
||||
responses=not_implemented_response,
|
||||
tags=["Providers"],
|
||||
)
|
||||
async def test_provider(_: ProviderTestRequest) -> ProviderTestResponse:
|
||||
not_implemented("providers.test")
|
||||
async def test_provider(request: ProviderTestRequest) -> ProviderTestResponse:
|
||||
provider_or_404(request.provider_id)
|
||||
return await container.providers.test(request.provider_id, request.model)
|
||||
|
||||
|
||||
# Tasks
|
||||
|
||||
@@ -0,0 +1,181 @@
|
||||
import asyncio
|
||||
|
||||
from app.agent.permissions import PermissionMode
|
||||
from app.agent.tools import ToolExecutionContext
|
||||
from app.container import build_container
|
||||
from app.contracts import (
|
||||
AgentEventType,
|
||||
AgentRunCreateRequest,
|
||||
AgentRunStatus,
|
||||
ToolCall,
|
||||
)
|
||||
|
||||
|
||||
def run(coroutine):
|
||||
return asyncio.run(coroutine)
|
||||
|
||||
|
||||
def test_mock_provider_completes_agent_run() -> None:
|
||||
async def scenario() -> None:
|
||||
container = build_container()
|
||||
created = await container.agent.create_run(
|
||||
AgentRunCreateRequest(
|
||||
input="hello",
|
||||
provider_id="mock",
|
||||
model="mock-1",
|
||||
)
|
||||
)
|
||||
|
||||
completed = await container.agent.wait(created.run_id)
|
||||
events = [event async for event in container.agent.events(created.run_id)]
|
||||
|
||||
assert completed.status == AgentRunStatus.completed
|
||||
assert completed.output == "Mock response: hello"
|
||||
assert events[0].event == AgentEventType.run_started
|
||||
assert events[-1].event == AgentEventType.run_completed
|
||||
|
||||
run(scenario())
|
||||
|
||||
|
||||
def test_agent_calls_registered_tool_and_records_result() -> None:
|
||||
async def scenario() -> None:
|
||||
container = build_container()
|
||||
created = await container.agent.create_run(
|
||||
AgentRunCreateRequest(
|
||||
input='/tool system.echo {"text":"hello tool"}',
|
||||
provider_id="mock",
|
||||
model="mock-1",
|
||||
allowed_tools=["system.echo"],
|
||||
)
|
||||
)
|
||||
|
||||
completed = await container.agent.wait(created.run_id)
|
||||
|
||||
assert completed.status == AgentRunStatus.completed
|
||||
assert completed.current_step == 2
|
||||
assert completed.tool_results[0].success is True
|
||||
assert completed.tool_results[0].output == {"text": "hello tool"}
|
||||
assert completed.output is not None
|
||||
assert "Tool result received" in completed.output
|
||||
|
||||
run(scenario())
|
||||
|
||||
|
||||
def test_tool_arguments_are_validated() -> None:
|
||||
async def scenario() -> None:
|
||||
container = build_container()
|
||||
result = await container.tools.execute(
|
||||
ToolCall(
|
||||
tool_call_id="call_invalid",
|
||||
name="math.add",
|
||||
arguments={"left": 1},
|
||||
),
|
||||
ToolExecutionContext(run_id="run_test"),
|
||||
)
|
||||
|
||||
assert result.success is False
|
||||
assert result.error_code == "TOOL_ARGUMENT_INVALID"
|
||||
|
||||
run(scenario())
|
||||
|
||||
|
||||
def test_permission_confirmation_resumes_agent() -> None:
|
||||
async def scenario() -> None:
|
||||
container = build_container()
|
||||
protected_tool = container.tools.get("system.echo")
|
||||
protected_tool.definition.permission = "notes.write"
|
||||
container.permissions.policy.set_rule("notes.write", PermissionMode.confirm)
|
||||
|
||||
created = await container.agent.create_run(
|
||||
AgentRunCreateRequest(
|
||||
input='/tool system.echo {"text":"approved"}',
|
||||
provider_id="mock",
|
||||
model="mock-1",
|
||||
allowed_tools=["system.echo"],
|
||||
tool_timeout_seconds=2,
|
||||
)
|
||||
)
|
||||
|
||||
request_id = None
|
||||
async with asyncio.timeout(2):
|
||||
async for event in container.agent.events(created.run_id):
|
||||
if event.event == AgentEventType.permission_required:
|
||||
request_id = str(event.data["request_id"])
|
||||
break
|
||||
|
||||
assert request_id is not None
|
||||
assert container.agent.resolve_permission(
|
||||
created.run_id, request_id, "allow_once"
|
||||
)
|
||||
completed = await container.agent.wait(created.run_id)
|
||||
assert completed.status == AgentRunStatus.completed
|
||||
assert completed.tool_results[0].success is True
|
||||
|
||||
run(scenario())
|
||||
|
||||
|
||||
def test_step_limit_stops_repeated_agent_loop() -> None:
|
||||
async def scenario() -> None:
|
||||
container = build_container()
|
||||
created = await container.agent.create_run(
|
||||
AgentRunCreateRequest(
|
||||
input='/tool system.echo {"text":"one step"}',
|
||||
provider_id="mock",
|
||||
model="mock-1",
|
||||
allowed_tools=["system.echo"],
|
||||
max_steps=1,
|
||||
)
|
||||
)
|
||||
|
||||
completed = await container.agent.wait(created.run_id)
|
||||
assert completed.status == AgentRunStatus.failed
|
||||
assert completed.error_code == "MAX_STEPS_EXCEEDED"
|
||||
assert len(completed.tool_results) == 1
|
||||
|
||||
run(scenario())
|
||||
|
||||
|
||||
def test_token_budget_stops_agent_run() -> None:
|
||||
async def scenario() -> None:
|
||||
container = build_container()
|
||||
created = await container.agent.create_run(
|
||||
AgentRunCreateRequest(
|
||||
input="hello budget",
|
||||
provider_id="mock",
|
||||
model="mock-1",
|
||||
token_budget=1,
|
||||
)
|
||||
)
|
||||
|
||||
completed = await container.agent.wait(created.run_id)
|
||||
assert completed.status == AgentRunStatus.failed
|
||||
assert completed.error_code == "TOKEN_BUDGET_EXCEEDED"
|
||||
|
||||
run(scenario())
|
||||
|
||||
|
||||
def test_cancelling_permission_wait_cancels_run() -> None:
|
||||
async def scenario() -> None:
|
||||
container = build_container()
|
||||
protected_tool = container.tools.get("system.echo")
|
||||
protected_tool.definition.permission = "notes.write"
|
||||
created = await container.agent.create_run(
|
||||
AgentRunCreateRequest(
|
||||
input='/tool system.echo {"text":"cancel"}',
|
||||
provider_id="mock",
|
||||
model="mock-1",
|
||||
allowed_tools=["system.echo"],
|
||||
tool_timeout_seconds=10,
|
||||
)
|
||||
)
|
||||
|
||||
async with asyncio.timeout(2):
|
||||
while container.agent.get_run(created.run_id).status != AgentRunStatus.waiting_permission:
|
||||
await asyncio.sleep(0)
|
||||
|
||||
cancelled = await container.agent.cancel(created.run_id)
|
||||
await container.agent.wait(created.run_id)
|
||||
assert cancelled.status == AgentRunStatus.cancelled
|
||||
assert container.agent.get_run(created.run_id).cancelled is True
|
||||
|
||||
run(scenario())
|
||||
@@ -28,7 +28,7 @@ def test_query_shells_are_empty_and_typed() -> None:
|
||||
assert notes.page.limit == 20
|
||||
assert skills.items == []
|
||||
assert plugins.items == []
|
||||
assert providers.items == []
|
||||
assert [provider.provider_id for provider in providers.items] == ["mock"]
|
||||
assert index.status == "idle"
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,206 @@
|
||||
# AI Core 与 Agent Core 开发说明
|
||||
|
||||
> 本文档用于团队开发和模块联调,记录当前已经落地的核心边界与使用方式。
|
||||
|
||||
## 当前实现
|
||||
|
||||
当前已经建立第一条可运行链路:
|
||||
|
||||
```text
|
||||
FastAPI
|
||||
→ Provider Registry
|
||||
→ Agent Runtime
|
||||
→ Permission Manager
|
||||
→ Tool Registry
|
||||
→ Agent Trace / SSE
|
||||
```
|
||||
|
||||
对应代码:
|
||||
|
||||
```text
|
||||
backend/app/
|
||||
├── providers/
|
||||
│ ├── base.py Provider Protocol 与统一 Turn
|
||||
│ ├── registry.py Provider 注册、发现、模型列表和连接测试
|
||||
│ └── mock.py 离线开发 Provider
|
||||
├── agent/
|
||||
│ ├── runtime.py Agent Loop、限制、取消、Trace 和 SSE
|
||||
│ ├── tools.py Tool 注册、参数校验、隔离执行和结果转换
|
||||
│ ├── permissions.py 权限策略、确认请求和会话授权
|
||||
│ └── builtin_tools.py 无副作用的内置开发 Tool
|
||||
└── container.py AI Core 依赖组装
|
||||
```
|
||||
|
||||
Router 只负责 HTTP/SSE 与错误转换,不实现 Agent、Tool 或 Provider 业务逻辑。
|
||||
|
||||
## 模块边界
|
||||
|
||||
当前实现属于范涵宇负责的 AI Core / Agent Core:
|
||||
|
||||
- Provider 抽象与注册;
|
||||
- Agent Run 生命周期;
|
||||
- Tool Registry;
|
||||
- Tool 参数校验与执行隔离;
|
||||
- Permission;
|
||||
- Step、Timeout、Token Budget、取消;
|
||||
- 内存 Trace 与 SSE;
|
||||
- 公共 Contract 和 API 接入。
|
||||
|
||||
以下内容保持接口,不在本模块实现:
|
||||
|
||||
- Note、NoteBlock、Markdown Parser:由 Knowledge Core 提供;
|
||||
- FTS5、Vector、RRF、Reranker、Citation:由 Retrieval Core 提供;
|
||||
- 文件系统和 API Key 明文读取:由 Rust Host 提供;
|
||||
- Skill、Plugin 生命周期:后续在 Extension Core 中实现。
|
||||
|
||||
## 开发 Provider
|
||||
|
||||
默认注册离线 Provider:
|
||||
|
||||
```text
|
||||
provider_id = mock
|
||||
model = mock-1
|
||||
```
|
||||
|
||||
它支持:
|
||||
|
||||
```text
|
||||
chat
|
||||
tool_calling
|
||||
streaming
|
||||
```
|
||||
|
||||
普通 Chat 请求:
|
||||
|
||||
```json
|
||||
{
|
||||
"provider_id": "mock",
|
||||
"model": "mock-1",
|
||||
"messages": [
|
||||
{"role": "user", "content": "hello"}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
`POST /api/chat` 返回 ModelEvent SSE。
|
||||
|
||||
## Agent Run
|
||||
|
||||
创建普通 Agent Run:
|
||||
|
||||
```json
|
||||
{
|
||||
"input": "hello",
|
||||
"provider_id": "mock",
|
||||
"model": "mock-1",
|
||||
"max_steps": 10
|
||||
}
|
||||
```
|
||||
|
||||
请求:
|
||||
|
||||
```text
|
||||
POST /api/agent/runs
|
||||
```
|
||||
|
||||
创建后通过以下接口读取状态和事件:
|
||||
|
||||
```text
|
||||
GET /api/agent/runs/{run_id}
|
||||
GET /api/agent/runs/{run_id}/events
|
||||
POST /api/agent/runs/{run_id}/cancel
|
||||
```
|
||||
|
||||
当前 Run 与 Trace 保存在内存中,AI Core 重启后清空。后续数据库层接入时替换 Repository,不改变 API Contract。
|
||||
|
||||
## Tool Calling
|
||||
|
||||
当前注册两个无副作用开发 Tool:
|
||||
|
||||
```text
|
||||
system.echo
|
||||
math.add
|
||||
```
|
||||
|
||||
Mock Provider 使用下面的开发语法产生 Tool Call:
|
||||
|
||||
```text
|
||||
/tool system.echo {"text":"hello tool"}
|
||||
/tool math.add {"left":1,"right":2}
|
||||
```
|
||||
|
||||
Agent Run 需要显式声明 `allowed_tools`:
|
||||
|
||||
```json
|
||||
{
|
||||
"input": "/tool math.add {\"left\":1,\"right\":2}",
|
||||
"provider_id": "mock",
|
||||
"model": "mock-1",
|
||||
"allowed_tools": ["math.add"]
|
||||
}
|
||||
```
|
||||
|
||||
Tool 参数由独立 Pydantic Model 再次校验。Tool 的异常、非法参数、超时和权限拒绝统一转换为 `ToolResult`,不会直接打断 API 进程。
|
||||
|
||||
## Permission
|
||||
|
||||
Permission Policy 当前支持:
|
||||
|
||||
```text
|
||||
allow
|
||||
confirm
|
||||
deny
|
||||
```
|
||||
|
||||
需要确认时,Agent 状态进入 `waiting_permission`,并发出 `PermissionRequired` 事件。前端使用:
|
||||
|
||||
```text
|
||||
POST /api/agent/runs/{run_id}/permissions/{request_id}
|
||||
```
|
||||
|
||||
提交以下决策之一:
|
||||
|
||||
```json
|
||||
{"decision":"allow_once"}
|
||||
{"decision":"allow_session"}
|
||||
{"decision":"deny"}
|
||||
```
|
||||
|
||||
默认需要确认的高影响权限包括 `notes.write`、`notes.delete`、`network.request` 和 `secrets.use`。
|
||||
|
||||
## Knowledge / Retrieval 接入约定
|
||||
|
||||
其他成员完成服务后,通过注册 Tool 接入 Agent,不让 Agent Runtime 直接依赖具体实现:
|
||||
|
||||
```python
|
||||
tool_registry.register(
|
||||
definition=tool_definition,
|
||||
arguments_model=arguments_model,
|
||||
executor=executor,
|
||||
)
|
||||
```
|
||||
|
||||
建议第一批接入:
|
||||
|
||||
```text
|
||||
notes.search
|
||||
notes.read
|
||||
notes.create
|
||||
notes.update
|
||||
notes.list
|
||||
notes.move
|
||||
rag.search
|
||||
tasks.create
|
||||
tasks.update
|
||||
tasks.list
|
||||
```
|
||||
|
||||
写操作 Executor 调用 Knowledge Core Service,不直接访问 SQLite;检索 Executor 调用 Retrieval Core Service,不直接拼接 FTS5 或 sqlite-vec SQL。
|
||||
|
||||
## 当前限制与下一步
|
||||
|
||||
- Provider 目前只有完全离线的 Mock 实现;下一步实现 OpenAI-Compatible 与 Ollama Adapter。
|
||||
- Run/Trace 暂存内存;下一步抽象 Repository 并接入 SQLite。
|
||||
- Permission 已有核心等待/恢复机制,前端确认 UI 尚未联调。
|
||||
- Note/RAG Tool 等待对应模块 Service 接入。
|
||||
- Skill/Plugin 将复用现有 Tool Registry 和 Permission Manager。
|
||||
+5
-4
@@ -151,9 +151,10 @@ RunFailed
|
||||
RunCancelled
|
||||
```
|
||||
|
||||
## 当前壳子行为
|
||||
## 当前实现状态
|
||||
|
||||
- 列表、搜索、索引状态等只读接口返回符合 Contract 的空结果或 `idle` 状态。
|
||||
- Chat 与 Agent Events 返回符合 SSE 格式的 `NOT_IMPLEMENTED` 事件。
|
||||
- 需要数据库、文件、模型或 Runtime 的操作统一返回 `501`。
|
||||
- Chat、Agent Run、Agent Events、Tool 列表、Provider 列表、模型列表和连接测试已经接入 AI Core。
|
||||
- 默认提供 `mock/mock-1` 离线 Provider,以及 `system.echo`、`math.add` 开发 Tool。
|
||||
- Notes、Search、Skills、Plugins、Tasks、Media、Index 等尚未接入业务服务的接口继续返回空结果、`idle` 或 `501`。
|
||||
- 需要尚未接入的数据库、文件或扩展 Runtime 的操作统一返回 `501`。
|
||||
- 接入业务模块时保持当前路径和 Contract,不在 Router 中直接实现数据库、Provider 或 Agent 逻辑。
|
||||
|
||||
Reference in New Issue
Block a user