实现 AI Core 与 Agent Core 基础功能

- 更新 README 描述从后端壳子到 AI Core/Agent Core
- 添加 ToolCall 和 ToolResult 数据结构定义
- 扩展 AgentRun 模型增加输出、错误码、工具调用结果等字段
- 添加 mock 提供商类型支持
- 实现聊天、代理运行、工具调用和提供商管理的核心路由逻辑
- 集成容器化依赖注入和错误处理机制
- 更新 API 接口契约和文档说明
This commit is contained in:
2026-08-27 13:55:00 +08:00
parent ce155b27f4
commit b71984d951
18 changed files with 1369 additions and 40 deletions
+3 -1
View File
@@ -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`
+12
View File
@@ -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",
]
+46
View File
@@ -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,
)
+80
View File
@@ -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)
+346
View File
@@ -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
+108
View File
@@ -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),
)
+49
View File
@@ -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()
+22
View File
@@ -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
View File
@@ -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(
+11
View File
@@ -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",
]
+30
View File
@@ -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]: ...
+127
View File
@@ -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."
+55
View File
@@ -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
View File
@@ -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
+181
View File
@@ -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())
+1 -1
View File
@@ -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"